{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n\n<h3 style=\"text-align:center; background-color:#C8FF33;padding:40px;border-radius: 30px;\">\n<p style=\"text-align: left\">See also this notebook:</p>\n    <p style=\"text-align: left\"><b><a href=\"https://www.kaggle.com/code/soumya9977/hubmap-multiorgan-segmentation-1-3-data-prep\"> &nbsp; HuBMAP multiOrgan Segmentation 1/3 [data prep]</a></b></p>\n    <p style=\"text-align: left\"><b>* HuBMAP: Vanilla Unet + W&B + Pytorch 2/3 [train]</b></p>\n    <p style=\"text-align: left\"><b><a href=\"https://www.kaggle.com/code/soumya9977/hubmap-vanilla-unet-pytorch-3-3-inference\"> &nbsp; HuBMAP: Vanilla Unet Pytorch 3/3 [inference]</a></b></p>\n</h3>\n\n\n## Please _DO_ upvote!\n","metadata":{}},{"cell_type":"markdown","source":"### Changelog\n\n|| Version | Comments | LB |\n|---|  --- | --- | --- |\n||8| Vanilla Unet w/ 256x256, 40 epoch, no aug, no lr sch, no pre/post processing, L2 norm | `0.04` |\n|**Best**| --| Tiled data, Vanilla Unet w/ 256x256, 40 epoch, no aug, no lr sch, no pre/post processing, L2 norm | `0.37` |","metadata":{}},{"cell_type":"code","source":"!pip install -q tifffile","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:48:42.436055Z","iopub.execute_input":"2022-07-24T08:48:42.436420Z","iopub.status.idle":"2022-07-24T08:48:54.112160Z","shell.execute_reply.started":"2022-07-24T08:48:42.436389Z","shell.execute_reply":"2022-07-24T08:48:54.111097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import, SEED, Config:","metadata":{}},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\nimport warnings\nimport cv2\nimport matplotlib.pyplot as plt\nimport json\nimport gc\nimport time\nfrom tqdm import tqdm\nimport random\nfrom collections import defaultdict\nfrom IPython.display import display\n\n#Pytorch Imports\nimport torch\nimport torchvision # torch package for vision related things\nimport torch.nn.functional as F  # Parameterless functions, like (some) activation functions\nimport torchvision.datasets as datasets  # Standard datasets\nimport torchvision.transforms as transforms  # Transformations we can perform on our dataset for augmentation\nfrom torch import optim  # For optimizers like SGD, Adam, etc.\nfrom torch import nn  # All neural network modules\nfrom torch.utils.data import Dataset, DataLoader  # Gives easier dataset managment by creating mini batches etc.\nfrom tqdm import tqdm  # For nice progress bar!\nfrom torchvision.transforms import Resize\nfrom torchaudio.transforms import MelSpectrogram, AmplitudeToDB\nfrom torch.optim import lr_scheduler\nfrom tifffile import imread\n\nfrom albumentations.pytorch.transforms import ToTensorV2\nimport albumentations as A\nfrom sklearn.model_selection import train_test_split\n\n# For colored terminal text\nfrom colorama import Fore, Back, Style\nc_ = Fore.CYAN\nsr_ = Style.RESET_ALL\nb_ = Fore.BLUE\n\n# For descriptive error messages\nos.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"\n\nwarnings.filterwarnings(\"ignore\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-24T08:48:54.117981Z","iopub.execute_input":"2022-07-24T08:48:54.120307Z","iopub.status.idle":"2022-07-24T08:48:54.463805Z","shell.execute_reply.started":"2022-07-24T08:48:54.120267Z","shell.execute_reply":"2022-07-24T08:48:54.462887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Hyperparameters etc.\n# Hyperparameters\nCONFIG = {\n    \"in_channels\" :3,\n    \"num_classes\": 1,\n    \"BATCH_SIZE\" : 8,\n    \"NUM_EPOCHS\" : 40,\n    \"n_accumulate\": 1,\n    \"competition\": \"HuBMAP-Kaggle\", # HuBMAP-Kaggle\n    \"model_name\": \"Vanilla_Unet\",\n    \"LEARNING_RATE\": 1e-4,\n    \"DEVICE\": \"cuda\" if torch.cuda.is_available() else \"cpu\", \n    \"AUG\": \"No\",\n    \"SEED\": 42,\n    \"opt\": 'Adam',\n    \"Normalization\": \"L2\",\n    \"img_size\": 256\n}","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:48:54.465140Z","iopub.execute_input":"2022-07-24T08:48:54.465544Z","iopub.status.idle":"2022-07-24T08:48:54.726082Z","shell.execute_reply.started":"2022-07-24T08:48:54.465506Z","shell.execute_reply":"2022-07-24T08:48:54.722883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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    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    \nset_seed(CONFIG[\"SEED\"])","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:48:54.728476Z","iopub.execute_input":"2022-07-24T08:48:54.731782Z","iopub.status.idle":"2022-07-24T08:48:55.044312Z","shell.execute_reply.started":"2022-07-24T08:48:54.731712Z","shell.execute_reply":"2022-07-24T08:48:55.043387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# WandB:","metadata":{}},{"cell_type":"code","source":"import wandb\n\ntry:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    api_key = user_secrets.get_secret(\"WANDB\")\n    wandb.login(key=api_key)\n    anonymous = None\nexcept:\n    anonymous = \"must\"\n    print('To use your W&B account,\\nGo to Add-ons -> Secrets and provide your W&B access token. Use the Label name as WANDB. \\nGet your W&B access token from here: https://wandb.ai/authorize')\n# wandb.init(project=\"PogChamp2 Baseline\")\nrun = wandb.init(project=CONFIG['competition'], \n                 config=CONFIG,\n                 job_type='Train',\n                 tags=['semantic segmentation', CONFIG['model_name']],\n                 anonymous='must',\n                 name = \"Vanilla_Unet_1\",\n                 notes = \"\")","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:48:55.047729Z","iopub.execute_input":"2022-07-24T08:48:55.048014Z","iopub.status.idle":"2022-07-24T08:49:03.562992Z","shell.execute_reply.started":"2022-07-24T08:48:55.047988Z","shell.execute_reply":"2022-07-24T08:49:03.561921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Tabular:","metadata":{}},{"cell_type":"code","source":"DIR = \"../input/hubmap-organ-segmentation/\"\ntrain_df = pd.read_csv(os.path.join(DIR,\"train.csv\"))\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:03.567548Z","iopub.execute_input":"2022-07-24T08:49:03.569753Z","iopub.status.idle":"2022-07-24T08:49:03.944356Z","shell.execute_reply.started":"2022-07-24T08:49:03.569710Z","shell.execute_reply":"2022-07-24T08:49:03.943449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_DIR = DIR + \"train_images/\"\nfunc = lambda x: TRAIN_DIR + str(x) + \".tiff\"\ntrain_df[\"img_path\"] = train_df[\"id\"].apply(func)\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:03.948814Z","iopub.execute_input":"2022-07-24T08:49:03.951167Z","iopub.status.idle":"2022-07-24T08:49:03.978729Z","shell.execute_reply.started":"2022-07-24T08:49:03.951130Z","shell.execute_reply":"2022-07-24T08:49:03.977684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MASK_DIR = \"../input/hubmap-hpa-multiorgan-segmentation-1-3-data-prep/train_masks_np/\"\nfunc1 = lambda x: MASK_DIR + str(x) + \".npy\"\ntrain_df[\"mask_path\"] = train_df[\"id\"].apply(func1)\nprint(train_df[\"mask_path\"].iloc[9])\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:03.983291Z","iopub.execute_input":"2022-07-24T08:49:03.985632Z","iopub.status.idle":"2022-07-24T08:49:04.020835Z","shell.execute_reply.started":"2022-07-24T08:49:03.985597Z","shell.execute_reply":"2022-07-24T08:49:04.019701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, test_df = train_test_split(train_df, test_size=0.2, random_state=42)\n\ntrain_df.__len__(), len(test_df)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:04.025038Z","iopub.execute_input":"2022-07-24T08:49:04.027572Z","iopub.status.idle":"2022-07-24T08:49:04.042324Z","shell.execute_reply.started":"2022-07-24T08:49:04.027533Z","shell.execute_reply":"2022-07-24T08:49:04.040824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%timeit img_temp = imread(train_df[\"img_path\"].iloc[0])","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:04.049943Z","iopub.execute_input":"2022-07-24T08:49:04.052195Z","iopub.status.idle":"2022-07-24T08:49:10.443393Z","shell.execute_reply.started":"2022-07-24T08:49:04.052158Z","shell.execute_reply":"2022-07-24T08:49:10.442235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pytorch Data Class:","metadata":{}},{"cell_type":"code","source":"mean = np.array([0.7720342, 0.74582646, 0.76392896])\nstd = np.array([0.24745085, 0.26182273, 0.25782376])\n\n\ndef img2tensor(img,dtype:np.dtype=np.float32):\n    if img.ndim==2 : img = np.expand_dims(img,2)\n    img = np.transpose(img,(2,0,1))\n    return torch.from_numpy(img.astype(dtype, copy=False))\n\n\n\nclass hubmap_data(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.img_paths = df[\"img_path\"].to_numpy()\n        self.mask_paths = df[\"mask_path\"].to_numpy()\n        self.transform = transform\n        \n        \n    def __getitem__(self, index):\n        # load data from the pre-processed npy files\n        img_file = self.img_paths[index]\n#         print(img_file)\n        mask_file = self.mask_paths[index]\n        img = imread(img_file)#cv2.cvtColor(cv2.imread(img_file), cv2.COLOR_BGR2RGB)\n#         img = imread(img_file) #np.moveaxis(imread(img_file), -1, 0) #torch.from_numpy(np.moveaxis(imread(img_file), -1, 0))\n        # semantic = torch.from_numpy(np.load(self.data_path + '/label/{:d}.npy'.format(index)))\n#         mask = torch.from_numpy(np.moveaxis(np.load(mask_file), -1, 0))\n        mask = np.load(mask_file)#np.moveaxis(np.load(mask_file), -1, 0) #torch.from_numpy(np.moveaxis(np.load(mask_file), -1, 0))\n        \n        \n#       print(depth.shape)\n        if self.transform is not None:\n            transformed = self.transform(image=img, mask=mask)\n            img = transformed[\"image\"]\n            mask = transformed[\"mask\"]\n\n            if mask.shape[-1] == 1:\n                mask = mask.permute(2,0,1)\n                \n                \n        return img2tensor((img/255.0 - mean)/std),img2tensor(mask)\n\n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:10.444736Z","iopub.execute_input":"2022-07-24T08:49:10.445422Z","iopub.status.idle":"2022-07-24T08:49:10.457542Z","shell.execute_reply.started":"2022-07-24T08:49:10.445382Z","shell.execute_reply":"2022-07-24T08:49:10.456403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = A.Compose(\n    [\n        A.Resize(256, 256, p=1),\n#         A.HorizontalFlip(p=1.0),\n#         A.RandomCrop(224, 224, p=0.3),\n#         A.RandomSizedCrop(cfg[\"min_max_height\"], cfg[\"height\"], cfg[\"width\"], cfg[\"w2h_ratio\"], cfg[\"interpolation\"], p=1), \n#         ToTensorV2(),\n    ]\n)\n\ntest_transform = A.Compose(\n    [\n        A.Resize(256, 256, p=1.0),\n#         ToTensorV2(),\n    ]\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:10.458994Z","iopub.execute_input":"2022-07-24T08:49:10.459810Z","iopub.status.idle":"2022-07-24T08:49:10.468543Z","shell.execute_reply.started":"2022-07-24T08:49:10.459711Z","shell.execute_reply":"2022-07-24T08:49:10.467269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = hubmap_data(train_df,train_transform)\ntest_data = hubmap_data(test_df, test_transform)\n\nprint(train_data[10][0].shape, train_data[10][1].shape)\n# test_data[10][0].shape, test_data[10][1].shape","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:10.470183Z","iopub.execute_input":"2022-07-24T08:49:10.470593Z","iopub.status.idle":"2022-07-24T08:49:10.742662Z","shell.execute_reply.started":"2022-07-24T08:49:10.470538Z","shell.execute_reply":"2022-07-24T08:49:10.741472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp_img = train_data[10][0]\ntemp_mask = train_data[10][1]\n# plt.imshow(temp_mask.detach().numpy())\n# plt.imshow(temp_img.permute(1,2,0).detach().numpy())\nimg_trfm = ((temp_img.permute(1,2,0)*std + mean)*255.0).numpy().astype(np.uint8)\n# plt.imshow(img_trfm)\n# temp_img.permute(1,2,0).detach().numpy()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:10.744312Z","iopub.execute_input":"2022-07-24T08:49:10.744721Z","iopub.status.idle":"2022-07-24T08:49:10.927811Z","shell.execute_reply.started":"2022-07-24T08:49:10.744683Z","shell.execute_reply":"2022-07-24T08:49:10.926797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Unet Model:","metadata":{}},{"cell_type":"code","source":"class double_conv(nn.Module):\n    \"\"\"(conv => BN => ReLU) * 2\"\"\"\n\n    def __init__(self, in_ch, out_ch):\n        super(double_conv, self).__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n        )\n\n    def forward(self, x):\n        x = self.conv(x)\n        return x\n\n\nclass inconv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(inconv, self).__init__()\n        self.conv = double_conv(in_ch, out_ch)\n\n    def forward(self, x):\n        x = self.conv(x)\n        return x\n\n\nclass down(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(down, self).__init__()\n        self.mpconv = nn.Sequential(nn.MaxPool2d(2), double_conv(in_ch, out_ch))\n\n    def forward(self, x):\n        x = self.mpconv(x)\n        return x\n\n\nclass up(nn.Module):\n    def __init__(self, in_ch, out_ch, bilinear=True):\n        super(up, self).__init__()\n\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode=\"bilinear\", align_corners=True)\n        else:\n            self.up = nn.ConvTranspose2d(in_ch // 2, in_ch // 2, 2, stride=2)\n\n        self.conv = double_conv(in_ch, out_ch)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n\n        # input is CHW\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n\n        x1 = F.pad(x1, (diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2))\n        \n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\n\nclass outconv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(outconv, self).__init__()\n        self.conv = nn.Conv2d(in_ch, out_ch, 1)\n\n    def forward(self, x):\n        x = self.conv(x)\n        return x\n\n\nclass UNet(nn.Module):\n    def __init__(self, n_channels, n_classes):\n        super(UNet, self).__init__()\n        self.inc = inconv(n_channels, 64)\n        self.down1 = down(64, 128)\n        self.down2 = down(128, 256)\n        self.down3 = down(256, 512)\n        self.down4 = down(512, 512)\n        self.up1 = up(1024, 256, False)\n        self.up2 = up(512, 128, False)\n        self.up3 = up(256, 64, False)\n        self.up4 = up(128, 64, False)\n        self.outc = outconv(64, n_classes)\n#         self.linear = nn.Linear(256, 128*256)\n\n    def forward(self, x):\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n        x = self.up1(x5, x4)\n        x = self.up2(x, x3)\n        x = self.up3(x, x2)\n        x = self.up4(x, x1)\n        x = self.outc(x)\n#         x = self.linear(x)\n        \n        return torch.sigmoid(x)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:10.929422Z","iopub.execute_input":"2022-07-24T08:49:10.929812Z","iopub.status.idle":"2022-07-24T08:49:10.952981Z","shell.execute_reply.started":"2022-07-24T08:49:10.929774Z","shell.execute_reply":"2022-07-24T08:49:10.951873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model check:","metadata":{}},{"cell_type":"code","source":"train_on_gpu = torch.cuda.is_available()\nmodel = UNet(n_channels=CONFIG[\"in_channels\"], n_classes=CONFIG[\"num_classes\"])\nif train_on_gpu:\n    model.cuda()\n\nmodel(temp_img.unsqueeze(0).cuda()).shape","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:10.954540Z","iopub.execute_input":"2022-07-24T08:49:10.955467Z","iopub.status.idle":"2022-07-24T08:49:11.141398Z","shell.execute_reply.started":"2022-07-24T08:49:10.955419Z","shell.execute_reply":"2022-07-24T08:49:11.140293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DataLoader","metadata":{}},{"cell_type":"code","source":"# REMARK shuffle = False [result might change with shuffle =True]\ntrain_loader = torch.utils.data.DataLoader(\n               dataset=train_data,\n               batch_size=CONFIG[\"BATCH_SIZE\"],\n               shuffle=False,\n               num_workers=2)\n\ntest_loader = torch.utils.data.DataLoader(\n              dataset=test_data,\n              batch_size=CONFIG[\"BATCH_SIZE\"],\n              shuffle=False,\n              num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:11.143305Z","iopub.execute_input":"2022-07-24T08:49:11.143750Z","iopub.status.idle":"2022-07-24T08:49:11.150018Z","shell.execute_reply.started":"2022-07-24T08:49:11.143713Z","shell.execute_reply":"2022-07-24T08:49:11.149016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss, Optimizer & Metric:","metadata":{}},{"cell_type":"code","source":"#PyTorch\nclass DiceLoss(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(DiceLoss, 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 = F.sigmoid(inputs)       \n        \n        #flatten label and prediction tensors\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()                            \n        dice = (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n        \n        return 1 - dice\n\nloss_fn = DiceLoss() #Dice()\noptimizer = optim.Adam(model.parameters(), lr=CONFIG[\"LEARNING_RATE\"],weight_decay=1e-5)","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:11.151583Z","iopub.execute_input":"2022-07-24T08:49:11.152220Z","iopub.status.idle":"2022-07-24T08:49:11.163986Z","shell.execute_reply.started":"2022-07-24T08:49:11.152183Z","shell.execute_reply":"2022-07-24T08:49:11.162683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\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-07-24T08:49:11.165662Z","iopub.execute_input":"2022-07-24T08:49:11.166051Z","iopub.status.idle":"2022-07-24T08:49:11.176796Z","shell.execute_reply.started":"2022-07-24T08:49:11.166015Z","shell.execute_reply":"2022-07-24T08:49:11.175741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# model save class:","metadata":{}},{"cell_type":"code","source":"model_path = \"./model_weights\"\nif not os.path.exists(model_path):\n    os.makedirs(model_path)\n    \ndef save_model(model, optimizer, criterion, epoch):\n    \"\"\"\n    Function to save the trained model to disk.\n    \"\"\"\n    print(f\"\\n Saving model at {epoch}th epoch\")\n    fname = f'{model_path}/{CONFIG[\"model_name\"]}-{CONFIG[\"LEARNING_RATE\"]}-{CONFIG[\"AUG\"]}-{epoch}-{CONFIG[\"BATCH_SIZE\"]}.pth'\n    torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'loss': criterion,\n                }, fname )\n    return fname","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:11.178076Z","iopub.execute_input":"2022-07-24T08:49:11.178558Z","iopub.status.idle":"2022-07-24T08:49:11.188910Z","shell.execute_reply.started":"2022-07-24T08:49:11.178520Z","shell.execute_reply":"2022-07-24T08:49:11.187920Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train & Validation Function:","metadata":{}},{"cell_type":"code","source":"def train_fn(loader, model, optimizer, loss_fn,epoch):\n    bar = tqdm(enumerate(loader), total=len(loader))\n    \n    \n    running_loss = 0.0\n    dataset_size = 0\n    \n    for batch_idx, (data, targets) in bar:\n        data = data.to(device=CONFIG[\"DEVICE\"])\n        targets = targets.float().to(device=CONFIG[\"DEVICE\"])\n\n        # forward\n        predictions = model(data)\n        loss = loss_fn(predictions, targets)\n\n        \n        optimizer.zero_grad()\n        # backward\n        loss.backward()\n        # Update Weights\n        optimizer.step()\n        \n        # Calculate Loss\n        running_loss += loss.item()\n        \n        epoch_loss = running_loss / len(loader)\n#         running_loss += (loss.item() * BATCH_SIZE)\n#         dataset_size += BATCH_SIZE\n        \n#         epoch_loss = running_loss / dataset_size\n        # update tqdm loop\n        bar.set_postfix(Epoch=epoch,loss=epoch_loss)\n    \n        \n    gc.collect()\n    \n    return epoch_loss","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:11.190289Z","iopub.execute_input":"2022-07-24T08:49:11.190749Z","iopub.status.idle":"2022-07-24T08:49:11.203075Z","shell.execute_reply.started":"2022-07-24T08:49:11.190713Z","shell.execute_reply":"2022-07-24T08:49:11.202168Z"},"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 = losses[CFG.loss]\n    \n    val_scores = []\n    \n    pbar = tqdm(enumerate(dataloader), total=len(dataloader))\n    for step, (images, masks) in pbar:\n        images  = images.to(CONFIG[\"DEVICE\"], dtype=torch.float)\n        masks   = masks.to(CONFIG[\"DEVICE\"], dtype=torch.float)\n        \n        \n        y_pred  = model(images)\n        loss    = loss_fn(y_pred, masks)\n        \n        running_loss += (loss.item() * CONFIG[\"BATCH_SIZE\"])\n        dataset_size += CONFIG[\"BATCH_SIZE\"]\n        \n        epoch_loss = running_loss / dataset_size\n        \n        y_pred = nn.Sigmoid()(y_pred)\n        val_dice = dice_coef(masks, y_pred).cpu().detach().numpy()\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}',\n                        lr=f'{current_lr:0.5f}',\n                        gpu_memory=f'{mem:0.2f} GB')\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-07-24T08:49:11.204435Z","iopub.execute_input":"2022-07-24T08:49:11.204912Z","iopub.status.idle":"2022-07-24T08:49:11.217686Z","shell.execute_reply.started":"2022-07-24T08:49:11.204875Z","shell.execute_reply":"2022-07-24T08:49:11.216614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train loop:","metadata":{}},{"cell_type":"code","source":"history = defaultdict(list)\nbest_val_dice = -float('inf')\nbest_val_jaccard = -float('inf')\nstart = time.time()\n\n\nfor epoch in range(CONFIG[\"NUM_EPOCHS\"]):\n    \n    print(f'{sr_}===================== Epoch: [{epoch+1}/{CONFIG[\"NUM_EPOCHS\"]}] =====================')\n\n    # train and validation loop\n    train_loss = train_fn(train_loader, model, optimizer, loss_fn, epoch+1)\n    val_loss, val_scores = valid_one_epoch(model, test_loader,CONFIG[\"DEVICE\"], epoch+1)\n    val_dice, val_jaccard = val_scores\n    \n    \n    # save model\n    model_file = save_model(model, optimizer, loss_fn, epoch+1)\n    \n    \n    # weight and baise Log the metrics\n    wandb.log({\"Train Loss\": train_loss})\n    wandb.log({\"Valid Loss\": val_loss})\n    wandb.log({\"Valid dice\": val_dice})\n    wandb.log({\"Valid jaccard\": val_jaccard})\n\n    # logging\n    history[\"epoch\"].append(epoch+1)\n    history['Train_Loss'].append(train_loss)\n    history['Valid_Loss'].append(val_loss)\n    history['Valid_jaccard'].append(val_jaccard)\n    history['Valid_dice'].append(val_dice)\n    history[\"model_file\"].append(model_file)\n\n    \n    # print loss and scores\n    print(f\"\\n Final Train Loss: {train_loss:0.4f}  |  Final Val Loss: {val_loss:0.4f}\") \n    \n    print(f'Valid Dice: {val_dice:0.4f} | Valid Jaccard: {val_jaccard:0.4f}')\n    \n    \n    # print best val score\n    if best_val_dice < val_dice:\n        print(f\"{b_}val dice increased: {best_val_dice} ---> {val_dice}\")\n        best_val_dice = val_dice\n    \n    if best_val_jaccard < val_jaccard:\n        print(f\"{b_}val jaccard increased: {best_val_jaccard} ---> {val_jaccard}\")\n        best_val_jaccard = val_jaccard\n\n        \n        \n# print time\nend = time.time()\ntime_elapsed = end - start\navg_time_per_epoch = time_elapsed/CONFIG[\"NUM_EPOCHS\"] \nsummary_df = pd.DataFrame.from_dict(history)\n\nprint(f'\\n ====================== [Training Summary] ====================== \\n')\nprint('Training complete in {:.0f}h {:.0f}m {:.0f}s'.format(\n    time_elapsed // 3600, (time_elapsed % 3600) // 60, (time_elapsed % 3600) % 60))\n\nprint('avg time per [train + val] epoch {:.0f}h {:.0f}m {:.0f}s'.format(\n    avg_time_per_epoch // 3600, (avg_time_per_epoch % 3600) // 60, (avg_time_per_epoch % 3600) % 60))\n\nprint(\"Best: ~~~~~~ dice: {:.4f}  || jaccard {:.4f} || val_loss {:.4f} || train_loss {:.4f} ~~~~~~\".format(\n    best_val_dice, best_val_jaccard, min(history['Valid_Loss']), min(history['Train_Loss'])))\n\n\nsummary_df.to_csv('training_summary.csv')\nprint(\"Saved training summary...\")\n\nprint(f'\\n =============================================================== \\n')\n\ndisplay(summary_df)\n\n# plotting training and validation loss\nplt.figure(figsize=(15,8))\n\nfig1 = plt.subplot(2,2,1)\nfig1.plot(history['epoch'], history['Train_Loss'], color='lime', marker='>')\nfig1.set_title(\"train loss [dice]\")\nfig1.set(xlabel='epoch', ylabel='train loss')\n\nfig2 = plt.subplot(2,2,2)\nfig2.plot(history['epoch'], history['Valid_Loss'], color='cyan',  marker='>')\nfig2.set_title(\"val loss [dice]\")\nfig2.set(xlabel='epoch', ylabel='val loss')\n\nfig3 = plt.subplot(2,2,3)\nfig3.plot(history['epoch'], history['Valid_dice'], color='orange',  marker='>')\nfig3.set_title(\"Val Dice\")\nfig3.set(xlabel='epoch', ylabel='val dice')\n\nfig4 = plt.subplot(2,2,4)\nfig4.plot(history['epoch'], history['Valid_jaccard'], color='lightcoral',  marker='>')\nfig4.set_title(\"val jaccard\")\nfig4.set(xlabel='epoch', ylabel='val jaccard')\n\nplt.show();","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:49:11.220623Z","iopub.execute_input":"2022-07-24T08:49:11.221238Z","iopub.status.idle":"2022-07-24T08:53:14.713441Z","shell.execute_reply.started":"2022-07-24T08:49:11.221211Z","shell.execute_reply":"2022-07-24T08:53:14.712444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model File Check:","metadata":{}},{"cell_type":"code","source":"!ls -la ./model_weights","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:53:14.714995Z","iopub.execute_input":"2022-07-24T08:53:14.716011Z","iopub.status.idle":"2022-07-24T08:53:15.437067Z","shell.execute_reply.started":"2022-07-24T08:53:14.715968Z","shell.execute_reply":"2022-07-24T08:53:15.435728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}