{"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":"# mode_ = 'TRAIN'\nmode_ = 'TEST'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport cv2\n\ndef is_image(I):\n    \"\"\"\n    Is I an image.\n    \"\"\"\n    if not isinstance(I, np.ndarray):\n        return False\n    if not I.ndim == 3:\n        return False\n    return True\n\n\ndef is_uint8_image(I):\n    \"\"\"\n    Is I a uint8 image.\n    \"\"\"\n    if not is_image(I):\n        return False\n    if I.dtype != np.uint8:\n        return False\n    return True\n\n\ndef standardize(I, percentile=95):\n    \"\"\"\n    Transform image I to standard brightness.\n    Modifies the luminosity channel such that a fixed percentile is saturated.\n\n    :param I: Image uint8 RGB.\n    :param percentile: Percentile for luminosity saturation. At least (100 - percentile)% of pixels should be fully luminous (white).\n    :return: Image uint8 RGB with standardized brightness.\n    \"\"\"\n    assert is_uint8_image(I), \"Image should be RGB uint8.\"\n    I_LAB = cv2.cvtColor(I, cv2.COLOR_RGB2LAB)\n    L_float = I_LAB[:, :, 0].astype(float)\n    p = np.percentile(L_float, percentile)\n    I_LAB[:, :, 0] = np.clip(255 * L_float / p, 0, 255).astype(np.uint8)\n    I = cv2.cvtColor(I_LAB, cv2.COLOR_LAB2RGB)\n    return I\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef rle_encode_less_memory(img):\n    #the image should be transposed\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)\n\n# https://www.kaggle.com/paulorzp/rle-functions-run-length-encode-decode\ndef mask2rle(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels= img.T.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n \ndef rle2mask(mask_rle, shape=(1600,256)):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (width,height) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n\n    '''\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n#     print(starts.size)\n    img = np.zeros(shape[0]*shape[1], dtype=np.int8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n#         img[lo] = 1\n#         img[hi] = 1\n    retImg = img.reshape(shape).T\n    print(np.where(retImg==1))\n    return retImg\n\ndef dict2mask(mask_dict, shape=(1600,256)):\n    '''\n    aaa\n\n    '''\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8).reshape(shape)\n    for xy in mask_dict:\n#         img[xy[1]-1][xy[0]-1] = 1\n        img[xy[1]][xy[0]] = 1\n    \n    return img","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MASK_FOLDER = '../input/maskdata3/'\nTRAIN_FOLDER = '/kaggle/input/hubmap-organ-segmentation/'\nTRAIN_IMAGE_FOLDER = '../input/transformedimage2/transformed_image/'\n\nif mode_ == 'TRAIN':\n    IMAGE_FOLDER = 'train_images/'\n#     pretrain_dir = '/kaggle/input/swintinysmall22kpretrained/'\n    pretrain_dir = '/kaggle/input/pretrained2/'\n    LABELS = TRAIN_FOLDER+'train.csv'\nelse:\n    IMAGE_FOLDER = 'test_images/'\n#     pretrain_dir = '/kaggle/input/pretrained2/'\n    pretrain_dir = '/kaggle/input/hubmapmodel/'\n    LABELS = TRAIN_FOLDER+'test.csv'\n    SUBMIT = TRAIN_FOLDER+'sample_submission.csv'\n   ######## TEST OF TEST 2 #########\n#     IMAGE_FOLDER = 'train_images/'\n#     LABELS = TRAIN_FOLDER+'train.csv'\n    \n# TRAIN = TRAIN_FOLDER+IMAGE_FOLDER\nTRAIN = TRAIN_IMAGE_FOLDER\nMASKS = MASK_FOLDER\n\nroot_dir = '/kaggle/working/'\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# make DataLoader\nimport os\nimport cv2\nimport time\nimport random\nimport sys\nimport math\n\nimport matplotlib.pyplot as plt\n\nimport torch\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import RandomSampler \nfrom torch.utils.data import SequentialSampler\n\nimport pandas as pd\nfrom sklearn.model_selection import KFold\n\nimport numpy as np\n\nimport warnings\n# warnings.filterwarnings('ignore')\n\nskip_id = 10078\ndef make_fold(fold=0):\n    df = pd.read_csv(TRAIN_FOLDER + 'train.csv')\n    df = df[df['id'] != skip_id]\n\n    num_fold = 5\n    skf = KFold(n_splits=num_fold, shuffle=True,random_state=42)\n\n    df.loc[:,'fold']=-1\n    for f,(t_idx, v_idx) in enumerate(skf.split(X=df['id'], y=df['organ'])):\n        df.iloc[v_idx,-1]=f\n\n    #check\n    if 0:\n        for f in range(num_fold):\n            train_df=df[df.fold!=f].reset_index(drop=True)\n            valid_df=df[df.fold==f].reset_index(drop=True)\n\n            print('fold %d'%f)\n            t = train_df.organ.value_counts().to_dict()\n            v = valid_df.organ.value_counts().to_dict()\n            for k in ['kidney', 'prostate', 'largeintestine', 'spleen', 'lung']:\n                print('%32s %3d (%0.3f)  %3d (%0.3f)'%(k,t.get(k,0),t.get(k,0)/len(train_df),v.get(k,0),v.get(k,0)/len(valid_df)))\n\n            print('')\n            zz=0\n\n    train_df=df[df.fold!=fold].reset_index(drop=True)\n    valid_df=df[df.fold==fold].reset_index(drop=True)\n    return train_df,valid_df\n\ndef do_random_crop(image, mask):\n    h, w, _ = image.shape\n    crop_size=[math.floor(h/2),math.floor(w/2)]\n\n    # 0~(400-224)の間で画像のtop, leftを決める\n    top = np.random.randint(0, h - crop_size[0])\n    left = np.random.randint(0, w - crop_size[1])\n\n    # top, leftから画像のサイズである224を足して、bottomとrightを決める\n    bottom = top + crop_size[0]\n    right = left + crop_size[1]\n\n    # 決めたtop, bottom, left, rightを使って画像を抜き出す\n    image = image[top:bottom, left:right, :]\n    mask = mask[top:bottom, left:right, :]\n    return image, mask\n\ndef do_random_flip(image, mask):\n    if np.random.rand()>0.5:\n        image = cv2.flip(image,0)\n        mask = cv2.flip(mask,0)\n    if np.random.rand()>0.5:\n        image = cv2.flip(image,1)\n        mask = cv2.flip(mask,1)\n    if np.random.rand()>0.5:\n        # change order: before, 768,768,3: x,y,z -> y,x,z\n        image = image.transpose(1,0,2)\n        mask = mask.transpose(1,0)\n    # メモリ内の連続した配列（ndim> = 1）を返します（C順）。\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\n\ndef do_luminosity_standarize(image):\n    percentile=95\n    \"\"\"\n    Transform image I to standard brightness.\n    Modifies the luminosity channel such that a fixed percentile is saturated.\n\n    :param I: Image uint8 RGB.\n    :param percentile: Percentile for luminosity saturation. At least (100 - percentile)% of pixels should be fully luminous (white).\n    :return: Image uint8 RGB with standardized brightness.\n    \"\"\"\n    image = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)\n    L_float = image[:, :, 0].astype(float)\n    p = np.percentile(L_float, percentile)\n    image[:, :, 0] = np.clip(255 * L_float / p, 0, 255).astype(np.uint8)\n    image = cv2.cvtColor(image, cv2.COLOR_LAB2RGB)\n    \n    L_float = 0\n    p = 0\n    \n    return image\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#         lambda image, mask: do_random_crop(image, mask),\n    ], 1): image, mask = fn(image, mask)\n        \n    return image, mask\n\ndef valid_augment5(image, mask, organ):\n    #image, mask  = do_crop(image, mask, image_size, xy=(None,None))\n    return image, mask\n\ndef test_augment5t(image, organ):\n    #image, mask  = do_crop(image, mask, image_size, xy=(None,None))\n    image = do_luminosity_standarize(image)\n    return image\n\ndef image_to_tensor(image, mode='bgr'): #image mode\n    if mode=='bgr':\n        image = image[:,:,::-1]\n    x = image\n    x = x.transpose(2,0,1)\n    x = np.ascontiguousarray(x)\n    x = torch.tensor(x, dtype=torch.float)\n    return x\n\ndef mask_to_tensor(mask):\n    x = mask\n    x = torch.tensor(x, dtype=torch.float)\n    return x\n\nimage_size = 768\nclass HubmapDatasetTest(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 = df.id.astype(str).values\n        self.fnames = [fname for fname in os.listdir(TRAIN_FOLDER+IMAGE_FOLDER) 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    def __str__(self):\n        string = ''\n        string += '\\tlen = %d\\n' % len(self)\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_FOLDER+IMAGE_FOLDER,fname)), cv2.COLOR_BGR2RGB)\n        image = image.astype(np.float32)/255\n        # image size scalse for HuBMAP image\n        if(d.data_source=='Hubmap'):\n            scale = d.pixel_size/0.4 * (image.shape[1]/3000)\n            image = cv2.resize(image,dsize=(int(scale*3000),int(scale*3000)),interpolation=cv2.INTER_LINEAR)\n        # INTER_AREA: ピクセル領域の関係を利用したリサンプリング。画像を大幅に縮小する場合、モアレを避けることができる手法。画像を拡大する場合は、INTER_NEARESTと同様になる\n        image = cv2.resize(image,dsize=(image_size,image_size),interpolation=cv2.INTER_AREA)\n\n        if(d.data_source=='Hubmap'):\n            if self.augment is not None:\n                image = self.augment(image, organ)\n\n        r ={}\n        r['index']= index\n        r['id'] = fname\n        r['organ'] = torch.tensor([organ], dtype=torch.long)\n        r['image'] = image_to_tensor(image)\n#         r['mask' ] = mask_to_tensor(mask)\n        return r\n\n# ====================================\n\ntest_df = pd.read_csv(LABELS)\ntest_dataset = HubmapDatasetTest(test_df, test_augment5t)\n\n# divide into 2 dataset according to HPA or Hubmap\ntest_HPA_df = test_df[test_df['data_source']=='HPA'].reset_index()\ntest_Hubmap_df = test_df[test_df['data_source']=='Hubmap'].reset_index()\n\ntest_HPA_dataset = HubmapDatasetTest(test_HPA_df, test_augment5t)\ntest_Hubmap_dataset = HubmapDatasetTest(test_Hubmap_df, test_augment5t)\n\ntensor_list = ['mask', 'image', 'organ']\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\ndef null_collate_test(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['organ'] = d['organ'].reshape(-1)\n    return d\n\n#     test_loader = DataLoader(\n#         test_dataset,\n#         sampler = SequentialSampler(test_dataset),\n#         batch_size  = 8,\n#         drop_last   = False,\n#         num_workers = 4,\n#         pin_memory  = False,\n#         collate_fn = null_collate,\n#     )\n\ntest_HPA_loader = DataLoader(\n    test_HPA_dataset,\n    sampler = SequentialSampler(test_HPA_dataset),\n#         batch_size  = 8,\n    batch_size  = 1,\n    drop_last   = False,\n#         num_workers = 4,\n    num_workers = 0,\n    pin_memory  = False,\n    collate_fn = null_collate_test,\n)\n\ntest_Hubmap_loader = DataLoader(\n    test_Hubmap_dataset,\n    sampler = SequentialSampler(test_Hubmap_dataset),\n#         batch_size  = 8,\n    batch_size  = 1,\n    drop_last   = False,\n#         num_workers = 4,\n    num_workers = 0,\n    pin_memory  = False,\n    collate_fn = null_collate_test,\n)\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import nn\nimport torch.cuda.amp as amp\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\n\nfrom itertools import repeat\nimport collections.abc\nimport gc","metadata":{},"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        pretrained_original=dict(\n#             checkpoint = pretrain_dir+'fold4_00010080.model.pth',\n#             checkpoint_stained = pretrain_dir+'fold4_stained_00012915.model.pth',\n            checkpoint = pretrain_dir+'fold1_00008190.model.pth',\n            checkpoint_stained = pretrain_dir+'fold1_stained_00008505.model.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_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#             checkpoint = '/kaggle/input/swintinysmall22kpretrained/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_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_count":null,"outputs":[]},{"cell_type":"code","source":"def _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\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        \ndef trunc_normal_(tensor, mean=0., std=1., a=-2., b=2.):\n    return _no_grad_trunc_normal_(tensor, mean, std, a, b)\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\nclass PatchEmbed(nn.Module):\n    r\"\"\" Image to Patch Embedding\n\n    Args:\n        patch_size (int): Patch token size. Default: 4.\n        in_chans (int): Number of input image channels. Default: 3.\n        embed_dim (int): Number of linear projection output channels. Default: 96.\n        norm_layer (nn.Module, optional): Normalization layer. Default: None\n    \"\"\"\n    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        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_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_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_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        self.patch_embed = PatchEmbed(\n            patch_size=patch_size,\n            in_chans=in_chans,\n            embed_dim=embed_dim,\n            norm_layer=norm_layer if patch_norm else None\n        )\n        self.pos_drop = nn.Dropout(p=drop_rate)\n\n        # stochastic depth\n        dpr = np.linspace(0, drop_path_rate, sum(depths)).tolist() # stochastic depth decay rule\n\n        # build layers\n        self.layers = nn.ModuleList()\n        for i in range(self.num_layers):\n            layer = BasicLayer(\n                dim=int(embed_dim * 2 ** i),\n                depth=depths[i],\n                num_heads=num_heads[i],\n                window_size=window_size,\n                mlp_ratio=self.mlp_ratio,\n                qkv_bias=qkv_bias, qk_scale=qk_scale,\n                drop=drop_rate,\n                attn_drop=attn_drop_rate,\n                drop_path=dpr[sum(depths[:i]):sum(depths[:i + 1])],\n                norm_layer=norm_layer,\n                downsample=PatchMerging if (i < self.num_layers - 1) else None,\n            )\n            self.layers.append(layer)\n\n        #---\n        # add a norm layer for each output\n        self.out_norm = nn.ModuleList(\n            [ out_norm(int(embed_dim * 2 ** i)) for i in range(self.num_layers)]\n        )\n\n        #---\n        self.apply(self._init_weights)\n\n    def _init_weights(self, m):\n        if isinstance(m, nn.Linear):\n            trunc_normal_(m.weight, std=.02)\n            if isinstance(m, nn.Linear) and m.bias is not None:\n                nn.init.constant_(m.bias, 0)\n        elif isinstance(m, nn.LayerNorm):\n            nn.init.constant_(m.bias, 0)\n            nn.init.constant_(m.weight, 1.0)\n\n    def forward(self, x):\n        x = self.patch_embed(x)\n        Wh, Ww = x.size(2), x.size(3)\n\n        #positional encode?\n        x = x.flatten(2).transpose(1, 2)\n        x = self.pos_drop(x)\n\n        outs = []\n        for i in range(self.num_layers):\n            x_out, H, W, x, Wh, Ww = self.layers[i](x, Wh, Ww)\n            out = x_out.view(-1, H, W, int(self.embed_dim * 2 ** i)).permute(0, 3, 1, 2).contiguous()\n            out = self.out_norm[i](out)\n            outs.append(out)\n\n        return outs","metadata":{},"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_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    )\n\nclass UPerDecoder(nn.Module):\n    def __init__(self,\n        in_dim=[256, 512, 1024, 2048],\n        ppm_pool_scale=[1, 2, 3, 6],\n        ppm_dim=512,\n        fpn_out_dim=256\n    ):\n        super(UPerDecoder, self).__init__()\n\n        # PPM ----\n        dim = in_dim[-1]\n        ppm_pooling = []\n        ppm_conv = []\n\n        for scale in ppm_pool_scale:\n            ppm_pooling.append(\n                nn.AdaptiveAvgPool2d(scale)\n            )\n            ppm_conv.append(\n                nn.Sequential(\n                    nn.Conv2d(dim, ppm_dim, kernel_size=1, bias=False),\n                    nn.BatchNorm2d(ppm_dim),\n                    nn.ReLU(inplace=True)\n                )\n            )\n        self.ppm_pooling   = nn.ModuleList(ppm_pooling)\n        self.ppm_conv      = nn.ModuleList(ppm_conv)\n        self.ppm_out = conv3x3_bn_relu(dim + len(ppm_pool_scale)*ppm_dim, fpn_out_dim, 1)\n\n        # FPN ----\n        fpn_in = []\n        for i in range(0, len(in_dim)-1):  # skip the top layer\n            fpn_in.append(\n                nn.Sequential(\n                    nn.Conv2d(in_dim[i], fpn_out_dim, kernel_size=1, bias=False),\n                    nn.BatchNorm2d(fpn_out_dim),\n                    nn.ReLU(inplace=True)\n                )\n            )\n        self.fpn_in = nn.ModuleList(fpn_in)\n\n        fpn_out = []\n        for i in range(len(in_dim) - 1):  # skip the top layer\n            fpn_out.append(\n                conv3x3_bn_relu(fpn_out_dim, fpn_out_dim, 1),\n            )\n        self.fpn_out = nn.ModuleList(fpn_out)\n\n        self.fpn_fuse = nn.Sequential(\n            conv3x3_bn_relu(len(in_dim) * fpn_out_dim, fpn_out_dim, 1),\n        )\n\n    def forward(self, feature):\n        f = feature[-1]\n        pool_shape = f.shape[2:]\n\n        ppm_out = [f]\n        for pool, conv in zip(self.ppm_pooling, self.ppm_conv):\n            p = pool(f)\n            p = F.interpolate(p, size=pool_shape, mode='bilinear', align_corners=False)\n            p = conv(p)\n            ppm_out.append(p)\n        ppm_out = torch.cat(ppm_out, 1)\n        down = self.ppm_out(ppm_out)\n\n        fpn_out = [down]\n        for i in reversed(range(len(feature) - 1)):\n            lateral = feature[i]\n            lateral = self.fpn_in[i](lateral) # lateral branch\n            down = F.interpolate(down, size=lateral.shape[2:], mode='bilinear', align_corners=False) # top-down branch\n            down = down + lateral\n            fpn_out.append(self.fpn_out[i](down))\n\n        fpn_out.reverse() # [P2 - P5]\n        fusion_shape = fpn_out[0].shape[2:]\n        fusion = [fpn_out[0]]\n        for i in range(1, len(fpn_out)):\n            fusion.append(\n                F.interpolate( fpn_out[i], fusion_shape, mode='bilinear', align_corners=False)\n            )\n        x = self.fpn_fuse( torch.cat(fusion, 1))\n\n        return x, fusion\n    \nclass LayerNorm2d(nn.Module):\n    def __init__(self, dim, eps=1e-6):\n        super().__init__()\n        self.weight = nn.Parameter(torch.ones(dim))\n        self.bias = nn.Parameter(torch.zeros(dim))\n        self.eps = eps\n\n    def forward(self, x):\n        u = x.mean(1, keepdim=True)\n        s = (x - u).pow(2).mean(1, keepdim=True)\n        x = (x - u) / torch.sqrt(s + self.eps)\n        x = self.weight[:, None, None] * x + self.bias[:, None, None]\n        return x\n    \ndef criterion_aux_loss(logit, mask):\n    mask = F.interpolate(mask,size=logit.shape[-2:], mode='nearest')\n    loss = F.binary_cross_entropy_with_logits(logit,mask)\n    return loss\n\n\n\nclass Net(nn.Module):\n\n    def load_pretrain( self,):\n        checkpoint = cfg[self.arch]['checkpoint']\n        print('loading %s ...'%checkpoint)\n#         print(torch.load(checkpoint, map_location=lambda storage, loc: storage))\n        checkpoint = torch.load(checkpoint, map_location=lambda storage, loc: storage)['model']\n        if 0:\n            skip = ['relative_coords_table','relative_position_index']\n            filtered={}\n            for k,v in checkpoint.items():\n                if any([s in k for s in skip ]): continue\n                filtered[k]=v\n            checkpoint = filtered\n        print(self.encoder.load_state_dict(checkpoint,strict=False))  #True\n        \n    def load_pretrain_stained( self,):\n        checkpoint = cfg[self.arch]['checkpoint_stained']\n        print('loading %s ...'%checkpoint)\n#         print(torch.load(checkpoint, map_location=lambda storage, loc: storage))\n        checkpoint = torch.load(checkpoint, map_location=lambda storage, loc: storage)['model']\n        if 0:\n            skip = ['relative_coords_table','relative_position_index']\n            filtered={}\n            for k,v in checkpoint.items():\n                if any([s in k for s in skip ]): continue\n                filtered[k]=v\n            checkpoint = filtered\n        print(self.encoder.load_state_dict(checkpoint,strict=False))  #True\n        \n    def __init__( self,):\n        super(Net, self).__init__()\n        self.output_type = ['inference', 'loss']\n\n        self.rgb = RGB()\n        self.arch = 'pretrained_original'\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    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":{},"execution_count":null,"outputs":[]},{"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\n\ndef validate(net, valid_loader):\n\n    valid_num = 0\n    idx_set = []\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#                 batch['sex'] = batch['sex'].cuda()\n#                 batch['age'] = batch['age'].cuda()\n\n                output = net(batch)\n                loss0  = output['bce_loss'].mean()\n        \n        idx_set.append(batch['index'])\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    if(mode_ == 'TRAIN'):\n        return [dice, loss,  0, 0]\n    else:\n        return idx_set, valid_probability\n\ndef tensor_to_mask(msk, id_):\n    mask_mean = np.mean(msk)\n    mask_mean_of_mean = np.mean(msk[msk>mask_mean])\n    msk = msk>mask_mean_of_mean\n    msk = msk.astype('uint8')\n    \n    img_height = int(test_df.loc[test_df['id']==int(id_),'img_height'].values[0])\n    img_width = int(test_df.loc[test_df['id']==int(id_),'img_width'].values[0])\n    msk  = cv2.resize(msk, dsize=(img_width,img_height),interpolation=cv2.INTER_NEAREST)\n    return msk\n\ndef test(net, t_loader, submission_df):\n    net = net.eval()\n#     start_timer = time.time()\n    for t, batch in enumerate(t_loader):\n        net.output_type = ['inference']\n        with torch.no_grad():\n            with amp.autocast(enabled = is_amp):\n                batch_size = len(batch['index'])\n                batch['image'] = batch['image'].cuda()\n#                 batch['mask' ] = batch['mask' ].cuda()\n                batch['organ'] = batch['organ'].cuda()\n\n                output = net(batch)\n\n        probability  = output['probability'].data.cpu().detach().numpy()\n        for b in range(batch_size):\n            id_ = (batch['id'][b].split('.')[0])\n            msk = tensor_to_mask(probability[b,0], id_)\n            record = pd.Series([id_, mask2rle(msk)], index=submission_df.columns)\n            submission_df = submission_df.append(record, ignore_index=True)\n\n            del id_, msk\n            gc.collect()\n\n        del probability\n        gc.collect()\n#         torch.cuda.empty_cache()\n\n#         del  batch['image'], batch\n#         gc.collect()\n#         torch.cuda.empty_cache()\n    \n    return submission_df","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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\n\ndef get_learning_rate(optimizer):\n    return optimizer.param_groups[0]['lr']","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\n# from kaggle\nfold = 0\n\nout_dir = root_dir + '/result/upernet-swin-v1-tiny-aux5-768/fold-%d' % (fold)\n\nif('swin_tiny_patch4_window7_224_22k.pth' in cfg['pretrained_original']['checkpoint']):\n    initial_checkpoint = None\nelse:\n    initial_checkpoint = cfg['pretrained_original']['checkpoint']\n    initial_checkpoint_stained = cfg['pretrained_original']['checkpoint_stained']\n# initial_checkpoint = cfg['pretrained_original']['checkpoint']\n\nstart_lr   = 5e-5 #0.0001\nbatch_size = 8 #32 #32\n\n### setup  ----------------------------------------\nfor f in ['checkpoint','train','valid','backup'] : os.makedirs(out_dir +'/'+f, exist_ok=True)\n\n### net ----------------------------------------\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## optimiser ----------------------------------\n# if 0: ##freeze\nif 1: ##freeze\n#     for p in net.stem.parameters():   p.requires_grad = False\n    for p in net.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            \nfreeze_bn(net)\n\n#-----------------------------------------------\n\n# optimizer = torch.optim.AdamW(filter(lambda p: p.requires_grad, net.parameters()),lr=start_lr)\n\nif(mode_ == 'TRAIN'):\n#     num_iteration = 1000*len(train_loader)\n    num_iteration = 45*len(train_loader)\n    iter_log   = len(train_loader)*3 #479\nelse:\n    num_iteration = 1\n#     iter_log   = len(test_loader)*1\n    iter_log   = len(test_HPA_loader)*1\niter_valid = iter_log\niter_save  = iter_log\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#-----------------------------------------------\n# valid_loss = np.zeros(4,np.float32)\n# train_loss = np.zeros(2,np.float32)\n# batch_loss = np.zeros_like(train_loss)\n# sum_train_loss = np.zeros_like(train_loss)\n# sum_train = 0\n\nstart_timer = time.time()\niteration = start_iteration\nepoch = start_epoch\nrate = 0\n\n# submission_df = pd.read_csv(SUBMIT).fillna('')\n# submission_df['id'] = submission_df['id'].astype(int)\n\ncols = ['id', 'rle']\nsubmission_df = pd.DataFrame(index=[], columns=cols)\n# test_df_ = pd.read_csv(LABELS)\n\nimport matplotlib.pyplot as plt\n# import gc\n\n# while iteration < (start_iteration+num_iteration):\nif(mode_ == 'TRAIN'):\n    for t, batch in enumerate(train_loader):\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/fold4_stained_%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        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        # learning rate schduler ------------\n        rate = get_learning_rate(optimizer)\n\n        # one iteration update  -------------\n        batch_size = len(batch['index'])\n        batch['image'] = batch['image'].half().cuda()\n        batch['mask' ] = batch['mask' ].half().cuda()\n        batch['organ'] = batch['organ'].cuda()\n\n        net.train()\n        net.output_type = ['loss']\n        if 1:\n            with amp.autocast(enabled = is_amp):\n                output = net(batch)\n                loss0  = output['bce_loss'].mean()\n                loss1  = output['aux2_loss'].mean()\n\n            optimizer.zero_grad()\n            scaler.scale(loss0+0.2*loss1).backward()\n\n            scaler.unscale_(optimizer)\n            scaler.step(optimizer)\n            scaler.update()\n\n        # print statistics  --------\n        batch_loss[:2] = [loss0.item(),loss1.item()]\n        sum_train_loss += batch_loss\n        sum_train += 1\n        if t % 100 == 0:\n            train_loss = sum_train_loss / (sum_train + 1e-12)\n            sum_train_loss[...] = 0\n            sum_train = 0\n\n        print('\\r', end='', flush=True)\n        print(message(mode='print'), end='', flush=True)\n        epoch += 1 / len(train_loader)\n        iteration += 1\n\n    torch.cuda.empty_cache()\n\nelse:\n#     print(test_HPA_df)\n    if(test_HPA_df['id'].size>0):\n        \n        submission_df = test(net, test_HPA_loader, submission_df)\n        \n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"###############################\n### repeat again for Hubmap ###\n###############################\n\n# del scaler\n# del net\n# gc.collect()\n# torch.cuda.empty_cache()\n\n# scaler = amp.GradScaler(enabled = is_amp)\n# net = 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_stained()\n\n## optimiser ----------------------------------\n# if 0: ##freeze\nif 1: ##freeze\n#     for p in net.stem.parameters():   p.requires_grad = False\n    for p in net.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            \nfreeze_bn(net)\n\n#-----------------------------------------------\n\n# optimizer = torch.optim.AdamW(filter(lambda p: p.requires_grad, net.parameters()),lr=start_lr)\n\nif(mode_ == 'TRAIN'):\n#     num_iteration = 1000*len(train_loader)\n    num_iteration = 45*len(train_loader)\n    iter_log   = len(train_loader)*3 #479\nelse:\n    num_iteration = 1\n#     iter_log   = len(test_loader)*1\n    iter_log   = len(test_Hubmap_loader)*1\niter_valid = iter_log\niter_save  = iter_log\n\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if(test_Hubmap_df['id'].size>0):\n#     print(test_Hubmap_df)\n    submission_df = test(net, test_Hubmap_loader, submission_df)\n        \n########## reorder as test.csv ########## \nl_order = test_df['id'].astype(str).values.tolist()\nsubmission_df['order'] = submission_df['id'].astype(str).apply(lambda x: l_order.index(x) if x in l_order else -1)\n# print(submission_df)\nsubmission_df = submission_df.sort_values('order')\nsubmission_df = submission_df.drop(\"order\", axis=1)\n# print(submission_df)\nsubmission_df['id'] = submission_df['id'].astype(str)\nsubmission_df.to_csv('submission.csv',index=False)        \n","metadata":{},"execution_count":null,"outputs":[]}]}