{"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 numpy as np\nimport pandas as pd\nimport os, glob\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\nfrom skimage import color\nimport cv2\nfrom torch.utils.data import Dataset, DataLoader\nimport pandas as pd\nimport numpy as np\nimport os\nimport cv2\nimport gc\nimport random\nfrom albumentations import *\nfrom sklearn.model_selection import KFold\nimport matplotlib.pyplot as plt\nfrom lovasz import lovasz_hinge\nfrom fastai.vision.all import *\n\n# unet model\n# from unet_models import R2AttU_Net\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nSEED = 2020\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    #the following line gives ~10% speedup\n    #but may lead to some stochasticity in the results \n    torch.backends.cudnn.benchmark = True\n    \nseed_everything(SEED)","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:03.980179Z","iopub.execute_input":"2022-11-05T07:01:03.981214Z","iopub.status.idle":"2022-11-05T07:01:07.856762Z","shell.execute_reply.started":"2022-11-05T07:01:03.981054Z","shell.execute_reply":"2022-11-05T07:01:07.855755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\nconfig = {\n    'split_seed_list':[0],\n#     'FOLD_LIST':[0,1,2,3], \n#     'model_path':'../input/hubmap-new-03-03/',\n    'model_name':'seresnext101',\n    'num_classes':1,\n#     'resolution':1024, #(1024,1024),(512,512),\n#     'input_resolution':320, #(320,320), #(256,256), #(512,512), #(384,384)\n    'deepsupervision':True, # always false for inference\n    'clfhead':False,\n    'clf_threshold':0.5,\n    'small_mask_threshold':0, #256*256*0.03, #512*512*0.03,\n    'mask_threshold':0.5,\n#     'pad_size':256, #(64,64), #(256,256), #(128,128)\n    \n#     'tta':3,\n#     'test_batch_size':12,\n    \n#     'FP16':False,\n    'num_workers':4,\n    'device':torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"),\n}\n\ndevice = config['device']","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:20.626058Z","iopub.execute_input":"2022-11-05T07:01:20.626949Z","iopub.status.idle":"2022-11-05T07:01:20.718684Z","shell.execute_reply.started":"2022-11-05T07:01:20.626911Z","shell.execute_reply":"2022-11-05T07:01:20.717559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Process the raw images and annotations to create a training dataset\nThis competition is about developing a segmentation problem to identify regions with glomeruli in human kidney tissue images.\n* train set: 15 images\n* (public) test set: 5 images\n* (private) test set: idk haha need to check... more than 5 images for sure\n* Each of them has ~50k pixel size and is saved as a high-resolution tiff image.\n\nThe training set and public test set also includes annotations in both RLE-encoded and unencoded (JSON) forms (can use either one), denoting segmentations of glomeruli.\n\n**The expected prediction for a given image is an RLE-encoded mask containing ALL objects in the image (i.e. not just glomeruli). The mask, as mentioned in the Evaluation page, should be binary when encoded - with 0 indicating the lack of a masked pixel, and 1 indicating a masked pixel.**","metadata":{}},{"cell_type":"code","source":"DATASET_FOLDER = \"/kaggle/input/hubmap-kidney-segmentation\"\nTRAIN_DATA = '../input/hubmap-kidney-segmentation/train/'\nLABELS = '../input/hubmap-kidney-segmentation/train.csv'\nos.listdir(DATASET_FOLDER)","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:21.845378Z","iopub.execute_input":"2022-11-05T07:01:21.845756Z","iopub.status.idle":"2022-11-05T07:01:21.856846Z","shell.execute_reply.started":"2022-11-05T07:01:21.845724Z","shell.execute_reply":"2022-11-05T07:01:21.855765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_listdir = os.listdir(os.path.join(DATASET_FOLDER, \"train\")) \n# TODO whats the difference between anatomical-strcture.json and the other .json?\ntrain_listdir.sort()\nprint(train_listdir)","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:22.134618Z","iopub.execute_input":"2022-11-05T07:01:22.134972Z","iopub.status.idle":"2022-11-05T07:01:22.161254Z","shell.execute_reply.started":"2022-11-05T07:01:22.134942Z","shell.execute_reply":"2022-11-05T07:01:22.160079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Extra data source \n\nCan also get extra data from another competition: https://www.kaggle.com/competitions/hubmap-organ-segmentation/data. Since this dataset includes 5 different organs, we can filter to only the images that show \"kidney\" (since the organ type is recorded in train.csv)\n\nTODO: Not sure whether we can use these images for training as well as a form of data augmentation?","metadata":{}},{"cell_type":"code","source":"DATASET_FOLDER_2 = \"/kaggle/input/hubmap-organ-segmentation/\"\nimage_paths_2 = sorted(glob.glob(os.path.join(DATASET_FOLDER_2, 'train_images/*.tiff')))\nprint(f\"Number of images: {len(image_paths_2)}\")","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:22.649920Z","iopub.execute_input":"2022-11-05T07:01:22.650359Z","iopub.status.idle":"2022-11-05T07:01:22.773988Z","shell.execute_reply.started":"2022-11-05T07:01:22.650318Z","shell.execute_reply":"2022-11-05T07:01:22.772930Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Code by Ihelon https://www.kaggle.com/code/ihelon/illustrations-kumapi390-eda\n\nn_cols = 4\n\nfor ind, image_path in enumerate(image_paths_2[:4]): # show the first 4 example images\n    image = cv2.imread(image_path)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    if ind % n_cols == 0:\n        plt.figure(figsize=(16, 5))\n    plt.subplot(1, n_cols, ind % n_cols + 1)\n    plt.imshow(image)\n    plt.axis(\"off\")\n    if ind % n_cols == n_cols - 1:\n        plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:22.896952Z","iopub.execute_input":"2022-11-05T07:01:22.897829Z","iopub.status.idle":"2022-11-05T07:01:28.948682Z","shell.execute_reply.started":"2022-11-05T07:01:22.897780Z","shell.execute_reply":"2022-11-05T07:01:28.947449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Obtaining the masks using the json annotations (i.e. option 1)","metadata":{}},{"cell_type":"code","source":"import json\nwith open(os.path.join(TRAIN_DATA, \"0486052bb.json\"), \"r\") as f: \n    # this is the json annotation file for one of the TIFF files (i.e. one of the medical iamges)\n    parsed = json.load(f)","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:28.950294Z","iopub.execute_input":"2022-11-05T07:01:28.950677Z","iopub.status.idle":"2022-11-05T07:01:28.970880Z","shell.execute_reply.started":"2022-11-05T07:01:28.950642Z","shell.execute_reply":"2022-11-05T07:01:28.970218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(parsed))\nfor p in parsed: # parsed is a json list, p is a dict\n    print(p.keys())\n    break # same for all 130 dictionaries in the json list","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:28.972365Z","iopub.execute_input":"2022-11-05T07:01:28.973013Z","iopub.status.idle":"2022-11-05T07:01:28.979214Z","shell.execute_reply.started":"2022-11-05T07:01:28.972976Z","shell.execute_reply":"2022-11-05T07:01:28.978115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* Type and id are useless (same for all TIFF files)\n* geometry: contains a Polygon with coordinates for the feature's enclosing volume\n* properties: including the name and color of the feature in the image.","metadata":{}},{"cell_type":"code","source":"parsed[0][\"type\"] # means that the annotation is referring to a \"feature\" in the medical image \n# same for all parsed[i][\"type\"] for all i in the json list of 130 dicts","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:28.982347Z","iopub.execute_input":"2022-11-05T07:01:28.983120Z","iopub.status.idle":"2022-11-05T07:01:28.991946Z","shell.execute_reply.started":"2022-11-05T07:01:28.983083Z","shell.execute_reply":"2022-11-05T07:01:28.990973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"parsed[0][\"id\"] # means that the annotation contains a Polygon (i think)\n# same for all parsed[i][\"id\"] for all i in the json list of 130 dicts","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:28.993579Z","iopub.execute_input":"2022-11-05T07:01:28.994293Z","iopub.status.idle":"2022-11-05T07:01:29.004489Z","shell.execute_reply.started":"2022-11-05T07:01:28.994258Z","shell.execute_reply":"2022-11-05T07:01:29.003637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(parsed[0][\"geometry\"])","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:29.005915Z","iopub.execute_input":"2022-11-05T07:01:29.006605Z","iopub.status.idle":"2022-11-05T07:01:29.016686Z","shell.execute_reply.started":"2022-11-05T07:01:29.006572Z","shell.execute_reply":"2022-11-05T07:01:29.015925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"How do you use a Polygon? I think its a series of <x,y> points that you can draw on the image to segment out a section of the image","metadata":{}},{"cell_type":"code","source":"parsed[0][\"properties\"]\n# same for all parsed[i][\"type\"] for all i in the json list of 130 dicts","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:29.018462Z","iopub.execute_input":"2022-11-05T07:01:29.019242Z","iopub.status.idle":"2022-11-05T07:01:29.028505Z","shell.execute_reply.started":"2022-11-05T07:01:29.019203Z","shell.execute_reply":"2022-11-05T07:01:29.027439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"There are over 600,000 glomeruli in each human kidney (Nyengaard, 1992). Normal glomeruli typically range from 100-350μm in diameter with a roughly spherical shape (Kannan, 2019).\n\nTeams are invited to develop segmentation algorithms that identify glomeruli in the PAS stained microscopy data.","metadata":{}},{"cell_type":"markdown","source":"## Obtaining the masks using the RLE annotations in the dataframe (i.e. option 2) \nThe RLE Compression algorithm is the compression used in the TIFF 6.0.","metadata":{}},{"cell_type":"code","source":"df_masks = pd.read_csv(LABELS).set_index('id')\ndf_masks.head()","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:29.030024Z","iopub.execute_input":"2022-11-05T07:01:29.030798Z","iopub.status.idle":"2022-11-05T07:01:29.442132Z","shell.execute_reply.started":"2022-11-05T07:01:29.030761Z","shell.execute_reply":"2022-11-05T07:01:29.441087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* id is the tiff image file name\n* encoding is a long-ass string that gives a mask of the image (i.e. masking out all the parts of the kidney that are not the glomeruli)","metadata":{}},{"cell_type":"code","source":"#functions to convert encoding to mask and mask to encoding\ndef enc2mask(encs, shape):\n    '''\n    mask_rle: run-length as string formated (start length)\n    img_shape: (width,height) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n    '''\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for m,enc in enumerate(encs):\n        if isinstance(enc,np.float) and np.isnan(enc): continue\n        s = enc.split()\n        for i in range(len(s)//2):\n            start = int(s[2*i]) - 1\n            length = int(s[2*i+1])\n            img[start:start+length] = 1 + m\n    return img.reshape(shape).T\ndef mask2enc(mask, n=1):\n    pixels = mask.T.flatten()\n    encs = []\n    for i in range(1,n+1):\n        p = (pixels == i).astype(np.int8)\n        if p.sum() == 0: encs.append(np.nan)\n        else:\n            p = np.concatenate([[0], p, [0]])\n            runs = np.where(p[1:] != p[:-1])[0] + 1\n            runs[1::2] -= runs[::2]\n            encs.append(' '.join(str(x) for x in runs))\n    return encs","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:29.443684Z","iopub.execute_input":"2022-11-05T07:01:29.444085Z","iopub.status.idle":"2022-11-05T07:01:29.454272Z","shell.execute_reply.started":"2022-11-05T07:01:29.444052Z","shell.execute_reply":"2022-11-05T07:01:29.453046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Save the tiff images and the RLE encoding masks as 256x256 pngs, to create training dataset","metadata":{}},{"cell_type":"markdown","source":"# Now that we have the 256x256 images and coresponding masks as png files, we can load this data as a pytorch Dataset object","metadata":{}},{"cell_type":"code","source":"TRAIN_IMAGES = '../input/hubmap-256x256/train/'\nMASK_IMAGES = '../input/hubmap-256x256/masks/'\n\nNUM_WORKERS = 4 # for pytorch DataLoader\n\nNFOLDS = 2","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:29.457588Z","iopub.execute_input":"2022-11-05T07:01:29.458154Z","iopub.status.idle":"2022-11-05T07:01:29.469950Z","shell.execute_reply.started":"2022-11-05T07:01:29.458119Z","shell.execute_reply":"2022-11-05T07:01:29.469019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://www.kaggle.com/iafoss/256x256-images\ndef img2tensor(img,dtype:np.dtype=np.float32):\n    if img.ndim==2 : img = np.expand_dims(img,2)\n    img = np.transpose(img,(2,0,1))\n    return torch.from_numpy(img.astype(dtype, copy=False))\n\n\nmean = np.array([0.65459856,0.48386562,0.69428385])\nstd = np.array([0.15167958,0.23584107,0.13146145])\n\nclass HuBMAPDataset(Dataset): # torch.utils.data.Dataset object\n    def __init__(self, fold=0, train=True, tfms=None):\n        ids = pd.read_csv(LABELS).id.values\n        print(f\"Current fold: {fold}/{NFOLDS - 1}\")\n        kf = KFold(n_splits=NFOLDS,random_state=SEED,shuffle=True)\n        ids = set(ids[list(kf.split(ids))[fold][0 if train else 1]])\n        self.fnames = [fname for fname in os.listdir(TRAIN_IMAGES) if fname.split('_')[0] in ids]\n        self.train = train\n        self.tfms = tfms\n        \n    def __len__(self):\n        return len(self.fnames)\n    \n    def __getitem__(self, idx):\n        fname = self.fnames[idx]\n        img = cv2.cvtColor(cv2.imread(os.path.join(TRAIN_IMAGES,fname)), cv2.COLOR_BGR2RGB)\n        mask = cv2.imread(os.path.join(MASK_IMAGES,fname),cv2.IMREAD_GRAYSCALE)\n        if self.tfms is not None:\n            augmented = self.tfms(image=img,mask=mask)\n            img,mask = augmented['image'],augmented['mask']\n        return img2tensor((img/255.0 - mean)/std),img2tensor(mask)\n    \n# Get augementations\ndef get_aug(p=1.0):\n    return Compose([\n        HorizontalFlip(),\n        VerticalFlip(),\n        RandomRotate90(),\n        ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, \n                         border_mode=cv2.BORDER_REFLECT),\n        OneOf([\n            OpticalDistortion(p=0.3),\n            GridDistortion(p=.1),\n            IAAPiecewiseAffine(p=0.3),\n        ], p=0.3),\n        OneOf([\n            HueSaturationValue(10,15,10),\n            CLAHE(clip_limit=2),\n            RandomBrightnessContrast(),            \n        ], p=0.3),\n    ], p=p)","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:30.303675Z","iopub.execute_input":"2022-11-05T07:01:30.304063Z","iopub.status.idle":"2022-11-05T07:01:30.320179Z","shell.execute_reply.started":"2022-11-05T07:01:30.304030Z","shell.execute_reply":"2022-11-05T07:01:30.318918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#example of train images with masks\nds = HuBMAPDataset(tfms=get_aug())\ndl = DataLoader(ds,batch_size=64,shuffle=False,num_workers=NUM_WORKERS) # torch.utils.data.Dataloader object\nimgs,masks = next(iter(dl))\n\nplt.figure(figsize=(16,16))\nfor i,(img,mask) in enumerate(zip(imgs,masks)):\n    img = ((img.permute(1,2,0)*std + mean)*255.0).numpy().astype(np.uint8)\n    plt.subplot(8,8,i+1)\n    plt.imshow(img,vmin=0,vmax=255)\n    plt.imshow(mask.squeeze().numpy(), alpha=0.2)\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\n    \ndel ds,dl,imgs,masks","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:01:30.820947Z","iopub.execute_input":"2022-11-05T07:01:30.821694Z","iopub.status.idle":"2022-11-05T07:01:55.102742Z","shell.execute_reply.started":"2022-11-05T07:01:30.821657Z","shell.execute_reply":"2022-11-05T07:01:55.101745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Modelling\n## Model option 1: UNet with fastai (TODO: implement attention gates)\nhttps://www.kaggle.com/code/phoebezhouhuixin/hubmap-pytorch-fast-ai-starter/edit  \nAttention gates: https://github.com/tureckova/Abdomen-CT-Image-Segmentation\n\n### What is UneXt50?\nThe model used in this kernel is based on a U-shape network (UneXt50, see image below), which I used in Severstal and Understanding Clouds competitions. The idea of a U-shape network is coming from a [Unet](https://arxiv.org/pdf/1505.04597.pdf) architecture proposed in 2015 for medical images: **the encoder part creates a representation of features at different levels, while the decoder combines the features and generates a prediction as a segmentation mask.** The skip connections between encoder and decoder allow us to utilize features from the intermediate conv layers of the encoder effectively, without a need for the information to go the full way through entire encoder and decoder. This is especially important to link the predicted mask to the specific pixels of the detected object. Later people realized that **ImageNet pretrained computer vision models could drastically improve the quality of a segmentation model** because of optimized architecture of the encoder, high encoder capacity (in contrast to one used in the original Unet), and the power of the transfer learning.\n\nThere are several important things that must be added to a Unet network, however, to make it able to reach competitive results with current state of the art approaches. First, it is **Feature Pyramid Network (FPN)**: Provides additional skip connection between different upscaling blocks of the decoder and the output layer. So, the **final prediction is produced based on the concatenation of U-net output with resized outputs of the intermediate layers**. These skip-connections provide a shortcut for gradient flow improving model performance and convergence speed. Since intermediate layers have many channels, their upscaling and use as an input for the final layer would introduce a significant overhead in terms of the computational time and memory. Therefore, 3x3+3x3 convolutions are applied (factorization) before the resize to reduce the number of channels.\n\nAnother very important thing is the **Atrous Spatial Pyramid Pooling (ASPP) block** added between encoder and decoder. The flaw of the traditional U-shape networks is resulted by a small receptive field. Therefore, if a model needs to make a decision about a segmentation of a large object, especially for a large image resolution, it can get confused being able to look only into parts of the object. A way to increase the receptive field and enable interactions between different parts of the image is use of a block combining convolutions with different dilations ([Atrous convolutions](https://arxiv.org/pdf/1606.00915.pdf) with various rates in ASPP block). While the original paper uses 6,12,18 rates, they may be customized for a particular task and a particular image resolution to maximize the performance. One more thing I added is using group convolutions in ASPP block to reduce the number of model parameters.\n\nFinally, the decoder upscaling blocks are based on [pixel shuffle](https://arxiv.org/pdf/1609.05158.pdf) rather than transposed convolution used in the first Unet models. It allows to avoid artifacts in the produced masks. And I use [semisupervised Imagenet pretrained ResNeXt50](https://github.com/facebookresearch/semi-supervised-ImageNet1K-models) model as a backbone. In Pytorch it provides the performance of EfficientNet B2-B3 with much faster convergence for the computational cost and GPU RAM requirements of EfficientNet B0 (though, in TF EfficientNet is highly optimized and may be a good thing to use).\n\n![](https://i.ibb.co/z5KxDzm/Une-Xt50-1.png)","metadata":{}},{"cell_type":"code","source":"# class FPN(nn.Module):\n#     def __init__(self, input_channels:list, output_channels:list):\n#         super().__init__()\n#         self.convs = nn.ModuleList(\n#             [nn.Sequential(nn.Conv2d(in_ch, out_ch*2, kernel_size=3, padding=1),\n#              nn.ReLU(inplace=True), nn.BatchNorm2d(out_ch*2),\n#              nn.Conv2d(out_ch*2, out_ch, kernel_size=3, padding=1))\n#             for in_ch, out_ch in zip(input_channels, output_channels)])\n        \n#     def forward(self, xs:list, last_layer):\n#         hcs = [F.interpolate(c(x),scale_factor=2**(len(self.convs)-i),mode='bilinear') \n#                for i,(c,x) in enumerate(zip(self.convs, xs))]\n#         hcs.append(last_layer)\n#         return torch.cat(hcs, dim=1)\n\n# class UnetBlock(Module): # decoder block\n#     def __init__(self, up_in_c:int, x_in_c:int, nf:int=None, blur:bool=False,\n#                  self_attention:bool=False, **kwargs):\n#         super().__init__()\n#         self.shuf = PixelShuffle_ICNR(up_in_c, up_in_c//2, blur=blur, **kwargs)\n#         self.bn = nn.BatchNorm2d(x_in_c)\n#         ni = up_in_c//2 + x_in_c\n#         nf = nf if nf is not None else max(up_in_c//2,32)\n#         self.conv1 = ConvLayer(ni, nf, norm_type=None, **kwargs)\n#         self.conv2 = ConvLayer(nf, nf, norm_type=None,\n#             xtra=SelfAttention(nf) if self_attention else None, **kwargs)\n#         self.relu = nn.ReLU(inplace=True)\n\n#     def forward(self, up_in:Tensor, left_in:Tensor) -> Tensor:\n#         s = left_in\n#         up_out = self.shuf(up_in)\n#         cat_x = self.relu(torch.cat([up_out, self.bn(s)], dim=1))\n#         return self.conv2(self.conv1(cat_x))\n        \n# class _ASPPModule(nn.Module):\n#     def __init__(self, inplanes, planes, kernel_size, padding, dilation, groups=1):\n#         super().__init__()\n#         self.atrous_conv = nn.Conv2d(inplanes, planes, kernel_size=kernel_size,\n#                 stride=1, padding=padding, dilation=dilation, bias=False, groups=groups)\n#         self.bn = nn.BatchNorm2d(planes)\n#         self.relu = nn.ReLU()\n\n#         self._init_weight()\n\n#     def forward(self, x):\n#         x = self.atrous_conv(x)\n#         x = self.bn(x)\n\n#         return self.relu(x)\n\n#     def _init_weight(self):\n#         for m in self.modules():\n#             if isinstance(m, nn.Conv2d):\n#                 torch.nn.init.kaiming_normal_(m.weight)\n#             elif isinstance(m, nn.BatchNorm2d):\n#                 m.weight.data.fill_(1)\n#                 m.bias.data.zero_()\n\n# class ASPP(nn.Module):\n#     def __init__(self, inplanes=512, mid_c=256, dilations=[6, 12, 18, 24], out_c=None):\n#         super().__init__()\n#         self.aspps = [_ASPPModule(inplanes, mid_c, 1, padding=0, dilation=1)] + \\\n#             [_ASPPModule(inplanes, mid_c, 3, padding=d, dilation=d,groups=4) for d in dilations]\n#         self.aspps = nn.ModuleList(self.aspps)\n#         self.global_pool = nn.Sequential(nn.AdaptiveMaxPool2d((1, 1)),\n#                         nn.Conv2d(inplanes, mid_c, 1, stride=1, bias=False),\n#                         nn.BatchNorm2d(mid_c), nn.ReLU())\n#         out_c = out_c if out_c is not None else mid_c\n#         self.out_conv = nn.Sequential(nn.Conv2d(mid_c*(2+len(dilations)), out_c, 1, bias=False),\n#                                     nn.BatchNorm2d(out_c), nn.ReLU(inplace=True))\n#         self.conv1 = nn.Conv2d(mid_c*(2+len(dilations)), out_c, 1, bias=False)\n#         self._init_weight()\n\n#     def forward(self, x):\n#         x0 = self.global_pool(x)\n#         xs = [aspp(x) for aspp in self.aspps]\n#         x0 = F.interpolate(x0, size=xs[0].size()[2:], mode='bilinear', align_corners=True)\n#         x = torch.cat([x0] + xs, dim=1)\n#         return self.out_conv(x)\n    \n#     def _init_weight(self):\n#         for m in self.modules():\n#             if isinstance(m, nn.Conv2d):\n#                 torch.nn.init.kaiming_normal_(m.weight)\n#             elif isinstance(m, nn.BatchNorm2d):\n#                 m.weight.data.fill_(1)\n#                 m.bias.data.zero_()\n\n# class UneXt50(nn.Module):\n#     def __init__(self, stride=1, **kwargs):\n#         super().__init__()\n#         #encoder\n#         m = torch.hub.load('facebookresearch/semi-supervised-ImageNet1K-models',\n#                            'resnext50_32x4d_ssl') # change this\n#         self.enc0 = nn.Sequential(m.conv1, m.bn1, nn.ReLU(inplace=True))\n#         self.enc1 = nn.Sequential(nn.MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1),\n#                             m.layer1) #256\n#         self.enc2 = m.layer2 #512\n#         self.enc3 = m.layer3 #1024\n#         self.enc4 = m.layer4 #2048\n#         #aspp with customized dilatations\n#         self.aspp = ASPP(2048,256,out_c=512,dilations=[stride*1,stride*2,stride*3,stride*4])\n#         self.drop_aspp = nn.Dropout2d(0.5)\n#         #decoder\n#         self.dec4 = UnetBlock(512,1024,256)\n#         self.dec3 = UnetBlock(256,512,128)\n#         self.dec2 = UnetBlock(128,256,64)\n#         self.dec1 = UnetBlock(64,64,32)\n#         self.fpn = FPN([512,256,128,64],[16]*4)\n#         self.drop = nn.Dropout2d(0.1)\n#         self.final_conv = ConvLayer(32+16*4, 1, ks=1, norm_type=None, act_cls=None)\n        \n#     def forward(self, x):\n#         enc0 = self.enc0(x)\n#         enc1 = self.enc1(enc0)\n#         enc2 = self.enc2(enc1)\n#         enc3 = self.enc3(enc2)\n#         enc4 = self.enc4(enc3)\n#         enc5 = self.aspp(enc4)\n#         dec3 = self.dec4(self.drop_aspp(enc5),enc3) # this is the input to self.dec3\n#         dec2 = self.dec3(dec3,enc2) # this is the input to self.dec2\n#         dec1 = self.dec2(dec2,enc1) # this is the input to self.dec1\n#         dec0 = self.dec1(dec1,enc0) # this is the input to FPN (as the `last_layer` argument)\n#         x = self.fpn([enc5, dec3, dec2, dec1], dec0)\n#         x = self.final_conv(self.drop(x))\n#         x = F.interpolate(x,scale_factor=2,mode='bilinear')\n#         return x\n\n# #split the model to encoder and decoder for fast.ai\n# split_layers = lambda m: [list(m.enc0.parameters())+list(m.enc1.parameters())+\n#                 list(m.enc2.parameters())+list(m.enc3.parameters())+\n#                 list(m.enc4.parameters()),\n                          \n#                 list(m.aspp.parameters())+list(m.dec4.parameters())+\n#                 list(m.dec3.parameters())+list(m.dec2.parameters())+\n#                 list(m.dec1.parameters())+list(m.fpn.parameters())+\n#                 list(m.final_conv.parameters())]","metadata":{"execution":{"iopub.status.busy":"2022-11-05T05:19:01.344310Z","iopub.execute_input":"2022-11-05T05:19:01.344808Z","iopub.status.idle":"2022-11-05T05:19:01.362389Z","shell.execute_reply.started":"2022-11-05T05:19:01.344771Z","shell.execute_reply":"2022-11-05T05:19:01.361068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Number 1 solution model (seresnext)","metadata":{}},{"cell_type":"code","source":"# !wget http://data.lip6.fr/cadene/pretrainedmodels/se_resnext101_32x4d-3b2fe3d8.pth  --no-check-certificate","metadata":{"execution":{"iopub.status.busy":"2022-11-05T05:43:30.714528Z","iopub.execute_input":"2022-11-05T05:43:30.715024Z","iopub.status.idle":"2022-11-05T05:52:43.940827Z","shell.execute_reply.started":"2022-11-05T05:43:30.714980Z","shell.execute_reply":"2022-11-05T05:52:43.939531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn, optim\nimport torch.nn.functional as F\nimport sys\npackage_dir = \"../input/pretrainedmodels/pretrained-models.pytorch-master/\"\nsys.path.insert(0, package_dir)\nimport pretrainedmodels\n\n\ndef conv3x3(in_channel, out_channel): #not change resolusion\n    return nn.Conv2d(in_channel,out_channel,\n                      kernel_size=3,stride=1,padding=1,dilation=1,bias=False)\n\ndef conv1x1(in_channel, out_channel): #not change resolution\n    return nn.Conv2d(in_channel,out_channel,\n                      kernel_size=1,stride=1,padding=0,dilation=1,bias=False)\n\ndef init_weight(m):\n    classname = m.__class__.__name__\n    if classname.find('Conv') != -1:\n        #nn.init.xavier_uniform_(m.weight, gain=1)\n        #nn.init.xavier_normal_(m.weight, gain=1)\n        #nn.init.kaiming_uniform_(m.weight, mode='fan_in', nonlinearity='relu')\n        nn.init.kaiming_normal_(m.weight, mode='fan_in', nonlinearity='relu')\n        #nn.init.orthogonal_(m.weight, gain=1)\n        if m.bias is not None:\n            m.bias.data.zero_()\n    elif classname.find('Batch') != -1:\n        m.weight.data.normal_(1,0.02)\n        m.bias.data.zero_()\n    elif classname.find('Linear') != -1:\n        nn.init.orthogonal_(m.weight, gain=1)\n        if m.bias is not None:\n            m.bias.data.zero_()\n    elif classname.find('Embedding') != -1:\n        nn.init.orthogonal_(m.weight, gain=1)\n\n        \nclass cSEBlock(nn.Module):\n    def __init__(self, c, feat):\n        super().__init__()\n        self.attention_fc = nn.Linear(feat,1, bias=False)\n        self.bias         = nn.Parameter(torch.zeros((1,c,1), requires_grad=True))\n        self.sigmoid      = nn.Sigmoid()\n        self.dropout      = nn.Dropout2d(0.1)\n        \n    def forward(self,inputs):\n        batch,c,h,w = inputs.size()\n        x = inputs.view(batch,c,-1)\n        x = self.attention_fc(x) + self.bias\n        x = x.view(batch,c,1,1)\n        x = self.sigmoid(x)\n        x = self.dropout(x)\n        return inputs * x\n\nclass sSEBlock(nn.Module):\n    def __init__(self, c, h, w):\n        super().__init__()\n        self.attention_fc = nn.Linear(c,1, bias=False).apply(init_weight)\n        self.bias         = nn.Parameter(torch.zeros((1,h,w,1), requires_grad=True))\n        self.sigmoid      = nn.Sigmoid()\n        \n    def forward(self,inputs):\n        batch,c,h,w = inputs.size()\n        x = torch.transpose(inputs, 1,2) #(*,c,h,w)->(*,h,c,w)\n        x = torch.transpose(x, 2,3) #(*,h,c,w)->(*,h,w,c)\n        x = self.attention_fc(x) + self.bias\n        x = torch.transpose(x, 2,3) #(*,h,w,1)->(*,h,1,w)\n        x = torch.transpose(x, 1,2) #(*,h,1,w)->(*,1,h,w)\n        x = self.sigmoid(x)\n        return inputs * x\n    \nclass scSEBlock(nn.Module):\n    def __init__(self, c, h, w):\n        super().__init__()\n        self.cSE = cSEBlock(c,h*w)\n        self.sSE = sSEBlock(c,h,w)\n    \n    def forward(self, inputs):\n        x1 = self.cSE(inputs)\n        x2 = self.sSE(inputs)\n        return x1+x2\n    \n    \n# class SpatialAttention2d(nn.Module):\n#     def __init__(self, in_channel):\n#         super().__init__()\n#         self.squeeze = conv1x1(in_channel,1).apply(init_weight)\n#         self.sigmoid = nn.Sigmoid()\n        \n#     def forward(self, inputs):\n#         x = self.squeeze(inputs)\n#         x = self.sigmoid(x)\n#         return inputs * x\n    \n    \n# class GAB(nn.Module):\n#     def __init__(self, in_channel, reduction=4):\n#         super().__init__()\n#         self.global_avgpool = nn.AdaptiveAvgPool2d(1)\n#         self.conv1 = conv1x1(in_channel, in_channel//reduction).apply(init_weight)\n#         self.conv2 = conv1x1(in_channel//reduction, in_channel).apply(init_weight)\n#         self.relu  = nn.ReLU(True)\n#         self.sigmoid = nn.Sigmoid()\n        \n#     def forward(self, inputs):\n#         x = self.global_avgpool(inputs)\n#         x = self.relu(self.conv1(x))\n#         x = self.sigmoid(self.conv2(x))\n#         return inputs * x\n\n    \n# class scSEBlock2(nn.Module):\n#     def __init__(self, in_channel, reduction=4):\n#         super().__init__()\n#         self.cSE = GAB(in_channel, reduction)\n#         self.sSE = SpatialAttention2d(in_channel)\n    \n#     def forward(self, inputs):\n#         x1 = self.cSE(inputs)\n#         x2 = self.sSE(inputs)\n#         return x1+x2\n    \n\nclass Attention(nn.Module):\n    def __init__(self, channels):\n        super().__init__()\n        self.channels = channels\n        self.theta    = nn.utils.spectral_norm(conv1x1(channels, channels//8)).apply(init_weight)\n        self.phi      = nn.utils.spectral_norm(conv1x1(channels, channels//8)).apply(init_weight)\n        self.g        = nn.utils.spectral_norm(conv1x1(channels, channels//2)).apply(init_weight)\n        self.o        = nn.utils.spectral_norm(conv1x1(channels//2, channels)).apply(init_weight)\n        self.gamma    = nn.Parameter(torch.tensor(0.), requires_grad=True)\n        \n    def forward(self, inputs):\n        batch,c,h,w = inputs.size()\n        theta = self.theta(inputs) #->(*,c/8,h,w)\n        phi   = F.max_pool2d(self.phi(inputs), [2,2]) #->(*,c/8,h/2,w/2)\n        g     = F.max_pool2d(self.g(inputs), [2,2]) #->(*,c/2,h/2,w/2)\n        \n        theta = theta.view(batch, self.channels//8, -1) #->(*,c/8,h*w)\n        phi   = phi.view(batch, self.channels//8, -1) #->(*,c/8,h*w/4)\n        g     = g.view(batch, self.channels//2, -1) #->(*,c/2,h*w/4)\n        \n        beta = F.softmax(torch.bmm(theta.transpose(1,2), phi), -1) #->(*,h*w,h*w/4)\n        o    = self.o(torch.bmm(g, beta.transpose(1,2)).view(batch,self.channels//2,h,w)) #->(*,c,h,w)\n        return self.gamma*o + inputs\n    \n    \n\nclass ChannelAttentionModule(nn.Module):\n    def __init__(self, in_channel, reduction):\n        super().__init__()\n        self.global_maxpool = nn.AdaptiveMaxPool2d(1)\n        self.global_avgpool = nn.AdaptiveAvgPool2d(1) \n        self.fc = nn.Sequential(\n            conv1x1(in_channel, in_channel//reduction).apply(init_weight),\n            nn.ReLU(True),\n            conv1x1(in_channel//reduction, in_channel).apply(init_weight)\n        )\n        \n    def forward(self, inputs):\n        x1 = self.global_maxpool(inputs)\n        x2 = self.global_avgpool(inputs)\n        x1 = self.fc(x1)\n        x2 = self.fc(x2)\n        x  = torch.sigmoid(x1 + x2)\n        return x\n    \n    \nclass SpatialAttentionModule(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.conv3x3 = conv3x3(2,1).apply(init_weight)\n        \n    def forward(self, inputs):\n        x1,_ = torch.max(inputs, dim=1, keepdim=True)\n        x2 = torch.mean(inputs, dim=1, keepdim=True)\n        x  = torch.cat([x1,x2], dim=1)\n        x  = self.conv3x3(x)\n        x  = torch.sigmoid(x)\n        return x\n    \n    \nclass CBAM(nn.Module):\n    def __init__(self, in_channel, reduction):\n        super().__init__()\n        self.channel_attention = ChannelAttentionModule(in_channel, reduction)\n        self.spatial_attention = SpatialAttentionModule()\n        \n    def forward(self, inputs):\n        x = inputs * self.channel_attention(inputs)\n        x = x * self.spatial_attention(x)\n        return x\n    \n    \nclass CenterBlock(nn.Module):\n    def __init__(self, in_channel, out_channel):\n        super().__init__()\n        self.conv = conv3x3(in_channel, out_channel).apply(init_weight)\n        \n    def forward(self, inputs):\n        x = self.conv(inputs)\n        return x\n\n\nclass DecodeBlock(nn.Module):\n    def __init__(self, in_channel, out_channel, upsample):\n        super().__init__()\n        self.bn1 = nn.BatchNorm2d(in_channel).apply(init_weight)\n        self.upsample = nn.Sequential()\n        if upsample:\n            self.upsample.add_module('upsample',nn.Upsample(scale_factor=2, mode='nearest'))\n        self.conv3x3_1 = conv3x3(in_channel, in_channel).apply(init_weight)\n        self.bn2 = nn.BatchNorm2d(in_channel).apply(init_weight)\n        self.conv3x3_2 = conv3x3(in_channel, out_channel).apply(init_weight)\n        self.cbam = CBAM(out_channel, reduction=16)\n        self.conv1x1   = conv1x1(in_channel, out_channel).apply(init_weight)\n        \n    def forward(self, inputs):\n        x  = F.relu(self.bn1(inputs))\n        x  = self.upsample(x)\n        x  = self.conv3x3_1(x)\n        x  = self.conv3x3_2(F.relu(self.bn2(x)))\n        x  = self.cbam(x)\n        x += self.conv1x1(self.upsample(inputs)) #shortcut\n        return x\n    \n    \n#U-Net ResNet34 + CBAM + hypercolumns + deepsupervision\nclass UNET_RESNET34(nn.Module):\n    def __init__(self, resolution, deepsupervision, clfhead, load_weights=True):\n        super().__init__()\n        h,w = resolution\n        self.deepsupervision = deepsupervision\n        self.clfhead = clfhead\n        \n        #encoder\n        model_name = 'resnet34' #26M\n        resnet34 = pretrainedmodels.__dict__['resnet34'](num_classes=1000,pretrained=None)\n        if load_weights:\n            resnet34.load_state_dict(torch.load(f'../../../pretrainedmodels_weight/{model_name}.pth'))\n        self.conv1   = resnet34.conv1 #(*,3,h,w)->(*,64,h/2,w/2)\n        self.bn1     = resnet34.bn1\n        self.maxpool = resnet34.maxpool #->(*,64,h/4,w/4)\n        self.layer1  = resnet34.layer1 #->(*,64,h/4,w/4) \n        self.layer2  = resnet34.layer2 #->(*,128,h/8,w/8) \n        self.layer3  = resnet34.layer3 #->(*,256,h/16,w/16) \n        self.layer4  = resnet34.layer4 #->(*,512,h/32,w/32) \n        \n        #center\n        self.center  = CenterBlock(512,512) #->(*,512,h/32,w/32) \n        \n        #decoder\n        self.decoder4 = DecodeBlock(512+512,64, upsample=True) #->(*,64,h/16,w/16) \n        self.decoder3 = DecodeBlock(64+256,64, upsample=True) #->(*,64,h/8,w/8) \n        self.decoder2 = DecodeBlock(64+128,64,  upsample=True) #->(*,64,h/4,w/4) \n        self.decoder1 = DecodeBlock(64+64,64,   upsample=True) #->(*,64,h/2,w/2)\n        self.decoder0 = DecodeBlock(64,64, upsample=True) #->(*,64,h,w) \n        \n        #upsample\n        self.upsample4 = nn.Upsample(scale_factor=16, mode='bilinear', align_corners=True)\n        self.upsample3 = nn.Upsample(scale_factor=8, mode='bilinear', align_corners=True)\n        self.upsample2 = nn.Upsample(scale_factor=4, mode='bilinear', align_corners=True)\n        self.upsample1 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        \n        #deep supervision\n        self.deep4 = conv1x1(64,1).apply(init_weight)\n        self.deep3 = conv1x1(64,1).apply(init_weight)\n        self.deep2 = conv1x1(64,1).apply(init_weight)\n        self.deep1 = conv1x1(64,1).apply(init_weight)\n        \n        #final conv\n        self.final_conv = nn.Sequential(\n            conv3x3(320,64).apply(init_weight),\n            nn.ELU(True),\n            conv1x1(64,1).apply(init_weight)\n        )\n        \n        #clf head\n        self.avgpool = nn.AdaptiveAvgPool2d(1)\n        self.clf = nn.Sequential(\n            nn.BatchNorm1d(512).apply(init_weight),\n            nn.Linear(512,512).apply(init_weight),\n            nn.ELU(True),\n            nn.BatchNorm1d(512).apply(init_weight),\n            nn.Linear(512,1).apply(init_weight)\n        )\n        \n    def forward(self, inputs):\n        #encoder\n        x0 = F.relu(self.bn1(self.conv1(inputs))) #->(*,64,h/2,w/2) \n        x0 = self.maxpool(x0) #->(*,64,h/4,w/4)\n        x1 = self.layer1(x0) #->(*,64,h/4,w/4)\n        x2 = self.layer2(x1) #->(*,128,h/8,w/8)\n        x3 = self.layer3(x2) #->(*,256,h/16,w/16)\n        x4 = self.layer4(x3) #->(*,512,h/32,w/32)\n        \n        #clf head\n        logits_clf = self.clf(self.avgpool(x4).squeeze(-1).squeeze(-1)) #->(*,1)\n        if config['clf_threshold'] is not None:\n            if (torch.sigmoid(logits_clf)>config['clf_threshold']).sum().item()==0:\n                bs,_,h,w = inputs.shape\n                logits = torch.zeros((bs,1,h,w))\n                if self.clfhead:\n                    if self.deepsupervision:\n                        return logits,_,_\n                    else:\n                        return logits,_\n                else:\n                    if self.deepsupervision:\n                        return logits,_\n                    else:\n                        return logits\n        \n        #center\n        y5 = self.center(x4) #->(*,512,h/32,w/32)\n        \n        #decoder\n        y4 = self.decoder4(torch.cat([x4,y5], dim=1)) #->(*,64,h/16,w/16)\n        y3 = self.decoder3(torch.cat([x3,y4], dim=1)) #->(*,64,h/8,w/8)\n        y2 = self.decoder2(torch.cat([x2,y3], dim=1)) #->(*,64,h/4,w/4)\n        y1 = self.decoder1(torch.cat([x1,y2], dim=1)) #->(*,64,h/2,w/2)\n        y0 = self.decoder0(y1) #->(*,64,h,w)\n        \n        #hypercolumns\n        y4 = self.upsample4(y4) #->(*,64,h,w)\n        y3 = self.upsample3(y3) #->(*,64,h,w)\n        y2 = self.upsample2(y2) #->(*,64,h,w)\n        y1 = self.upsample1(y1) #->(*,64,h,w)\n        hypercol = torch.cat([y0,y1,y2,y3,y4], dim=1)\n        \n        #final conv\n        logits = self.final_conv(hypercol) #->(*,1,h,w)\n        \n        #clf head\n        logits_clf = self.clf(self.avgpool(x4).squeeze(-1).squeeze(-1)) #->(*,1)\n        \n        if self.clfhead:\n            if self.deepsupervision:\n                s4 = self.deep4(y4)\n                s3 = self.deep3(y3)\n                s2 = self.deep2(y2)\n                s1 = self.deep1(y1)\n                logits_deeps = [s4,s3,s2,s1]\n                return logits, logits_deeps, logits_clf\n            else:\n                return logits, logits_clf\n        else:\n            if self.deepsupervision:\n                s4 = self.deep4(y4)\n                s3 = self.deep3(y3)\n                s2 = self.deep2(y2)\n                s1 = self.deep1(y1)\n                logits_deeps = [s4,s3,s2,s1]\n                return logits, logits_deeps\n            else:\n                return logits\n\n        \n#U-Net SeResNext50 + CBAM + hypercolumns + deepsupervision\nclass UNET_SERESNEXT50(nn.Module):\n    def __init__(self, resolution, deepsupervision, clfhead, load_weights=True):\n        super().__init__()\n        h,w = resolution\n        self.deepsupervision = deepsupervision\n        self.clfhead = clfhead\n        \n        #encoder\n        model_name = 'se_resnext50_32x4d' #26M\n        seresnext50 = pretrainedmodels.__dict__[model_name](pretrained=None)\n        if load_weights:\n            seresnext50.load_state_dict(torch.load(f'../../../pretrainedmodels_weight/{model_name}.pth'))\n        \n        self.encoder0 = nn.Sequential(\n            seresnext50.layer0.conv1, #(*,3,h,w)->(*,64,h/2,w/2)\n            seresnext50.layer0.bn1,\n            seresnext50.layer0.relu1,\n        )\n        self.encoder1 = nn.Sequential(\n            seresnext50.layer0.pool, #->(*,64,h/4,w/4)\n            seresnext50.layer1 #->(*,256,h/4,w/4)\n        )\n        self.encoder2 = seresnext50.layer2 #->(*,512,h/8,w/8)\n        self.encoder3 = seresnext50.layer3 #->(*,1024,h/16,w/16)\n        self.encoder4 = seresnext50.layer4 #->(*,2048,h/32,w/32)\n        \n        #center\n        self.center  = CenterBlock(2048,512) #->(*,512,h/32,w/32) 10,16\n        \n        #decoder\n        self.decoder4 = DecodeBlock(512+2048,64, upsample=True) #->(*,64,h/16,w/16) 20,32\n        self.decoder3 = DecodeBlock(64+1024,64, upsample=True) #->(*,64,h/8,w/8) 40,64\n        self.decoder2 = DecodeBlock(64+512,64,  upsample=True) #->(*,64,h/4,w/4) 80,128\n        self.decoder1 = DecodeBlock(64+256,64,   upsample=True) #->(*,64,h/2,w/2) 160,256\n        self.decoder0 = DecodeBlock(64,64, upsample=True) #->(*,64,h,w) 320,512\n        \n        #upsample\n        self.upsample4 = nn.Upsample(scale_factor=16, mode='bilinear', align_corners=True)\n        self.upsample3 = nn.Upsample(scale_factor=8, mode='bilinear', align_corners=True)\n        self.upsample2 = nn.Upsample(scale_factor=4, mode='bilinear', align_corners=True)\n        self.upsample1 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        \n        #deep supervision\n        self.deep4 = conv1x1(64,1).apply(init_weight)\n        self.deep3 = conv1x1(64,1).apply(init_weight)\n        self.deep2 = conv1x1(64,1).apply(init_weight)\n        self.deep1 = conv1x1(64,1).apply(init_weight)\n        \n        #final conv\n        self.final_conv = nn.Sequential(\n            conv3x3(320,64).apply(init_weight),\n            nn.ELU(True),\n            conv1x1(64,1).apply(init_weight)\n        )\n        \n        #clf head\n        self.avgpool = nn.AdaptiveAvgPool2d(1)\n        self.clf = nn.Sequential(\n            nn.BatchNorm1d(2048).apply(init_weight),\n            nn.Linear(2048,512).apply(init_weight),\n            nn.ELU(True),\n            nn.BatchNorm1d(512).apply(init_weight),\n            nn.Linear(512,1).apply(init_weight)\n        )\n        \n    def forward(self, inputs):\n        #encoder\n        x0 = self.encoder0(inputs) #->(*,64,h/2,w/2) 160,256\n        x1 = self.encoder1(x0) #->(*,256,h/4,w/4)\n        x2 = self.encoder2(x1) #->(*,512,h/8,w/8)\n        x3 = self.encoder3(x2) #->(*,1024,h/16,w/16)\n        x4 = self.encoder4(x3) #->(*,2048,h/32,w/32)\n        \n        #clf head\n        logits_clf = self.clf(self.avgpool(x4).squeeze(-1).squeeze(-1)) #->(*,1)\n        if config['clf_threshold'] is not None:\n            if (torch.sigmoid(logits_clf)>config['clf_threshold']).sum().item()==0:\n                bs,_,h,w = inputs.shape\n                logits = torch.zeros((bs,1,h,w))\n                if self.clfhead:\n                    if self.deepsupervision:\n                        return logits,_,_\n                    else:\n                        return logits,_\n                else:\n                    if self.deepsupervision:\n                        return logits,_\n                    else:\n                        return logits\n        \n        #center\n        y5 = self.center(x4) #->(*,320,h/32,w/32)\n        \n        #decoder\n        y4 = self.decoder4(torch.cat([x4,y5], dim=1)) #->(*,64,h/16,w/16)\n        y3 = self.decoder3(torch.cat([x3,y4], dim=1)) #->(*,64,h/8,w/8)\n        y2 = self.decoder2(torch.cat([x2,y3], dim=1)) #->(*,64,h/4,w/4)\n        y1 = self.decoder1(torch.cat([x1,y2], dim=1)) #->(*,64,h/2,w/2) 160,256\n        y0 = self.decoder0(y1) #->(*,64,h,w) 320,512\n        \n        #hypercolumns\n        y4 = self.upsample4(y4) #->(*,64,h,w)\n        y3 = self.upsample3(y3) #->(*,64,h,w)\n        y2 = self.upsample2(y2) #->(*,64,h,w)\n        y1 = self.upsample1(y1) #->(*,64,h,w)\n        hypercol = torch.cat([y0,y1,y2,y3,y4], dim=1)\n        \n        #final conv\n        logits = self.final_conv(hypercol) #->(*,4,h,w)\n        \n        #clf head\n        logits_clf = self.clf(self.avgpool(x4).squeeze(-1).squeeze(-1)) #->(*,1)\n        \n        if self.clfhead:\n            if self.deepsupervision:\n                s4 = self.deep4(y4)\n                s3 = self.deep3(y3)\n                s2 = self.deep2(y2)\n                s1 = self.deep1(y1)\n                logits_deeps = [s4,s3,s2,s1]\n                return logits, logits_deeps, logits_clf\n            else:\n                return logits, logits_clf\n        else:\n            if self.deepsupervision:\n                s4 = self.deep4(y4)\n                s3 = self.deep3(y3)\n                s2 = self.deep2(y2)\n                s1 = self.deep1(y1)\n                logits_deeps = [s4,s3,s2,s1]\n                return logits, logits_deeps\n            else:\n                return logits\n    \n\n#U-Net SeResNext101 + CBAM + hypercolumns + deepsupervision\nclass UNET_SERESNEXT101(nn.Module):\n    def __init__(self, resolution, deepsupervision, clfhead, load_weights=True):\n        super().__init__()\n        h,w = resolution\n        self.deepsupervision = deepsupervision\n        self.clfhead = clfhead\n        \n        #encoder\n        model_name = 'se_resnext101_32x4d'\n        seresnext101 = pretrainedmodels.__dict__[model_name](pretrained=None)\n        if load_weights:\n            print('using pretrained model')\n#             seresnext101.load_state_dict(torch.load(f'../../../pretrainedmodels_weight/{model_name}.pth'))            \n            seresnext101.load_state_dict(torch.load('../input/pretrained-models-zh/se_resnext101_32x4d-3b2fe3d8.pth'))\n\n            \n            \n        \n        self.encoder0 = nn.Sequential(\n            seresnext101.layer0.conv1, #(*,3,h,w)->(*,64,h/2,w/2)\n            seresnext101.layer0.bn1,\n            seresnext101.layer0.relu1,\n        )\n        self.encoder1 = nn.Sequential(\n            seresnext101.layer0.pool, #->(*,64,h/4,w/4)\n            seresnext101.layer1 #->(*,256,h/4,w/4)\n        )\n        self.encoder2 = seresnext101.layer2 #->(*,512,h/8,w/8)\n        self.encoder3 = seresnext101.layer3 #->(*,1024,h/16,w/16)\n        self.encoder4 = seresnext101.layer4 #->(*,2048,h/32,w/32)\n        \n        #center\n        self.center  = CenterBlock(2048,512) #->(*,512,h/32,w/32)\n        \n        #decoder\n        self.decoder4 = DecodeBlock(512+2048,64, upsample=True) #->(*,64,h/16,w/16)\n        self.decoder3 = DecodeBlock(64+1024,64, upsample=True) #->(*,64,h/8,w/8)\n        self.decoder2 = DecodeBlock(64+512,64,  upsample=True) #->(*,64,h/4,w/4) \n        self.decoder1 = DecodeBlock(64+256,64,   upsample=True) #->(*,64,h/2,w/2) \n        self.decoder0 = DecodeBlock(64,64, upsample=True) #->(*,64,h,w) \n        \n        #upsample\n        self.upsample4 = nn.Upsample(scale_factor=16, mode='bilinear', align_corners=True)\n        self.upsample3 = nn.Upsample(scale_factor=8, mode='bilinear', align_corners=True)\n        self.upsample2 = nn.Upsample(scale_factor=4, mode='bilinear', align_corners=True)\n        self.upsample1 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        \n        #deep supervision\n        self.deep4 = conv1x1(64,1).apply(init_weight)\n        self.deep3 = conv1x1(64,1).apply(init_weight)\n        self.deep2 = conv1x1(64,1).apply(init_weight)\n        self.deep1 = conv1x1(64,1).apply(init_weight)\n        \n        #final conv\n        self.final_conv = nn.Sequential(\n            conv3x3(320,64).apply(init_weight),\n            nn.ELU(True),\n            conv1x1(64,1).apply(init_weight)\n        )\n        \n        #clf head\n        self.avgpool = nn.AdaptiveAvgPool2d(1)\n        self.clf = nn.Sequential(\n            nn.BatchNorm1d(2048).apply(init_weight),\n            nn.Linear(2048,512).apply(init_weight),\n            nn.ELU(True),\n            nn.BatchNorm1d(512).apply(init_weight),\n            nn.Linear(512,1).apply(init_weight)\n        )\n        \n    def forward(self, inputs):\n        #encoder\n        x0 = self.encoder0(inputs) #->(*,64,h/2,w/2)\n        x1 = self.encoder1(x0) #->(*,256,h/4,w/4)\n        x2 = self.encoder2(x1) #->(*,512,h/8,w/8)\n        x3 = self.encoder3(x2) #->(*,1024,h/16,w/16)\n        x4 = self.encoder4(x3) #->(*,2048,h/32,w/32)\n        \n        #clf head\n        logits_clf = self.clf(self.avgpool(x4).squeeze(-1).squeeze(-1)) #->(*,1)\n        if config['clf_threshold'] is not None:\n            if (torch.sigmoid(logits_clf)>config['clf_threshold']).sum().item()==0:\n                bs,_,h,w = inputs.shape\n                logits = torch.zeros((bs,1,h,w))\n                if self.clfhead:\n                    if self.deepsupervision:\n                        return logits,_,_\n                    else:\n                        return logits,_\n                else:\n                    if self.deepsupervision:\n                        return logits,_\n                    else:\n                        return logits\n        \n        #center\n        y5 = self.center(x4) #->(*,320,h/32,w/32)\n        \n        #decoder\n        y4 = self.decoder4(torch.cat([x4,y5], dim=1)) #->(*,64,h/16,w/16)\n        y3 = self.decoder3(torch.cat([x3,y4], dim=1)) #->(*,64,h/8,w/8)\n        y2 = self.decoder2(torch.cat([x2,y3], dim=1)) #->(*,64,h/4,w/4)\n        y1 = self.decoder1(torch.cat([x1,y2], dim=1)) #->(*,64,h/2,w/2) \n        y0 = self.decoder0(y1) #->(*,64,h,w)\n        \n        #hypercolumns\n        y4 = self.upsample4(y4) #->(*,64,h,w)\n        y3 = self.upsample3(y3) #->(*,64,h,w)\n        y2 = self.upsample2(y2) #->(*,64,h,w)\n        y1 = self.upsample1(y1) #->(*,64,h,w)\n        hypercol = torch.cat([y0,y1,y2,y3,y4], dim=1)\n        \n        #final conv\n        logits = self.final_conv(hypercol) #->(*,1,h,w)\n        \n        #clf head\n        logits_clf = self.clf(self.avgpool(x4).squeeze(-1).squeeze(-1)) #->(*,1)\n        \n        if self.clfhead:\n            if self.deepsupervision:\n                s4 = self.deep4(y4)\n                s3 = self.deep3(y3)\n                s2 = self.deep2(y2)\n                s1 = self.deep1(y1)\n                logits_deeps = [s4,s3,s2,s1]\n                return logits, logits_deeps, logits_clf\n            else:\n                return logits, logits_clf\n        else:\n            if self.deepsupervision:\n                s4 = self.deep4(y4)\n                s3 = self.deep3(y3)\n                s2 = self.deep2(y2)\n                s1 = self.deep1(y1)\n                logits_deeps = [s4,s3,s2,s1]\n                return logits, logits_deeps\n            else:\n                return logits    \n\n    \ndef build_model(resolution, deepsupervision, clfhead, load_weights):\n    model_name = config['model_name']\n    if model_name=='resnet34':\n        model = UNET_RESNET34(resolution, deepsupervision, clfhead, load_weights)\n    elif model_name=='seresnext50':\n        model = UNET_SERESNEXT50(resolution, deepsupervision, clfhead, load_weights)\n    elif model_name=='seresnext101':\n        model = UNET_SERESNEXT101(resolution, deepsupervision, clfhead, load_weights)\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:04:05.015278Z","iopub.execute_input":"2022-11-05T07:04:05.015769Z","iopub.status.idle":"2022-11-05T07:04:07.544320Z","shell.execute_reply.started":"2022-11-05T07:04:05.015727Z","shell.execute_reply":"2022-11-05T07:04:07.543006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#split the model to encoder and decoder for fast.ai\nsplit_layers = lambda m: [list(m.encoder0.parameters())+list(m.encoder1.parameters())+\n                list(m.encoder2.parameters())+list(m.encoder3.parameters())+\n                list(m.encoder4.parameters()),\n                          \n                list(m.center.parameters())+list(m.decoder4.parameters())+\n                list(m.decoder3.parameters())+list(m.decoder2.parameters())+\n                list(m.decoder1.parameters())+list(m.decoder0.parameters())+\n                list(m.final_conv.parameters()), list(m.deep4.parameters()),\n                list(m.deep3.parameters()), list(m.deep2.parameters()),\n                list(m.deep1.parameters())]","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:04:07.550173Z","iopub.execute_input":"2022-11-05T07:04:07.553366Z","iopub.status.idle":"2022-11-05T07:04:07.563478Z","shell.execute_reply.started":"2022-11-05T07:04:07.553321Z","shell.execute_reply":"2022-11-05T07:04:07.562466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Loss function and performance metric\nThe loss that works the best for image segmentation in most of the cases is [Lovász loss](https://arxiv.org/pdf/1705.08790.pdf), a differentiable surrogate of IoU. However, ReLU in it must be replaced by (ELU + 1), like I did [here](https://www.kaggle.com/iafoss/lovasz). Another trick is consideration of a symmetric Lovász loss: consider not only a predicted segmentation and a provided mask but also the inverse prediction and the inverse mask (predict mask for negative case).","metadata":{}},{"cell_type":"code","source":"def symmetric_lovasz(outputs, targets):\n    return 0.5*(lovasz_hinge(outputs, targets) + lovasz_hinge(-outputs, 1.0 - targets))\n\nclass Dice_soft(Metric):\n    def __init__(self, axis=1): \n        self.axis = axis \n    def reset(self): self.inter,self.union = 0,0\n    def accumulate(self, learn):\n        pred,targ = flatten_check(torch.sigmoid(learn.pred), learn.y)\n        self.inter += (pred*targ).float().sum().item()\n        self.union += (pred+targ).float().sum().item()\n    @property\n    def value(self): return 2.0 * self.inter/self.union if self.union > 0 else None\n    \n# dice with automatic threshold selection\nclass Dice_th(Metric):\n    def __init__(self, ths=np.arange(0.1,0.9,0.05), axis=1): \n        self.axis = axis\n        self.ths = ths\n        \n    def reset(self): \n        self.inter = torch.zeros(len(self.ths))\n        self.union = torch.zeros(len(self.ths))\n        \n    def accumulate(self, learn):\n        pred,targ = flatten_check(torch.sigmoid(learn.pred), learn.y)\n        for i,th in enumerate(self.ths):\n            p = (pred > th).float()\n            self.inter[i] += (p*targ).float().sum().item()\n            self.union[i] += (p+targ).float().sum().item()\n\n    @property\n    def value(self):\n        dices = torch.where(self.union > 0.0, \n                2.0*self.inter/self.union, torch.zeros_like(self.union))\n        return dices.max()","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:04:08.616059Z","iopub.execute_input":"2022-11-05T07:04:08.617314Z","iopub.status.idle":"2022-11-05T07:04:08.637196Z","shell.execute_reply.started":"2022-11-05T07:04:08.617270Z","shell.execute_reply":"2022-11-05T07:04:08.636048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model evaluation class","metadata":{}},{"cell_type":"code","source":"#iterator like wrapper that returns predicted and gt masks\nclass Model_pred:\n    def __init__(self, model, dl, tta:bool=True, half:bool=False):\n        self.model = model\n        self.dl = dl\n        self.tta = tta\n        self.half = half\n        \n    def __iter__(self):\n        self.model.eval()\n        name_list = self.dl.dataset.fnames\n        count=0\n        with torch.no_grad():\n#             for idx, (x,y) in enumerate(iter(self.dl)):\n#                 if idx != len(self.dl) - 1:\n#                     print(idx)\n#                     continue\n                x = x.cuda()\n                if self.half: x = x.half()\n                p = self.model(x)\n                py = torch.sigmoid(p).detach()\n#                 print('py', py.get_device())\n                if self.tta:\n                    #x,y,xy flips as TTA\n                    flips = [[-1],[-2],[-2,-1]]\n                    for f in flips:\n                        p = self.model(torch.flip(x,f))\n                        p = torch.flip(p,f)\n#                         print('flip', torch.sigmoid(p).detach().get_device())\n                        if not p.is_cuda: \n                            py += torch.sigmoid(p).cuda().detach()\n                        else:\n                            py += torch.sigmoid(p).detach()  # bug\n                    py /= (1+len(flips))\n                if y is not None and len(y.shape)==4 and py.shape != y.shape:\n                    py = F.upsample(py, size=(y.shape[-2],y.shape[-1]), mode=\"bilinear\")\n                py = py.permute(0,2,3,1).float().cpu()\n                batch_size = len(py)\n                for i in range(batch_size):\n                    taget = y[i].detach().cpu() if y is not None else None\n                    yield py[i],taget,name_list[count]\n                    count += 1\n                    \n    def __len__(self):\n        return len(self.dl.dataset)\n    \nclass Dice_th_pred(Metric):\n    def __init__(self, ths=np.arange(0.1,0.9,0.01), axis=1): \n        self.axis = axis\n        self.ths = ths\n        self.reset()\n        \n    def reset(self): \n        self.inter = torch.zeros(len(self.ths))\n        self.union = torch.zeros(len(self.ths))\n        \n    def accumulate(self,p,t):\n        pred,targ = flatten_check(p, t)\n        for i,th in enumerate(self.ths):\n            p = (pred > th).float()\n            self.inter[i] += (p*targ).float().sum().item()\n            self.union[i] += (p+targ).float().sum().item()\n\n    @property\n    def value(self):\n        dices = torch.where(self.union > 0.0, 2.0*self.inter/self.union, \n                            torch.zeros_like(self.union))\n        return dices\n    \ndef save_img(data,name,out):\n    data = data.float().cpu().numpy()\n    img = cv2.imencode('.png',(data*255).astype(np.uint8))[1]\n    out.writestr(name, img)","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:04:09.393766Z","iopub.execute_input":"2022-11-05T07:04:09.396896Z","iopub.status.idle":"2022-11-05T07:04:09.450044Z","shell.execute_reply.started":"2022-11-05T07:04:09.396856Z","shell.execute_reply":"2022-11-05T07:04:09.448953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## TransUNet","metadata":{}},{"cell_type":"code","source":"# # transunet\n\n# ! git clone https://github.com/Beckschen/TransUNet.git\n# # ! rm -rf TransUNet\n\n# ! pip install ml-collections \n\n# from TransUNet.networks.vit_seg_modeling import CONFIGS as CONFIGS_ViT_seg\n# from TransUNet.networks.vit_seg_modeling import VisionTransformer as ViT_seg\n\n# IMG_SIZE = 256\n# VIT_PATCHES_SIZE = 16\n# MODEL_NAME = 'R50-ViT-B_16'\n\n# # start with smallest model \n# config_vit = CONFIGS_ViT_seg[MODEL_NAME]\n# config_vit.patches.grid = (int(IMG_SIZE / VIT_PATCHES_SIZE), int(IMG_SIZE / VIT_PATCHES_SIZE))\n# config_vit.n_classes = 1\n# config_vit.n_skip = 3\n# config_vit.pretrained_path = \"R50+ViT-B_16.npz\"\n\n# ! wget https://storage.googleapis.com/vit_models/imagenet21k/R50+ViT-B_16.npz\n    \n# split_layers = lambda m: [list(m.transformer.parameters()),\n#                 list(m.decoder.parameters())]","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:04:10.162468Z","iopub.execute_input":"2022-11-05T07:04:10.163053Z","iopub.status.idle":"2022-11-05T07:04:10.169239Z","shell.execute_reply.started":"2022-11-05T07:04:10.163008Z","shell.execute_reply":"2022-11-05T07:04:10.168071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Perform training","metadata":{}},{"cell_type":"code","source":"# # without kfold\n# fold = 0\n# batch_size = 16\n# nfolds = 2\n# fp16 = MixedPrecision()\n# dice = Dice_th_pred(np.arange(0.2,0.7,0.01))\n\n# ds_t = HuBMAPDataset(fold=fold, train=True, tfms=get_aug())\n# ds_v = HuBMAPDataset(fold=fold, train=False)\n# data = ImageDataLoaders.from_dsets(ds_t,ds_v,bs=batch_size, # ImageDataLoaders is a fastai thing - TODO switch to pytorch?\n#             num_workers=NUM_WORKERS,pin_memory=True).cuda()\n\n# # model = ViT_seg(config_vit, img_size=IMG_SIZE, num_classes=1).cuda()\n# model = build_model(\n# #     resolution=(256, 256),\n#     resolution=(None,None),\n#     deepsupervision=False,\n#     clfhead=False,\n#     load_weights=False  # seems like pretrainedmodels_weight is a private repo\n# ).cuda()\n# # model.load_from(weights=np.load(config_vit.pretrained_path))\n# learn = Learner(data, model, loss_func=symmetric_lovasz,# fastai again - switch to pytorch?\n#             metrics=[Dice_soft(),Dice_th()], \n#             splitter=split_layers)#.to_fp16(clip=0.5)\n\n# #start with training the head\n# # learn.freeze_to(-1) #doesn't work\n# # for param in learn.opt.param_groups[0]['params']:\n# #     param.requires_grad = False\n# # learn.fit_one_cycle(1, lr_max=0.5e-2)\n\n# #continue training full model\n# # learn.unfreeze()\n# learn.fit_one_cycle(2, lr_max=slice(2e-4,2e-3),\n#     cbs=[SaveModelCallback(monitor='dice_th',comp=np.greater), GradientClip, fp16])\n# torch.save(learn.model.state_dict(),f'model_{fold}.pth')\n\n# #model evaluation on val and saving the masks\n# mp = Model_pred(learn.model,learn.dls.loaders[1])\n# with zipfile.ZipFile('val_masks_tta.zip', 'a') as out:\n#     for p in progress_bar(mp):\n#         dice.accumulate(p[0],p[1])\n#         save_img(p[0],p[2],out)\n# gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:04:10.848873Z","iopub.execute_input":"2022-11-05T07:04:10.849503Z","iopub.status.idle":"2022-11-05T07:04:10.856825Z","shell.execute_reply.started":"2022-11-05T07:04:10.849456Z","shell.execute_reply":"2022-11-05T07:04:10.855736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # TODO how many epochs is this? no idea how fastai works\n# from fastai.callback.fp16 import MixedPrecision\n# from fastai.callback.training import GradientClip\n# batch_size = 16\n# nfolds = 3\n# fp16 = MixedPrecision()\n\n# dice = Dice_th_pred(np.arange(0.2,0.7,0.01))\n# for fold in range(nfolds):\n#     print(f'### FOLD {fold} ###')\n#     ds_t = HuBMAPDataset(fold=fold, train=True, tfms=get_aug())\n#     ds_v = HuBMAPDataset(fold=fold, train=False)\n#     data = ImageDataLoaders.from_dsets(ds_t,ds_v,bs=batch_size, # ImageDataLoaders is a fastai thing - TODO switch to pytorch?\n#                 num_workers=NUM_WORKERS,pin_memory=True).cuda()\n    \n#     ### change model here ###\n# #     model = UneXt50().cuda()\n# #     model = ViT_seg(config_vit, img_size=IMG_SIZE, num_classes=1).cuda()\n# #     model.load_from(weights=np.load(config_vit.pretrained_path))\n\n#     model = build_model(\n#         resolution=(None,None),\n#         deepsupervision=False,  # try True but will have tensor dimension misalign bug\n#         clfhead=False,\n#         load_weights=True  # seems like pretrainedmodels_weight is a private repo\n#     ).cuda()\n#     learn = Learner(data, model, loss_func=symmetric_lovasz,# fastai again - switch to pytorch?\n#                 metrics=[Dice_soft(),Dice_th()], \n#                 splitter=split_layers)#.to_fp16(clip=0.5)\n#     #start with training the head\n#     learn.freeze_to(-1) #doesn't work\n#     for param in learn.opt.param_groups[0]['params']:\n#         param.requires_grad = False\n#     learn.fit_one_cycle(2, lr_max=0.5e-2)\n\n#     #continue training full model\n#     learn.unfreeze()\n#     learn.fit_one_cycle(4, lr_max=slice(2e-4,2e-3),\n#         cbs=[SaveModelCallback(monitor='dice_th',comp=np.greater), GradientClip, fp16])\n#     torch.save(learn.model.state_dict(),f'model_{fold}.pth')\n    \n#     #model evaluation on val and saving the masks\n# #     mp = Model_pred(learn.model.cuda(),learn.dls.loaders[1])\n# # #     with zipfile.ZipFile('val_masks_tta.zip', 'a') as out:    \n# #     out = zipfile.ZipFile('val_masks_tta.zip', 'a')\n# #     for p in progress_bar(mp):\n# #         dice.accumulate(p[0],p[1])\n# #         save_img(p[0],p[2],out)\n#     gc.collect()\n","metadata":{"execution":{"iopub.status.busy":"2022-11-05T07:04:11.271061Z","iopub.execute_input":"2022-11-05T07:04:11.271536Z","iopub.status.idle":"2022-11-05T09:26:30.964371Z","shell.execute_reply.started":"2022-11-05T07:04:11.271497Z","shell.execute_reply":"2022-11-05T09:26:30.962362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dices = dice.value\n# noise_ths = dice.ths\n# best_dice = dices.max()\n# best_thr = noise_ths[dices.argmax()]\n# plt.figure(figsize=(8,4))\n# plt.plot(noise_ths, dices, color='blue')\n# plt.vlines(x=best_thr, ymin=dices.min(), ymax=dices.max(), colors='black')\n# d = dices.max() - dices.min()\n# plt.text(noise_ths[-1]-0.1, best_dice-0.1*d, f'DICE = {best_dice:.3f}', fontsize=12);\n# plt.text(noise_ths[-1]-0.1, best_dice-0.2*d, f'TH = {best_thr:.3f}', fontsize=12);\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-05T09:43:43.095701Z","iopub.execute_input":"2022-11-05T09:43:43.096189Z","iopub.status.idle":"2022-11-05T09:43:43.372756Z","shell.execute_reply.started":"2022-11-05T09:43:43.096125Z","shell.execute_reply":"2022-11-05T09:43:43.371612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ","metadata":{}}]}