{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":7117294,"sourceType":"datasetVersion","datasetId":4104676,"isSourceIdPinned":true},{"sourceId":7141413,"sourceType":"datasetVersion","datasetId":4121860},{"sourceId":7148274,"sourceType":"datasetVersion","datasetId":4126872}],"dockerImageVersionId":30588,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-07T17:10:37.834128Z","iopub.execute_input":"2023-12-07T17:10:37.834519Z","iopub.status.idle":"2023-12-07T17:10:38.207847Z","shell.execute_reply.started":"2023-12-07T17:10:37.834477Z","shell.execute_reply":"2023-12-07T17:10:38.206964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q segmentation_models_pytorch","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:10:38.209931Z","iopub.execute_input":"2023-12-07T17:10:38.210406Z","iopub.status.idle":"2023-12-07T17:10:57.960276Z","shell.execute_reply.started":"2023-12-07T17:10:38.210373Z","shell.execute_reply":"2023-12-07T17:10:57.95892Z"},"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","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:10:57.962217Z","iopub.execute_input":"2023-12-07T17:10:57.963283Z","iopub.status.idle":"2023-12-07T17:11:04.743254Z","shell.execute_reply.started":"2023-12-07T17:10:57.963236Z","shell.execute_reply":"2023-12-07T17:11:04.742442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# File preparation\n1. get all training / validation images/masks directory\n2. don't spend too much time on validation ","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:11:04.744409Z","iopub.execute_input":"2023-12-07T17:11:04.74474Z","iopub.status.idle":"2023-12-07T17:11:04.749152Z","shell.execute_reply.started":"2023-12-07T17:11:04.744714Z","shell.execute_reply":"2023-12-07T17:11:04.748155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nbase_path = '/kaggle/input/blood-vessel-segmentation/train'  \ntrain_base_path = '/kaggle/input/patched-sennet-kidney-1-data'\n\n# train_dataset = ['kidney_1_dense'] #,'kidney_2']\nval_img = 'kidney_3_sparse'\nval_mask = 'kidney_3_sparse'\n\nimage_train_files = []\nlabel_train_files = []\n\n\n# for dataset in train_dataset:\n\n# images_path = os.path.join(train_base_path, dataset, 'images')\n# labels_path = os.path.join(train_base_path, dataset, 'labels')\n# images_path = '/kaggle/input/800x800kidney2/train_k1_patch800_img'\n# labels_path = '/kaggle/input/800x800kidney2/train_k1_patch800_msk'\nimages_path = '/kaggle/input/patched-kidney-1/train_k1_patch800_img'\nlabels_path = '/kaggle/input/patched-kidney-1/train_k1_patch800_msk'\nimage_files = sorted([os.path.join(images_path, f) for f in os.listdir(images_path) if f.endswith('.tif')])\nlabel_files = sorted([os.path.join(labels_path, f) for f in os.listdir(labels_path) if f.endswith('.tif')])\nimage_train_files.extend(image_files)\nlabel_train_files.extend(label_files)\nimage_train_files = image_train_files[3000:]\nlabel_files = label_files[3000:]\nprint(f'len of image path {len(image_train_files)}')\nX_train, X_val, y_train, y_val = train_test_split(image_train_files, label_files, test_size=0.3)\n\nimages_val_path = os.path.join(base_path, val_img, 'images')\nlabels_val_path = os.path.join(base_path, val_mask, 'labels')\nimage_val_files = sorted([os.path.join(images_val_path, f) for f in os.listdir(images_val_path) if f.endswith('.tif')])\nlabel_val_files = sorted([os.path.join(labels_val_path, f) for f in os.listdir(labels_val_path) if f.endswith('.tif')])\n# image_val_files = image_val_files[1000:1500]\n# label_val_files = label_val_files[1000:1500]\nprint(f\"len of val path {len(image_val_files)}\")\nprint(len(label_val_files))","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:11:04.75223Z","iopub.execute_input":"2023-12-07T17:11:04.752525Z","iopub.status.idle":"2023-12-07T17:11:05.647938Z","shell.execute_reply.started":"2023-12-07T17:11:04.752485Z","shell.execute_reply":"2023-12-07T17:11:05.646744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    backbone = \"resnext50_32x4d\"\n    train_bs = 1\n    valid_bs = 1\n    img_size = [800,800]\n    epochs = 10\n    lr = 1e-3\n    over_lap = 0.3\n    patch_size = 800\n    bin_path = '/kaggle/input/resnext-k1-800/resnext_k1_800.bin'\n\n    num_classes   = 1\n    device        = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    data_transforms = {\n        \"train\": A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.ShiftScaleRotate(p=0.5),\n        ], p=1.0),\n        \n        \"valid\": A.Compose([\n        ], p=1.0)\n    }\n    ","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:12:29.478598Z","iopub.execute_input":"2023-12-07T17:12:29.479023Z","iopub.status.idle":"2023-12-07T17:12:29.487655Z","shell.execute_reply.started":"2023-12-07T17:12:29.478989Z","shell.execute_reply":"2023-12-07T17:12:29.486414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataloader:\nthis dataloader return the original size for testing ","metadata":{}},{"cell_type":"code","source":"def load_img(path):\n    img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n    img = np.tile(img[...,None], [1, 1, 3]) # gray to rgb\n    img = img.astype('float32') # original is uint16\n    mx = np.max(img)\n    if mx:\n        img/=mx # scale image to [0, 1]\n    return img\n\ndef load_msk(path):\n    msk = cv2.imread(path, cv2.IMREAD_UNCHANGED) \n    msk = msk.astype('float32')\n    msk/=255.0\n    return msk","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:12:29.489663Z","iopub.execute_input":"2023-12-07T17:12:29.48999Z","iopub.status.idle":"2023-12-07T17:12:29.500154Z","shell.execute_reply.started":"2023-12-07T17:12:29.489963Z","shell.execute_reply":"2023-12-07T17:12:29.49903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BuildDataset(torch.utils.data.Dataset):\n    def __init__(self, img_paths, msk_paths=[], transforms=None):\n        self.img_paths  = img_paths\n        self.msk_paths  = msk_paths\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.img_paths)\n    \n    def __getitem__(self, index):\n        img_path  = self.img_paths[index]\n        img = load_img(img_path)\n        \n        if len(self.msk_paths)>0:\n            msk_path = self.msk_paths[index]\n            msk = load_msk(msk_path)\n            if self.transforms:\n                data = self.transforms(image=img, mask=msk)\n                img  = data['image']\n                msk  = data['mask']\n            img = np.transpose(img, (2, 0, 1))\n            return torch.tensor(img), torch.tensor(msk)\n        else:\n            orig_size = img.shape\n            if self.transforms:\n                data = self.transforms(image=img)\n                img  = data['image']\n            img = np.transpose(img, (2, 0, 1))\n            return torch.tensor(img), torch.tensor(np.array([orig_size[0], orig_size[1]]))","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:12:29.501651Z","iopub.execute_input":"2023-12-07T17:12:29.502661Z","iopub.status.idle":"2023-12-07T17:12:29.514536Z","shell.execute_reply.started":"2023-12-07T17:12:29.502621Z","shell.execute_reply":"2023-12-07T17:12:29.513471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#X_train, X_val, y_train, y_val\n# image_val_files = sorted([os.path.join(images_val_path, f) for f in os.listdir(images_val_path) if f.endswith('.tif')])\n# label_val_files\ntrain_dataset = BuildDataset(X_train, y_train, transforms=CFG.data_transforms['train'])\nvalid_dataset = BuildDataset(X_val, y_val, transforms=None)\npsudo_test_dataset = BuildDataset(image_val_files,label_val_files,transforms=None)\ntrain_loader = DataLoader(train_dataset, batch_size=CFG.train_bs, num_workers=0, shuffle=True, pin_memory=True, drop_last=False)\nvalid_loader = DataLoader(valid_dataset, batch_size=CFG.valid_bs, num_workers=0, shuffle=False, pin_memory=True)\npsudo_test_loader = DataLoader(psudo_test_dataset, batch_size=1, num_workers=0, shuffle=False, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:12:29.517087Z","iopub.execute_input":"2023-12-07T17:12:29.517788Z","iopub.status.idle":"2023-12-07T17:12:29.531502Z","shell.execute_reply.started":"2023-12-07T17:12:29.517747Z","shell.execute_reply":"2023-12-07T17:12:29.530361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"resnet: https://arxiv.org/abs/1611.05431","metadata":{}},{"cell_type":"code","source":"class attention_block(nn.Module):\n    def __init__(self, F_g, F_l, n_coefficients):\n        super(attention_block, self).__init__()\n\n        self.bypass = nn.Sequential(\n            nn.Conv2d(F_g, n_coefficients, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(n_coefficients)\n        )\n\n        self.upsample = nn.Sequential(\n            nn.Conv2d(F_l, n_coefficients, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(n_coefficients)\n        )\n\n        self.psi = nn.Sequential(\n            nn.Conv2d(n_coefficients, 1, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(1),\n            nn.Sigmoid()\n        )\n\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, gate, skip_connection):\n        \"\"\"\n        :param gate: gating signal from previous layer\n        :param skip_connection: activation from corresponding encoder layer\n        :return: output activations\n        \"\"\"\n        g1 = self.upsample(gate)\n        x1 = self.bypass(skip_connection)\n        psi = self.relu(g1 + x1)\n        psi = self.psi(psi)\n        out = skip_connection * psi\n        return out\n\nclass MyUNet(nn.Module):\n    def contracting_block(self, in_channels, out_channels, kernel_size=3):\n        block = torch.nn.Sequential(\n                    torch.nn.Conv2d(kernel_size=kernel_size, in_channels=in_channels, out_channels=out_channels, padding=1),\n                    torch.nn.ReLU(),\n                    torch.nn.BatchNorm2d(out_channels),\n                    torch.nn.Conv2d(kernel_size=kernel_size, in_channels=out_channels, out_channels=out_channels, padding=1),\n                    torch.nn.ReLU(),\n                )\n        return block\n    def expansive_block(self, in_channels, mid_channel, out_channels, kernel_size=3):\n            block = torch.nn.Sequential(\n                    torch.nn.Conv2d(kernel_size=kernel_size, in_channels=in_channels, out_channels=mid_channel, padding=1),\n                    torch.nn.ReLU(),\n                    torch.nn.BatchNorm2d(mid_channel),\n                    torch.nn.Conv2d(kernel_size=kernel_size, in_channels=mid_channel, out_channels=mid_channel, padding=1),\n                    torch.nn.ReLU(),\n                    torch.nn.BatchNorm2d(mid_channel),\n                    torch.nn.ConvTranspose2d(in_channels=mid_channel, out_channels=out_channels, kernel_size=3, stride=2, padding=1, output_padding=1)\n            )\n            return  block\n\n    def final_block(self, in_channels, mid_channel, out_channels, kernel_size=1):\n            block = torch.nn.Sequential(\n                    \n                    torch.nn.Conv2d(kernel_size=3, in_channels=in_channels, out_channels=mid_channel, padding=1),\n                    torch.nn.ReLU(),\n                    torch.nn.BatchNorm2d(mid_channel),\n                    torch.nn.Conv2d(kernel_size=kernel_size, in_channels=mid_channel, out_channels=out_channels, padding=0),\n                    torch.nn.ReLU(),\n                    torch.nn.BatchNorm2d(out_channels),\n            )\n            return  block\n\n    def __init__(self, in_channel, out_channel):\n        super(MyUNet, self).__init__()\n        #Encode\n        self.conv_encode1 = self.contracting_block(in_channels=in_channel, out_channels=64)\n        self.conv_maxpool1 = torch.nn.MaxPool2d(kernel_size=2)\n        self.conv_encode2 = self.contracting_block(64, 128)\n        self.conv_maxpool2 = torch.nn.MaxPool2d(kernel_size=2)\n        self.conv_encode3 = self.contracting_block(128, 256)\n        self.conv_maxpool3 = torch.nn.MaxPool2d(kernel_size=2)\n        # Bottleneck\n        self.bottleneck = torch.nn.Sequential(\n                            torch.nn.Conv2d(kernel_size=3, in_channels=256, out_channels=512, padding=1),\n                            torch.nn.ReLU(),\n                            torch.nn.BatchNorm2d(512),\n                            torch.nn.Conv2d(kernel_size=3, in_channels=512, out_channels=512, padding=1),\n                            torch.nn.ReLU(),\n                            torch.nn.BatchNorm2d(512),\n                            torch.nn.ConvTranspose2d(in_channels=512, out_channels=256, kernel_size=3, stride=2, padding=1, output_padding=1)\n                            )\n\n        # Decode\n        self.attention3 = attention_block(256, 256, 128)\n        self.conv_decode3 = self.expansive_block(512, 256, 128)\n        self.attention2 = attention_block(128, 128, 64)\n        self.conv_decode2 = self.expansive_block(256, 128, 64)\n        self.attention1 = attention_block(64, 64, 32)\n        self.final_layer = self.final_block(128, 64, out_channel)\n    \n    def forward(self, x):\n        # Encode\n        encode_block1 = self.conv_encode1(x)\n        encode_pool1 = self.conv_maxpool1(encode_block1)\n        encode_block2 = self.conv_encode2(encode_pool1)\n        encode_pool2 = self.conv_maxpool2(encode_block2)\n        encode_block3 = self.conv_encode3(encode_pool2)\n        encode_pool3 = self.conv_maxpool3(encode_block3)\n        # Bottleneck\n        bottle_neck1 = self.bottleneck(encode_pool3)\n        # Decode\n        att_3 = self.attention3(bottle_neck1, encode_block3) \n        decode_block1 = torch.cat((bottle_neck1, att_3), 1)\n        cat_layer2 = self.conv_decode3(decode_block1)\n        \n        att_2 = self.attention2(cat_layer2, encode_block2) \n        decode_block2 = torch.cat((cat_layer2, att_2), 1)\n        cat_layer1 = self.conv_decode2(decode_block2)\n        \n        \n        att_1 = self.attention1(cat_layer1, encode_block1) \n        decode_block3 = torch.cat((cat_layer1, att_1), 1)\n        final_layer = self.final_layer(decode_block3)\n        return final_layer\ndef build_model(backbone, num_classes, device):\n    model = MyUNet(in_channel=3, out_channel=num_classes)\n    model.to(device)\n    return model\n\ndef load_model(backbone, num_classes, device, path):\n    model = build_model(backbone, num_classes, device)\n    model.load_state_dict(torch.load(path))\n    return model\n\nmodel = build_model(CFG.backbone, \n                   1, \n                   CFG.device, )\n# model = build_model(\n#     CFG.backbone,\n#     num_classes = 1,\n#     device = CFG.device\n#     )","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:12:29.602119Z","iopub.execute_input":"2023-12-07T17:12:29.602488Z","iopub.status.idle":"2023-12-07T17:12:29.722748Z","shell.execute_reply.started":"2023-12-07T17:12:29.60246Z","shell.execute_reply":"2023-12-07T17:12:29.721856Z"},"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.unsqueeze(1).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.unsqueeze(1).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":"2023-12-07T17:12:29.724421Z","iopub.execute_input":"2023-12-07T17:12:29.724756Z","iopub.status.idle":"2023-12-07T17:12:29.734086Z","shell.execute_reply.started":"2023-12-07T17:12:29.724728Z","shell.execute_reply":"2023-12-07T17:12:29.732955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = smp.losses.DiceLoss(mode='binary')\n","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:12:29.735363Z","iopub.execute_input":"2023-12-07T17:12:29.735803Z","iopub.status.idle":"2023-12-07T17:12:29.744999Z","shell.execute_reply.started":"2023-12-07T17:12:29.735775Z","shell.execute_reply":"2023-12-07T17:12:29.744046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_one_epoch(model, optimizer, scheduler, dataloader, device, epoch):\n    model.train()\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        images = images.to(device, dtype=torch.float)\n        masks  = masks.to(device, dtype=torch.float)\n        \n        batch_size = images.size(0)\n    \n        y_pred = model(images)\n        loss   = criterion(y_pred, masks)\n        loss.backward()\n        optimizer.step()\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":{"execution":{"iopub.status.busy":"2023-12-07T17:12:29.747235Z","iopub.execute_input":"2023-12-07T17:12:29.749061Z","iopub.status.idle":"2023-12-07T17:12:29.760795Z","shell.execute_reply.started":"2023-12-07T17:12:29.749031Z","shell.execute_reply":"2023-12-07T17:12:29.759842Z"},"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        images  = images.to(device, dtype=torch.float)\n        masks   = masks.to(device, dtype=torch.float)\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":{"execution":{"iopub.status.busy":"2023-12-07T17:12:29.76225Z","iopub.execute_input":"2023-12-07T17:12:29.762587Z","iopub.status.idle":"2023-12-07T17:12:29.776466Z","shell.execute_reply.started":"2023-12-07T17:12:29.762557Z","shell.execute_reply":"2023-12-07T17:12:29.775578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# pseudo test on performance","metadata":{}},{"cell_type":"code","source":"\ndef patch_image(img, patch_size, model = None, over_lap=0.2):\n    \"\"\"\n    Splits the image into patches with overlap.\n\n    \"\"\"\n    shape = img.shape\n\n    height, width = shape[2],shape[3]\n\n    stride = patch_size * (1 - over_lap)\n    num_patches = np.ceil(np.array([height, width]) / stride).astype(np.int64)\n    starts = [np.int64(np.linspace(0, width - patch_size, num_patches[1])),\n              np.int64(np.linspace(0, height - patch_size, num_patches[0]))]\n    patches = []\n    for y in starts[1]:\n        for x in starts[0]:\n            if model != None: \n                patch_img = img[:,:,y:y + patch_size, x:x + patch_size]\n#                 print(type(patch_img),\" and shape is \",patch_img.shape)\n#                 print(f'inside the patch function: size of patched image: {np.shape(patch_img)}')\n                patches.append(patch_img)\n    patches = torch.cat(patches,dim = 0)\n    pred = model(patches)\n    return pred\n\n\ndef combine_patches_torch(patches, original_shape, patch_size, over_lap=0.1):\n    height, width = original_shape[2],original_shape[3]\n    stride = int(patch_size * (1 - over_lap))\n    combined = np.zeros((height, width), dtype=np.float32)\n    weight = np.zeros((height, width), dtype=np.float32)\n\n    num_patches_y = np.ceil(height / stride).astype(np.int64)\n    num_patches_x = np.ceil(width / stride).astype(np.int64)\n\n    starts_y = np.linspace(0, height - patch_size, num_patches_y).astype(np.int64)\n    starts_x = np.linspace(0, width - patch_size, num_patches_x).astype(np.int64)\n\n    patch_idx = 0\n    for y in starts_y:\n        for x in starts_x:\n            \n            patch = patches[patch_idx].detach().cpu()\n            patch = patch.numpy().astype(np.float32)\n#             print(f'inside the combine function: type of combine = {type(combined)}, shape of patches = {np.shape(patch)}')\n            # with torch, I cannot add different sized tensor together \n            combined[y:y + patch_size, x:x + patch_size] += patch.squeeze()\n            weight[y:y + patch_size, x:x + patch_size] += 1.0\n            patch_idx += 1\n\n    # Avoid division by zero\n    weight[weight == 0] = 1.0\n    combined = combined / weight\n    combined = torch.from_numpy(combined)\n    combined = combined.unsqueeze(0).unsqueeze(0)\n    \n    return combined\n\n","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:12:29.777967Z","iopub.execute_input":"2023-12-07T17:12:29.778444Z","iopub.status.idle":"2023-12-07T17:12:29.79432Z","shell.execute_reply.started":"2023-12-07T17:12:29.778407Z","shell.execute_reply":"2023-12-07T17:12:29.793403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Sanity check","metadata":{}},{"cell_type":"code","source":"sample_ids = [random.randint(0, len(train_dataset)) for _ in range(20)]\nfor id in sample_ids:\n    img, msk =  train_dataset[id]\n    print(img.shape)\n    print(msk.shape)\n    img = img.permute((1, 2, 0)).numpy()*255.0\n    img = img.astype('uint8')\n    msk = (msk).numpy().astype('uint8')\n#     print(img.shape)\n#     print(msk.shape)\n    plt.figure(figsize=(9, 4))\n    plt.subplot(1,3,1)\n    plt.imshow(img)\n    plt.subplot(1,3,2)\n    plt.imshow(msk)\n    plt.show()\n#     msks = patch_image(msk,512)\n#     for i in msks: \n#         print(i.shape)\n#     ori_shape =msk.shape\n#     c_m = combine_patches(msks,original_shape = ori_shape,patch_size = 512)\n#     plt.subplot(1,3,3)\n#     plt.imshow(c_m)\n#     plt.show()\n    \n    ","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:12:29.795606Z","iopub.execute_input":"2023-12-07T17:12:29.795913Z","iopub.status.idle":"2023-12-07T17:12:39.226608Z","shell.execute_reply.started":"2023-12-07T17:12:29.795887Z","shell.execute_reply":"2023-12-07T17:12:39.225574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef test_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        images  = images.to(device, dtype=torch.float)\n        masks   = masks.to(device, dtype=torch.float)\n        \n        batch_size = images.size(0)\n        ori_shape = images.shape\n        patches = patch_image(images,patch_size=CFG.patch_size,over_lap = CFG.over_lap,model = model)\n        y_pred  = combine_patches_torch(patches,ori_shape,patch_size=CFG.patch_size,over_lap = CFG.over_lap)\n        y_pred=y_pred.to(device, dtype=torch.float)\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      \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":"2023-12-07T17:12:39.2278Z","iopub.execute_input":"2023-12-07T17:12:39.228076Z","iopub.status.idle":"2023-12-07T17:12:39.241391Z","shell.execute_reply.started":"2023-12-07T17:12:39.228052Z","shell.execute_reply":"2023-12-07T17:12:39.240353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from colorama import Fore, Back, Style #?\nc_  = Fore.GREEN\nsr_ = Style.RESET_ALL","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:12:39.242589Z","iopub.execute_input":"2023-12-07T17:12:39.242873Z","iopub.status.idle":"2023-12-07T17:12:39.254046Z","shell.execute_reply.started":"2023-12-07T17:12:39.242844Z","shell.execute_reply":"2023-12-07T17:12:39.25316Z"},"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#         test_loss, test_scores = test_one_epoch(model,\n#                psudo_test_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#         print(f'pseudo test loss: {test_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":{"execution":{"iopub.status.busy":"2023-12-07T17:12:39.256625Z","iopub.execute_input":"2023-12-07T17:12:39.256964Z","iopub.status.idle":"2023-12-07T17:12:39.270383Z","shell.execute_reply.started":"2023-12-07T17:12:39.256939Z","shell.execute_reply":"2023-12-07T17:12:39.269257Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=CFG.lr)\nscheduler = None\nmodel, history = run_training(model, optimizer, scheduler,\n                                device=CFG.device,\n                                num_epochs=CFG.epochs)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:12:39.271744Z","iopub.execute_input":"2023-12-07T17:12:39.272097Z","iopub.status.idle":"2023-12-07T17:12:56.92987Z","shell.execute_reply.started":"2023-12-07T17:12:39.272061Z","shell.execute_reply":"2023-12-07T17:12:56.928691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_loss, test_scores = test_one_epoch(model,\n               psudo_test_loader,\n               device=CFG.device,\n               epoch=1)\nprint(test_loss)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T17:12:56.930797Z","iopub.status.idle":"2023-12-07T17:12:56.931127Z","shell.execute_reply.started":"2023-12-07T17:12:56.930968Z","shell.execute_reply":"2023-12-07T17:12:56.930983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}