{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":6917177,"sourceType":"datasetVersion","datasetId":3889865},{"sourceId":7001173,"sourceType":"datasetVersion","datasetId":3881967},{"sourceId":7029230,"sourceType":"datasetVersion","datasetId":4027203}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# !nvcc --version\n# # needed to get CutMix from TV 0.16\n# !pip uninstall -q -y torchaudio\n# # !conda install pytorch==2.1.0 torchvision==0.16.0 cudatoolkit=11.6 -c pytorch\n# !pip install -q torch torchvision torchdata -U -f https://download.pytorch.org/whl/cu118\n# !pip list | grep torch","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-05T17:58:22.504153Z","iopub.execute_input":"2023-12-05T17:58:22.504494Z","iopub.status.idle":"2023-12-05T17:58:22.509609Z","shell.execute_reply.started":"2023-12-05T17:58:22.504463Z","shell.execute_reply":"2023-12-05T17:58:22.508761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%load_ext autoreload\n%autoreload 2","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-14T02:19:01.912579Z","iopub.execute_input":"2023-12-14T02:19:01.912934Z","iopub.status.idle":"2023-12-14T02:19:01.941153Z","shell.execute_reply.started":"2023-12-14T02:19:01.912905Z","shell.execute_reply":"2023-12-14T02:19:01.940402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os, glob\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nDATASET_FOLDER = \"/kaggle/input/UBC-OCEAN/\"\nDATASET_TILES = \"/kaggle/input/ubc-ocean-tiles-w-masks-2048px-scale-0-25/\"\nDATASET_IMAGES = os.path.join(DATASET_TILES, \"train_images\")\nDATASET_MASKS = os.path.join(DATASET_TILES, \"train_masks\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-05T17:58:22.551313Z","iopub.execute_input":"2023-12-05T17:58:22.551592Z","iopub.status.idle":"2023-12-05T17:58:22.901865Z","shell.execute_reply.started":"2023-12-05T17:58:22.551545Z","shell.execute_reply":"2023-12-05T17:58:22.900891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Checkout training labels","metadata":{}},{"cell_type":"code","source":"df_train = pd.read_csv(os.path.join(DATASET_FOLDER, \"train.csv\"))\n# labels = list(df_train[\"label\"].unique())\nprint(f\"Dataset/train size: {len(df_train)}\")\ndisplay(df_train.head())","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-05T17:58:22.903083Z","iopub.execute_input":"2023-12-05T17:58:22.903577Z","iopub.status.idle":"2023-12-05T17:58:22.956841Z","shell.execute_reply.started":"2023-12-05T17:58:22.903543Z","shell.execute_reply":"2023-12-05T17:58:22.955935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inds = list(map(os.path.basename, glob.glob(os.path.join(DATASET_IMAGES, \"*\"))))\nprint(f\"{inds=}\")\ndf_train = df_train[df_train[\"image_id\"].astype(str).isin(inds)]\nprint(f\"Dataset/train size: {len(df_train)}\")\ndisplay(df_train.head())","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-12-05T17:58:22.959297Z","iopub.execute_input":"2023-12-05T17:58:22.959705Z","iopub.status.idle":"2023-12-05T17:58:23.017842Z","shell.execute_reply.started":"2023-12-05T17:58:22.95967Z","shell.execute_reply":"2023-12-05T17:58:23.017003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_= df_train[[\"label\"]].value_counts().plot.pie(autopct='%1.1f%%', ylabel=\"label\", figsize=(3,3))","metadata":{"execution":{"iopub.status.busy":"2023-12-05T17:58:23.019011Z","iopub.execute_input":"2023-12-05T17:58:23.019262Z","iopub.status.idle":"2023-12-05T17:58:23.22702Z","shell.execute_reply.started":"2023-12-05T17:58:23.019239Z","shell.execute_reply":"2023-12-05T17:58:23.225525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Show some samples 🖼️ per class\n\n**TODO** filter only tiles with tumor based on mask","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nnb_samples = 6\nn, m = len(np.unique(df_train['label'])), nb_samples\nfig, axarr = plt.subplots(nrows=n, ncols=m, figsize=(m * 2, n * 2))\nfor ilb, (lb, df_) in enumerate(df_train.groupby('label')):\n    img_ids = list(df_['image_id'])\n    for i in range(m):\n        if i == 0:\n            axarr[ilb, i].set_title(f\"{lb} #{len(df_)}\")\n        ls_imgs = glob.glob(os.path.join(DATASET_IMAGES, str(img_ids[i]), \"*.png\"))\n        img_path = ls_imgs[0]\n        img = plt.imread(img_path)\n        mask = np.sum(img[..., :3], axis=2) == 0\n        img[mask, :] = 255\n        axarr[ilb, i].imshow(img)\n        # axarr[ilb, i].set_xticks([])\n        # axarr[ilb, i].set_yticks([])\n_= plt.axis('off')","metadata":{"execution":{"iopub.status.busy":"2023-12-05T17:58:23.228488Z","iopub.execute_input":"2023-12-05T17:58:23.228882Z","iopub.status.idle":"2023-12-05T17:58:32.801108Z","shell.execute_reply.started":"2023-12-05T17:58:23.228848Z","shell.execute_reply":"2023-12-05T17:58:32.800157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data pre-processing\n\n### Color 🦩 normalizations","metadata":{}},{"cell_type":"code","source":"from PIL import Image\nfrom joblib import Parallel, delayed\nfrom tqdm.auto import tqdm\n\ndef _color_means(img_path):\n    img = np.array(Image.open(img_path))\n    mask = np.sum(img[..., :3], axis=2) == 0\n    img[mask, :] = 255\n    if np.max(img) > 1.5:\n        img = img / 255.0\n    clr_mean = {i: np.mean(img[..., i]) for i in range(3)}\n    clr_std = {i: np.std(img[..., i]) for i in range(3)}\n    return clr_mean, clr_std\n\n# os.path.join(DATASET_SMALL_FOLDER, \"train_images\")\nls_images = glob.glob(os.path.join(DATASET_IMAGES, \"*\", \"*.png\"))\nclr_mean_std = Parallel(n_jobs=os.cpu_count())(delayed(_color_means)(fn) for fn in tqdm(ls_images[:10_000]))","metadata":{"execution":{"iopub.status.busy":"2023-12-05T18:00:11.794975Z","iopub.execute_input":"2023-12-05T18:00:11.795452Z","iopub.status.idle":"2023-12-05T18:01:38.530756Z","shell.execute_reply.started":"2023-12-05T18:00:11.795413Z","shell.execute_reply":"2023-12-05T18:01:38.529629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_color_mean = pd.DataFrame([c[0] for c in clr_mean_std]).describe()\ndisplay(img_color_mean.T)\nimg_color_std = pd.DataFrame([c[1] for c in clr_mean_std]).describe()\ndisplay(img_color_std.T)\n\nimg_color_mean = list(img_color_mean.T[\"mean\"])\nimg_color_std = list(img_color_std.T[\"mean\"])\nprint(f\"{img_color_mean=}\\n{img_color_std=}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-05T18:01:38.532919Z","iopub.execute_input":"2023-12-05T18:01:38.533758Z","iopub.status.idle":"2023-12-05T18:01:38.631158Z","shell.execute_reply.started":"2023-12-05T18:01:38.533723Z","shell.execute_reply":"2023-12-05T18:01:38.630224Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset & DataModule\n\nCreating standard PyTorch dataset to define how the data shall be loaded and set representations. We define the sample pair as:\n- RGB image\n- one-hot lable encding\n\nA DataModule standardizes the training, val, test splits, data preparation and transforms. The main advantage is consistent data splits, data preparation and transforms across models.","metadata":{}},{"cell_type":"code","source":"# !pip install -U -q pytorch-lightning","metadata":{"execution":{"iopub.status.busy":"2023-12-05T18:01:38.632601Z","iopub.execute_input":"2023-12-05T18:01:38.632956Z","iopub.status.idle":"2023-12-05T18:01:38.658712Z","shell.execute_reply.started":"2023-12-05T18:01:38.632924Z","shell.execute_reply":"2023-12-05T18:01:38.657544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport torch\nimport random\nfrom PIL import Image\nfrom torch.utils.data import Dataset\n\nclass CancerTilesDataset(Dataset):\n\n    def __init__(\n        self,\n        df_data,\n        path_img_dir: str,\n        path_mask_dir: str,\n        transforms = None,\n        split: float = 0.9,\n        mode: str = 'train',\n        labels_lut = None,\n        one_hot: bool = True,\n        tumor_thr: float = 0.3,\n        #white_thr: int = 225,\n        #thr_max_bg: float = 0.2,\n    ):\n        assert os.path.isdir(path_img_dir)\n        self.path_img_dir = path_img_dir\n        self.path_mask_dir = path_mask_dir\n        self.transforms = transforms\n        self.split = split\n        self.mode = mode\n        self.tumor_thr = tumor_thr\n        #self.white_thr = white_thr\n        #self.thr_max_bg = thr_max_bg\n\n        self.one_hot = one_hot\n        self.data = df_data\n        self.labels_unique = sorted(self.data[\"label\"].unique()) + [\"Other\"]\n        self.labels_lut = labels_lut or {lb: i for i, lb in enumerate(self.labels_unique)}\n        # shuffle data\n        ls_img = sorted(glob.glob(os.path.join(self.path_img_dir, \"*\", \"*.png\")))\n        random.Random(42).shuffle(ls_img)\n        self.imgs = [(os.path.basename(os.path.dirname(p)), os.path.basename(p))\n                     for p in ls_img]\n\n        # split dataset\n        assert 0.0 <= self.split <= 1.0\n        frac = int(self.split * len(self.imgs))\n        self.imgs = self.imgs[:frac] if mode == 'train' else self.imgs[frac:]\n        #self.labels = list(self.data['label'])\n\n    @property\n    def num_classes(self) -> int:\n        return len(self.labels_lut)\n\n    def to_one_hot(self, label: str) -> tuple:\n        one_hot = [0] * self.num_classes\n        one_hot[self.labels_lut[label]] = 1\n        return tuple(one_hot)\n\n    def __getitem__(self, idx: int) -> tuple:\n        image_id, tile = self.imgs[idx]\n        img_path = os.path.join(self.path_img_dir, image_id, tile)\n        mask_path = os.path.join(self.path_mask_dir, image_id, tile)\n        \n        img = np.array(Image.open(img_path))[..., :3]\n        black_bg = np.sum(img, axis=2) == 0\n        img[black_bg, :] = 255\n        \n        mask = np.array(Image.open(mask_path))[..., :3]\n        tumor_mask = float(np.sum(mask == 1)) / np.prod(mask.shape)\n        tumor_type = self.data.loc[self.data['image_id'] == int(image_id), \"label\"]\n        #print(tumor_mask, tumor_type)\n        lb = tumor_type.item() if tumor_mask > self.tumor_thr else \"Other\"\n        # TODO: assume as multilabel problem\n        if self.one_hot:\n            lb = self.to_one_hot(lb)\n        else:\n            lb = self.labels_lut[lb]\n        lb = torch.tensor(lb).to(int)\n\n        # augmentation\n        if self.transforms:\n            img = self.transforms(Image.fromarray(img))\n        #print(f\"img dim: {img.shape}\")\n        return img, lb\n\n    def __len__(self) -> int:\n        return len(self.imgs)\n\n# ==============================\n# ==============================\n\ndataset = CancerTilesDataset(df_train, DATASET_IMAGES, DATASET_MASKS)\n\n# quick view\nfig, axes = plt.subplots(nrows=3, ncols=3, figsize=(10, 10))\nfor i in range(9):\n    img, lb = dataset[i]\n    ax = axes[i // 3, i % 3]\n    ax.imshow(img)\n    ax.set_title(f\"{lb}\\n dims: {img.shape} /{np.max(img)}\")\nfig.tight_layout()","metadata":{"execution":{"iopub.status.busy":"2023-12-05T18:01:38.661407Z","iopub.execute_input":"2023-12-05T18:01:38.661749Z","iopub.status.idle":"2023-12-05T18:01:43.652365Z","shell.execute_reply.started":"2023-12-05T18:01:38.661722Z","shell.execute_reply":"2023-12-05T18:01:43.651234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let us define some standard image augmentaion procedures and color normalizations...","metadata":{}},{"cell_type":"code","source":"from torchvision import transforms as T\nfrom torchvision.transforms import InterpolationMode\n\nTRAIN_TRANSFORM = T.Compose([\n    T.CenterCrop(512),\n    #T.RandomResizedCrop(512, interpolation=InterpolationMode.BICUBIC, antialias=True),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\n    T.RandomRotation(degrees=10),\n    T.ColorJitter(),\n    T.ToTensor(),\n    #T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n    T.Normalize(img_color_mean, img_color_std),  # custom\n])\n\nVALID_TRANSFORM = T.Compose([\n    T.CenterCrop(512),\n    T.ToTensor(),\n    #T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n    T.Normalize(img_color_mean, img_color_std),  # custom\n])","metadata":{"execution":{"iopub.status.busy":"2023-12-05T18:01:43.653705Z","iopub.execute_input":"2023-12-05T18:01:43.65413Z","iopub.status.idle":"2023-12-05T18:01:43.984893Z","shell.execute_reply.started":"2023-12-05T18:01:43.654094Z","shell.execute_reply":"2023-12-05T18:01:43.984111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The DataModule include creating training and validation dataset with given split and feading it to particular data loaders...","metadata":{}},{"cell_type":"code","source":"import operator\nimport torchvision\nfrom lightning_utilities import compare_version\n\nTV_GE_0_16 = compare_version(\"torchvision\", operator.ge, \"0.16.0\")\nprint(TV_GE_0_16)","metadata":{"execution":{"iopub.status.busy":"2023-12-05T18:01:43.985975Z","iopub.execute_input":"2023-12-05T18:01:43.986242Z","iopub.status.idle":"2023-12-05T18:01:44.056413Z","shell.execute_reply.started":"2023-12-05T18:01:43.986219Z","shell.execute_reply":"2023-12-05T18:01:44.055557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import multiprocessing as mproc\nimport pytorch_lightning as pl\nfrom torch.utils.data import DataLoader\n\nclass CancerSubtypeDM(pl.LightningDataModule):\n\n    def __init__(\n        self,\n        df_data,\n        path_img_dir: str,\n        path_mask_dir: str,\n        batch_size: int = 32,\n        num_workers: int = None,\n        train_transforms = TRAIN_TRANSFORM,\n        valid_transforms = VALID_TRANSFORM,\n        one_hot: bool = True,\n    ):\n        super().__init__()\n        self.df_data = df_data\n        self.path_img_dir = path_img_dir\n        self.path_mask_dir = path_mask_dir\n        self.batch_size = batch_size\n        self.num_workers = num_workers or mproc.cpu_count()\n        self.train_dataset = None\n        self.valid_dataset = None\n        self.train_transforms = train_transforms\n        self.valid_transforms = valid_transforms\n        self.one_hot = one_hot\n\n    def prepare_data(self):\n        pass\n\n    @property\n    def num_classes(self) -> int:\n        assert self.train_dataset and self.valid_dataset\n        return len(set(self.train_dataset.labels_unique + self.valid_dataset.labels_unique))\n\n    def setup(self, stage=None):\n        self.train_dataset = CancerTilesDataset(\n            self.df_data, self.path_img_dir, self.path_mask_dir,\n            mode='train', transforms=self.train_transforms, one_hot=self.one_hot)\n        print(f\"training dataset: {len(self.train_dataset)}\")\n        self.valid_dataset = CancerTilesDataset(\n            self.df_data, self.path_img_dir, self.path_mask_dir,\n            mode='valid', transforms=self.valid_transforms, one_hot=self.one_hot,\n            # as validation is subsampled it may happen that some labels are missing\n            # and so created one-hot-encoding vector will be sorter\n            labels_lut=self.train_dataset.labels_lut)\n        print(f\"validation dataset: {len(self.valid_dataset)}\")\n\n    def train_dataloader(self):\n        return DataLoader(\n            self.train_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=True,\n        )\n\n    def val_dataloader(self):\n        return DataLoader(\n            self.valid_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=False,\n        )\n\n    def test_dataloader(self):\n        pass\n\n# ==============================\n# ==============================\n\ndm = CancerSubtypeDM(\n    df_train, DATASET_IMAGES, DATASET_MASKS,\n    batch_size=18, one_hot=not TV_GE_0_16\n)\ndm.setup()\nprint(dm.num_classes)\n\n# quick view\nfig = plt.figure(figsize=(3, 9))\nfor imgs, lbs in dm.train_dataloader():\n    print(f'batch labels: {lbs if TV_GE_0_16 else torch.sum(lbs, axis=0)}')\n    print(f'image size: {imgs[0].shape}')\n    for i in range(3):\n        ax = fig.add_subplot(3, 1, i + 1)\n        #print(np.rollaxis(imgs[i].numpy(), 0, 3).shape)\n        ax.imshow(np.rollaxis(imgs[i].numpy(), 0, 3))\n        ax.set_title(lbs[i])\n    break","metadata":{"execution":{"iopub.status.busy":"2023-12-05T18:01:44.057602Z","iopub.execute_input":"2023-12-05T18:01:44.057888Z","iopub.status.idle":"2023-12-05T18:01:50.716784Z","shell.execute_reply.started":"2023-12-05T18:01:44.057864Z","shell.execute_reply":"2023-12-05T18:01:50.715803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CNN Model\n\nWe start with some standard CNN models taken from TIMM.\nThen we define Ligthning module including training and validation step and configure optimizer/schedular.\n\n- **Schedulers's example**: https://www.kaggle.com/code/isbhargav/guide-to-pytorch-learning-rate-scheduling\n- **Schedulers explained**: https://towardsdatascience.com/a-visual-guide-to-learning-rate-schedulers-in-pytorch-24bbb262c863#5407\n- **TIMM models**: https://github.com/huggingface/pytorch-image-models/blob/main/results/results-imagenet.csv","metadata":{}},{"cell_type":"code","source":"!pip install -q lion-pytorch adan-pytorch torch_optimizer","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-12-05T18:01:50.718451Z","iopub.execute_input":"2023-12-05T18:01:50.719535Z","iopub.status.idle":"2023-12-05T18:02:04.681121Z","shell.execute_reply.started":"2023-12-05T18:01:50.719499Z","shell.execute_reply":"2023-12-05T18:02:04.679834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\nimport torch\nfrom adan_pytorch import Adan\nfrom lion_pytorch import Lion\nfrom torch_optimizer import AdaBound, RAdam, Yogi\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torchmetrics.classification import MulticlassAccuracy, MulticlassF1Score\nfrom torchvision.transforms import v2\n\nclass LitCancerSubtype(pl.LightningModule):\n\n    def __init__(self, net, lr: float = 1e-4):\n        super().__init__()\n        self.net = net\n        self.arch = net.pretrained_cfg.get('architecture')\n        self.num_classes = net.num_classes\n        self.train_accuracy = MulticlassAccuracy(num_classes=self.num_classes)\n        self.val_accuracy = MulticlassAccuracy(num_classes=self.num_classes)\n        self.val_f1_score = MulticlassF1Score(num_classes=self.num_classes)\n        self.learn_rate = lr\n        if TV_GE_0_16:\n            # expriment with CutMix | MixUp with torch 2.1+\n            # https://pytorch.org/vision/stable/auto_examples/transforms/plot_cutmix_mixup.html#where-to-use-mixup-and-cutmix\n            cutmix = v2.CutMix(num_classes=net.num_classes)\n            mixup = v2.MixUp(num_classes=net.num_classes)\n            self.cutmix_or_mixup = v2.RandomChoice([cutmix, mixup])\n\n    def forward(self, x):\n        y = F.softmax(self.net(x))\n        if y.isnan().any():\n            y = torch.ones_like(y) / self.num_classes\n        return y\n\n    def compute_loss(self, y_hat, y):\n        return F.cross_entropy(y_hat, y.to(y_hat.dtype))\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        if TV_GE_0_16:\n            x, y = self.cutmix_or_mixup(x, y)\n        y_hat = self(x)\n        #print(f\"{lbs=} ?= {y_hat=}\")\n        loss = self.compute_loss(y_hat, y)\n        #print(f\"{y=} ?= {y_hat=} -> {loss=}\")\n        self.log(\"train_loss\", loss, logger=True, prog_bar=True)\n        if not TV_GE_0_16:\n            y = torch.argmax(y, axis=1)\n        #print(f\"{lb=} ?= {y_hat=} -> {self.train_accuracy(y_hat, lbs)}\")\n        self.log(\"train_acc\", self.train_accuracy(y_hat, y), logger=True, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        x, y = batch\n        y_hat = self(x)\n        if not TV_GE_0_16:\n            loss = self.compute_loss(y_hat, y)\n            self.log(\"valid_loss\", loss, logger=True, prog_bar=False)\n            y = torch.argmax(y, axis=1)\n        self.log(\"valid_acc\", self.val_accuracy(y_hat, y), logger=True, prog_bar=False)\n        self.log(\"valid_f1\", self.val_f1_score(y_hat, y), logger=True, prog_bar=True)\n\n    def configure_optimizers(self):\n        optimizer = AdaBound(self.parameters(), lr=self.learn_rate)\n        #optimizer = RAdam(self.parameters(), lr=self.learn_rate)\n        #optimizer = torch.optim.AdamW(self.parameters(), lr=self.learn_rate)\n        #optimizer = Lion(self.parameters(), lr=self.learn_rate, weight_decay=1e-2)\n        #optimizer = Adan(self.parameters(), lr=self.learn_rate * 10, betas=(0.02, 0.08, 0.01), weight_decay=0.02)\n        \n        #scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        #    optimizer, T_max=self.trainer.max_epochs, eta_min=1e-6, verbose=True)\n        scheduler = torch.optim.lr_scheduler.CyclicLR(\n          optimizer, base_lr=self.learn_rate, max_lr=self.learn_rate * 5,\n          step_size_up=5, cycle_momentum=False, mode=\"triangular2\", verbose=True)\n        #scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        #    optimizer, max_lr=self.learn_rate * 5, steps_per_epoch=1, epochs=self.trainer.max_epochs)\n        return [optimizer], [scheduler]\n\n# ==============================\n# ==============================\n\nnet = timm.create_model('tf_efficientnetv2_s_in21ft1k', pretrained=True, num_classes=dm.num_classes)\n# net = timm.create_model('maxvit_tiny_tf_512', pretrained=True, num_classes=dm.num_classes)\nmodel = LitCancerSubtype(net=net, lr=1e-4)\n# print(model)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-12-05T18:02:04.682727Z","iopub.execute_input":"2023-12-05T18:02:04.683043Z","iopub.status.idle":"2023-12-05T18:02:13.229728Z","shell.execute_reply.started":"2023-12-05T18:02:04.683018Z","shell.execute_reply":"2023-12-05T18:02:13.228836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training\n\nWe use Pytorch Lightning which allow us to drop all the boilet plate code and simplify all training just to use/call Trainer...","metadata":{}},{"cell_type":"code","source":"logger = pl.loggers.CSVLogger(save_dir='logs/', name=model.arch)\nnb_epochs = 10 if torch.cuda.is_available() else 2\n\n# ==============================\n\ntrainer = pl.Trainer(\n#     accelerator=\"cuda\",\n#     devices=2,\n    # fast_dev_run=True,\n    # callbacks=[swa],\n    logger=logger,\n    max_epochs=nb_epochs,\n    precision=16,\n    accumulate_grad_batches=14,\n    #val_check_interval=0.5,\n)\n\n# ==============================\n\n# trainer.tune(model, datamodule=dm)\ntrainer.fit(model=model, datamodule=dm)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-12-05T18:02:13.233055Z","iopub.execute_input":"2023-12-05T18:02:13.23376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Quick visualization of the training process...","metadata":{}},{"cell_type":"code","source":"import seaborn as sn\n\nmetrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\ndel metrics[\"step\"]\nmetrics.set_index(\"epoch\", inplace=True)\n# display(metrics.dropna(axis=1, how=\"all\").head())\ng = sn.relplot(data=metrics, kind=\"line\")\nplt.gcf().set_size_inches(12, 4)\n# plt.gca().set_yscale('log')\nplt.grid()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Save the model!","metadata":{}},{"cell_type":"code","source":"trainer.save_checkpoint(\"image_classification_model.pt\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Inference coming in https://www.kaggle.com/code/jirkaborovec/cancer-subtype-lightning-torch-inference-tiles","metadata":{}}]}