{"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":7246771,"sourceType":"datasetVersion","datasetId":3770816},{"sourceId":7294550,"sourceType":"datasetVersion","datasetId":4023955},{"sourceId":7300920,"sourceType":"datasetVersion","datasetId":4162804}],"dockerImageVersionId":30588,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Train Notebook: https://www.kaggle.com/hustzx/2nd-0-61-train-abmil-dsmil-transmil","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-30T14:03:28.919981Z","iopub.execute_input":"2023-12-30T14:03:28.920614Z","iopub.status.idle":"2023-12-30T14:04:54.670326Z","shell.execute_reply.started":"2023-12-30T14:03:28.920587Z","shell.execute_reply":"2023-12-30T14:04:54.669193Z"}}},{"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":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport pandas as pd\nimport numpy as np\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 = 5_000_000_000\nos.environ['VIPS_CONCURRENCY'] = '4'\nos.environ['VIPS_DISC_THRESHOLD'] = '15gb'","metadata":{"execution":{"iopub.status.busy":"2023-12-30T14:04:54.672646Z","iopub.execute_input":"2023-12-30T14:04:54.673501Z","iopub.status.idle":"2023-12-30T14:04:59.885571Z","shell.execute_reply.started":"2023-12-30T14:04:54.673455Z","shell.execute_reply":"2023-12-30T14:04:59.884608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    os.environ['PYTHONHASHSEED'] = str(seed)\n#     random.seed(seed)\n#     np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.backends.cudnn.deterministic = True\nseed_everything(42)","metadata":{"execution":{"iopub.status.busy":"2023-12-30T14:04:59.886762Z","iopub.execute_input":"2023-12-30T14:04:59.887037Z","iopub.status.idle":"2023-12-30T14:04:59.960744Z","shell.execute_reply.started":"2023-12-30T14:04:59.887012Z","shell.execute_reply":"2023-12-30T14:04:59.959683Z"},"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-30T14:04:59.963176Z","iopub.execute_input":"2023-12-30T14:04:59.963495Z","iopub.status.idle":"2023-12-30T14:05:31.840384Z","shell.execute_reply.started":"2023-12-30T14:04:59.963466Z","shell.execute_reply":"2023-12-30T14:05:31.839449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SingleWSIDataset(Dataset):\n    def __init__(self, data_path: str, wsi_name: str, patch_size: int, ratio, mode: str):\n        super().__init__()\n        self.data_path = data_path\n        self.wsi_name = wsi_name\n        self.ratio = ratio\n        assert mode in ['train', 'test']\n        self.mode = mode\n        self.wsi = pyvips.Image.new_from_file(os.path.join(data_path, f'{mode}_images', wsi_name + '.png'))\n        self.img = Image.open(os.path.join(data_path, f'{mode}_images', wsi_name + '.png'))\n        self.is_tma = self.wsi.height < 5000 and self.wsi.width < 5000\n        self.patch_size = patch_size\n        self.transform = T.Compose([T.ToTensor(), T.Resize((224, 224), antialias=True), T.Normalize(mean=[0.2585, 0.2556, 0.2506], std=[0.229, 0.224, 0.225])])\n        self.cor_list = self.get_patch()\n\n    def get_patch(self):\n        cor_list = []\n        if self.is_tma:\n            thumbnail = self.wsi\n        else:\n            thumbnail = pyvips.Image.new_from_file(os.path.join(self.data_path, f'{self.mode}_thumbnails', self.wsi_name + '_thumbnail.png'))\n        wsi_width, wsi_height = self.wsi.width, self.wsi.height\n        thu_width, thu_height = thumbnail.width, thumbnail.height\n        h_r, w_r = wsi_height / thu_height, wsi_width / thu_width\n        down_h, down_w = int(self.patch_size / h_r), int(self.patch_size / w_r)\n        cors = [(x, y) for y in range(0, thu_height, down_h) for x in range(0, thu_width, down_w)]\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) < min(down_h, thu_height - y) * min(down_w, thu_width - x) * 0.7 or len(cor_list) == 0 or self.is_tma:\n                cor_list.append((int(x * w_r), int(y * h_r)))\n        if self.is_tma:\n            return cor_list\n        if self.wsi.height < 40000 and self.wsi.width < 40000:\n            R_ratio = 0.8\n        elif self.wsi.height < 80000 and self.wsi.width < 80000:\n            R_ratio = 0.6\n        else:\n            R_ratio = 0.5\n        random.shuffle(cor_list)\n        cor_list = cor_list[:max(int(len(cor_list) * R_ratio), 1)]\n        return cor_list\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.img.crop((x, y, min(x + self.patch_size, self.img.width - 1), min(y + self.patch_size, self.img.height - 1)))\n        tile = self.transform(tile)\n        return tile","metadata":{"execution":{"iopub.status.busy":"2023-12-30T14:05:31.841809Z","iopub.execute_input":"2023-12-30T14:05:31.842121Z","iopub.status.idle":"2023-12-30T14:05:31.860399Z","shell.execute_reply.started":"2023-12-30T14:05:31.842088Z","shell.execute_reply":"2023-12-30T14:05:31.859516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"abmil1_ckpts = [\n        ('/kaggle/input/202312-ubc/submit29_sample/submit29_sample/wsi_vitp16_abmil_0.5_21ep_foldall_soup.pth', 0.6),\n        ('/kaggle/input/202312-ubc/submit28/submit28/wsi_vitp16_abmil_0.5_21ep_foldall_soup.pth',0.3),\n        ('/kaggle/input/checkpoints/final_wsi_vitp16_abmil_1.0_15ep.pth',0.1)\n    ]\n\n\nabmil2_ckpts = [\n    ('/kaggle/input/202312-ubc/submit29_sample/submit29_sample/wsi_vitp8_abmil_0.5_21ep_foldall_soup.pth', 0.5),\n    ('/kaggle/input/202312-ubc/submit28/submit28/wsi_vitp8_abmil_0.5_21ep_foldall_soup.pth', 0.4),\n    ('/kaggle/input/checkpoints/final_wsi_vitp8_abmil_1.0_15ep.pth',0.1)\n]\n\n\ndsmil1_ckpts = [\n    ('/kaggle/input/202312-ubc/submit29_sample/submit29_sample/wsi_vitp16_dsmil_0.5_21ep_foldall_soup.pth', 0.6),\n    ('/kaggle/input/202312-ubc/submit28/submit28/wsi_vitp16_dsmil_0.5_21ep_foldall_soup.pth',0.3),\n    ('/kaggle/input/checkpoints/final_wsi_vitp16_dsmil_1.0_15ep.pth',0.1) \n]\n\n\ndsmil2_ckpts = [\n    ('/kaggle/input/202312-ubc/submit29_sample/submit29_sample/wsi_vitp8_dsmil_0.5_21ep_foldall_soup.pth', 0.5),\n    ('/kaggle/input/202312-ubc/submit28/submit28/wsi_vitp8_dsmil_0.5_21ep_foldall_soup.pth', 0.4),\n    ('/kaggle/input/checkpoints/final_wsi_vitp8_dsmil_1.0_15ep.pth',0.1) \n]\n\ntransmil1_ckpts = [\n    ('/kaggle/input/202312-ubc/submit29_sample/submit29_sample/wsi_vitp16_transmil_0.5_21ep_foldall_soup.pth', 0.6),\n    ('/kaggle/input/202312-ubc/submit28/submit28/wsi_vitp16_transmil_0.5_21ep_foldall_soup.pth',0.3),\n    ('/kaggle/input/checkpoints/final_wsi_vitp16_transmil_1.0_15ep.pth',0.1) \n]\n\n\ntransmil2_ckpts = [\n    ('/kaggle/input/202312-ubc/submit29_sample/submit29_sample/wsi_vitp8_transmil_0.5_21ep_foldall_soup.pth', 0.5),\n    ('/kaggle/input/202312-ubc/submit28/submit28/wsi_vitp8_transmil_0.5_21ep_foldall_soup.pth', 0.4),\n    ('/kaggle/input/checkpoints/final_wsi_vitp8_transmil_1.0_15ep.pth',0.1) \n]\n","metadata":{"execution":{"iopub.status.busy":"2023-12-30T14:05:31.861616Z","iopub.execute_input":"2023-12-30T14:05:31.861899Z","iopub.status.idle":"2023-12-30T14:05:31.877372Z","shell.execute_reply.started":"2023-12-30T14:05:31.861874Z","shell.execute_reply":"2023-12-30T14:05:31.876471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.cuda.amp import autocast\n\ndata_dir = '/kaggle/input/UBC-OCEAN'\nbase_dir = '/kaggle/input/202312-ubc/base_ckpt/base_ckpt'\nckpt_dir = '/kaggle/input/202312-ubc/ckpt/ckpt'\ntest_csv = pd.read_csv(os.path.join(data_dir, 'test.csv'))\n# test_csv = pd.read_csv(os.path.join(data_dir, 'train.csv')).sample(5)\nprint(test_csv.head())\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))}\ndevice1 = torch.device('cuda:0')\ndevice2 = torch.device('cuda:1')\n\nvit_p16 = VisionTransformer(patch_size=16, embed_dim=384, num_heads=6, num_classes=0)\nvit_p16.load_state_dict(torch.load(f'{base_dir}/dino_vit_small_patch16_ep200.torch', map_location='cpu'))\nvit_p16 = vit_p16.to(device1)\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(f'{base_dir}/dino_vit_small_patch8_ep200.torch', map_location='cpu'))\nvit_p8 = vit_p8.to(device2)\nvit_p8.eval()\n\n# transmil1 = TransMIL(384, num_classes)\n# transmil1.load_state_dict(torch.load('/kaggle/input/checkpoints/wsi_vitp16_transmil_0.3_15ep.pth', map_location='cpu'))\n# transmil1 = transmil1.to(device1)\n# transmil1.eval()\n\n# transmil2 = TransMIL(384, num_classes)\n# transmil2.load_state_dict(torch.load('/kaggle/input/checkpoints/wsi_vitp8_transmil_0.3_15ep.pth', map_location='cpu'))\n# transmil2 = transmil2.to(device2)\n# transmil2.eval()\n\nsubmission = pd.DataFrame(columns=['image_id', 'label'])\n\nfrom tqdm import tqdm,trange\nwith torch.no_grad():\n    for i in trange(len(test_csv)):\n        sample = test_csv.iloc[i]\n        name = str(sample['image_id'])\n        wsidataset = SingleWSIDataset(data_dir, name, 512, 0.5, 'test')\n#         wsidataset = SingleWSIDataset(data_dir, name, 512, 0.5, 'train')\n        loader1 = DataLoader(wsidataset, 512, True, num_workers=0, pin_memory=True)\n        loader2 = DataLoader(wsidataset, 512, True, num_workers=0, pin_memory=True)\n        \n        ############################################### 特征提取部分 ###############################################\n        features1 = torch.empty((0, 384)).to(device1, non_blocking=True)\n        features2 = torch.empty((0, 384)).to(device2, non_blocking=True)    \n        for tiles1, tiles2 in zip(loader1, loader2):\n            tiles1 = tiles1.to(device1, non_blocking=True)\n            tiles2 = tiles2.to(device2, non_blocking=True)\n            with autocast():\n                feat1 = vit_p16(tiles1)\n                features1 = torch.cat([features1, feat1])\n                feat2 = vit_p8(tiles2)\n                features2 = torch.cat([features2, feat2])\n        ##########################################################################################################\n        \n        ################################################# 推理部分 ################################################\n     \n        abmil1_scores = 0\n        abmil2_scores = 0\n        dsmil1_scores = 0\n        dsmil2_scores = 0\n        transmil1_scores = 0\n        transmil2_scores = 0\n        \n        for index in range(len(abmil1_ckpts)):\n            # 定义模型\n            abmil1 = ABMIL(384, 512, 128, num_classes)\n            abmil2 = ABMIL(384, 512, 128, num_classes)\n            dsmil1 = DSMIL(IClassifier(384, num_classes), BClassifier(384, num_classes))\n            dsmil2 = DSMIL(IClassifier(384, num_classes), BClassifier(384, num_classes))\n            transmil1 = TransMIL(384, num_classes)\n            transmil2 = TransMIL(384, num_classes)\n\n\n            abmil1.load_state_dict(torch.load(f'{abmil1_ckpts[index][0]}', map_location='cpu'))\n            abmil1 = abmil1.to(device1)\n            abmil1.eval()\n\n            abmil2.load_state_dict(torch.load(f'{abmil2_ckpts[index][0]}', map_location='cpu'))\n            abmil2 = abmil2.to(device2)\n            abmil2.eval()\n\n            dsmil1.load_state_dict(torch.load(f'{dsmil1_ckpts[index][0]}', map_location='cpu'))\n            dsmil1 = dsmil1.to(device1)\n            dsmil1.eval()\n\n            dsmil2.load_state_dict(torch.load(f'{dsmil2_ckpts[index][0]}', map_location='cpu'))\n            dsmil2 = dsmil2.to(device2)\n            dsmil2.eval()\n            \n            transmil1.load_state_dict(torch.load(f'{transmil1_ckpts[index][0]}', map_location='cpu'))\n            transmil1 = transmil1.to(device1)\n            transmil1.eval()\n\n            transmil2 = TransMIL(384, num_classes)\n            transmil2.load_state_dict(torch.load(f'{transmil2_ckpts[index][0]}', map_location='cpu'))\n            transmil2 = transmil2.to(device2)\n            transmil2.eval()\n\n            \n            with autocast():\n                abmil_score1 = torch.softmax(abmil1(features1), 1)\n                abmil_score2 = torch.softmax(abmil2(features2), 1)\n\n                classes, bag_prediction, _, _ = dsmil1(features1)\n                max_prediction, idx = torch.max(classes, 0, True)\n                dsmil_score1 = 0.5 * torch.softmax(max_prediction, 1) + 0.5 * torch.softmax(bag_prediction, 1)\n                classes, bag_prediction, _, _ = dsmil2(features2)\n                max_prediction, idx = torch.max(classes, 0, True)\n                dsmil_score2 = 0.5 * torch.softmax(max_prediction, 1) + 0.5 * torch.softmax(bag_prediction, 1)\n\n                transmil_score1 = torch.softmax(transmil1(features1.unsqueeze(0)), 1).cpu()\n                transmil_score2 = torch.softmax(transmil2(features2.unsqueeze(0)), 1).cpu()\n        \n            \n            abmil1_scores += abmil_score1.cpu()* abmil1_ckpts[index][1]\n            abmil2_scores += abmil_score2.cpu()* abmil2_ckpts[index][1]\n            dsmil1_scores += dsmil_score1.cpu()* dsmil1_ckpts[index][1]\n            dsmil2_scores += dsmil_score2.cpu()* dsmil2_ckpts[index][1]\n            transmil1_scores += transmil_score1.cpu()* transmil1_ckpts[index][1]\n            transmil2_scores += transmil_score2.cpu()* transmil2_ckpts[index][1]\n\n            \n            del abmil1,abmil2,dsmil1,dsmil2,abmil_score1,abmil_score2,dsmil_score1,dsmil_score2,transmil_score1,transmil_score2\n        ##########################################################################################################\n        \n        score = (abmil1_scores+abmil2_scores+dsmil1_scores+dsmil2_scores+transmil1_scores+transmil2_scores)/6\n        pred = torch.max(score, 1)\n        pred_label = label_names[pred.indices]\n        submission.loc[len(submission)] = [name, pred_label]\n\nsubmission.to_csv('submission.csv', index=False)\npd.read_csv('submission.csv').head()","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-12-30T14:05:31.878615Z","iopub.execute_input":"2023-12-30T14:05:31.878921Z","iopub.status.idle":"2023-12-30T14:06:18.534881Z","shell.execute_reply.started":"2023-12-30T14:05:31.878884Z","shell.execute_reply":"2023-12-30T14:06:18.533915Z"},"trusted":true},"execution_count":null,"outputs":[]}]}