{"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":7029230,"sourceType":"datasetVersion","datasetId":4027203}],"dockerImageVersionId":30588,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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","execution":{"iopub.status.busy":"2023-12-05T02:55:30.528344Z","iopub.execute_input":"2023-12-05T02:55:30.528969Z","iopub.status.idle":"2023-12-05T02:55:30.876425Z","shell.execute_reply.started":"2023-12-05T02:55:30.528924Z","shell.execute_reply":"2023-12-05T02:55:30.875687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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[:9000]))","metadata":{"execution":{"iopub.status.busy":"2023-12-05T02:55:30.877967Z","iopub.execute_input":"2023-12-05T02:55:30.87834Z","iopub.status.idle":"2023-12-05T02:57:01.169071Z","shell.execute_reply.started":"2023-12-05T02:55:30.878314Z","shell.execute_reply":"2023-12-05T02:57:01.168122Z"},"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=}\")\ndf_train = pd.read_csv(os.path.join(DATASET_FOLDER, \"train.csv\"))","metadata":{"execution":{"iopub.status.busy":"2023-12-05T02:57:01.170302Z","iopub.execute_input":"2023-12-05T02:57:01.170593Z","iopub.status.idle":"2023-12-05T02:57:01.269639Z","shell.execute_reply.started":"2023-12-05T02:57:01.170565Z","shell.execute_reply":"2023-12-05T02:57:01.268797Z"},"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    split: float = 0.90\n\n    def __init__(\n        self,\n        df_data,\n        path_img_dir: str,\n        path_mask_dir: str,\n        transforms = None,\n        mode: str = 'train',\n        labels_lut = None,\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.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.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                # 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        labels = self.to_one_hot(lb)\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.imgs)\n\n\ndataset = CancerTilesDataset(df_train, DATASET_IMAGES, DATASET_MASKS)","metadata":{"execution":{"iopub.status.busy":"2023-12-05T02:57:01.271788Z","iopub.execute_input":"2023-12-05T02:57:01.27209Z","iopub.status.idle":"2023-12-05T02:57:04.582346Z","shell.execute_reply.started":"2023-12-05T02:57:01.272066Z","shell.execute_reply":"2023-12-05T02:57:04.581471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import transforms as T\nfrom torchvision.transforms import InterpolationMode\n\nTRAIN_TRANSFORM = T.Compose([\n    T.CenterCrop(384),\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(384),\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-05T02:57:04.583448Z","iopub.execute_input":"2023-12-05T02:57:04.583847Z","iopub.status.idle":"2023-12-05T02:57:04.992036Z","shell.execute_reply.started":"2023-12-05T02:57:04.583822Z","shell.execute_reply":"2023-12-05T02:57:04.990993Z"},"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    ):\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\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)\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,\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\ndm = CancerSubtypeDM(df_train, DATASET_IMAGES, DATASET_MASKS, batch_size=32)\ndm.setup()\nprint(dm.num_classes)\n","metadata":{"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2023-12-05T02:57:04.993254Z","iopub.execute_input":"2023-12-05T02:57:04.993545Z","iopub.status.idle":"2023-12-05T02:57:08.729127Z","shell.execute_reply.started":"2023-12-05T02:57:04.993521Z","shell.execute_reply":"2023-12-05T02:57:08.72799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q lion-pytorch adan-pytorch torch_optimizer","metadata":{"execution":{"iopub.status.busy":"2023-12-05T02:57:08.730782Z","iopub.execute_input":"2023-12-05T02:57:08.731687Z","iopub.status.idle":"2023-12-05T02:57:21.434757Z","shell.execute_reply.started":"2023-12-05T02:57:08.731646Z","shell.execute_reply":"2023-12-05T02:57:21.433529Z"},"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\nimport pytorch_lightning as pl\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), dim=1)\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.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 = 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]\nnet = timm.create_model('vit_base_patch16_384.augreg_in21k_ft_in1k', 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=2e-4)\n# print(model)","metadata":{"execution":{"iopub.status.busy":"2023-12-05T02:57:21.436345Z","iopub.execute_input":"2023-12-05T02:57:21.43663Z","iopub.status.idle":"2023-12-05T02:57:28.327431Z","shell.execute_reply.started":"2023-12-05T02:57:21.436603Z","shell.execute_reply":"2023-12-05T02:57:28.326598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"logger = pl.loggers.CSVLogger(save_dir='logs/', name=model.arch)\nnb_epochs = 50 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-mixed\",\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":{"execution":{"iopub.status.busy":"2023-12-05T02:57:28.328755Z","iopub.execute_input":"2023-12-05T02:57:28.329456Z","iopub.status.idle":"2023-12-05T02:57:47.033228Z","shell.execute_reply.started":"2023-12-05T02:57:28.329401Z","shell.execute_reply":"2023-12-05T02:57:47.031769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-05T02:57:47.036482Z","iopub.execute_input":"2023-12-05T02:57:47.036851Z","iopub.status.idle":"2023-12-05T02:57:48.256031Z","shell.execute_reply.started":"2023-12-05T02:57:47.036817Z","shell.execute_reply":"2023-12-05T02:57:48.254541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.save_checkpoint(\"image_classification_model.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-12-05T02:57:48.257149Z","iopub.status.idle":"2023-12-05T02:57:48.257829Z","shell.execute_reply.started":"2023-12-05T02:57:48.257639Z","shell.execute_reply":"2023-12-05T02:57:48.257658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n","metadata":{},"execution_count":null,"outputs":[]}]}