{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":1760030,"sourceType":"datasetVersion","datasetId":1046169},{"sourceId":6774553,"sourceType":"datasetVersion","datasetId":3898019},{"sourceId":7314103,"sourceType":"datasetVersion","datasetId":4023955}],"dockerImageVersionId":30588,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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!pip list | grep pyvips","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-31T16:33:15.405512Z","iopub.execute_input":"2023-12-31T16:33:15.406016Z","iopub.status.idle":"2023-12-31T16:34:41.81444Z","shell.execute_reply.started":"2023-12-31T16:33:15.40599Z","shell.execute_reply":"2023-12-31T16:34:41.813349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport pandas as pd\nimport numpy as np\nfrom joblib import Parallel, delayed\nimport pyvips\nimport random\nimport time\n\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as T\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.multiprocessing\n\nfrom timm.models.vision_transformer import VisionTransformer\n\nfrom PIL import ImageFile, Image\n\ntorch.multiprocessing.set_sharing_strategy('file_system')\nImageFile.LOAD_TRUNCATED_IMAGES = True\nImage.MAX_IMAGE_PIXELS = None\nos.environ['VIPS_CONCURRENCY'] = '4'\nos.environ['VIPS_DISC_THRESHOLD'] = '15gb'\n\ndef set_seed(seed):\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    random.seed(seed)\n\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\nset_seed(2024)","metadata":{"execution":{"iopub.status.busy":"2023-12-31T16:34:41.816658Z","iopub.execute_input":"2023-12-31T16:34:41.817405Z","iopub.status.idle":"2023-12-31T16:34:47.448621Z","shell.execute_reply.started":"2023-12-31T16:34:41.817364Z","shell.execute_reply":"2023-12-31T16:34:47.447845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ABMIL(nn.Module):\n    def __init__(self, in_dim, feat_dim, attn_dim, num_classes):\n        super().__init__()\n        self.downlinear = nn.Sequential(nn.Linear(in_dim, feat_dim), nn.ReLU())\n        self.attention_V = nn.Sequential(nn.Linear(feat_dim, attn_dim), nn.Tanh())\n        self.attention_U = nn.Sequential(nn.Linear(feat_dim, attn_dim), nn.Sigmoid())\n        self.attention_weights = nn.Linear(attn_dim, 1)\n        self.classifier = nn.Linear(feat_dim, num_classes)\n\n        self.apply(self._init_weights)\n\n    def _init_weights(self, m):\n        if isinstance(m, nn.Linear):\n            nn.init.normal_(m.weight, std=.02)\n            if isinstance(m, nn.Linear) and m.bias is not None:\n                nn.init.constant_(m.bias, 0)\n\n    def forward(self, x):\n        x = self.downlinear(x)\n\n        A_V = self.attention_V(x)\n        A_U = self.attention_U(x)\n        A = self.attention_weights(A_V * A_U)\n        A = torch.transpose(A, 1, 0)\n        A = torch.softmax(A, dim=1)\n        x = torch.mm(A, x)\n\n        scores = self.classifier(x)\n\n        return scores\n\n\nclass FCLayer(nn.Module):\n    def __init__(self, in_size, out_size=1):\n        super(FCLayer, self).__init__()\n        self.fc = nn.Sequential(nn.Linear(in_size, out_size))\n\n    def forward(self, feats):\n        x = self.fc(feats)\n        return feats, x\n\n\nclass IClassifier(nn.Module):\n    def __init__(self, feature_size, output_class):\n        super(IClassifier, self).__init__()\n\n        self.fc = nn.Linear(feature_size, output_class)\n\n    def forward(self, feats):\n        c = self.fc(feats.view(feats.shape[0], -1))  # N x C\n        return feats.view(feats.shape[0], -1), c\n\n\nclass BClassifier(nn.Module):\n    def __init__(self, input_size, output_class, dropout_v=0.0, nonlinear=True, passing_v=False):  # K, L, N\n        super(BClassifier, self).__init__()\n        if nonlinear:\n            self.q = nn.Sequential(nn.Linear(input_size, 128), nn.ReLU(), nn.Linear(128, 128), nn.Tanh())\n        else:\n            self.q = nn.Linear(input_size, 128)\n        if passing_v:\n            self.v = nn.Sequential(\n                nn.Dropout(dropout_v),\n                nn.Linear(input_size, input_size),\n                nn.ReLU()\n            )\n        else:\n            self.v = nn.Identity()\n\n        ### 1D convolutional layer that can handle multiple class (including binary)\n        self.fcc = nn.Conv1d(output_class, output_class, kernel_size=input_size)\n\n    def forward(self, feats, c):  # N x K, N x C\n        device = feats.device\n        V = self.v(feats)  # N x V, unsorted\n        Q = self.q(feats).view(feats.shape[0], -1)  # N x Q, unsorted\n\n        # handle multiple classes without for loop\n        _, m_indices = torch.sort(c, 0, descending=True)  # sort class scores along the instance dimension, m_indices in shape N x C\n        m_feats = torch.index_select(feats, dim=0, index=m_indices[0, :])  # select critical instances, m_feats in shape C x K\n        q_max = self.q(m_feats)  # compute queries of critical instances, q_max in shape C x Q\n        A = torch.mm(Q, q_max.transpose(0, 1))  # compute inner product of Q to each entry of q_max, A in shape N x C, each column contains unnormalized attention scores\n        A = F.softmax(A / torch.sqrt(torch.tensor(Q.shape[1], dtype=torch.float32, device=device)), 0)  # normalize attention scores, A in shape N x C,\n        B = torch.mm(A.transpose(0, 1), V)  # compute bag representation, B in shape C x V\n\n        B = B.view(1, B.shape[0], B.shape[1])  # 1 x C x V\n        C = self.fcc(B)  # 1 x C x 1\n        C = C.view(1, -1)\n        return C, A, B\n\n\nclass DSMIL(nn.Module):\n    def __init__(self, i_classifier, b_classifier):\n        super(DSMIL, self).__init__()\n        self.i_classifier = i_classifier\n        self.b_classifier = b_classifier\n\n    def forward(self, x):\n        feats, classes = self.i_classifier(x)\n        prediction_bag, A, B = self.b_classifier(feats, classes)\n\n        return classes, prediction_bag, A, B\n\n!pip install /kaggle/input/einops-030/einops-0.3.0-py2.py3-none-any.whl\nfrom einops import rearrange, reduce\nfrom torch import einsum\nfrom math import ceil\n\n\ndef exists(val):\n    return val is not None\n\n\ndef moore_penrose_iter_pinv(x, iters=6):\n    device = x.device\n\n    abs_x = torch.abs(x)\n    col = abs_x.sum(dim=-1)\n    row = abs_x.sum(dim=-2)\n    z = rearrange(x, '... i j -> ... j i') / (torch.max(col) * torch.max(row))\n\n    I = torch.eye(x.shape[-1], device=device)\n    I = rearrange(I, 'i j -> () i j')\n\n    for _ in range(iters):\n        xz = x @ z\n        z = 0.25 * z @ (13 * I - (xz @ (15 * I - (xz @ (7 * I - xz)))))\n\n    return z\n\n\nclass NystromAttention(nn.Module):\n    def __init__(\n            self,\n            dim,\n            dim_head=64,\n            heads=8,\n            num_landmarks=256,\n            pinv_iterations=6,\n            residual=True,\n            residual_conv_kernel=33,\n            eps=1e-8,\n            dropout=0.\n    ):\n        super().__init__()\n        self.eps = eps\n        inner_dim = heads * dim_head\n\n        self.num_landmarks = num_landmarks\n        self.pinv_iterations = pinv_iterations\n\n        self.heads = heads\n        self.scale = dim_head ** -0.5\n        self.to_qkv = nn.Linear(dim, inner_dim * 3, bias=False)\n\n        self.to_out = nn.Sequential(\n            nn.Linear(inner_dim, dim),\n            nn.Dropout(dropout)\n        )\n\n        self.residual = residual\n        if residual:\n            kernel_size = residual_conv_kernel\n            padding = residual_conv_kernel // 2\n            self.res_conv = nn.Conv2d(heads, heads, (kernel_size, 1), padding=(padding, 0), groups=heads, bias=False)\n\n    def forward(self, x, mask=None, return_attn=False):\n        b, n, _, h, m, iters, eps = *x.shape, self.heads, self.num_landmarks, self.pinv_iterations, self.eps\n\n        # pad so that sequence can be evenly divided into m landmarks\n\n        remainder = n % m\n        if remainder > 0:\n            padding = m - (n % m)\n            x = F.pad(x, (0, 0, padding, 0), value=0)\n\n            if exists(mask):\n                mask = F.pad(mask, (padding, 0), value=False)\n\n        # derive query, keys, values\n\n        q, k, v = self.to_qkv(x).chunk(3, dim=-1)\n        q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h=h), (q, k, v))\n\n        # set masked positions to 0 in queries, keys, values\n\n        if exists(mask):\n            mask = rearrange(mask, 'b n -> b () n')\n            q, k, v = map(lambda t: t * mask[..., None], (q, k, v))\n\n        q = q * self.scale\n\n        # generate landmarks by sum reduction, and then calculate mean using the mask\n\n        l = ceil(n / m)\n        landmark_einops_eq = '... (n l) d -> ... n d'\n        q_landmarks = reduce(q, landmark_einops_eq, 'sum', l=l)\n        k_landmarks = reduce(k, landmark_einops_eq, 'sum', l=l)\n\n        # calculate landmark mask, and also get sum of non-masked elements in preparation for masked mean\n\n        divisor = l\n        if exists(mask):\n            mask_landmarks_sum = reduce(mask, '... (n l) -> ... n', 'sum', l=l)\n            divisor = mask_landmarks_sum[..., None] + eps\n            mask_landmarks = mask_landmarks_sum > 0\n\n        # masked mean (if mask exists)\n\n        q_landmarks /= divisor\n        k_landmarks /= divisor\n\n        # similarities\n\n        einops_eq = '... i d, ... j d -> ... i j'\n        sim1 = einsum(einops_eq, q, k_landmarks)\n        sim2 = einsum(einops_eq, q_landmarks, k_landmarks)\n        sim3 = einsum(einops_eq, q_landmarks, k)\n\n        # masking\n\n        if exists(mask):\n            mask_value = -torch.finfo(q.dtype).max\n            sim1.masked_fill_(~(mask[..., None] * mask_landmarks[..., None, :]), mask_value)\n            sim2.masked_fill_(~(mask_landmarks[..., None] * mask_landmarks[..., None, :]), mask_value)\n            sim3.masked_fill_(~(mask_landmarks[..., None] * mask[..., None, :]), mask_value)\n\n        # eq (15) in the paper and aggregate values\n\n        attn1, attn2, attn3 = map(lambda t: t.softmax(dim=-1), (sim1, sim2, sim3))\n        attn2_inv = moore_penrose_iter_pinv(attn2, iters)\n\n        out = (attn1 @ attn2_inv) @ (attn3 @ v)\n\n        # add depth-wise conv residual of values\n\n        if self.residual:\n            out += self.res_conv(v)\n\n        # merge and combine heads\n\n        out = rearrange(out, 'b h n d -> b n (h d)', h=h)\n        out = self.to_out(out)\n        out = out[:, -n:]\n\n        if return_attn:\n            attn = attn1 @ attn2_inv @ attn3\n            return out, attn\n\n        return out\n\n\nclass TransLayer(nn.Module):\n    def __init__(self, norm_layer=nn.LayerNorm, dim=512):\n        super().__init__()\n        self.norm = norm_layer(dim)\n        self.attn = NystromAttention(\n            dim=dim,\n            dim_head=dim // 8,\n            heads=8,\n            num_landmarks=dim // 2,  # number of landmarks\n            pinv_iterations=6,  # number of moore-penrose iterations for approximating pinverse. 6 was recommended by the paper\n            residual=True,  # whether to do an extra residual with the value or not. supposedly faster convergence if turned on\n            dropout=0.1\n        )\n\n    def forward(self, x):\n        x = x + self.attn(self.norm(x))\n\n        return x\n\n\nclass PPEG(nn.Module):\n    def __init__(self, dim=512):\n        super(PPEG, self).__init__()\n        self.proj = nn.Conv2d(dim, dim, 7, 1, 7 // 2, groups=dim)\n        self.proj1 = nn.Conv2d(dim, dim, 5, 1, 5 // 2, groups=dim)\n        self.proj2 = nn.Conv2d(dim, dim, 3, 1, 3 // 2, groups=dim)\n\n    def forward(self, x, H, W):\n        B, _, C = x.shape\n        cls_token, feat_token = x[:, 0], x[:, 1:]\n        cnn_feat = feat_token.transpose(1, 2).view(B, C, H, W)\n        x = self.proj(cnn_feat) + cnn_feat + self.proj1(cnn_feat) + self.proj2(cnn_feat)\n        x = x.flatten(2).transpose(1, 2)\n        x = torch.cat((cls_token.unsqueeze(1), x), dim=1)\n        return x\n\n\nclass TransMIL(nn.Module):\n    def __init__(self, in_dim, n_classes):\n        super(TransMIL, self).__init__()\n        self.pos_layer = PPEG(dim=512)\n        self._fc1 = nn.Sequential(nn.Linear(in_dim, 512), nn.ReLU())\n        self.cls_token = nn.Parameter(torch.randn(1, 1, 512))\n        self.n_classes = n_classes\n        self.layer1 = TransLayer(dim=512)\n        self.layer2 = TransLayer(dim=512)\n        self.norm = nn.LayerNorm(512)\n        self._fc2 = nn.Linear(512, self.n_classes)\n\n    def forward(self, h):\n        h = h.float()  # [B, n, 1024]\n\n        h = self._fc1(h)  # [B, n, 512]\n\n        # ---->pad\n        H = h.shape[1]\n        _H, _W = int(np.ceil(np.sqrt(H))), int(np.ceil(np.sqrt(H)))\n        add_length = _H * _W - H\n        h = torch.cat([h, h[:, :add_length, :]], dim=1)  # [B, N, 512]\n\n        # ---->cls_token\n        B = h.shape[0]\n        cls_tokens = self.cls_token.expand(B, -1, -1).to(h.device)\n        h = torch.cat((cls_tokens, h), dim=1)\n\n        # ---->Translayer x1\n        h = self.layer1(h)  # [B, N, 512]\n\n        # ---->PPEG\n        h = self.pos_layer(h, _H, _W)  # [B, N, 512]\n\n        # ---->Translayer x2\n        h = self.layer2(h)  # [B, N, 512]\n\n        # ---->cls_token\n        h = self.norm(h)[:, 0]\n\n        # ---->predict\n        logits = self._fc2(h)  # [B, n_classes]\n\n        return logits","metadata":{"execution":{"iopub.status.busy":"2023-12-31T16:34:47.44983Z","iopub.execute_input":"2023-12-31T16:34:47.450111Z","iopub.status.idle":"2023-12-31T16:35:18.917865Z","shell.execute_reply.started":"2023-12-31T16:34:47.450086Z","shell.execute_reply":"2023-12-31T16:35:18.916745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Union\n\ndef get_valid_cor_list(path, patch_size:int, shuffle:bool, ratio:Union[float, int]):\n    ImageFile.LOAD_TRUNCATED_IMAGES = True\n    Image.MAX_IMAGE_PIXELS = None\n    os.environ['VIPS_CONCURRENCY'] = '4'\n    os.environ['VIPS_DISC_THRESHOLD'] = '15gb'\n    \n    wsi = pyvips.Image.new_from_file(path[0])   \n    wsi_width, wsi_height = wsi.width, wsi.height\n    is_tma = wsi_width < 5000 and wsi_height < 5000\n    \n    thumbnail = pyvips.Image.new_from_file(path[1]) if not is_tma else wsi.copy()\n    thu_width, thu_height = thumbnail.width, thumbnail.height\n    w_r, h_r = wsi_width / thu_width, wsi_height / thu_height\n    down_w, down_h = int(patch_size / w_r), int(patch_size / h_r)\n    \n    cors = [(x, y) for y in range(0, thu_height, down_h) for x in range(0, thu_width, down_w)]\n    cor_list = []\n    for x, y in cors:\n        tile = thumbnail.crop(x, y, min(down_w, thu_width - x), min(down_h, thu_height - y)).numpy()[..., :3]\n        black_bg = np.mean(tile, axis=2) < 20\n        tile[black_bg, :] = 255\n        mask_bg = np.mean(tile, axis=2) > 235\n        if np.sum(mask_bg) < tile.shape[0] * tile.shape[1] * 0.5 or is_tma:\n            cor_list.append((int(x * w_r), int(y * h_r)))\n\n    if shuffle:\n        random.shuffle(cor_list)\n        \n    if 0 < ratio <= 1:\n        cor_list = cor_list[:max(int(len(cor_list) * ratio), 1)]\n    elif ratio > 1:\n        cor_list = cor_list[:min(len(cor_list), ratio)]\n        \n    del wsi, thumbnail\n    gc.collect()\n\n    return cor_list\n\nclass SingleWSIDataset(Dataset):\n    def __init__(self, path, cor_list:list, patch_size:int):\n        super().__init__()\n        self.wsi = pyvips.Image.new_from_file(path)\n        self.cor_list = cor_list\n        self.patch_size = patch_size\n        self.transform = T.Compose([T.ToTensor(), \n                                    T.Resize((224, 224), antialias=True), \n                                    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])\n\n    def __len__(self):\n        return len(self.cor_list)\n\n    def __getitem__(self, idx):\n        x, y = self.cor_list[idx]\n        tile = self.wsi.crop(x, y, min(self.patch_size, self.wsi.width - x), min(self.patch_size, self.wsi.height - y)).numpy()[..., :3]\n        tile = self.transform(tile)\n        return tile","metadata":{"execution":{"iopub.status.busy":"2023-12-31T16:35:18.921161Z","iopub.execute_input":"2023-12-31T16:35:18.921826Z","iopub.status.idle":"2023-12-31T16:35:18.938684Z","shell.execute_reply.started":"2023-12-31T16:35:18.921785Z","shell.execute_reply":"2023-12-31T16:35:18.937764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_core = 2\nmode = 'test'\ndata_dir = '/kaggle/input/UBC-OCEAN'\ntest_csv = pd.read_csv(os.path.join(data_dir, f'{mode}.csv'))#.iloc[:8]\nlabel_names = ['CC', 'EC', 'HGSC', 'LGSC', 'MC', 'Other']\nnum_classes = len(label_names)\nlabel_dict = {label_names[i]: i for i in range(len(label_names))}\npatch_size = 512\nratio = 1024\n\npath_list = []\ntotal_cor_list_list = []\nstart = time.time()\nfor i in range(len(test_csv)):\n    sample = test_csv.iloc[i]\n    name = str(sample['image_id'])\n    print(name)\n    path_list.append([os.path.join(data_dir, f'{mode}_images', name + '.png'), os.path.join(data_dir, f'{mode}_thumbnails', name + '_thumbnail.png')])\n\n    if len(path_list) == num_core or i == len(test_csv) - 1:\n        if len(path_list) == 1:\n            cor_list_list = [get_valid_cor_list(path_list[0], patch_size, True, ratio)]\n        else:\n            cor_list_list = Parallel(n_jobs=len(path_list))(delayed(get_valid_cor_list)(path, patch_size, True, ratio) for path in path_list)\n        \n        total_cor_list_list.extend(cor_list_list)\n        path_list = []\n        \n        del cor_list_list\n        gc.collect()\n\nassert len(total_cor_list_list) == len(test_csv)\nend = time.time()\nprint('提取坐标用时', end - start)","metadata":{"execution":{"iopub.status.busy":"2023-12-31T16:35:18.940146Z","iopub.execute_input":"2023-12-31T16:35:18.940868Z","iopub.status.idle":"2023-12-31T16:35:23.606029Z","shell.execute_reply.started":"2023-12-31T16:35:18.940842Z","shell.execute_reply":"2023-12-31T16:35:23.605171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.cuda.amp import autocast\n\ndevice = torch.device('cuda:0')\n\nvit_p16 = VisionTransformer(patch_size=16, embed_dim=384, num_heads=6, num_classes=0)\nvit_p16.load_state_dict(torch.load('/kaggle/input/checkpoints/dino_vit_small_patch16_ep200.torch', map_location='cpu'))\nvit_p16 = vit_p16.to(device)\nvit_p16 = nn.DataParallel(vit_p16)\nvit_p16.eval()\n\nvit_p8 = VisionTransformer(patch_size=8, embed_dim=384, num_heads=6, num_classes=0)\nvit_p8.load_state_dict(torch.load('/kaggle/input/checkpoints/dino_vit_small_patch8_ep200.torch', map_location='cpu'))\nvit_p8 = vit_p8.to(device)\nvit_p8 = nn.DataParallel(vit_p8)\nvit_p8.eval()\n\nabmil1 = ABMIL(384, 512, 128, num_classes)\nabmil1.load_state_dict(torch.load('/kaggle/input/checkpoints/transform_wsi_vitp16_abmil_1.0_15ep.pth', map_location='cpu'))\nabmil1 = abmil1.to(device)\nabmil1.eval()\n\nabmil2 = ABMIL(384, 512, 128, num_classes)\nabmil2.load_state_dict(torch.load('/kaggle/input/checkpoints/transform_wsi_vitp8_abmil_1.0_15ep.pth', map_location='cpu'))\nabmil2 = abmil2.to(device)\nabmil2.eval()\n\ndsmil1 = DSMIL(IClassifier(384, num_classes), BClassifier(384, num_classes))\ndsmil1.load_state_dict(torch.load('/kaggle/input/checkpoints/transform_wsi_vitp16_dsmil_1.0_15ep.pth', map_location='cpu'))\ndsmil1 = dsmil1.to(device)\ndsmil1.eval()\n\ndsmil2 = DSMIL(IClassifier(384, num_classes), BClassifier(384, num_classes))\ndsmil2.load_state_dict(torch.load('/kaggle/input/checkpoints/transform_wsi_vitp8_dsmil_1.0_15ep.pth', map_location='cpu'))\ndsmil2 = dsmil2.to(device)\ndsmil2.eval()\n\ntransmil1 = TransMIL(384, num_classes)\ntransmil1.load_state_dict(torch.load('/kaggle/input/checkpoints/transform_wsi_vitp16_transmil_1.0_15ep.pth', map_location='cpu'))\ntransmil1 = transmil1.to(device)\ntransmil1.eval()\n\ntransmil2 = TransMIL(384, num_classes)\ntransmil2.load_state_dict(torch.load('/kaggle/input/checkpoints/transform_wsi_vitp8_transmil_1.0_15ep.pth', map_location='cpu'))\ntransmil2 = transmil2.to(device)\ntransmil2.eval()\n\nsubmission = pd.DataFrame(columns=['image_id', 'label'])\n\nwith torch.no_grad(), autocast():\n    for i in range(len(test_csv)):\n        start = time.time()\n        \n        sample = test_csv.iloc[i]\n        name = str(sample['image_id'])\n        \n        path = os.path.join(data_dir, f'{mode}_images', name + '.png')\n        \n        wsidataset = SingleWSIDataset(path, total_cor_list_list[i], patch_size)\n        loader1 = DataLoader(wsidataset, 1024, True, num_workers=0, pin_memory = True)\n        loader2 = DataLoader(wsidataset, 1024, True, num_workers=0, pin_memory = True)\n\n        features1 = torch.empty((0, 384)).to(device, non_blocking=True)\n        features2 = torch.empty((0, 384)).to(device, non_blocking=True)\n        \n        for tiles1 in loader1:\n            tiles1 = tiles1.to(device, non_blocking=True)\n            feat1 = vit_p16(tiles1)\n            features1 = torch.cat([features1, feat1])\n\n            del tiles1, feat1\n            gc.collect()\n            torch.cuda.empty_cache()\n        \n        for tiles2 in loader2:\n            tiles2 = tiles2.to(device, non_blocking=True)\n            feat2 = vit_p8(tiles2)\n            features2 = torch.cat([features2, feat2])\n\n            del tiles2, feat2\n            gc.collect()\n            torch.cuda.empty_cache()\n        \n        abmil_scores1 = torch.softmax(abmil1(features1), 1)\n        abmil_scores2 = torch.softmax(abmil2(features2), 1)\n\n        classes, bag_prediction, _, _ = dsmil1(features1)\n        max_prediction, index = torch.max(classes, 0, True)\n        dsmil_scores1 = 0.5 * torch.softmax(max_prediction, 1) + 0.5 * torch.softmax(bag_prediction, 1)\n        classes, bag_prediction, _, _ = dsmil2(features2)\n        max_prediction, index = torch.max(classes, 0, True)\n        dsmil_scores2 = 0.5 * torch.softmax(max_prediction, 1) + 0.5 * torch.softmax(bag_prediction, 1)\n\n        transmil_scores1 = torch.softmax(transmil1(features1.unsqueeze(0)), 1)\n        transmil_scores2 = torch.softmax(transmil2(features2.unsqueeze(0)), 1)\n\n        pred = torch.max((abmil_scores1.cpu() + abmil_scores2.cpu() + \n                          dsmil_scores1.cpu() + dsmil_scores2.cpu() + \n                          transmil_scores1.cpu() + transmil_scores2.cpu())/6, 1)\n        pred_label = label_names[pred.indices]\n\n        submission.loc[len(submission)] = [int(name), pred_label]\n        \n        del wsidataset, loader1, loader2, features1, features2\n        gc.collect()\n        torch.cuda.empty_cache()\n        \n        end = time.time()\n        print('finish', name, end - start)\n\nsubmission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-12-31T16:35:23.607105Z","iopub.execute_input":"2023-12-31T16:35:23.607382Z","iopub.status.idle":"2023-12-31T16:36:15.136768Z","shell.execute_reply.started":"2023-12-31T16:35:23.607358Z","shell.execute_reply":"2023-12-31T16:36:15.135681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.read_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-12-31T16:36:15.139043Z","iopub.execute_input":"2023-12-31T16:36:15.139424Z","iopub.status.idle":"2023-12-31T16:36:15.156165Z","shell.execute_reply.started":"2023-12-31T16:36:15.139391Z","shell.execute_reply":"2023-12-31T16:36:15.154926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from torch.cuda.amp import autocast\n\n# num_core = 4\n# mode = 'train'\n# data_dir = '/kaggle/input/UBC-OCEAN'\n# test_csv = pd.read_csv(os.path.join(data_dir, f'{mode}.csv')).iloc[:5]\n# label_names = ['CC', 'EC', 'HGSC', 'LGSC', 'MC', 'Other']\n# num_classes = len(label_names)\n# label_dict = {label_names[i]: i for i in range(len(label_names))}\n# patch_size = 512\n# device = torch.device('cuda:0')\n\n# vit_p16 = VisionTransformer(patch_size=16, embed_dim=384, num_heads=6, num_classes=0)\n# vit_p16.load_state_dict(torch.load('/kaggle/input/checkpoints/dino_vit_small_patch16_ep200.torch', map_location='cpu'))\n# vit_p16 = vit_p16.to(device)\n# # vit_p16 = nn.DataParallel(vit_p16)\n# vit_p16.eval()\n\n# vit_p8 = VisionTransformer(patch_size=8, embed_dim=384, num_heads=6, num_classes=0)\n# vit_p8.load_state_dict(torch.load('/kaggle/input/checkpoints/dino_vit_small_patch8_ep200.torch', map_location='cpu'))\n# vit_p8 = vit_p8.to(device)\n# # vit_p8 = nn.DataParallel(vit_p8)\n# vit_p8.eval()\n\n# abmil1 = ABMIL(384, 512, 128, num_classes)\n# abmil1.load_state_dict(torch.load('/kaggle/input/checkpoints/final_wsi_vitp16_abmil_1.0_15ep.pth', map_location='cpu'))\n# abmil1 = abmil1.to(device)\n# abmil1.eval()\n\n# abmil2 = ABMIL(384, 512, 128, num_classes)\n# abmil2.load_state_dict(torch.load('/kaggle/input/checkpoints/final_wsi_vitp8_abmil_1.0_15ep.pth', map_location='cpu'))\n# abmil2 = abmil2.to(device)\n# abmil2.eval()\n\n# dsmil1 = DSMIL(IClassifier(384, num_classes), BClassifier(384, num_classes))\n# dsmil1.load_state_dict(torch.load('/kaggle/input/checkpoints/final_wsi_vitp16_dsmil_1.0_15ep.pth', map_location='cpu'))\n# dsmil1 = dsmil1.to(device)\n# dsmil1.eval()\n\n# dsmil2 = DSMIL(IClassifier(384, num_classes), BClassifier(384, num_classes))\n# dsmil2.load_state_dict(torch.load('/kaggle/input/checkpoints/final_wsi_vitp8_dsmil_1.0_15ep.pth', map_location='cpu'))\n# dsmil2 = dsmil2.to(device)\n# dsmil2.eval()\n\n# transmil1 = TransMIL(384, num_classes)\n# transmil1.load_state_dict(torch.load('/kaggle/input/checkpoints/final_wsi_vitp16_transmil_1.0_15ep.pth', map_location='cpu'))\n# transmil1 = transmil1.to(device)\n# transmil1.eval()\n\n# transmil2 = TransMIL(384, num_classes)\n# transmil2.load_state_dict(torch.load('/kaggle/input/checkpoints/final_wsi_vitp8_transmil_1.0_15ep.pth', map_location='cpu'))\n# transmil2 = transmil2.to(device)\n# transmil2.eval()\n\n# submission = pd.DataFrame(columns=['image_id', 'label'])\n\n# path_list = []\n# name_list = []\n# with torch.no_grad(), autocast():\n#     for i in range(len(test_csv)):\n#         sample = test_csv.iloc[i]\n#         name = str(sample['image_id'])\n#         print('name', name)\n#         name_list.append(name)\n#         path_list.append([os.path.join(data_dir, f'{mode}_images', name + '.png'), os.path.join(data_dir, f'{mode}_thumbnails', name + '_thumbnail.png')])\n\n#         if len(path_list) == num_core or i == len(test_csv) - 1:\n#             start = time.time()\n#             if len(path_list) == 1:\n#                 total_cor_list_list = [get_valid_cor_list(path_list[0], patch_size, True, 3072)]\n#             else:\n#                 total_cor_list_list = Parallel(n_jobs=len(path_list))(delayed(get_valid_cor_list)(path, patch_size, True, 3072) for path in path_list)\n            \n#             for j in range(len(total_cor_list_list)):\n#                 wsidataset = SingleWSIDataset(path_list[j][0], total_cor_list_list[j], patch_size)\n#                 loader1 = DataLoader(wsidataset, 512, True, pin_memory = True)\n#                 loader2 = DataLoader(wsidataset, 512, True, pin_memory = True)\n                \n#                 features1 = torch.empty((0, 384)).to(device, non_blocking=True)\n#                 features2 = torch.empty((0, 384)).to(device, non_blocking=True)\n\n#                 for tiles1, tiles2 in zip(loader1, loader2):\n#                     tiles1 = tiles1.to(device, non_blocking=True)\n#                     feat1 = vit_p16(tiles1)\n#                     features1 = torch.cat([features1, feat1])\n\n#                     tiles2 = tiles2.to(device, non_blocking=True)\n#                     feat2 = vit_p8(tiles2)\n#                     features2 = torch.cat([features2, feat2])\n                    \n#                     del tiles1, tiles2, feat1, feat2\n#                     gc.collect()\n#                     torch.cuda.empty_cache()\n\n#                 abmil_scores1 = torch.softmax(abmil1(features1), 1)\n#                 abmil_scores2 = torch.softmax(abmil2(features2), 1)\n\n#                 classes, bag_prediction, _, _ = dsmil1(features1)\n#                 max_prediction, index = torch.max(classes, 0, True)\n#                 dsmil_scores1 = 0.5 * torch.softmax(max_prediction, 1) + 0.5 * torch.softmax(bag_prediction, 1)\n#                 classes, bag_prediction, _, _ = dsmil2(features2)\n#                 max_prediction, index = torch.max(classes, 0, True)\n#                 dsmil_scores2 = 0.5 * torch.softmax(max_prediction, 1) + 0.5 * torch.softmax(bag_prediction, 1)\n\n#                 transmil_scores1 = torch.softmax(transmil1(features1.unsqueeze(0)), 1)\n#                 transmil_scores2 = torch.softmax(transmil2(features2.unsqueeze(0)), 1)\n\n#                 pred = torch.max((abmil_scores1.cpu() + abmil_scores2.cpu() + \n#                                   dsmil_scores1.cpu() + dsmil_scores2.cpu() + \n#                                   transmil_scores1.cpu() + transmil_scores2.cpu())/6, 1)\n#                 pred_label = label_names[pred.indices]\n\n#                 submission.loc[len(submission)] = [name_list[j], pred_label]\n                \n#                 del wsidataset, loader1, loader2, features1, features2\n#                 gc.collect()\n#                 torch.cuda.empty_cache()\n            \n#             name_list = []\n#             path_list = []\n        \n#             del total_cor_list_list\n#             gc.collect()\n\n#             end = time.time()\n#             print('finish', end - start)\n# submission.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-12-31T16:36:15.157588Z","iopub.execute_input":"2023-12-31T16:36:15.15793Z","iopub.status.idle":"2023-12-31T16:36:15.169543Z","shell.execute_reply.started":"2023-12-31T16:36:15.1579Z","shell.execute_reply":"2023-12-31T16:36:15.168641Z"},"trusted":true},"execution_count":null,"outputs":[]}]}