{"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":6774400,"sourceType":"datasetVersion","datasetId":3895136},{"sourceId":6774553,"sourceType":"datasetVersion","datasetId":3898019},{"sourceId":7036947,"sourceType":"datasetVersion","datasetId":3983823},{"sourceId":151484971,"sourceType":"kernelVersion"}],"dockerImageVersionId":30559,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Cancer🔬 Classification: Inference with ⚡`lightning`\n\n**It is continuation of Training: https://www.kaggle.com/code/jirkaborovec/cancer-subtype-tiles-w-lightning-timm-models**","metadata":{}},{"cell_type":"code","source":"!ls /kaggle/input/pyvips-python-and-deb-package-gpu\n# intall the deb packages\n!yes | dpkg -i --force-depends /kaggle/input/pyvips-python-and-deb-package-gpu/linux_packages/archives/*.deb\n# install the python wrapper\n!pip install pyvips -f /kaggle/input/pyvips-python-and-deb-package-gpu/python_packages/ --no-index","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-23T16:10:32.187031Z","iopub.execute_input":"2023-11-23T16:10:32.188045Z","iopub.status.idle":"2023-11-23T16:11:35.789253Z","shell.execute_reply.started":"2023-11-23T16:10:32.187996Z","shell.execute_reply":"2023-11-23T16:11:35.788038Z"},"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/\"\nIMAGES_FOLDER = \"./test_tiles\"\n\nos.environ['VIPS_CONCURRENCY'] = '4'\nos.environ['VIPS_DISC_THRESHOLD'] = '15gb'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-23T16:11:35.791521Z","iopub.execute_input":"2023-11-23T16:11:35.791856Z","iopub.status.idle":"2023-11-23T16:11:35.797935Z","shell.execute_reply.started":"2023-11-23T16:11:35.791824Z","shell.execute_reply":"2023-11-23T16:11:35.796882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data for inference\n\n**Note, we canot extract all tiles from all imges because of unsufficient space/storage**\n\nThis needs porting several code (classes) fromm the training notebook:\n\n- extracting the tiles from whole image\n- validation augmention, mainly color mean & STD","metadata":{}},{"cell_type":"code","source":"df_train = pd.read_csv(os.path.join(DATASET_FOLDER, \"train.csv\"))\nlabels = sorted(df_train[\"label\"].unique())\nprint(f\"{labels=}\")","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-23T16:11:35.799207Z","iopub.execute_input":"2023-11-23T16:11:35.799595Z","iopub.status.idle":"2023-11-23T16:11:35.848043Z","shell.execute_reply.started":"2023-11-23T16:11:35.799561Z","shell.execute_reply":"2023-11-23T16:11:35.847065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pyvips\nimport numpy as np\nimport random\nfrom PIL import Image\n\ndef extract_image_tiles(\n    p_img, folder, size: int = 2048, scale: float = 0.5,\n    drop_thr: float = 0.6, white_thr: int = 240, max_samples: int = 50\n) -> list:\n    name, _ = os.path.splitext(os.path.basename(p_img))\n    im = pyvips.Image.new_from_file(p_img)\n    w = h = size\n    # https://stackoverflow.com/a/47581978/4521646\n    idxs = [(y, y + h, x, x + w) for y in range(0, im.height, h) for x in range(0, im.width, w)]\n    # random subsample\n    max_samples = max_samples if isinstance(max_samples, int) else int(len(idxs) * max_samples)\n    random.shuffle(idxs)\n    files = []\n    for y, y_, x, x_ in idxs:\n        # https://libvips.github.io/pyvips/vimage.html#pyvips.Image.crop\n        tile = im.crop(x, y, min(w, im.width - x), min(h, im.height - y)).numpy()[..., :3]\n        if tile.shape[:2] != (h, w):\n            tile_ = tile\n            tile_size = (h, w) if tile.ndim == 2 else (h, w, tile.shape[2])\n            tile = np.zeros(tile_size, dtype=tile.dtype)\n            tile[:tile_.shape[0], :tile_.shape[1], ...] = tile_\n        black_bg = np.sum(tile, axis=2) == 0\n        tile[black_bg, :] = 255\n        mask_bg = np.mean(tile, axis=2) > white_thr\n        if np.sum(mask_bg) >= (np.prod(mask_bg.shape) * drop_thr):\n            #print(f\"skip almost empty tile: {k:06}_{int(x_ / w)}-{int(y_ / h)}\")\n            continue\n        p_img = os.path.join(folder, f\"{int(x_ / w)}-{int(y_ / h)}.png\")\n        # print(tile.shape, tile.dtype, tile.min(), tile.max())\n        new_size = int(size * scale), int(size * scale)\n        Image.fromarray(tile).resize(new_size, Image.LANCZOS).save(p_img)\n        files.append(p_img)\n        # need to set counter check as some empty tiles could be skipped earlier\n        if len(files) >= max_samples:\n            break\n    return files","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-23T16:11:35.851382Z","iopub.execute_input":"2023-11-23T16:11:35.851703Z","iopub.status.idle":"2023-11-23T16:11:35.865093Z","shell.execute_reply.started":"2023-11-23T16:11:35.851677Z","shell.execute_reply":"2023-11-23T16:11:35.864164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def extract_prune_tiles(\n    path_img: str, folder: str, size: int = 2048, scale: float = 0.25,\n    drop_thr: float = 0.6, max_samples: int = 30\n) -> str:\n    print(f\"processing: {path_img}\")\n    name, _ = os.path.splitext(os.path.basename(path_img))\n    folder = os.path.join(folder, name)\n    os.makedirs(folder, exist_ok=True)\n    tiles = extract_image_tiles(\n        path_img, folder, size=size, scale=scale,\n        drop_thr=drop_thr, max_samples=max_samples)\n    return folder","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-23T16:11:35.866353Z","iopub.execute_input":"2023-11-23T16:11:35.866721Z","iopub.status.idle":"2023-11-23T16:11:35.877905Z","shell.execute_reply.started":"2023-11-23T16:11:35.866688Z","shell.execute_reply":"2023-11-23T16:11:35.877081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import transforms as T\n\nimg_color_mean = [0.8721593659261734, 0.7799686061900686, 0.8644588534918227]\nimg_color_std = [0.08258995918115268, 0.10991684444009092, 0.06839816226731532]\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-11-23T16:11:35.87903Z","iopub.execute_input":"2023-11-23T16:11:35.879371Z","iopub.status.idle":"2023-11-23T16:11:35.891512Z","shell.execute_reply.started":"2023-11-23T16:11:35.879329Z","shell.execute_reply":"2023-11-23T16:11:35.890634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset\n\nclass TilesFolderDataset(Dataset):\n\n    def __init__(\n        self,\n        folder: str,\n        image_ext: str =  '.png',\n        transforms = None\n    ):\n        assert os.path.isdir(folder)\n        self.transforms = transforms\n        self.imgs = glob.glob(os.path.join(folder, \"*\" + image_ext))\n\n    def __getitem__(self, idx: int) -> tuple:\n        img_path = self.imgs[idx]\n        assert os.path.isfile(img_path), f\"missing: {img_path}\"\n        img = np.array(Image.open(img_path))[..., :3]\n        # filter background\n        mask = np.sum(img, axis=2) == 0\n        img[mask, :] = 255\n        if np.max(img) < 1.5:\n            img = np.clip(img * 255, 0, 255).astype(np.uint8)\n        # augmentation\n        if self.transforms:\n            img = self.transforms(Image.fromarray(img))\n        #print(f\"img dim: {img.shape}\")\n        return img\n\n    def __len__(self) -> int:\n        return len(self.imgs)","metadata":{"execution":{"iopub.status.busy":"2023-11-23T16:11:35.892499Z","iopub.execute_input":"2023-11-23T16:11:35.892795Z","iopub.status.idle":"2023-11-23T16:11:35.903293Z","shell.execute_reply.started":"2023-11-23T16:11:35.892744Z","shell.execute_reply":"2023-11-23T16:11:35.902354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ls = sorted(glob.glob(os.path.join(DATASET_FOLDER, \"test_images\", '*.png')))\nprint(f\"found images: {len(ls)}\")\nfolder_tiles = extract_prune_tiles(ls[0], IMAGES_FOLDER, size=2048, scale=0.25)\ndataset = TilesFolderDataset(folder_tiles)\nprint(f\"found tiles: {len(dataset)}\")\n\n# quick view\nfig, axes = plt.subplots(nrows=3, ncols=3, figsize=(10, 10))\nfor i in range(9):\n    img = dataset[i]\n    axes[i // 3, i % 3].imshow(img)\nfig.tight_layout()","metadata":{"execution":{"iopub.status.busy":"2023-11-23T16:11:35.904487Z","iopub.execute_input":"2023-11-23T16:11:35.904788Z","iopub.status.idle":"2023-11-23T16:12:13.777519Z","shell.execute_reply.started":"2023-11-23T16:11:35.904758Z","shell.execute_reply":"2023-11-23T16:12:13.776504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"import timm\nimport torch\nimport torchvision\nimport pytorch_lightning as pl\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom pathlib import Path\nimport torchvision.models as models\n\n# Load the EfficientNet model with pre-trained weights\n#net = timm.create_model('efficientnet_b1', pretrained=False)\n#net.load_state_dict(torch.load(weights_path))\n\nclass LitCancerSubtype(pl.LightningModule):\n\n    def __init__(self, net):\n        super().__init__()\n        self.net = net\n        self.arch = net.pretrained_cfg.get('architecture')\n        self.num_classes = net.num_classes\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# ==============================\n# ==============================\n\nPATH_CKPT = (\n    \"/kaggle/input/cancer-subtype-tiles-w-lightning-timm-models/\"\n    \"image_classification_model.pt\"\n)\n\nckpt = torch.load(PATH_CKPT, map_location=torch.device('cpu'))\n\n# see: https://pytorch.org/vision/stable/models.html\nnet = timm.create_model(\n    'maxvit_tiny_tf_512', pretrained=False, num_classes=len(labels))\nmodel = LitCancerSubtype(net=net)\nmodel.load_state_dict(ckpt['state_dict'], strict=False)\nprint(model)","metadata":{"execution":{"iopub.status.busy":"2023-11-21T06:03:20.270579Z","iopub.status.idle":"2023-11-21T06:03:20.271063Z","shell.execute_reply.started":"2023-11-21T06:03:20.270832Z","shell.execute_reply":"2023-11-21T06:03:20.270853Z"}}},{"cell_type":"markdown","source":"## CNN Model\n\nWe start with some stanrd CNN models taken from torch vision.","metadata":{}},{"cell_type":"code","source":"import timm\nimport torch\nimport torchvision\nimport pytorch_lightning as pl\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom pathlib import Path\nimport torchvision.models as models\n\nclass LitCancerSubtype(pl.LightningModule):\n\n    def __init__(self, net):\n        super().__init__()\n        self.net = net\n        self.arch = net.pretrained_cfg.get('architecture')\n        self.num_classes = net.num_classes\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# ==============================\n# ==============================\nimport io\n\nbuffer = io.BytesIO()\n\nweights_path = Path(\n    \"/kaggle/input/efficientnet/maxvit_xlarge_tf_512.in21k_ft_in1k.pt\"\n)\nweight = torch.load(weights_path, map_location=torch.device('cpu'))\ntorch.save(weight.state_dict(), buffer)\nbuffer.seek(0)\nweight_load = torch.load(buffer)\n\n\nnet = timm.create_model(\n    'maxvit_xlarge_tf_512', pretrained = False, num_classes=len(labels))\nmodel = LitCancerSubtype(net=net)\nmodel.load_state_dict(weight_load, strict=False)\nprint(model)","metadata":{"execution":{"iopub.status.busy":"2023-11-23T17:47:14.359158Z","iopub.execute_input":"2023-11-23T17:47:14.359594Z","iopub.status.idle":"2023-11-23T17:47:42.471211Z","shell.execute_reply.started":"2023-11-23T17:47:14.359559Z","shell.execute_reply":"2023-11-23T17:47:42.470245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls /kaggle/input/efficientnet","metadata":{"execution":{"iopub.status.busy":"2023-11-23T17:21:04.971811Z","iopub.execute_input":"2023-11-23T17:21:04.972221Z","iopub.status.idle":"2023-11-23T17:21:06.062147Z","shell.execute_reply.started":"2023-11-23T17:21:04.97219Z","shell.execute_reply":"2023-11-23T17:21:06.060968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test data & submission\n\nlest load sample submission and add append images we can predict","metadata":{}},{"cell_type":"code","source":"df_test = pd.read_csv(os.path.join(DATASET_FOLDER, \"test.csv\"))\n# default label\ndf_test['label'] = ['HGSC'] * len(df_test)\n# labels = list(df_train[\"label\"].unique())\nprint(f\"Dataset/test size: {len(df_test)}\")\ndisplay(df_test.head())","metadata":{"execution":{"iopub.status.busy":"2023-11-23T17:48:01.610971Z","iopub.execute_input":"2023-11-23T17:48:01.611839Z","iopub.status.idle":"2023-11-23T17:48:01.627646Z","shell.execute_reply.started":"2023-11-23T17:48:01.611801Z","shell.execute_reply":"2023-11-23T17:48:01.626779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"import shutil\nfrom torch.utils.data import DataLoader\n\nmodel.eval()\nmodel = model.cuda()\n\nsubmission = []\nfor _, row in df_test.iterrows():\n    row = dict(row)\n    # prepare data - cut and load tiles\n    folder_tiles = extract_prune_tiles(\n        os.path.join(DATASET_FOLDER, \"test_images\", f\"{str(row['image_id'])}.png\"),\n        IMAGES_FOLDER, size=2048, scale=0.25)\n    dataset = TilesFolderDataset(folder_tiles, transforms=VALID_TRANSFORM)\n    if not len(dataset):\n        print (f\"seem no tiles were cut for `{folder_tiles}`\")\n        submission.append(row)\n        continue\n    dataloader = DataLoader(dataset, batch_size=4, num_workers=10, shuffle=False)\n    # iterate over images and collect predictions\n    preds = []\n    for imgs in dataloader:\n        #print(f\"{imgs.shape}\")\n        with torch.no_grad():\n            pred = model(imgs.cuda())\n        preds += pred.cpu().numpy().tolist()\n    print(f\"Sum contribution from all tiles: {np.sum(preds, axis=0)}\")\n    print(f\"Max contribution over all tiles: {np.max(preds, axis=0)}\")\n    # decide label\n    lb = np.argmax(np.sum(preds, axis=0))\n    row['label'] = labels[lb]\n    print(row)\n    submission.append(row)\n    # cleaning\n    #shutil.rmtree(folder_tiles)\n    os.system(f\"rm -rf {folder_tiles}\")\n\ndf_sub = pd.DataFrame(submission)","metadata":{"execution":{"iopub.status.busy":"2023-11-23T17:48:03.858188Z","iopub.execute_input":"2023-11-23T17:48:03.858939Z","iopub.status.idle":"2023-11-23T17:48:49.636783Z","shell.execute_reply.started":"2023-11-23T17:48:03.858906Z","shell.execute_reply":"2023-11-23T17:48:49.635642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Finalize - export submission","metadata":{}},{"cell_type":"code","source":"display(df_sub.head())\ndf_sub[[\"image_id\", \"label\"]].to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-23T17:49:03.275101Z","iopub.execute_input":"2023-11-23T17:49:03.275485Z","iopub.status.idle":"2023-11-23T17:49:03.28984Z","shell.execute_reply.started":"2023-11-23T17:49:03.275456Z","shell.execute_reply":"2023-11-23T17:49:03.288857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"!ls /","metadata":{"execution":{"iopub.status.busy":"2023-11-23T08:43:36.88658Z","iopub.status.idle":"2023-11-23T08:43:36.886925Z","shell.execute_reply.started":"2023-11-23T08:43:36.886757Z","shell.execute_reply":"2023-11-23T08:43:36.886774Z"}}},{"cell_type":"markdown","source":"timm.list_models()","metadata":{"execution":{"iopub.status.busy":"2023-11-23T16:28:30.000138Z","iopub.execute_input":"2023-11-23T16:28:30.00057Z","iopub.status.idle":"2023-11-23T16:28:30.029738Z","shell.execute_reply.started":"2023-11-23T16:28:30.000535Z","shell.execute_reply":"2023-11-23T16:28:30.0288Z"}}},{"cell_type":"markdown","source":"import timm\ntransformer = timm.create_model('swinv2_cr_tiny_ns_224', pretrained = True)","metadata":{"execution":{"iopub.status.busy":"2023-11-23T16:35:01.070419Z","iopub.execute_input":"2023-11-23T16:35:01.071426Z","iopub.status.idle":"2023-11-23T16:35:04.940562Z","shell.execute_reply.started":"2023-11-23T16:35:01.071385Z","shell.execute_reply":"2023-11-23T16:35:04.93979Z"}}},{"cell_type":"markdown","source":"!ls","metadata":{"execution":{"iopub.status.busy":"2023-11-23T16:35:14.515973Z","iopub.execute_input":"2023-11-23T16:35:14.516334Z","iopub.status.idle":"2023-11-23T16:35:15.585284Z","shell.execute_reply.started":"2023-11-23T16:35:14.516306Z","shell.execute_reply":"2023-11-23T16:35:15.584036Z"}}},{"cell_type":"markdown","source":"import torch\ntorch.save(transformer, \"swinv2_cr_tiny_ns_224.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-11-23T16:35:49.714531Z","iopub.execute_input":"2023-11-23T16:35:49.715341Z","iopub.status.idle":"2023-11-23T16:35:49.887104Z","shell.execute_reply.started":"2023-11-23T16:35:49.715294Z","shell.execute_reply":"2023-11-23T16:35:49.886237Z"}}},{"cell_type":"markdown","source":"import os\n!ls ","metadata":{"execution":{"iopub.status.busy":"2023-11-22T08:45:10.722995Z","iopub.execute_input":"2023-11-22T08:45:10.723793Z","iopub.status.idle":"2023-11-22T08:45:11.682736Z","shell.execute_reply.started":"2023-11-22T08:45:10.723764Z","shell.execute_reply":"2023-11-22T08:45:11.681746Z"}}},{"cell_type":"markdown","source":"from IPython.display import FileLink, FileLinks\ndisplay(FileLink('swinv2_cr_tiny_ns_224.pt'))","metadata":{"execution":{"iopub.status.busy":"2023-11-23T16:36:02.39441Z","iopub.execute_input":"2023-11-23T16:36:02.394792Z","iopub.status.idle":"2023-11-23T16:36:02.401538Z","shell.execute_reply.started":"2023-11-23T16:36:02.39476Z","shell.execute_reply":"2023-11-23T16:36:02.400669Z"}}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}