{"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":7299425,"sourceType":"datasetVersion","datasetId":4234446}],"dockerImageVersionId":30627,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"%%capture --no-stderr\n!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":{"execution":{"iopub.status.busy":"2023-12-28T23:24:04.790776Z","iopub.execute_input":"2023-12-28T23:24:04.791471Z","iopub.status.idle":"2023-12-28T23:25:06.112346Z","shell.execute_reply.started":"2023-12-28T23:24:04.79144Z","shell.execute_reply":"2023-12-28T23:25:06.111016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport copy\nimport time\nimport random\nimport shutil\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom tqdm import tqdm\nimport concurrent.futures\nfrom concurrent.futures import FIRST_COMPLETED\nfrom concurrent.futures import ThreadPoolExecutor\nfrom sklearn.metrics import balanced_accuracy_score, accuracy_score\nfrom sklearn.model_selection import train_test_split, StratifiedKFold\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom PIL import Image\nfrom collections import Counter\nimport timm\nimport pyvips\ntqdm.pandas()\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport torch\nimport torch.nn as nn\nfrom torchvision import models\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\nos.environ['VIPS_CONCURRENCY'] = '1'\nos.environ['VIPS_DISC_THRESHOLD'] = '15gb'\npyvips.cache_set_max(0)\n\nROOT = \"../input/\"\nDATA_PATH = \"../input/UBC-OCEAN/\"\nIMAGES_FOLDER = \"./test_tiles\"\n\n!mkdir -p /kaggle/temp/test_tiles","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:06.114696Z","iopub.execute_input":"2023-12-28T23:25:06.115087Z","iopub.status.idle":"2023-12-28T23:25:13.741902Z","shell.execute_reply.started":"2023-12-28T23:25:06.11505Z","shell.execute_reply":"2023-12-28T23:25:13.740673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONFIG = {\n    \"seed\": 42,\n    \"model_name\": \"maxvit_tiny_tf_512\",\n    \"checkpoint_path\": ROOT + \"basic-model-maxvit-512/models/fold_0.bin\",\n    \"num_classes\": 5,\n    \"device\": torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\"),\n    \"apex\": True\n}","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:13.743703Z","iopub.execute_input":"2023-12-28T23:25:13.744716Z","iopub.status.idle":"2023-12-28T23:25:13.775858Z","shell.execute_reply.started":"2023-12-28T23:25:13.744677Z","shell.execute_reply":"2023-12-28T23:25:13.77494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def extract_image_tiles(p_img, folder, size = 1024, scale = 0.25,drop_thr = 0.6, white_thr = 240):\n    max_samples = 30\n    name, _ = os.path.splitext(os.path.basename(p_img))\n    im = pyvips.Image.new_from_file(p_img)\n    if im.width > 5000 and im.height > 5000:\n        im = im.resize(0.5)\n    w = h = size\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.shuffle(idxs)\n    n_samples = 0\n    for y, y_, x, x_ in idxs:\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            continue\n        p_img = os.path.join(folder, f\"{int(x_ / w)}-{int(y_ / h)}.png\")\n        new_size = int(size * scale), int(size * scale)\n        Image.fromarray(tile).resize(new_size, Image.LANCZOS).save(p_img)\n        n_samples += 1\n        if n_samples >= max_samples:\n            break\n    del black_bg, tile, mask_bg, im\n    gc.collect()\n    return ","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:13.778084Z","iopub.execute_input":"2023-12-28T23:25:13.778412Z","iopub.status.idle":"2023-12-28T23:25:13.799707Z","shell.execute_reply.started":"2023-12-28T23:25:13.778386Z","shell.execute_reply":"2023-12-28T23:25:13.798876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def extract_prune_tiles(path_img, folder, size = 1024, scale = 0.25, drop_thr = 0.6):\n    idx = int(path_img.split('/')[-1].split('.')[0])\n    name = os.path.splitext(os.path.basename(path_img))[0]\n    folder = os.path.join(folder, name)\n    os.makedirs(folder, exist_ok=True)\n    extract_image_tiles(path_img, folder, size=size, scale=scale, drop_thr=drop_thr)\n    return idx, folder","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:13.800795Z","iopub.execute_input":"2023-12-28T23:25:13.801097Z","iopub.status.idle":"2023-12-28T23:25:13.814294Z","shell.execute_reply.started":"2023-12-28T23:25:13.801073Z","shell.execute_reply":"2023-12-28T23:25:13.813463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_resolution_wsi(row):\n    wsi_id = row.image_id\n    wsi_path = os.path.join(DATA_PATH, \"test_images\", f\"{wsi_id}.png\")\n    im = pyvips.Image.new_from_file(wsi_path)\n    return (im.width, im.height)","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:13.81533Z","iopub.execute_input":"2023-12-28T23:25:13.815615Z","iopub.status.idle":"2023-12-28T23:25:13.824371Z","shell.execute_reply.started":"2023-12-28T23:25:13.815592Z","shell.execute_reply":"2023-12-28T23:25:13.823472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Patch Loader    \nclass WSIPatchLoader(Dataset):\n    \"\"\"\n    Dataloader for iterating through all patches in a WSI\n    \"\"\"    \n    def __init__(self, folder, transform=None):\n        self.imgs = glob(os.path.join(folder, \"*.png\"))\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.imgs)\n\n    def __getitem__(self, idx):\n        im_path = self.imgs[idx]\n        im = cv2.cvtColor(cv2.imread(im_path), cv2.COLOR_BGR2RGB)\n        if self.transform:\n            im = self.transform(image=im)['image']\n        return im","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:13.825564Z","iopub.execute_input":"2023-12-28T23:25:13.825838Z","iopub.status.idle":"2023-12-28T23:25:13.837393Z","shell.execute_reply.started":"2023-12-28T23:25:13.825805Z","shell.execute_reply":"2023-12-28T23:25:13.83653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = nn.Parameter(torch.ones(1)*p)\n        self.eps = eps\n\n    def forward(self, x):\n        return self.gem(x, p=self.p, eps=self.eps)\n        \n    def gem(self, x, p=3, eps=1e-6):\n        return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1./p)\n        \n    def __repr__(self):\n        return self.__class__.__name__ + \\\n                '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + \\\n                ', ' + 'eps=' + str(self.eps) + ')'","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:13.838586Z","iopub.execute_input":"2023-12-28T23:25:13.838846Z","iopub.status.idle":"2023-12-28T23:25:13.851213Z","shell.execute_reply.started":"2023-12-28T23:25:13.838824Z","shell.execute_reply":"2023-12-28T23:25:13.850359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Model    \nclass UBCCancerSubtype(nn.Module):\n    def __init__(self, n_class=5, pretrained=False, checkpoint_path=None):\n        super(UBCCancerSubtype, self).__init__()\n        self.backbone = timm.create_model('maxvit_tiny_tf_512', pretrained=pretrained)\n        in_features = self.backbone.head.fc.in_features\n        self.backbone.head.fc = nn.Linear(in_features, n_class)\n        self.softmax = nn.Softmax(dim=1)\n        self.num_classes = n_class\n        if checkpoint_path is not None:\n            self.initialize_model(checkpoint_path)\n        \n    def initialize_model(self, checkpoint_path):\n        ckpt = torch.load(checkpoint_path, map_location=torch.device('cpu'))\n        ckpt = {k.removeprefix(\"backbone.\"): v for k, v in ckpt.items()}\n        self.backbone.load_state_dict(ckpt)\n        \n    def forward(self, x):\n        Y_prob = self.backbone(x)\n        if Y_prob.isnan().any():\n            Y_prob = torch.ones(self.num_classes) / self.num_classes\n        return Y_prob","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:13.852312Z","iopub.execute_input":"2023-12-28T23:25:13.852619Z","iopub.status.idle":"2023-12-28T23:25:13.86139Z","shell.execute_reply.started":"2023-12-28T23:25:13.85259Z","shell.execute_reply":"2023-12-28T23:25:13.86049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef get_prediction(folder, model, device, data_transforms):\n    dataset = WSIPatchLoader(folder, transform=data_transforms)\n    if not len(dataset):\n        return 0\n    dataloader = DataLoader(dataset, batch_size=4, num_workers=2, shuffle=False)\n    # iterate over patches and collect predictions\n    preds = []\n    for imgs in dataloader:\n        imgs = imgs.to(device)\n        pred = model(imgs)\n        preds.extend(pred.cpu().detach().numpy().tolist())\n    # decide label\n    lb = np.argmax(np.sum(preds, axis=0))\n    if os.path.isdir(folder):\n        shutil.rmtree(folder)\n    return lb.item()","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:13.864807Z","iopub.execute_input":"2023-12-28T23:25:13.865669Z","iopub.status.idle":"2023-12-28T23:25:13.875778Z","shell.execute_reply.started":"2023-12-28T23:25:13.865644Z","shell.execute_reply":"2023-12-28T23:25:13.874929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(DATA_PATH + 'test.csv')\n\ntest = df[['image_id']] # for submission\n# get width and height of each wsi\ndf[['image_width', 'image_height']] = df.progress_apply(get_resolution_wsi, axis=1, result_type='expand')\ndf[\"is_tma\"] = (df[\"image_width\"] <= 5000) & (df[\"image_height\"] <= 5000)\n# get slides coordinates\ndf[\"wsi_path\"] = df[\"image_id\"].apply(lambda idx:os.path.join(DATA_PATH, \"test_images\", f'{idx}.png'))\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:13.876741Z","iopub.execute_input":"2023-12-28T23:25:13.877006Z","iopub.status.idle":"2023-12-28T23:25:13.95153Z","shell.execute_reply.started":"2023-12-28T23:25:13.876974Z","shell.execute_reply":"2023-12-28T23:25:13.95065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[\"size\"] = df[\"image_width\"] * df[\"image_height\"]\ndf_sorted = df.copy(deep=True)\ndf_sorted.sort_values(by=\"size\", ascending=False, inplace=True)\ndf_large = df_sorted.loc[df_sorted[\"size\"]>=(2*10**9)]\ndf_fast = df_sorted.loc[df_sorted[\"size\"]<(2*10**9)]\ndf_large.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:13.952675Z","iopub.execute_input":"2023-12-28T23:25:13.952947Z","iopub.status.idle":"2023-12-28T23:25:13.964847Z","shell.execute_reply.started":"2023-12-28T23:25:13.952923Z","shell.execute_reply":"2023-12-28T23:25:13.963996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UBCCancerSubtype(n_class=CONFIG[\"num_classes\"], checkpoint_path=CONFIG[\"checkpoint_path\"])\nmodel = model.to(CONFIG[\"device\"])\nmodel.eval()\n\n#Data Transformation\nimg_color_mean = [0.8540517447797603, 0.7541362654418763, 0.8512977343452348]\nimg_color_std = [0.08107490432889232, 0.11650886216683837, 0.06458838921081513]\ndata_transforms = {\n    \"valid\": A.Compose([\n#         A.Resize(CONFIG['img_size'], CONFIG['img_size']),\n        A.Normalize(mean=img_color_mean, std=img_color_std, max_pixel_value=255.0, p=1.0),\n        ToTensorV2()], p=1.)\n}","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:13.965856Z","iopub.execute_input":"2023-12-28T23:25:13.966199Z","iopub.status.idle":"2023-12-28T23:25:18.714146Z","shell.execute_reply.started":"2023-12-28T23:25:13.966171Z","shell.execute_reply":"2023-12-28T23:25:18.713305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_preds = {\"image_id\": [], \"label\": []}\n\nwith torch.no_grad():\n    with concurrent.futures.ProcessPoolExecutor(max_workers=2) as executor:\n        remaining_futures = {executor.submit(extract_prune_tiles, row.wsi_path, IMAGES_FOLDER, 2048, 0.25) for _, row in df_large.iterrows()}\n        while remaining_futures:\n            done, remaining_futures = concurrent.futures.wait(remaining_futures, return_when=FIRST_COMPLETED)\n            for fut in done:\n                idx, folder = fut.result()\n                all_preds[\"image_id\"].append(idx)\n                all_preds[\"label\"].append(get_prediction(folder, model, CONFIG['device'], data_transforms[\"valid\"]))","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:18.71527Z","iopub.execute_input":"2023-12-28T23:25:18.715573Z","iopub.status.idle":"2023-12-28T23:25:18.724174Z","shell.execute_reply.started":"2023-12-28T23:25:18.715548Z","shell.execute_reply":"2023-12-28T23:25:18.723011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with torch.no_grad():\n    with concurrent.futures.ProcessPoolExecutor(max_workers=3) as executor:\n        remaining_futures = {executor.submit(extract_prune_tiles, row.wsi_path, IMAGES_FOLDER, 2048, 0.25) for _, row in df_fast.iterrows()}\n        while remaining_futures:\n            done, remaining_futures = concurrent.futures.wait(remaining_futures, return_when=FIRST_COMPLETED)\n            for fut in done:\n                idx, folder = fut.result()\n                all_preds[\"image_id\"].append(idx)\n                all_preds[\"label\"].append(get_prediction(folder, model, CONFIG['device'], data_transforms[\"valid\"]))","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:18.725327Z","iopub.execute_input":"2023-12-28T23:25:18.725586Z","iopub.status.idle":"2023-12-28T23:25:18.737429Z","shell.execute_reply.started":"2023-12-28T23:25:18.725564Z","shell.execute_reply":"2023-12-28T23:25:18.736531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame(all_preds)\nlabel_map = {\"HGSC\":0, \"LGSC\":1, \"EC\":2, \"CC\":3, \"MC\":4}#, \"Other\":5}\nlabel_map_inv = {v:k for k,v in label_map.items()}\nsubmission[\"label\"] = submission[\"label\"].map(label_map_inv)\n\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:18.738559Z","iopub.execute_input":"2023-12-28T23:25:18.739262Z","iopub.status.idle":"2023-12-28T23:25:18.756651Z","shell.execute_reply.started":"2023-12-28T23:25:18.73923Z","shell.execute_reply":"2023-12-28T23:25:18.75575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = test.merge(submission, on=\"image_id\", how='left')\nsubmission[\"label\"] = submission[\"label\"].fillna(\"HGSC\")\nwhile \"submission.csv\" not in os.listdir(\"/kaggle/working\"):\n    submission.to_csv('/kaggle/working/submission.csv', index = False)\n\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-28T23:25:18.757685Z","iopub.execute_input":"2023-12-28T23:25:18.75795Z","iopub.status.idle":"2023-12-28T23:25:18.77915Z","shell.execute_reply.started":"2023-12-28T23:25:18.757928Z","shell.execute_reply":"2023-12-28T23:25:18.778234Z"},"trusted":true},"execution_count":null,"outputs":[]}]}