{"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":153189206,"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":{"execution":{"iopub.status.busy":"2023-12-04T03:14:15.248393Z","iopub.execute_input":"2023-12-04T03:14:15.248682Z","iopub.status.idle":"2023-12-04T03:14:15.260949Z","shell.execute_reply.started":"2023-12-04T03:14:15.248655Z","shell.execute_reply":"2023-12-04T03:14:15.259646Z"}}},{"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-12-04T06:01:25.369968Z","iopub.execute_input":"2023-12-04T06:01:25.370376Z","iopub.status.idle":"2023-12-04T06:02:28.375322Z","shell.execute_reply.started":"2023-12-04T06:01:25.370341Z","shell.execute_reply":"2023-12-04T06:02:28.374189Z"},"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/\"\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-12-04T06:02:28.377593Z","iopub.execute_input":"2023-12-04T06:02:28.377893Z","iopub.status.idle":"2023-12-04T06:02:28.384175Z","shell.execute_reply.started":"2023-12-04T06:02:28.377865Z","shell.execute_reply":"2023-12-04T06:02:28.383115Z"},"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":"df_train = pd.read_csv(os.path.join(DATASET_FOLDER, \"train.csv\"))\nlabels = sorted(df_train[\"label\"].unique()) + [\"Other\"]\nprint(f\"{labels=}\")\ndel df_train","metadata":{"execution":{"iopub.status.busy":"2023-12-04T06:02:28.385389Z","iopub.execute_input":"2023-12-04T06:02:28.385658Z","iopub.status.idle":"2023-12-04T06:02:28.434967Z","shell.execute_reply.started":"2023-12-04T06:02:28.385633Z","shell.execute_reply":"2023-12-04T06:02:28.434074Z"},"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, size: int = 2048, scale: float = 0.5,\n    drop_thr: float = 0.6, white_thr: int = 245, max_samples: int = 50\n) -> list:\n    im = pyvips.Image.new_from_file(p_img)\n    height = im.height\n    width = im.width \n    print(f\"height1:{height},width1:{width}\")\n    size1 = max(height,width)\n    scale = float(20480/size1)\n    print(f\"size1:{size1},scale:{scale}\")\n    if scale<1.0:\n        im = im.resize(scale)\n    height = im.height\n    width = im.width \n    print(f\"height2:{height},width2:{width}\")\n    size1 = max(height,width)\n    size = size1//10\n    w = h = size\n    scale = float(2048/w)\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    images = []\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        # print(tile.shape, tile.dtype, tile.min(), tile.max())\n        new_size = int(size * scale), int(size * scale)\n        images.append(np.array(\n            Image.fromarray(tile).resize(new_size, Image.LANCZOS)\n        ))\n        # need to set counter check as some empty tiles could be skipped earlier\n        if len(images) >= max_samples:\n            break\n    return images","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-04T06:02:28.437971Z","iopub.execute_input":"2023-12-04T06:02:28.438258Z","iopub.status.idle":"2023-12-04T06:02:28.453065Z","shell.execute_reply.started":"2023-12-04T06:02:28.438234Z","shell.execute_reply":"2023-12-04T06:02:28.452198Z"},"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-04T06:02:28.454258Z","iopub.execute_input":"2023-12-04T06:02:28.45461Z","iopub.status.idle":"2023-12-04T06:02:28.46785Z","shell.execute_reply.started":"2023-12-04T06:02:28.454579Z","shell.execute_reply":"2023-12-04T06:02:28.466911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset\n\nclass TilesImageDataset(Dataset):\n\n    def __init__(\n        self,\n        img_path: str,\n        size: int = 2048,\n        scale: float = 0.25,\n        drop_thr: float = 0.6,\n        max_samples: int = 30,\n        transforms = None\n    ):\n        assert os.path.isfile(img_path)\n        self.transforms = transforms\n        max_samples = 1\n        self.imgs = extract_image_tiles(\n            img_path, size=size, scale=scale,\n            drop_thr=drop_thr, max_samples=max_samples)\n\n    def __getitem__(self, idx: int) -> tuple:\n        img = self.imgs[idx]\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-04T06:02:28.469097Z","iopub.execute_input":"2023-12-04T06:02:28.469489Z","iopub.status.idle":"2023-12-04T06:02:28.480438Z","shell.execute_reply.started":"2023-12-04T06:02:28.469455Z","shell.execute_reply":"2023-12-04T06:02:28.479593Z"},"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)}\")\ndataset = TilesImageDataset(ls[0])\nprint(f\"found tiles: {len(dataset)}\")\n\n# # quick view\n# fig, axes = plt.subplots(nrows=3, ncols=3, figsize=(10, 10))\n# for i in range(9):\n#     img = dataset[i]\n#     axes[i // 3, i % 3].imshow(img)\n# fig.tight_layout()","metadata":{"execution":{"iopub.status.busy":"2023-12-04T06:02:28.481407Z","iopub.execute_input":"2023-12-04T06:02:28.481672Z","iopub.status.idle":"2023-12-04T06:02:43.056604Z","shell.execute_reply.started":"2023-12-04T06:02:28.481649Z","shell.execute_reply":"2023-12-04T06:02:43.055638Z"},"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\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.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# ==============================\n# ==============================\n\nPATH_CKPT = (\n    \"/kaggle/input/cancer-subtype-tiles-masks-w-lightning-timm/\"\n    \"image_classification_model.pt\"\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('tf_efficientnetv2_s_in21ft1k', pretrained=False, num_classes=len(labels))\nmodel = LitCancerSubtype(net=net)\nmodel.load_state_dict(ckpt['state_dict'])\n# print(model)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-12-04T06:02:43.05807Z","iopub.execute_input":"2023-12-04T06:02:43.058724Z","iopub.status.idle":"2023-12-04T06:02:43.706606Z","shell.execute_reply.started":"2023-12-04T06:02:43.058687Z","shell.execute_reply":"2023-12-04T06:02:43.705662Z"},"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-04T06:02:43.707839Z","iopub.execute_input":"2023-12-04T06:02:43.708138Z","iopub.status.idle":"2023-12-04T06:02:43.72208Z","shell.execute_reply.started":"2023-12-04T06:02:43.708112Z","shell.execute_reply":"2023-12-04T06:02:43.721237Z"},"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-04T06:02:43.725096Z","iopub.execute_input":"2023-12-04T06:02:43.725395Z","iopub.status.idle":"2023-12-04T06:02:44.724536Z","shell.execute_reply.started":"2023-12-04T06:02:43.72537Z","shell.execute_reply":"2023-12-04T06:02:44.723342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"import scipy\nfrom torch.utils.data import DataLoader\nfrom joblib.externals.loky.backend.context import get_context\n\ndef infer_single_image(\n    idx_row, model, device=\"cuda\", other_thr: float = 0.9, max_samples: int = 40\n) -> dict:\n    row = dict(idx_row[1])\n    # prepare data - cut and load tiles\n    img_path = os.path.join(DATASET_FOLDER, \"test_images\", f\"{str(row['image_id'])}.png\")\n    print(f\"processing: {img_path}\")\n    dataset = TilesImageDataset(\n        img_path, size=2048, scale=0.25, max_samples=max_samples, transforms=VALID_TRANSFORM\n    )\n    if not len(dataset):\n        print (f\"seem no tiles were cut for `{row['image_id']}`\")\n        return row\n    preds = []\n    model = model.to(device)\n    dataloader = DataLoader(\n        dataset, batch_size=4, num_workers=2, shuffle=False,\n        # see: https://github.com/pytorch/pytorch/issues/44687#issuecomment-790842173\n        multiprocessing_context=get_context('loky')\n    )\n    # iterate over images and collect predictions | \n    for imgs in dataloader:\n        #print(f\"{imgs.shape}\")\n        with torch.no_grad():\n            pred = model(imgs.to(device))\n        preds += pred.cpu().numpy().tolist()\n    probs = scipy.special.softmax(preds, axis=1)\n    #print(f\"Softmax on sum of all tiles: {probs}\")\n    print(f\"Sum contrinution from all tiles: {np.sum(probs, axis=0)}\")\n    print(f\"Max contribution over all tiles: {np.max(probs, axis=0)}\")\n    # decide label\n    probs_agg = np.sum(probs, axis=0) / np.sum(probs)\n    if probs_agg[-1] > other_thr:\n        # if other is more then 90% use it\n        row['label'] = \"Other\"\n    else:\n        lb = np.argmax(probs_agg[:-1])\n        row['label'] = labels[lb]\n    print(row)\n    return row","metadata":{"execution":{"iopub.status.busy":"2023-12-04T06:02:44.726552Z","iopub.execute_input":"2023-12-04T06:02:44.72748Z","iopub.status.idle":"2023-12-04T06:02:44.738962Z","shell.execute_reply.started":"2023-12-04T06:02:44.727435Z","shell.execute_reply":"2023-12-04T06:02:44.738075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.auto import tqdm\nfrom joblib import Parallel, delayed\n\nmodel.eval()\n\nif len(df_test) > 1:\n    submission = Parallel(n_jobs=2)(\n        delayed(infer_single_image)\n        (idx_row, model=model, device=\"cuda\", max_samples=50, other_thr=0.98)\n        for idx_row in tqdm(df_test.iterrows(), total=len(df_test))\n    )\nelse:\n    submission = [\n        infer_single_image(idx_row, model)\n        for idx_row in tqdm(df_test.iterrows(), total=len(df_test))\n    ]\n\ndf_sub = pd.DataFrame(submission)","metadata":{"execution":{"iopub.status.busy":"2023-12-04T06:02:44.740254Z","iopub.execute_input":"2023-12-04T06:02:44.74056Z","iopub.status.idle":"2023-12-04T06:03:03.497918Z","shell.execute_reply.started":"2023-12-04T06:02:44.740535Z","shell.execute_reply":"2023-12-04T06:03:03.496855Z"},"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-12-04T06:03:03.499655Z","iopub.execute_input":"2023-12-04T06:03:03.500038Z","iopub.status.idle":"2023-12-04T06:03:03.515Z","shell.execute_reply.started":"2023-12-04T06:03:03.500001Z","shell.execute_reply":"2023-12-04T06:03:03.514099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! head submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-12-04T06:03:03.516435Z","iopub.execute_input":"2023-12-04T06:03:03.516786Z","iopub.status.idle":"2023-12-04T06:03:04.547884Z","shell.execute_reply.started":"2023-12-04T06:03:03.516753Z","shell.execute_reply":"2023-12-04T06:03:04.546659Z"},"trusted":true},"execution_count":null,"outputs":[]}]}