{"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 PIL.Image\nimport matplotlib.pyplot as plt\n\nimport 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","metadata":{"execution":{"iopub.status.busy":"2022-08-12T06:57:03.071704Z","iopub.execute_input":"2022-08-12T06:57:03.072378Z","iopub.status.idle":"2022-08-12T06:57:06.426555Z","shell.execute_reply.started":"2022-08-12T06:57:03.072288Z","shell.execute_reply":"2022-08-12T06:57:06.425467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_encode_less_memory(img):\n    #the image should be transposed\n    #img = cv2.imread(img)\n    pixels = img.T.flatten()\n    \n    # This simplified method requires first and last pixel to be zero\n    pixels[0] = 0\n    pixels[-1] = 0\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    runs[1::2] -= runs[::2]\n    \n    return ' '.join(str(x) for x in runs)","metadata":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.428646Z","iopub.execute_input":"2022-08-12T06:57:06.429239Z","iopub.status.idle":"2022-08-12T06:57:06.436424Z","shell.execute_reply.started":"2022-08-12T06:57:06.429201Z","shell.execute_reply":"2022-08-12T06:57:06.434557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(\"../input/hubmap-organ-segmentation/test.csv\")\n\ntest_path = \"../input/hubmap-organ-segmentation/test_images\"\n\n#if not os.path.exists('weights'):\n    #os.makedirs('weights')\n#pretrain_dir = os.getcwd() + \"/weights\"","metadata":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.437973Z","iopub.execute_input":"2022-08-12T06:57:06.438339Z","iopub.status.idle":"2022-08-12T06:57:06.478326Z","shell.execute_reply.started":"2022-08-12T06:57:06.438304Z","shell.execute_reply":"2022-08-12T06:57:06.47746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.482069Z","iopub.execute_input":"2022-08-12T06:57:06.482357Z","iopub.status.idle":"2022-08-12T06:57:06.503733Z","shell.execute_reply.started":"2022-08-12T06:57:06.482331Z","shell.execute_reply":"2022-08-12T06:57:06.502444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.504975Z","iopub.execute_input":"2022-08-12T06:57:06.506105Z","iopub.status.idle":"2022-08-12T06:57:06.518937Z","shell.execute_reply.started":"2022-08-12T06:57:06.506071Z","shell.execute_reply":"2022-08-12T06:57:06.517853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.520634Z","iopub.execute_input":"2022-08-12T06:57:06.521051Z","iopub.status.idle":"2022-08-12T06:57:06.540332Z","shell.execute_reply.started":"2022-08-12T06:57:06.521018Z","shell.execute_reply":"2022-08-12T06:57:06.539316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_size = 768\n# image_size = 392\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        print('Image name.......................',fname)\n        mask = cv2.imread(os.path.join(MASKS,fname.split('.')[0]+'.png'),cv2.IMREAD_GRAYSCALE)\n        print(mask)\n        plt.imshow(mask)\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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.543945Z","iopub.execute_input":"2022-08-12T06:57:06.544236Z","iopub.status.idle":"2022-08-12T06:57:06.559904Z","shell.execute_reply.started":"2022-08-12T06:57:06.544212Z","shell.execute_reply":"2022-08-12T06:57:06.558728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.561589Z","iopub.execute_input":"2022-08-12T06:57:06.561979Z","iopub.status.idle":"2022-08-12T06:57:06.580578Z","shell.execute_reply.started":"2022-08-12T06:57:06.561945Z","shell.execute_reply":"2022-08-12T06:57:06.579617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.582085Z","iopub.execute_input":"2022-08-12T06:57:06.5826Z","iopub.status.idle":"2022-08-12T06:57:06.600624Z","shell.execute_reply.started":"2022-08-12T06:57:06.582565Z","shell.execute_reply":"2022-08-12T06:57:06.599368Z"},"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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.604966Z","iopub.execute_input":"2022-08-12T06:57:06.605224Z","iopub.status.idle":"2022-08-12T06:57:06.624048Z","shell.execute_reply.started":"2022-08-12T06:57:06.605201Z","shell.execute_reply":"2022-08-12T06:57:06.623068Z"},"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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.625613Z","iopub.execute_input":"2022-08-12T06:57:06.625971Z","iopub.status.idle":"2022-08-12T06:57:06.645086Z","shell.execute_reply.started":"2022-08-12T06:57:06.625938Z","shell.execute_reply":"2022-08-12T06:57:06.64404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.646652Z","iopub.execute_input":"2022-08-12T06:57:06.647146Z","iopub.status.idle":"2022-08-12T06:57:06.669329Z","shell.execute_reply.started":"2022-08-12T06:57:06.647111Z","shell.execute_reply":"2022-08-12T06:57:06.668298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.672619Z","iopub.execute_input":"2022-08-12T06:57:06.67302Z","iopub.status.idle":"2022-08-12T06:57:06.682142Z","shell.execute_reply.started":"2022-08-12T06:57:06.672995Z","shell.execute_reply":"2022-08-12T06:57:06.681198Z"},"trusted":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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.683559Z","iopub.execute_input":"2022-08-12T06:57:06.684395Z","iopub.status.idle":"2022-08-12T06:57:06.701666Z","shell.execute_reply.started":"2022-08-12T06:57:06.684362Z","shell.execute_reply":"2022-08-12T06:57:06.700674Z"},"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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.703663Z","iopub.execute_input":"2022-08-12T06:57:06.704766Z","iopub.status.idle":"2022-08-12T06:57:06.714664Z","shell.execute_reply.started":"2022-08-12T06:57:06.704732Z","shell.execute_reply":"2022-08-12T06:57:06.713777Z"},"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.arch = 'swin_small_patch4_window7_224_22k'\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        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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.71632Z","iopub.execute_input":"2022-08-12T06:57:06.717267Z","iopub.status.idle":"2022-08-12T06:57:06.731166Z","shell.execute_reply.started":"2022-08-12T06:57:06.717222Z","shell.execute_reply":"2022-08-12T06:57:06.729999Z"},"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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.732746Z","iopub.execute_input":"2022-08-12T06:57:06.733791Z","iopub.status.idle":"2022-08-12T06:57:06.747368Z","shell.execute_reply.started":"2022-08-12T06:57:06.733748Z","shell.execute_reply":"2022-08-12T06:57:06.746448Z"},"trusted":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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.749824Z","iopub.execute_input":"2022-08-12T06:57:06.750557Z","iopub.status.idle":"2022-08-12T06:57:06.761061Z","shell.execute_reply.started":"2022-08-12T06:57:06.750522Z","shell.execute_reply":"2022-08-12T06:57:06.759946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pretrain_dir = \"./weights\"","metadata":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.762509Z","iopub.execute_input":"2022-08-12T06:57:06.762976Z","iopub.status.idle":"2022-08-12T06:57:06.771506Z","shell.execute_reply.started":"2022-08-12T06:57:06.762944Z","shell.execute_reply":"2022-08-12T06:57:06.770535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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        \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":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.774973Z","iopub.execute_input":"2022-08-12T06:57:06.775686Z","iopub.status.idle":"2022-08-12T06:57:06.786337Z","shell.execute_reply.started":"2022-08-12T06:57:06.775658Z","shell.execute_reply":"2022-08-12T06:57:06.785397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_size = 768\n\norgan_threshold = {\n    'kidney': 0.225,\n    'prostate': 0.225,\n    'largeintestine': 0.225,\n    'spleen': 0.225,\n    'lung': 0.2,\n}\n\n\n# ----\nresult = {\n    'id': [],\n    'probability': [],\n    'rle': [],\n}\ndef load_model(path):\n    model =  Net().cuda()\n    model.output_type = ['inference']\n    model.load_state_dict(torch.load(path)['state_dict'],strict=False)\n    model.eval()\n    return model\nmodels = []\nckpt_paths = [f'../input/swin-small/v7_model_small_bt8_fold.pth',\n              '../input/swin-small/v7_model_small_bt8_fold1.pth',\n              '../input/swin-small/v7_model_small_bt8_fold2.pth',\n              '../input/swin-small/v7_model_small_bt8_fold3.pth',\n              '../input/swin-small/v7_model_small_bt8_fold4.pth',\n             ]\nfor path in ckpt_paths:\n    models.append(load_model(path))","metadata":{"execution":{"iopub.status.busy":"2022-08-12T06:57:06.787495Z","iopub.execute_input":"2022-08-12T06:57:06.787842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#start_timer = timer()\nfor person_id in test_df['id'].unique():\n    organ_cat = test_df[test_df['id'] == person_id]['organ'].values[0]\n    image = cv2.imread('../input/hubmap-organ-segmentation/test_images/%d.tiff' % person_id,\n                       cv2.IMREAD_COLOR)\n    image = image.astype(np.float32) / 255\n    H, W, _ = image.shape\n    # image = cv2.cvtColor(read_tiff(tiff_file),cv2.COLOR_RGB2BGR)\n    #d = test_df[test_df['id'] == person_id]['pixel_size'].values[0]\n    s = test_df[test_df['id'] == person_id]['pixel_size'].values[0] / 0.4 * (image_size / test_df[test_df['id'] == person_id]['img_height'].values[0])\n    \n    #heighttt = test_df[test_df['id'] == person_id]['img_height'].values[0]\n    #s = d / 0.4 * (image_size / heighttt)\n#     if s >= 1:\n#         image = cv2.resize(image, dsize=(image_size,image_size), interpolation=cv2.INTER_AREA)\n#     else:\n#         image = cv2.resize(image, dsize=None, fx=s, fy=s, interpolation=cv2.INTER_LINEAR)\n    image = cv2.resize(image, dsize=(image_size,image_size), interpolation=cv2.INTER_AREA)\n\n    # we already probe that all test image are fixed size (within 0.8,1.2 of normalised size)\n    # hence we just use resize. else use padding to 32 for variable size\n    print(\"S:\",s)\n    image = image_to_tensor(image)\n    image = image.cuda()\n    batch = {\n        'image':\n            torch.stack([\n                image,\n                torch.flip(image, [1]),\n                torch.flip(image, [2]),\n            ]),  # simple TTA\n    }\n\n    with amp.autocast(enabled=is_amp):\n        py = None\n        for net in models:\n            output = net(batch)  # data_parallel(net, batch) #\n            # probability += output['probability']\n            probability = F.interpolate(output['probability'], size=(H, W), mode='bilinear', align_corners=False)\n            # undo TTA\n            probability[1] = torch.flip(probability[1], [1])\n            probability[2] = torch.flip(probability[2], [2])\n            if py is None:\n                py = probability.detach()\n            else:\n                py += probability.detach()\n        py /= len(models)\n    py = py.float().data.cpu().numpy().mean(0)[0]\n    p = py > organ_threshold[organ_cat]\n    \n    #widthhh=test_df[test_df['id'] == person_id]['img_width'].values[0]\n    #heighttt=test_df[test_df['id'] == person_id]['img_height'].values[0]\n    mask = cv2.resize(1*p, (H,W),interpolation=cv2.INTER_NEAREST)\n    \n    #rle = rle_encode_less_memory(1*p)\n    rle = rle_encode_less_memory(mask)\n    \n    result['rle'].append(rle)\n    #result['probability'].append(probability)\n    result['id'].append(person_id)\n    #print('\\r', t, end='', flush=True)\nprint('')\n# ---\n\nsubmit_df = pd.DataFrame({'id': result['id'], 'rle': result['rle']})\nprint(submit_df)\nprint('submit_df ok!')\nprint('')\n\nsubmit_df.to_csv('submission.csv', index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(mask)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}