{"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":7171567,"sourceType":"datasetVersion","datasetId":4143602},{"sourceId":7219928,"sourceType":"datasetVersion","datasetId":4178748}],"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-12-30T06:23:41.382212Z","iopub.execute_input":"2023-12-30T06:23:41.382737Z","iopub.status.idle":"2023-12-30T06:24:44.810374Z","shell.execute_reply.started":"2023-12-30T06:23:41.382709Z","shell.execute_reply":"2023-12-30T06:24:44.809172Z"},"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\"\nBAG_SIZE = 6\nMAX_SAMPLE_PER_IMAGE = 75 #BAG_SIZE*5\nINPUT_SIZE = 512\nTILE_SIZE = 1024\nBATCH_SIZE = 2\nNUM_WORKERS = 2\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-30T06:24:44.812243Z","iopub.execute_input":"2023-12-30T06:24:44.812565Z","iopub.status.idle":"2023-12-30T06:24:45.164871Z","shell.execute_reply.started":"2023-12-30T06:24:44.812537Z","shell.execute_reply":"2023-12-30T06:24:45.164112Z"},"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 = ['CC', 'EC', 'HGSC', 'LGSC','MC','Stroma','Necrosis']\nprint(f\"{LABELS=}\")","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-30T06:24:45.166051Z","iopub.execute_input":"2023-12-30T06:24:45.166697Z","iopub.status.idle":"2023-12-30T06:24:45.185004Z","shell.execute_reply.started":"2023-12-30T06:24:45.16665Z","shell.execute_reply":"2023-12-30T06:24:45.18404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:24:45.187696Z","iopub.execute_input":"2023-12-30T06:24:45.188157Z","iopub.status.idle":"2023-12-30T06:24:45.20597Z","shell.execute_reply.started":"2023-12-30T06:24:45.18813Z","shell.execute_reply":"2023-12-30T06:24:45.204936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.groupby('image_id').get_group(4)['is_tma'].item()","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:24:45.207116Z","iopub.execute_input":"2023-12-30T06:24:45.207448Z","iopub.status.idle":"2023-12-30T06:24:45.224545Z","shell.execute_reply.started":"2023-12-30T06:24:45.207417Z","shell.execute_reply":"2023-12-30T06:24:45.223764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TH_TMA_FILE_SIZE   = 1.5\n!du -s -m /kaggle/input/UBC-OCEAN/test_images/* > amount.csv\ndf_size = pd.read_csv(\"amount.csv\", delimiter='\\t', header=None).rename(columns={0:\"size_mb\", 1:\"file\"})\ndf_size[\"image_id\"] = df_size[\"file\"].apply(lambda x: int(os.path.basename(x).split(\".\")[0] ) )\ndf_size[\"size_mb_log10\"] = np.log10(df_size[\"size_mb\"].values)\ndf_size[\"is_tma\"] = df_size[\"size_mb_log10\"] < TH_TMA_FILE_SIZE\n# df_size.set_index('image_id',inplace = True)\ndf_size","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:24:45.225565Z","iopub.execute_input":"2023-12-30T06:24:45.225846Z","iopub.status.idle":"2023-12-30T06:24:46.195801Z","shell.execute_reply.started":"2023-12-30T06:24:45.225821Z","shell.execute_reply":"2023-12-30T06:24:46.194774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pyvips\nimport numpy as np\nimport random\nfrom PIL import Image\nimport gc\n\ndef drop_bg(tile, area_thresh = 0.6, white_thresh = 220):\n    drop_black = False\n    drop_white = False\n    mask_bg = np.sum(tile, axis=2) == 0\n    if np.sum(mask_bg) >= (np.prod(mask_bg.shape) * area_thresh): # too much black\n        drop_black = True\n    tile[mask_bg, :] = 255\n    mask_bg = np.mean(tile, axis=2) > white_thresh\n    if np.sum(mask_bg) >= (np.prod(mask_bg.shape) * area_thresh):# too much white\n        drop_white = True\n#     if drop_black:\n#         print('too many blacks')\n#     if drop_white:\n#         print('too many whites')\n    return drop_black or drop_white\n\ndef evaluate_tumor_quality(model, dataloader):\n    model.eval()\n    model = model.to('cuda')\n    preds = []\n    for batch in dataloader:\n        pred = model(batch.cuda())\n        preds.append(pred.detach().cpu())\n        \n    \n    preds = torch.cat(preds)\n#     print(preds.shape)\n    preds = torch.softmax(preds, axis = 1)\n    labels = torch.argmax(preds, axis = 1).numpy()\n    proba = torch.max(preds, axis=1).values.numpy()\n#     print(labels, proba)\n    has_tumor = [l in range(5) for l, p in zip(labels, proba)]\n#     print(list(zip(labels, proba, has_tumor)))\n    model = model.to('cpu')\n    return has_tumor\n\ndef extract_image_tiles(\n    img_wsi,\n    tile_size = 1024,\n    resize_shape = None,\n    max_sample=None,\n    check_drop_bg  =True,\n    patch_classifier_model = None\n):\n    im = img_wsi# pyvips.Image.new_from_file(p_img)\n    w = h = tile_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    print('total patches', len(idxs))\n    \n    np.random.seed(42)\n    if max_sample:\n        np.random.shuffle(idxs)\n        idxs = idxs[:max_sample]\n    \n    dataset_bg_det = BG_Det_Dataset(im, idxs, tile_size = tile_size)\n    dataloader_bg_det = DataLoader(\n            dataset_bg_det, batch_size=4, num_workers=0, shuffle=False,\n        )\n    del im\n    for _ in range(5):\n        gc.collect()\n    bool_indices = []\n    tiles = []\n    for tiles_batch, batch_bool in tqdm(dataloader_bg_det):\n        bool_indices+=list(batch_bool)\n        tiles.append(tiles_batch)\n    tiles = np.concatenate(tiles)\n    tile_list_no_bg = np.array(tiles)[~np.array(bool_indices)]\n    print('patches after filtering background', len(tile_list_no_bg))\n    \n    if len(tile_list_no_bg) == 0:\n        return tiles[:BAG_SIZE*4]\n\n    dataset = PatchesDataset(\n        [Image.fromarray(tile) for tile in tile_list_no_bg], \n        patch_classifier_model['preprocess']\n    )\n    dataloader = DataLoader(\n            dataset, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS, shuffle=False,\n        )\n    has_tumor = evaluate_tumor_quality(patch_classifier_model['model'], dataloader)\n    \n    tile_list_no_bg_pc = np.array(tile_list_no_bg)[has_tumor]\n    print('patches after filtering background and patch classification', len(tile_list_no_bg_pc))\n\n    if len(tile_list_no_bg_pc) < BAG_SIZE:\n        print('patches got removed due to bg removal and patch classifier, adding random bags')\n#         tile_list_no_bg_pc = list(tile_list_no_bg_pc)\n#         tile_list_no_bg_pc += list(tile_list_no_bg[:BAG_SIZE*4])\n#         print(np.array(tile_list_no_bg_pc).shape)\n        \n    \n    del dataset, dataloader\n    for _ in range(5):\n        gc.collect()\n#     for img in tile_list_no_bg_pc[:10]:\n#         plt.imshow(img)\n#         plt.show()\n    return tile_list_no_bg_pc","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-30T06:29:05.803817Z","iopub.execute_input":"2023-12-30T06:29:05.804228Z","iopub.status.idle":"2023-12-30T06:29:05.8243Z","shell.execute_reply.started":"2023-12-30T06:29:05.804185Z","shell.execute_reply":"2023-12-30T06:29:05.823262Z"},"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.Resize((INPUT_SIZE,INPUT_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-30T06:29:06.461168Z","iopub.execute_input":"2023-12-30T06:29:06.462127Z","iopub.status.idle":"2023-12-30T06:29:06.468699Z","shell.execute_reply.started":"2023-12-30T06:29:06.46209Z","shell.execute_reply":"2023-12-30T06:29:06.467508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset\n\nclass TilesDataset(Dataset):\n\n    def __init__(\n        self,\n        tiles,\n        bag_size,\n        transforms\n    ):\n        self.transforms = transforms\n        self.imgs = tiles\n        self.available_bags = len(self.imgs)//bag_size if len(self.imgs)//bag_size>0 else 1\n        self.bag_size = bag_size\n        self.imgs_bags = [self.imgs[i*bag_size:i*bag_size+bag_size] for i in range(self.available_bags) ]\n\n    def __getitem__(self, idx: int) -> tuple:\n        bag = self.imgs_bags[idx]\n        # augmentation\n        if self.transforms:\n            bag =[self.transforms(Image.fromarray(item))  for item in bag]\n        #print(f\"img dim: {img.shape}\")\n#         print(bag)\n        return torch.stack(bag)#, torch.tensor(labels).to(int)\n\n\n    def __len__(self) -> int:\n        return self.available_bags\n    \nclass PatchesDataset(Dataset):\n    def __init__(\n        self,\n        tiles,\n        transforms\n    ):\n        self.transforms = transforms\n        self.tiles = tiles\n\n    def __getitem__(self, idx: int) -> tuple:\n        tile = self.tiles[idx]\n        tile = self.transforms(tile)\n        return tile\n\n    def __len__(self) -> int:\n        return len(self.tiles)\n    \nclass BG_Det_Dataset(Dataset):\n    def __init__(\n        self,\n        im,\n        idxs,\n        tile_size\n    ):\n        self.im = im\n        self.idxs = idxs\n        self.tile_size = tile_size\n\n    def __getitem__(self, idx: int) -> tuple:\n        y, y_, x, x_ = self.idxs[idx]\n        # make tile\n        w = h = self.tile_size\n        tile = self.im.crop(x, y, min(w, self.im.width - x), min(h, self.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        drop_bool = drop_bg(tile.copy())\n#         if drop_bool == True:\n#             tile = []\n        return tile, drop_bool\n    def __len__(self) -> int:\n        return len(self.idxs)","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:29:07.024285Z","iopub.execute_input":"2023-12-30T06:29:07.024659Z","iopub.status.idle":"2023-12-30T06:29:07.040079Z","shell.execute_reply.started":"2023-12-30T06:29:07.02463Z","shell.execute_reply":"2023-12-30T06:29:07.03913Z"},"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\nenc= timm.create_model(\n    'tiny_vit_21m_512.dist_in22k_ft_in1k', pretrained=False, num_classes=7)\nclass Model(nn.Module):\n    def __init__(self,enc):\n        super().__init__()\n        self.enc = enc\n    def forward(self, x):\n        return self.enc(x)\nmodel_pc = Model(enc)\nmodel_data = torch.load('/kaggle/input/2023-12-09-17-07-51/tiny_vit_21m_512.dist_in22k_ft_in1k/version_0/checkpoints/epoch=4-step=9925.ckpt')\nmodel_pc.load_state_dict(model_data['state_dict'])\n\npatch_classifier_model = {\n    'model': model_pc,\n    'preprocess': VALID_TRANSFORM\n}","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:24:49.767986Z","iopub.execute_input":"2023-12-30T06:24:49.768446Z","iopub.status.idle":"2023-12-30T06:25:15.529637Z","shell.execute_reply.started":"2023-12-30T06:24:49.768421Z","shell.execute_reply":"2023-12-30T06:25:15.528768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:25:15.530764Z","iopub.execute_input":"2023-12-30T06:25:15.531378Z","iopub.status.idle":"2023-12-30T06:25:15.806385Z","shell.execute_reply.started":"2023-12-30T06:25:15.531348Z","shell.execute_reply":"2023-12-30T06:25:15.805455Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nclass AdaptiveConcatPool2d(torch.nn.Module):\n    \"Layer that concats `AdaptiveAvgPool2d` and `AdaptiveMaxPool2d`\"\n    def __init__(self, size=None):\n        super().__init__()\n        self.size = size or 1\n        self.ap = torch.nn.AdaptiveAvgPool2d(self.size)\n        self.mp = torch.nn.AdaptiveMaxPool2d(self.size)\n    def forward(self, x): return torch.cat([self.mp(x), self.ap(x)], 1)\n\n\nclass Model_BOW(nn.Module):\n    def __init__(self, enc, feature_dim,\n                 num_class = 5):\n        super().__init__()\n    #         self.conv_bag = torch.nn.Conv2d(in_channels=12*3,out_channels=3,kernel_size=3,padding=1)\n    #         self.net = net\n    #         self.arch = net.pretrained_cfg.get('architecture')\n        self.num_classes = num_class\n        self.enc = enc\n        self.feature_dim = feature_dim\n        self.head = nn.Sequential(\n            AdaptiveConcatPool2d(),\n            torch.nn.Flatten(),\n            nn.Linear(2*feature_dim,512),\n            torch.nn.ReLU(),\n            #torch.nn.Mish(),\n            torch.nn.LayerNorm(512),\n            nn.Linear(512,256),\n            torch.nn.ReLU(),\n            torch.nn.Mish(),\n            torch.nn.LayerNorm(256),\n            #torch.nn.Dropout(0.5),\n            torch.nn.Linear(256,self.num_classes)\n        )\n\n\n        \n    def forward(self, x, x_large):\n        batch,bag,c,h,w = x.shape\n\n        x = x.view(batch*bag, c,h,w)\n        #x: bs*N x C x 4 x 4\n        x = self.enc(x)\n        _,c_out,h_out,w_out = x.shape\n        #concatenate the output for tiles into a single map\n        x = x.view(batch,bag,c_out,h_out,w_out)\n        #x: bsxN x C x 4 x 4\n        #print(x.shape)\n        x_l = self.enc(x_large)\n        #print(x_l.shape)\n        x_l = x_l.unsqueeze(1)\n        x=torch.cat([x,x_l], axis=1)\n\n        #x_l = x_l.repeat(1,bag,1,1,1)\n        #x = x+x_l # add?\n\n        x = x.permute(0,2,1,3,4).contiguous() # bs,C,N,H,W\n        #print(x.shape)\n        x = x.view(batch,c_out,h_out*(bag+1),w_out)\n        #print(x.shape)\n        x = self.head(x)\n\n        #x: bs x n\n        return x\n\nnet= timm.create_model(\n    'maxvit_tiny_tf_512', pretrained=False, num_classes=5)\nnet = nn.Sequential(*list(net.children())[:-1])\nmodel_bow=Model_BOW(enc=net, feature_dim=512)\nmodel_data = torch.load('/kaggle/input/2023-12-14-20-32-38/fold_0/maxvit_tiny_tf_512/version_0/checkpoints/epoch=99-step=3400.ckpt')\nmodel_bow.load_state_dict(model_data['state_dict'])","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:25:15.807591Z","iopub.execute_input":"2023-12-30T06:25:15.807902Z","iopub.status.idle":"2023-12-30T06:25:21.104673Z","shell.execute_reply.started":"2023-12-30T06:25:15.807875Z","shell.execute_reply":"2023-12-30T06:25:21.103779Z"},"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-12-30T06:25:21.105747Z","iopub.execute_input":"2023-12-30T06:25:21.106015Z","iopub.status.idle":"2023-12-30T06:25:21.121939Z","shell.execute_reply.started":"2023-12-30T06:25:21.105993Z","shell.execute_reply":"2023-12-30T06:25:21.121004Z"},"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-30T06:25:21.123188Z","iopub.execute_input":"2023-12-30T06:25:21.123494Z","iopub.status.idle":"2023-12-30T06:25:22.136537Z","shell.execute_reply.started":"2023-12-30T06:25:21.123469Z","shell.execute_reply":"2023-12-30T06:25:22.135556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"def infer_single_image(model, patch_classifier_model, row, train):\n    row = dict(row)\n    # prepare data - cut and load tiles\n    if train:\n        split = 'train'\n    else:\n        split = 'test'\n    path = os.path.join(DATASET_FOLDER, \"{}_images\".format(split), f\"{str(row['image_id'])}.png\")\n    img_wsi = pyvips.Image.new_from_file(path)\n\n    tiles = extract_image_tiles(\n        img_wsi,\n        tile_size=TILE_SIZE, \n        max_sample=MAX_SAMPLE_PER_IMAGE,\n        patch_classifier_model = patch_classifier_model\n    )\n#     print(tiles[0].shape)\n\n    del img_wsi\n    for _ in range(5):\n        gc.collect()\n    \n    if len(tiles)==0:\n        row['label'] = 'Other'\n    else:\n        \n        img_global = Image.open(os.path.join(\n            '/kaggle/input/UBC-OCEAN/{}_thumbnails'.format(split),\n            '{}_thumbnail.png'.format(row['image_id'])\n        ))\n    #     img_global = np.array(img_global)\n        img_global = data_transforms_valid(image = np.array(img_global))['image']\n        img_global = img_global.unsqueeze(0)\n\n\n        dataset = TilesDataset(tiles, bag_size=BAG_SIZE,transforms=VALID_TRANSFORM)\n        dataloader = DataLoader(\n                dataset, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS, shuffle=False,\n        #                            multiprocessing_context=get_context('loky')\n            )\n        print('bags', len(dataset))\n        # iterate over images and collect predictions\n        preds = []\n        model = model.to('cuda')\n        for imgs in dataloader:\n            #print(f\"{imgs.shape}\")\n            img_global_batch = torch.cat([img_global.clone()]*len(imgs))\n            img_global_batch = img_global_batch.cuda()\n            with torch.no_grad():\n                pred = model(imgs.cuda(), img_global_batch)\n            preds.append(pred.detach().cpu())\n        model = model.to('cpu')\n        # decide label\n        preds = torch.cat(preds)\n        print(preds.shape)\n        preds = torch.mean(preds, axis=0)\n        preds = torch.softmax(preds.view(-1), axis=0)\n        print(preds)\n        lb = np.argmax(preds)\n        row['label'] = LABELS[lb]\n\n        del dataset, dataloader\n        for _ in range(5):\n            gc.collect()\n\n    return row\n","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:29:18.154209Z","iopub.execute_input":"2023-12-30T06:29:18.154607Z","iopub.status.idle":"2023-12-30T06:29:18.167834Z","shell.execute_reply.started":"2023-12-30T06:29:18.154577Z","shell.execute_reply":"2023-12-30T06:29:18.166597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def infer_single_image_tma(model, row, train):\n    \n    model = model.to('cuda')\n    model.eval()\n    row = dict(row)\n    if train:\n        split = 'train'\n    else: \n        split = 'test'\n    path = os.path.join('/kaggle/input/UBC-OCEAN/{}_images/{}.png'.format(split, row['image_id']))\n    img = Image.open(path)\n    img = np.array(img)\n    pred = np.zeros((1,7))\n    for tx in rotate_tx:\n        img_tx = tx(image = img)['image']   \n        img_tx = data_transforms_valid(image = img_tx)['image']\n        img_tx = img_tx.unsqueeze(0)\n        #print(f\"{imgs.shape}\")\n        with torch.no_grad():\n#             print(model(img_tx.cuda()).detach().cpu().numpy())\n            pred += model(img_tx.cuda()).detach().cpu().numpy()\n    # decide label\n    model = model.to('cpu')\n    lb = np.argmax(pred)\n    row['label'] = LABELS[lb]\n    return row","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:29:20.052767Z","iopub.execute_input":"2023-12-30T06:29:20.05316Z","iopub.status.idle":"2023-12-30T06:29:20.061701Z","shell.execute_reply.started":"2023-12-30T06:29:20.053127Z","shell.execute_reply":"2023-12-30T06:29:20.060649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nitem = (2048, 512)\ndata_transforms_valid = A.Compose([\n            A.PadIfNeeded(item[0], item[0]),\n            A.CenterCrop(item[0], item[0]),\n            A.Resize(item[1], item[1]),\n            A.Normalize(\n                mean = img_color_mean, #[0.485, 0.456, 0.406], \n                std = img_color_std, #[0.229, 0.224, 0.225], \n                max_pixel_value=255.0, \n                p=1.0\n            ),\n            ToTensorV2()], p=1.)\n\ntx90 = A.Compose(\n    [\n        A.Rotate(limit=(90,90),always_apply=True, p=1.0),\n#         A.Resize(4,4)\n    ], p=1.0)\ntx180 = A.Compose(\n    [\n        A.Rotate(limit=(90,90),always_apply=True, p=1.0),\n#         A.Resize(4,4)\n    ]*2, p=1.0)\ntx270 = A.Compose(\n    [\n        A.Rotate(limit=(90,90),always_apply=True, p=1.0),\n#         A.Resize(4,4)\n    ]*3, p=1.0)\nrotate_tx = [tx90, tx180, tx270]","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:29:20.571954Z","iopub.execute_input":"2023-12-30T06:29:20.572346Z","iopub.status.idle":"2023-12-30T06:29:20.581655Z","shell.execute_reply.started":"2023-12-30T06:29:20.572316Z","shell.execute_reply":"2023-12-30T06:29:20.580616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\nfrom torch.utils.data import DataLoader\nfrom joblib.externals.loky.backend.context import get_context\nfrom tqdm.auto import tqdm \nimport time\n\nmodel_bow.eval()\nTRAIN = False\n\nsubmission = []\nif TRAIN:\n    df_select = df_train.copy()\nelse:\n    df_select = df_test.copy()\n    \n# df_select = df_select[50:100]    \n\nfor i, row in tqdm(df_select.iterrows(), total = len(df_select)):\n    print(row.to_list())\n    if TRAIN:\n        is_tma = df_select.groupby('image_id').get_group(row['image_id'])['is_tma'].item()\n    else:\n        is_tma = df_size.groupby('image_id').get_group(row['image_id'])['is_tma'].item()\n    if is_tma:\n        row = infer_single_image(model_bow, patch_classifier_model, row, TRAIN)\n\n        row = infer_single_image_tma(model_pc, row, TRAIN)\n        if row['label'] in ['Stroma', 'Necrosis']:\n            row['label'] = 'Other'\n        submission.append(dict(row))\n    else:\n        start = time.time()\n        row = infer_single_image(model_bow, patch_classifier_model, row, TRAIN)\n#         print(row)\n        submission.append(dict(row))\n        print(time.time()-start)\n    print(row)\ndf_sub = pd.DataFrame(submission)\n","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:30:21.751156Z","iopub.execute_input":"2023-12-30T06:30:21.752095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from sklearn import metrics\n\n# print(metrics.balanced_accuracy_score(df_train[~df_train['is_tma']].label.values, df_sub[~df_train['is_tma']].label.values))\n\n# print(metrics.confusion_matrix(df_train[~df_train['is_tma']].label.values, df_sub[~df_train['is_tma']].label.values))","metadata":{"execution":{"iopub.status.busy":"2023-12-30T06:26:02.699611Z","iopub.execute_input":"2023-12-30T06:26:02.699898Z","iopub.status.idle":"2023-12-30T06:26:02.704446Z","shell.execute_reply.started":"2023-12-30T06:26:02.699873Z","shell.execute_reply":"2023-12-30T06:26:02.703541Z"},"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-30T06:26:02.70546Z","iopub.execute_input":"2023-12-30T06:26:02.705728Z","iopub.status.idle":"2023-12-30T06:26:03.759688Z","shell.execute_reply.started":"2023-12-30T06:26:02.705705Z","shell.execute_reply":"2023-12-30T06:26:03.758644Z"},"trusted":true},"execution_count":null,"outputs":[]}]}