{"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":"!cp -r ../input/pytorch-segmentation-models-lib/ ./\n!pip config set global.disable-pip-version-check true\n!pip install -q ./pytorch-segmentation-models-lib/timm-0.4.12-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:01.474874Z","iopub.execute_input":"2022-09-15T17:40:01.475797Z","iopub.status.idle":"2022-09-15T17:40:15.509032Z","shell.execute_reply.started":"2022-09-15T17:40:01.475262Z","shell.execute_reply":"2022-09-15T17:40:15.507893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp -r ../input/einops-041-wheel/ ./\n!pip config set global.disable-pip-version-check true\n!pip install -q ../input/einops-041-wheel/einops-0.4.1-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:15.511486Z","iopub.execute_input":"2022-09-15T17:40:15.511876Z","iopub.status.idle":"2022-09-15T17:40:27.059122Z","shell.execute_reply.started":"2022-09-15T17:40:15.511832Z","shell.execute_reply":"2022-09-15T17:40:27.057913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('../input/hubmap-submit-06') \nsys.path.append('../input/hubmap-submit-06/[third_party]')   \nsys.path.append('../input/hubmap-submit-06/[third_party]/timm')\n\nsys.path.insert(0, '../input/coat-5layer')\nsys.path.insert(0, '../input/coatsmall5level')\nsys.path.insert(0, '../input/coatsmall')\n","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:27.061342Z","iopub.execute_input":"2022-09-15T17:40:27.061742Z","iopub.status.idle":"2022-09-15T17:40:27.068420Z","shell.execute_reply.started":"2022-09-15T17:40:27.061697Z","shell.execute_reply":"2022-09-15T17:40:27.067342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\nimport importlib\nfrom timeit import default_timer as timer\n\nimport torch\nimport torch.cuda.amp as amp\nimport torch.nn.functional as F\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport tifffile \nimport cv2\nimport os\nimport gc\nfrom tqdm.notebook import tqdm\nimport rasterio\nfrom rasterio.windows import Window\n\nfrom fastai.vision.all import *\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A # Augmentations\nimport warnings\nwarnings.filterwarnings(\"ignore\")\nimport glob\nimport copy\nprint('import ok\\n')","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:27.071575Z","iopub.execute_input":"2022-09-15T17:40:27.072361Z","iopub.status.idle":"2022-09-15T17:40:32.519194Z","shell.execute_reply.started":"2022-09-15T17:40:27.072323Z","shell.execute_reply":"2022-09-15T17:40:32.517344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport matplotlib.pyplot as plt\nfrom skimage import io,img_as_float\nimport cv2\nfrom sklearn.model_selection import StratifiedKFold, KFold, StratifiedGroupKFold\n# import rasterio\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport random\nimport time\nimport copy\nfrom collections import defaultdict\nimport gc\n\n#SMP\n# import segmentation_models_pytorch as smp\n\n#HUGGINGFACE \nimport transformers \n\n# SWA\nfrom torch.optim.swa_utils import AveragedModel, SWALR\n\n\n#PYTORCH\nimport torch\nfrom torch.utils.data import Dataset,DataLoader\nfrom torch.optim import lr_scheduler\nimport torch.nn as nn\nfrom tqdm import tqdm\nfrom torchmetrics.functional import dice\nimport torch.optim as optim\nfrom torchvision.utils import make_grid\nfrom torchvision.io import read_image\nfrom pathlib import Path\nimport torchvision.transforms.functional as F\nfrom torch.nn import functional as nnF\nfrom torch.cuda import amp\n\nimport argparse\nfrom yaml import parse\nfrom colorama import Fore, Back, Style\nc_  = Fore.GREEN\nlo_  = Fore.RED\nyel_ = Fore.BLUE\nsr_ = Style.RESET_ALL\n\nplt.rcParams[\"savefig.bbox\"] = 'tight'","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:32.520858Z","iopub.execute_input":"2022-09-15T17:40:32.521500Z","iopub.status.idle":"2022-09-15T17:40:33.256742Z","shell.execute_reply.started":"2022-09-15T17:40:32.521461Z","shell.execute_reply":"2022-09-15T17:40:33.255743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport cv2\n\nclass dotdict(dict):\n\t__setattr__ = dict.__setitem__\n\t__delattr__ = dict.__delitem__\n\t\n\tdef __getattr__(self, name):\n\t\ttry:\n\t\t\treturn self[name]\n\t\texcept KeyError:\n\t\t\traise AttributeError(name)\n\n\n#--- helper ----------\ndef time_to_str(t, mode='min'):\n\tif mode=='min':\n\t\tt  = int(t)/60\n\t\thr = t//60\n\t\tmin = t%60\n\t\treturn '%2d hr %02d min'%(hr,min)\n\t\n\telif mode=='sec':\n\t\tt   = int(t)\n\t\tmin = t//60\n\t\tsec = t%60\n\t\treturn '%2d min %02d sec'%(min,sec)\n\t\n\telse:\n\t\traise NotImplementedError\n\t\ndef image_show(name, image, type='bgr', resize=1):\n\tif type == 'rgb': image = np.ascontiguousarray(image[:,:,::-1])\n\tH,W = image.shape[0:2]\n\t\n\tcv2.namedWindow(name, cv2.WINDOW_GUI_NORMAL)  #WINDOW_NORMAL #WINDOW_GUI_EXPANDED\n\tcv2.imshow(name, image) #.astype(np.uint8))\n\tcv2.resizeWindow(name, round(resize*W), round(resize*H))","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.258734Z","iopub.execute_input":"2022-09-15T17:40:33.259151Z","iopub.status.idle":"2022-09-15T17:40:33.268575Z","shell.execute_reply.started":"2022-09-15T17:40:33.259113Z","shell.execute_reply":"2022-09-15T17:40:33.267376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys, os\n\nsys.path.append('../input/hubmap-submit-06/[third_party]')\nsys.path.append('../[third_party]')\n#sys.path.append('/root/share1/kaggle/2022/hubmap-organ-segmentation/code/hubmap-dummy-02/[submit]/[third_party]')\n\n# from my_lib_kv import *\n\nimport tifffile as tiff\nimport json\nimport cv2\nimport pandas as pd\nimport math\nimport numpy as np\n\n\n\n##--------------------------------------------------------------------------------------\norgan_meta = dotdict(\n\tkidney = dotdict(\n\t\tlabel = 1,\n\t\tum    = 0.5000,\n\t\tftu   ='glomeruli',\n\t),\n\tprostate = dotdict(\n\t\tlabel = 2,\n\t\tum    = 6.2630,\n\t\tftu   ='glandular acinus',\n\t),\n\tlargeintestine = dotdict(\n\t\tlabel = 3,\n\t\tum    = 0.2290,\n\t\tftu   ='crypt',\n\t),\n\tspleen = dotdict(\n\t\tlabel = 4,\n\t\tum    = 0.4945,\n\t\tftu   ='white pulp',\n\t),\n\tlung = dotdict(\n\t\tlabel = 5,\n\t\tum    = 0.7562,\n\t\tftu   ='alveolus',\n\t),\n)\n\n\n\norgan_to_label = {k: organ_meta[k].label for k in organ_meta.keys()}\nlabel_to_organ = {v:k for k,v in organ_to_label.items()}\nnum_organ=5\n#['kidney', 'prostate', 'largeintestine', 'spleen', 'lung']\n\n\ndef read_tiff(image_file, mode='rgb'):\n\timage = tiff.imread(image_file)\n\timage = image.squeeze()\n\tif image.shape[0] == 3:\n\t\timage = image.transpose(1, 2, 0)\n\tif mode=='bgr':\n\t\timage = image[:,:,::-1]\n\timage = np.ascontiguousarray(image)\n\treturn image\n\ndef read_json_as_list(json_file):\n\twith open(json_file) as f:\n\t\tj = json.load(f)\n\treturn j\n\n\n# # --- rle ---------------------------------\ndef rle_decode(rle, height, width , fill=255, dtype=np.uint8):\n\ts = rle.split()\n\tstart, length = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n\tstart -= 1\n\tmask = np.zeros(height*width, dtype=dtype)\n\tfor i, l in zip(start, length):\n\t\tmask[i:i+l] = fill\n\tmask = mask.reshape(width,height).T\n\tmask = np.ascontiguousarray(mask)\n\treturn mask\n\n\ndef rle_encode(mask):\n\tm = mask.T.flatten()\n\tm = np.concatenate([[0], m, [0]])\n\trun = np.where(m[1:] != m[:-1])[0] + 1\n\trun[1::2] -= run[::2]\n\trle =  ' '.join(str(r) for r in run)\n\treturn rle\n\n\n#\n# # --draw ------------------------------------------\ndef mask_to_inner_contour(mask):\n\tmask = mask>0.5\n\tpad = np.lib.pad(mask, ((1, 1), (1, 1)), 'reflect')\n\tcontour = mask & (\n\t\t\t(pad[1:-1,1:-1] != pad[:-2,1:-1]) \\\n\t\t\t| (pad[1:-1,1:-1] != pad[2:,1:-1]) \\\n\t\t\t| (pad[1:-1,1:-1] != pad[1:-1,:-2]) \\\n\t\t\t| (pad[1:-1,1:-1] != pad[1:-1,2:])\n\t)\n\treturn contour\n\n\ndef draw_contour_overlay(image, mask, color=(0,0,255), thickness=1):\n\tcontour =  mask_to_inner_contour(mask)\n\tif thickness==1:\n\t\timage[contour] = color\n\telse:\n\t\tr = max(1,thickness//2)\n\t\tfor y,x in np.stack(np.where(contour)).T:\n\t\t\tcv2.circle(image, (x,y), r, color, lineType=cv2.LINE_4 )\n\treturn image\n\ndef result_to_overlay(image, mask, probability=None, **kwargs):\n \n\t\n\tH,W,C= image.shape\n\tif mask is None:\n\t\tmask = np.zeros((H,W),np.float32)\n\tif probability is None:\n\t\tprobability = np.zeros((H,W),np.float32)\n\t\t\n\to1 = np.zeros((H,W,3),np.float32)\n\to1[:,:,2] = mask\n\to1[:,:,1] = probability\n\t\n\to2 = image.copy()\n\to2 = o2*0.5\n\to2[:,:,1] += 0.5*probability\n\to2 = draw_contour_overlay(o2, mask, color=(0,0,1), thickness=max(3,int(7*H/1500)))\n\t\n\t#---\n\to2,image,o1 = [(m*255).astype(np.uint8) for m in [o2,image,o1]]\n\tif kwargs.get('dice_score',-1)>=0:\n\t\tdraw_shadow_text(o2,'dice=%0.5f'%kwargs.get('dice_score'),(20,80),2.5,(255,255,255),5)\n\tif kwargs.get('d',None) is not None:\n\t\td = kwargs.get('d')\n\t\tdraw_shadow_text(o2,d['id'],(20,140),1.5,(255,255,255),3)\n\t\tdraw_shadow_text(o2,d.organ+'(%s)'%(organ_meta[d.organ].ftu),(20,190),1.5,(255,255,255),3)\n\t\tdraw_shadow_text(o2,'%0.1f um'%(d.pixel_size),(20,240),1.5,(255,255,255),3)\n\t\ts100 = int(100/d.pixel_size)\n\t\tsx,sy = W-s100-40,H-80\n\t\tcv2.rectangle(o2,(sx,sy),(sx+s100,sy+s100//2),(0,0,0),-1)\n\t\tdraw_shadow_text(o2,'100um',(sx+8,sy+40),1,(255,255,255),2)\n\t\tpass\n\t\n\t#draw_shadow_text(image,'input',(5,15),0.6,(1,1,1),1)\n\t#draw_shadow_text(im_paste,'predict',(5,15),0.6,(1,1,1),1)\n\n\toverlay = np.hstack([o2,image,o1])\n\treturn overlay\n\n# --lb metric ------------------------------------------\n# https://www.kaggle.com/competitions/hubmap-organ-segmentation/overview/supervised-ml-evaluation\n\ndef compute_dice_score(probability, mask):\n\tN = len(probability)\n\tp = probability.reshape(N,-1)\n\tt = mask.reshape(N,-1)\n\t\n\tp = p>0.5\n\tt = t>0.5\n\tuion = p.sum(-1) + t.sum(-1)\n\toverlap = (p*t).sum(-1)\n\tdice = 2*overlap/(uion+0.0001)\n\treturn dice","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.270436Z","iopub.execute_input":"2022-09-15T17:40:33.271150Z","iopub.status.idle":"2022-09-15T17:40:33.301216Z","shell.execute_reply.started":"2022-09-15T17:40:33.271112Z","shell.execute_reply":"2022-09-15T17:40:33.300148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport pdb\n\nclass MixUpSample(nn.Module):\n\tdef __init__( self, scale_factor=2):\n\t\tsuper().__init__()\n\t\tself.mixing = nn.Parameter(torch.tensor(0.5))\n\t\tself.scale_factor = scale_factor\n\t\n\tdef forward(self, x):\n\t\tx = self.mixing *F.interpolate(x, scale_factor=self.scale_factor, mode='bilinear', align_corners=False) \\\n\t\t\t+ (1-self.mixing )*F.interpolate(x, scale_factor=self.scale_factor, mode='nearest')\n\t\treturn x\n\n#https://github.com/lhoyer/DAFormer/blob/master/mmseg/models/decode_heads/daformer_head.py\ndef Conv2dBnReLU(in_channel, out_channel, kernel_size=3, padding=1,stride=1, dilation=1):\n\treturn nn.Sequential(\n\t\tnn.Conv2d(in_channel, out_channel, kernel_size=kernel_size, padding=padding, stride=stride, dilation=dilation, bias=False),\n\t\tnn.BatchNorm2d(out_channel),\n\t\tnn.ReLU(inplace=True),\n\t)\n\nclass ASPP(nn.Module):\n\t\n\tdef __init__(self,\n\t\t\t\t in_channel,\n\t\t\t\t channel,\n\t\t\t\t dilation,\n\t\t\t\t ):\n\t\tsuper(ASPP, self).__init__()\n\t\t\n\t\tself.conv = nn.ModuleList()\n\t\tfor d in dilation:\n\t\t\tself.conv.append(\n\t\t\t\tConv2dBnReLU(\n\t\t\t\t\tin_channel,\n\t\t\t\t\tchannel,\n\t\t\t\t\tkernel_size=1 if d == 1 else 3,\n\t\t\t\t\tdilation=d,\n\t\t\t\t\tpadding=0 if d == 1 else d,\n\t\t\t\t)\n\t\t\t)\n\t\t\n\t\tself.out = Conv2dBnReLU(\n\t\t\tlen(dilation) * channel,\n\t\t\tchannel,\n\t\t\tkernel_size=3,\n\t\t\tpadding=1,\n\t\t\t)\n\t\n\tdef forward(self, x):\n\t\taspp = []\n\t\tfor conv in self.conv:\n\t\t\taspp.append(conv(x))\n\t\taspp = torch.cat(aspp, dim=1)\n\t\tout = self.out(aspp)\n\t\treturn out\n\n#DepthwiseSeparable\nclass DSConv2d(nn.Module):\n\tdef __init__(self,\n\t\t\t\t in_channel,\n\t\t\t\t out_channel,\n\t\t\t\t kernel_size,\n\t\t\t\t stride   = 1,\n\t\t\t\t padding  = 0,\n\t\t\t\t dilation = 1\n\t\t):\n\t\tsuper().__init__()\n\t\t\n\t\tself.depthwise = nn.Sequential(\n\t\t\tnn.Conv2d( in_channel, in_channel, kernel_size, stride=stride, padding=padding, dilation=dilation),\n\t\t\tnn.BatchNorm2d(in_channel),\n\t\t\tnn.ReLU(inplace=True)\n\t\t)\n\t\t\n\t\tself.pointwise = nn.Sequential(\n\t\t\tnn.Conv2d( in_channel, out_channel, kernel_size=1, stride=1, padding=0),\n\t\t\tnn.BatchNorm2d(out_channel),\n\t\t\tnn.ReLU(inplace=True)\n\t\t)\n\t\n\tdef forward(self, x):\n\t\tx = self.depthwise(x)\n\t\tx = self.pointwise(x)\n\t\treturn x\n\nclass DSASPP(nn.Module):\n\t\n\tdef __init__(self,\n\t\t\t\t in_channel,\n\t\t\t\t channel,\n\t\t\t\t dilation,\n\t\t\t\t ):\n\t\tsuper(DSASPP, self).__init__()\n\t\t\n\t\tself.conv = nn.ModuleList()\n\t\tfor d in dilation:\n\t\t\tif d == 1:\n\t\t\t\tself.conv.append(\n\t\t\t\t\tConv2dBnReLU(\n\t\t\t\t\t\tin_channel,\n\t\t\t\t\t\tchannel,\n\t\t\t\t\t\tkernel_size=1 if d == 1 else 3,\n\t\t\t\t\t\tdilation=d,\n\t\t\t\t\t\tpadding=0 if d == 1 else d,\n\t\t\t\t\t)\n\t\t\t\t)\n\t\t\telse:\n\t\t\t\tself.conv.append(\n\t\t\t\t\tDSConv2d(\n\t\t\t\t\t\tin_channel,\n\t\t\t\t\t\tchannel,\n\t\t\t\t\t\tkernel_size=3,\n\t\t\t\t\t\tdilation=d,\n\t\t\t\t\t\tpadding=d,\n\t\t\t\t\t)\n\t\t\t\t)\n\t\t\n\t\tself.out = Conv2dBnReLU(\n\t\t\tlen(dilation) * channel,\n\t\t\tchannel,\n\t\t\tkernel_size=3,\n\t\t\tpadding=1,\n\t\t)\n\t \n\tdef forward(self, x):\n\t\taspp = []\n\t\tfor conv in self.conv:\n\t\t\taspp.append(conv(x))\n\t\taspp = torch.cat(aspp, dim=1)\n\t\tout = self.out(aspp)\n\t\treturn out\n\n\t\n##############################################################################################33\n\nclass DaformerDecoder(nn.Module):\n\tdef __init__(\n\t\t\tself,\n\t\t\tencoder_dim = [32, 64, 160, 256],\n\t\t\tdecoder_dim = 256,\n\t\t\tdilation = [1, 6, 12, 18],\n\t\t\tuse_bn_mlp  = True,\n\t\t\tfuse = 'conv3x3',\n\t):\n\t\tsuper().__init__()\n\t\tself.mlp = nn.ModuleList([\n\t\t\tnn.Sequential(\n\t\t\t\t# Conv2dBnReLU(dim, decoder_dim, 1, padding=0), #follow mmseg to use conv-bn-relu\n\t\t\t\t*(\n\t\t\t\t  ( nn.Conv2d(dim, decoder_dim, 1, padding= 0,  bias=False),\n\t\t\t\t\tnn.BatchNorm2d(decoder_dim),\n\t\t\t\t\tnn.ReLU(inplace=True),\n\t\t\t\t)if use_bn_mlp else\n\t\t\t\t  ( nn.Conv2d(dim, decoder_dim, 1, padding= 0,  bias=True),)\n\t\t\t\t),\n\t\t\t\t\n\t\t\t\tMixUpSample(2**i) if i!=0 else nn.Identity(),\n\t\t\t) for i, dim in enumerate(encoder_dim)])\n\t  \n\t\tif fuse=='conv1x1':\n\t\t\tself.fuse = nn.Sequential(\n\t\t\t\tnn.Conv2d(len(encoder_dim) * decoder_dim, decoder_dim, 1, padding=0, bias=False),\n\t\t\t\tnn.BatchNorm2d(decoder_dim),\n\t\t\t\tnn.ReLU(inplace=True),\n\t\t\t)\n\t\t\n\t\tif fuse=='conv3x3':\n\t\t\tself.fuse = nn.Sequential(\n\t\t\t\tnn.Conv2d(len(encoder_dim) * decoder_dim, decoder_dim, 3, padding=1, bias=False),\n\t\t\t\tnn.BatchNorm2d(decoder_dim),\n\t\t\t\tnn.ReLU(inplace=True),\n\t\t\t)\n\t\t\n\t\tif fuse=='aspp':\n\t\t\tself.fuse = ASPP(\n\t\t\t\tdecoder_dim*len(encoder_dim),\n\t\t\t\tdecoder_dim,\n\t\t\t\tdilation,\n\t\t\t)\n\t\t\t\n\t\tif fuse=='ds-aspp':\n\t\t\tself.fuse = DSASPP(\n\t\t\t\tdecoder_dim*len(encoder_dim),\n\t\t\t\tdecoder_dim,\n\t\t\t\tdilation,\n\t\t\t)\n\t\t\n\t\n\tdef forward(self, feature):\n\t\t\n\t\tout = []\n\t\tfor i,f in enumerate(feature):\n\t\t\tf = self.mlp[i](f)\n\t\t\tout.append(f)\n\t\t\t# pdb.set_trace()\n\t\t\t#print(f.shape)\n\t\tx = self.fuse(torch.cat(out, dim = 1))\n\t\treturn x, out\n\n\nclass daformer_conv3x3 (DaformerDecoder):\n\tdef __init__(self, **kwargs):\n\t\tsuper(daformer_conv3x3, self).__init__(\n\t\t\tfuse = 'conv3x3',\n\t\t\t**kwargs\n\t\t)\nclass daformer_conv1x1 (DaformerDecoder):\n\tdef __init__(self, **kwargs):\n\t\tsuper(daformer_conv1x1, self).__init__(\n\t\t\tfuse = 'conv1x1',\n\t\t\t**kwargs\n\t\t)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.303062Z","iopub.execute_input":"2022-09-15T17:40:33.303678Z","iopub.status.idle":"2022-09-15T17:40:33.333896Z","shell.execute_reply.started":"2022-09-15T17:40:33.303642Z","shell.execute_reply":"2022-09-15T17:40:33.332925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CoAT-5level\n\n# https://github.com/mlpc-ucsd/CoaT/blob/main/src/models/coat.py\n\n\"\"\"\nCoaT architecture.\n\nModified from timm/models/vision_transformer.py\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom timm.data import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD\nfrom timm.models.layers import DropPath, to_2tuple, trunc_normal_\nfrom timm.models.registry import register_model\n\nfrom einops import rearrange\nfrom functools import partial\nfrom torch import nn, einsum\nimport pdb\n\n# __all__ = [\n# \t\"coat_tiny\",\n# \t\"coat_mini\",\n# \t\"coat_small\",\n# \t\"coat_lite_tiny\",\n# \t\"coat_lite_mini\",\n# \t\"coat_lite_small\"\n# ]\n\n\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\n#---------------------------------------------\n#https://github.com/facebookresearch/ConvNeXt/blob/main/models/convnext.py#L15\nclass LayerNorm2d(nn.Module):\n\tdef __init__(self, dim, eps=1e-6):\n\t\tsuper().__init__()\n\t\tself.dim = dim\n\t\tself.weight = nn.Parameter(torch.ones(dim))\n\t\tself.bias = nn.Parameter(torch.zeros(dim))\n\t\tself.eps = eps\n\t\n\tdef forward(self, x):\n\t\tbatch_size,C,H,W = x.shape\n\t\t#assert C==self.dim, 'C=%d, self.dim=%d'%(C,self.dim)\n\t\t#print('C=%d, self.dim=%d'%(C,self.dim))\n\t\t\n\t\tu = x.mean(1, keepdim=True)\n\t\ts = (x - u).pow(2).mean(1, keepdim=True)\n\t\tx = (x - u) / torch.sqrt(s + self.eps)\n\t\tx = self.weight[:, None, None] * x + self.bias[:, None, None]\n\t\treturn x\n#---------------------------------------------\n\ndef _cfg_coat(url='', **kwargs):\n\treturn {\n\t\t'url': url,\n\t\t'num_classes': 1000, 'input_size': (3, 224, 224), 'pool_size': None,\n\t\t'crop_pct': .9, 'interpolation': 'bicubic',\n\t\t'mean': IMAGENET_DEFAULT_MEAN, 'std': IMAGENET_DEFAULT_STD,\n\t\t'first_conv': 'patch_embed.proj', 'classifier': 'head',\n\t\t**kwargs\n\t}\n\n\nclass Mlp(nn.Module):\n\t\"\"\" Feed-forward network (FFN, a.k.a. MLP) class. \"\"\"\n\tdef __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):\n\t\tsuper().__init__()\n\t\tout_features = out_features or in_features\n\t\thidden_features = hidden_features or in_features\n\t\tself.fc1 = nn.Linear(in_features, hidden_features)\n\t\tself.act = act_layer()\n\t\tself.fc2 = nn.Linear(hidden_features, out_features)\n\t\tself.drop = nn.Dropout(drop)\n\t\n\tdef forward(self, x):\n\t\tx = self.fc1(x)\n\t\tx = self.act(x)\n\t\tx = self.drop(x)\n\t\tx = self.fc2(x)\n\t\tx = self.drop(x)\n\t\treturn x\n\n\nclass ConvRelPosEnc(nn.Module):\n\t\"\"\" Convolutional relative position encoding. \"\"\"\n\tdef __init__(self, Ch, h, window):\n\t\t\"\"\"\n\t\tInitialization.\n\t\t\tCh: Channels per head.\n\t\t\th: Number of heads.\n\t\t\twindow: Window size(s) in convolutional relative positional encoding. It can have two forms:\n\t\t\t\t\t1. An integer of window size, which assigns all attention heads with the same window size in ConvRelPosEnc.\n\t\t\t\t\t2. A dict mapping window size to #attention head splits (e.g. {window size 1: #attention head split 1, window size 2: #attention head split 2})\n\t\t\t\t\t   It will apply different window size to the attention head splits.\n\t\t\"\"\"\n\t\tsuper().__init__()\n\t\t\n\t\tif isinstance(window, int):\n\t\t\twindow = {window: h}                                                         # Set the same window size for all attention heads.\n\t\t\tself.window = window\n\t\telif isinstance(window, dict):\n\t\t\tself.window = window\n\t\telse:\n\t\t\traise ValueError()\n\t\t\n\t\tself.conv_list = nn.ModuleList()\n\t\tself.head_splits = []\n\t\tfor cur_window, cur_head_split in window.items():\n\t\t\tdilation = 1                                                                 # Use dilation=1 at default.\n\t\t\tpadding_size = (cur_window + (cur_window - 1) * (dilation - 1)) // 2         # Determine padding size. Ref: https://discuss.pytorch.org/t/how-to-keep-the-shape-of-input-and-output-same-when-dilation-conv/14338\n\t\t\tcur_conv = nn.Conv2d(cur_head_split*Ch, cur_head_split*Ch,\n\t\t\t                     kernel_size=(cur_window, cur_window),\n\t\t\t                     padding=(padding_size, padding_size),\n\t\t\t                     dilation=(dilation, dilation),\n\t\t\t                     groups=cur_head_split*Ch,\n\t\t\t                     )\n\t\t\tself.conv_list.append(cur_conv)\n\t\t\tself.head_splits.append(cur_head_split)\n\t\tself.channel_splits = [x*Ch for x in self.head_splits]\n\t\n\tdef forward(self, q, v, size):\n\t\tB, h, N, Ch = q.shape\n\t\tH, W = size\n\t\tassert N == 1 + H * W\n\t\t\n\t\t# Convolutional relative position encoding.\n\t\tq_img = q[:,:,1:,:]                                                              # Shape: [B, h, H*W, Ch].\n\t\tv_img = v[:,:,1:,:]                                                              # Shape: [B, h, H*W, Ch].\n\t\t\n\t\tv_img = rearrange(v_img, 'B h (H W) Ch -> B (h Ch) H W', H=H, W=W)               # Shape: [B, h, H*W, Ch] -> [B, h*Ch, H, W].\n\t\tv_img_list = torch.split(v_img, self.channel_splits, dim=1)                      # Split according to channels.\n\t\tconv_v_img_list = [conv(x) for conv, x in zip(self.conv_list, v_img_list)]\n\t\tconv_v_img = torch.cat(conv_v_img_list, dim=1)\n\t\tconv_v_img = rearrange(conv_v_img, 'B (h Ch) H W -> B h (H W) Ch', h=h)          # Shape: [B, h*Ch, H, W] -> [B, h, H*W, Ch].\n\t\t\n\t\tEV_hat_img = q_img * conv_v_img\n\t\tzero = torch.zeros((B, h, 1, Ch), dtype=q.dtype, layout=q.layout, device=q.device)\n\t\tEV_hat = torch.cat((zero, EV_hat_img), dim=2)                                # Shape: [B, h, N, Ch].\n\t\t\n\t\treturn EV_hat\n\n\nclass FactorAtt_ConvRelPosEnc(nn.Module):\n\t\"\"\" Factorized attention with convolutional relative position encoding class. \"\"\"\n\tdef __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0., shared_crpe=None):\n\t\tsuper().__init__()\n\t\tself.num_heads = num_heads\n\t\thead_dim = dim // num_heads\n\t\tself.scale = qk_scale or head_dim ** -0.5\n\t\t\n\t\tself.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)\n\t\tself.attn_drop = nn.Dropout(attn_drop)                                           # Note: attn_drop is actually not used.\n\t\tself.proj = nn.Linear(dim, dim)\n\t\tself.proj_drop = nn.Dropout(proj_drop)\n\t\t\n\t\t# Shared convolutional relative position encoding.\n\t\tself.crpe = shared_crpe\n\t\n\tdef forward(self, x, size):\n\t\tB, N, C = x.shape\n\t\t\n\t\t# Generate Q, K, V.\n\t\tqkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)  # Shape: [3, B, h, N, Ch].\n\t\tq, k, v = qkv[0], qkv[1], qkv[2]                                                 # Shape: [B, h, N, Ch].\n\t\t\n\t\t# Factorized attention.\n\t\tk_softmax = k.softmax(dim=2)                                                     # Softmax on dim N.\n\t\tk_softmax_T_dot_v = einsum('b h n k, b h n v -> b h k v', k_softmax, v)          # Shape: [B, h, Ch, Ch].\n\t\tfactor_att        = einsum('b h n k, b h k v -> b h n v', q, k_softmax_T_dot_v)  # Shape: [B, h, N, Ch].\n\t\t\n\t\t# Convolutional relative position encoding.\n\t\tcrpe = self.crpe(q, v, size=size)                                                # Shape: [B, h, N, Ch].\n\t\t\n\t\t# Merge and reshape.\n\t\tx = self.scale * factor_att + crpe\n\t\tx = x.transpose(1, 2).reshape(B, N, C)                                           # Shape: [B, h, N, Ch] -> [B, N, h, Ch] -> [B, N, C].\n\t\t\n\t\t# Output projection.\n\t\tx = self.proj(x)\n\t\tx = self.proj_drop(x)\n\t\t\n\t\treturn x                                                                         # Shape: [B, N, C].\n\n\nclass ConvPosEnc(nn.Module):\n\t\"\"\" Convolutional Position Encoding.\n\t\tNote: This module is similar to the conditional position encoding in CPVT.\n\t\"\"\"\n\tdef __init__(self, dim, k=3):\n\t\tsuper(ConvPosEnc, self).__init__()\n\t\tself.proj = nn.Conv2d(dim, dim, k, 1, k//2, groups=dim)\n\t\n\tdef forward(self, x, size):\n\t\tB, N, C = x.shape\n\t\tH, W = size\n\t\tassert N == 1 + H * W\n\t\t\n\t\t# Extract CLS token and image tokens.\n\t\tcls_token, img_tokens = x[:, :1], x[:, 1:]                                       # Shape: [B, 1, C], [B, H*W, C].\n\t\t\n\t\t# Depthwise convolution.\n\t\tfeat = img_tokens.transpose(1, 2).view(B, C, H, W)\n\t\tx = self.proj(feat) + feat\n\t\tx = x.flatten(2).transpose(1, 2)\n\t\t\n\t\t# Combine with CLS token.\n\t\tx = torch.cat((cls_token, x), dim=1)\n\t\t\n\t\treturn x\n\n\nclass SerialBlock(nn.Module):\n\t\"\"\" Serial block class.\n\t\tNote: In this implementation, each serial block only contains a conv-attention and a FFN (MLP) module. \"\"\"\n\tdef __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,\n\t             drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm,\n\t             shared_cpe=None, shared_crpe=None):\n\t\tsuper().__init__()\n\t\t\n\t\t# Conv-Attention.\n\t\tself.cpe = shared_cpe\n\t\t\n\t\tself.norm1 = norm_layer(dim)\n\t\tself.factoratt_crpe = FactorAtt_ConvRelPosEnc(\n\t\t\tdim, num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop,\n\t\t\tshared_crpe=shared_crpe)\n\t\tself.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\t\t\n\t\t# MLP.\n\t\tself.norm2 = norm_layer(dim)\n\t\tmlp_hidden_dim = int(dim * mlp_ratio)\n\t\tself.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)\n\t\n\tdef forward(self, x, size):\n\t\t# Conv-Attention.\n\t\tx = self.cpe(x, size)                  # Apply convolutional position encoding.\n\t\tcur = self.norm1(x)\n\t\tcur = self.factoratt_crpe(cur, size)   # Apply factorized attention and convolutional relative position encoding.\n\t\tx = x + self.drop_path(cur)\n\t\t\n\t\t# MLP.\n\t\tcur = self.norm2(x)\n\t\tcur = self.mlp(cur)\n\t\tx = x + self.drop_path(cur)\n\t\t\n\t\treturn x\n\n\nclass ParallelBlock(nn.Module):\n\t\"\"\" Parallel block class. \"\"\"\n\tdef __init__(self, dims, num_heads, mlp_ratios=[], qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,\n\t             drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm,\n\t             shared_cpes=None, shared_crpes=None):\n\t\tsuper().__init__()\n\t\t\n\t\t# Conv-Attention.\n\t\tself.cpes = shared_cpes\n\t\t\n\t\tself.norm12 = norm_layer(dims[1])\n\t\tself.norm13 = norm_layer(dims[2])\n\t\tself.norm14 = norm_layer(dims[3])\n\t\tself.norm15 = norm_layer(dims[4])\n\n\n\t\tself.factoratt_crpe2 = FactorAtt_ConvRelPosEnc(\n\t\t\tdims[1], num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop,\n\t\t\tshared_crpe=shared_crpes[1]\n\t\t)\n\t\tself.factoratt_crpe3 = FactorAtt_ConvRelPosEnc(\n\t\t\tdims[2], num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop,\n\t\t\tshared_crpe=shared_crpes[2]\n\t\t)\n\t\tself.factoratt_crpe4 = FactorAtt_ConvRelPosEnc(\n\t\t\tdims[3], num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop,\n\t\t\tshared_crpe=shared_crpes[3]\n\t\t)\n\t\tself.factoratt_crpe5 = FactorAtt_ConvRelPosEnc(\n\t\t\tdims[4], num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop,\n\t\t\tshared_crpe=shared_crpes[4]\n\t\t)\n\n\n\t\tself.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\t\t\n\t\t# MLP.\n\t\tself.norm22 = norm_layer(dims[1])\n\t\tself.norm23 = norm_layer(dims[2])\n\t\tself.norm24 = norm_layer(dims[3])\n\t\tself.norm25 = norm_layer(dims[4])\n\n\t\tassert dims[1] == dims[2] == dims[3] ==dims[4]                             # In parallel block, we assume dimensions are the same and share the linear transformation.\n\t\tassert mlp_ratios[1] == mlp_ratios[2] == mlp_ratios[3]\n\t\tmlp_hidden_dim = int(dims[1] * mlp_ratios[1])\n\t\tself.mlp2 = self.mlp3 = self.mlp4 =self.mlp5= Mlp(in_features=dims[1], hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)\n\t\n\tdef upsample(self, x, output_size, size):\n\t\t\"\"\" Feature map up-sampling. \"\"\"\n\t\treturn self.interpolate(x, output_size=output_size, size=size)\n\t\n\tdef downsample(self, x, output_size, size):\n\t\t\"\"\" Feature map down-sampling. \"\"\"\n\t\treturn self.interpolate(x, output_size=output_size, size=size)\n\t\n\tdef interpolate(self, x, output_size, size):\n\t\t\"\"\" Feature map interpolation. \"\"\"\n\t\tB, N, C = x.shape\n\t\tH, W = size\n\t\tassert N == 1 + H * W\n\t\t\n\t\tcls_token  = x[:, :1, :]\n\t\timg_tokens = x[:, 1:, :]\n\t\t\n\t\timg_tokens = img_tokens.transpose(1, 2).reshape(B, C, H, W)\n\t\timg_tokens = F.interpolate(img_tokens, size=output_size, mode='bilinear')  # FIXME: May have alignment issue.\n\t\timg_tokens = img_tokens.reshape(B, C, -1).transpose(1, 2)\n\t\t\n\t\tout = torch.cat((cls_token, img_tokens), dim=1)\n\t\t\n\t\treturn out\n\t\n\tdef forward(self, x1, x2, x3, x4, x5,sizes):\n\t\t_, (H2, W2), (H3, W3), (H4, W4),(H5,W5) = sizes\n\t\t\n\t\t# Conv-Attention.\n\t\tx2 = self.cpes[1](x2, size=(H2, W2))  # Note: x1 is ignored.\n\t\tx3 = self.cpes[2](x3, size=(H3, W3))\n\t\tx4 = self.cpes[3](x4, size=(H4, W4))\n\t\tx5 = self.cpes[4](x5, size=(H5, W5))\n\t\t\n\t\tcur2 = self.norm12(x2)\n\t\tcur3 = self.norm13(x3)\n\t\tcur4 = self.norm14(x4)\n\t\tcur5 = self.norm15(x5)\n\n\t\tcur2 = self.factoratt_crpe2(cur2, size=(H2,W2))\n\t\tcur3 = self.factoratt_crpe3(cur3, size=(H3,W3))\n\t\tcur4 = self.factoratt_crpe4(cur4, size=(H4,W4))\n\t\tcur5 = self.factoratt_crpe4(cur5, size=(H5,W5))\n\n\n\t\tupsample3_2 = self.upsample(cur3, output_size=(H2,W2), size=(H3,W3))\n\t\tupsample4_3 = self.upsample(cur4, output_size=(H3,W3), size=(H4,W4))\n\t\tupsample4_2 = self.upsample(cur4, output_size=(H2,W2), size=(H4,W4))\n\t\tdownsample2_3 = self.downsample(cur2, output_size=(H3,W3), size=(H2,W2))\n\t\tdownsample3_4 = self.downsample(cur3, output_size=(H4,W4), size=(H3,W3))\n\t\tdownsample2_4 = self.downsample(cur2, output_size=(H4,W4), size=(H2,W2))\n\t\tupsample5_2 = self.upsample(cur5, output_size=(H2,W2), size=(H5,W5))\n\t\tupsample5_3 = self.upsample(cur5, output_size=(H3,W3), size=(H5,W5))\n\t\tdownsample3_5 = self.downsample(cur3, output_size=(H5,W5), size=(H3,W3))\n\t\tupsample5_4 = self.upsample(cur5, output_size=(H4,W4), size=(H5,W5))\n\t\tdownsample2_5 = self.downsample(cur2, output_size=(H5,W5), size=(H2,W2))\n\t\tdownsample4_5 = self.downsample(cur4, output_size=(H5,W5), size=(H4,W4))\n        \n     \n\n\t\t# cur2 = cur2  + upsample3_2   + upsample4_2+  upsample5_2\n\t\t# cur3 = cur3  + upsample4_3   + downsample2_3 + upsample5_3\n\t\t# cur4 = cur4  + downsample3_4 + downsample2_4 + upsample5_4 \n\t\t# cur5 = cur5  + downsample3_5 + downsample2_5 + downsample4_5 4\n        # not use all connetction\n\t\tcur2 = cur2  + upsample3_2   + upsample4_2\n\t\tcur3 = cur3  + upsample4_3   + downsample2_3\n\t\tcur4 = cur4  + upsample5_4   + downsample2_4\n\t\tcur5 = cur5  + downsample4_5 + downsample2_5\n\n\n\t\tx2 = x2 + self.drop_path(cur2)\n\t\tx3 = x3 + self.drop_path(cur3)\n\t\tx4 = x4 + self.drop_path(cur4)\n\t\tx5 = x5 + self.drop_path(cur5)\n\t\t\n\t\t# MLP.\n\t\tcur2 = self.norm22(x2)\n\t\tcur3 = self.norm23(x3)\n\t\tcur4 = self.norm24(x4)\n\t\tcur5 = self.norm25(x5)\n\n\t\tcur2 = self.mlp2(cur2)\n\t\tcur3 = self.mlp3(cur3)\n\t\tcur4 = self.mlp4(cur4)\n\t\tcur5 = self.mlp5(cur5)\n\n\t\tx2 = x2 + self.drop_path(cur2)\n\t\tx3 = x3 + self.drop_path(cur3)\n\t\tx4 = x4 + self.drop_path(cur4)\n\t\tx5 = x5 + self.drop_path(cur5)\n\t\t\n\t\treturn x1, x2, x3, x4, x5\n\n\nclass PatchEmbed(nn.Module):\n\t\"\"\" Image to Patch Embedding \"\"\"\n\tdef __init__(self, patch_size=16, in_chans=3, embed_dim=768):\n\t\tsuper().__init__()\n\t\tpatch_size = to_2tuple(patch_size)\n\t\t\n\t\tself.patch_size = patch_size\n\t\tself.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n\t\tself.norm = nn.LayerNorm(embed_dim)\n\t\n\tdef forward(self, x):\n\t\t_, _, H, W = x.shape\n\t\tout_H, out_W = H // self.patch_size[0], W // self.patch_size[1]\n\t\t\n\t\tx = self.proj(x).flatten(2).transpose(1, 2)\n\t\tout = self.norm(x)\n\t\t\n\t\treturn out, (out_H, out_W)\n\n\nclass CoaT(nn.Module):\n\t\"\"\" CoaT class. \"\"\"\n\tdef __init__(self, patch_size=16, in_chans=3, embed_dims=[0, 0, 0, 0],\n\t             serial_depths=[0, 0, 0, 0], parallel_depth=0,\n\t             num_heads=0, mlp_ratios=[0, 0, 0, 0], qkv_bias=True, qk_scale=None, drop_rate=0., attn_drop_rate=0.,\n\t             drop_path_rate=0.,\n\t             norm_layer=partial(nn.LayerNorm, eps=1e-6),\n\t             return_interm_layers=True,\n\t             out_features=['x1_nocls','x2_nocls','x3_nocls','x4_nocls','x5_nocls'],\n\t             crpe_window={3:2, 5:3, 7:3},\n\t             pretrain=None,\n\t             out_norm = nn.Identity, #use nn.Identity, nn.BatchNorm2d, LayerNorm2d\n\t             **kwargs):\n\t\tsuper().__init__()\n\t\tself.return_interm_layers = return_interm_layers\n\t\tself.pretrain     = pretrain\n\t\tself.embed_dims   = embed_dims\n\t\tself.out_features = out_features\n\t\t#self.num_classes  = num_classes\n\t\t\n\t\t# Patch embeddings.\n\t\tself.patch_embed1 = PatchEmbed(patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dims[0])\n\t\tself.patch_embed2 = PatchEmbed(patch_size=2, in_chans=embed_dims[0], embed_dim=embed_dims[1])\n\t\tself.patch_embed3 = PatchEmbed(patch_size=2, in_chans=embed_dims[1], embed_dim=embed_dims[2])\n\t\tself.patch_embed4 = PatchEmbed(patch_size=2, in_chans=embed_dims[2], embed_dim=embed_dims[3])\n\t\tself.patch_embed5 = PatchEmbed(patch_size=2, in_chans=embed_dims[3], embed_dim=embed_dims[4])\n\t\t\n\t\t# Class tokens.\n\t\tself.cls_token1 = nn.Parameter(torch.zeros(1, 1, embed_dims[0]))\n\t\tself.cls_token2 = nn.Parameter(torch.zeros(1, 1, embed_dims[1]))\n\t\tself.cls_token3 = nn.Parameter(torch.zeros(1, 1, embed_dims[2]))\n\t\tself.cls_token4 = nn.Parameter(torch.zeros(1, 1, embed_dims[3]))\n\t\tself.cls_token5 = nn.Parameter(torch.zeros(1, 1, embed_dims[4]))\n\t\t\n\t\t# Convolutional position encodings.\n\t\tself.cpe1 = ConvPosEnc(dim=embed_dims[0], k=3)\n\t\tself.cpe2 = ConvPosEnc(dim=embed_dims[1], k=3)\n\t\tself.cpe3 = ConvPosEnc(dim=embed_dims[2], k=3)\n\t\tself.cpe4 = ConvPosEnc(dim=embed_dims[3], k=3)\n\t\tself.cpe5 = ConvPosEnc(dim=embed_dims[4], k=3)\n\t\t\n\t\t# Convolutional relative position encodings.\n\t\tself.crpe1 = ConvRelPosEnc(Ch=embed_dims[0] // num_heads, h=num_heads, window=crpe_window)\n\t\tself.crpe2 = ConvRelPosEnc(Ch=embed_dims[1] // num_heads, h=num_heads, window=crpe_window)\n\t\tself.crpe3 = ConvRelPosEnc(Ch=embed_dims[2] // num_heads, h=num_heads, window=crpe_window)\n\t\tself.crpe4 = ConvRelPosEnc(Ch=embed_dims[3] // num_heads, h=num_heads, window=crpe_window)\n\t\tself.crpe5 = ConvRelPosEnc(Ch=embed_dims[4] // num_heads, h=num_heads, window=crpe_window)\n\n\t\t# Enable stochastic depth.\n\t\tdpr = drop_path_rate\n\t\t\n\t\t# Serial blocks 1.\n\t\tself.serial_blocks1 = nn.ModuleList([\n\t\t\tSerialBlock(\n\t\t\t\tdim=embed_dims[0], num_heads=num_heads, mlp_ratio=mlp_ratios[0], qkv_bias=qkv_bias, qk_scale=qk_scale,\n\t\t\t\tdrop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr, norm_layer=norm_layer,\n\t\t\t\tshared_cpe=self.cpe1, shared_crpe=self.crpe1\n\t\t\t)\n\t\t\tfor _ in range(serial_depths[0])]\n\t\t)\n\t\t\n\t\t# Serial blocks 2.\n\t\tself.serial_blocks2 = nn.ModuleList([\n\t\t\tSerialBlock(\n\t\t\t\tdim=embed_dims[1], num_heads=num_heads, mlp_ratio=mlp_ratios[1], qkv_bias=qkv_bias, qk_scale=qk_scale,\n\t\t\t\tdrop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr, norm_layer=norm_layer,\n\t\t\t\tshared_cpe=self.cpe2, shared_crpe=self.crpe2\n\t\t\t)\n\t\t\tfor _ in range(serial_depths[1])]\n\t\t)\n\t\t\n\t\t# Serial blocks 3.\n\t\tself.serial_blocks3 = nn.ModuleList([\n\t\t\tSerialBlock(\n\t\t\t\tdim=embed_dims[2], num_heads=num_heads, mlp_ratio=mlp_ratios[2], qkv_bias=qkv_bias, qk_scale=qk_scale,\n\t\t\t\tdrop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr, norm_layer=norm_layer,\n\t\t\t\tshared_cpe=self.cpe3, shared_crpe=self.crpe3\n\t\t\t)\n\t\t\tfor _ in range(serial_depths[2])]\n\t\t)\n\t\t\n\t\t# Serial blocks 4.\n\t\tself.serial_blocks4 = nn.ModuleList([\n\t\t\tSerialBlock(\n\t\t\t\tdim=embed_dims[3], num_heads=num_heads, mlp_ratio=mlp_ratios[3], qkv_bias=qkv_bias, qk_scale=qk_scale,\n\t\t\t\tdrop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr, norm_layer=norm_layer,\n\t\t\t\tshared_cpe=self.cpe4, shared_crpe=self.crpe4\n\t\t\t)\n\t\t\tfor _ in range(serial_depths[3])]\n\t\t)\n\n\t\t# Serial blocks 5.\n\t\tself.serial_blocks5 = nn.ModuleList([\n\t\t\tSerialBlock(\n\t\t\t\tdim=embed_dims[4], num_heads=num_heads, mlp_ratio=mlp_ratios[4], qkv_bias=qkv_bias, qk_scale=qk_scale,\n\t\t\t\tdrop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr, norm_layer=norm_layer,\n\t\t\t\tshared_cpe=self.cpe4, shared_crpe=self.crpe4\n\t\t\t)\n\t\t\tfor _ in range(serial_depths[4])]\n\t\t)\n\t\t\n\t\t# Parallel blocks.\n\t\tself.parallel_depth = parallel_depth\n\t\tif self.parallel_depth > 0:\n\t\t\tself.parallel_blocks = nn.ModuleList([\n\t\t\t\tParallelBlock(\n\t\t\t\t\tdims=embed_dims, num_heads=num_heads, mlp_ratios=mlp_ratios, qkv_bias=qkv_bias, qk_scale=qk_scale,\n\t\t\t\t\tdrop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr, norm_layer=norm_layer,\n\t\t\t\t\tshared_cpes=[self.cpe1, self.cpe2, self.cpe3, self.cpe4,  self.cpe5],\n\t\t\t\t\tshared_crpes=[self.crpe1, self.crpe2, self.crpe3, self.crpe4,  self.cpe5]\n\t\t\t\t)\n\t\t\t\tfor _ in range(parallel_depth)]\n\t\t\t)\n\t\t\n\t\t# Classification head(s).\n\t\t# if not self.return_interm_layers:\n\t\t# \tself.norm1 = norm_layer(embed_dims[0])\n\t\t# \tself.norm2 = norm_layer(embed_dims[1])\n\t\t# \tself.norm3 = norm_layer(embed_dims[2])\n\t\t# \tself.norm4 = norm_layer(embed_dims[3])\n\t\t#\n\t\t# \tif self.parallel_depth > 0:                                  # CoaT series: Aggregate features of last three scales for classification.\n\t\t# \t\tassert embed_dims[1] == embed_dims[2] == embed_dims[3]\n\t\t# \t\tself.aggregate = torch.nn.Conv1d(in_channels=3, out_channels=1, kernel_size=1)\n\t\t# \t\tself.head = nn.Linear(embed_dims[3], num_classes)\n\t\t# \telse:\n\t\t# \t\tself.head = nn.Linear(embed_dims[3], num_classes)        # CoaT-Lite series: Use feature of last scale for classification.\n\t\t#---\n\t\t# add a norm layer for each output\n\t\tself.out_norm = nn.ModuleList(\n\t\t\t[ out_norm(embed_dims[i]) for i in range(4)]\n\t\t)\n\t\t\n\t\t# Initialize weights.\n\t\ttrunc_normal_(self.cls_token1, std=.02)\n\t\ttrunc_normal_(self.cls_token2, std=.02)\n\t\ttrunc_normal_(self.cls_token3, std=.02)\n\t\ttrunc_normal_(self.cls_token4, std=.02)\n\t\ttrunc_normal_(self.cls_token5, std=.02)\n\t\tself.apply(self._init_weights)\n\t\n\tdef _init_weights(self, m):\n\t\tif isinstance(m, nn.Linear):\n\t\t\ttrunc_normal_(m.weight, std=.02)\n\t\t\tif isinstance(m, nn.Linear) and m.bias is not None:\n\t\t\t\tnn.init.constant_(m.bias, 0)\n\t\telif isinstance(m, nn.LayerNorm):\n\t\t\tnn.init.constant_(m.bias, 0)\n\t\t\tnn.init.constant_(m.weight, 1.0)\n\t\n\t@torch.jit.ignore\n\tdef no_weight_decay(self):\n\t\treturn {'cls_token1', 'cls_token2', 'cls_token3', 'cls_token4'}\n\t\n\t# def get_classifier(self):\n\t# \treturn self.head\n\t\n\t# def reset_classifier(self, num_classes, global_pool=''):\n\t# \tself.num_classes = num_classes\n\t# \tself.head = nn.Linear(self.embed_dim, num_classes) if num_classes > 0 else nn.Identity()\n\t\n\tdef insert_cls(self, x, cls_token):\n\t\t\"\"\" Insert CLS token. \"\"\"\n\t\tcls_tokens = cls_token.expand(x.shape[0], -1, -1)\n\t\tx = torch.cat((cls_tokens, x), dim=1)\n\t\treturn x\n\t\n\tdef remove_cls(self, x):\n\t\t\"\"\" Remove CLS token. \"\"\"\n\t\treturn x[:, 1:, :]\n\t\n\tdef forward(self, x0):\n\t\tB = x0.shape[0]\n\t\t\n\t\t\n\t\t# Serial blocks 1.\n\t\tx1, (H1, W1) = self.patch_embed1(x0)\n\t\tcls = self.cls_token1#torch.zeros_like(self.cls_token1)#self.cls_token1\n\t\tx1 = self.insert_cls(x1, cls)\n\t\tfor blk in self.serial_blocks1:\n\t\t\tx1 = blk(x1, size=(H1, W1))\n\t\tx1_nocls = self.remove_cls(x1)\n\t\tx1_nocls = x1_nocls.reshape(B, H1, W1, -1).permute(0, 3, 1, 2).contiguous()\n\t\t\n\t\t# Serial blocks 2.\n\t\tx2, (H2, W2) = self.patch_embed2(x1_nocls)\n\t\tcls = self.cls_token2# torch.zeros_like(self.cls_token2)#self.cls_token2#\n\t\tx2 = self.insert_cls(x2,cls)\n\t\tfor blk in self.serial_blocks2:\n\t\t\tx2 = blk(x2, size=(H2, W2))\n\t\tx2_nocls = self.remove_cls(x2)\n\t\tx2_nocls = x2_nocls.reshape(B, H2, W2, -1).permute(0, 3, 1, 2).contiguous()\n\t\t\n\t\t# Serial blocks 3.\n\t\tx3, (H3, W3) = self.patch_embed3(x2_nocls)\n\t\tcls = self.cls_token3#torch.zeros_like(self.cls_token3)# self.cls_token3\n\t\tx3 = self.insert_cls(x3, cls)\n\t\tfor blk in self.serial_blocks3:\n\t\t\tx3 = blk(x3, size=(H3, W3))\n\t\tx3_nocls = self.remove_cls(x3)\n\t\tx3_nocls = x3_nocls.reshape(B, H3, W3, -1).permute(0, 3, 1, 2).contiguous()\n\t\t\n\t\t# Serial blocks 4.\n\t\tx4, (H4, W4) = self.patch_embed4(x3_nocls)\n\t\tcls = self.cls_token5#torch.zeros_like(self.cls_token4)#self.cls_token4\n\t\tx4 = self.insert_cls(x4, cls)\n\t\tfor blk in self.serial_blocks4:\n\t\t\tx4 = blk(x4, size=(H4, W4))\n\t\tx4_nocls = self.remove_cls(x4)\n\t\tx4_nocls = x4_nocls.reshape(B, H4, W4, -1).permute(0, 3, 1, 2).contiguous()\n\n\t\t# Serial blocks 5.\n\t\tx5, (H5, W5) = self.patch_embed4(x4_nocls)\n\t\tcls = self.cls_token4#torch.zeros_like(self.cls_token4)#self.cls_token4\n\t\tx5 = self.insert_cls(x5, cls)\n\t\tfor blk in self.serial_blocks5:\n\t\t\tx5 = blk(x5, size=(H5, W5))\n\t\tx5_nocls = self.remove_cls(x5)\n\t\tx5_nocls = x5_nocls.reshape(B, H5, W5, -1).permute(0, 3, 1, 2).contiguous()\n\t\t\n\t\t# Only serial blocks: Early return.\n\t\tif self.parallel_depth == 0:\n\t\t\tx1_nocls = self.out_norm[0](x1_nocls)\n\t\t\tx2_nocls = self.out_norm[1](x2_nocls)\n\t\t\tx3_nocls = self.out_norm[2](x3_nocls)\n\t\t\tx4_nocls = self.out_norm[3](x4_nocls)\n\t\t\treturn [x1_nocls,x2_nocls,x3_nocls,x4_nocls]\n\t\t \n\t \n\t\t\t\n\t\t\n\t\t# Parallel blocks.\n\t\tfor blk in self.parallel_blocks:\n\t\t\tx1, x2, x3, x4,x5 = blk(x1, x2, x3, x4,x5, sizes=[(H1, W1), (H2, W2), (H3, W3), (H4, W4), (H5, W5)])\n\t\t# pdb.set_trace()\n\t\t# remove cls and return feature for seg\n\t\tif self.return_interm_layers:       # Return intermediate features for down-stream tasks (e.g. Deformable DETR and Detectron2).\n\t\t\tfeat_out = {}\n\t\t\tif 'x1_nocls' in self.out_features:\n\t\t\t\tx1_nocls = self.remove_cls(x1)\n\t\t\t\tx1_nocls = x1_nocls.reshape(B, H1, W1, -1).permute(0, 3, 1, 2).contiguous()\n\t\t\t\tfeat_out['x1_nocls'] = x1_nocls\n\t\t\tif 'x2_nocls' in self.out_features:\n\t\t\t\tx2_nocls = self.remove_cls(x2)\n\t\t\t\tx2_nocls = x2_nocls.reshape(B, H2, W2, -1).permute(0, 3, 1, 2).contiguous()\n\t\t\t\tfeat_out['x2_nocls'] = x2_nocls\n\t\t\tif 'x3_nocls' in self.out_features:\n\t\t\t\tx3_nocls = self.remove_cls(x3)\n\t\t\t\tx3_nocls = x3_nocls.reshape(B, H3, W3, -1).permute(0, 3, 1, 2).contiguous()\n\t\t\t\tfeat_out['x3_nocls'] = x3_nocls\n\t\t\tif 'x4_nocls' in self.out_features:\n\t\t\t\tx4_nocls = self.remove_cls(x4)\n\t\t\t\tx4_nocls = x4_nocls.reshape(B, H4, W4, -1).permute(0, 3, 1, 2).contiguous()\n\t\t\t\tfeat_out['x4_nocls'] = x4_nocls\n\t\t\tif 'x5_nocls' in self.out_features:\n\t\t\t\tx5_nocls = self.remove_cls(x5)\n\t\t\t\tx5_nocls = x5_nocls.reshape(B, H5, W5, -1).permute(0, 3, 1, 2).contiguous()\n\t\t\t\tfeat_out['x5_nocls'] = x5_nocls\n\t\t\tfeat_out = list(feat_out.values())\n\t\t\treturn feat_out\n\t\telse:\n\t\t\tx2 = self.norm2(x2)\n\t\t\tx3 = self.norm3(x3)\n\t\t\tx4 = self.norm4(x4)\n\t\t\tx2_cls = x2[:, :1]              # Shape: [B, 1, C].\n\t\t\tx3_cls = x3[:, :1]\n\t\t\tx4_cls = x4[:, :1]\n\t\t\tmerged_cls = torch.cat((x2_cls, x3_cls, x4_cls), dim=1)       # Shape: [B, 3, C].\n\t\t\tmerged_cls = self.aggregate(merged_cls).squeeze(dim=1)        # Shape: [B, C].\n\t\t\treturn merged_cls\n\t\n \n\n\n\n#CoaT.\n#@register_model\n# def coat_tiny(**kwargs):\n# \tmodel = CoaT(patch_size=4, embed_dims=[152, 152, 152, 152], serial_depths=[2, 2, 2, 2], parallel_depth=6, num_heads=8, mlp_ratios=[4, 4, 4, 4], **kwargs)\n# \tmodel.default_cfg = _cfg_coat()\n# \treturn model\n#\n# #@register_model\n# def coat_mini(**kwargs):\n# \tmodel = CoaT(patch_size=4, embed_dims=[152, 216, 216, 216], serial_depths=[2, 2, 2, 2], parallel_depth=6, num_heads=8, mlp_ratios=[4, 4, 4, 4], **kwargs)\n# \tmodel.default_cfg = _cfg_coat()\n# \treturn model\n#\n# #@register_model\n# def coat_small(**kwargs):\n# \tmodel = CoaT(patch_size=4, embed_dims=[152, 320, 320, 320,320], serial_depths=[2, 2, 2, 2,2], parallel_depth=6, num_heads=8, mlp_ratios=[4,4, 4, 4, 4], **kwargs)\n# \tmodel.default_cfg = _cfg_coat()\n# \treturn model\n# #\n# # CoaT-Lite.\n# #@register_model\n# def coat_lite_tiny(**kwargs):\n# \tmodel = CoaT(patch_size=4, embed_dims=[64, 128, 256, 320], serial_depths=[2, 2, 2, 2], parallel_depth=0, num_heads=8, mlp_ratios=[8, 8, 4, 4], **kwargs)\n# \tmodel.default_cfg = _cfg_coat()\n# \treturn model\n#\n# #@register_model\n# def coat_lite_mini(**kwargs):\n# \tmodel = CoaT(patch_size=4, embed_dims=[64, 128, 320, 512], serial_depths=[2, 2, 2, 2], parallel_depth=0, num_heads=8, mlp_ratios=[8, 8, 4, 4], **kwargs)\n# \tmodel.default_cfg = _cfg_coat()\n# \treturn model\n\n#@register_model\nclass coat_lite_small (CoaT):\n\tdef __init__(self, **kwargs):\n\t\tsuper(coat_lite_small, self).__init__(\n\t        patch_size=4, embed_dims=[64, 128, 320, 512], serial_depths=[2, 2, 2, 2],\n\t\t\tparallel_depth=6, num_heads=8, mlp_ratios=[8,8,4,4], **kwargs)\n\t \n\n#@register_model\nclass coat_lite_medium (CoaT):\n\tdef __init__(self, **kwargs):\n\t\tsuper(coat_lite_medium, self).__init__(\n\t\t\tpatch_size=4, embed_dims=[128, 256, 320, 512],\n\t\t\tserial_depths=[3, 6, 10, 8],\n\t\t    parallel_depth=0, num_heads=8, mlp_ratios=[4, 4, 4, 4],\n\t\t    pretrain = 'coat_lite_medium_384x384_f9129688.pth',\n\t\t\t**kwargs)\n\n#@register_model\nclass coat_parallel_small_plus (CoaT):\n\tdef __init__(self, **kwargs):\n\t\tsuper(coat_parallel_small_plus, self).__init__(\n\t\t\tpatch_size=4, embed_dims=[152, 320, 320, 320,320],\n\t\t\tserial_depths=[2, 2, 2, 2,2],\n\t\t\tparallel_depth=6, num_heads=8, mlp_ratios=[4, 4, 4, 4,4], \n\t\t\tpretrain = 'coat_small_7479cf9b.pth',\n\t\t\t**kwargs)\n\n#@register_model \nclass coat_parallel_small_plus1 (CoaT):\n\tdef __init__(self, **kwargs):\n\t\tsuper(coat_parallel_small_plus1, self).__init__(\n\t\t\tpatch_size=4, embed_dims=[152, 320, 320, 320, 320],\n\t\t\tserial_depths=[2, 2, 2, 2, 2],\n\t\t\tparallel_depth=6, num_heads=8, mlp_ratios=[4, 4, 4, 4, 4], \n\t\t\tpretrain = 'coat_small_7479cf9b.pth',\n\t\t\t**kwargs)\n\n            \n#@register_model\nclass coat_parallel_small (CoaT):\n\tdef __init__(self, **kwargs):\n\t\tsuper(coat_parallel_small, self).__init__(\n\t\t\tpatch_size=4, embed_dims=[152, 320, 320, 320],\n\t\t\tserial_depths=[2, 2, 2, 2],\n\t\t\tparallel_depth=6, num_heads=8, mlp_ratios=[4, 4, 4, 4], \n\t\t\tpretrain = 'coat_small_7479cf9b.pth',\n\t\t\t**kwargs)\n\t\t\t\n\t\t\t\n\t\t\t\n\n\nif 0:\n\tnet = coat_parallel_small_plus1()\n\t# print(net)\n\tx = torch.rand(2,3,786,786)\n\tout = net(x)\n\t# pdb.set_trace()\n\tprint([f.shape for f in out])\n\t\n\t# checkpoint = '/root/share1/data/pretrain_model/coat_lite_medium_384x384_f9129688.pth'\n\t# checkpoint = torch.load(checkpoint, map_location=lambda storage, loc: storage)\n\t# print(list(checkpoint.keys()))\n\t# #['epoch', 'arch', 'state_dict', 'optimizer', 'version', 'args', 'amp', 'metric']\n\t# state_dict = checkpoint['model']\n\t# print( net.load_state_dict(state_dict,strict=False) )","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.335833Z","iopub.execute_input":"2022-09-15T17:40:33.336277Z","iopub.status.idle":"2022-09-15T17:40:33.460639Z","shell.execute_reply.started":"2022-09-15T17:40:33.336162Z","shell.execute_reply":"2022-09-15T17:40:33.459596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nimport pdb\n\n\n#################################################################\n\nclass RGB(nn.Module):\n\tIMAGE_RGB_MEAN = [0.485, 0.456, 0.406]  # [0.5, 0.5, 0.5]\n\tIMAGE_RGB_STD = [0.229, 0.224, 0.225]  # [0.5, 0.5, 0.5]\n\t\n\tdef __init__(self, ):\n\t\tsuper(RGB, self).__init__()\n\t\tself.register_buffer('mean', torch.zeros(1, 3, 1, 1))\n\t\tself.register_buffer('std', torch.ones(1, 3, 1, 1))\n\t\tself.mean.data = torch.FloatTensor(self.IMAGE_RGB_MEAN).view(self.mean.shape)\n\t\tself.std.data = torch.FloatTensor(self.IMAGE_RGB_STD).view(self.std.shape)\n\t\n\tdef forward(self, x):\n\t\tx = (x - self.mean) / self.std\n\t\treturn x\n\n\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\n\nclass Net(nn.Module):\n\t# def load_pretrain( self,):\n\t# \tpath =  '/home/jupyter/share/WRQ/Hubmap/input/coat_lite_medium/coat_lite_medium_384x384_f9129688.pth'\n\t# \tcheckpoint = torch.load(path, map_location=lambda storage, loc: storage)\n\t# \tself.state_dict = checkpoint['model']\n\t# \tself.encoder.load_state_dict(checkpoint,strict=False)\n\n\n\tdef __init__(self,\n\t             \n\t             decoder=daformer_conv1x1,   \n\t\t\t\t \n\t             encoder_cfg={},\n\t             decoder_cfg={},\n\t             ):  # decoder = daformer_conv3x3,   for coat-medium\n\t\tsuper(Net, self).__init__()\n\t\tdecoder_dim = decoder_cfg.get('decoder_dim', 320)\n\t\t\n\t\t# ----\n\t\tself.output_type = ['inference', 'loss']\n\t\t\n\n\t\tself.rgb = RGB()\n    \n\t\t# encoder=coat_lite_medium()\n\t\tself.encoder=coat_parallel_small_plus1()\n\n\n\t\t\n\t\t\n\t\tencoder_dim = self.encoder.embed_dims\n\t\t# pdb.set_trace()\n\t\t# [64, 128, 320, 512]\t\n\t\t# [128, 256, 320, 512],\n\t\t\n\t\tself.decoder = decoder(\n\t\t\tencoder_dim=encoder_dim,\n\t\t\tdecoder_dim=decoder_dim,\n\t\t)\n\t\tself.logit = nn.Sequential(\n\t\t\tnn.Conv2d(decoder_dim, 1, kernel_size=1),\n\t\t)\n\n\t\tself.aux = nn.ModuleList([\n            nn.Conv2d(decoder_dim, 1, kernel_size=1, padding=0) for i in range(4)\n        ])\n\t\tself.avgpool = nn.AdaptiveAvgPool2d((1,1))\n\t\t# self.cls_head = nn.Linear(320,5,bias = False)\n\t\tself.cls_head = nn.Sequential(\n\t\t\t\t\t\t\tnn.BatchNorm1d(320).apply(init_weight),\n\t\t\t\t\t\t\tnn.Linear(320, 128 ).apply(init_weight),\n\t\t\t\t\t\t\t\n\t\t\t\t\t\t\tnn.ReLU(inplace=True),\n\t\t\t\t\t\t\tnn.BatchNorm1d(128).apply(init_weight),\n\t\t\t\t\t\t\tnn.Linear(128, 5).apply(init_weight)\n\t\t\t\t\t\t)\n\t \n\t\n\n\tdef forward(self, batch):\n\t\t\n\t\tx = batch['image']\n\t\t# organ = batch['organ']\n\t\t# mask = batch['mask']\n\t\tnum_class = 5\n        \n# \t\tx = self.rgb(x) #Uncomment if you want to normalise to Imagenet data color\n# \t\tprint(x)\n\t\tB, C, H, W = x.shape\n\t\t# pdb.set_trace()\n        \n\t\tencoder = self.encoder(x)\n\t\t# print([f.shape for f in encoder])\n\t\t# cls_feature = encoder[-1]\n\t\t# cls_feature = self.avgpool(cls_feature)\n\t\t# cls_feature = cls_feature.view(cls_feature.size(0), -1)\n\t\t# cls_feature = self.cls_head(cls_feature)        # (batch,num_class)\n\n\n\t\t# pdb.set_trace()\n\t\tlast, decoder = self.decoder(encoder)\n\t\t# pdb.set_trace()\n\t\tlogit = self.logit(last)\n\t\tlogit = F.interpolate(logit, size=None, scale_factor=4, mode='bilinear', align_corners=False)\n\t\t# print(logit.shape)\n\t\t# pdb.set_trace()\n\t\toutput = {}\n\t\tif 'loss' in self.output_type:\n\t\t\toutput['bce_loss'] = F.binary_cross_entropy_with_logits(logit,batch['mask']) # reduction='mean' # REDUCTION: \"none\" means the loss instance will return the full array of per-sample losses.\n\t\t\t# pdb.set_trace()\n\t\t\t# output[\"label_loss\"] = F.nll_loss(F.log_softmax(cls_feature,dim=1), label)\n\t\t\t# output[\"label_loss\"] = nn.CrossEntropyLoss()(cls_feature,label)\n\t\t\tfor i in range(4):\n\t\t\t\toutput['aux%d_loss'%i] = criterion_aux_loss(self.aux[i](decoder[i]),batch['mask'])\n# \t\t\tprint('BCE_Loss_Shape: ',output['bce_loss'].shape,output['bce_loss'].type,type(output['bce_loss'].shape))     \n\t\t\tprobability_from_logit = torch.sigmoid(logit)\n\t\t\toutput['probability'] = probability_from_logit\n\t\t\t               \n\t\tif 'inference' in self.output_type:\n\t\t\t # TO TRY : ----Calculate Dice, Jaccard here based on organ give the threshold. ) 0.5 for all, lung needs 0.1-----    \n\t\t\t #  Split the logit into batch size -> then check which position organ = lung -> Change somehow either that array ka threshold to 0.1 or scale up values to 0.5Th\n\t\t\t           \n\t\t\tprobability_from_logit = torch.sigmoid(logit)\n\t\t\toutput['probability'] = probability_from_logit\n\t\t\n\t\treturn output\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) # reduction='mean'\n    return loss\n\n \n\n# def one_hot(label, depth = 5):\n\n#     y = torch.zeros(label.size(0), depth)\n#     idx = torch.tensor(label).view(-1, 1).cuda()\n#     y.scatter_(dim = 1, index = idx, value = 1)\n#     return y\n \n# def criterion_multi_binary_cross_entropy(logit, mask, organ):\n# \tlogit = F.interpolate(logit, size=None, scale_factor=4, mode='bilinear', align_corners=False)\n# \tbatch_size, C, H, W = logit.shape\n\t\n# \tlabel = mask.long() * organ.reshape(batch_size, 1, 1, 1)\n# \tonehot = torch.zeros((batch_size, num_organ + 1, H, W)).to(mask)\n# \tonehot = onehot.scatter(1, label, 1)\n# \t# onehot[:,0] = 1-onehot[:,0]\n\t\n# \tloss = F.binary_cross_entropy_with_logits(logit, onehot)\n# \treturn loss\n\n\ndef run_check_net():\n\tbatch_size = 2\n\timage_size = 768\n\t\n\t# ---\n\tbatch = {\n\t\t'image': torch.from_numpy(np.random.uniform(-1, 1, (batch_size, 3, image_size, image_size))).float(),\n\t\t'mask': torch.from_numpy(np.random.choice(2, (batch_size, 1, image_size, image_size))).float(),\n\t\t'organ': torch.from_numpy(np.random.choice(5, (batch_size, 1))).long(),\n\t}\n\tbatch = {k: v for k, v in batch.items()}\n    \n# \tprint(batch.column)\n# \tprint(batch['image'].dtype)\n    \n\tnet = Net()\n\t# net.load_pretrain()\n\twith torch.no_grad():\n\t\twith torch.cuda.amp.autocast(enabled=True):\n\t\t\toutput = net(batch)\n\t\n\tprint('batch')\n\tfor k, v in batch.items():\n\t\tprint('%32s :' % k, v.shape)\n\t\n\tprint('output')\n\tfor k, v in output.items():\n\t\tif 'loss' not in k:\n\t\t\tprint('%32s :' % k, v.shape)\n\tfor k, v in output.items():\n\t\tif 'loss' in k:\n\t\t\tprint('%32s :' % k, v.item())\n\n\n# # # main #################################################################\n# if __name__ == '__main__':\n# \trun_check_net()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.465866Z","iopub.execute_input":"2022-09-15T17:40:33.466141Z","iopub.status.idle":"2022-09-15T17:40:33.493456Z","shell.execute_reply.started":"2022-09-15T17:40:33.466115Z","shell.execute_reply":"2022-09-15T17:40:33.492434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# run_check_net()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.494921Z","iopub.execute_input":"2022-09-15T17:40:33.495627Z","iopub.status.idle":"2022-09-15T17:40:33.509438Z","shell.execute_reply.started":"2022-09-15T17:40:33.495591Z","shell.execute_reply":"2022-09-15T17:40:33.508321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    def __init__(self,n_fold = 5,seed = 42,batch_size = 4,debug=False,kaggle = True):\n        \n        # Directories\n        self.root_dir = \"../input/hubmap-organ-segmentation\" if kaggle else '/media/harish/4tb_hdd/mustaffa/saharsh_hubmap/hubmap-organ-segmentation'\n        self.IMAGES   = '../input/hubmap-stain-norm-original/train_images_stain_normalised' \n        self.MASKS    = '../input/hubmap-hpa-2022-png-dataset/train_masks_png'\n        \n        self.exp_name      = 'Hubmap-Coat-training'\n        \n        self.fold_no       = 0\n        self.n_fold        = n_fold\n        self.folds_to_run  = [0,1,2,3,4]\n        self.seed          = seed\n        # step2: data\n        self.batch_size    = batch_size\n        self.train_bs      = batch_size\n        self.valid_bs      = batch_size*2\n        self.debug         = debug\n        self.img_size      = [768, 768] #..[768,768] # [1617,1617] -> try with batch size of 1\n        \n        # step3: model  \n        self.num_classes   = 1\n        self.device        = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")        \n        self.backbone      = 'coat-small'\n        self.model_name    = \"CoAT\"\n        \n        # step4: optimizer\n        self.epochs        = 30\n        self.lr            = 1e-4 #as per LRRT 4.64e-4 for DICEBCE #0.00085\n        self.lr_drop      = 8\n        self.optimizer     = 'Adam'\n        self.weight_decay  = 1e-5\n        \n        # step5: scheduler\n        self.scheduler     = \"ReduceLROnPlateau\" # ReduceLROnPlateau OR CosineAnnealingLR\n        self.min_lr        = 1e-6 # 0.00005\n        self.T_max         = self.epochs + 50 #int(280/self.batch_size*self.epochs)+50\n#         self.T_0           = 25\n        self.warmup_epochs = 0\n        self.wd            = 1e-6\n        self.n_accumulate  = 1 #max(1, 32//self.batch_size)       \n       \n        \n        # step6: infer\n#         self.thr = 0.4\n        self.thr = {\n            \"Hubmap\":{\n                \"kidney\" : 0.4,\n                \"prostate\":0.25,\n                \"largeintestine\":0.3,\n                \"spleen\":0.25,\n                \"lung\":0.05,  \n            },\n\n             \"HPA\":{\n                \"kidney\" : 0.5,\n                \"prostate\":0.5,\n                \"largeintestine\":0.5,\n                \"spleen\":0.5,\n                \"lung\":0.1,\n\n             },}\n#         self.tta = True\n        \n    def display(self):\n        print(f\"{self.exp_name}\")\n        print(f\"debug is {self.debug}\")        \n        print(f\"Batch size is  {self.batch_size}\")\n        print(f\"img_size is    {self.img_size}\")\n        print(f\"fold_no is     {self.fold_no}\")\n        print(f\"Model is       {self.model_name}_{self.backbone}\")\n        print(f\"epochs is      {self.epochs}\")\n        print(f\"Optimiser is   {self.optimizer}\")\n        print(f\"Scheduler is   {self.scheduler}\")\n        print(f\"LR is          {self.lr}\")\n        \n    \n\ndef set_seed(seed = 42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    print(f\"Setting seed as {seed}\")\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    print('> SEEDING DONE')\n\n\ndef initialize_config(debug=False,batch_size=4):\n    cfg = CFG(batch_size = batch_size,debug = debug)    \n    set_seed(cfg.seed)\n    cfg.display()\n    return cfg","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.510818Z","iopub.execute_input":"2022-09-15T17:40:33.511305Z","iopub.status.idle":"2022-09-15T17:40:33.529175Z","shell.execute_reply.started":"2022-09-15T17:40:33.511269Z","shell.execute_reply":"2022-09-15T17:40:33.527889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_folds(cfg = None):\n    \n    train_csv = pd.read_csv(os.path.join(cfg.root_dir,\"train.csv\"))\n    train_csv.set_index('id',inplace = True)\n    \n    images_path = os.listdir(cfg.IMAGES)\n    images_path = sorted(images_path)\n    masks_path =  os.listdir(cfg.MASKS)\n    masks_path = sorted(masks_path)\n    \n    ids  = [filename[:-4] for filename in masks_path]\n    \n#     files = [filename[:-4] for filename in masks_path]\n#     ids = [f[:-5] for f in files]\n\n    organs      = [train_csv['organ'][int(idx)] for idx in ids]\n    img_height  = [train_csv['img_height'][int(idx)] for idx in ids]\n    img_width   = [train_csv['img_width'][int(idx)] for idx in ids]\n    data_source = [train_csv['data_source'][int(idx)] for idx in ids] \n    rle_mask    = [train_csv['rle'][int(idx)] for idx in ids] \n    images_path = [os.path.join(cfg.IMAGES,f) for f in images_path]\n    masks_path  = [os.path.join(cfg.MASKS,f) for f in masks_path] \n    \n\n    maping = {\n        'id': ids,\n        'organ':organs,\n        'image_path':images_path,\n        'mask_path': masks_path,\n        'data_source': data_source,\n        'rle': rle_mask,\n        'img_height': img_height,\n        'img_width': img_width,\n\n    }\n    df = pd.DataFrame.from_dict(maping)    \n    skf = KFold(n_splits=cfg.n_fold, shuffle=True,random_state=cfg.seed)\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    return df\n\ndef get_optimizer(cfg,optimizer_name= 'Adam'):\n    if optimizer_name == 'Adam':\n        optimizer = optim.Adam(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\n    \n    elif optimizer_name == 'AdamW':\n        optimizer = optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\n        \n    return optimizer\n\n\ndef get_scheduler(cfg,optimizer,df):\n    \n    if len(df[df['fold'] == cfg.folds_to_run[0]]) % cfg.batch_size != 0:\n        num_steps = len(df[df['fold'] != cfg.folds_to_run[0]]) // cfg.batch_size + 1\n       \n    else:\n        num_steps = len(df[df['fold'] != cfg.folds_to_run[0]]) // cfg.batch_size\n\n    if cfg.scheduler == 'CosineAnnealingLR':\n        scheduler = lr_scheduler.CosineAnnealingLR(optimizer,T_max=cfg.T_max,eta_min=cfg.min_lr)\n        \n    elif cfg.scheduler == 'ReduceLROnPlateau':\n        scheduler = lr_scheduler.ReduceLROnPlateau( optimizer  = optimizer, verbose=True, factor=0.7,mode=\"max\",patience=2, threshold=0.001,threshold_mode = 'abs',min_lr =cfg.min_lr)\n        \n    elif cfg.scheduler == 'StepLR':\n        scheduler = lr_scheduler.StepLR(optimizer,cfg.lr_drop ,gamma=0.1)\n        \n    else:\n        scheduler = None\n        \n    return scheduler\n","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.531180Z","iopub.execute_input":"2022-09-15T17:40:33.531720Z","iopub.status.idle":"2022-09-15T17:40:33.549666Z","shell.execute_reply.started":"2022-09-15T17:40:33.531679Z","shell.execute_reply":"2022-09-15T17:40:33.548603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Augmentations**","metadata":{}},{"cell_type":"code","source":"# --------- Augmnetation functions LB = 0.6-------\ndef do_random_flip(image, mask):\n    if np.random.rand()>0.5:\n        image = cv2.flip(image,0)\n        mask = cv2.flip(mask,0)\n    if np.random.rand()>0.5:\n        image = cv2.flip(image,1)\n        mask = cv2.flip(mask,1)\n    if np.random.rand()>0.5:\n        image = image.transpose(1,0,2)\n        mask = mask.transpose(1,0)\n    \n    image = np.ascontiguousarray(image)\n    mask = np.ascontiguousarray(mask)\n    return image, mask\n\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\n#----------------------Below combined aug functions--------------------\n\ndef valid_augment5(image, mask):\n  \n    return image, mask\n\ndef train_augment5b(image, mask):\n    \n    more_transform = A.Compose([        \n         A.OneOf([ \n             A.ElasticTransform(p=1.0, alpha=image.shape[1]*3, sigma=image.shape[1] * 0.07, alpha_affine=image.shape[1] * 0.09),\n             A.GridDistortion(num_steps=5, distort_limit=0.05, p=1.0),\n              ], p=0.25), \n#          A.CoarseDropout(max_holes=8, max_height=image.shape[0]//20, max_width=image.shape[1]//20,\n#                          min_holes=5, fill_value=0, mask_fill_value=0, p=0.25),\n        ], p=1.0)\n\n    \n    \n    image, mask = do_random_flip(image, mask)\n    image, mask = do_random_rot90(image, mask)      \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]) # Remove and try\n    ], 1): 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]),# scale=[0.50, 2.0] or 0.7,1.5 or 0.50,1.5\n    ], 1): image, mask = fn(image, mask)\n\n    augmented = more_transform(image= (image*255).astype(np.uint8), mask=mask)\n    image = augmented['image'].astype(np.float32)/255\n    mask = augmented['mask']\n        \n    return image, mask\n","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.551277Z","iopub.execute_input":"2022-09-15T17:40:33.551693Z","iopub.status.idle":"2022-09-15T17:40:33.577626Z","shell.execute_reply.started":"2022-09-15T17:40:33.551657Z","shell.execute_reply":"2022-09-15T17:40:33.576670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Dataset**\nprepare_loaders -> Prepares the dataset and dataloader, Trainer API needs dataset","metadata":{}},{"cell_type":"code","source":"###############################################################\n##### Data prepare is from original dataset shape which is beetween (0,1)\n###############################################################\n\ndef read_tiff(path, scale=None, verbose=0): \n    image = tifffile.imread(path)\n    if len(image.shape) == 5:\n        image = image.squeeze().transpose(1, 2, 0)\n    \n    if verbose:\n        print(f\"[{path}] Image shape: {image.shape}\")\n    \n    if scale:\n        new_size = (image.shape[1] // scale, image.shape[0] // scale)\n        image = cv2.resize(image, new_size)\n        if verbose:\n            print(f\"[{path}] Resized Image shape: {image.shape}\")\n        \n    mx = np.max(image)\n    image = image.astype(np.float32)\n    if mx:\n        image /= mx # scale image to [0, 1]\n    return image\n        \n#https://www.kaggle.com/bguberfain/memory-aware-rle-encoding\n#with transposed mask\n\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\ndef rle_encode(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\n# Ref: https://www.kaggle.com/code/paulorzp/rle-functions-run-lenght-encode-decode/script\ndef rle_decode(mask_rle, shape):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height,width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\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    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape).T  # Needed to align to RLE direction\n\n\n# def build_transforms(CFG):\n#     data_transforms = {\n#         \"valid_test\": A.Compose([\n#             A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n#             ], p=1.0)\n#         }\n#     return data_transforms\n\n\n# class build_dataset(Dataset):\n#     def __init__(self, df, label=True, transforms=None):\n#         self.df = df\n#         self.label = label\n#         self.transforms = transforms\n        \n#     def __len__(self):\n#         return len(self.df)\n    \n#     def __getitem__(self, index):\n        \n#         img_path   = self.df.loc[index, 'image_path']\n#         img_height = self.df.loc[index, 'img_height']\n#         img_width  = self.df.loc[index, 'img_width']\n#         organs     = self.df.loc[index, 'organ']\n#         id_        = self.df.loc[index, 'id']\n#         img        = read_tiff(img_path)\n#         sours      = self.df.loc[index,'data_source']\n        \n#         if self.label:\n#             rle_mask = self.df.loc[index, 'rle']\n#             mask = rle_decode(rle_mask, (img_height, img_width))\n#             # pdb.set_trace() \n#             if self.transforms:\n#                 data = self.transforms(image=img, mask=mask)\n#                 img  = data['image']\n#                 mask  = data['mask']\n            \n#             mask = np.expand_dims(mask, axis=0)\n#             img = np.transpose(img, (2, 0, 1))\n#             # mask = np.transpose(mask, (2, 0, 1))\n            \n#             return torch.tensor(img), torch.tensor(mask)\n        \n#         else:    # resize for infer\n#             if self.transforms:\n#                 data = self.transforms(image=img)\n#                 img  = data['image']\n                \n#             img = np.transpose(img, (2, 0, 1))   #(c, h, w)\n#             return torch.tensor(img), img_height, img_width,id_,organs,sours\n# #             return torch.tensor(img), img_height, img_width, id_,organ\n","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.580307Z","iopub.execute_input":"2022-09-15T17:40:33.581039Z","iopub.status.idle":"2022-09-15T17:40:33.596575Z","shell.execute_reply.started":"2022-09-15T17:40:33.581003Z","shell.execute_reply":"2022-09-15T17:40:33.595536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HuBMAPDataset(Dataset):\n    def __init__(self, df= None, cfg = None, transforms = None, mode='train'):\n        self.df = df\n        self.mode = mode\n        self.cfg = cfg\n        self.transforms = transforms\n#         print(len(df))\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        \n        try: \n#             print('Reading image')\n            img_path   = self.df.loc[index, 'image_path']            \n            img_height = self.df.loc[index, 'img_height']\n            img_width  = self.df.loc[index, 'img_width']\n            organ     = self.df.loc[index, 'organ']\n            id_        = self.df.loc[index, 'id']            \n            sours      = self.df.loc[index,'data_source'] \n            mask_path  = self.df.loc[index,'mask_path'] \n            image      = read_tiff(img_path)\n            image_size = cfg.img_size[0]\n            \n            if self.mode=='train':\n                #rle_mask = self.df.loc[index, 'rle']                \n                #mask = rle_decode(rle_mask, (img_height, img_width))\n                mask = cv2.imread(mask_path,cv2.IMREAD_GRAYSCALE)\n                mask  = (mask/255).astype(np.uint8)\n#                 print('Read mask',mask.shape)\n                # pdb.set_trace() \n                image = cv2.resize(image,dsize=(image_size,image_size),interpolation=cv2.INTER_LINEAR)\n                mask  = cv2.resize(mask, dsize=(image_size,image_size),interpolation=cv2.INTER_LINEAR)\n#                 print('After resize,', image.shape, mask.shape,np.unique(mask))\n                \n                if self.transforms:                   \n                    image, mask = train_augment5b(image, mask)\n                else:\n                    image, mask = valid_augment5(image, mask)\n               \n                # -------------- Sanity check------------------------\n#                 plt.imshow(image)\n#                 print('4 - Shapes: img,mask',image.shape,mask.shape,np.unique(mask))\n                # ---------------------------------------------------\n\n                mask = np.expand_dims(mask, axis=0)\n                image = np.transpose(image, (2, 0, 1))\n                image = image.astype(np.float32)\n                mask  = mask.astype(np.float32)\n                # mask = np.transpose(mask, (2, 0, 1))         \n\n                return torch.tensor(image), torch.tensor(mask), organ\n\n            else:    # resize for infer\n                if self.transforms:\n                    data = self.transforms(image=img)\n                    img  = data['image']\n\n                img = np.transpose(img, (2, 0, 1))   #(c, h, w)      \n                return torch.tensor(img), img_height, img_width,id_,organs,sours\n            \n        except:\n            print(index,\"error occured in dataloader/dataset\")\n            return None\n\n","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.598270Z","iopub.execute_input":"2022-09-15T17:40:33.598692Z","iopub.status.idle":"2022-09-15T17:40:33.615114Z","shell.execute_reply.started":"2022-09-15T17:40:33.598658Z","shell.execute_reply":"2022-09-15T17:40:33.613948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_loaders(fold,df,cfg, debug=False):\n    \n    train_df = df.query(\"fold!=@fold\").reset_index(drop=True)\n    valid_df = df.query(\"fold==@fold\").reset_index(drop=True)\n    \n    if debug:\n#         train_df = train_df.head(4*5)\n        train_df =  train_df.iloc[:100, :] # CHECKING TIME\n        valid_df = valid_df.head(4*5)\n    \n    \n    train_dataset = HuBMAPDataset(train_df, transforms=True,cfg=cfg,mode = 'train') \n    valid_dataset = HuBMAPDataset(valid_df, transforms=False,cfg=cfg, mode = 'train') \n    \n    train_loader = DataLoader(train_dataset, batch_size=cfg.train_bs if not cfg.debug else 20, \n                            num_workers=2, shuffle=True, pin_memory=True, drop_last=False)\n    valid_loader = DataLoader(valid_dataset, batch_size=cfg.valid_bs if not cfg.debug else 20, \n                            num_workers=2, shuffle=False, pin_memory=True)\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.616968Z","iopub.execute_input":"2022-09-15T17:40:33.617980Z","iopub.status.idle":"2022-09-15T17:40:33.629405Z","shell.execute_reply.started":"2022-09-15T17:40:33.617942Z","shell.execute_reply":"2022-09-15T17:40:33.628420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Loss**","metadata":{}},{"cell_type":"code","source":"'''\nEp 1:\nDice value:  tensor(0.6109, device='cuda:0')\nBCE value:  tensor(0.7966, device='cuda:0')\n'''\n# import segmentation_models_pytorch as smp\n# class CustomLoss(nn.Module):\n#     def __init__(self):\n#         super(CustomLoss,self).__init__()\n#         self.diceloss = smp.losses.DiceLoss(mode='binary')\n#         self.binloss = smp.losses.SoftBCEWithLogitsLoss(reduction = 'mean' , smooth_factor = 0.1)\n\n#     def forward(self, output, mask):\n#         dice = self.diceloss(output,mask)\n#         bce = self.binloss(output, mask)\n#         loss = bce * 0.3 + dice * 0.7\n#         return loss\n    \nclass DiceBCELoss(nn.Module):\n\n    def __init__(self, weight=None, size_average=True):\n        super(DiceBCELoss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        \n        #comment out if your model contains a sigmoid or equivalent activation layer\n#         inputs = nnF.sigmoid(inputs)       \n        \n        #flatten label and prediction tensors\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)        \n        \n        BCE = nnF.binary_cross_entropy_with_logits(inputs, targets, reduction='mean')\n#         return BCE\n    \n        intersection = (inputs * targets).sum()                            \n        dice_loss = 1 - (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n#         print('Dice value: ',dice_loss)\n#         print('BCE value: ',BCE)\n        Dice_BCE = BCE + dice_loss\n        \n        return Dice_BCE\n","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:40:33.631039Z","iopub.execute_input":"2022-09-15T17:40:33.631917Z","iopub.status.idle":"2022-09-15T17:40:33.643195Z","shell.execute_reply.started":"2022-09-15T17:40:33.631882Z","shell.execute_reply":"2022-09-15T17:40:33.642160Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Metrics**","metadata":{}},{"cell_type":"code","source":"def dice_coef(y_true, y_pred, thr=0.5, dim=(2,3), epsilon=0.001):\n    y_true = y_true.to(torch.float32)\n    y_pred = (y_pred>thr).to(torch.float32)\n    inter = (y_true*y_pred).sum(dim=dim)\n    den = y_true.sum(dim=dim) + y_pred.sum(dim=dim)\n    dice = ((2*inter+epsilon)/(den+epsilon)).mean(dim=(1,0))\n    return dice\n\n# --lb metric ------------------------------------------\n# https://www.kaggle.com/competitions/hubmap-organ-segmentation/overview/supervised-ml-evaluation\n\ndef compute_dice_score(mask, probability):\n\tN = len(probability)\n\tp = probability.reshape(N,-1)\n\tt = mask.reshape(N,-1)\n\t\n\tp = p>0.5\n\tt = t>0.5\n\tuion = p.sum(-1) + t.sum(-1)\n\toverlap = (p*t).sum(-1)\n\tdice = 2*overlap/(uion+0.0001)\n\treturn dice\n\n\ndef iou_coef(y_true, y_pred, thr=0.5, dim=(2,3), epsilon=0.001):\n    y_true = y_true.to(torch.float32)\n    y_pred = (y_pred>thr).to(torch.float32)\n    inter = (y_true*y_pred).sum(dim=dim)\n    union = (y_true + y_pred - y_true*y_pred).sum(dim=dim)\n    iou = ((inter+epsilon)/(union+epsilon)).mean(dim=(1,0))\n    return iou","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:47:20.317291Z","iopub.execute_input":"2022-09-15T17:47:20.317667Z","iopub.status.idle":"2022-09-15T17:47:20.329519Z","shell.execute_reply.started":"2022-09-15T17:47:20.317635Z","shell.execute_reply":"2022-09-15T17:47:20.328383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Model**","metadata":{}},{"cell_type":"code","source":"def build_model():\n\n    model = Net()\n    model.output_type = [\"loss\"]\n#     model.float()\n    return model\n\ndef load_model(path):\n    model = Net()\n    model.output_type = [\"inference\"]\n    model.load_state_dict(torch.load(path))\n    model.eval()\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:47:21.029735Z","iopub.execute_input":"2022-09-15T17:47:21.030375Z","iopub.status.idle":"2022-09-15T17:47:21.035960Z","shell.execute_reply.started":"2022-09-15T17:47:21.030339Z","shell.execute_reply":"2022-09-15T17:47:21.034800Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Train epoch**","metadata":{}},{"cell_type":"code","source":"'''\nFLOW OF TRAINING:\n\n1. Load batch of data (add to gpu if accessible)\n2. forward pass\n   |_ output = model (input)\n   |_ loss   = criterion (output,target)\n   \n3. backward pass + optimize (only in training)\n   |_ loss.backward()\n   |_ optimizer.step()\n   \n4. Calculate statistics\n   |_ running_loss += loss.item()*batch_size\n\nreturn epoch_loss = running_loss/len(dataloader)\n'''\n\ndef train_one_epoch(cfg,model, optimizer, scheduler, dataloader, device, epoch):\n    model.train()\n    scaler = amp.GradScaler()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    criterion = DiceBCELoss() #CustomLoss() # DiceBCELoss() \n    \n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    for step, (images, masks, organs) in pbar:  #Shapes: img: (8,3,512,512)\n#         loss     = 0\n#         aux_loss = 0\n#         bce_loss = 0\n        \n        # get a batch of inputs\n        images = images.to(device, dtype=torch.float)\n        masks  = masks.to(device, dtype=torch.float)\n        #masks = masks.unsqueeze(1)\n        batch_size = images.size(0)\n        \n        send_batch = {\n            'image': images,\n            'mask': masks,\n            'organ': organs,\n        }\n        \n        # ------------- forward ---------------\n        with amp.autocast(enabled=True):\n            output   = model(send_batch)   \n            y_pred = output['probability']\n#             print('Prob_logits: shape,min,max',y_pred.shape,torch.min(y_pred),torch.max(y_pred))\n            loss = criterion(masks,y_pred)\n        \n#             output   = model(send_batch)     \n#             bce_loss = output['bce_loss']\n#             y_pred   = output['probability']\n#             loss = bce_loss\n#             for i in range(4):\n#                 aux_loss += output['aux%d_loss'%i] \n                \n#             aux_loss = aux_loss/4.0 \n# #             print(f'Loss is: BCE_Loss = {bce_loss}, Aux_Loss = {aux_loss}')\n#             loss = 0.75*bce_loss + 0.25*aux_loss\n\n        # -------- backward + optimize --------\n#         print('Model wt. gradient = ',model.weight.grad.abs().sum())\n#         print('Model wt.          = ',model.layername.weight.shape)\n#         print(f'Loss is: BCE_Loss = {bce_loss}, Aux_Loss = {aux_loss}')\n        \n        scaler.scale(loss).backward() #/cfg.n_accumulate\n    \n        if ((step+1)%cfg.n_accumulate==0 or (step+1)==len(dataloader)):\n            scaler.step(optimizer)\n            scaler.update()\n\n            # zero the parameter gradients\n            optimizer.zero_grad()\n            \n            if cfg.scheduler == 'CosineAnnealingLR' and cfg.scheduler: # 'CosineAnnealingLR' or 'StepLR' location\n                scheduler.step()\n       \n        # statistics\n        running_loss += (loss.item() * batch_size) #pytorch loss.item gives average loss of the batch\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss / dataset_size\n        \n        mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n        current_lr = optimizer.param_groups[0]['lr']\n        pbar.set_postfix(train_loss=f'{epoch_loss:0.4f}', lr=f'{current_lr:0.6f}',gpu_mem=f'{mem:0.2f} GB')\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return epoch_loss","metadata":{"execution":{"iopub.status.busy":"2022-09-15T18:10:44.961506Z","iopub.execute_input":"2022-09-15T18:10:44.961882Z","iopub.status.idle":"2022-09-15T18:10:44.975295Z","shell.execute_reply.started":"2022-09-15T18:10:44.961846Z","shell.execute_reply":"2022-09-15T18:10:44.974095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef valid_one_epoch(model, dataloader, device, epoch):\n    model.eval()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    criterion = DiceBCELoss()\n    \n    val_scores = []\n    \n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid ')\n    for step, (images, masks, organs) in pbar:       \n#         loss     = 0\n#         aux_loss = 0\n#         bce_loss = 0\n        \n        images  = images.to(device, dtype=torch.float)\n        masks   = masks.to(device, dtype=torch.float)\n#         masks   = masks.unsqueeze(1)\n#         print(\"Masks: \",masks.shape)\n        batch_size = images.size(0)\n        \n        send_batch = {\n            'image': images,\n            'mask': masks,\n            'organ': organs,\n        }\n        output   = model(send_batch)   \n        y_pred = output['probability']      \n        \n        loss = criterion(masks,y_pred)\n#         output = model(send_batch) \n#         y_pred = output['probability']\n#         bce_loss = output['bce_loss']\n#         loss = bce_loss\n\n#         for i in range(4):\n#             aux_loss += output['aux%d_loss'%i]         \n#         aux_loss = aux_loss/4.0\n# #         print(f'Loss is: BCE_Loss = {bce_loss}, Aux_Loss = {aux_loss}')\n#         loss = 0.75*bce_loss + 0.25*aux_loss\n        \n#         outputs = model(images)\n#         logits  = outputs.logits \n#         logits = nn.Sigmoid()(logits)\n#         logits = nn.functional.interpolate(logits, size=masks.shape[-2:], mode=\"nearest\")\n# #         logits = upsample_obj(logits)\n# #         print(\"Logits: \",logits.shape)\n# #         logits  = nn.functional.interpolate(logits, size=masks.shape[-2:], mode=\"bilinear\", align_corners=False) #upsample to masks size\n#         y_pred  = logits\n#         loss    = criterion(y_pred, masks)\n        \n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss / dataset_size\n        \n        \n#         print(\"Ypred: \",y_pred.shape)\n        val_dice = dice_coef(masks, y_pred).cpu().detach().numpy()\n#         val_dice    = compute_dice_score(masks, y_pred).cpu().detach().numpy()\n#         print('Val Dice: ',val_dice)\n        val_jaccard = iou_coef(masks, y_pred).cpu().detach().numpy()\n        val_scores.append([val_dice, val_jaccard])\n        \n        mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n        current_lr = optimizer.param_groups[0]['lr']\n        pbar.set_postfix(valid_loss=f'{epoch_loss:0.4f}',lr=f'{current_lr:0.6f}',gpu_memory=f'{mem:0.2f} GB')\n    \n    val_scores  = np.mean(val_scores, axis=0)\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return epoch_loss, val_scores","metadata":{"execution":{"iopub.status.busy":"2022-09-15T18:10:45.348582Z","iopub.execute_input":"2022-09-15T18:10:45.348928Z","iopub.status.idle":"2022-09-15T18:10:45.360180Z","shell.execute_reply.started":"2022-09-15T18:10:45.348896Z","shell.execute_reply":"2022-09-15T18:10:45.358856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Run Training**","metadata":{}},{"cell_type":"code","source":"import copy\ndef run_training(cfg,fold, model, optimizer, scheduler, device, num_epochs):\n    # To automatically log gradients\n    \n    if torch.cuda.is_available():\n        print(\"cuda: {}\\n\".format(torch.cuda.get_device_name()))\n    \n    start = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_dice      = -np.inf\n    best_epoch     = -1\n    history = defaultdict(list)\n  \n    for epoch in range(1, num_epochs + 1): \n        gc.collect()\n        print(f'Epoch {epoch}/{num_epochs}', end='')\n        train_loss = train_one_epoch(cfg,model, optimizer, scheduler,dataloader=train_loader,device=cfg.device, epoch=epoch)        \n        val_loss, val_scores = valid_one_epoch(model, valid_loader,device=cfg.device, epoch=epoch)        \n        val_dice, val_jaccard = val_scores\n        \n        #ReduceLR scheduler\n        if cfg.scheduler =='ReduceLROnPlateau': \n#             print('not none scheduler...')\n            scheduler.step(val_dice) #Monitor metric\n            \n        \n        history['Fold'].append(fold)\n        history['Epoch'].append(epoch)\n        history['Train Loss'].append(train_loss)\n        history['Valid Loss'].append(val_loss)\n        history['Valid Dice'].append(val_dice)\n        history['Valid Jaccard'].append(val_jaccard)\n        \n        pd.DataFrame(history).to_csv('./CoAT_fold_0_768_lr-1e-4_12ep.csv', index=False)\n        \n        print(f'Valid Dice: {val_dice:0.4f} | Valid Jaccard: {val_jaccard:0.4f}')\n        print(f'{lo_}Valid Loss: {val_loss:0.4f} | Train Loss: {train_loss:0.4f}{sr_}')\n        # deep copy the model\n        if val_dice >= best_dice:\n            \n            print(f\"{c_}Valid Dice Score Improved ({best_dice:0.4f} ---> {val_dice:0.4f}){sr_}\")\n            best_dice    = val_dice\n            best_jaccard = val_jaccard\n            best_epoch   = epoch\n            \n            \n            best_model_wts = copy.deepcopy(model.state_dict())\n            PATH = f\"CoAT_fold_0_768_lr-1e-4_12ep-best_epoch-{best_epoch}.bin\"            \n            \n#         if epoch > 7:\n            torch.save(model.state_dict(), PATH)\n            # Save a model file from the current directory\n            print(f\"Model Saved{sr_}\")\n            \n        last_model_wts = copy.deepcopy(model.state_dict())\n        PATH = f\"CoAT_fold_0_768_lr-1e-4_12ep-last_epoch.bin\"\n        torch.save(model.state_dict(), PATH)\n            \n        print(); print(f'{sr_}')\n        \n        # ---------Add patience factor ------------\n        \n        # -----------------------------------------\n    \n    end = time.time()\n    time_elapsed = end - start\n    print('Training complete in {:.0f}h {:.0f}m {:.0f}s'.format(\n        time_elapsed // 3600, (time_elapsed % 3600) // 60, (time_elapsed % 3600) % 60))\n    print(\"Best Score: {:.4f}\".format(best_dice))\n    \n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    \n    return model, history","metadata":{"execution":{"iopub.status.busy":"2022-09-15T18:10:46.472999Z","iopub.execute_input":"2022-09-15T18:10:46.473389Z","iopub.status.idle":"2022-09-15T18:10:46.486666Z","shell.execute_reply.started":"2022-09-15T18:10:46.473352Z","shell.execute_reply":"2022-09-15T18:10:46.485515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Check Data**","metadata":{}},{"cell_type":"code","source":"fold = 0\ncfg  = initialize_config(debug=False,batch_size=2)\ndf   = create_folds(cfg=cfg)\n# # pd.set_option('display.max_colwidth', None)\n# display(df.head(10))\nprint(df.columns)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T18:10:46.807292Z","iopub.execute_input":"2022-09-15T18:10:46.807931Z","iopub.status.idle":"2022-09-15T18:10:46.974104Z","shell.execute_reply.started":"2022-09-15T18:10:46.807897Z","shell.execute_reply":"2022-09-15T18:10:46.973006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check Tensor shapes\ni = 0\nfor i in range(5):\n    train_loader, valid_loader = prepare_loaders(fold=fold,df = df,cfg=cfg)\n    batch = next(iter(train_loader))\n    images, labels,organ = batch\n#     print(images.shape, labels.shape, type(images), type(labels), images.dtype, labels.dtype)\n\n    # Sanity check sample n=1 \n    testImg = images[0]\n    testMsk = labels[0].squeeze()\n    print(\"Img: \",testImg.shape, testImg.dtype, type(testImg))\n    print(\"Mask: \",testMsk.shape,testMsk.dtype, type(testMsk))\n\n    # Plot exmaple mask \n    x=testMsk #.permute(1,2,0)\n    x=x[:,:]\n\n\n    plt.figure(figsize=(8,8))\n    plt.subplot(1, 2, 1)\n    plt.imshow((testImg.permute(1,2,0)), interpolation='none')\n    plt.subplot(1, 2, 2)\n    plt.imshow(testImg.permute(1,2,0), 'gray', interpolation='none')\n    plt.imshow(x, 'jet', interpolation='none', alpha=0.7)\n    plt.show()\n    del testImg,x,testMsk, batch, images, labels, organ","metadata":{"execution":{"iopub.status.busy":"2022-09-15T18:10:47.210264Z","iopub.execute_input":"2022-09-15T18:10:47.210852Z","iopub.status.idle":"2022-09-15T18:10:47.215701Z","shell.execute_reply.started":"2022-09-15T18:10:47.210820Z","shell.execute_reply":"2022-09-15T18:10:47.214615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **MAIN**","metadata":{}},{"cell_type":"code","source":"# Initlaisations\n\n# mean = np.array([0.485,0.456,0.406])\n# std = np.array([0.229,0.224,0.225])\nupsample_obj = MixUpSample()\n\nfor fold in range(1):\n    print(f'#'*15*2)\n    print(f'### Fold: {fold}')\n    print(f'#'*15*2)\n    \n    cfg                        = initialize_config(debug=False,batch_size=4)\n    df                         = create_folds(cfg=cfg)\n    train_loader, valid_loader = prepare_loaders(fold=fold,df = df,cfg=cfg, debug = False)\n    model                      = build_model().to(cfg.device) \n    optimizer                  = get_optimizer(cfg,optimizer_name=cfg.optimizer)\n    scheduler                  = get_scheduler(cfg,optimizer,df)\n    model, history             = run_training(cfg,fold, model, optimizer, scheduler,device=cfg.device,num_epochs=cfg.epochs)\n    \n    \n    #Plot 1\n    #plt.subplot(3, 1, 1)\n    plt.plot(history['Epoch'], history['Train Loss'], 'r--')\n    plt.plot(history['Epoch'], history['Valid Loss'], 'b-')\n    plt.legend(['Training Loss', 'Valid Loss'])\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.savefig('./Epoch_vs_Loss.png')\n    plt.show()\n\n    #Plot 2\n    #plt.subplot(3,1, 2)\n    plt.plot(history['Epoch'], history['Valid Dice'])\n    plt.xlabel('Epoch')\n    plt.ylabel('Valid Dice')\n    plt.savefig('./Epoch_vs_Valid-Dice.png')\n    plt.show()\n\n    #Plot 3\n    #plt.subplot(3,1, 3)\n    plt.plot(history['Valid Dice'], history['Train Loss'] )\n    plt.xlabel('Valid Dice') \n    plt.ylabel('Train Loss')\n    plt.savefig('./Valid-Dice_vs_Loss.png')\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2022-09-15T18:10:48.136189Z","iopub.execute_input":"2022-09-15T18:10:48.136569Z","iopub.status.idle":"2022-09-15T18:11:11.025516Z","shell.execute_reply.started":"2022-09-15T18:10:48.136538Z","shell.execute_reply":"2022-09-15T18:11:11.023899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# del model\n# torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:47:11.436887Z","iopub.execute_input":"2022-09-15T17:47:11.437362Z","iopub.status.idle":"2022-09-15T17:47:11.626601Z","shell.execute_reply.started":"2022-09-15T17:47:11.437320Z","shell.execute_reply":"2022-09-15T17:47:11.624771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Debug Data:\n'''\nProb_logits: shape,min,max torch.Size([4, 1, 768, 768]) tensor(0.1708, device='cuda:0', dtype=torch.float16, grad_fn=<MinBackward1>) tensor(0.7876, device='cuda:0', dtype=torch.float16, grad_fn=<MaxBackward1>)\nProb_logits: shape,min,max torch.Size([4, 1, 768, 768]) tensor(1.0729e-06, device='cuda:0', dtype=torch.float16,\n       grad_fn=<MinBackward1>) tensor(0.9644, device='cuda:0', dtype=torch.float16, grad_fn=<MaxBackward1>)\nVal Dice:  1.0639091e-07\nValid :  67%|██████▋   | 2/3 [00:10<00:04,  4.69s/it, gpu_memory=13.23 GB, lr=0.000100, valid_loss=1.5616]\nVal Dice:  0.0012314679\nValid : 100%|██████████| 3/3 [00:13<00:00,  4.40s/it, gpu_memory=13.23 GB, lr=0.000100, valid_loss=1.5571]\nVal Dice:  0.007386328\n'''\n","metadata":{},"execution_count":null,"outputs":[]}]}