{"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":152945614,"sourceType":"kernelVersion"}],"dockerImageVersionId":30559,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Inference\n\nCode adapted from https://www.kaggle.com/code/jirkaborovec/cancer-subtype-lit-torch-infer-tiles-parallel","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\n\n\n# # intall the deb packages\n# !yes | dpkg -i --force-depends /kaggle/input/pyvips-python-and-deb-package/linux_packages/archives/*.deb\n# # install the python wrapper\n# !pip install pyvips -f /kaggle/input/pyvips-python-and-deb-package/python_packages/ --no-index","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-01T19:52:23.21931Z","iopub.execute_input":"2023-12-01T19:52:23.21962Z","iopub.status.idle":"2023-12-01T19:53:26.763023Z","shell.execute_reply.started":"2023-12-01T19:52:23.219594Z","shell.execute_reply":"2023-12-01T19:53:26.762022Z"},"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/\"\nDATASET_IMAGES = \"/kaggle/input/UBC-OCEAN/test_images\"\nMODEL_PATH = \"/kaggle/input/ubc-ocean-training-labeled-tiles-224-maxvit/image_classification_model.pt\"\nMODEL_NAME = \"maxvit_rmlp_base_rw_224.sw_in12k_ft_in1k\"\nTILE_SIZE = 224\nSCALE = 0.175\nWHITE_THR = 240\nDROP_THR = 0.5 #0.6 #0.8\nTMA_THR = 5000 # TMA images usually have length around ~3000 pixels\nMAX_SAMPLES = 50\nDEVICE = \"cuda\"\nBATCH_SIZE = 16\n\nos.environ['VIPS_CONCURRENCY'] = '4'\nos.environ['VIPS_DISC_THRESHOLD'] = '20gb'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-01T20:47:09.445891Z","iopub.execute_input":"2023-12-01T20:47:09.446297Z","iopub.status.idle":"2023-12-01T20:47:09.453493Z","shell.execute_reply.started":"2023-12-01T20:47:09.446267Z","shell.execute_reply":"2023-12-01T20:47:09.452477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data for inference","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=}\")\ndel df_train","metadata":{"execution":{"iopub.status.busy":"2023-12-01T20:47:17.863857Z","iopub.execute_input":"2023-12-01T20:47:17.864731Z","iopub.status.idle":"2023-12-01T20:47:17.874966Z","shell.execute_reply.started":"2023-12-01T20:47:17.864694Z","shell.execute_reply":"2023-12-01T20:47:17.873953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pyvips\nimport numpy as np\nimport random\nfrom typing import List, Tuple, Union\n\n\ndef is_tma(p_img: pyvips.vimage.Image, tma_thr: int) -> bool:\n    # Determine whether image is TMA by checking its longer dimension\n    img = pyvips.Image.new_from_file(p_img)\n    return max(img.width, img.height) < tma_thr\n    \n\ndef load_and_resize_img(p_img: str, scale: float) -> pyvips.vimage.Image:\n    if not os.path.isfile(p_img):\n        return None\n    img = pyvips.Image.new_from_file(p_img)\n#     img = img.resize(scale, kernel='lanczos2')\n    img = pyvips.Image.thumbnail(p_img, int(scale * img.width)) # Faster than resize\n    return img.copy_memory() # Needed when opening image using thumbnail\n\n\ndef find_subintervals(\n    binary_list: List[Union[int, float]], \n    target_digit: int,\n) -> List[Tuple[int, int]]:\n    subintervals = []\n    start = None\n\n    for i, digit in enumerate(binary_list):\n        if digit == target_digit:\n            if start is None:\n                start = i\n        elif start is not None:\n            subintervals.append((start, i - 1))\n            start = None\n\n    if start is not None:\n        subintervals.append((start, len(binary_list) - 1))\n\n    return subintervals\n\n\ndef get_subimg_x_intervals(img: pyvips.vimage.Image) -> List[Tuple[int, int]]:\n    is_empty_column = (img.rot90().fliphor().bandfold().bandmean() == 0).numpy().squeeze() / 255\n    return find_subintervals(is_empty_column, target_digit=0)\n\n\ndef get_tile_img(\n    img: pyvips.vimage.Image,\n    x: int,\n    y: int,\n    tile_size: int,\n) -> pyvips.vimage.Image:\n    return img.crop(x, y, min(tile_size, img.width - x), min(tile_size, img.height - y)\n       ).gravity('north-west', tile_size, tile_size, extend='black')\n\n    \ndef is_background_tile(\n    tile: pyvips.vimage.Image,\n    white_thr: int,\n    drop_thr: float,\n) -> bool:\n    mean_tile = tile.bandmean()\n    mask_bg = (mean_tile == 0).bandjoin(mean_tile > white_thr).bandor()    \n    return (mask_bg.avg() / 255) > drop_thr\n\n\ndef get_tiles_wsi(\n    img: pyvips.vimage.Image,\n    x_intervals: List[Tuple[int, int]], \n    tile_size: int,\n    white_thr: int,\n    drop_thr: float,\n    max_samples: int = None,  # Set to None to use all tiles\n) -> List[np.ndarray]:    \n    idxs = [\n        (x, y)\n        for x_interval in x_intervals\n        for y in range(0, img.height, tile_size)\n        for x in range(x_interval[0], x_interval[1] + 1, tile_size)\n    ]\n    random.seed(42)\n    random.shuffle(idxs)\n    \n    tile_imgs = []\n    for x, y in idxs:\n        tile_img = get_tile_img(img, x, y, tile_size)\n        if not is_background_tile(tile_img, white_thr, drop_thr):\n            tile_imgs.append(\n                tile_img.numpy()[..., :3]\n            )\n            if len(tile_imgs) == max_samples:\n                break\n    return tile_imgs            \n               \n\ndef get_tiles_tma(\n    img: pyvips.vimage.Image,\n    tile_size: int,\n) -> List[np.ndarray]:\n    # Get center crop of the image\n    w, h = img.width, img.height\n    x = 0 if w < tile_size else (w - tile_size) // 2\n    y = 0 if h < tile_size else (h - tile_size) // 2\n    tile_img = img.crop(x, y, min(tile_size, w - x), min(tile_size, h - y)\n        ).gravity('centre', tile_size, tile_size, extend='copy')\n    return [tile_img.numpy()[..., :3]]\n            \n            \ndef make_tiles_numpy(\n    p_img: str, \n    tile_size: int,\n    scale: float, \n    white_thr: float, \n    drop_thr: float,\n    tma_thr: int,\n    max_samples: int = None,  # Max number of tiles for WSI image\n) -> List[np.ndarray]:\n    tma = is_tma(p_img, tma_thr)\n    if tma:\n        scale *= 0.5   # TMA has 40x magnification while WSI has 20x magnification\n        img = load_and_resize_img(p_img, scale)\n        return get_tiles_tma(img, tile_size)\n    else:\n        img = load_and_resize_img(p_img, scale)\n        x_intervals = get_subimg_x_intervals(img)\n        return get_tiles_wsi(img, x_intervals, tile_size, white_thr, drop_thr, max_samples)","metadata":{"execution":{"iopub.status.busy":"2023-12-01T20:47:19.193346Z","iopub.execute_input":"2023-12-01T20:47:19.193737Z","iopub.status.idle":"2023-12-01T20:47:19.216588Z","shell.execute_reply.started":"2023-12-01T20:47:19.193708Z","shell.execute_reply":"2023-12-01T20:47:19.215655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import transforms as T\n\nimg_color_mean=[0.8029574609001011, 0.6753532015392826, 0.8150805152175007]\nimg_color_std=[0.09240988133071351, 0.11661346690553148, 0.06439091956270869]\n\nVALID_TRANSFORM = T.Compose([\n    T.CenterCrop(TILE_SIZE),\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-01T20:47:22.64466Z","iopub.execute_input":"2023-12-01T20:47:22.645399Z","iopub.status.idle":"2023-12-01T20:47:22.650986Z","shell.execute_reply.started":"2023-12-01T20:47:22.645368Z","shell.execute_reply":"2023-12-01T20:47:22.649996Z"},"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 = 224,\n        scale: float = 0.175,\n        white_thr: int = 240,\n        drop_thr: float = 0.5,\n        tma_thr: int = 5000,\n        max_samples: int = 50,\n        transforms = None\n    ):\n        assert os.path.isfile(img_path)\n        self.transforms = transforms\n        self.imgs = make_tiles_numpy(\n            img_path,\n            tile_size=size,\n            scale=scale, \n            white_thr=white_thr, \n            drop_thr=drop_thr,\n            tma_thr=tma_thr,\n            max_samples=max_samples,\n        )\n\n    def __getitem__(self, idx: int) -> Union[torch.Tensor, np.ndarray]:\n        img = self.imgs[idx]\n        # filter background\n        black_bg = np.sum(img, axis=2) == 0\n        img[black_bg, :] = 255\n        # augmentation\n        if self.transforms:\n            img = self.transforms(Image.fromarray(img))\n        return img\n\n    def __len__(self) -> int:\n        return len(self.imgs)","metadata":{"execution":{"iopub.status.busy":"2023-12-01T20:47:29.715272Z","iopub.execute_input":"2023-12-01T20:47:29.716008Z","iopub.status.idle":"2023-12-01T20:47:29.726599Z","shell.execute_reply.started":"2023-12-01T20:47:29.715973Z","shell.execute_reply":"2023-12-01T20:47:29.725307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\nls = sorted(glob.glob(os.path.join(DATASET_FOLDER, \"test_images\", '*.png')))\nprint(f\"found images: {len(ls)}\")\ndataset = TilesImageDataset(\n    ls[0],\n    size=TILE_SIZE,\n    scale=SCALE,\n    white_thr=WHITE_THR,\n    drop_thr=DROP_THR,\n    tma_thr=TMA_THR,\n    max_samples=MAX_SAMPLES,\n    transforms=None,\n)\nprint(f\"found tiles: {len(dataset)}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-01T20:47:32.907423Z","iopub.execute_input":"2023-12-01T20:47:32.908287Z","iopub.status.idle":"2023-12-01T20:47:53.892274Z","shell.execute_reply.started":"2023-12-01T20:47:32.908254Z","shell.execute_reply":"2023-12-01T20:47:53.891359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 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-12-01T20:48:14.42367Z","iopub.execute_input":"2023-12-01T20:48:14.424024Z","iopub.status.idle":"2023-12-01T20:48:16.595536Z","shell.execute_reply.started":"2023-12-01T20:48:14.423997Z","shell.execute_reply":"2023-12-01T20:48:16.594382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CNN Model","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):\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\nckpt = torch.load(MODEL_PATH, map_location=torch.device('cpu'))\n\n# see: https://pytorch.org/vision/stable/models.html\nnet = timm.create_model(MODEL_NAME, pretrained=False, num_classes=len(labels))\nmodel = LitCancerSubtype(net=net)\nmodel.load_state_dict(ckpt['state_dict'])\nprint(model)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-12-01T20:48:19.325335Z","iopub.execute_input":"2023-12-01T20:48:19.325727Z","iopub.status.idle":"2023-12-01T20:48:22.696401Z","shell.execute_reply.started":"2023-12-01T20:48:19.32569Z","shell.execute_reply":"2023-12-01T20:48:22.695462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test data & submission\n","metadata":{}},{"cell_type":"code","source":"df_test = pd.read_csv(os.path.join(DATASET_FOLDER, \"test.csv\"))\n# default label\ndf_test['label'] = [''] * len(df_test)\nprint(f\"Dataset/test size: {len(df_test)}\")\ndisplay(df_test.head())","metadata":{"execution":{"iopub.status.busy":"2023-12-01T20:48:38.263106Z","iopub.execute_input":"2023-12-01T20:48:38.263783Z","iopub.status.idle":"2023-12-01T20:48:38.277689Z","shell.execute_reply.started":"2023-12-01T20:48:38.263739Z","shell.execute_reply":"2023-12-01T20:48:38.276575Z"},"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-01T20:48:49.715735Z","iopub.execute_input":"2023-12-01T20:48:49.716636Z","iopub.status.idle":"2023-12-01T20:48:50.80006Z","shell.execute_reply.started":"2023-12-01T20:48:49.716603Z","shell.execute_reply":"2023-12-01T20:48:50.798782Z"},"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\ndash = \"-\"*100\n\ndef infer_single_image(\n    idx_row, model, device=\"cuda\",\n) -> dict:\n    row = dict(idx_row[1])\n    # prepare data - cut and load tiles\n    img_path = os.path.join(DATASET_IMAGES, f\"{str(row['image_id'])}.png\")\n    dataset = TilesImageDataset(\n        img_path,\n        size=TILE_SIZE,\n        scale=SCALE,\n        white_thr=WHITE_THR,\n        drop_thr=DROP_THR,\n        tma_thr=TMA_THR,\n        max_samples=MAX_SAMPLES,\n        transforms=VALID_TRANSFORM,\n    )\n    print(f\"processing #{idx_row[0]+1} {img_path} with {len(dataset)} tiles\")\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=BATCH_SIZE, 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.extend(\n            pred.cpu().numpy().tolist()\n        )\n    probs = scipy.special.softmax(preds, axis=1)\n#     print(f\"{probs=}\")\n    print(f\"Mean prob over all tiles: {np.mean(probs, axis=0)}\")\n    print(f\"Max prob over all tiles: {np.max(probs, axis=0)}\")\n    # decide label\n    probs_agg = np.mean(probs, axis=0)\n    lb = np.argmax(probs_agg)\n    row['label'] = labels[lb]\n    print(row)\n    print(dash)\n    return row","metadata":{"execution":{"iopub.status.busy":"2023-12-01T20:57:14.642975Z","iopub.execute_input":"2023-12-01T20:57:14.643781Z","iopub.status.idle":"2023-12-01T20:57:14.655298Z","shell.execute_reply.started":"2023-12-01T20:57:14.643747Z","shell.execute_reply":"2023-12-01T20:57:14.654145Z"},"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=DEVICE)\n        for idx_row in tqdm(df_test.iterrows(), total=len(df_test))\n    )\nelse:\n    submission = [\n        infer_single_image(idx_row, model, device=DEVICE)\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-01T20:57:15.372064Z","iopub.execute_input":"2023-12-01T20:57:15.372872Z","iopub.status.idle":"2023-12-01T20:57:42.502492Z","shell.execute_reply.started":"2023-12-01T20:57:15.372838Z","shell.execute_reply":"2023-12-01T20:57:42.501352Z"},"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-01T20:57:54.760324Z","iopub.execute_input":"2023-12-01T20:57:54.760738Z","iopub.status.idle":"2023-12-01T20:57:54.77582Z","shell.execute_reply.started":"2023-12-01T20:57:54.760706Z","shell.execute_reply":"2023-12-01T20:57:54.774912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! head submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-12-01T20:57:55.045123Z","iopub.execute_input":"2023-12-01T20:57:55.045532Z","iopub.status.idle":"2023-12-01T20:57:56.130692Z","shell.execute_reply.started":"2023-12-01T20:57:55.045497Z","shell.execute_reply":"2023-12-01T20:57:56.129445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}