{"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":1760030,"sourceType":"datasetVersion","datasetId":1046169},{"sourceId":6774553,"sourceType":"datasetVersion","datasetId":3898019},{"sourceId":7246771,"sourceType":"datasetVersion","datasetId":3770816},{"sourceId":7314103,"sourceType":"datasetVersion","datasetId":4023955},{"sourceId":7337080,"sourceType":"datasetVersion","datasetId":4259584},{"sourceId":7337566,"sourceType":"datasetVersion","datasetId":4259908},{"sourceId":7340806,"sourceType":"datasetVersion","datasetId":4162804}],"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":"2024-01-06T05:16:28.409717Z","iopub.execute_input":"2024-01-06T05:16:28.410003Z","iopub.status.idle":"2024-01-06T05:17:56.251042Z","shell.execute_reply.started":"2024-01-06T05:16:28.409977Z","shell.execute_reply":"2024-01-06T05:17:56.24991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ","metadata":{}},{"cell_type":"markdown","source":"## Import","metadata":{}},{"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\n\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))}\n# num_classes = 2\n# labels = [0, 1]\n# label_dict = {labels[i]: i for i in range(len(labels))}\nin_dim = 384\nseed = 42\napply_focal = False\n\ndef ensure_dir(path):\n    if not os.path.exists(path):\n        os.makedirs(path)\n\n\ndef 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(seed)\n","metadata":{"execution":{"iopub.status.busy":"2024-01-06T05:17:56.252963Z","iopub.execute_input":"2024-01-06T05:17:56.2533Z","iopub.status.idle":"2024-01-06T05:18:01.447975Z","shell.execute_reply.started":"2024-01-06T05:17:56.253271Z","shell.execute_reply":"2024-01-06T05:18:01.446952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss","metadata":{}},{"cell_type":"code","source":"\nclass LabelSmoothingCrossEntropy(nn.Module):\n    def __init__(self, smoothing=0.1):\n        super(LabelSmoothingCrossEntropy, self).__init__()\n        self.smoothing = smoothing\n    \n    def forward(self, input, target):\n        log_probs = F.log_softmax(input, dim=-1)\n\n        # 创建一个与真实标签相同大小的张量，其中包含每个类别的平滑标签\n        confidence = 1.0 - self.smoothing\n        with torch.no_grad():\n            true_dist = torch.zeros_like(log_probs)\n            true_dist.fill_(self.smoothing / (input.size(1) - 1))\n            true_dist.scatter_(1, target.data.unsqueeze(1), confidence)\n\n        # 计算交叉熵\n        loss = (-true_dist * log_probs).sum(dim=-1).mean()\n        return loss\n\nlabel_smooth = LabelSmoothingCrossEntropy(0.1)\n\nclass FocalLoss(nn.Module):\n    '''Multi-class Focal loss implementation'''\n\n    def __init__(self, gamma=2, weight=None, ignore_index=-100):\n        super(FocalLoss, self).__init__()\n        self.gamma = gamma\n        self.weight = weight\n        self.ignore_index = ignore_index\n\n    def forward(self, input, target):\n        \"\"\"\n        :param input: torch.Tensor, shape=[N, C]\n        :param target: torch.Tensor, shape=[N, ]\n        \"\"\"\n        logpt = F.log_softmax(input, dim=1)\n        pt = torch.exp(logpt)\n        logpt = (1 - pt) ** self.gamma * logpt\n        loss = F.nll_loss(logpt, target, self.weight, ignore_index=self.ignore_index)\n        return loss\n\nfocal_loss = FocalLoss()","metadata":{"execution":{"iopub.status.busy":"2024-01-06T05:18:01.449367Z","iopub.execute_input":"2024-01-06T05:18:01.449763Z","iopub.status.idle":"2024-01-06T05:18:01.460718Z","shell.execute_reply.started":"2024-01-06T05:18:01.449737Z","shell.execute_reply":"2024-01-06T05:18:01.459604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"# 1 Dataset and transform\nclass 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 isinstance(self.ratio, float):\n            features = features[:int(len(features) * self.ratio)]\n        elif isinstance(self.ratio, int):\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\n","metadata":{"execution":{"iopub.status.busy":"2024-01-06T05:18:01.463117Z","iopub.execute_input":"2024-01-06T05:18:01.463464Z","iopub.status.idle":"2024-01-06T05:18:01.48213Z","shell.execute_reply.started":"2024-01-06T05:18:01.463432Z","shell.execute_reply":"2024-01-06T05:18:01.48139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Models","metadata":{}},{"cell_type":"code","source":"\n# 2 Model\nclass 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 MILNet(nn.Module):\n    def __init__(self, i_classifier, b_classifier):\n        super(MILNet, 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\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\n","metadata":{"execution":{"iopub.status.busy":"2024-01-06T05:18:01.483435Z","iopub.execute_input":"2024-01-06T05:18:01.483695Z","iopub.status.idle":"2024-01-06T05:18:33.517955Z","shell.execute_reply.started":"2024-01-06T05:18:01.483672Z","shell.execute_reply":"2024-01-06T05:18:33.516965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"# 0 Raw data loading\ntrain_csv = pd.read_csv(os.path.join('/kaggle/input/202312-ubc/', 'train_revised_extra.csv'))\nEC = train_csv[train_csv['label'].isin(['EC','CC'])]\nother = train_csv[train_csv['label'].isin(['LGSC','MC','Other'])]\n# Do upsample\ntrain_csv = pd.concat([train_csv] + [EC] * 3 + [other] * 5).sample(frac=1.0)\nprint(train_csv['label'].value_counts())\nfrom sklearn.model_selection import StratifiedKFold\ntrain_csv['tag'] = train_csv['label'] + '_' + train_csv['is_tma'].astype(str)\n\n# skf = StratifiedKFold(n_splits=5,shuffle=True, random_state=seed)\n# for i, (_, val_index) in enumerate(skf.split(train_csv, train_csv['tag'])):\n#     train_csv.loc[val_index, \"fold\"] = i\n\n\n\ndef train_fold(fold=0, model_type=16, mil_type='abmil', ratio=0.3, epochs = 21, test = True):\n    TASK = f'UBC_fold_{fold}, model_type_{model_type}, mil_type_{mil_type},ratio_{ratio}'\n    print(f\"====={TASK}=======\")\n    if test:\n        train_data = train_csv[train_csv[\"fold\"] != fold]\n        val_data = train_csv[train_csv[\"fold\"] == fold]\n    else:\n        fold='all'\n        train_data = train_csv\n        val_data = train_csv\n    print(train_data['label'].value_counts())\n    print(val_data['label'].value_counts())\n    accumulate = True\n    feature_dir = f'/kaggle/input/ubc-ocean-vit-p{model_type}-wsi-features'\n    mil_model_name = f'wsi_vitp{model_type}_{mil_type}_{ratio}_{epochs}ep'\n    device = torch.device('cuda:0')\n    output_dir = f'./ckpt/{mil_type}_patch{model_type}/'\n    ensure_dir(output_dir)\n    # 3 Training\n    train_dataset = WSIFeatDataset(train_data, feature_dir, ratio, 0) if test else WSIFeatDataset(train_csv, feature_dir, ratio, 0)\n    train_loader = DataLoader(train_dataset, batch_size=1, shuffle=True, num_workers=4, pin_memory=True)\n    val_dataset = WSIFeatDataset(val_data, feature_dir, ratio, 0)\n    val_loader = DataLoader(val_dataset, batch_size=1, shuffle=False, num_workers=4, pin_memory=True)\n\n    if mil_type == 'abmil':\n        model = ABMIL(in_dim, 512, 128, num_classes)\n    elif mil_type == 'transmil':\n        model = TransMIL(384, num_classes)\n    else:\n        model = MILNet(IClassifier(in_dim, num_classes), BClassifier(in_dim, num_classes))\n\n    optimizer = optim.Adam(model.parameters(), 5e-4, weight_decay=5e-4)\n    scheduler = lr_scheduler.CosineAnnealingLR(optimizer, epochs, 5e-5)\n\n    model = model.to(device)\n\n    max_acc = 0.6\n    balanced_acc = 0.\n    for epoch in range(1, epochs+1):\n        loss_sum = 0.\n        n = 0\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 in ['abmil']:\n                scores = model(features)\n                if apply_focal:\n                    loss = focal_loss(scores, label)\n                else:\n                    loss = F.cross_entropy(scores, label)\n                    # loss = label_smooth(scores, label)\n            elif mil_type == 'transmil':\n                scores = model(features.unsqueeze(0))\n                if apply_focal:\n                    loss = focal_loss(scores, label)\n                else:\n                    loss = F.cross_entropy(scores, label)\n                    # loss = label_smooth(scores, label)\n            else:\n                classes, bag_prediction, _, _ = model(features)\n                max_prediction, index = torch.max(classes, 0, True)\n                if apply_focal:\n                    loss_bag = focal_loss(bag_prediction, label)\n                    loss_max = focal_loss(max_prediction.view(1, -1), label)\n                    loss = 0.5 * loss_bag + 0.5 * loss_max\n                else:\n                    loss_bag = F.cross_entropy(bag_prediction, label)\n                    loss_max = F.cross_entropy(max_prediction.view(1, -1), label)\n                    # loss_bag = label_smooth(bag_prediction, label)\n                    # loss_max = label_smooth(max_prediction.view(1, -1), label)\n                    loss = 0.5 * loss_bag + 0.5 * loss_max\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 in ['abmil']:\n                        scores = model(features)\n                    elif mil_type == 'transmil':\n                        scores = model(features.unsqueeze(0))\n                    else:\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\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                    torch.save(model.state_dict(), f\"{output_dir}/{mil_model_name}_fold{fold}_{epoch}_val{round(balanced_acc,3)}.pth\")\n\n        scheduler.step()\n\n        if not test:\n            if epoch > 8 and epoch%3 == 0:\n                torch.save(model.state_dict(), f\"{output_dir}/{mil_model_name}_fold{fold}_{epoch}.pth\")\n\n# # 划分验证集，五折\n# for model_type in [8, 16]:\n#     for mil_type in ['abmil', 'dsmil']:\n#         for label_type in label_names:\n#             for fold in range(0, 3):\n#                 train_fold(fold=fold, model_type=model_type, mil_type=mil_type, label_type=label_type)\n\n#不划分验证集，全量训练\nfor model_type in [8, 16]:\n    for mil_type in ['abmil', 'dsmil', 'transmil']:\n        train_fold(fold='all', model_type=model_type, mil_type=mil_type, ratio=0.5, test = False)","metadata":{"execution":{"iopub.status.busy":"2024-01-06T05:18:33.519404Z","iopub.execute_input":"2024-01-06T05:18:33.519671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Soup","metadata":{}},{"cell_type":"code","source":"\nimport torch\nfrom tqdm import tqdm\n\nclass ModelUtils:\n    @staticmethod\n    def model_soup(ckpts, save_name):\n        torch.device('cpu')\n        print('--------model_average---------')\n        avg = None\n        num = 0\n        for path in tqdm(ckpts):\n            num += 1\n            states = torch.load(path, map_location=torch.device('cpu'))\n            if 'model_state_dict' in states.keys():\n                states = states['model_state_dict']\n            if avg is None:\n                avg = states\n            else:\n                for k in avg.keys():\n                    if 'position' not in k:\n                        avg[k] += states[k]\n                    else:\n                        print(k)\n        # average\n        for k in avg.keys():\n            if 'position' not in k:\n                avg[k] /= num\n            else:\n                print(k)\n        torch.save(avg, save_name)\n\n\n\n\nckpts1 = [f'ckpt/abmil_patch16/wsi_vitp16_abmil_0.5_21ep_foldall_{i}.pth' for i in [12, 15, 18, 21]]\nsave_name1 = 'wsi_vitp16_abmil_0.5_21ep_foldall_soup.pth'\nModelUtils.model_soup(ckpts1, save_name1)\n\nckpts2 = [f'ckpt/abmil_patch8/wsi_vitp8_abmil_0.5_21ep_foldall_{i}.pth' for i in [12, 15, 18, 21]]\nsave_name2 = 'wsi_vitp8_abmil_0.5_21ep_foldall_soup.pth'\nModelUtils.model_soup(ckpts2, save_name2)\n\nckpts3 = [f'ckpt/dsmil_patch16/wsi_vitp16_dsmil_0.5_21ep_foldall_{i}.pth' for i in [12, 15, 18, 21]]\nsave_name3 = 'wsi_vitp16_dsmil_0.5_21ep_foldall_soup.pth'\nModelUtils.model_soup(ckpts3, save_name3)\n\nckpts4 = [f'ckpt/dsmil_patch8/wsi_vitp8_dsmil_0.5_21ep_foldall_{i}.pth' for i in [12, 15, 18, 21]]\nsave_name4 = 'wsi_vitp8_dsmil_0.5_21ep_foldall_soup.pth'\nModelUtils.model_soup(ckpts4, save_name4)\n\nckpts5 = [f'ckpt/transmil_patch16/wsi_vitp16_transmil_0.5_21ep_foldall_{i}.pth' for i in [12, 15, 18, 21]]\nsave_name5 = 'wsi_vitp16_transmil_0.5_21ep_foldall_soup.pth'\nModelUtils.model_soup(ckpts5, save_name5)\n\nckpts6 = [f'ckpt/transmil_patch8/wsi_vitp8_transmil_0.5_21ep_foldall_{i}.pth' for i in [12, 15, 18, 21]]\nsave_name6 = 'wsi_vitp8_transmil_0.5_21ep_foldall_soup.pth'\nModelUtils.model_soup(ckpts6, save_name6)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}