{"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":34547,"databundleVersionId":3897958,"sourceType":"competition"},{"sourceId":3848167,"sourceType":"datasetVersion","datasetId":2289605}],"dockerImageVersionId":30559,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n        break\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-03-09T03:14:39.253032Z","iopub.execute_input":"2024-03-09T03:14:39.253289Z","iopub.status.idle":"2024-03-09T03:14:47.460808Z","shell.execute_reply.started":"2024-03-09T03:14:39.253265Z","shell.execute_reply":"2024-03-09T03:14:47.459910Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q segmentation_models_pytorch","metadata":{"execution":{"iopub.status.busy":"2024-03-09T03:14:51.438932Z","iopub.execute_input":"2024-03-09T03:14:51.439412Z","iopub.status.idle":"2024-03-09T03:15:10.020397Z","shell.execute_reply.started":"2024-03-09T03:14:51.439380Z","shell.execute_reply":"2024-03-09T03:15:10.019291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import argparse\nimport copy\nimport gc\nimport glob\nimport os\nimport random\nimport sys\nimport time\nfrom collections import defaultdict\nimport numpy as np\nimport pandas as pd\nimport torch\nimport cv2\nfrom skimage import io\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom fastprogress import progress_bar\nfrom matplotlib import pyplot as plt\nfrom sklearn.model_selection import KFold\nfrom torch import optim\nfrom fastai.learner import Metric\nfrom fastai.torch_core import flatten_check\nfrom itertools import chain\nfrom timm.models.layers import to_2tuple, DropPath, trunc_normal_\nimport torch.utils.checkpoint as checkpoint\nimport torch.nn.functional as F\nfrom torchvision import models\nfrom torch.cuda import amp\nfrom torch.optim import lr_scheduler\nfrom tqdm import tqdm\nfrom segmentation_models_pytorch.base import modules as md\nfrom timm.models.layers import *\n\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2024-03-09T03:15:16.843155Z","iopub.execute_input":"2024-03-09T03:15:16.844053Z","iopub.status.idle":"2024-03-09T03:15:24.234878Z","shell.execute_reply.started":"2024-03-09T03:15:16.844017Z","shell.execute_reply":"2024-03-09T03:15:24.234083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    def __init__(self, fold=0., seed=2024, train_bs=32, debug=False, pretrained_path=None):\n        self.n_fold = 5\n        self.fold = fold\n        self.seed = seed\n        self.train_bs = train_bs\n        self.valid_bs = train_bs  # may need to change this\n        self.debug = debug\n        self.img_size = [256, 256]\n        self.exp_name = 'Hubmap256-training'\n        self.epochs = 1000\n        self.lr = 2e-3\n        self.pretrained_path = pretrained_path\n        self.scheduler = \"CosineAnnealingLR\"\n        self.min_lr = 2e-4\n        self.T_max = int(30000 / self.train_bs * self.epochs) + 50\n        self.T_0 = 25\n        self.warmup_epochs = 10\n        self.wd = 1e-6\n        self.n_accumulate = max(1, 32 // self.train_bs)\n        self.num_classes = 1\n        self.device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n        self.backbone = \"swinunet\"\n        self.model_name = \"swinunet\"\n        self.model_urls = {\n            \"swinv2_tiny_window16_256\": \"./swinv2_tiny_patch4_window8_256.pth\",\n            \"swinv2_small_window8_256\": \"./swinv2_small_patch4_window8_256.pth\",\n            \"swinv2_small_window16_256\": \"./swinv2_small_patch4_window16_256.pth\",\n            \"swinv2_base_window16_256\": \"./swinv2_base_patch4_window16_256.pth\",\n        }\n\n        self.size = \"swinv2_base_window16_256\"\n\n        self.load_best_model = False\n\n        self.train_dataset = \"hap\"  # only all, hap, hubmap\n\n        self.dice_dataset = \"hap\"  # only all, hap, hubmap\n\n        self.only_dice = 0\n\n    def display(self):\n        print(f\"{self.exp_name}\")\n        print(f\"debug is {self.debug}\")\n        print(f\"seed is {self.seed}\")\n        print(f\"train_bs is {self.train_bs}\")\n        print(f\"img_size is {self.img_size}\")\n        print(f\"fold_no is {self.fold}\")\n        print(f\"backbone is {self.backbone}\")\n        print(f\"epochs is {self.epochs}\")\n        print(f\"Pretrained: {self.pretrained_path}\")","metadata":{"execution":{"iopub.status.busy":"2024-03-09T03:16:35.719486Z","iopub.execute_input":"2024-03-09T03:16:35.720416Z","iopub.status.idle":"2024-03-09T03:16:35.731604Z","shell.execute_reply.started":"2024-03-09T03:16:35.720381Z","shell.execute_reply":"2024-03-09T03:16:35.730521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TrainDataset(Dataset):\n    def __init__(self, graph_list=None, cfg=None, transforms=None, mode='train'):\n        self.graph_list = graph_list\n        # remove that one faulty image from train_csv\n        self.mode = mode\n        self.cfg = cfg\n        self.transforms = transforms\n        if cfg.train_dataset == \"hap\":\n            prefix = \"../input/hubmap-2022-256x256/\"\n        elif cfg.train_dataset == \"hubmap\":\n            prefix = \"../hubmap-256x256/\"\n        else:\n            prefix = \"../all_256/\"\n        self.image_paths = [prefix + \"train/\" + i for i in graph_list]\n        self.mask_paths = [prefix + \"masks/\" + i.replace(\"train\", \"mask\") for i in graph_list]\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n\n        img = io.imread(self.image_paths[idx])\n        mask = cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE)\n        if self.transforms:\n            data = self.transforms(image=img, mask=mask)\n            img = data['image']\n            mask = data['mask']\n        img = np.transpose(img, (2, 0, 1)) / 255.0\n        return torch.tensor(img), torch.tensor(mask)\n\n\nclass DiceDataset(Dataset):\n    def __init__(self, graph_list=None, cfg=None):\n        self.graph_list = graph_list\n        # remove that one faulty image from train_csv\n        self.cfg = cfg\n        if cfg.train_dataset == \"hap\":\n            prefix = \"../input/hubmap-2022-256x256/\"\n        elif cfg.train_dataset == \"hubmap\":\n            prefix = \"../hubmap-256x256/\"\n        else:\n            prefix = \"../all_256/\"\n        self.image_paths = [prefix + \"train/\" + i for i in graph_list]\n        self.mask_paths = [prefix + \"masks/\" + i.replace(\"train\", \"mask\") for i in graph_list]\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n\n        img = io.imread(self.image_paths[idx])\n        mask = cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE)\n        img = np.transpose(img, (2, 0, 1)) / 255.0\n        return torch.tensor(img), torch.tensor(mask)\n\n\ndef get_transforms(train=True, cfg=None):\n    data_transforms = {\n        \"train\": A.Compose([\n            #       A.Resize(*cfg.img_size, interpolation=cv2.INTER_NEAREST),\n            A.HorizontalFlip(p=0.5),\n            #         A.VerticalFlip(p=0.5),\n            A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.05, rotate_limit=10, p=0.5),\n            A.OneOf([\n                A.GridDistortion(num_steps=5, distort_limit=0.05, p=1.0),\n                # #             A.OpticalDistortion(distort_limit=0.05, shift_limit=0.05, p=1.0),\n                A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=1.0)\n            ], p=0.25),\n            A.CoarseDropout(max_holes=8, max_height=cfg.img_size[0] // 20, max_width=cfg.img_size[1] // 20,\n                            min_holes=5, fill_value=0, mask_fill_value=0, p=0.5),\n        ], p=1.0),\n\n        \"valid\": A.Compose([\n            #             A.Resize(*cfg.img_size, interpolation=cv2.INTER_NEAREST),\n        ], p=1.0)\n    }\n    if train == True:\n        return data_transforms[\"train\"]\n    else:\n        return data_transforms['valid']\n\n\ndef prepare_train_loaders(fold, df, cfg, debug=False):\n    train_list = df[df['fold'] == fold].reset_index(drop=True)[\"graph_name\"].values\n    valid_list = df[df['fold'] != fold].reset_index(drop=True)[\"graph_name\"].values\n\n    if debug:\n        train_list = train_list[:20]\n        valid_list = valid_list[:20]\n\n    train_dataset = TrainDataset(train_list, transforms=get_transforms(train=True, cfg=cfg), cfg=cfg, mode='train')\n    valid_dataset = TrainDataset(valid_list, transforms=get_transforms(train=False, cfg=cfg), cfg=cfg, mode='valid')\n\n    #     print(get_statistics(train_dataset))\n\n    train_loader = DataLoader(train_dataset, batch_size=cfg.train_bs if not cfg.debug else 20,\n                              num_workers=0, shuffle=True, pin_memory=True, drop_last=False)\n    valid_loader = DataLoader(valid_dataset, batch_size=cfg.valid_bs if not cfg.debug else 20,\n                              num_workers=0, shuffle=True, pin_memory=True)\n\n    return train_loader, valid_loader\n\ndef prepare_valid_loaders(cfg):\n    if cfg.train_dataset == \"hap\":\n        prefix = \"../input/hubmap-2022-256x256/\"\n    elif cfg.train_dataset == \"hubmap\":\n        prefix = \"../hubmap-256x256/\"\n    else:\n        prefix = \"../all_256/\"\n    dice_graph_path_list = glob.glob(prefix + \"train/*\")\n    dice_graph_name_list = [i[i.rindex(\"/\") + 1:] for i in dice_graph_path_list]\n    dice_dataset = DiceDataset(dice_graph_name_list, cfg=cfg)\n    dice_loader = DataLoader(dice_dataset,  num_workers=0, shuffle=True, batch_size=1, pin_memory=True)\n    return dice_loader","metadata":{"execution":{"iopub.status.busy":"2024-03-09T03:16:36.690940Z","iopub.execute_input":"2024-03-09T03:16:36.691295Z","iopub.status.idle":"2024-03-09T03:16:36.716094Z","shell.execute_reply.started":"2024-03-09T03:16:36.691265Z","shell.execute_reply":"2024-03-09T03:16:36.715162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DiceScore(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(DiceScore, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        # comment out if your model contains a sigmoid or equivalent activation layer\n        inputs = F.sigmoid(inputs)\n\n        # flatten label and prediction tensors\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n\n        intersection = (inputs * targets).sum()\n        dice = (2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth)\n\n        return dice\n\n\nclass DiceBCELoss(nn.Module):\n    # Formula Given above.\n    def __init__(self, weight=None, size_average=True):\n        super(DiceBCELoss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        # comment out if your model contains a sigmoid or equivalent activation layer\n        #         inputs = nnF.sigmoid(inputs)\n\n        # flatten label and prediction tensors\n\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n\n        BCE = F.binary_cross_entropy_with_logits(inputs, targets, reduction='mean')\n\n        inputs = F.sigmoid(inputs)\n        intersection = (inputs * targets).sum()\n        dice_loss = 1 - (2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth)\n\n        Dice_BCE = BCE + dice_loss\n\n        return Dice_BCE\n\n\nclass Dice_th_pred(Metric):\n    def __init__(self, ths=np.arange(0.1, 0.9, 0.01), axis=1):\n        self.axis = axis\n        self.ths = ths\n        self.reset()\n\n    def reset(self):\n        self.inter = torch.zeros(len(self.ths))\n        self.union = torch.zeros(len(self.ths))\n\n    def accumulate(self, p, t):\n        pred, targ = flatten_check(p, t)\n        for i, th in enumerate(self.ths):\n            p = (pred > th).float()\n            self.inter[i] += (p * targ).float().sum().item()\n            self.union[i] += (p + targ).float().sum().item()\n\n    @property\n    def value(self):\n        dices = torch.where(self.union > 0.0, 2.0 * self.inter / self.union,\n                            torch.zeros_like(self.union))\n        return dices","metadata":{"execution":{"iopub.status.busy":"2024-03-09T03:16:37.253963Z","iopub.execute_input":"2024-03-09T03:16:37.254585Z","iopub.status.idle":"2024-03-09T03:16:37.268076Z","shell.execute_reply.started":"2024-03-09T03:16:37.254556Z","shell.execute_reply":"2024-03-09T03:16:37.267091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Mlp(nn.Module):\n    def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):\n        super().__init__()\n        out_features = out_features or in_features\n        hidden_features = hidden_features or in_features\n        self.fc1 = nn.Linear(in_features, hidden_features)\n        self.act = act_layer()\n        self.fc2 = nn.Linear(hidden_features, out_features)\n        self.drop = nn.Dropout(drop)\n\n    def forward(self, x):\n        x = self.fc1(x)\n        x = self.act(x)\n        x = self.drop(x)\n        x = self.fc2(x)\n        x = self.drop(x)\n        return x\n\n\ndef window_partition(x, window_size):\n    \"\"\"\n    Args:\n        x: (B, H, W, C)\n        window_size (int): window size\n    Returns:\n        windows: (num_windows*B, window_size, window_size, C)\n    \"\"\"\n\n    B, H, W, C = x.shape\n    x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)\n    windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)\n    return windows\n\n\ndef window_reverse(windows, window_size, H, W):\n    \"\"\"\n    Args:\n        windows: (num_windows*B, window_size, window_size, C)\n        window_size (int): Window size\n        H (int): Height of image\n        W (int): Width of image\n    Returns:\n        x: (B, H, W, C)\n    \"\"\"\n    B = int(windows.shape[0] / (H * W / window_size / window_size))\n    x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)\n    x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)\n    return x\n\n\nclass WindowAttention(nn.Module):\n    r\"\"\" Window based multi-head self attention (W-MSA) module with relative position bias.\n    It supports both of shifted and non-shifted window.\n    Args:\n        dim (int): Number of input channels.\n        window_size (tuple[int]): The height and width of the window.\n        num_heads (int): Number of attention heads.\n        qkv_bias (bool, optional):  If True, add a learnable bias to query, key, value. Default: True\n        attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0\n        proj_drop (float, optional): Dropout ratio of output. Default: 0.0\n        pretrained_window_size (tuple[int]): The height and width of the window in pre-training.\n    \"\"\"\n\n    def __init__(self, dim, window_size, num_heads, qkv_bias=True, attn_drop=0., proj_drop=0.,\n                 pretrained_window_size=[0, 0]):\n\n        super().__init__()\n        self.dim = dim\n        self.window_size = window_size  # Wh, Ww\n        self.pretrained_window_size = pretrained_window_size\n        self.num_heads = num_heads\n\n        self.logit_scale = nn.Parameter(torch.log(10 * torch.ones((num_heads, 1, 1))), requires_grad=True).to(cfg.device)\n\n        # mlp to generate continuous relative position bias\n        self.cpb_mlp = nn.Sequential(nn.Linear(2, 512, bias=True),\n                                     nn.ReLU(inplace=True),\n                                     nn.Linear(512, num_heads, bias=False))\n\n        # get relative_coords_table\n        relative_coords_h = torch.arange(-(self.window_size[0] - 1), self.window_size[0], dtype=torch.float32)\n        relative_coords_w = torch.arange(-(self.window_size[1] - 1), self.window_size[1], dtype=torch.float32)\n        relative_coords_table = torch.stack(\n            torch.meshgrid([relative_coords_h,\n                            relative_coords_w])).permute(1, 2, 0).contiguous().unsqueeze(0)  # 1, 2*Wh-1, 2*Ww-1, 2\n        if pretrained_window_size[0] > 0:\n            relative_coords_table[:, :, :, 0] /= (pretrained_window_size[0] - 1)\n            relative_coords_table[:, :, :, 1] /= (pretrained_window_size[1] - 1)\n        else:\n            relative_coords_table[:, :, :, 0] /= (self.window_size[0] - 1)\n            relative_coords_table[:, :, :, 1] /= (self.window_size[1] - 1)\n        relative_coords_table *= 8  # normalize to -8, 8\n        relative_coords_table = torch.sign(relative_coords_table) * torch.log2(\n            torch.abs(relative_coords_table) + 1.0) / np.log2(8)\n\n        self.register_buffer(\"relative_coords_table\", relative_coords_table)\n\n        # get pair-wise relative position index for each token inside the window\n        coords_h = torch.arange(self.window_size[0])\n        coords_w = torch.arange(self.window_size[1])\n        coords = torch.stack(torch.meshgrid([coords_h, coords_w]))  # 2, Wh, Ww\n        coords_flatten = torch.flatten(coords, 1)  # 2, Wh*Ww\n        relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :]  # 2, Wh*Ww, Wh*Ww\n        relative_coords = relative_coords.permute(1, 2, 0).contiguous()  # Wh*Ww, Wh*Ww, 2\n        relative_coords[:, :, 0] += self.window_size[0] - 1  # shift to start from 0\n        relative_coords[:, :, 1] += self.window_size[1] - 1\n        relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1\n        relative_position_index = relative_coords.sum(-1)  # Wh*Ww, Wh*Ww\n        self.register_buffer(\"relative_position_index\", relative_position_index)\n\n        self.qkv = nn.Linear(dim, dim * 3, bias=False)\n        if qkv_bias:\n            self.q_bias = nn.Parameter(torch.zeros(dim))\n            self.v_bias = nn.Parameter(torch.zeros(dim))\n        else:\n            self.q_bias = None\n            self.v_bias = None\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj = nn.Linear(dim, dim)\n        self.proj_drop = nn.Dropout(proj_drop)\n        self.softmax = nn.Softmax(dim=-1)\n\n    def forward(self, x, mask=None):\n        \"\"\"\n        Args:\n            x: input features with shape of (num_windows*B, N, C)\n            mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None\n        \"\"\"\n        B_, N, C = x.shape\n        qkv_bias = None\n        if self.q_bias is not None:\n            qkv_bias = torch.cat((self.q_bias, torch.zeros_like(self.v_bias, requires_grad=False), self.v_bias))\n        qkv = F.linear(input=x, weight=self.qkv.weight, bias=qkv_bias)\n        qkv = qkv.reshape(B_, N, 3, self.num_heads, -1).permute(2, 0, 3, 1, 4)\n        q, k, v = qkv[0], qkv[1], qkv[2]  # make torchscript happy (cannot use tensor as tuple)\n\n        # cosine attention\n        attn = (F.normalize(q, dim=-1) @ F.normalize(k, dim=-1).transpose(-2, -1))\n        logit_scale = torch.clamp(self.logit_scale, max=torch.log(torch.tensor(1. / 0.01, device='cuda'))).exp()\n        attn = attn * logit_scale\n\n        relative_position_bias_table = self.cpb_mlp(self.relative_coords_table).view(-1, self.num_heads)\n        relative_position_bias = relative_position_bias_table[self.relative_position_index.view(-1)].view(\n            self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1)  # Wh*Ww,Wh*Ww,nH\n        relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous()  # nH, Wh*Ww, Wh*Ww\n        relative_position_bias = 16 * torch.sigmoid(relative_position_bias)\n        attn = attn + relative_position_bias.unsqueeze(0)\n\n        if mask is not None:\n            nW = mask.shape[0]\n            attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)\n            attn = attn.view(-1, self.num_heads, N, N)\n            attn = self.softmax(attn)\n        else:\n            attn = self.softmax(attn)\n\n        attn = self.attn_drop(attn)\n\n        x = (attn @ v).transpose(1, 2).reshape(B_, N, C)\n        x = self.proj(x)\n        x = self.proj_drop(x)\n        return x\n\n    def extra_repr(self) -> str:\n        return f'dim={self.dim}, window_size={self.window_size}, ' \\\n               f'pretrained_window_size={self.pretrained_window_size}, num_heads={self.num_heads}'\n\n    def flops(self, N):\n        # calculate flops for 1 window with token length of N\n        flops = 0\n        # qkv = self.qkv(x)\n        flops += N * self.dim * 3 * self.dim\n        # attn = (q @ k.transpose(-2, -1))\n        flops += self.num_heads * N * (self.dim // self.num_heads) * N\n        #  x = (attn @ v)\n        flops += self.num_heads * N * N * (self.dim // self.num_heads)\n        # x = self.proj(x)\n        flops += N * self.dim * self.dim\n        return flops\n\n\nclass SwinTransformerBlock(nn.Module):\n    r\"\"\" Swin Transformer Block.\n    Args:\n        dim (int): Number of input channels.\n        input_resolution (tuple[int]): Input resulotion.\n        num_heads (int): Number of attention heads.\n        window_size (int): Window size.\n        shift_size (int): Shift size for SW-MSA.\n        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.\n        qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True\n        drop (float, optional): Dropout rate. Default: 0.0\n        attn_drop (float, optional): Attention dropout rate. Default: 0.0\n        drop_path (float, optional): Stochastic depth rate. Default: 0.0\n        act_layer (nn.Module, optional): Activation layer. Default: nn.GELU\n        norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm\n        pretrained_window_size (int): Window size in pre-training.\n    \"\"\"\n\n    def __init__(self, dim, input_resolution, num_heads, window_size=7, shift_size=0,\n                 mlp_ratio=4., qkv_bias=True, drop=0., attn_drop=0., drop_path=0.,\n                 act_layer=nn.GELU, norm_layer=nn.LayerNorm, pretrained_window_size=0):\n        super().__init__()\n        self.dim = dim\n        self.input_resolution = input_resolution\n        self.num_heads = num_heads\n        self.window_size = window_size\n        self.shift_size = shift_size\n        self.mlp_ratio = mlp_ratio\n        if min(self.input_resolution) <= self.window_size:\n            # if window size is larger than input resolution, we don't partition windows\n            self.shift_size = 0\n            self.window_size = min(self.input_resolution)\n        assert 0 <= self.shift_size < self.window_size, \"shift_size must in 0-window_size\"\n\n        self.norm1 = norm_layer(dim)\n        self.attn = WindowAttention(\n            dim, window_size=to_2tuple(self.window_size), num_heads=num_heads,\n            qkv_bias=qkv_bias, attn_drop=attn_drop, proj_drop=drop,\n            pretrained_window_size=to_2tuple(pretrained_window_size))\n\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n        self.norm2 = norm_layer(dim)\n        mlp_hidden_dim = int(dim * mlp_ratio)\n        self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)\n\n        if self.shift_size > 0:\n            # calculate attention mask for SW-MSA\n            H, W = self.input_resolution\n            img_mask = torch.zeros((1, H, W, 1))  # 1 H W 1\n            h_slices = (slice(0, -self.window_size),\n                        slice(-self.window_size, -self.shift_size),\n                        slice(-self.shift_size, None))\n            w_slices = (slice(0, -self.window_size),\n                        slice(-self.window_size, -self.shift_size),\n                        slice(-self.shift_size, None))\n            cnt = 0\n            for h in h_slices:\n                for w in w_slices:\n                    img_mask[:, h, w, :] = cnt\n                    cnt += 1\n\n            mask_windows = window_partition(img_mask, self.window_size)  # nW, window_size, window_size, 1\n            mask_windows = mask_windows.view(-1, self.window_size * self.window_size)\n            attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)\n            attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))\n        else:\n            attn_mask = None\n\n        self.register_buffer(\"attn_mask\", attn_mask)\n\n    def forward(self, x):\n        H, W = self.input_resolution\n        B, L, C = x.shape\n        assert L == H * W, \"input feature has wrong size\"\n\n        shortcut = x\n        x = x.view(B, H, W, C)\n\n        # cyclic shift\n        if self.shift_size > 0:\n            shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))\n        else:\n            shifted_x = x\n\n        # partition windows\n        x_windows = window_partition(shifted_x, self.window_size)  # nW*B, window_size, window_size, C\n        x_windows = x_windows.view(-1, self.window_size * self.window_size, C)  # nW*B, window_size*window_size, C\n\n        # W-MSA/SW-MSA\n        attn_windows = self.attn(x_windows, mask=self.attn_mask)  # nW*B, window_size*window_size, C\n\n        # merge windows\n        attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)\n        shifted_x = window_reverse(attn_windows, self.window_size, H, W)  # B H' W' C\n\n        # reverse cyclic shift\n        if self.shift_size > 0:\n            x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))\n        else:\n            x = shifted_x\n        x = x.view(B, H * W, C)\n        x = shortcut + self.drop_path(self.norm1(x))\n\n        # FFN\n        x = x + self.drop_path(self.norm2(self.mlp(x)))\n\n        return x\n\n    def extra_repr(self) -> str:\n        return f\"dim={self.dim}, input_resolution={self.input_resolution}, num_heads={self.num_heads}, \" \\\n               f\"window_size={self.window_size}, shift_size={self.shift_size}, mlp_ratio={self.mlp_ratio}\"\n\n    def flops(self):\n        flops = 0\n        H, W = self.input_resolution\n        # norm1\n        flops += self.dim * H * W\n        # W-MSA/SW-MSA\n        nW = H * W / self.window_size / self.window_size\n        flops += nW * self.attn.flops(self.window_size * self.window_size)\n        # mlp\n        flops += 2 * H * W * self.dim * self.dim * self.mlp_ratio\n        # norm2\n        flops += self.dim * H * W\n        return flops\n\n\nclass PatchMerging(nn.Module):\n    r\"\"\" Patch Merging Layer.\n    Args:\n        input_resolution (tuple[int]): Resolution of input feature.\n        dim (int): Number of input channels.\n        norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm\n    \"\"\"\n\n    def __init__(self, input_resolution, dim, norm_layer=nn.LayerNorm):\n        super().__init__()\n        self.input_resolution = input_resolution\n        self.dim = dim\n        self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)\n        self.norm = norm_layer(2 * dim)\n\n    def forward(self, x):\n        \"\"\"\n        x: B, H*W, C\n        \"\"\"\n        H, W = self.input_resolution\n        B, L, C = x.shape\n        assert L == H * W, \"input feature has wrong size\"\n        assert H % 2 == 0 and W % 2 == 0, f\"x size ({H}*{W}) are not even.\"\n\n        x = x.view(B, H, W, C)\n\n        x0 = x[:, 0::2, 0::2, :]  # B H/2 W/2 C\n        x1 = x[:, 1::2, 0::2, :]  # B H/2 W/2 C\n        x2 = x[:, 0::2, 1::2, :]  # B H/2 W/2 C\n        x3 = x[:, 1::2, 1::2, :]  # B H/2 W/2 C\n        x = torch.cat([x0, x1, x2, x3], -1)  # B H/2 W/2 4*C\n        x = x.view(B, -1, 4 * C)  # B H/2*W/2 4*C\n\n        x = self.reduction(x)\n        x = self.norm(x)\n\n        return x\n\n    def extra_repr(self) -> str:\n        return f\"input_resolution={self.input_resolution}, dim={self.dim}\"\n\n    def flops(self):\n        H, W = self.input_resolution\n        flops = (H // 2) * (W // 2) * 4 * self.dim * 2 * self.dim\n        flops += H * W * self.dim // 2\n        return flops\n\n\nclass BasicLayer(nn.Module):\n    \"\"\" A basic Swin Transformer layer for one stage.\n    Args:\n        dim (int): Number of input channels.\n        input_resolution (tuple[int]): Input resolution.\n        depth (int): Number of blocks.\n        num_heads (int): Number of attention heads.\n        window_size (int): Local window size.\n        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.\n        qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True\n        drop (float, optional): Dropout rate. Default: 0.0\n        attn_drop (float, optional): Attention dropout rate. Default: 0.0\n        drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0\n        norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm\n        downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None\n        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.\n        pretrained_window_size (int): Local window size in pre-training.\n    \"\"\"\n\n    def __init__(self, dim, input_resolution, depth, num_heads, window_size,\n                 mlp_ratio=4., qkv_bias=True, drop=0., attn_drop=0.,\n                 drop_path=0., norm_layer=nn.LayerNorm, downsample=None, use_checkpoint=False,\n                 pretrained_window_size=0):\n\n        super().__init__()\n        self.dim = dim\n        self.input_resolution = input_resolution\n        self.depth = depth\n        self.use_checkpoint = use_checkpoint\n\n        # build blocks\n        self.blocks = nn.ModuleList([\n            SwinTransformerBlock(dim=dim, input_resolution=input_resolution,\n                                 num_heads=num_heads, window_size=window_size,\n                                 shift_size=0 if (i % 2 == 0) else window_size // 2,\n                                 mlp_ratio=mlp_ratio,\n                                 qkv_bias=qkv_bias,\n                                 drop=drop, attn_drop=attn_drop,\n                                 drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,\n                                 norm_layer=norm_layer,\n                                 pretrained_window_size=pretrained_window_size)\n            for i in range(depth)])\n\n        # patch merging layer\n        if downsample is not None:\n            self.downsample = downsample(input_resolution, dim=dim, norm_layer=norm_layer)\n        else:\n            self.downsample = None\n\n    def forward(self, x):\n        for blk in self.blocks:\n            if self.use_checkpoint:\n                x = checkpoint.checkpoint(blk, x)\n            else:\n                x = blk(x)\n        if self.downsample is not None:\n            x = self.downsample(x)\n        return x\n\n    def extra_repr(self) -> str:\n        return f\"dim={self.dim}, input_resolution={self.input_resolution}, depth={self.depth}\"\n\n    def flops(self):\n        flops = 0\n        for blk in self.blocks:\n            flops += blk.flops()\n        if self.downsample is not None:\n            flops += self.downsample.flops()\n        return flops\n\n    def _init_respostnorm(self):\n        for blk in self.blocks:\n            nn.init.constant_(blk.norm1.bias, 0)\n            nn.init.constant_(blk.norm1.weight, 0)\n            nn.init.constant_(blk.norm2.bias, 0)\n            nn.init.constant_(blk.norm2.weight, 0)\n\n\nclass PatchEmbed(nn.Module):\n    r\"\"\" Image to Patch Embedding\n    Args:\n        img_size (int): Image size.  Default: 224.\n        patch_size (int): Patch token size. Default: 4.\n        in_chans (int): Number of input image channels. Default: 3.\n        embed_dim (int): Number of linear projection output channels. Default: 96.\n        norm_layer (nn.Module, optional): Normalization layer. Default: None\n    \"\"\"\n\n    def __init__(self, img_size=224, patch_size=4, in_chans=3, embed_dim=96, norm_layer=None):\n        super().__init__()\n        img_size = to_2tuple(img_size)\n        patch_size = to_2tuple(patch_size)\n        patches_resolution = [img_size[0] // patch_size[0], img_size[1] // patch_size[1]]\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.patches_resolution = patches_resolution\n        self.num_patches = patches_resolution[0] * patches_resolution[1]\n\n        self.in_chans = in_chans\n        self.embed_dim = embed_dim\n\n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n        if norm_layer is not None:\n            self.norm = norm_layer(embed_dim)\n        else:\n            self.norm = None\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        # FIXME look at relaxing size constraints\n        assert H == self.img_size[0] and W == self.img_size[1], \\\n            f\"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]}).\"\n        x = self.proj(x).flatten(2).transpose(1, 2)  # B Ph*Pw C\n        if self.norm is not None:\n            x = self.norm(x)\n        return x\n\n    def flops(self):\n        Ho, Wo = self.patches_resolution\n        flops = Ho * Wo * self.embed_dim * self.in_chans * (self.patch_size[0] * self.patch_size[1])\n        if self.norm is not None:\n            flops += Ho * Wo * self.embed_dim\n        return flops\n\n\nclass SwinTransformerV2(nn.Module):\n    r\"\"\" Swin Transformer\n        A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows`  -\n          https://arxiv.org/pdf/2103.14030\n    Args:\n        img_size (int | tuple(int)): Input image size. Default 224\n        patch_size (int | tuple(int)): Patch size. Default: 4\n        in_chans (int): Number of input image channels. Default: 3\n        num_classes (int): Number of classes for classification head. Default: 1000\n        embed_dim (int): Patch embedding dimension. Default: 96\n        depths (tuple(int)): Depth of each Swin Transformer layer.\n        num_heads (tuple(int)): Number of attention heads in different layers.\n        window_size (int): Window size. Default: 7\n        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4\n        qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True\n        drop_rate (float): Dropout rate. Default: 0\n        attn_drop_rate (float): Attention dropout rate. Default: 0\n        drop_path_rate (float): Stochastic depth rate. Default: 0.1\n        norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.\n        ape (bool): If True, add absolute position embedding to the patch embedding. Default: False\n        patch_norm (bool): If True, add normalization after patch embedding. Default: True\n        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False\n        pretrained_window_sizes (tuple(int)): Pretrained window sizes of each layer.\n    \"\"\"\n\n    def __init__(self, img_size=224, patch_size=4, in_chans=3, num_classes=1000,\n                 embed_dim=96, depths=[2, 2, 6, 2], num_heads=[3, 6, 12, 24],\n                 window_size=7, mlp_ratio=4., qkv_bias=True,\n                 drop_rate=0., attn_drop_rate=0., drop_path_rate=0.1,\n                 norm_layer=nn.LayerNorm, ape=False, patch_norm=True,\n                 use_checkpoint=False, pretrained_window_sizes=[0, 0, 0, 0], **kwargs):\n        super().__init__()\n\n        self.num_classes = num_classes\n        self.num_layers = len(depths)\n        self.embed_dim = embed_dim\n        self.ape = ape\n        self.patch_norm = patch_norm\n        self.num_features = int(embed_dim * 2 ** (self.num_layers - 1))\n        self.mlp_ratio = mlp_ratio\n\n        # split image into non-overlapping patches\n        self.patch_embed = PatchEmbed(\n            img_size=img_size, patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim,\n            norm_layer=norm_layer if self.patch_norm else None)\n        num_patches = self.patch_embed.num_patches\n        patches_resolution = self.patch_embed.patches_resolution\n        self.patches_resolution = patches_resolution\n\n        # absolute position embedding\n        if self.ape:\n            self.absolute_pos_embed = nn.Parameter(torch.zeros(1, num_patches, embed_dim))\n            trunc_normal_(self.absolute_pos_embed, std=.02)\n\n        self.pos_drop = nn.Dropout(p=drop_rate)\n\n        # stochastic depth\n        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))]  # stochastic depth decay rule\n\n        # build layers\n        self.layers = nn.ModuleList()\n        for i_layer in range(self.num_layers):\n            layer = BasicLayer(dim=int(embed_dim * 2 ** i_layer),\n                               input_resolution=(patches_resolution[0] // (2 ** i_layer),\n                                                 patches_resolution[1] // (2 ** i_layer)),\n                               depth=depths[i_layer],\n                               num_heads=num_heads[i_layer],\n                               window_size=window_size,\n                               mlp_ratio=self.mlp_ratio,\n                               qkv_bias=qkv_bias,\n                               drop=drop_rate, attn_drop=attn_drop_rate,\n                               drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])],\n                               norm_layer=norm_layer,\n                               downsample=PatchMerging if (i_layer < self.num_layers - 1) else None,\n                               use_checkpoint=use_checkpoint,\n                               pretrained_window_size=pretrained_window_sizes[i_layer])\n            self.layers.append(layer)\n\n        self.norm = norm_layer(self.num_features)\n        self.avgpool = nn.AdaptiveAvgPool1d(1)\n        self.head = nn.Linear(self.num_features, num_classes) if num_classes > 0 else nn.Identity()\n\n        self.apply(self._init_weights)\n        for bly in self.layers:\n            bly._init_respostnorm()\n\n    def _init_weights(self, m):\n        if isinstance(m, nn.Linear):\n            trunc_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        elif isinstance(m, nn.LayerNorm):\n            nn.init.constant_(m.bias, 0)\n            nn.init.constant_(m.weight, 1.0)\n\n    @torch.jit.ignore\n    def no_weight_decay(self):\n        return {'absolute_pos_embed'}\n\n    @torch.jit.ignore\n    def no_weight_decay_keywords(self):\n        return {\"cpb_mlp\", \"logit_scale\", 'relative_position_bias_table'}\n\n    def forward_features(self, x):\n        x = self.patch_embed(x)\n        if self.ape:\n            x = x + self.absolute_pos_embed\n        x = self.pos_drop(x)\n\n        for layer in self.layers:\n            x = layer(x)\n\n        x = self.norm(x)  # B L C\n        x = self.avgpool(x.transpose(1, 2))  # B C 1\n        x = torch.flatten(x, 1)\n        return x\n\n    def extra_features(self, x):\n        x = self.patch_embed(x)\n        if self.ape:\n            x = x + self.absolute_pos_embed\n        x = self.pos_drop(x)\n        feature = []\n\n        for layer in self.layers:\n            x = layer(x)\n            bs, n, f = x.shape\n            h = int(n ** 0.5)\n\n            feature.append(x.view(-1, h, h, f).permute(0, 3, 1, 2).contiguous())\n        return feature\n\n    def get_unet_feature(self, x):\n        x = self.patch_embed(x)\n        if self.ape:\n            x = x + self.absolute_pos_embed\n        x = self.pos_drop(x)\n        bs, n, f = x.shape\n        h = int(n ** 0.5)\n        feature = [x.view(-1, h, h, f).permute(0, 3, 1, 2).contiguous()]\n\n        for layer in self.layers:\n            x = layer(x)\n            bs, n, f = x.shape\n            h = int(n ** 0.5)\n\n            feature.append(x.view(-1, h, h, f).permute(0, 3, 1, 2).contiguous())\n        return feature\n\n    def forward(self, x):\n        x = self.forward_features(x)\n        x = self.head(x)\n        return x\n\n    def flops(self):\n        flops = 0\n        flops += self.patch_embed.flops()\n        for i, layer in enumerate(self.layers):\n            flops += layer.flops()\n        flops += self.num_features * self.patches_resolution[0] * self.patches_resolution[1] // (2 ** self.num_layers)\n        flops += self.num_features * self.num_classes\n        return flops\n\n\ndef swin_v2(size, img_size=256, in_22k=False, config=None, pretrained=False, **kwargs):\n    if size == \"swinv2_tiny_window16_256\":\n        model = SwinTransformerV2(img_size=img_size, window_size=16, embed_dim=96, depths=[2, 2, 6, 2],\n                                  num_heads=[3, 6, 12, 24], **kwargs)\n        if pretrained:\n            checkpoint = torch.load(config.model_urls[size])[\"model\"]\n            if img_size != 256:\n                del checkpoint[\"layers.0.blocks.0.attn.relative_coords_table\"]\n                del checkpoint[\"layers.0.blocks.0.attn.relative_position_index\"]\n                del checkpoint[\"layers.0.blocks.1.attn_mask\"]\n                del checkpoint[\"layers.0.blocks.1.attn.relative_coords_table\"]\n                del checkpoint[\"layers.0.blocks.1.attn.relative_position_index\"]\n                del checkpoint[\"layers.1.blocks.0.attn.relative_coords_table\"]\n                del checkpoint[\"layers.1.blocks.0.attn.relative_position_index\"]\n                del checkpoint[\"layers.1.blocks.1.attn_mask\"]\n                del checkpoint[\"layers.1.blocks.1.attn.relative_coords_table\"]\n                del checkpoint[\"layers.1.blocks.1.attn.relative_position_index\"]\n                del checkpoint[\"layers.2.blocks.0.attn.relative_coords_table\"]\n                del checkpoint[\"layers.2.blocks.0.attn.relative_position_index\"]\n                del checkpoint[\"layers.2.blocks.1.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.1.attn.relative_coords_table\"]\n                del checkpoint[\"layers.2.blocks.1.attn.relative_position_index\"]\n                del checkpoint[\"layers.2.blocks.2.attn.relative_coords_table\"]\n                del checkpoint[\"layers.2.blocks.2.attn.relative_position_index\"]\n                del checkpoint[\"layers.2.blocks.3.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.3.attn.relative_coords_table\"]\n                del checkpoint[\"layers.2.blocks.3.attn.relative_position_index\"]\n                del checkpoint[\"layers.2.blocks.4.attn.relative_coords_table\"]\n                del checkpoint[\"layers.2.blocks.4.attn.relative_position_index\"]\n                del checkpoint[\"layers.2.blocks.5.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.5.attn.relative_coords_table\"]\n                del checkpoint[\"layers.2.blocks.5.attn.relative_position_index\"]\n                del checkpoint[\"layers.3.blocks.0.attn.relative_coords_table\"]\n                del checkpoint[\"layers.3.blocks.0.attn.relative_position_index\"]\n                del checkpoint[\"layers.3.blocks.1.attn.relative_coords_table\"]\n                del checkpoint[\"layers.3.blocks.1.attn.relative_position_index\"]\n            model.load_state_dict(checkpoint, strict=False)\n        else:\n            pass\n        \n        \n    elif size == \"swinv2_small_window8_256\":\n        model = SwinTransformerV2(img_size=img_size, window_size=8, embed_dim=96, depths=[2, 2, 18, 2],\n                                  num_heads=[3, 6, 12, 24], **kwargs)\n        \n        if pretrained:\n            checkpoint = torch.load(config.model_urls[size])[\"model\"]\n            if img_size != 256:\n                del checkpoint[\"layers.0.blocks.1.attn_mask\"]\n                del checkpoint[\"layers.1.blocks.1.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.1.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.3.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.5.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.7.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.9.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.11.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.13.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.15.attn_mask\"]\n                del checkpoint[\"layers.2.blocks.17.attn_mask\"]\n\n            model.load_state_dict(checkpoint, strict=False)\n        \n        else:\n            pass\n        \n    elif size == \"swinv2_small_window16_256\":\n        model = SwinTransformerV2(img_size=img_size, window_size=16, embed_dim=96, depths=[2, 2, 18, 2],\n                                  num_heads=[3, 6, 12, 24], **kwargs)\n        if pretrained:\n            checkpoint = torch.load(config.model_urls[size])[\"model\"]\n            if img_size != 256:\n                del checkpoint[\"layers.0.blocks.1.attn_mask\"]\n                del checkpoint[\"layers.1.blocks.1.attn_mask\"]\n                del checkpoint[\"layers.3.blocks.0.attn.relative_coords_table\"]\n                del checkpoint[\"layers.3.blocks.0.attn.relative_position_index\"]\n                del checkpoint[\"layers.3.blocks.1.attn.relative_coords_table\"]\n                del checkpoint[\"layers.3.blocks.1.attn.relative_position_index\"]\n            model.load_state_dict(checkpoint, strict=False)\n        else:\n            pass\n    elif size == \"swinv2_base_window16_256\":\n        model = SwinTransformerV2(img_size=img_size, window_size=16, embed_dim=128, depths=[2, 2, 18, 2],\n                                  num_heads=[4, 8, 16, 32], **kwargs)\n        if pretrained:\n            checkpoint = torch.load(config.model_urls[size])[\"model\"]\n            if img_size != 256:\n                del checkpoint[\"layers.0.blocks.1.attn_mask\"]\n                del checkpoint[\"layers.1.blocks.1.attn_mask\"]\n                del checkpoint[\"layers.3.blocks.0.attn.relative_coords_table\"]\n                del checkpoint[\"layers.3.blocks.0.attn.relative_position_index\"]\n                del checkpoint[\"layers.3.blocks.1.attn.relative_coords_table\"]\n                del checkpoint[\"layers.3.blocks.1.attn.relative_position_index\"]\n            model.load_state_dict(checkpoint, strict=False)\n        else:\n            pass\n    \n    model = model.to(cfg.device)\n    return model\n\n\nclass PSPModule(nn.Module):\n    # In the original inmplementation they use precise RoI pooling\n    # Instead of using adaptative average pooling\n    def __init__(self, in_channels, bin_sizes=[1, 2, 4, 6]):\n        super(PSPModule, self).__init__()\n        out_channels = in_channels // len(bin_sizes)\n        self.stages = nn.ModuleList([self._make_stages(in_channels, out_channels, b_s)\n                                     for b_s in bin_sizes])\n        self.bottleneck = nn.Sequential(\n            nn.Conv2d(in_channels + (out_channels * len(bin_sizes)), in_channels,\n                      kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(in_channels),\n            nn.ReLU(inplace=True),\n            nn.Dropout2d(0.1)\n        )\n\n    def _make_stages(self, in_channels, out_channels, bin_sz):\n        prior = nn.AdaptiveAvgPool2d(output_size=bin_sz)\n        conv = nn.Conv2d(in_channels, out_channels, kernel_size=1, bias=False)\n        bn = nn.BatchNorm2d(out_channels)\n        relu = nn.ReLU(inplace=True)\n        return nn.Sequential(prior, conv, bn, relu)\n\n    def forward(self, features):\n        h, w = features.size()[2], features.size()[3]\n        pyramids = [features]\n        pyramids.extend([F.interpolate(stage(features), size=(h, w), mode='bilinear',\n                                       align_corners=True) for stage in self.stages])\n        output = self.bottleneck(torch.cat(pyramids, dim=1))\n        return output\n\n\nclass ResNet(nn.Module):\n    def __init__(self, in_channels=3, output_stride=16, backbone='resnet101', pretrained=True):\n        super(ResNet, self).__init__()\n        model = getattr(models, backbone)(pretrained)\n        if not pretrained or in_channels != 3:\n            self.initial = nn.Sequential(\n                nn.Conv2d(in_channels, 64, 7, stride=2, padding=3, bias=False),\n                nn.BatchNorm2d(64),\n                nn.ReLU(inplace=True),\n                nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n            )\n            # initialize_weights(self.initial)\n        else:\n            self.initial = nn.Sequential(*list(model.children())[:4])\n\n        self.layer1 = model.layer1\n        self.layer2 = model.layer2\n        self.layer3 = model.layer3\n        self.layer4 = model.layer4\n\n        if output_stride == 16:\n            s3, s4, d3, d4 = (2, 1, 1, 2)\n        elif output_stride == 8:\n            s3, s4, d3, d4 = (1, 1, 2, 4)\n\n        if output_stride == 8:\n            for n, m in self.layer3.named_modules():\n                if 'conv1' in n and (backbone == 'resnet34' or backbone == 'resnet18'):\n                    m.dilation, m.padding, m.stride = (d3, d3), (d3, d3), (s3, s3)\n                elif 'conv2' in n:\n                    m.dilation, m.padding, m.stride = (d3, d3), (d3, d3), (s3, s3)\n                elif 'downsample.0' in n:\n                    m.stride = (s3, s3)\n\n        for n, m in self.layer4.named_modules():\n            if 'conv1' in n and (backbone == 'resnet34' or backbone == 'resnet18'):\n                m.dilation, m.padding, m.stride = (d4, d4), (d4, d4), (s4, s4)\n            elif 'conv2' in n:\n                m.dilation, m.padding, m.stride = (d4, d4), (d4, d4), (s4, s4)\n            elif 'downsample.0' in n:\n                m.stride = (s4, s4)\n\n    def forward(self, x):\n        x = self.initial(x)\n        x1 = self.layer1(x)\n        print(\"x1\", x1.shape)\n        x2 = self.layer2(x1)\n        print(\"x2\", x2.shape)\n        x3 = self.layer3(x2)\n        print(\"x3\", x3.shape)\n        x4 = self.layer4(x3)\n        print(\"x4\", x4.shape)\n\n        return [x1, x2, x3, x4]\n\n\ndef up_and_add(x, y):\n    return F.interpolate(x, size=(y.size(2), y.size(3)), mode='bilinear', align_corners=True) + y\n\n\nclass FPN_fuse(nn.Module):\n    def __init__(self, feature_channels=[256, 512, 1024, 2048], fpn_out=256):\n        super(FPN_fuse, self).__init__()\n        assert feature_channels[0] == fpn_out\n        self.conv1x1 = nn.ModuleList([nn.Conv2d(ft_size, fpn_out, kernel_size=1)\n                                      for ft_size in feature_channels[1:]])\n        self.smooth_conv = nn.ModuleList([nn.Conv2d(fpn_out, fpn_out, kernel_size=3, padding=1)]\n                                         * (len(feature_channels) - 1))\n        self.conv_fusion = nn.Sequential(\n            nn.Conv2d(len(feature_channels) * fpn_out, fpn_out, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(fpn_out),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, features):\n        features[1:] = [conv1x1(feature) for feature, conv1x1 in zip(features[1:], self.conv1x1)]  ##\n        P = [up_and_add(features[i], features[i - 1]) for i in reversed(range(1, len(features)))]\n        P = [smooth_conv(x) for smooth_conv, x in zip(self.smooth_conv, P)]\n        P = list(reversed(P))\n        P.append(features[-1])  # P = [P1, P2, P3, P4]\n        H, W = P[0].size(2), P[0].size(3)\n        P[1:] = [F.interpolate(feature, size=(H, W), mode='bilinear', align_corners=True) for feature in P[1:]]\n\n        x = self.conv_fusion(torch.cat((P), dim=1))\n        return x\n\n\nclass UperNet_swin(nn.Module):\n    # Implementing only the object path\n    def __init__(self, size=\"swinv2_small_window16_256\", config=None, img_size=256, num_classes=1, in_channels=3, pretrained=True):\n        super(UperNet_swin, self).__init__()\n\n        self.backbone = swin_v2(size=size, img_size=img_size, config=config)\n        if size.split(\"_\")[1] in [\"small\", \"tiny\"]:\n            feature_channels = [192, 384, 768, 768]\n        elif size.split(\"_\")[1] in [\"base\"]:\n            feature_channels = [256, 512, 1024, 1024]\n        self.PPN = PSPModule(feature_channels[-1])\n        self.FPN = FPN_fuse(feature_channels, fpn_out=feature_channels[0])\n        self.head = nn.Conv2d(feature_channels[0], num_classes, kernel_size=3, padding=1)\n\n    def forward(self, x):\n        input_size = (x.size()[2], x.size()[3])\n\n        features = self.backbone.extra_features(x)\n        features[-1] = self.PPN(features[-1])\n        x = self.head(self.FPN(features))\n\n        x = F.interpolate(x, size=input_size, mode='bilinear')\n        return x\n\n    def get_backbone_params(self):\n        return self.backbone.parameters()\n\n    def get_decoder_params(self):\n        return chain(self.PPN.parameters(), self.FPN.parameters(), self.head.parameters())\n\n    def freeze_bn(self):\n        for module in self.modules():\n            if isinstance(module, nn.BatchNorm2d): module.eval()\n\n\nclass DecoderBlock(nn.Module):\n    def __init__(\n            self,\n            in_channels,\n            skip_channels,\n            out_channels,\n            use_batchnorm=True,\n            attention_type=None,\n    ):\n        super().__init__()\n        self.conv1 = md.Conv2dReLU(\n            in_channels + skip_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        if attention_type == \"cbam\":\n            self.attention1 = CbamModule(channels=in_channels + skip_channels)\n        else:\n            self.attention1 = md.Attention(attention_type, in_channels=in_channels + skip_channels)\n        self.conv2 = md.Conv2dReLU(\n            out_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        if attention_type == \"cbam\":\n            self.attention2 = CbamModule(channels=out_channels)\n        else:\n            self.attention2 = md.Attention(attention_type, in_channels=out_channels)\n        self.in_channels = in_channels\n        self.out_channels = out_channels\n        self.skip_channels = skip_channels\n\n    def forward(self, x, skip=None):\n        if skip is None:\n            x = F.interpolate(x, scale_factor=2, mode=\"nearest\")\n        else:\n            if x.shape[-1] != skip.shape[-1]:\n                x = F.interpolate(x, scale_factor=2, mode=\"nearest\")\n        if skip is not None:\n            # print(x.shape,skip.shape)\n            x = torch.cat([x, skip], dim=1)\n            x = self.attention1(x)\n        x = self.conv1(x)\n        x = self.conv2(x)\n        x = self.attention2(x)\n        return x\n\n\nclass CenterBlock(nn.Sequential):\n    def __init__(self, in_channels, out_channels, use_batchnorm=True):\n        conv1 = md.Conv2dReLU(\n            in_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        conv2 = md.Conv2dReLU(\n            out_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n        super().__init__(conv1, conv2)\n\n\nclass UnetDecoder(nn.Module):\n    def __init__(\n            self,\n            encoder_channels,\n            decoder_channels,\n            n_blocks=5,\n            use_batchnorm=True,\n            attention_type=None,\n            center=False,\n    ):\n        super().__init__()\n\n        if n_blocks != len(decoder_channels):\n            raise ValueError(\n                \"Model depth is {}, but you provide `decoder_channels` for {} blocks.\".format(\n                    n_blocks, len(decoder_channels)\n                )\n            )\n\n        # remove first skip with same spatial resolution\n        encoder_channels = encoder_channels[1:]\n        # reverse channels to start from head of encoder\n        encoder_channels = encoder_channels[::-1]\n\n        # computing blocks input and output channels\n        head_channels = encoder_channels[0]\n        in_channels = [head_channels] + list(decoder_channels[:-1])\n        skip_channels = list(encoder_channels[1:]) + [0]\n\n        out_channels = decoder_channels\n\n        if center:\n            self.center = CenterBlock(head_channels, head_channels, use_batchnorm=use_batchnorm)\n        else:\n            self.center = nn.Identity()\n\n        # combine decoder keyword arguments\n        kwargs = dict(use_batchnorm=use_batchnorm, attention_type=attention_type)\n        blocks = [\n            DecoderBlock(in_ch, skip_ch, out_ch, **kwargs)\n            for in_ch, skip_ch, out_ch in zip(in_channels, skip_channels, out_channels)\n        ]\n        self.blocks = nn.ModuleList(blocks)\n\n    def forward(self, *features):\n\n        features = features[1:]  # remove first skip with same spatial resolution\n        features = features[::-1]  # reverse channels to start from head of encoder\n\n        head = features[0]\n        skips = features[1:]\n\n        x = self.center(head)\n        for i, decoder_block in enumerate(self.blocks):\n            skip = skips[i] if i < len(skips) else None\n            x = decoder_block(x, skip)\n            # y_i = self.upsample1(y_i)\n        # hypercol = torch.cat([y0,y1,y2,y3,y4], dim=1)\n\n        return x\n\n\nclass SegmentationHead(nn.Sequential):\n    def __init__(self, in_channels, out_channels, kernel_size=3, upsampling=1):\n        conv2d = nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, padding=kernel_size // 2)\n        upsampling = nn.UpsamplingBilinear2d(scale_factor=upsampling) if upsampling > 1 else nn.Identity()\n        super().__init__(conv2d, upsampling)\n\n\nclass unet_swin(nn.Module):\n\n    def __init__(\n            self, config, size=\"small\", img_size=256  # \"base\" \"large\"\n    ):\n        super().__init__()\n\n        self.encoder = swin_v2(size=size, img_size=img_size, config=config)\n\n        if size.split(\"_\")[1] in [\"small\", \"tiny\"]:\n            feature_channels = (3, 192, 384, 768, 768)\n        elif size.split(\"_\")[1] in [\"base\"]:\n            feature_channels = (3, 256, 512, 1024, 1024)\n        self.decoder = UnetDecoder(encoder_channels=feature_channels, n_blocks=4, decoder_channels=(512, 256, 128, 64),\n                                   attention_type=None)\n\n        self.segmentation_head = SegmentationHead(in_channels=64, out_channels=1, kernel_size=3, upsampling=4\n                                                  )\n\n    def forward(self, input):\n        encoder_featrue = self.encoder.get_unet_feature(input)\n        decoder_output = self.decoder(*encoder_featrue)\n        masks = self.segmentation_head(decoder_output)\n\n        return masks","metadata":{"execution":{"iopub.status.busy":"2024-03-09T03:16:38.353636Z","iopub.execute_input":"2024-03-09T03:16:38.354009Z","iopub.status.idle":"2024-03-09T03:16:38.532089Z","shell.execute_reply.started":"2024-03-09T03:16:38.353979Z","shell.execute_reply":"2024-03-09T03:16:38.531104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef valid_one_epoch(model, dataloader, device, epoch, optimizer, cfg):\n    model.eval()\n    total_val_loss = 0\n    total_val_score = 0\n\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid ')\n\n    for step, (images, masks) in pbar:\n        images = images.to(device, dtype=torch.float)\n        masks = masks.to(device, dtype=torch.float)\n\n        y_pred = model(images)\n        criterion = DiceBCELoss()\n        loss = criterion(y_pred, masks)\n        dice_score = DiceScore()(y_pred, masks).detach().item()\n\n        loss = loss.detach().item()\n        total_val_loss += loss\n        total_val_score += dice_score\n\n    print(f'\\nTesting epoch {epoch} ')\n    print(f'Total DiceBCE loss: {total_val_loss / len(dataloader):.4f}')\n    print(f'Total average Dice Score: {total_val_score / len(dataloader):.4f}')\n\n    torch.cuda.empty_cache()\n    gc.collect()\n\n    return total_val_loss / len(dataloader), total_val_score / len(dataloader)\n\n\ndef train_one_epoch(model, optimizer, scheduler, dataloader, device, epoch, cfg):\n    \n    model.train()\n    scaler = amp.GradScaler()\n    total_loss = 0\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ', leave=False)\n    data_size = 0\n    total_dice_score = 0\n    \n    for step, (images, masks) in pbar:\n        images = images.to(device, dtype=torch.float)\n        masks = masks.to(device, dtype=torch.float)\n\n        batch_size = images.size(0)\n        data_size += batch_size\n\n        with amp.autocast(enabled=True):\n            y_pred = model(images)\n            criterion = DiceBCELoss()\n            loss = criterion(y_pred, masks)\n            dice_score = DiceScore()(y_pred, masks).detach().item()\n\n        scaler.scale(loss / cfg.n_accumulate).backward()\n\n        if ((step + 1) % cfg.n_accumulate == 0 or (step + 1) == len(dataloader)):\n\n            scaler.step(optimizer)\n            scaler.update()\n            # zero the parameter gradients\n            optimizer.zero_grad()\n\n            if scheduler is not None:\n                scheduler.step()\n\n        loss = loss.detach().item()\n        total_loss += loss\n        total_dice_score += dice_score\n\n        pbar.set_postfix(desc=f'Loss={loss:.4f} DiceScore= {dice_score:.4f}  Batch_id={step}')\n\n    print(f'\\nTraining epoch {epoch} ')\n    print(f'Total DiceBCE loss: {total_loss / len(dataloader):.4f}')\n    print(f'Total average Dice Score: {total_dice_score / len(dataloader):.4f}')\n\n    torch.cuda.empty_cache()\n    gc.collect()\n\n    return (total_loss / len(dataloader), total_dice_score / len(dataloader))\n\n\ndef build_model(cfg):\n    model = unet_swin(img_size=256, size=cfg.size, config=cfg)\n    model.to(cfg.device)\n    return model\n\n\ndef load_model(path, cfg=None):\n    model = build_model(cfg)\n    model.load_state_dict(torch.load(path))\n    model.eval()\n    return model\n\n\ndef save_img(data, name, out):\n    data = data.float().cpu().numpy()\n    img = cv2.imencode('.png', (data * 255).astype(np.uint8))[1]\n    out.writestr(name, img)\n\n\ndef fetch_scheduler(optimizer, cfg):\n    if cfg.scheduler == 'CosineAnnealingLR':\n        scheduler = lr_scheduler.CosineAnnealingLR(optimizer, T_max=cfg.T_max,\n                                                   eta_min=cfg.min_lr)\n    elif cfg.scheduler == 'CosineAnnealingWarmRestarts':\n        scheduler = lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=cfg.T_0,\n                                                             eta_min=cfg.min_lr)\n    elif cfg.scheduler == 'ReduceLROnPlateau':\n        scheduler = lr_scheduler.ReduceLROnPlateau(optimizer,\n                                                   mode='min',\n                                                   factor=0.1,\n                                                   patience=7,\n                                                   threshold=0.0001,\n                                                   min_lr=cfg.min_lr, )\n    elif cfg.scheduer == 'ExponentialLR':\n        scheduler = lr_scheduler.ExponentialLR(optimizer, gamma=0.85)\n\n    else:\n        scheduler = None\n\n    return scheduler\n\n\nclass Model_pred:\n    def __init__(self, model, dl, tta: bool = True, half: bool = False, config=None):\n        self.model = model\n        self.dl = dl\n        self.tta = tta\n        self.half = half\n        self.config = config\n\n    def __iter__(self):\n        self.model.eval()\n        name_list = self.dl.dataset.graph_list\n        count = 0\n        with torch.no_grad():\n            for x, y in iter(self.dl):\n                if self.config.device != \"cpu\":\n                    x = x.to(self.config.device)\n                if self.half:\n                    x = x.half()\n                x = x.type(torch.float)\n                p = self.model(x)\n                py = torch.sigmoid(p).detach()\n                if self.tta:\n                    # x,y,xy flips as TTA\n                    flips = [[-1], [-2], [-2, -1]]\n                    for f in flips:\n                        p = self.model(torch.flip(x, f))\n                        p = torch.flip(p, f)\n                        py += torch.sigmoid(p).detach()\n                    py /= (1 + len(flips))\n                if y is not None and len(y.shape) == 4 and py.shape != y.shape:\n                    py = F.upsample(py, size=(y.shape[-2], y.shape[-1]), mode=\"bilinear\")\n                py = py.permute(0, 2, 3, 1).float().cpu()\n                batch_size = len(py)\n                for i in range(batch_size):\n                    taget = y[i].detach().cpu() if y is not None else None\n                    yield py[i], taget, name_list[count]\n                    count += 1\n\n    def __len__(self):\n        return len(self.dl.dataset)","metadata":{"execution":{"iopub.status.busy":"2024-03-09T03:16:38.533984Z","iopub.execute_input":"2024-03-09T03:16:38.534293Z","iopub.status.idle":"2024-03-09T03:16:38.563452Z","shell.execute_reply.started":"2024-03-09T03:16:38.534267Z","shell.execute_reply":"2024-03-09T03:16:38.562541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def score(weight_path):\n    score_lindex = weight_path.rindex(\"_\") + 1\n    score_rindex = weight_path.rindex(\".\")\n    return float(weight_path[score_lindex:score_rindex])\n\ndef set_seed(seed=42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    print(f\"Setting seed as {seed}\")\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    print('> SEEDING DONE')\n\n\ndef initialise_config(debug=False, train_bs=4, fold=0, pretrained_path=None):\n    cfg = CFG(fold=fold, train_bs=train_bs, debug=debug, pretrained_path=pretrained_path)\n    set_seed(cfg.seed)\n    return cfg\n\n\ndef create_folds(cfg=None):\n    if cfg.train_dataset == \"hap\":\n        image_name_list = [ i[i.rindex(\"/\"):]for i in glob.glob(\"../input/hubmap-2022-256x256/train/*.png\")]\n    elif cfg.train_dataset == \"hubmap\":\n        image_name_list = [i[i.rindex(\"/\"):] for i in glob.glob(\"../hubmap-256x256/train/*.png\")]\n    else:\n        image_name_list = [i[i.rindex(\"/\"):] for i in glob.glob(\"../all_256/train/*.png\")]\n\n    df = pd.DataFrame({\"graph_name\":image_name_list})\n    skf = KFold(n_splits=cfg.n_fold, shuffle=True, random_state=cfg.seed)\n    for fold, idxes in enumerate(skf.split(range(len(df)))):\n        df.loc[idxes[1], 'fold'] = fold\n    return df\n\n\ndef run_training(model, optimizer, scheduler, device, num_epochs, fold, train_loader, valid_loader, cfg):\n    if device != \"cpu\":\n        print(\"cuda: {}\\n\".format(torch.cuda.get_device_name()))\n\n    start = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_dice = 0\n    history = defaultdict(list)\n\n    if cfg.only_dice != 1:\n        for epoch in range(cfg.epochs):\n\n            print(f'Epoch {epoch}/{num_epochs}', end='')\n            train_loss, train_score = train_one_epoch(model, optimizer, scheduler,\n                                                      dataloader=train_loader,\n                                                      device=cfg.device, epoch=epoch, cfg=cfg)\n\n            val_loss, val_score = valid_one_epoch(model, valid_loader,\n                                                  device=cfg.device,\n                                                  epoch=epoch,\n                                                  optimizer=optimizer, cfg=cfg)\n\n            history['epoch'].append(epoch)\n            history['Train Loss'].append(train_loss)\n            history['Valid Loss'].append(val_loss)\n            history['Valid Scores'].append(val_score)\n\n            print(f'Train Loss: {train_loss} | Valid Loss: {val_loss}')\n            print(f'Train Score: {train_score} | Valid Dice Score: {val_score}')\n\n            # deep copy the model\n            if val_score >= best_dice:\n#                 os.system(f\"rm models/fold_{fold}/{cfg.size}_*\")\n                print(f\"Valid Score Improved ({best_dice:0.4f} ---> {val_score:0.4f})\")\n                best_dice = val_score\n                best_model_wts = copy.deepcopy(model.state_dict())\n                PATH = f\"Models/fold_{fold}/{cfg.size}_{val_score:0.4f}.pth\"\n                torch.save(model.state_dict(), PATH)\n\n                print(f\"Model Saved\")\n\n            print()\n            print()\n\n        end = time.time()\n        time_elapsed = end - start\n        print('Training complete in {:.0f}h {:.0f}m {:.0f}s'.format(\n            time_elapsed // 3600, (time_elapsed % 3600) // 60, (time_elapsed % 3600) % 60))\n\n        model.load_state_dict(best_model_wts)\n\n        plt.subplot(1, 2, 1, frameon=False)\n        plt.title(f'fold_{fold}_train_loss')\n        plt.xlabel('Epoch')\n        plt.plot(history['epoch'], history['Train Loss'], \"r\")\n\n        plt.subplot(1, 2, 2, frameon=False)\n        plt.title(f'fold_{fold}_test_dice')\n        plt.xlabel('Epoch')\n        plt.plot(history['epoch'], history['Valid Scores'], \"b\")\n\n        plt.savefig(f\"Models/fold_{fold}/metric_fold_{fold}.jpg\")\n        plt.close()\n\n    dice_loader = prepare_valid_loaders(cfg)\n    mp = Model_pred(model, dice_loader, config=cfg)\n    dice = Dice_th_pred(np.arange(0.2, 0.7, 0.01))\n    for p in progress_bar(mp):\n        dice.accumulate(p[0], p[1])\n    # save_img(p[0], p[2], out)\n    gc.collect()\n    dices = dice.value\n    noise_ths = dice.ths\n    best_dice = dices.max()\n    best_thr = noise_ths[dices.argmax()]\n    plt.figure(figsize=(8, 4))\n    plt.plot(noise_ths, dices, color='blue')\n    plt.vlines(x=best_thr, ymin=dices.min(), ymax=dices.max(), colors='black')\n    d = dices.max() - dices.min()\n    plt.text(noise_ths[-1] - 0.1, best_dice - 0.1 * d, f'DICE = {best_dice:.3f}', fontsize=12)\n    plt.text(noise_ths[-1] - 0.1, best_dice - 0.2 * d, f'TH = {best_thr:.3f}', fontsize=12)\n    plt.savefig(f'Models/fold_{fold}/save.jpg')\n    plt.close()\n\n    weight_path = glob.glob(f\"Models/fold_{fold}/{cfg.size}*.pth\")[0]\n    down_index = weight_path.rindex(\"_\")\n    new_weight_path = weight_path[:down_index] + f\"_{best_thr:.3f}\" + weight_path[down_index:]\n    os.rename(weight_path, new_weight_path)\n\n    return model, history","metadata":{"execution":{"iopub.status.busy":"2024-03-09T03:33:58.211040Z","iopub.execute_input":"2024-03-09T03:33:58.211417Z","iopub.status.idle":"2024-03-09T03:33:58.240470Z","shell.execute_reply.started":"2024-03-09T03:33:58.211389Z","shell.execute_reply":"2024-03-09T03:33:58.239563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main(cfg):\n    cfg.display()\n    print(f'#' * 30)\n    print(f'### Fold: {cfg.fold}')\n    print(f'#' * 30)\n\n    train_loader, valid_loader = prepare_train_loaders(fold=cfg.fold,\n                                                     df=create_folds(cfg),\n                                                     debug=cfg.debug,\n                                                     cfg=cfg)\n    \n    if cfg.load_best_model:\n        models = glob.glob(f\"Models/fold_{cfg.fold}/{cfg.size}_*.pth\")\n        models = sorted(models, key=lambda i: score(i), reverse=True)\n        model = load_model(models[0], cfg=cfg).to(cfg.device)\n        print(\"Load Pretrained Model: \" + models[0])\n    elif cfg.pretrained_path is None:\n        model = unet_swin(img_size=256, size=cfg.size, config=cfg).to(cfg.device)\n    else:\n        model = load_model(cfg.pretrained_path, cfg=cfg).to(cfg.device)\n        print(\"Load pretrained Model: \" + cfg.pretrained_path)\n\n    optimizer = optim.Adam(model.parameters(), lr=cfg.lr, weight_decay=cfg.wd)\n    scheduler = fetch_scheduler(optimizer, cfg=cfg)\n    run_training(model, optimizer, scheduler,\n                     device=cfg.device,\n                     num_epochs=cfg.epochs, fold=cfg.fold,\n                     train_loader=train_loader,\n                     valid_loader=valid_loader,\n                     cfg=cfg)","metadata":{"execution":{"iopub.status.busy":"2024-03-09T03:34:06.685662Z","iopub.execute_input":"2024-03-09T03:34:06.686076Z","iopub.status.idle":"2024-03-09T03:34:06.696084Z","shell.execute_reply.started":"2024-03-09T03:34:06.686044Z","shell.execute_reply":"2024-03-09T03:34:06.694956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg = initialise_config(train_bs=8, fold=0)\nif not os.path.exists('Models'):\n    print(f'******Create Models folder******')\n    os.mkdir('Models')\n    \nif not os.path.exists(f\"Models/fold_{cfg.fold}\"):\n    os.mkdir(f\"Models/fold_{cfg.fold}\")\n\nmain(cfg)","metadata":{"execution":{"iopub.status.busy":"2024-03-09T03:36:06.843394Z","iopub.execute_input":"2024-03-09T03:36:06.843795Z","iopub.status.idle":"2024-03-09T03:46:27.956271Z","shell.execute_reply.started":"2024-03-09T03:36:06.843763Z","shell.execute_reply":"2024-03-09T03:46:27.955381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls Models/fold_0/","metadata":{"execution":{"iopub.status.busy":"2024-03-09T03:47:01.755531Z","iopub.execute_input":"2024-03-09T03:47:01.756230Z","iopub.status.idle":"2024-03-09T03:47:02.749505Z","shell.execute_reply.started":"2024-03-09T03:47:01.756197Z","shell.execute_reply":"2024-03-09T03:47:02.748421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}