{"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":6898597,"sourceType":"datasetVersion","datasetId":3962744},{"sourceId":7176659,"sourceType":"datasetVersion","datasetId":3958369},{"sourceId":153762364,"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,"execution":{"iopub.status.busy":"2023-12-13T04:00:12.994743Z","iopub.execute_input":"2023-12-13T04:00:12.995017Z","iopub.status.idle":"2023-12-13T04:01:17.299844Z","shell.execute_reply.started":"2023-12-13T04:00:12.994985Z","shell.execute_reply":"2023-12-13T04:01:17.298671Z"},"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'\n\n!mkdir -p /kaggle/temp/test_tiles","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-13T04:01:17.301896Z","iopub.execute_input":"2023-12-13T04:01:17.302218Z","iopub.status.idle":"2023-12-13T04:01:18.604203Z","shell.execute_reply.started":"2023-12-13T04:01:17.302189Z","shell.execute_reply":"2023-12-13T04:01:18.602955Z"},"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())","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-13T04:01:18.605903Z","iopub.execute_input":"2023-12-13T04:01:18.606653Z","iopub.status.idle":"2023-12-13T04:01:18.624789Z","shell.execute_reply.started":"2023-12-13T04:01:18.606621Z","shell.execute_reply":"2023-12-13T04:01:18.624048Z"},"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))\nlabels = sorted(df_train[\"label\"].unique())\nprint(f\"{labels=}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-13T04:01:18.626865Z","iopub.execute_input":"2023-12-13T04:01:18.627181Z","iopub.status.idle":"2023-12-13T04:01:18.839948Z","shell.execute_reply.started":"2023-12-13T04:01:18.627156Z","shell.execute_reply":"2023-12-13T04:01:18.838953Z"},"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) form the training notebook:\n\n- extracting the tiles from whole image\n- validation argmention, mainly color mean & STD","metadata":{}},{"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, idxs","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-13T04:01:18.841518Z","iopub.execute_input":"2023-12-13T04:01:18.842155Z","iopub.status.idle":"2023-12-13T04:01:19.190063Z","shell.execute_reply.started":"2023-12-13T04:01:18.842117Z","shell.execute_reply":"2023-12-13T04:01:19.189269Z"},"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-12-13T04:01:19.191064Z","iopub.execute_input":"2023-12-13T04:01:19.191331Z","iopub.status.idle":"2023-12-13T04:01:19.197715Z","shell.execute_reply.started":"2023-12-13T04:01:19.191306Z","shell.execute_reply":"2023-12-13T04:01:19.196696Z"},"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-12-13T04:01:19.198834Z","iopub.execute_input":"2023-12-13T04:01:19.199107Z","iopub.status.idle":"2023-12-13T04:01:22.653827Z","shell.execute_reply.started":"2023-12-13T04:01:19.199083Z","shell.execute_reply":"2023-12-13T04:01:22.652851Z"},"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-12-13T04:01:22.655247Z","iopub.execute_input":"2023-12-13T04:01:22.655775Z","iopub.status.idle":"2023-12-13T04:01:22.666269Z","shell.execute_reply.started":"2023-12-13T04:01:22.655743Z","shell.execute_reply":"2023-12-13T04:01:22.66523Z"},"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)}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-13T04:01:22.667697Z","iopub.execute_input":"2023-12-13T04:01:22.668018Z","iopub.status.idle":"2023-12-13T04:01:55.195723Z","shell.execute_reply.started":"2023-12-13T04:01:22.667981Z","shell.execute_reply":"2023-12-13T04:01:55.194764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\n\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 = '/kaggle/input/ptlight-base/logs/maxvit_base_tf_512/version_0/checkpoints/epoch=59-step=480.ckpt'\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_base_tf_512', pretrained=False, num_classes=len(labels))\nmodel = LitCancerSubtype(net=net)\nmodel.load_state_dict(ckpt['state_dict'])","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-12-13T04:01:55.198339Z","iopub.execute_input":"2023-12-13T04:01:55.198628Z","iopub.status.idle":"2023-12-13T04:02:18.586349Z","shell.execute_reply.started":"2023-12-13T04:01:55.198604Z","shell.execute_reply":"2023-12-13T04:02:18.584955Z"},"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'] = ['Other'] * 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-12-13T04:06:21.930441Z","iopub.execute_input":"2023-12-13T04:06:21.930848Z","iopub.status.idle":"2023-12-13T04:06:21.94723Z","shell.execute_reply.started":"2023-12-13T04:06:21.930815Z","shell.execute_reply":"2023-12-13T04:06:21.946228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cat /kaggle/input/UBC-OCEAN/sample_submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-12-13T04:06:23.08568Z","iopub.execute_input":"2023-12-13T04:06:23.086041Z","iopub.status.idle":"2023-12-13T04:06:24.10097Z","shell.execute_reply.started":"2023-12-13T04:06:23.086012Z","shell.execute_reply":"2023-12-13T04:06:24.099975Z"},"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()\nsoftmax = torch.nn.Softmax(dim=1)\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.extend(pred.cpu().numpy())\n        \n    # Calculate mean probability for each class\n    mean_probs = np.mean(preds, axis=0)\n    max_mean_prob = np.max(mean_probs)\n    \n    #print(f\"Mean probabilities: {mean_probs}\")\n    #print(f\"Maximum mean probability: {max_mean_prob}\")\n    \n    # Decide if it's an outlier\n    if max_mean_prob < 0.5:\n        #print(\"Outlier detected\")\n        row['label'] = 'Other'  # Assign to outlier or a specific class\n    else:\n        lb = np.argmax(mean_probs)\n        row['label'] = labels[lb]\n\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-12-13T04:06:25.826519Z","iopub.execute_input":"2023-12-13T04:06:25.826948Z","iopub.status.idle":"2023-12-13T04:07:04.77805Z","shell.execute_reply.started":"2023-12-13T04:06:25.826898Z","shell.execute_reply":"2023-12-13T04:07:04.776814Z"},"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)\n\n! head submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-12-13T04:07:04.780538Z","iopub.execute_input":"2023-12-13T04:07:04.780966Z","iopub.status.idle":"2023-12-13T04:07:05.821538Z","shell.execute_reply.started":"2023-12-13T04:07:04.780915Z","shell.execute_reply":"2023-12-13T04:07:05.820525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}