{"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":7305416,"sourceType":"datasetVersion","datasetId":3904916},{"sourceId":6884642,"sourceType":"datasetVersion","datasetId":3955404},{"sourceId":7081912,"sourceType":"datasetVersion","datasetId":3962274},{"sourceId":7291678,"sourceType":"datasetVersion","datasetId":3886390},{"sourceId":7292275,"sourceType":"datasetVersion","datasetId":4102617},{"sourceId":7292870,"sourceType":"datasetVersion","datasetId":4104024},{"sourceId":7294015,"sourceType":"datasetVersion","datasetId":4104039},{"sourceId":7294018,"sourceType":"datasetVersion","datasetId":4104126},{"sourceId":7298971,"sourceType":"datasetVersion","datasetId":4030506},{"sourceId":117061919,"sourceType":"kernelVersion"}],"dockerImageVersionId":30588,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport math\nimport copy\nimport time\nimport random\nimport glob\nfrom matplotlib import pyplot as plt\n\n# For data manipulation\nimport numpy as np\nimport pandas as pd\n\n# Pytorch Imports\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda import amp\nimport torchvision\n\n# Utils\nimport joblib\nfrom tqdm import tqdm\nfrom collections import defaultdict\n\n# Sklearn Imports\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import StratifiedKFold\n\n# For Image Models\nimport timm\nimport multiprocessing as mp\nfrom joblib import Parallel, delayed\n\nfrom multiprocessing import Process, Queue, Manager, Lock\n\n# Albumentations for augmentations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# For colored terminal text\nfrom colorama import Fore, Back, Style\nb_ = Fore.BLUE\nsr_ = Style.RESET_ALL\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# For descriptive error messages\nos.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"","metadata":{"execution":{"iopub.status.busy":"2023-12-26T11:40:54.697494Z","iopub.execute_input":"2023-12-26T11:40:54.697906Z","iopub.status.idle":"2023-12-26T11:41:00.97953Z","shell.execute_reply.started":"2023-12-26T11:40:54.69787Z","shell.execute_reply":"2023-12-26T11:41:00.978726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.system(\"dpkg -i --force-depends /kaggle/input/pyvips-python-and-deb-package-gpu/linux_packages/archives/*.deb\") \nos.system(\"pip install pyvips -f /kaggle/input/pyvips-python-and-deb-package-gpu/python_packages/ --no-index\") ","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/einops-download/einops-0.6.0-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2023-12-26T11:41:58.465038Z","iopub.execute_input":"2023-12-26T11:41:58.465326Z","iopub.status.idle":"2023-12-26T11:42:29.958376Z","shell.execute_reply.started":"2023-12-26T11:41:58.465301Z","shell.execute_reply":"2023-12-26T11:42:29.957278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(42)","metadata":{"execution":{"iopub.status.busy":"2023-12-26T11:42:29.959843Z","iopub.execute_input":"2023-12-26T11:42:29.960224Z","iopub.status.idle":"2023-12-26T11:42:29.972308Z","shell.execute_reply.started":"2023-12-26T11:42:29.960188Z","shell.execute_reply":"2023-12-26T11:42:29.971566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ['VIPS_DISC_THRESHOLD'] = '30gb'\n\nimport pyvips\n\ndef vips2numpy(vi):\n    format_to_dtype = {\n       'uchar': np.uint8,\n       'char': np.int8,\n       'ushort': np.uint16,\n       'short': np.int16,\n       'uint': np.uint32,\n       'int': np.int32,\n       'float': np.float32,\n       'double': np.float64,\n       'complex': np.complex64,\n       'dpcomplex': np.complex128,\n    }\n\n    return np.ndarray(\n        buffer=vi.write_to_memory(),\n        dtype=format_to_dtype[vi.format],\n        shape=[vi.height, vi.width, vi.bands])\n\n\ndef tile(img, sz, N):\n    rgb_mask = np.sum(img, axis=2) == 0\n    img[rgb_mask] = [255, 255, 255]\n    del rgb_mask\n    shape = img.shape\n    pad0,pad1 = (sz - shape[0]%sz)%sz, (sz - shape[1]%sz)%sz\n    img = np.pad(img,[[pad0//2,pad0-pad0//2],[pad1//2,pad1-pad1//2],[0,0]],constant_values=255)\n    img = img.reshape(img.shape[0]//sz,sz,img.shape[1]//sz,sz,3)\n    img = img.transpose(0,2,1,3,4).reshape(-1,sz,sz,3)\n     \n    if len(img) < 16:\n        img = np.pad(img,[[0,16-len(img)],[0,0],[0,0],[0,0]],constant_values=255)\n    \n    idxs = np.argsort(img.reshape(img.shape[0],-1).sum(-1))[:N] # pick up Top N dark tiles\n    #idxs = np.argsort(img.std(axis=(1,2)).max(1))[-N:][::-1]\n    \n    img = img[idxs]\n    return img","metadata":{"execution":{"iopub.status.busy":"2023-12-26T11:42:29.974803Z","iopub.execute_input":"2023-12-26T11:42:29.975176Z","iopub.status.idle":"2023-12-26T11:42:30.285067Z","shell.execute_reply.started":"2023-12-26T11:42:29.975149Z","shell.execute_reply":"2023-12-26T11:42:30.284315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_DIR = '/kaggle/input/UBC-OCEAN'\n# TEST_DIR = '/kaggle/input/UBC-OCEAN/test_thumbnails'\nALT_TEST_DIR = '/kaggle/input/UBC-OCEAN/test_images'\nTEST_THUMBNAILS_ROOT_DIR = '/kaggle/input/UBC-OCEAN/test_thumbnails'","metadata":{"execution":{"iopub.status.busy":"2023-12-26T11:42:30.286194Z","iopub.execute_input":"2023-12-26T11:42:30.286551Z","iopub.status.idle":"2023-12-26T11:42:30.291486Z","shell.execute_reply.started":"2023-12-26T11:42:30.286516Z","shell.execute_reply":"2023-12-26T11:42:30.290561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir /kaggle/working/temp","metadata":{"execution":{"iopub.status.busy":"2023-12-26T11:42:30.292739Z","iopub.execute_input":"2023-12-26T11:42:30.293082Z","iopub.status.idle":"2023-12-26T11:42:31.238686Z","shell.execute_reply.started":"2023-12-26T11:42:30.293051Z","shell.execute_reply":"2023-12-26T11:42:31.237597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def crop_dup(ori_img):\n    gray = cv2.cvtColor(ori_img,cv2.COLOR_BGR2GRAY)\n    _,thresh = cv2.threshold(gray,1,255,cv2.THRESH_BINARY)\n\n    contours,hierarchy = cv2.findContours(thresh,cv2.RETR_EXTERNAL,cv2.CHAIN_APPROX_SIMPLE)\n    cnt_sort = sorted(contours, key = cv2.contourArea)[-6:]\n    max_area = cv2.contourArea(cnt_sort[-1])\n    x,y,w,h = cv2.boundingRect(cnt_sort[-1])\n    max_crop = ori_img[y:y+h,x:x+w]\n    max_std = 0\n    \n    for cnt in cnt_sort:\n        area = cv2.contourArea(cnt)\n        if area >= 0.9*max_area:\n            x,y,w,h = cv2.boundingRect(cnt)\n            crop = ori_img[y:y+h,x:x+w]\n            std = crop.std()\n            if std > max_std:\n                max_std = std\n                max_crop = crop\n        \n    return max_crop\n    \nIMAGE_SIZE = [518]\ncrop_size = 518\nN=32\n\ndef crop_image_into_tiles(img, num_tiles_side):\n    # Lấy kích thước của hình ảnh gốc\n    img_height, img_width, _ = img.shape\n    \n    # Tính toán chiều rộng và chiều cao của mỗi ô\n    tile_width = int(3 / 4 * img_width)\n    tile_height = int(3 / 4 * img_height)\n    \n    # Tính toán stride\n    stride_x = int(img_width / 12)\n    stride_y = int(img_height / 12)\n    \n    tiles = []\n    \n    for i in range(num_tiles_side):\n        for j in range(num_tiles_side):\n            # Tính toán tọa độ của ô hiện tại\n            left = int(i * stride_x)\n            upper = int(j * stride_y)\n            right = left + tile_width\n            lower = upper + tile_height\n            \n            # Cắt hình ảnh để lấy ô hiện tại\n            tile = img[max(0, upper):min(img_height, lower), max(0, left):min(img_width, right), :]\n            \n            # Thêm ô vào danh sách\n            tiles.append(tile)\n    \n    return tiles\n\ndef get_indices(n, total):\n    if 2*n <= total:\n        indices = range(0, 2*n, 2)\n        return indices\n    if total <= n:\n        indices = range(0, n, 1)\n        return indices\n   \n    c = 0; indices = []\n    for i in range(0, total, 2):\n        indices.append(i)\n        c += 1\n    for i in range(1, total, 2):\n        indices.append(i)\n        c += 1\n        if c >= n:\n            break\n    indices = sorted(indices)\n    return indices\n\ndef center_crop(image, height, width):\n    image_height, image_width, c = image.shape\n    if height <= image_height and width <= image_width:\n#         print(\"1\")\n        pad_w = (image_width-width)//2\n        pad_h = (image_height-height)//2\n        image_crop = image[pad_h:height+pad_h,pad_w:width+pad_w,:].copy()\n    elif height <= image_height and width > image_width:\n#         print(\"2\")\n        pad_h = (image_height-height)//2\n        image_crop = image[pad_h:height+pad_h,:,:].copy()\n        pad_w = width - image_width\n        image_crop = np.pad(image_crop,[[0,0],[pad_w//2,pad_w-pad_w//2],[0,0]],constant_values=255)\n    elif height > image_height and width <= image_width:\n#         print(\"3\")\n        pad_w = (image_width-width)//2\n        image_crop = image[:,pad_w:width+pad_w,:].copy()\n        pad_h = height - image_height\n        image_crop = np.pad(image_crop,[[pad_h//2,pad_h-pad_h//2],[0,0],[0,0]],constant_values=255)\n        \n    return image_crop\ndef read_worker(q_in_predictor, df):\n    for index, row in df.iterrows():\n        image_id = str(row[\"image_id\"])\n        image_path = os.path.join(ALT_TEST_DIR, image_id + \".png\")\n        thumbnail_path = os.path.join(TEST_THUMBNAILS_ROOT_DIR, image_id + \"_thumbnail.png\")\n        image = pyvips.Image.new_from_file(image_path)\n        image_height, image_width = image.height, image.width\n#         image_width = row['image_width']\n#         image_height = row['image_height']\n\n        if image_width < 4001 or image_height < 4001:\n            is_tma = True\n        else:\n            is_tma = False\n        batch_518 = []  \n#         is_tma = True\n        if is_tma == True:\n            id_dir_518 = os.path.join(\"/kaggle/working/518\", str(image_id))\n            os.makedirs(id_dir_518, exist_ok=True)\n            \n            image = cv2.imread(image_path)\n            image_height, image_width, _ = image.shape\n            image = cv2.resize(image, (image_width//8, image_height//8), interpolation=cv2.INTER_LANCZOS4)\n            new_image_height, new_image_width, _ = image.shape\n            if(new_image_height > 518*1.7 or new_image_width > 518*1.7):\n                images = tile(image, sz=crop_size, N=N)\n                n_tiles = len(images)\n\n                id_dir_518 = os.path.join(\"/kaggle/working/518\", str(image_id))\n                os.makedirs(id_dir_518, exist_ok=True)\n                for idx in range(0, 16):\n                    img_out_path_518 = os.path.join(id_dir_518, str(idx) + \".jpg\")\n                    img = images[idx]\n                    img_518 = cv2.cvtColor(img, cv2.COLOR_RGB2BGR) \n                    cv2.imwrite(img_out_path_518, img_518, [cv2.IMWRITE_JPEG_QUALITY, 100])\n                    batch_518.append(img_out_path_518)\n            else:\n                if(new_image_height > 518 or new_image_width > 518):\n                    image = center_crop(image, 518, 518)\n                else:\n                    shape = image.shape\n                    pad0,pad1 = 518 - shape[0], 518 - shape[1]\n                    image = np.pad(image,[[pad0//2,pad0-pad0//2],[pad1//2,pad1-pad1//2],[0,0]],constant_values=255)\n\n                img_out_path_518 = os.path.join(id_dir_518, str(0) + \".jpg\")\n                cv2.imwrite(img_out_path_518, image, [cv2.IMWRITE_JPEG_QUALITY, 100])\n                batch_518.append(img_out_path_518)\n                images = crop_image_into_tiles(image, 4)[:16]\n                for idx in range(0, 15):\n                    img_out_path_518 = os.path.join(id_dir_518, str(idx+1) + \".jpg\")\n                    img = images[idx]\n                    shape = img.shape\n                    pad0,pad1 = 518 - shape[0], 518 - shape[1]\n                    img = np.pad(img,[[pad0//2,pad0-pad0//2],[pad1//2,pad1-pad1//2],[0,0]],constant_values=255)\n                    cv2.imwrite(img_out_path_518, img, [cv2.IMWRITE_JPEG_QUALITY, 100])\n                    batch_518.append(img_out_path_518)\n        else:\n            t1 = time.time()\n            image = pyvips.Image.thumbnail(image_path, image_width//4, height=image_height//4, size=\"force\")\n            image = vips2numpy(image)\n#             image = cv2.resize(image, (101, 102), interpolation=cv2.INTER_LANCZOS4)\n            image = crop_dup(image)\n            t2 = time.time()\n            print(\"Process time: \", t2 - t1)\n            \n            height, width, c = image.shape\n            print(f\"{index}. Input width: {width} height: {height}\")\n            images = tile(image, sz=crop_size, N=N)\n            n_tiles = len(images)\n            indices = range(0, n_tiles)\n\n            id_dir_518 = os.path.join(\"/kaggle/working/518\", str(image_id))\n            os.makedirs(id_dir_518, exist_ok=True)\n            for idx in indices:\n                img_out_path_518 = os.path.join(id_dir_518, str(idx) + \".jpg\")\n                img = images[idx]\n#                     img_518 = cv2.resize(img, (518, 518), interpolation=cv2.INTER_LANCZOS4)\n                img_518 = cv2.cvtColor(img, cv2.COLOR_RGB2BGR) \n                cv2.imwrite(img_out_path_518, img_518, [cv2.IMWRITE_JPEG_QUALITY, 100])\n                batch_518.append(img_out_path_518)\n\n        batch_queue = {\"image_id\": image_id, \"batch_518\": batch_518, \"is_tma\": is_tma}\n        q_in_predictor.put(batch_queue, block=True, timeout=None)\n    q_in_predictor.put(None, block=True, timeout=None)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-12-26T11:46:36.259299Z","iopub.execute_input":"2023-12-26T11:46:36.259755Z","iopub.status.idle":"2023-12-26T11:46:36.318364Z","shell.execute_reply.started":"2023-12-26T11:46:36.259715Z","shell.execute_reply":"2023-12-26T11:46:36.317063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\nimport torch.nn.functional as F\nfrom torch.nn.parameter import Parameter\nimport math\nfrom timm.layers.adaptive_avgmax_pool import SelectAdaptivePool2d\nfrom timm.models._manipulate import checkpoint_seq\n\n  \n\nclass SelfAttentionPooling(nn.Module):\n    def __init__(self, input_dim):\n        super(SelfAttentionPooling, self).__init__()\n        self.W = nn.Linear(input_dim, 1)\n        \n    def forward(self, x):\n        att_w = nn.functional.softmax(self.W(x).squeeze(dim=-1), dim=-1).unsqueeze(dim=-1)\n        x = torch.sum(x * att_w, dim=1)\n        return x\n\n\nclass 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) + ')'\n    \nclass WsiNet(nn.Module):\n    def __init__(self, back_bone, device_id):\n        super().__init__()\n        self.back_bone = back_bone\n        self.model = timm.create_model(back_bone, num_classes=5, pretrained=False, in_chans=3)\n        # print(self.model)\n        if \"coat_lite\" not in back_bone:\n            self.model.set_grad_checkpointing()\n\n        if \"swin_large\" in back_bone:\n            input_fmt = \"NHWC\"\n            self.model.head = nn.Identity()\n        elif \"convnextv2_large\" in back_bone:\n            input_fmt = \"NCHW\"\n            self.model.head = nn.Identity()\n        elif \"convnextv2_base\" in back_bone:\n            input_fmt = \"NCHW\"\n            self.model.head = nn.Identity()\n        elif \"efficientnetv2_b0\" in back_bone:\n            input_fmt = \"NCHW\"\n            self.model.global_pool = nn.Identity()\n            self.model.classifier = nn.Identity()\n        elif \"coat_lite\" in back_bone:\n            input_fmt = \"NCHW\"\n            self.model.head = nn.Identity()\n        elif \"vit\" in back_bone:\n            self.model.fc_norm = nn.Identity()\n            self.model.head_drop = nn.Identity()\n            self.model.head = nn.Identity()\n\n\n        self.IMAGENET_DEFAULT_MEAN = torch.tensor([0.485, 0.456, 0.406]).to(device_id)\n        self.IMAGENET_DEFAULT_STD = torch.tensor([0.229, 0.224, 0.225]).to(device_id)\n        if \"coat_lite\" not in back_bone and \"vit\" not in back_bone:\n            self.global_pool = GeM()\n        if \"swin_large\" in back_bone:\n            self.num_feature = 1536\n        elif \"convnextv2_large\" in back_bone:\n            self.num_feature = 1536\n        elif \"convnextv2_base\" in back_bone:\n            self.num_feature = 1024\n        elif \"efficientnetv2_b0\" in back_bone:\n            self.num_feature = 1280\n        elif \"coat_lite\" in back_bone:\n            self.num_feature = 512\n        elif \"vit\" in back_bone:\n            self.num_feature = 384\n\n        self.atten_pooling = SelfAttentionPooling(self.num_feature)\n        self.fc = nn.Linear(self.num_feature, 5)\n\n    def forward(self, x):\n        B = x.shape[0]\n        N_TILE = x.shape[1]\n        x = x.reshape(B*N_TILE, x.shape[2], x.shape[3], x.shape[4])\n        x = x.transpose(1, 2).transpose(1, 3).contiguous()\n        x = x/255.0\n        x = (x - self.IMAGENET_DEFAULT_MEAN[None,:, None, None])/self.IMAGENET_DEFAULT_STD[None,:, None, None]\n        features = self.model(x)\n        \n\n        # features = features.transpose(1, 2).transpose(1, 3).contiguous()\n        \n        # features = self.global_pool(features)\n        \n        features = features.reshape(B, N_TILE, -1)\n        features = self.atten_pooling(features)\n        logit_class = self.fc(features)\n\n        return logit_class","metadata":{"execution":{"iopub.status.busy":"2023-12-26T11:42:31.280044Z","iopub.execute_input":"2023-12-26T11:42:31.280288Z","iopub.status.idle":"2023-12-26T11:42:31.301178Z","shell.execute_reply.started":"2023-12-26T11:42:31.280267Z","shell.execute_reply":"2023-12-26T11:42:31.300348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\nLABEL_MAP = {0: \"CC\", 1: \"EC\", 2: \"HGSC\", 3: \"LGSC\", 4: \"MC\", 5: \"Other\", 6: \"AAA\"}\ndef predict_worker(q_in):\n    df_sub = pd.read_csv(f\"{ROOT_DIR}/sample_submission.csv\")\n    wsi_weights = [1/5, 1/5, 1/5, 1/5, 1/5]\n#     tma_weights = [1/10, 1/10, 1/10, 1/10, 1/10, 1/10, 1/10, 1/10, 1/10, 1/10]\n    wsi_weight_paths = [\n        \"/kaggle/input/ubctmamodels/fold_0.pth\",\n    \"/kaggle/input/ubctmamodels/fold_1.pth\",\n    \"/kaggle/input/ubctmamodels/fold_2.pth\",\n    \"/kaggle/input/ubctmamodels/fold_3.pth\",\n    \"/kaggle/input/ubctmamodels/fold_4.pth\"]\n        \n\n    wsi_model_types = [\"vit\", \"vit\", \"vit\", \"vit\", \"vit\"]\n#     tma_model_types = [\"convnextv2_large\", \"convnextv2_large\", \"convnextv2_large\", \"convnextv2_large\", \"convnextv2_large\",\n#                       \"swin_large\", \"swin_large\", \"swin_large\", \"swin_large\", \"swin_large\"]\n    wsi_models = []\n    tma_models = []\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    for weight_path, model_type in zip(wsi_weight_paths, wsi_model_types):\n        if model_type == \"swin_large\":\n            model = WsiNet(\"swin_large_patch4_window12_384.ms_in22k\", device)\n        elif model_type == \"convnextv2_large\":\n            model = WsiNet(\"convnextv2_large.fcmae_ft_in22k_in1k_384\", device)\n        elif model_type == \"convnextv2_base\":\n            model = WsiNet(\"convnextv2_base.fcmae_ft_in22k_in1k_384\", device)\n        elif model_type == \"eff_b0\":\n            model = WsiNet(\"efficientnet_b0.ra_in1k\", device)\n        elif model_type == \"coat\":\n            model = WsiNet(\"coat_lite_medium_384.in1k\", device)\n        elif model_type == \"vit\":\n            model = WsiNet(\"vit_small_patch14_reg4_dinov2.lvd142m\", device)\n        model.load_state_dict(torch.load(weight_path))\n        model.to(device)\n        model.eval()\n        wsi_models.append(model)\n \n    \n    preds = []\n    image_id_list = []\n    count = 0\n    with torch.no_grad():\n        while True:\n            batch_queue = q_in.get()\n\n            if batch_queue is None:\n                count += 1\n                if count == 3:\n                    break\n                continue\n            t1 = time.time()\n            image_id = batch_queue[\"image_id\"]\n            paths_518 = batch_queue[\"batch_518\"]\n            batch_518_1 = []\n            batch_518_2 = []\n            \n            batch_518_flip = []\n            is_tma = batch_queue[\"is_tma\"]\n            \n            if is_tma == True:\n                for path in paths_518:\n                    img_518 = cv2.imread(path)\n                    img_518 = cv2.cvtColor(img_518, cv2.COLOR_BGR2RGB) \n                    batch_518_1.append(img_518)\n                    img_518_flip = cv2.flip(img_518, 1)\n                    batch_518_flip.append(img_518_flip)\n            else:\n                n_tiles = len(paths_518)\n                for i in range(0, 16):\n                    path = paths_518[i]\n                    img_518 = cv2.imread(path)\n                    img_518 = cv2.cvtColor(img_518, cv2.COLOR_BGR2RGB) \n                    batch_518_1.append(img_518)\n                indices = get_indices(16, n_tiles)\n                for i in indices:\n                    path = paths_518[i]\n                    img_518 = cv2.imread(path)\n                    img_518 = cv2.cvtColor(img_518, cv2.COLOR_BGR2RGB) \n                    batch_518_2.append(img_518)\n                    img_518_flip = cv2.flip(img_518, 1)\n                    batch_518_flip.append(img_518_flip)\n\n            batch_518_1 = np.array(batch_518_1)\n            batch_518_flip = np.array(batch_518_flip)\n            batch_518_1 = torch.from_numpy(batch_518_1).to(\"cuda:0\").float()\n            batch_518_1 = torch.unsqueeze(batch_518_1, 0)\n            batch_518_flip = torch.from_numpy(batch_518_flip).to(\"cuda:0\").float()\n            batch_518_flip = torch.unsqueeze(batch_518_flip, 0)\n            id_dir_518 = os.path.join(\"/kaggle/working/518\", str(image_id))\n            shutil.rmtree(id_dir_518) \n            if is_tma != True:\n                batch_518_2 = np.array(batch_518_2)\n                batch_518_2 = torch.from_numpy(batch_518_2).to(\"cuda:0\").float()\n                batch_518_2 = torch.unsqueeze(batch_518_2, 0)\n\n            if is_tma == True:\n                ensemble_probs = torch.zeros((1, 5)).to(device)\n                for i, model in enumerate(wsi_models):\n                    model.eval()\n                    model_type = wsi_model_types[i]\n                    if model_type == \"vit\":\n                        logit_class = model(batch_518_1)\n                        logit_class_flip = model(batch_518_flip)\n                        probs = torch.sigmoid(logit_class)\n                        probs_flip = torch.sigmoid(logit_class_flip)\n                        ensemble_probs += wsi_weights[i]*(probs+probs_flip)/2\n            else:\n                ensemble_probs = torch.zeros((1, 5)).to(device)\n                for i, model in enumerate(wsi_models):\n                    model.eval()\n                    model_type = wsi_model_types[i]\n                    if model_type == \"vit\":\n                        logit_class_1 = model(batch_518_1)\n                        logit_class_2 = model(batch_518_2)\n                        logit_class_flip = model(batch_518_flip)\n                        probs_1 = torch.sigmoid(logit_class_1)\n                        probs_2 = torch.sigmoid(logit_class_2)\n                        probs_flip = torch.sigmoid(logit_class_flip)\n                        print(\"Model {}: \".format(i), probs_1, probs_2, probs_flip)\n                        ensemble_probs += wsi_weights[i]*(probs_1 + probs_2 + probs_flip)/3\n\n            scored, predicted = torch.max(ensemble_probs, 1)\n            predicted = predicted.detach().cpu().numpy()\n            image_id_list.append(image_id)\n\n            t2 = time.time()\n            print(\"===Predict time: \", t2 - t1, scored[0])\n\n#             if scored[0] < 0.5:\n#                 preds.append(5)\n#             else:\n#                 preds.append(6)\n            preds.append(predicted[0])\n            print(image_id, LABEL_MAP[predicted[0]])   \n\n            \n    pred_labels = []\n    for i in range(len(preds)):\n        pred_labels.append(LABEL_MAP[preds[i]])\n\n    df_sub[\"image_id\"] = image_id_list\n    df_sub[\"label\"] = pred_labels\n    df_sub.to_csv(\"submission.csv\", index=False)\n    df_sub","metadata":{"execution":{"iopub.status.busy":"2023-12-26T11:46:41.61025Z","iopub.execute_input":"2023-12-26T11:46:41.610678Z","iopub.status.idle":"2023-12-26T11:46:41.677895Z","shell.execute_reply.started":"2023-12-26T11:46:41.610638Z","shell.execute_reply":"2023-12-26T11:46:41.676914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"q_in_predictor = Queue(maxsize=4000)\npredict_process = Process(target=predict_worker, args=(q_in_predictor, ))\npredict_process.start()\n\ndf = pd.read_csv(f\"{ROOT_DIR}/test.csv\")\ndf['label'] = 0 # dummy\ntotal_images = len(df)         \nread_0_process = Process(target=read_worker, args=(q_in_predictor, df.iloc[:total_images//3,:], ))\nread_1_process = Process(target=read_worker, args=(q_in_predictor, df.iloc[total_images//3:2*total_images//3,:], ))\nread_2_process = Process(target=read_worker, args=(q_in_predictor, df.iloc[2*total_images//3:,:], ))\nread_0_process.start()\nread_1_process.start()\nread_2_process.start()\n\nread_0_process.join()\nread_1_process.join()\nread_2_process.join()\npredict_process.join()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}