{"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":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":7052893,"sourceType":"datasetVersion","datasetId":4059285},{"sourceId":7052895,"sourceType":"datasetVersion","datasetId":4059287},{"sourceId":7052899,"sourceType":"datasetVersion","datasetId":4059290},{"sourceId":7052901,"sourceType":"datasetVersion","datasetId":4059291}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Training on tile dataset of size 224 using maxvit\nA model is trained using only tiles that are annotated as \"Tumor\". The label of the tile is given by one of the 5 labels: CC, EC, HGSC, LGSC, MC. \n\nCode adapted from https://www.kaggle.com/code/jirkaborovec/cancer-subtype-tiles-masks-w-lightning-timm/","metadata":{}},{"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_IMAGES = sorted(\n    glob.glob(\"/kaggle/input/tilessize224scale0-175part*/UBC-OCEAN_train_tiles_size224_scale0.175/batch*/*/*/Tumor/*.png\")\n)\nTILE_SIZE = 224\nBATCH_SIZE = 24","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-30T00:36:22.501938Z","iopub.execute_input":"2023-11-30T00:36:22.502367Z","iopub.status.idle":"2023-11-30T00:36:27.580933Z","shell.execute_reply.started":"2023-11-30T00:36:22.50233Z","shell.execute_reply":"2023-11-30T00:36:27.580154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Checkout some labels","metadata":{}},{"cell_type":"code","source":"df_train = pd.read_csv(os.path.join(DATASET_FOLDER, \"train.csv\"))\nprint(f\"Dataset/train size (full images): {len(df_train)}\")\nprint(f\"Dataset/train size (tiles): {len(DATASET_IMAGES)}\")\ndisplay(df_train.head())","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-30T00:36:27.582822Z","iopub.execute_input":"2023-11-30T00:36:27.583435Z","iopub.status.idle":"2023-11-30T00:36:27.613212Z","shell.execute_reply.started":"2023-11-30T00:36:27.5834Z","shell.execute_reply":"2023-11-30T00:36:27.612361Z"},"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), title=\"Image-level labels\")","metadata":{"execution":{"iopub.status.busy":"2023-11-30T00:36:27.614478Z","iopub.execute_input":"2023-11-30T00:36:27.615062Z","iopub.status.idle":"2023-11-30T00:36:27.818065Z","shell.execute_reply.started":"2023-11-30T00:36:27.615028Z","shell.execute_reply":"2023-11-30T00:36:27.81687Z"},"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\n# import random\n# from joblib import Parallel, delayed\n# from tqdm.auto import tqdm\n\n# def _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#     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# random.seed(42)\n# img_paths = random.sample(DATASET_IMAGES, 10000)\n# clr_mean_std = Parallel(n_jobs=os.cpu_count())(delayed(_color_means)(fn) for fn in tqdm(img_paths))","metadata":{"execution":{"iopub.status.busy":"2023-11-30T00:36:27.820806Z","iopub.execute_input":"2023-11-30T00:36:27.821199Z","iopub.status.idle":"2023-11-30T00:36:27.826184Z","shell.execute_reply.started":"2023-11-30T00:36:27.821147Z","shell.execute_reply":"2023-11-30T00:36:27.825146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# img_color_mean = pd.DataFrame([c[0] for c in clr_mean_std]).describe()\n# display(img_color_mean.T)\n# img_color_std = pd.DataFrame([c[1] for c in clr_mean_std]).describe()\n# display(img_color_std.T)\n\n# img_color_mean = list(img_color_mean.T[\"mean\"])\n# img_color_std = list(img_color_std.T[\"mean\"])\n# print(f\"{img_color_mean=}\\n{img_color_std=}\")","metadata":{"execution":{"iopub.status.busy":"2023-11-30T00:36:27.827458Z","iopub.execute_input":"2023-11-30T00:36:27.827796Z","iopub.status.idle":"2023-11-30T00:36:27.83746Z","shell.execute_reply.started":"2023-11-30T00:36:27.82777Z","shell.execute_reply":"2023-11-30T00:36:27.836146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_color_mean=[0.8029574609001011, 0.6753532015392826, 0.8150805152175007]\nimg_color_std=[0.09240988133071351, 0.11661346690553148, 0.06439091956270869]","metadata":{"execution":{"iopub.status.busy":"2023-11-30T00:36:27.838963Z","iopub.execute_input":"2023-11-30T00:36:27.840456Z","iopub.status.idle":"2023-11-30T00:36:27.847502Z","shell.execute_reply.started":"2023-11-30T00:36:27.840407Z","shell.execute_reply":"2023-11-30T00:36:27.845888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset & DataModule","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport random\nfrom PIL import Image\nfrom torch.utils.data import Dataset\nfrom typing import List\n\nclass CancerTilesDataset(Dataset):\n    split: float = 0.90\n\n    def __init__(\n        self,\n        img_df,\n        tile_paths: List[str] = [],\n        transforms = None,\n        mode: str = 'train',\n        labels_lut = None,\n    ):\n        self.transforms = transforms\n        self.mode = mode\n        self.labels_unique = sorted(img_df[\"label\"].unique())\n        self.labels_lut = labels_lut or {lb: i for i, lb in enumerate(self.labels_unique)}\n\n        # split images\n        assert 0.0 <= self.split <= 1.0\n        # Stratified sampling by labels\n        img_df_train = img_df.groupby('label', group_keys=False).apply(lambda x: x.sample(frac=self.split, random_state=42))\n        img_df_test = img_df[~img_df['image_id'].isin(img_df_train['image_id'])]\n        self.img_df = img_df_train if mode == 'train' else img_df_test\n        \n        # Get tile_paths and labels within the split\n        self.tile_paths = []\n        self.labels = []\n        for tile_path in tile_paths:\n            image_id = tile_path.split('/')[-4]\n            if int(image_id) in self.img_df['image_id'].values:\n                self.tile_paths.append(tile_path)\n                label = self.img_df.loc[self.img_df['image_id'] == int(image_id), \"label\"].item()\n                self.labels.append(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        tile_path = self.tile_paths[idx]\n        tile = np.array(Image.open(tile_path))[..., :3]\n        black_bg = np.sum(tile, axis=2) == 0\n        tile[black_bg, :] = 255\n        \n        label = self.to_one_hot(self.labels[idx])\n\n        # augmentation\n        if self.transforms:\n            tile = self.transforms(Image.fromarray(tile))\n        return tile, torch.tensor(label).to(int)\n\n    def __len__(self) -> int:\n        return len(self.tile_paths)\n\n# ==============================\n# ==============================\n\ndataset = CancerTilesDataset(df_train, DATASET_IMAGES)\n\n# quick view\nnp.random.seed(42)\nfig, axes = plt.subplots(nrows=3, ncols=3, figsize=(10, 10))\nfor i, img_idx in enumerate(np.random.randint(len(dataset), size=9)):\n    img, lb = dataset[img_idx]\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-11-30T00:36:27.84982Z","iopub.execute_input":"2023-11-30T00:36:27.850843Z","iopub.status.idle":"2023-11-30T00:36:38.17415Z","shell.execute_reply.started":"2023-11-30T00:36:27.850797Z","shell.execute_reply":"2023-11-30T00:36:38.172863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_= pd.DataFrame(dataset.labels).value_counts().plot.pie(autopct='%1.1f%%', ylabel=\"label\", figsize=(3,3), title=\"Tile-level labels\")","metadata":{"execution":{"iopub.status.busy":"2023-11-30T00:36:38.175385Z","iopub.execute_input":"2023-11-30T00:36:38.175791Z","iopub.status.idle":"2023-11-30T00:36:38.300137Z","shell.execute_reply.started":"2023-11-30T00:36:38.175764Z","shell.execute_reply":"2023-11-30T00:36:38.299234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let us define some standard image augmentation 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(TILE_SIZE),\n    #T.RandomResizedCrop(TILE_SIZE, interpolation=InterpolationMode.BICUBIC, antialias=True),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\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(TILE_SIZE),\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-11-30T00:36:38.301388Z","iopub.execute_input":"2023-11-30T00:36:38.30216Z","iopub.status.idle":"2023-11-30T00:36:38.731916Z","shell.execute_reply.started":"2023-11-30T00:36:38.302124Z","shell.execute_reply":"2023-11-30T00:36:38.731079Z"},"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 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        img_df,\n        tile_paths: List[str] = [],\n        batch_size: int = 32,\n        num_workers: int = None,\n        train_transforms = TRAIN_TRANSFORM,\n        valid_transforms = VALID_TRANSFORM\n    ):\n        super().__init__()\n        self.img_df = img_df\n        self.tile_paths = tile_paths\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\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.img_df, self.tile_paths, mode='train', transforms=self.train_transforms)\n        print(f\"training dataset: {len(self.train_dataset)}\")\n        self.valid_dataset = CancerTilesDataset(\n            self.img_df, self.tile_paths, mode='valid', transforms=self.valid_transforms,\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(df_train, DATASET_IMAGES, batch_size=BATCH_SIZE)\ndm.setup()\nprint(f\"Number of classes = {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: {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        ax.imshow(np.rollaxis(imgs[i].numpy(), 0, 3))\n        ax.set_title(lbs[i])\n    break","metadata":{"execution":{"iopub.status.busy":"2023-11-30T00:36:38.73476Z","iopub.execute_input":"2023-11-30T00:36:38.735043Z","iopub.status.idle":"2023-11-30T00:36:48.602524Z","shell.execute_reply.started":"2023-11-30T00:36:38.735018Z","shell.execute_reply":"2023-11-30T00:36:48.601482Z"},"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-11-30T00:36:48.603972Z","iopub.execute_input":"2023-11-30T00:36:48.604435Z","iopub.status.idle":"2023-11-30T00:37:02.249755Z","shell.execute_reply.started":"2023-11-30T00:36:48.604404Z","shell.execute_reply":"2023-11-30T00:37:02.248663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\nimport torch\nimport torchvision\nfrom adan_pytorch import Adan\nfrom lion_pytorch import Lion\nfrom torch_optimizer import AdaBound, RAdam, Yogi\n\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torchmetrics.classification import MulticlassAccuracy, MulticlassF1Score\n\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\n    def forward(self, x):\n        y = F.softmax(self.net(x))\n#         if y.isnan().any():\n#             y = torch.ones(self.num_classes) / self.num_classes\n        return y\n\n    def compute_loss(self, y_hat, y):\n        return F.cross_entropy(y_hat, y)\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        y_hat = self.net(x)\n        y = torch.argmax(y, axis=1)\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#         print(f\"{y=} ?= {y_hat=} -> {self.train_accuracy(y_hat, y)}\")\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.net(x)\n        y = torch.argmax(y, axis=1)\n        loss = self.compute_loss(y_hat, y)\n        self.log(\"valid_loss\", loss, logger=True, prog_bar=False)\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 * 10,\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\n# see: https://pytorch.org/vision/stable/models.html\n# net = timm.create_model('resnet50', pretrained=True, num_classes=dm.num_classes)\n# model = LitCancerSubtype(net=net, lr=1e-4)\nnet = timm.create_model('maxvit_rmlp_base_rw_224.sw_in12k_ft_in1k', pretrained=True, num_classes=dm.num_classes)\nmodel = LitCancerSubtype(net=net, lr=2e-5)\n# print(model)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-11-30T00:37:02.251358Z","iopub.execute_input":"2023-11-30T00:37:02.251693Z","iopub.status.idle":"2023-11-30T00:37:38.418847Z","shell.execute_reply.started":"2023-11-30T00:37:02.251662Z","shell.execute_reply":"2023-11-30T00:37:38.417972Z"},"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":"csv_logger = pl.loggers.CSVLogger(save_dir='logs/', name=model.arch)\nnb_epochs = 20 if torch.cuda.is_available() else 2\n\n# ==============================\n\ntrainer = pl.Trainer(\n    accelerator=\"gpu\",\n#     devices=2,\n#     fast_dev_run=True,\n    # callbacks=[swa],\n    logger=csv_logger,\n    max_epochs=nb_epochs,\n    max_time=\"00:05:00:00\",\n#     max_epochs=10,\n#     max_steps=1,\n    precision=\"16-mixed\",\n    accumulate_grad_batches=7,\n#     overfit_batches=10,\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-11-30T00:37:38.420073Z","iopub.execute_input":"2023-11-30T00:37:38.420731Z","iopub.status.idle":"2023-11-30T00:38:48.176403Z","shell.execute_reply.started":"2023-11-30T00:37:38.420695Z","shell.execute_reply":"2023-11-30T00:38:48.175149Z"},"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()\ndisplay(metrics.tail(1))","metadata":{"execution":{"iopub.status.busy":"2023-11-30T00:38:48.178189Z","iopub.execute_input":"2023-11-30T00:38:48.178589Z","iopub.status.idle":"2023-11-30T00:38:49.803784Z","shell.execute_reply.started":"2023-11-30T00:38:48.178548Z","shell.execute_reply":"2023-11-30T00:38:49.80215Z"},"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":{"execution":{"iopub.status.busy":"2023-11-30T00:38:49.804959Z","iopub.status.idle":"2023-11-30T00:38:49.805486Z","shell.execute_reply.started":"2023-11-30T00:38:49.805227Z","shell.execute_reply":"2023-11-30T00:38:49.80525Z"},"trusted":true},"execution_count":null,"outputs":[]}]}