{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import argparse\nimport os\nimport sys\nimport time\nimport numpy as np\nimport torch\nfrom torch import nn\nfrom torch.cuda import amp\nfrom torch.utils.data import DataLoader, RandomSampler, SequentialSampler","metadata":{"execution":{"iopub.status.busy":"2022-08-05T18:00:05.950594Z","iopub.execute_input":"2022-08-05T18:00:05.951397Z","iopub.status.idle":"2022-08-05T18:00:07.991141Z","shell.execute_reply.started":"2022-08-05T18:00:05.951319Z","shell.execute_reply":"2022-08-05T18:00:07.990094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os.path\n\n\nclass config:\n    def __init__(self, fold=0):\n        self.root_dir = \".\"\n        self.pretrain_dir = '../input/swin-tiny-small-22k-pretrained'\n        self.is_amp = True\n        self.TRAIN = '../input/hubmap-768x768/train/'\n        self.MASKS = '../input/hubmap-768x768/masks/'\n        self.LABELS = '../input/hubmap-organ-segmentation/train.csv'\n\n        self.fold = fold\n\n        self.out_dir = self.root_dir + f'/models/fold_{self.fold}'\n\n        if not os.path.exists(self.out_dir):\n            os.mkdir(self.out_dir)\n\n        self.initial_checkpoint = None\n\n        self.start_lr = 5e-5  # 0.0001\n        self.batch_size = 8  # 32 #32\n\n        self.checkpoint = dict(\n\n            # configs/_base_/models/upernet_swin.py\n            basic=dict(\n                swin=dict(\n                    embed_dim=96,\n                    depths=[2, 2, 6, 2],\n                    num_heads=[3, 6, 12, 24],\n                    window_size=7,\n                    mlp_ratio=4.,\n                    qkv_bias=True,\n                    qk_scale=None,\n                    drop_rate=0.,\n                    attn_drop_rate=0.,\n                    drop_path_rate=0.3,\n                    ape=False,\n                    patch_norm=True,\n                    out_indices=(0, 1, 2, 3),\n                    use_checkpoint=False\n                ),\n\n            ),\n\n            # configs/swin/upernet_swin_tiny_patch4_window7_512x512_160k_ade20k.py\n            swin_tiny_patch4_window7_224=dict(\n                checkpoint=self.pretrain_dir + '/swin_tiny_patch4_window7_224_22k.pth',\n\n                swin=dict(\n                    embed_dim=96,\n                    depths=[2, 2, 6, 2],\n                    num_heads=[3, 6, 12, 24],\n                    window_size=7,\n                    ape=False,\n                    drop_path_rate=0.3,\n                    patch_norm=True,\n                    use_checkpoint=False,\n                ),\n                upernet=dict(\n                    in_channels=[96, 192, 384, 768],\n                ),\n            ),\n\n            # /configs/swin/upernet_swin_small_patch4_window7_512x512_160k_ade20k.py\n            swin_small_patch4_window7_224_22k=dict(\n                checkpoint=self.pretrain_dir + '/swin_small_patch4_window7_224_22k.pth',\n\n                swin=dict(\n                    embed_dim=96,\n                    depths=[2, 2, 18, 2],\n                    num_heads=[3, 6, 12, 24],\n                    window_size=7,\n                    ape=False,\n                    drop_path_rate=0.3,\n                    patch_norm=True,\n                    use_checkpoint=False\n                ),\n                upernet=dict(\n                    in_channels=[96, 192, 384, 768],\n                ),\n            ),\n        )\n","metadata":{"execution":{"iopub.status.busy":"2022-08-05T18:00:07.993401Z","iopub.execute_input":"2022-08-05T18:00:07.994523Z","iopub.status.idle":"2022-08-05T18:00:08.028637Z","shell.execute_reply.started":"2022-08-05T18:00:07.994481Z","shell.execute_reply":"2022-08-05T18:00:08.027719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\n\nimport cv2\nimport numpy as np\n\n\ndef do_random_flip(image, mask):\n    if np.random.rand() > 0.5:\n        image = cv2.flip(image, 0)\n        mask = cv2.flip(mask, 0)\n    if np.random.rand() > 0.5:\n        image = cv2.flip(image, 1)\n        mask = cv2.flip(mask, 1)\n    if np.random.rand() > 0.5:\n        image = image.transpose(1, 0, 2)\n        mask = mask.transpose(1, 0)\n\n    image = np.ascontiguousarray(image)\n    mask = np.ascontiguousarray(mask)\n    return image, mask\n\n\ndef do_random_rot90(image, mask):\n    r = np.random.choice([\n        0,\n        cv2.ROTATE_90_CLOCKWISE,\n        cv2.ROTATE_90_COUNTERCLOCKWISE,\n        cv2.ROTATE_180,\n    ])\n    if r == 0:\n        return image, mask\n    else:\n        image = cv2.rotate(image, r)\n        mask = cv2.rotate(mask, r)\n        return image, mask\n\n\ndef do_random_contast(image, mask, mag=0.3):\n    alpha = 1 + random.uniform(-1, 1) * mag\n    image = image * alpha\n    image = np.clip(image, 0, 1)\n    return image, mask\n\n\ndef do_random_hsv(image, mask, mag=[0.15, 0.25, 0.25]):\n    image = (image * 255).astype(np.uint8)\n    hsv = cv2.cvtColor(image, cv2.COLOR_BGR2HSV)\n\n    h = hsv[:, :, 0].astype(np.float32)  # hue\n    s = hsv[:, :, 1].astype(np.float32)  # saturation\n    v = hsv[:, :, 2].astype(np.float32)  # value\n    h = (h * (1 + random.uniform(-1, 1) * mag[0])) % 180\n    s = s * (1 + random.uniform(-1, 1) * mag[1])\n    v = v * (1 + random.uniform(-1, 1) * mag[2])\n\n    hsv[:, :, 0] = np.clip(h, 0, 180).astype(np.uint8)\n    hsv[:, :, 1] = np.clip(s, 0, 255).astype(np.uint8)\n    hsv[:, :, 2] = np.clip(v, 0, 255).astype(np.uint8)\n    image = cv2.cvtColor(hsv, cv2.COLOR_HSV2BGR)\n    image = image.astype(np.float32) / 255\n    return image, mask\n\n\ndef do_random_noise(image, mask, mag=0.1):\n    height, width = image.shape[:2]\n    noise = np.random.uniform(-1, 1, (height, width, 1)) * mag\n    image = image + noise\n    image = np.clip(image, 0, 1)\n    return image, mask\n\n\ndef do_random_rotate_scale(image, mask, angle=30, scale=[0.8, 1.2]):\n    angle = np.random.uniform(-angle, angle)\n    scale = np.random.uniform(*scale) if scale is not None else 1\n\n    height, width = image.shape[:2]\n    center = (height // 2, width // 2)\n\n    transform = cv2.getRotationMatrix2D(center, angle, scale)\n    image = cv2.warpAffine(image, transform, (width, height), flags=cv2.INTER_LINEAR,\n                           borderMode=cv2.BORDER_CONSTANT, borderValue=(0, 0, 0))\n    mask = cv2.warpAffine(mask, transform, (width, height), flags=cv2.INTER_LINEAR,\n                          borderMode=cv2.BORDER_CONSTANT, borderValue=0)\n    return image, mask","metadata":{"execution":{"iopub.status.busy":"2022-08-05T18:00:08.033128Z","iopub.execute_input":"2022-08-05T18:00:08.035836Z","iopub.status.idle":"2022-08-05T18:00:08.205840Z","shell.execute_reply.started":"2022-08-05T18:00:08.035801Z","shell.execute_reply":"2022-08-05T18:00:08.204798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compute_dice_score(probability, mask):\n    N = len(probability)\n\n    probability = probability > 0.5\n    mask = mask > 0.5\n\n    p = probability.reshape(N, -1)\n    t = mask.reshape(N, -1)\n\n    uion = p.sum(-1) + t.sum(-1)\n    overlap = (p*t).sum(-1)\n    dice = 2*overlap/(uion+0.0001)\n    return dice","metadata":{"execution":{"iopub.status.busy":"2022-08-05T18:00:08.208503Z","iopub.execute_input":"2022-08-05T18:00:08.208851Z","iopub.status.idle":"2022-08-05T18:00:08.216141Z","shell.execute_reply.started":"2022-08-05T18:00:08.208816Z","shell.execute_reply":"2022-08-05T18:00:08.215143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nimport time\nimport warnings\nfrom itertools import repeat\nimport collections.abc\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom sklearn.model_selection import KFold\nfrom torch import nn\nfrom torch.cuda import amp\n\n\ndef image_to_tensor(image, mode='bgr'):  # image mode\n    if mode == 'bgr':\n        image = image[:, :, ::-1]\n    x = image\n    x = x.transpose(2, 0, 1)\n    x = np.ascontiguousarray(x)\n    x = torch.tensor(x, dtype=torch.float)\n    return x\n\n\ndef mask_to_tensor(mask):\n    x = mask\n    x = torch.tensor(x, dtype=torch.float)\n    return x\n\n\ntensor_list = ['mask', 'image', 'organ']\n\n\ndef null_collate(batch):\n    d = {}\n    key = batch[0].keys()\n    for k in key:\n        v = [b[k] for b in batch]\n        if k in tensor_list:\n            v = torch.stack(v)\n        d[k] = v\n\n    d['mask'] = d['mask'].unsqueeze(1)\n    d['organ'] = d['organ'].reshape(-1)\n    return d\n\n\ndef _ntuple(n):\n    def parse(x):\n        if isinstance(x, collections.abc.Iterable):\n            return x\n        return tuple(repeat(x, n))\n\n    return parse\n\n\nto_2tuple = _ntuple(2)\n\n\ndef _no_grad_trunc_normal_(tensor, mean, std, a, b):\n    # Cut & paste from PyTorch official master until it's in a few official releases - RW\n    # Method based on https://people.sc.fsu.edu/~jburkardt/presentations/truncated_normal.pdf\n    def norm_cdf(x):\n        # Computes standard normal cumulative distribution function\n        return (1. + math.erf(x / math.sqrt(2.))) / 2.\n\n    if (mean < a - 2 * std) or (mean > b + 2 * std):\n        warnings.warn(\"mean is more than 2 std from [a, b] in nn.init.trunc_normal_. \"\n                      \"The distribution of values may be incorrect.\",\n                      stacklevel=2)\n\n\ndef trunc_normal_(tensor, mean=0., std=1., a=-2., b=2.):\n    return _no_grad_trunc_normal_(tensor, mean, std, a, b)\n\n\ndef drop_path(x, drop_prob: float = 0., training: bool = False, scale_by_keep: bool = True):\n    \"\"\"Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).\n    This is the same as the DropConnect impl I created for EfficientNet, etc networks, however,\n    the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...\n    See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for\n    changing the layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use\n    'survival rate' as the argument.\n    \"\"\"\n    if drop_prob == 0. or not training:\n        return x\n    keep_prob = 1 - drop_prob\n    shape = (x.shape[0],) + (1,) * (x.ndim - 1)  # work with diff dim tensors, not just 2D ConvNets\n    random_tensor = x.new_empty(shape).bernoulli_(keep_prob)\n    if keep_prob > 0.0 and scale_by_keep:\n        random_tensor.div_(keep_prob)\n    return x * random_tensor\n\n\nclass DropPath(nn.Module):\n    \"\"\"Drop paths (Stochastic Depth) per sample  (when applied in main path of residual blocks).\n    \"\"\"\n\n    def __init__(self, drop_prob: float = 0., scale_by_keep: bool = True):\n        super(DropPath, self).__init__()\n        self.drop_prob = drop_prob\n        self.scale_by_keep = scale_by_keep\n\n    def forward(self, x):\n        return drop_path(x, self.drop_prob, self.training, self.scale_by_keep)\n\n    def extra_repr(self):\n        return f'drop_prob={round(self.drop_prob, 3):0.3f}'\n\n\nclass RGB(nn.Module):\n    IMAGE_RGB_MEAN = [0.485, 0.456, 0.406]  # [0.5, 0.5, 0.5]\n    IMAGE_RGB_STD = [0.229, 0.224, 0.225]  # [0.5, 0.5, 0.5]\n\n    def __init__(self, ):\n        super(RGB, self).__init__()\n        self.register_buffer('mean', torch.zeros(1, 3, 1, 1))\n        self.register_buffer('std', torch.ones(1, 3, 1, 1))\n        self.mean.data = torch.FloatTensor(self.IMAGE_RGB_MEAN).view(self.mean.shape)\n        self.std.data = torch.FloatTensor(self.IMAGE_RGB_STD).view(self.std.shape)\n\n    def forward(self, x):\n        x = (x - self.mean) / self.std\n        return x\n\n\ndef message(batch_loss=0, train_loss=0, iteration=0, iter_save=0, rate=0, epoch=0, valid_loss=None,\n            start_timer=None,\n            mode='print'):\n    asterisk = ' '\n    if mode == 'print':\n        loss = batch_loss\n    if mode == 'log':\n        loss = train_loss\n        if (iteration % iter_save == 0): asterisk = '*'\n\n    text = \\\n        ('%0.2e   %08d%s %6.2f | ' % (rate, iteration, asterisk, epoch,)).replace('e-0', 'e-').replace('e+0', 'e+') + \\\n        '%4.3f  %4.3f  %4.4f  %4.3f   | ' % (*valid_loss,) + \\\n        '%4.3f  %4.3f   | ' % (*loss,) + \\\n        '%s' % ((time.time() - start_timer))\n\n    return text\n\n\ndef valid_augment5(image, mask, organ):\n    # image, mask  = do_crop(image, mask, image_size, xy=(None,None))\n    return image, mask\n\n\ndef train_augment5b(image, mask, organ):\n    image, mask = do_random_flip(image, mask)\n    image, mask = do_random_rot90(image, mask)\n\n    for fn in np.random.choice([\n        lambda image, mask: (image, mask),\n        lambda image, mask: do_random_noise(image, mask, mag=0.1),\n        lambda image, mask: do_random_contast(image, mask, mag=0.40),\n        lambda image, mask: do_random_hsv(image, mask, mag=[0.40, 0.40, 0])\n    ], 2): image, mask = fn(image, mask)\n\n    for fn in np.random.choice([\n        lambda image, mask: (image, mask),\n        lambda image, mask: do_random_rotate_scale(image, mask, angle=45, scale=[0.50, 2.0]),\n    ], 1): image, mask = fn(image, mask)\n\n    return image, mask\n\n\ndef make_fold(config, fold=0):\n    df = pd.read_csv(config.root_dir + '/../input/hubmap-organ-segmentation/train.csv')\n\n    num_fold = 5\n    skf = KFold(n_splits=num_fold, shuffle=True, random_state=42)\n\n    df.loc[:, 'fold'] = -1\n    for f, (t_idx, v_idx) in enumerate(skf.split(X=df['id'], y=df['organ'])):\n        df.iloc[v_idx, -1] = f\n\n    # #check\n    # if 0:\n    #     for f in range(num_fold):\n    #         train_df=df[df.fold!=f].reset_index(drop=True)\n    #         valid_df=df[df.fold==f].reset_index(drop=True)\n    #\n    #         print('fold %d'%f)\n    #         t = train_df.organ.value_counts().to_dict()\n    #         v = valid_df.organ.value_counts().to_dict()\n    #         for k in ['kidney', 'prostate', 'largeintestine', 'spleen', 'lung']:\n    #             print('%32s %3d (%0.3f)  %3d (%0.3f)'%(k,t.get(k,0),t.get(k,0)/len(train_df),v.get(k,0),v.get(k,0)/len(valid_df)))\n    #\n    #         print('')\n    #         zz=0\n\n    train_df = df[df.fold != fold].reset_index(drop=True)\n    valid_df = df[df.fold == fold].reset_index(drop=True)\n    return train_df, valid_df\n\n\ndef validate(net, valid_loader, config):\n    valid_num = 0\n    valid_probability = []\n    valid_mask = []\n    valid_loss = 0\n\n    net = net.eval()\n    start_timer = time.time()\n    for t, batch in enumerate(valid_loader):\n\n        net.output_type = ['loss', 'inference']\n        with torch.no_grad():\n            with amp.autocast(enabled=config.is_amp):\n                batch_size = len(batch['index'])\n                batch['image'] = batch['image'].cuda()\n                batch['mask'] = batch['mask'].cuda()\n                batch['organ'] = batch['organ'].cuda()\n\n                output = net(batch)\n                loss0 = output['bce_loss'].mean()\n\n        valid_probability.append(output['probability'].data.cpu().numpy())\n        valid_mask.append(batch['mask'].data.cpu().numpy())\n        valid_num += batch_size\n        valid_loss += batch_size * loss0.item()\n\n        # debug\n        # if 0:\n        #     organ = batch['organ'].data.cpu().numpy()\n        #     image = batch['image']\n        #     mask = batch['mask']\n        #     probability = output['probability']\n        #\n        #     for b in range(batch_size):\n        #         m = tensor_to_image(image[b])\n        #         t = tensor_to_mask(mask[b, 0])\n        #         p = tensor_to_mask(probability[b, 0])\n        #         overlay = result_to_overlay(m, t, p)\n        #\n        #         text = label_to_organ[organ[b]]\n        #         draw_shadow_text(overlay, text, (5, 15), 0.7, (1, 1, 1), 1)\n        #\n        #         image_show_norm('overlay', overlay, min=0, max=1, resize=1)\n        #         cv2.waitKey(0)\n\n        print('\\r %8d / %d  %s' % (valid_num, len(valid_loader.dataset), (time.time() - start_timer)), end='',\n              flush=True)\n\n    assert (valid_num == len(valid_loader.dataset))\n\n    probability = np.concatenate(valid_probability)\n    mask = np.concatenate(valid_mask)\n\n    loss = valid_loss / valid_num\n\n    dice = compute_dice_score(probability, mask)\n    dice = dice.mean()\n\n    return [dice, loss, 0, 0]","metadata":{"execution":{"iopub.status.busy":"2022-08-05T18:00:08.217871Z","iopub.execute_input":"2022-08-05T18:00:08.218614Z","iopub.status.idle":"2022-08-05T18:00:09.147561Z","shell.execute_reply.started":"2022-08-05T18:00:08.218574Z","shell.execute_reply":"2022-08-05T18:00:09.146125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset\n\nimage_size = 768\n\n\nclass HapDataset(Dataset):\n    def __init__(self, df, config, augment=None):\n\n        self.df = df\n        self.augment = augment\n        self.length = len(self.df)\n        ids = pd.read_csv(config.LABELS).id.astype(str).values\n        self.config = config\n        self.fnames = [fname for fname in os.listdir(config.TRAIN) if fname.split('_')[0] in ids]\n        self.organ_to_label = {'kidney': 0,\n                               'prostate': 1,\n                               'largeintestine': 2,\n                               'spleen': 3,\n                               'lung': 4}\n\n    def __str__(self):\n        string = ''\n        string += '\\tlen = %d\\n' % len(self)\n\n        d = self.df.organ.value_counts().to_dict()\n        for k in ['kidney', 'prostate', 'largeintestine', 'spleen', 'lung']:\n            string += '%24s %3d (%0.3f) \\n' % (k, d.get(k, 0), d.get(k, 0) / len(self.df))\n        return string\n\n    def __len__(self):\n        return self.length\n\n    def __getitem__(self, index):\n        fname = self.fnames[index]\n        d = self.df.iloc[index]\n        organ = self.organ_to_label[d.organ]\n\n        image = cv2.cvtColor(cv2.imread(os.path.join(self.config.TRAIN, fname)), cv2.COLOR_BGR2RGB)\n        mask = cv2.imread(os.path.join(self.config.MASKS, fname), cv2.IMREAD_GRAYSCALE)\n\n        image = image.astype(np.float32) / 255\n\n        s = d.pixel_size / 0.4 * (image_size / 3000)\n\n        if self.augment is not None:\n            image, mask = self.augment(image, mask, organ)\n\n        r = {}\n        r['index'] = index\n        r['id'] = fname\n        r['organ'] = torch.tensor([organ], dtype=torch.long)\n        r['image'] = image_to_tensor(image)\n        r['mask'] = mask_to_tensor(mask > 0.5)\n        return r","metadata":{"execution":{"iopub.status.busy":"2022-08-05T18:00:09.149240Z","iopub.execute_input":"2022-08-05T18:00:09.149890Z","iopub.status.idle":"2022-08-05T18:00:09.189829Z","shell.execute_reply.started":"2022-08-05T18:00:09.149842Z","shell.execute_reply":"2022-08-05T18:00:09.188683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\n\nclass PatchEmbed(nn.Module):\n    r\"\"\" Image to Patch Embedding\n\n    Args:\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,\n                 patch_size=4,\n                 in_chans=3,\n                 embed_dim=96,\n                 norm_layer=None\n                 ):\n        super().__init__()\n        patch_size = to_2tuple(patch_size)\n        self.patch_size = patch_size\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\n        # padding\n        if W % self.patch_size[1] != 0:\n            x = F.pad(x, (0, self.patch_size[1] - W % self.patch_size[1]))\n        if H % self.patch_size[0] != 0:\n            x = F.pad(x, (0, 0, 0, self.patch_size[0] - H % self.patch_size[0]))\n\n        x = self.proj(x)  # B C Wh Ww\n        if self.norm is not None:\n            Wh, Ww = x.size(2), x.size(3)\n            x = x.flatten(2).transpose(1, 2)\n            x = self.norm(x)\n            x = x.transpose(1, 2).view(-1, self.embed_dim, Wh, Ww)\n\n        return x\n\n\nclass PatchMerging(nn.Module):\n    r\"\"\" Patch Merging Layer.\n\n    Args:\n        dim (int): Number of input channels.\n        norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm\n    \"\"\"\n\n    def __init__(self, dim, norm_layer=nn.LayerNorm):\n        super().__init__()\n        self.dim = dim\n        self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)\n        self.norm = norm_layer(4 * dim)\n\n    def forward(self, x, H, W):\n        \"\"\"\n        Args:\n            x: Input feature, tensor size (B, H*W, C).\n            H, W: Spatial resolution of the input feature.\n        \"\"\"\n\n        B, L, C = x.shape\n        assert L == H * W, \"input feature has wrong size\"\n\n        x = x.view(B, H, W, C)\n        # padding\n        pad_input = (H % 2 == 1) or (W % 2 == 1)\n        if pad_input:\n            x = F.pad(x, (0, 0, 0, W % 2, 0, H % 2))\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.norm(x)\n        x = self.reduction(x)\n\n        return x\n\n\nclass BasicLayer(nn.Module):\n    \"\"\" A basic Swin Transformer layer for one stage.\n\n    Args:\n        dim (int): Number of input channels.\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        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.\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        fused_window_process (bool, optional): If True, use one kernel to fused window shift & window partition for acceleration, similar for the reversed part. Default: False\n    \"\"\"\n\n    def __init__(self,\n                 dim,\n                 depth,\n                 num_heads,\n                 window_size,\n                 mlp_ratio=4.,\n                 qkv_bias=True,\n                 qk_scale=None,\n                 drop=0.,\n                 attn_drop=0.,\n                 drop_path=0.,\n                 norm_layer=nn.LayerNorm,\n                 downsample=None,\n                 # use_checkpoint=False,\n                 ):\n        super().__init__()\n        self.window_size = window_size\n        self.shift_size = window_size // 2\n        self.depth = depth\n\n        self.blocks = nn.ModuleList([\n            SwinTransformerBlock(\n                dim=dim,\n                num_heads=num_heads,\n                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                qk_scale=qk_scale,\n                drop=drop,\n                attn_drop=attn_drop,\n                drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,\n                norm_layer=norm_layer,\n            )\n            for i in range(depth)\n        ])\n        # patch merging layer\n        if downsample is not None:\n            self.downsample = downsample(dim=dim, norm_layer=norm_layer)\n        else:\n            self.downsample = None\n\n    def forward(self, x, H, W):\n        \"\"\"\n        Args:\n            x: Input feature, tensor size (B, H*W, C).\n            H, W: Spatial resolution of the input feature.\n        \"\"\"\n\n        # calculate attention mask for SW-MSA ----\n        Hp = int(np.ceil(H / self.window_size)) * self.window_size\n        Wp = int(np.ceil(W / self.window_size)) * self.window_size\n        img_mask = torch.zeros((1, Hp, Wp, 1), device=x.device)  # 1 Hp Wp 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        # ------\n\n        for blk in self.blocks:\n            x = blk(x, H, W, attn_mask)\n\n        if self.downsample is not None:\n            x_down = self.downsample(x, H, W)\n            Wh, Ww = (H + 1) // 2, (W + 1) // 2\n            return x, H, W, x_down, Wh, Ww\n        else:\n            return x, H, W, x, H, W\n\n\nclass SwinTransformerBlock(nn.Module):\n    r\"\"\" Swin Transformer Block.\n\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        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.\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        fused_window_process (bool, optional): If True, use one kernel to fused window shift & window partition for acceleration, similar for the reversed part. Default: False\n    \"\"\"\n\n    def __init__(self,\n                 dim,\n                 num_heads,\n                 window_size=7,\n                 shift_size=0,\n                 mlp_ratio=4.,\n                 qkv_bias=True,\n                 qk_scale=None,\n                 drop=0.,\n                 attn_drop=0.,\n                 drop_path=0.,\n                 act_layer=nn.GELU,\n                 norm_layer=nn.LayerNorm,\n                 ):\n        super().__init__()\n        self.dim = dim\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        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,\n            window_size=to_2tuple(self.window_size),\n            num_heads=num_heads,\n            qkv_bias=qkv_bias,\n            qk_scale=qk_scale,\n            attn_drop=attn_drop,\n            proj_drop=drop,\n        )\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\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    def forward(self, x, H, W, mask_matrix):\n\n        B, L, C = x.shape\n        assert L == H * W, \"input feature has wrong size\"\n\n        shortcut = x\n        x = self.norm1(x)\n        x = x.view(B, H, W, C)\n\n        # pad feature maps to multiples of window size\n        pad_l = pad_t = 0\n        pad_r = (self.window_size - W % self.window_size) % self.window_size\n        pad_b = (self.window_size - H % self.window_size) % self.window_size\n        x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b))\n        _, Hp, Wp, _ = x.shape\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            attn_mask = mask_matrix\n        else:\n            shifted_x = x\n            attn_mask = None\n\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        attn_windows = self.attn(x_windows, mask=attn_mask)  # nW*B, window_size*window_size, C\n        attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)\n\n        # reverse cyclic shift ---\n        shifted_x = window_reverse(attn_windows, self.window_size, Hp, Wp)  # B H' W' C\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\n        if pad_r > 0 or pad_b > 0:\n            x = x[:, :H, :W, :].contiguous()\n        x = x.view(B, H * W, C)\n\n        # FFN\n        x = shortcut + self.drop_path(x)\n        x = x + self.drop_path(self.mlp(self.norm2(x)))\n        return x\n\n    def extra_repr(self) -> str:\n        return f\"dim={self.dim}, num_heads={self.num_heads}, \" \\\n               f\"window_size={self.window_size}, shift_size={self.shift_size}, mlp_ratio={self.mlp_ratio}\"\n\n\nclass SwinTransformerV1(nn.Module):\n    def __init__(self,\n                 pretrain_img_size=224,\n                 patch_size=4,\n                 in_chans=3,\n                 embed_dim=96,\n                 depths=[2, 2, 6, 2],\n                 num_heads=[3, 6, 12, 24],\n                 window_size=7,\n                 mlp_ratio=4.,\n                 qkv_bias=True,\n                 qk_scale=None,\n                 drop_rate=0.,\n                 attn_drop_rate=0.,\n                 drop_path_rate=0.1,\n                 norm_layer=nn.LayerNorm,\n                 patch_norm=True,\n                 out_norm=nn.Identity,  # use nn.Identity, nn.BatchNorm2d, LayerNorm2d\n                 **kwargs\n                 ):\n        super().__init__()\n        self.pretrain_img_size = pretrain_img_size\n        self.num_layers = len(depths)\n        self.embed_dim = embed_dim\n        self.mlp_ratio = mlp_ratio\n\n        self.patch_embed = PatchEmbed(\n            patch_size=patch_size,\n            in_chans=in_chans,\n            embed_dim=embed_dim,\n            norm_layer=norm_layer if patch_norm else None\n        )\n        self.pos_drop = nn.Dropout(p=drop_rate)\n\n        # stochastic depth\n        dpr = np.linspace(0, drop_path_rate, sum(depths)).tolist()  # stochastic depth decay rule\n\n        # build layers\n        self.layers = nn.ModuleList()\n        for i in range(self.num_layers):\n            layer = BasicLayer(\n                dim=int(embed_dim * 2 ** i),\n                depth=depths[i],\n                num_heads=num_heads[i],\n                window_size=window_size,\n                mlp_ratio=self.mlp_ratio,\n                qkv_bias=qkv_bias, qk_scale=qk_scale,\n                drop=drop_rate,\n                attn_drop=attn_drop_rate,\n                drop_path=dpr[sum(depths[:i]):sum(depths[:i + 1])],\n                norm_layer=norm_layer,\n                downsample=PatchMerging if (i < self.num_layers - 1) else None,\n            )\n            self.layers.append(layer)\n\n        # ---\n        # add a norm layer for each output\n        self.out_norm = nn.ModuleList(\n            [out_norm(int(embed_dim * 2 ** i)) for i in range(self.num_layers)]\n        )\n\n        # ---\n        self.apply(self._init_weights)\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    def forward(self, x):\n        x = self.patch_embed(x)\n        Wh, Ww = x.size(2), x.size(3)\n\n        # positional encode?\n        x = x.flatten(2).transpose(1, 2)\n        x = self.pos_drop(x)\n\n        outs = []\n        for i in range(self.num_layers):\n            x_out, H, W, x, Wh, Ww = self.layers[i](x, Wh, Ww)\n            out = x_out.view(-1, H, W, int(self.embed_dim * 2 ** i)).permute(0, 3, 1, 2).contiguous()\n            out = self.out_norm[i](out)\n            outs.append(out)\n\n        return outs\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\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        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set\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    \"\"\"\n\n    def __init__(self, dim, window_size, num_heads, qkv_bias=True, qk_scale=None, attn_drop=0., proj_drop=0.):\n        super().__init__()\n        self.dim = dim\n        self.window_size = window_size  # Wh, Ww\n        self.num_heads = num_heads\n        head_dim = dim // num_heads\n        self.scale = qk_scale or head_dim ** (-0.5)\n\n        # define a parameter table of relative position bias\n        self.relative_position_bias_table = nn.Parameter(\n            torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads))  # 2*Wh-1 * 2*Ww-1, nH\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=qkv_bias)\n        self.proj = nn.Linear(dim, dim)\n        self.softmax = nn.Softmax(dim=-1)\n\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj_drop = nn.Dropout(proj_drop)\n\n        trunc_normal_(self.relative_position_bias_table, std=.02)\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\n        B_, N, C = x.shape\n        qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads).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        q = q * self.scale\n        attn = (q @ k.transpose(-2, -1))\n\n        relative_position_bias = \\\n            self.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], self.num_heads)\n        # Wh*Ww,Wh*Ww,nH\n        relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous()  # nH, Wh*Ww, Wh*Ww\n\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        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}, num_heads={self.num_heads}'\n\n\ndef window_partition(x, window_size):\n    \"\"\"\n    Args:\n        x: (B, H, W, C)\n        window_size (int): window size\n\n    Returns:\n        windows: (num_windows*B, window_size, window_size, C)\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\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\ndef conv3x3_bn_relu(in_planes, out_planes, stride=1):\n    \"3x3 convolution + BN + relu\"\n    return nn.Sequential(\n        nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride, padding=1, bias=False),\n        nn.BatchNorm2d(out_planes),\n        nn.ReLU(inplace=True),\n    )\n\n\nclass UPerDecoder(nn.Module):\n    def __init__(self,\n                 in_dim=[256, 512, 1024, 2048],\n                 ppm_pool_scale=[1, 2, 3, 6],\n                 ppm_dim=512,\n                 fpn_out_dim=256\n                 ):\n        super(UPerDecoder, self).__init__()\n\n        # PPM ----\n        dim = in_dim[-1]\n        ppm_pooling = []\n        ppm_conv = []\n\n        for scale in ppm_pool_scale:\n            ppm_pooling.append(\n                nn.AdaptiveAvgPool2d(scale)\n            )\n            ppm_conv.append(\n                nn.Sequential(\n                    nn.Conv2d(dim, ppm_dim, kernel_size=1, bias=False),\n                    nn.BatchNorm2d(ppm_dim),\n                    nn.ReLU(inplace=True)\n                )\n            )\n        self.ppm_pooling = nn.ModuleList(ppm_pooling)\n        self.ppm_conv = nn.ModuleList(ppm_conv)\n        self.ppm_out = conv3x3_bn_relu(dim + len(ppm_pool_scale) * ppm_dim, fpn_out_dim, 1)\n\n        # FPN ----\n        fpn_in = []\n        for i in range(0, len(in_dim) - 1):  # skip the top layer\n            fpn_in.append(\n                nn.Sequential(\n                    nn.Conv2d(in_dim[i], fpn_out_dim, kernel_size=1, bias=False),\n                    nn.BatchNorm2d(fpn_out_dim),\n                    nn.ReLU(inplace=True)\n                )\n            )\n        self.fpn_in = nn.ModuleList(fpn_in)\n\n        fpn_out = []\n        for i in range(len(in_dim) - 1):  # skip the top layer\n            fpn_out.append(\n                conv3x3_bn_relu(fpn_out_dim, fpn_out_dim, 1),\n            )\n        self.fpn_out = nn.ModuleList(fpn_out)\n\n        self.fpn_fuse = nn.Sequential(\n            conv3x3_bn_relu(len(in_dim) * fpn_out_dim, fpn_out_dim, 1),\n        )\n\n    def forward(self, feature):\n        f = feature[-1]\n        pool_shape = f.shape[2:]\n\n        ppm_out = [f]\n        for pool, conv in zip(self.ppm_pooling, self.ppm_conv):\n            p = pool(f)\n            p = F.interpolate(p, size=pool_shape, mode='bilinear', align_corners=False)\n            p = conv(p)\n            ppm_out.append(p)\n        ppm_out = torch.cat(ppm_out, 1)\n        down = self.ppm_out(ppm_out)\n\n        fpn_out = [down]\n        for i in reversed(range(len(feature) - 1)):\n            lateral = feature[i]\n            lateral = self.fpn_in[i](lateral)  # lateral branch\n            down = F.interpolate(down, size=lateral.shape[2:], mode='bilinear', align_corners=False)  # top-down branch\n            down = down + lateral\n            fpn_out.append(self.fpn_out[i](down))\n\n        fpn_out.reverse()  # [P2 - P5]\n        fusion_shape = fpn_out[0].shape[2:]\n        fusion = [fpn_out[0]]\n        for i in range(1, len(fpn_out)):\n            fusion.append(\n                F.interpolate(fpn_out[i], fusion_shape, mode='bilinear', align_corners=False)\n            )\n        x = self.fpn_fuse(torch.cat(fusion, 1))\n\n        return x, fusion\n\n\nclass LayerNorm2d(nn.Module):\n    def __init__(self, dim, eps=1e-6):\n        super().__init__()\n        self.weight = nn.Parameter(torch.ones(dim))\n        self.bias = nn.Parameter(torch.zeros(dim))\n        self.eps = eps\n\n    def forward(self, x):\n        u = x.mean(1, keepdim=True)\n        s = (x - u).pow(2).mean(1, keepdim=True)\n        x = (x - u) / torch.sqrt(s + self.eps)\n        x = self.weight[:, None, None] * x + self.bias[:, None, None]\n        return x\n\n\ndef criterion_aux_loss(logit, mask):\n    mask = F.interpolate(mask, size=logit.shape[-2:], mode='nearest')\n    loss = F.binary_cross_entropy_with_logits(logit, mask)\n    return loss\n\n\nclass Net(nn.Module):\n\n    def load_pretrain(self):\n\n        checkpoint = self.config.checkpoint[self.arch]['checkpoint']\n        print('loading %s ...' % checkpoint)\n        checkpoint = torch.load(checkpoint, map_location=lambda storage, loc: storage)['model']\n        # if 0:\n        #     skip = ['relative_coords_table','relative_position_index']\n        #     filtered={}\n        #     for k,v in checkpoint.items():\n        #         if any([s in k for s in skip ]): continue\n        #         filtered[k]=v\n        #     checkpoint = filtered\n        print(self.encoder.load_state_dict(checkpoint, strict=False))  # True\n\n    def __init__(self, config):\n        super(Net, self).__init__()\n        self.config = config\n        self.output_type = ['inference', 'loss']\n\n        self.rgb = RGB()\n        self.arch = 'swin_tiny_patch4_window7_224'\n\n        self.encoder = SwinTransformerV1(\n            **{**(config.checkpoint['basic']['swin']), **(config.checkpoint[self.arch]['swin']),\n               **{'out_norm': LayerNorm2d}}\n        )\n        encoder_dim = config.checkpoint[self.arch]['upernet']['in_channels']\n        # [96, 192, 384, 768]\n\n        self.decoder = UPerDecoder(\n            in_dim=encoder_dim,\n            ppm_pool_scale=[1, 2, 3, 6],\n            ppm_dim=512,\n            fpn_out_dim=256\n        )\n\n        self.logit = nn.Sequential(\n            nn.Conv2d(256, 1, kernel_size=1)\n        )\n        self.aux = nn.ModuleList([\n            nn.Conv2d(256, 1, kernel_size=1, padding=0) for i in range(4)\n        ])\n\n    def forward(self, batch):\n        x = batch['image']\n        B, C, H, W = x.shape\n        x = self.rgb(x)\n        encoder = self.encoder(x)\n        last, decoder = self.decoder(encoder)\n        logit = self.logit(last)\n        logit = F.interpolate(logit, size=None, scale_factor=4, mode='bilinear', align_corners=False)\n\n        output = {}\n        if 'loss' in self.output_type:\n            output['bce_loss'] = F.binary_cross_entropy_with_logits(logit, batch['mask'])\n            for i in range(4):\n                output['aux%d_loss' % i] = criterion_aux_loss(self.aux[i](decoder[i]), batch['mask'])\n\n        if 'inference' in self.output_type:\n            output['probability'] = torch.sigmoid(logit)\n\n        return output\n\n\ndef run_check_net(config):\n    batch_size = 2\n    image_size = 512\n\n    # ---\n    batch = {\n        'image': torch.from_numpy(np.random.uniform(-1, 1, (batch_size, 3, image_size, image_size))).float(),\n        'mask': torch.from_numpy(np.random.choice(2, (batch_size, 1, image_size, image_size))).float(),\n        'organ': torch.from_numpy(np.random.choice(5, (batch_size))).long(),\n    }\n    batch = {k: v.cuda() for k, v in batch.items()}\n\n    net = Net(config).cuda()\n    net.load_pretrain()\n\n    with torch.no_grad():\n        with torch.cuda.amp.autocast(enabled=True):\n            output = net(batch)\n\n    print('batch')\n    for k, v in batch.items():\n        print('%32s :' % k, v.shape)\n\n    print('output')\n    for k, v in output.items():\n        if 'loss' not in k:\n            print('%32s :' % k, v.shape)\n    for k, v in output.items():\n        if 'loss' in k:\n            print('%32s :' % k, v.item())\n\n\nclass Mlp(nn.Module):\n    \"\"\" Multilayer perceptron.\"\"\"\n\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","metadata":{"execution":{"iopub.status.busy":"2022-08-05T18:00:09.192816Z","iopub.execute_input":"2022-08-05T18:00:09.193846Z","iopub.status.idle":"2022-08-05T18:00:09.474075Z","shell.execute_reply.started":"2022-08-05T18:00:09.193799Z","shell.execute_reply":"2022-08-05T18:00:09.472948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def freeze_bn(net):\n    for m in net.modules():\n        if isinstance(m, nn.BatchNorm2d):\n            m.eval()\n            m.weight.requires_grad = False\n            m.bias.requires_grad = False\n\n\ndef get_learning_rate(optimizer):\n    return optimizer.param_groups[0]['lr']","metadata":{"execution":{"iopub.status.busy":"2022-08-05T18:00:09.475569Z","iopub.execute_input":"2022-08-05T18:00:09.475903Z","iopub.status.idle":"2022-08-05T18:00:09.485162Z","shell.execute_reply.started":"2022-08-05T18:00:09.475870Z","shell.execute_reply":"2022-08-05T18:00:09.482954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main(config):\n\n    print('\\n--- [START %s] %s\\n\\n' % ('Swin', '-' * 64))\n    print('\\n')\n\n    ## dataset ----------------------------------------\n    print('** dataset setting **\\n')\n\n    train_df, valid_df = make_fold(config, config.fold)\n\n    train_dataset = HapDataset(train_df, config, train_augment5b)\n    valid_dataset = HapDataset(valid_df, config, valid_augment5)\n\n    train_loader = DataLoader(\n        train_dataset,\n        sampler=RandomSampler(train_dataset),\n        batch_size=config.batch_size,\n        drop_last=True,\n        num_workers=8,\n        pin_memory=False,\n        worker_init_fn=lambda id: np.random.seed(torch.initial_seed() // 2 ** 32 + id),\n        collate_fn=null_collate,\n    )\n\n    valid_loader = DataLoader(\n        valid_dataset,\n        sampler=SequentialSampler(valid_dataset),\n        batch_size=8,\n        drop_last=False,\n        num_workers=4,\n        pin_memory=False,\n        collate_fn=null_collate,\n    )\n\n    print('fold = %s\\n' % str(config.fold))\n    print('train_dataset : \\n%s\\n' % (train_dataset))\n    print('valid_dataset : \\n%s\\n' % (valid_dataset))\n    print('\\n')\n\n    ## net ----------------------------------------\n    print('** net setting **\\n')\n\n    scaler = amp.GradScaler(enabled=config.is_amp)\n    net = Net(config).cuda()\n\n    if config.initial_checkpoint is not None:\n        f = torch.load(config.initial_checkpoint, map_location=lambda storage, loc: storage)\n        start_iteration = f['iteration']\n        start_epoch = f['epoch']\n        state_dict = f['state_dict']\n        net.load_state_dict(state_dict, strict=False)  # True\n    else:\n        start_iteration = 0\n        start_epoch = 0\n        net.load_pretrain()\n\n    print('\\tinitial_checkpoint = %s\\n' % config.initial_checkpoint)\n    print('\\n')\n\n    ## optimiser ----------------------------------\n    # if 0:  ##freeze\n    #     for p in net.stem.parameters():   p.requires_grad = False\n    #     pass\n\n    # freeze_bn(net)\n\n    # -----------------------------------------------\n\n    optimizer = torch.optim.AdamW(filter(lambda p: p.requires_grad, net.parameters()), lr=config.start_lr)\n\n    print('optimizer\\n  %s\\n' % (optimizer))\n    print('\\n')\n\n\n    num_iteration = 6000 * len(train_loader)\n    print(f'\\nIteration Num: {num_iteration}')\n    iter_log = len(train_loader) * 3  # 479\n    iter_valid = iter_log\n    print(f'\\nIteration for Valid: {iter_valid}')\n    iter_save = iter_log\n\n    print('** start training here! **\\n')\n    print('   batch_size = %d \\n' % (config.batch_size))\n    print('                     |-------------- VALID---------|---- TRAIN/BATCH ----------------\\n')\n    print('rate     iter  epoch | dice   loss   tp     tn     | loss           | time           \\n')\n    print('-------------------------------------------------------------------------------------\\n')\n\n    valid_loss = np.zeros(4, np.float32)\n    train_loss = np.zeros(2, np.float32)\n    batch_loss = np.zeros_like(train_loss)\n    sum_train_loss = np.zeros_like(train_loss)\n    sum_train = 0\n\n    start_timer = time.time()\n    iteration = start_iteration\n    epoch = start_epoch\n    rate = 0\n    best_metric = 0\n\n    while iteration < num_iteration:\n        for t, batch in enumerate(train_loader):\n\n            if iteration % iter_save == 0:\n                if iteration != start_iteration and valid_loss[0] > best_metric:\n                    best_metric = valid_loss[0]\n                    os.system(\"rm \" + config.out_dir + f'/fold_{config.fold}_*.pth')\n                    torch.save({\n                        'state_dict': net.state_dict(),\n                        'iteration': iteration,\n                        'epoch': epoch,\n                    }, config.out_dir + f'/fold_{config.fold}_{valid_loss[0]:.4f}.pth')\n\n            if iteration % iter_valid == 0:\n                print(\"\\nFor validation\")\n                valid_loss = validate(net, valid_loader, config)\n\n            if (iteration % iter_log == 0) or (iteration % iter_valid == 0):\n                print('\\r', end='', flush=True)\n                print(message(batch_loss, train_loss, iteration, iter_save, rate, epoch, valid_loss,\n                              start_timer, mode='log') + '\\n')\n\n            # learning rate schduler ------------\n            rate = get_learning_rate(optimizer)\n\n            # one iteration update  -------------\n            batch_size = len(batch['index'])\n            batch['image'] = batch['image'].half().cuda()\n            batch['mask'] = batch['mask'].half().cuda()\n            batch['organ'] = batch['organ'].cuda()\n\n            net.train()\n            net.output_type = ['loss']\n            if 1:\n                with amp.autocast(enabled=config.is_amp):\n                    output = net(batch)\n                    loss0 = output['bce_loss'].mean()\n                    loss1 = output['aux2_loss'].mean()\n\n                optimizer.zero_grad()\n                scaler.scale(loss0 + 0.2 * loss1).backward()\n\n                scaler.unscale_(optimizer)\n                scaler.step(optimizer)\n                scaler.update()\n\n            # print statistics  --------\n            batch_loss[:2] = [loss0.item(), loss1.item()]\n            sum_train_loss += batch_loss\n            sum_train += 1\n            if t % 100 == 0:\n                train_loss = sum_train_loss / (sum_train + 1e-12)\n                sum_train_loss[...] = 0\n                sum_train = 0\n\n            print('\\r', end='', flush=True)\n            print(message(batch_loss, train_loss, iteration, iter_save, rate, epoch, valid_loss,\n                          start_timer, mode='print'), end='', flush=True)\n            epoch += 1 / len(train_loader)\n            iteration += 1\n\n        torch.cuda.empty_cache()\n\n    print('\\n')","metadata":{"execution":{"iopub.status.busy":"2022-08-05T18:00:09.488349Z","iopub.execute_input":"2022-08-05T18:00:09.488696Z","iopub.status.idle":"2022-08-05T18:00:09.521722Z","shell.execute_reply.started":"2022-08-05T18:00:09.488663Z","shell.execute_reply":"2022-08-05T18:00:09.520323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! mkdir models","metadata":{"execution":{"iopub.status.busy":"2022-08-05T18:00:09.526615Z","iopub.execute_input":"2022-08-05T18:00:09.528838Z","iopub.status.idle":"2022-08-05T18:00:10.718592Z","shell.execute_reply.started":"2022-08-05T18:00:09.528802Z","shell.execute_reply":"2022-08-05T18:00:10.717084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fold = 0\ncfg = config(fold)\nmain(cfg)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-05T18:00:10.720751Z","iopub.execute_input":"2022-08-05T18:00:10.721155Z"},"trusted":true},"execution_count":null,"outputs":[]}]}