{"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":"markdown","source":"<div style=\"height:200px;width:100%;margin: 0;\">\n    <img src=\"https://storage.googleapis.com/kaggle-competitions/kaggle/34547/logos/header.png?t=2022-02-15-22-37-27\" style=\"width:100%;\" />\n</div>","metadata":{}},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"credits\"><center>Credits</center></h3>","metadata":{}},{"cell_type":"markdown","source":"This is the reverse engineering of the [hengck23 discussion](https://www.kaggle.com/code/hengck23/lb-0-75-variable-size-swin-transformer-v1-and-v2).<br>\nPlease upvote both discussion/notebooks if you are planning to use Swin Transformers or any part of the code.\n\n**hengck23 owner Disclaimer**\n\n[1] the code is taken from a larger project and is by no means complete. It will has missing import modules, etc. But these are trival functions that you can ignore or fill in yourself.\n\n[2] you are free to use, modify the code for your own notebook or submission","metadata":{}},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"imports\"><center>Imports</center></h3>","metadata":{}},{"cell_type":"code","source":"import os\nimport cv2\nimport time\nimport random\n\nimport torch\nfrom torch import nn\nimport torch.cuda.amp as amp\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import RandomSampler \nfrom torch.utils.data import SequentialSampler\nimport torch.nn.functional as F\nfrom torchmetrics.functional import dice_score\nfrom torch.optim.lr_scheduler import StepLR\n\nis_amp = True\nimport logging\nimport pandas as pd\nfrom sklearn.model_selection import KFold\n\nimport numpy as np\nfrom itertools import repeat\nimport collections.abc\n\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"paths\"><center>Paths</center></h3>","metadata":{}},{"cell_type":"code","source":"!mkdir /kaggle/working/result\n!mkdir /kaggle/working/checkpoint\n\nroot_dir = '/kaggle/working/'\npretrain_dir = '/kaggle/input/swin-tiny-small-22k-pretrained/'\n\nTRAIN = '../input/hubmap-2022-256x256/train/'\nMASKS = '../input/hubmap-2022-256x256/masks/'\nLABELS = '../input/hubmap-organ-segmentation/train.csv'","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"additionals\"><center>Additionals</center></h3>","metadata":{}},{"cell_type":"code","source":"def 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\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    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    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(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","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"randoms\"><center>Random choice</center></h3>","metadata":{}},{"cell_type":"code","source":"def valid_augment5(image, mask, organ):\n    #image, mask  = do_crop(image, mask, image_size, xy=(None,None))\n    return image, mask\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","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"augmentations\"><center>Augmentations</center></h3>","metadata":{}},{"cell_type":"code","source":"def 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\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    \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\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\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\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":{"_kg_hide-input":true,"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"dataset\"><center>Dataset</center></h3>","metadata":{}},{"cell_type":"code","source":"image_size = 768\n\nclass HubmapDataset(Dataset):\n    def __init__(self, df, augment=None):\n\n        self.df = df\n        self.augment = augment\n        self.length = len(self.df)\n        ids = pd.read_csv(LABELS).id.astype(str).values\n        self.fnames = [fname for fname in os.listdir(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(TRAIN,fname)), cv2.COLOR_BGR2RGB)\n        mask = cv2.imread(os.path.join(MASKS,fname),cv2.IMREAD_GRAYSCALE)\n        \n        image = image.astype(np.float32)/255\n        mask  = mask.astype(np.float32)/255\n\n        s = d.pixel_size/0.4 * (image_size/3000)\n        image = cv2.resize(image,dsize=(image_size,image_size),interpolation=cv2.INTER_LINEAR)\n        mask  = cv2.resize(mask, dsize=(image_size,image_size),interpolation=cv2.INTER_LINEAR)\n\n        if self.augment is not None:\n            image, mask = self.augment(image, mask, organ)\n\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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"patching\"><center>Image Patching</center></h3>","metadata":{}},{"cell_type":"code","source":"class 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    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\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    \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\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","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"swin\"><center>Swin Transformer</center></h3>","metadata":{}},{"cell_type":"code","source":"class 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    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\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","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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    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\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\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\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}\"","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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\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\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\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","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"window\"><center>Window functionalities</center></h3>","metadata":{}},{"cell_type":"code","source":"class 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    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\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\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","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"ann\"><center>Upnet + Net + MLP</center></h3>","metadata":{}},{"cell_type":"code","source":"def 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    )","metadata":{"_kg_hide-input":true,"_kg_hide-output":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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    \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","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Net(nn.Module):\n\n    def load_pretrain( self,):\n\n        checkpoint = cfg[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\n    def __init__( self,):\n        super(Net, self).__init__()\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            ** {**cfg['basic']['swin'], **cfg[self.arch]['swin'],\n                **{'out_norm' : LayerNorm2d} }\n        )\n        encoder_dim =cfg[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\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","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_check_net():\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().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())","metadata":{"_kg_hide-input":true,"_kg_hide-output":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"configs\"><center>Configuration</center></h3>","metadata":{}},{"cell_type":"code","source":"cfg = 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 = 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 = 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    )","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"folds\"><center>Folds</center></h3>","metadata":{}},{"cell_type":"code","source":"def make_fold(fold=0):\n    df = pd.read_csv(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","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"dice_score\"><center>Competition Metric</center></h3>","metadata":{}},{"cell_type":"code","source":"def compute_dice_score(probability, mask):\n    N = len(probability)\n    p = probability.reshape(N,-1)\n    t = mask.reshape(N,-1)\n\n    p = p>0.5\n    t = t>0.5\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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"validation\"><center>Validation</center></h3>","metadata":{}},{"cell_type":"code","source":"def validate(net, valid_loader):\n\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 = is_amp):\n\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            pass\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='',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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"init\"><center>Initialization</center></h3>","metadata":{}},{"cell_type":"code","source":"def get_learning_rate(optimizer):\n    return optimizer.param_groups[0]['lr']","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fold = 0\n\nout_dir = root_dir + '/result/upernet-swin-v1-tiny-aux5-768/fold-%d' % (fold)\ninitial_checkpoint = None\n\nstart_lr   = 5e-5 #0.0001\nbatch_size = 8 #32 #32\n\n\n## setup  ----------------------------------------\nfor f in ['checkpoint','train','valid','backup'] : os.makedirs(out_dir +'/'+f, exist_ok=True)\n\n    \nlog = open(out_dir+'/log.train.txt',mode='a')\nlog.write('\\n--- [START %s] %s\\n\\n' % ('Swin', '-' * 64))\nlog.write('\\n')\n\n\n## dataset ----------------------------------------\nlog.write('** dataset setting **\\n')\n\ntrain_df, valid_df = make_fold(fold)\n\ntrain_dataset = HubmapDataset(train_df, train_augment5b)\nvalid_dataset = HubmapDataset(valid_df, valid_augment5)\n\ntrain_loader  = DataLoader(\n    train_dataset,\n    sampler = RandomSampler(train_dataset),\n    batch_size  = 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\nvalid_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\nlog.write('fold = %s\\n'%str(fold))\nlog.write('train_dataset : \\n%s\\n'%(train_dataset))\nlog.write('valid_dataset : \\n%s\\n'%(valid_dataset))\nlog.write('\\n')\n\n\n## net ----------------------------------------\nlog.write('** net setting **\\n')\n\nscaler = amp.GradScaler(enabled = is_amp)\nnet = Net().cuda()\n\nif initial_checkpoint is not None:\n    f = torch.load(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\nelse:\n    start_iteration = 0\n    start_epoch = 0\n    net.load_pretrain()\n\n\nlog.write('\\tinitial_checkpoint = %s\\n' % initial_checkpoint)\nlog.write('\\n')\n\n\n## optimiser ----------------------------------\nif 0: ##freeze\n    for p in net.stem.parameters():   p.requires_grad = False\n    pass\n\ndef 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#freeze_bn(net)\n\n#-----------------------------------------------\n\noptimizer = torch.optim.AdamW(filter(lambda p: p.requires_grad, net.parameters()),lr=start_lr)\n\nlog.write('optimizer\\n  %s\\n'%(optimizer))\nlog.write('\\n')\n\nnum_iteration = 1000*len(train_loader)\niter_log   = len(train_loader)*3 #479\niter_valid = iter_log\niter_save  = iter_log","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"training\"><center>Training</center></h3>","metadata":{}},{"cell_type":"code","source":"log.write('** start training here! **\\n')\nlog.write('   batch_size = %d \\n'%(batch_size))\nlog.write('                     |-------------- VALID---------|---- TRAIN/BATCH ----------------\\n')\nlog.write('rate     iter  epoch | dice   loss   tp     tn     | loss           | time           \\n')\nlog.write('-------------------------------------------------------------------------------------\\n')\n\nvalid_loss = np.zeros(4,np.float32)\ntrain_loss = np.zeros(2,np.float32)\nbatch_loss = np.zeros_like(train_loss)\nsum_train_loss = np.zeros_like(train_loss)\nsum_train = 0\n\nstart_timer = time.time()\niteration = start_iteration\nepoch = start_epoch\nrate = 0\n\nwhile iteration < num_iteration:\n    for t, batch in enumerate(train_loader):\n\n        if iteration%iter_save==0:\n            if iteration != start_iteration:\n                torch.save({\n                    'state_dict': net.state_dict(),\n                    'iteration': iteration,\n                    'epoch': epoch,\n                }, out_dir + '/checkpoint/%08d.model.pth' %  (iteration))\n                pass\n\n\n        if (iteration%iter_valid==0):\n            valid_loss = validate(net, valid_loader)\n            pass\n\n\n        if (iteration%iter_log==0) or (iteration%iter_valid==0):\n            print('\\r', end='', flush=True)\n            log.write(message(mode='log') + '\\n')\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\n        net.train()\n        net.output_type = ['loss']\n        if 1:\n            with amp.autocast(enabled = 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\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(mode='print'), end='', flush=True)\n        epoch += 1 / len(train_loader)\n        iteration += 1\n        \n    torch.cuda.empty_cache()\n    \nlog.write('\\n')\nlog.close()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}