{"metadata":{"accelerator":"GPU","colab":{"gpuType":"A100","provenance":[],"toc_visible":true,"machine_shape":"hm"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":6942858,"sourceType":"datasetVersion","datasetId":3987147},{"sourceId":7279644,"sourceType":"datasetVersion","datasetId":4220769},{"sourceId":102030226,"sourceType":"kernelVersion"},{"sourceId":157399369,"sourceType":"kernelVersion"}],"dockerImageVersionId":30626,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"papermill":{"default_parameters":{},"duration":176.912306,"end_time":"2023-12-20T11:08:49.002118","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2023-12-20T11:05:52.089812","version":"2.4.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction","metadata":{}},{"cell_type":"markdown","source":"First, the references (Thank you!): \n\n - Baseline: [[SenNet + HOA] Train - UNet simple baseline](https://www.kaggle.com/code/kashiwaba/sennet-hoa-inference-unet-simple-baseline)\n - Generating Sub-volumes: [2d-to-3d unet demo](https://www.kaggle.com/code/hengck23/2d-to-3d-unet-demo)\n - Augmentations: [PyTorch Dataset with Volumetric Augmentations](https://www.kaggle.com/code/limitz/pytorch-dataset-with-volumetric-augmentations)\n - Some more bits and pieces: [2.5d Cutting model baseline [training]](https://www.kaggle.com/code/yoyobar/2-5d-cutting-model-baseline-training)\n\nHello and good day! I'm new to kaggle contests and this is my first public notebook. Looking forward to learning more from your comments ^_^","metadata":{}},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"!pip install -q segmentation_models_pytorch","metadata":{"id":"d13d6cb4","outputId":"44f5dcbe-5f79-49d2-ae6f-3d4783eafbd5","papermill":{"duration":20.047858,"end_time":"2023-12-20T11:06:15.524995","exception":false,"start_time":"2023-12-20T11:05:55.477137","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T07:44:19.298029Z","iopub.execute_input":"2024-01-06T07:44:19.298802Z","iopub.status.idle":"2024-01-06T07:44:38.998072Z","shell.execute_reply.started":"2024-01-06T07:44:19.298762Z","shell.execute_reply":"2024-01-06T07:44:38.996924Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nfrom tqdm import tqdm\nimport pandas as pd\nimport numpy as np\nfrom glob import glob\nimport gc\nimport time\nfrom collections import defaultdict\nimport  matplotlib.pyplot as plt\nfrom matplotlib.patches import Rectangle\nimport copy\nimport cv2\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import lr_scheduler\nfrom torch.cuda import amp\nimport torch.optim as optim\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nfrom PIL import Image\nimport torchvision.transforms as T\n\nimport math\nimport torchvision.transforms.functional as TF\nimport torch.nn.functional as F\n\nfrom colorama import Fore, Back, Style\nc_  = Fore.GREEN\nsr_ = Style.RESET_ALL","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","id":"1d805e4b","papermill":{"duration":8.733611,"end_time":"2023-12-20T11:06:24.268823","exception":false,"start_time":"2023-12-20T11:06:15.535212","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T07:46:49.720064Z","iopub.execute_input":"2024-01-06T07:46:49.720470Z","iopub.status.idle":"2024-01-06T07:46:57.176638Z","shell.execute_reply.started":"2024-01-06T07:46:49.720438Z","shell.execute_reply":"2024-01-06T07:46:57.175650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    seed          = 42\n    debug         = True # set debug=False for Full Training\n#     exp_name      = 'baseline'\n#     comment       = 'unet-se_resnext50_32x4d-512x512'\n#     output_dir    = './'\n#     model_name    = 'Unet'\n#     backbone      = 'se_resnext50_32x4d'\n    train_bs      = 1\n    valid_bs      = 1\n    subvlm_size   = 256 # Height and Width of each sub-volume\n    subvlm_depth  = 64\n    overlap       = 16  # Overlap between sub-volumes\n    epochs        = 5\n    n_accumulate  = max(1, 4//train_bs)\n    lr            = 6e-5\n    scheduler     = 'OneCycleLR'\n#     min_lr        = 1e-7\n#     T_max         = int(2279/(train_bs*n_accumulate)*epochs)+50\n#     T_0           = 25\n#     warmup_epochs = 3\n#     wd            = 1e-6\n#     n_fold        = 5\n#     num_classes   = 1\n    device        = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n    gt_df = \"/kaggle/input/sennet-hoa-gt-data/gt.csv\"\n    data_root = \"/kaggle/input\"\n    train_groups = [\"kidney_1_dense\"]\n    valid_groups = [\"kidney_3_dense\"]\n    loss_func     = \"BCELoss\"\n\n#     data_transforms = {\n#         \"train\": T.Compose([\n#             T.ToTensor(),\n#         ]),\n#         \"valid\": T.Compose([\n#             T.ToTensor(),\n#         ]),\n#     }","metadata":{"id":"77636ad9","papermill":{"duration":0.080098,"end_time":"2023-12-20T11:06:24.359359","exception":false,"start_time":"2023-12-20T11:06:24.279261","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T07:47:04.487297Z","iopub.execute_input":"2024-01-06T07:47:04.487660Z","iopub.status.idle":"2024-01-06T07:47:04.519063Z","shell.execute_reply.started":"2024-01-06T07:47:04.487630Z","shell.execute_reply":"2024-01-06T07:47:04.518150Z"},"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)\nset_seed(CFG.seed)","metadata":{"id":"08809918","papermill":{"duration":0.023778,"end_time":"2023-12-20T11:06:24.393681","exception":false,"start_time":"2023-12-20T11:06:24.369903","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T07:47:06.252500Z","iopub.execute_input":"2024-01-06T07:47:06.252877Z","iopub.status.idle":"2024-01-06T07:47:06.264037Z","shell.execute_reply.started":"2024-01-06T07:47:06.252845Z","shell.execute_reply":"2024-01-06T07:47:06.263106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DataLoader","metadata":{"id":"8e7dfc1f","papermill":{"duration":0.009694,"end_time":"2023-12-20T11:06:24.413795","exception":false,"start_time":"2023-12-20T11:06:24.404101","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"### **Note**\n* I wrote this dataloader based on the [2d-to-3d demo notebook](https://www.kaggle.com/code/hengck23/2d-to-3d-unet-demo). Currently it returns subvolumes of shape CxDxHxW. \n* I volume-normalized the data as there were some overflow errors when directly casting to float16 and the existing ones did instance-based normalizing. [Link to Dataset](https://www.kaggle.com/datasets/tahseenislamsajon/sennet-hoa-train-normalized-png)\n* Only using RandomFlip for now as the dimensions used are not uniform. From the discussions [here](https://www.kaggle.com/competitions/blood-vessel-segmentation/discussion/456118#2557825), I inferred that 64x256x256 should be a minimum. So a good choice may be 256x256x256, but it will take too much memory. ","metadata":{}},{"cell_type":"code","source":"def norm_by_percentile(x, low=10, high=99.8, alpha=0.01):\n\txmin = np.percentile(x, low)\n\txmax = np.percentile(x, high)\n\tx = (x - xmin) / (xmax - xmin)\n\tif 1:\n\t\tx[x > 1] = (x[x > 1] - 1) * alpha + 1\n\t\tx[x < 0] = (x[x < 0]) * alpha\n\t# x = np.clip(x,0,1)\n\treturn x","metadata":{"papermill":{"duration":0.019188,"end_time":"2023-12-20T11:06:24.443049","exception":false,"start_time":"2023-12-20T11:06:24.423861","status":"completed"},"tags":[],"id":"88b83a69","execution":{"iopub.status.busy":"2024-01-02T09:43:20.796806Z","iopub.execute_input":"2024-01-02T09:43:20.797198Z","iopub.status.idle":"2024-01-02T09:43:20.803341Z","shell.execute_reply.started":"2024-01-02T09:43:20.797167Z","shell.execute_reply":"2024-01-02T09:43:20.802393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class OverlappingVolumeDataset(Dataset):\n    def __init__(self, images_paths, labels_paths, depth, size, overlap, transform=None):\n        \n        self.images_paths = sorted(images_paths)\n        self.labels_paths = sorted(labels_paths)\n        self.depth = depth\n        self.size = size\n        self.overlap = overlap\n        self.transform = transform\n\n        self.full_volume, self.full_label = self.load_full_volume()\n\n        # Calculate number of sub-volumes\n        D, H, W = self.full_volume.shape\n        num_D = int(np.ceil((D - overlap) / (depth - overlap)))\n        num_H = int(np.ceil((H - overlap) / (size - overlap)))\n        num_W = int(np.ceil((W - overlap) / (size - overlap)))\n\n        # Calculate sub-volume indices \n        self.zz = np.linspace(0, D-depth, num_D).astype(int).tolist()\n        self.yy = np.linspace(0, H-size, num_H).astype(int).tolist()\n        self.xx = np.linspace(0, W-size, num_W).astype(int).tolist()\n\n    def load_full_volume(self):\n        images = [cv2.imread(f, cv2.IMREAD_GRAYSCALE).astype('float16') for f in self.images_paths]\n        volume = np.stack(images)\n        # norm not being used here since the data i'm using is already normalized\n#         volume = norm_by_percentile(volume)\n        volume = volume / 255.0\n        del images\n        \n        labels = [cv2.imread(f, cv2.IMREAD_GRAYSCALE).astype('float16') for f in self.labels_paths]\n        label = np.stack(labels)\n        label = label / 255.0\n        del labels\n        \n        return volume, label\n\n    def __len__(self):\n        return len(self.zz) * len(self.yy) * len(self.xx)\n\n    def __getitem__(self, idx):\n        D_idx = idx // (len(self.yy) * len(self.xx))\n        H_idx = (idx // len(self.xx)) % len(self.yy)\n        W_idx = idx % len(self.xx)\n\n        z, y, x = self.zz[D_idx], self.yy[H_idx], self.xx[W_idx]\n        image = self.full_volume[z:z + self.depth, y:y + self.size, x:x + self.size]\n        label = self.full_label[z:z + self.depth, y:y + self.size, x:x + self.size]\n\n        image = np.expand_dims(image, axis = 0)\n        label = np.expand_dims(label, axis = 0)\n\n        image = torch.from_numpy(image).to(torch.float32)\n        label = torch.from_numpy(label).to(torch.float32)\n\n        if self.transform:\n            rng = torch.get_rng_state()\n            image = self.transform(image)\n            torch.set_rng_state(rng)\n            label = self.transform(label)\n\n        return image, label\n","metadata":{"id":"e98480f3","papermill":{"duration":0.028582,"end_time":"2023-12-20T11:06:24.482152","exception":false,"start_time":"2023-12-20T11:06:24.45357","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T07:54:02.488469Z","iopub.execute_input":"2024-01-06T07:54:02.488861Z","iopub.status.idle":"2024-01-06T07:54:02.505366Z","shell.execute_reply.started":"2024-01-06T07:54:02.488831Z","shell.execute_reply":"2024-01-06T07:54:02.504406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# The augmentations. source: https://www.kaggle.com/code/limitz/pytorch-dataset-with-volumetric-augmentations\n\nclass RandomRotationNd(nn.Module):\n    ''' This augmentation first permutes the dimensions as an initial rotation\n        to select the rotation axis, then rotates around the (fixed) z axis.\n        The result is zoomed in to remove empty space and finally permuted\n        once more to move randomize the rotation axis.\n    '''\n    def __init__(self, dims):\n        super().__init__()\n        self.dims = dims\n\n    def forward(self, x):\n        angle = torch.rand(1).item() * 360\n        keep = torch.arange(x.dim() - self.dims)\n        perm = -torch.randperm(self.dims)-1\n        x = x.clone().permute(*[k.item() for k in keep], *[p.item() for p in perm])\n        rad = math.pi * angle / 180\n        scale = abs(math.sin(rad)) + abs(math.cos(rad))\n        for i in range(0, x.shape[-3],8): # presumptuous\n            v = x[...,i:i+8,:,:]\n            w = v.view(-1, *v.shape[-3:])\n            w = TF.rotate(w, angle)\n            v = w.view(*v.shape)\n            x[...,i:i+8,:,:] = v\n        s = x.shape\n        x = F.interpolate(x, scale_factor=scale, mode=\"bilinear\")\n        x = TF.center_crop(x, s[-2:])\n        perm = -torch.randperm(self.dims)-1\n        x = x.permute(*[k.item() for k in keep], *[p.item() for p in perm])\n        return x\n\nclass RandomRot90Nd(nn.Module):\n    def __init__(self, dims):\n        super().__init__()\n        self.dims = dims\n\n    def forward(self, x):\n        dims = -torch.randperm(self.dims)[:2]-1\n        dims = [d.item() for d in dims]\n        rot = torch.randint(4, (1,)).item()\n        return x.rot90(rot, dims)\n\nclass RandomPermuteNd(nn.Module):\n    def __init__(self, dims):\n        super().__init__()\n        self.dims = dims\n\n    def forward(self, x):\n        perm = -torch.randperm(self.dims)-1\n        keep = torch.arange(x.dim() - self.dims)\n        return x.permute(*[k.item() for k in keep], *[p.item() for p in perm])\n\nclass RandomFlipNd(nn.Module):\n    def  __init__(self, dims, p=0.5):\n        super().__init__()\n        self.dims = dims\n        self.p = p\n\n    def forward(self, x):\n        for i in range(self.dims):\n            if torch.rand(1) < self.p:\n                x = x.flip(-i-1)\n        return x\n\nclass ToDevice(nn.Module):\n    ''' Sometimes it helps to move the tensor to the gpu before augmentations like\n        rotation. Note however that you need to set num_workers to 0 in the dataloader\n    '''\n    def __init__(self, device=None):\n        super().__init__()\n        self.device = device or (\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n    def forward(self, x):\n        return x.to(self.device)\n","metadata":{"id":"C_QiyIqzU5mA","execution":{"iopub.status.busy":"2024-01-06T07:55:20.839085Z","iopub.execute_input":"2024-01-06T07:55:20.839967Z","iopub.status.idle":"2024-01-06T07:55:20.858511Z","shell.execute_reply.started":"2024-01-06T07:55:20.839932Z","shell.execute_reply":"2024-01-06T07:55:20.857614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# chose RandomFlipNd only because the the 3 spatial dimensions are not the same\ntrain_transforms = T.Compose((ToDevice(), RandomFlipNd(3)))\n\nvalid_transforms = T.Compose((ToDevice(), ))","metadata":{"id":"r5fFyki3U5mB","execution":{"iopub.status.busy":"2024-01-06T07:55:40.006984Z","iopub.execute_input":"2024-01-06T07:55:40.007355Z","iopub.status.idle":"2024-01-06T07:55:40.012442Z","shell.execute_reply.started":"2024-01-06T07:55:40.007325Z","shell.execute_reply":"2024-01-06T07:55:40.011443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_groups = CFG.train_groups\nvalid_groups = CFG.valid_groups\ngt_df = pd.read_csv(CFG.gt_df)\ngt_df[\"img_path\"] = gt_df[\"img_path\"].apply(lambda x: os.path.join(CFG.data_root, x)).apply(lambda x: x.replace('blood-vessel-segmentation/train', 'sennet-hoa-train-normalized-png').replace('.tif', '.png'))\ngt_df[\"msk_path\"] = gt_df[\"msk_path\"].apply(lambda x: os.path.join(CFG.data_root, x)).apply(lambda x: x.replace('blood-vessel-segmentation/train', 'sennet-hoa-train-normalized-png').replace('.tif', '.png'))\ntrain_df = gt_df.query(\"group in @train_groups\").reset_index(drop=True)\nvalid_df = gt_df.query(\"group in @valid_groups\").reset_index(drop=True)\ntrain_img_paths = train_df[\"img_path\"].values.tolist()\ntrain_msk_paths = train_df[\"msk_path\"].values.tolist()\nvalid_img_paths = valid_df[\"img_path\"].values.tolist()\nvalid_msk_paths = valid_df[\"msk_path\"].values.tolist()\nif CFG.debug:\n    train_img_paths = train_img_paths[150:150+CFG.train_bs*CFG.subvlm_depth*3]\n    train_msk_paths = train_msk_paths[150:150+CFG.train_bs*CFG.subvlm_depth*3]\n    valid_img_paths = valid_img_paths[150:150+CFG.valid_bs*CFG.subvlm_depth*2]\n    valid_msk_paths = valid_msk_paths[150:150+CFG.valid_bs*CFG.subvlm_depth*2]","metadata":{"id":"69ea62fa","papermill":{"duration":1.32373,"end_time":"2023-12-20T11:06:25.871815","exception":false,"start_time":"2023-12-20T11:06:24.548085","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T07:55:41.931385Z","iopub.execute_input":"2024-01-06T07:55:41.932084Z","iopub.status.idle":"2024-01-06T07:55:43.299670Z","shell.execute_reply.started":"2024-01-06T07:55:41.932053Z","shell.execute_reply":"2024-01-06T07:55:43.298819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = OverlappingVolumeDataset(\n                    images_paths=train_img_paths,\n                    labels_paths=train_msk_paths,\n                    depth=CFG.subvlm_depth, size=CFG.subvlm_size, overlap=CFG.overlap,\n                    transform=train_transforms\n                )","metadata":{"id":"376126f0","papermill":{"duration":142.165637,"end_time":"2023-12-20T11:08:48.048006","exception":false,"start_time":"2023-12-20T11:06:25.882369","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T07:55:43.862496Z","iopub.execute_input":"2024-01-06T07:55:43.862848Z","iopub.status.idle":"2024-01-06T07:55:59.474601Z","shell.execute_reply.started":"2024-01-06T07:55:43.862821Z","shell.execute_reply":"2024-01-06T07:55:59.473774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_dataset = OverlappingVolumeDataset(\n                    images_paths=valid_img_paths,\n                    labels_paths=valid_msk_paths,\n                    depth=CFG.subvlm_depth, size=CFG.subvlm_size, overlap=CFG.overlap,\n                    transform=valid_transforms\n                )","metadata":{"id":"-W42iOY8rGRy","execution":{"iopub.status.busy":"2024-01-06T07:56:03.634831Z","iopub.execute_input":"2024-01-06T07:56:03.635258Z","iopub.status.idle":"2024-01-06T07:56:23.848401Z","shell.execute_reply.started":"2024-01-06T07:56:03.635222Z","shell.execute_reply":"2024-01-06T07:56:23.847594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(train_dataset, batch_size=CFG.train_bs, num_workers=0, shuffle=True, pin_memory=False)\nvalid_loader = DataLoader(valid_dataset, batch_size=CFG.valid_bs, num_workers=0, shuffle=False, pin_memory=False)","metadata":{"id":"7f389bc3","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T07:57:05.157152Z","iopub.execute_input":"2024-01-06T07:57:05.157528Z","iopub.status.idle":"2024-01-06T07:57:05.162986Z","shell.execute_reply.started":"2024-01-06T07:57:05.157497Z","shell.execute_reply":"2024-01-06T07:57:05.161909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Check Augmentations","metadata":{"id":"43e9fa96","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[]}},{"cell_type":"code","source":"import inline\n\nvolume, target = train_dataset[30]\n\nvolume = volume.sub(volume.mean()).div(volume.std().add(1e-5))\ninline.plot(torch.stack((inline.disp(volume[0]), inline.disp(target[0]))), width=10)\n\nprint(\"For show: more augmentation of the same subvolume\")\n# Showing the same subvolume, with random flips\nrot = RandomFlipNd(3)\nrng = torch.get_rng_state()\nvolumes = torch.stack([rot(volume)[0] for _ in range(4)])\ntorch.set_rng_state(rng)\ntargets = torch.stack([rot(target)[0] for _ in range(4)])\ninline.plot(volumes.mul(0.288).add(0.5)[:,[32]])\ninline.plot(volumes)\ninline.plot(targets)","metadata":{"id":"-iO6imhuU5mF","outputId":"eb36b3dd-e384-4223-aa03-289af93551ab","execution":{"iopub.status.busy":"2024-01-06T07:57:20.495586Z","iopub.execute_input":"2024-01-06T07:57:20.496517Z","iopub.status.idle":"2024-01-06T07:57:23.232672Z","shell.execute_reply.started":"2024-01-06T07:57:20.496481Z","shell.execute_reply":"2024-01-06T07:57:23.231761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model - Using Monai ","metadata":{"id":"b0a5b01f","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install -q monai[einops]","metadata":{"id":"6CLUi6XcU5mH","outputId":"e1da52dc-e00f-4ce6-be23-47a145379d09","execution":{"iopub.status.busy":"2024-01-06T07:58:03.442692Z","iopub.execute_input":"2024-01-06T07:58:03.443070Z","iopub.status.idle":"2024-01-06T07:58:16.683469Z","shell.execute_reply.started":"2024-01-06T07:58:03.443039Z","shell.execute_reply":"2024-01-06T07:58:16.682123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from monai.networks.nets import UNETR\n\nmodel = UNETR(\n    in_channels=1,\n    out_channels=1,\n    img_size=(64, 256, 256),  \n)\n\nmodel.to(CFG.device)","metadata":{"id":"gJoeiMKhU5mH","outputId":"b11923d3-9045-4d05-89c0-b51923a6ded3","execution":{"iopub.status.busy":"2024-01-06T07:58:35.514186Z","iopub.execute_input":"2024-01-06T07:58:35.514570Z","iopub.status.idle":"2024-01-06T07:59:14.332816Z","shell.execute_reply.started":"2024-01-06T07:58:35.514540Z","shell.execute_reply":"2024-01-06T07:59:14.331873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_check_net():\n    for subvolume, label in train_loader:\n        with torch.no_grad():\n            with torch.cuda.amp.autocast(enabled=True):\n                vessel = model(subvolume)\n\n        print(subvolume.shape)\n        print(label.shape)\n\n        print('output')\n        print(vessel.shape)\n\n#         del vessel, subvolume, label, net\n        break\n\nrun_check_net()","metadata":{"id":"2251afc6","outputId":"ac1aa593-4322-4786-fcf2-f8e6d64f346a","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T08:01:47.007691Z","iopub.execute_input":"2024-01-06T08:01:47.008112Z","iopub.status.idle":"2024-01-06T08:01:47.825528Z","shell.execute_reply.started":"2024-01-06T08:01:47.008080Z","shell.execute_reply":"2024-01-06T08:01:47.824567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss Function","metadata":{"id":"22f157a3","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[]}},{"cell_type":"code","source":"DiceLoss = smp.losses.DiceLoss(mode='binary')\nBCELoss = nn.BCEWithLogitsLoss()\ndef criterion(y_pred, y_true):\n    if CFG.loss_func == \"DiceLoss\":\n        return DiceLoss(y_pred, y_true)\n    elif CFG.loss_func == \"BCELoss\":\n        return BCELoss(y_pred, y_true)","metadata":{"id":"79d5c2be","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T08:01:50.539562Z","iopub.execute_input":"2024-01-06T08:01:50.539942Z","iopub.status.idle":"2024-01-06T08:01:50.545550Z","shell.execute_reply.started":"2024-01-06T08:01:50.539910Z","shell.execute_reply":"2024-01-06T08:01:50.544605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metrics","metadata":{"id":"a4ccf6d9","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[]}},{"cell_type":"code","source":"def dice_coef(y_true, y_pred, thr=0.5, dim=(2,3,4), 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,4), 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":{"id":"99471fcb","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T08:01:52.330781Z","iopub.execute_input":"2024-01-06T08:01:52.331164Z","iopub.status.idle":"2024-01-06T08:01:52.339629Z","shell.execute_reply.started":"2024-01-06T08:01:52.331135Z","shell.execute_reply":"2024-01-06T08:01:52.338466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimizer","metadata":{"id":"50712c0e","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[]}},{"cell_type":"code","source":"def fetch_scheduler(optimizer):\n    if CFG.scheduler == 'CosineAnnealingLR':\n        scheduler = lr_scheduler.CosineAnnealingLR(optimizer,T_max=CFG.T_max,\n                                                   eta_min=CFG.min_lr)\n    elif CFG.scheduler == 'CosineAnnealingWarmRestarts':\n        scheduler = lr_scheduler.CosineAnnealingWarmRestarts(optimizer,T_0=CFG.T_0,\n                                                             eta_min=CFG.min_lr)\n    elif CFG.scheduler == 'ReduceLROnPlateau':\n        scheduler = lr_scheduler.ReduceLROnPlateau(optimizer,\n                                                   mode='min',\n                                                   factor=0.1,\n                                                   patience=7,\n                                                   threshold=0.0001,\n                                                   min_lr=CFG.min_lr,)\n    elif CFG.scheduler == 'ExponentialLR':\n        scheduler = lr_scheduler.ExponentialLR(optimizer, gamma=0.85)\n\n    elif CFG.scheduler == 'OneCycleLR':\n        scheduler = lr_scheduler.OneCycleLR(optimizer, max_lr=0.01, steps_per_epoch=len(train_loader), epochs=CFG.epochs)\n\n    elif CFG.scheduler == None:\n        return None\n\n    return scheduler","metadata":{"id":"659feccf","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T08:01:54.660948Z","iopub.execute_input":"2024-01-06T08:01:54.661814Z","iopub.status.idle":"2024-01-06T08:01:54.669365Z","shell.execute_reply.started":"2024-01-06T08:01:54.661782Z","shell.execute_reply":"2024-01-06T08:01:54.668414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = optim.AdamW(model.parameters(), lr=CFG.lr)\nscheduler = fetch_scheduler(optimizer)","metadata":{"id":"3eec1afc","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T08:01:55.585881Z","iopub.execute_input":"2024-01-06T08:01:55.586712Z","iopub.status.idle":"2024-01-06T08:01:55.592410Z","shell.execute_reply.started":"2024-01-06T08:01:55.586680Z","shell.execute_reply":"2024-01-06T08:01:55.591490Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Check LR","metadata":{"id":"a733f214","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[]}},{"cell_type":"code","source":"_optimizer = optim.Adam(model.parameters(), lr=CFG.lr)\n_scheduler = fetch_scheduler(_optimizer)\nlr_list = []\nfor e in range(CFG.epochs):\n    for step in range(len(train_loader)):\n        lr_list.append(_optimizer.param_groups[0]['lr'])\n        if (step + 1) % CFG.n_accumulate == 0:\n            _optimizer.step()\n            _scheduler.step()\nplt.plot(np.array(range(len(lr_list))), np.array(lr_list))\nplt.show()\ndel _optimizer, _scheduler","metadata":{"id":"48e45b69","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T08:02:22.309596Z","iopub.execute_input":"2024-01-06T08:02:22.309988Z","iopub.status.idle":"2024-01-06T08:02:22.551735Z","shell.execute_reply.started":"2024-01-06T08:02:22.309955Z","shell.execute_reply":"2024-01-06T08:02:22.550841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Function","metadata":{"id":"627f7e22","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[]}},{"cell_type":"code","source":"def train_one_epoch(model, optimizer, scheduler, dataloader, device, epoch):\n    model.train()\n    scaler = amp.GradScaler()\n\n    dataset_size = 0\n    running_loss = 0.0\n\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    for step, (images, masks) in pbar:\n\n        batch_size = images.size(0)\n\n        with amp.autocast(enabled=True):\n            y_pred = model(images)\n            loss   = criterion(y_pred, masks)\n            loss   = loss / CFG.n_accumulate\n\n        scaler.scale(loss).backward()\n\n        if (step + 1) % CFG.n_accumulate == 0:\n            scaler.step(optimizer)\n            scaler.update()\n\n            # zero the parameter gradients\n            optimizer.zero_grad()\n\n            if scheduler is not None:\n                scheduler.step()\n\n        running_loss += (loss.item() * batch_size)\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( epoch=f'{epoch}',\n                          train_loss=f'{epoch_loss:0.4f}',\n                          lr=f'{current_lr:0.5f}',\n                          gpu_mem=f'{mem:0.2f} GB')\n    torch.cuda.empty_cache()\n    gc.collect()\n    return epoch_loss","metadata":{"id":"7f6ce287","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T08:02:25.988002Z","iopub.execute_input":"2024-01-06T08:02:25.988377Z","iopub.status.idle":"2024-01-06T08:02:25.998568Z","shell.execute_reply.started":"2024-01-06T08:02:25.988348Z","shell.execute_reply":"2024-01-06T08:02:25.997685Z"},"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\n    val_scores = []\n\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid ')\n    for step, (images, masks) in pbar:\n\n        batch_size = images.size(0)\n\n        y_pred  = model(images)\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        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    return epoch_loss, val_scores","metadata":{"id":"fea34bc8","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T08:02:27.262021Z","iopub.execute_input":"2024-01-06T08:02:27.262686Z","iopub.status.idle":"2024-01-06T08:02:27.272973Z","shell.execute_reply.started":"2024-01-06T08:02:27.262651Z","shell.execute_reply":"2024-01-06T08:02:27.271750Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_training(model, optimizer, scheduler, device, num_epochs):\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_loss      = 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(model, optimizer, scheduler,\n                                           dataloader=train_loader,\n                                           device=CFG.device, epoch=epoch)\n\n        val_loss, val_scores = valid_one_epoch(model, valid_loader,\n                                                 device=CFG.device,\n                                                 epoch=epoch)\n        val_dice, val_jaccard = val_scores\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        print(f'Valid Dice: {val_dice:0.4f} | Valid Jaccard: {val_jaccard:0.4f}')\n        print(f'Valid Loss: {val_loss}')\n\n        # deep copy the model\n        if val_loss <= best_loss:\n            print(f\"{c_}Valid loss Improved ({best_loss} ---> {val_loss})\")\n            best_dice    = val_dice\n            best_jaccard = val_jaccard\n            best_loss = val_loss\n            best_epoch   = epoch\n            best_model_wts = copy.deepcopy(model.state_dict())\n            PATH = \"best_epoch.bin\"\n            torch.save(model.state_dict(), PATH)\n            print(f\"Model Saved{sr_}\")\n\n        last_model_wts = copy.deepcopy(model.state_dict())\n        PATH = \"last_epoch.bin\"\n        torch.save(model.state_dict(), PATH)\n\n        print(); print()\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 Loss: {:.4f}\".format(best_loss))\n\n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    return model, history","metadata":{"id":"43fd938a","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T08:02:27.762645Z","iopub.execute_input":"2024-01-06T08:02:27.763391Z","iopub.status.idle":"2024-01-06T08:02:27.775231Z","shell.execute_reply.started":"2024-01-06T08:02:27.763359Z","shell.execute_reply":"2024-01-06T08:02:27.774272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{"id":"e7834d3e","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[]}},{"cell_type":"code","source":"model, history = run_training(model, optimizer, scheduler,\n                                device=CFG.device,\n                                num_epochs=CFG.epochs\n                             )","metadata":{"id":"e8ea151a","outputId":"ae503a28-1153-4ae9-da3e-ceef2ba9e6f3","papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-01-06T08:02:29.802989Z","iopub.execute_input":"2024-01-06T08:02:29.803355Z","iopub.status.idle":"2024-01-06T08:18:22.301347Z","shell.execute_reply.started":"2024-01-06T08:02:29.803327Z","shell.execute_reply":"2024-01-06T08:18:22.300419Z"},"trusted":true},"execution_count":null,"outputs":[]}]}