{"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":151304386,"sourceType":"kernelVersion"}],"dockerImageVersionId":30558,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Cancer🔬 Classification: Baseline with ⚡`lightning`\n\n**It is continuation of EDA: https://www.kaggle.com/code/jirkaborovec/cancer-subtype-explore-data-images**","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 = \"/kaggle/input/cancer-subtype-eda-load-wsi-segm-mask/train_thumbnails/\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-01T22:44:41.670936Z","iopub.execute_input":"2023-12-01T22:44:41.671219Z","iopub.status.idle":"2023-12-01T22:44:42.003891Z","shell.execute_reply.started":"2023-12-01T22:44:41.671193Z","shell.execute_reply":"2023-12-01T22:44:42.003109Z"},"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\"))\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-01T22:44:42.005453Z","iopub.execute_input":"2023-12-01T22:44:42.005829Z","iopub.status.idle":"2023-12-01T22:44:42.034131Z","shell.execute_reply.started":"2023-12-01T22:44:42.005803Z","shell.execute_reply":"2023-12-01T22:44:42.033216Z"},"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-01T22:44:42.03534Z","iopub.execute_input":"2023-12-01T22:44:42.035669Z","iopub.status.idle":"2023-12-01T22:44:42.223598Z","shell.execute_reply.started":"2023-12-01T22:44:42.035636Z","shell.execute_reply":"2023-12-01T22:44:42.222369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Show some samples 🖼️ per class\n\nNote that not all images has thumbnails","metadata":{}},{"cell_type":"code","source":"df_train[\"path_thumbnail\"] = df_train['image_id'].apply(lambda id: f\"{id}_crop.png\")\ndf_train[\"thumbnail_exists\"] = df_train['path_thumbnail'].apply(\n    lambda name: os.path.isfile(os.path.join(DATASET_IMAGES, name)))\ndisplay(df_train.head())\ndf_train = df_train[df_train[\"thumbnail_exists\"] == True]\nprint(f\"size: {len(df_train)}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-01T22:44:42.227296Z","iopub.execute_input":"2023-12-01T22:44:42.228126Z","iopub.status.idle":"2023-12-01T22:44:42.771162Z","shell.execute_reply.started":"2023-12-01T22:44:42.228075Z","shell.execute_reply":"2023-12-01T22:44:42.769854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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        img_path = os.path.join(DATASET_IMAGES, f\"{img_ids[i]}_crop.png\")\n        img = plt.imread(img_path)\n        axarr[ilb, i].imshow(img)\n        if i == 0:\n            axarr[ilb, i].set_title(f\"{lb} #{len(df_)}\")\n        axarr[ilb, i].set_xticks([])\n        axarr[ilb, i].set_yticks([])\n_= plt.axis('off')","metadata":{"execution":{"iopub.status.busy":"2023-12-01T22:44:42.772628Z","iopub.execute_input":"2023-12-01T22:44:42.773526Z","iopub.status.idle":"2023-12-01T22:45:17.004337Z","shell.execute_reply.started":"2023-12-01T22:44:42.773481Z","shell.execute_reply":"2023-12-01T22:45:17.003416Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data pre-processing\n\n### Sacling down images\n\nWe could not fit this huge image size to any NEt so just as offline process lets downscale it to about 1024x1024","metadata":{}},{"cell_type":"code","source":"def prune_image_rows_cols(img, thr=0.001):\n    # delete empty columns\n    for l in reversed(range(img.shape[1])):\n        if (np.sum(img[:, l]) / float(img.shape[0])) < thr:\n            img = np.delete(img, l, 1)\n    # delete empty rows\n    for l in reversed(range(img.shape[0])):\n        if (np.sum(img[l, :]) / float(img.shape[1])) < thr:\n            img = np.delete(img, l, 0)\n    return img","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-01T22:45:17.00535Z","iopub.execute_input":"2023-12-01T22:45:17.005613Z","iopub.status.idle":"2023-12-01T22:45:17.012131Z","shell.execute_reply.started":"2023-12-01T22:45:17.00559Z","shell.execute_reply":"2023-12-01T22:45:17.011155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prune_scale_image(img_path: str, out_dir: str, size: int = 1024) -> None:\n    img = np.array(Image.open(img_path))\n    img = prune_image_rows_cols(img)\n    mask = np.sum(img[..., :3], axis=2) == 0\n    img[mask, :] = 255\n    img = Image.fromarray(img)\n    img.thumbnail((size, size))\n    img.save(os.path.join(out_dir, os.path.basename(img_path)))","metadata":{"execution":{"iopub.status.busy":"2023-12-01T22:45:17.013548Z","iopub.execute_input":"2023-12-01T22:45:17.014137Z","iopub.status.idle":"2023-12-01T22:45:17.029763Z","shell.execute_reply.started":"2023-12-01T22:45:17.014105Z","shell.execute_reply":"2023-12-01T22:45:17.028753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nfrom joblib import Parallel, delayed\n\n! mkdir -p train_thumbnails\n! rm -f train_thumbnails/*.png\n\nls = glob.glob(os.path.join(DATASET_IMAGES, '*.png'))\nprint(f\"found images: {len(ls)}\")\n\n# for p_img in tqdm(ls):\n#     prune_scale_image(p_img, \"./train_thumbnails\")\n    \n_= Parallel(n_jobs=4)(\n    delayed(prune_scale_image)(p_img, \"./train_thumbnails\") for p_img in tqdm(ls)\n)\nls = glob.glob(os.path.join(\"./train_thumbnails\", '*.png'))\nprint(f\"found images: {len(ls)}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-01T22:45:17.03099Z","iopub.execute_input":"2023-12-01T22:45:17.031306Z","iopub.status.idle":"2023-12-01T22:48:28.110578Z","shell.execute_reply.started":"2023-12-01T22:45:17.031279Z","shell.execute_reply":"2023-12-01T22:48:28.109532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Color 🦩 normalizations","metadata":{}},{"cell_type":"code","source":"def _color_means(img_path):\n    img = np.array(Image.open(img_path))\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(\"./train_thumbnails\", \"*.png\"))\nclr_mean_std = Parallel(n_jobs=os.cpu_count())(delayed(_color_means)(fn) for fn in tqdm(ls_images))","metadata":{"execution":{"iopub.status.busy":"2023-12-01T22:48:28.112021Z","iopub.execute_input":"2023-12-01T22:48:28.112314Z","iopub.status.idle":"2023-12-01T22:48:38.910922Z","shell.execute_reply.started":"2023-12-01T22:48:28.11229Z","shell.execute_reply":"2023-12-01T22:48:38.910138Z"},"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-01T22:48:38.914475Z","iopub.execute_input":"2023-12-01T22:48:38.914789Z","iopub.status.idle":"2023-12-01T22:48:38.960331Z","shell.execute_reply.started":"2023-12-01T22:48:38.914764Z","shell.execute_reply":"2023-12-01T22:48:38.959484Z"},"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":"import os\nimport torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset\n\nclass CancerThumbnailDataset(Dataset):\n    split: float = 0.90\n\n    def __init__(\n        self,\n        df_data,\n        path_img_dir: str =  'train_thumbnails',\n        transforms = None,\n        mode: str = 'train',\n        labels_lut = None\n    ):\n        self.path_img_dir = path_img_dir\n        self.transforms = transforms\n        self.mode = mode\n\n        self.data = df_data\n        self.labels_unique = sorted(self.data[\"label\"].unique())\n        self.labels_lut = labels_lut or {lb: i for i, lb in enumerate(self.labels_unique)}\n        # shuffle data\n        self.data = self.data.sample(frac=1, random_state=42).reset_index(drop=True)\n\n        # split dataset\n        assert 0.0 <= self.split <= 1.0\n        frac = int(self.split * len(self.data))\n        self.data = self.data[:frac] if mode == 'train' else self.data[frac:]\n        self.img_names = [f\"{id}_crop.png\" for id in self.data[\"image_id\"]]\n        #print(f\"missing: {sum([not os.path.isfile(os.path.join(self.path_img_dir, im))\n        #                       for im in self.img_names])}\")\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        img_path = os.path.join(self.path_img_dir, self.img_names[idx])\n        assert os.path.isfile(img_path), f\"missing: {img_path}\"\n        img = plt.imread(img_path)[..., :3]\n        if np.max(img) < 1.5:\n            img = np.clip(img * 255, 0, 255).astype(np.uint8)\n        labels = self.to_one_hot(self.labels[idx])\n\n        # augmentation\n        if self.transforms:\n            img = self.transforms(Image.fromarray(img))\n        #print(f\"img dim: {img.shape}\")\n        return img, torch.tensor(labels).to(int)\n\n    def __len__(self) -> int:\n        return len(self.data)\n\n# ==============================\n# ==============================\n\ndataset = CancerThumbnailDataset(df_train)\n\n# quick view\nfig = plt.figure(figsize=(10, 8))\nfor i in range(9):\n    img, lb = dataset[i]\n    ax = fig.add_subplot(3, 3, i + 1, xticks=[], yticks=[])\n    ax.imshow(img)\n    ax.set_title(f\"{lb}\\n dims: {img.shape} /{np.max(img)}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-01T22:48:38.961956Z","iopub.execute_input":"2023-12-01T22:48:38.962597Z","iopub.status.idle":"2023-12-01T22:48:44.275123Z","shell.execute_reply.started":"2023-12-01T22:48:38.96256Z","shell.execute_reply":"2023-12-01T22:48:44.274235Z"},"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.RandomRotation(45, fill=255),\n    T.RandomCrop(512, pad_if_needed=True, padding_mode=\"reflect\"),\n    #T.RandomResizedCrop(512, interpolation=InterpolationMode.BICUBIC, antialias=True),\n    T.RandomHorizontalFlip(),\n    T.RandomVerticalFlip(),\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-01T22:48:44.276454Z","iopub.execute_input":"2023-12-01T22:48:44.276984Z","iopub.status.idle":"2023-12-01T22:48:44.558664Z","shell.execute_reply.started":"2023-12-01T22:48:44.276951Z","shell.execute_reply":"2023-12-01T22:48:44.557752Z"},"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        df_data,\n        path_img_dir: str = 'train_thumbnails',\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.df_data = df_data\n        self.path_img_dir = path_img_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\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 = CancerThumbnailDataset(\n            self.df_data, self.path_img_dir, mode='train', transforms=self.train_transforms)\n        print(f\"training dataset: {len(self.train_dataset)}\")\n        self.valid_dataset = CancerThumbnailDataset(\n            self.df_data, self.path_img_dir, 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, batch_size=8)\ndm.setup()\nprint(dm.num_classes)\n\n# quick view\nfig = plt.figure(figsize=(3, 7))\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, xticks=[], yticks=[])\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-01T22:48:44.562288Z","iopub.execute_input":"2023-12-01T22:48:44.562569Z","iopub.status.idle":"2023-12-01T22:48:56.503905Z","shell.execute_reply.started":"2023-12-01T22:48:44.562544Z","shell.execute_reply":"2023-12-01T22:48:56.502865Z"},"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-01T22:48:56.505394Z","iopub.execute_input":"2023-12-01T22:48:56.505737Z","iopub.status.idle":"2023-12-01T22:49:09.882935Z","shell.execute_reply.started":"2023-12-01T22:48:56.505685Z","shell.execute_reply":"2023-12-01T22:49:09.881809Z"},"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 import nn\nfrom torch.nn import functional as F\nfrom torch_optimizer import AdaBound, RAdam, Yogi\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_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        y_hat = self(x)\n        lbs = torch.argmax(y, axis=1)\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        #print(f\"{lb=} ?= {y_hat=} -> {self.train_accuracy(y_hat, lbs)}\")\n        self.log(\"train_acc\", self.train_accuracy(y_hat, lbs), 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        lbs = 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, lbs), logger=True, prog_bar=False)\n        self.log(\"valid_f1\", self.val_f1_score(y_hat, lbs), logger=True, prog_bar=True)\n\n    def configure_optimizers(self):\n        #optimizer = AdaBound(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, betas=(0.02, 0.08, 0.01), weight_decay=0.02)\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\n# see: https://pytorch.org/vision/stable/models.html\nnet = timm.create_model('tf_efficientnetv2_s_in21ft1k', 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-01T22:49:09.884519Z","iopub.execute_input":"2023-12-01T22:49:09.884846Z","iopub.status.idle":"2023-12-01T22:49:17.076306Z","shell.execute_reply.started":"2023-12-01T22:49:09.884818Z","shell.execute_reply":"2023-12-01T22:49:17.075256Z"},"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 = 60 if torch.cuda.is_available() else 2\n\n# ==============================\n\ntrainer = pl.Trainer(\n    # fast_dev_run=True,\n    # callbacks=[swa],\n    logger=logger,\n    max_epochs=nb_epochs,\n    precision=16,\n    accumulate_grad_batches=4,\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-01T22:49:17.077712Z","iopub.execute_input":"2023-12-01T22:49:17.078075Z","iopub.status.idle":"2023-12-01T22:50:14.400031Z","shell.execute_reply.started":"2023-12-01T22:49:17.078041Z","shell.execute_reply":"2023-12-01T22:50:14.394533Z"},"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":{"execution":{"iopub.status.busy":"2023-12-01T22:50:14.402302Z","iopub.execute_input":"2023-12-01T22:50:14.407755Z","iopub.status.idle":"2023-12-01T22:50:15.438577Z","shell.execute_reply.started":"2023-12-01T22:50:14.407678Z","shell.execute_reply":"2023-12-01T22:50:15.437637Z"},"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-12-01T22:50:15.43975Z","iopub.execute_input":"2023-12-01T22:50:15.440029Z","iopub.status.idle":"2023-12-01T22:50:16.251626Z","shell.execute_reply.started":"2023-12-01T22:50:15.440005Z","shell.execute_reply":"2023-12-01T22:50:16.25084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Inference coming in https://www.kaggle.com/code/jirkaborovec/cancer-subtype-lightning-torch-inference","metadata":{}}]}