{"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":7337080,"sourceType":"datasetVersion","datasetId":4259584},{"sourceId":7337566,"sourceType":"datasetVersion","datasetId":4259908},{"sourceId":7353661,"sourceType":"datasetVersion","datasetId":4023955},{"sourceId":1760030,"sourceType":"datasetVersion","datasetId":1046169},{"sourceId":6774553,"sourceType":"datasetVersion","datasetId":3898019}],"dockerImageVersionId":30627,"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\n!pip install /kaggle/input/einops-030/einops-0.3.0-py2.py3-none-any.whl","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-01-07T05:39:51.752116Z","iopub.execute_input":"2024-01-07T05:39:51.752758Z","iopub.status.idle":"2024-01-07T05:41:55.730152Z","shell.execute_reply.started":"2024-01-07T05:39:51.752729Z","shell.execute_reply":"2024-01-07T05:41:55.728854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\n\nfrom sklearn.metrics import balanced_accuracy_score\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim as optim\nimport torch.optim.lr_scheduler as lr_scheduler","metadata":{"execution":{"iopub.status.busy":"2024-01-07T05:41:55.73229Z","iopub.execute_input":"2024-01-07T05:41:55.732676Z","iopub.status.idle":"2024-01-07T05:41:59.580594Z","shell.execute_reply.started":"2024-01-07T05:41:55.732643Z","shell.execute_reply":"2024-01-07T05:41:59.579638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class WSIFeatDataset(Dataset):\n    def __init__(self, data_csv: pd.DataFrame, feature_dir: str, ratio, phase: int):\n        super().__init__()\n        self.data_csv = data_csv\n        self.feature_dir = feature_dir\n        self.ratio = ratio\n        assert phase in [0, 1]\n        self.phase = phase\n\n    def __len__(self):\n        return len(self.data_csv)\n\n    def __getitem__(self, idx):\n        sample = self.data_csv.iloc[idx]\n        file_name = str(sample['image_id'])\n\n        features = torch.load(os.path.join(self.feature_dir, file_name + '.pt'), map_location='cpu')\n        random.shuffle(features)\n        if 0 < self.ratio <= 1:\n            features = features[:int(len(features) * self.ratio)]\n        elif self.ratio > 1:\n            features = features[:min(len(features), self.ratio)]\n\n        if self.phase == 0:\n            label = torch.tensor(label_dict[sample['label']])\n            return file_name, features, label\n        else:\n            return file_name, features","metadata":{"execution":{"iopub.status.busy":"2024-01-07T05:41:59.581798Z","iopub.execute_input":"2024-01-07T05:41:59.58223Z","iopub.status.idle":"2024-01-07T05:41:59.592043Z","shell.execute_reply.started":"2024-01-07T05:41:59.582201Z","shell.execute_reply":"2024-01-07T05:41:59.59106Z"},"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\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(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":"2024-01-07T05:41:59.59466Z","iopub.execute_input":"2024-01-07T05:41:59.595185Z","iopub.status.idle":"2024-01-07T05:41:59.656375Z","shell.execute_reply.started":"2024-01-07T05:41:59.595151Z","shell.execute_reply":"2024-01-07T05:41:59.655273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = '/kaggle/input/UBC-OCEAN'\ntrain_csv = pd.read_csv(os.path.join('/kaggle/input/checkpoints', 'beifen.csv'))\ntrain_data = train_csv.iloc[np.r_[0:100, 200:536]].reset_index(drop=True)\nval_data = train_csv.iloc[100:200].reset_index(drop=True)\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))}\nepochs = 15\nin_dim = 384\nmodel_type = 16\nratio = 1.0\nmil_type = 'abmil'\naccumulate = True\ntest = False\nfeature_dir = f'/kaggle/input/ubc-ocean-vit-p{model_type}-wsi-features'\nmil_model_name = f'final_wsi_vitp{model_type}_{mil_type}_{ratio}_{epochs}ep.pth'\ndevice = torch.device('cuda:0')\n\ntrain_dataset = WSIFeatDataset(train_data, feature_dir, ratio, 0) if test else WSIFeatDataset(train_csv, feature_dir, ratio, 0)\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True, num_workers=4, pin_memory=True)\nval_dataset = WSIFeatDataset(val_data, feature_dir, ratio, 0)\nval_loader = DataLoader(val_dataset, batch_size=1, shuffle=True, num_workers=4, pin_memory=True)\n\nif mil_type == 'abmil':\n    model = ABMIL(in_dim, 512, 128, num_classes)\nelif mil_type == 'dsmil':\n    model = DSMIL(IClassifier(in_dim, num_classes), BClassifier(in_dim, num_classes))\nelif mil_type == 'transmil':\n    model = TransMIL(in_dim, num_classes)\n\noptimizer = optim.Adam(model.parameters(), 5e-4, weight_decay=5e-4)\nscheduler = lr_scheduler.CosineAnnealingLR(optimizer, epochs, 5e-5)\n\nmodel = model.to(device)\n\nmax_acc = 0.\nbalanced_acc = 0.\nfor epoch in range(1, epochs + 1):\n    loss_sum = 0.\n    n = 0\n\n    loop = tqdm(train_loader, total=len(train_loader))\n    model.train()\n    for file_name, features, label in loop:\n        label = label.to(device)\n        features = features.squeeze(0).to(device)\n\n        if mil_type == 'abmil':\n            scores = model(features)\n            loss = F.cross_entropy(scores, label)\n        elif mil_type == 'dsmil':\n            classes, bag_prediction, _, _ = model(features)\n            max_prediction, index = torch.max(classes, 0, True)\n            loss_bag = F.cross_entropy(bag_prediction, label)\n            loss_max = F.cross_entropy(max_prediction.view(1, -1), label)\n            loss = 0.5 * loss_bag + 0.5 * loss_max\n        elif mil_type == 'transmil':\n            scores = model(features.unsqueeze(0))\n            loss = F.cross_entropy(scores, label)\n\n        if accumulate:\n            loss = loss / 4\n            loss.backward(retain_graph=True)\n            if (n + 1) % 4 == 0 or (n + 1) == len(train_loader):\n                optimizer.step()\n                optimizer.zero_grad()\n        else:\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n\n        n += 1\n        loss_sum += loss.item()\n\n        loop.set_description(f'Train [{epoch}/{epochs}]')\n        loop.set_postfix(loss=loss.item(), loss_mean=loss_sum / n)\n\n    if test:\n        with torch.no_grad():\n            acc = 0\n            loss_val = 0\n            y_true = []\n            y_pred = []\n\n            loop = tqdm(val_loader, total=len(val_loader))\n            model.eval()\n            for file_name, features, label in loop:\n                label = label.to(device)\n                features = features.squeeze(0).to(device)\n\n                if mil_type == 'abmil':\n                    scores = model(features)\n                    scores = torch.softmax(scores, 1)\n                elif mil_type == 'dsmil':\n                    classes, bag_prediction, _, _ = model(features)\n                    max_prediction, index = torch.max(classes, 0, True)\n                    scores = 0.5 * torch.softmax(max_prediction, 1) + 0.5 * torch.softmax(bag_prediction, 1)\n                elif mil_type == 'transmil':\n                    scores = model(features.unsqueeze(0))\n                    scores = torch.softmax(scores, 1)\n\n                pred = torch.argmax(scores)\n\n                y_pred.append(pred.item())\n                y_true.append(label.item())\n\n                if pred == label.squeeze(0):\n                    acc += 1\n\n                loop.set_description(f'Val [{epoch}/{epochs}]')\n                loop.set_postfix(acc=acc / len(val_loader), max_acc=max_acc, balanced_acc=balanced_acc)\n\n            if acc / len(val_loader) > max_acc:\n                max_acc = acc / len(val_loader)\n            y_true = np.array(y_true)\n            y_pred = np.array(y_pred)\n            if balanced_accuracy_score(y_true, y_pred) > balanced_acc:\n                balanced_acc = balanced_accuracy_score(y_true, y_pred)\n\n    scheduler.step()\n\n    if not test:\n        torch.save(model.state_dict(), mil_model_name)","metadata":{"execution":{"iopub.status.busy":"2024-01-07T05:50:30.351836Z","iopub.execute_input":"2024-01-07T05:50:30.352252Z","iopub.status.idle":"2024-01-07T05:58:33.501255Z","shell.execute_reply.started":"2024-01-07T05:50:30.352216Z","shell.execute_reply":"2024-01-07T05:58:33.499972Z"},"trusted":true},"execution_count":null,"outputs":[]}]}