{"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":7141413,"sourceType":"datasetVersion","datasetId":4121860},{"sourceId":7149637,"sourceType":"datasetVersion","datasetId":4127856},{"sourceId":7139969,"sourceType":"datasetVersion","datasetId":4104676,"isSourceIdPinned":true}],"dockerImageVersionId":30587,"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-07T20:56:16.539621Z","iopub.execute_input":"2023-12-07T20:56:16.540250Z","iopub.status.idle":"2023-12-07T20:56:16.545518Z","shell.execute_reply.started":"2023-12-07T20:56:16.540205Z","shell.execute_reply":"2023-12-07T20:56:16.544488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q segmentation_models_pytorch","metadata":{"execution":{"iopub.status.busy":"2023-12-07T20:56:16.547088Z","iopub.execute_input":"2023-12-07T20:56:16.547383Z","iopub.status.idle":"2023-12-07T20:56:28.022993Z","shell.execute_reply.started":"2023-12-07T20:56:16.547358Z","shell.execute_reply":"2023-12-07T20:56:28.021634Z"},"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-07T20:56:28.025444Z","iopub.execute_input":"2023-12-07T20:56:28.025750Z","iopub.status.idle":"2023-12-07T20:56:28.033067Z","shell.execute_reply.started":"2023-12-07T20:56:28.025721Z","shell.execute_reply":"2023-12-07T20:56:28.032191Z"},"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-07T20:56:28.034137Z","iopub.execute_input":"2023-12-07T20:56:28.034415Z","iopub.status.idle":"2023-12-07T20:56:28.044467Z","shell.execute_reply.started":"2023-12-07T20:56:28.034391Z","shell.execute_reply":"2023-12-07T20:56:28.043673Z"},"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/800x800kidney2/train_k1_patch800_img'\nlabels_path = '/kaggle/input/800x800kidney2/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:12000]\nlabel_files = label_files[3000:12000]\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-07T20:56:28.046306Z","iopub.execute_input":"2023-12-07T20:56:28.046567Z","iopub.status.idle":"2023-12-07T20:56:28.153109Z","shell.execute_reply.started":"2023-12-07T20:56:28.046545Z","shell.execute_reply":"2023-12-07T20:56:28.152294Z"},"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 = 6\n    valid_bs = 24\n    img_size = [800,800]\n    epochs = 2\n    lr = 1e-3\n    over_lap = 0.2\n    patch_size = 800\n    bin_path = '/kaggle/input/resnext-k1-2/resnext_800k12.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        ], p=1.0),\n        \n        \"valid\": A.Compose([\n        ], p=1.0)\n    }\n    ","metadata":{"execution":{"iopub.status.busy":"2023-12-07T20:56:28.154138Z","iopub.execute_input":"2023-12-07T20:56:28.154421Z","iopub.status.idle":"2023-12-07T20:56:28.160801Z","shell.execute_reply.started":"2023-12-07T20:56:28.154397Z","shell.execute_reply":"2023-12-07T20:56:28.159887Z"},"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-07T20:56:28.162367Z","iopub.execute_input":"2023-12-07T20:56:28.162826Z","iopub.status.idle":"2023-12-07T20:56:28.174259Z","shell.execute_reply.started":"2023-12-07T20:56:28.162793Z","shell.execute_reply":"2023-12-07T20:56:28.173540Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# img = load_img('/kaggle/input/blood-vessel-segmentation/train/kidney_1_dense/images/1000.tif')\n# msk = load_msk('/kaggle/input/blood-vessel-segmentation/train/kidney_1_dense/labels/1000.tif')\n# print(img.shape)\n# plt.figure(figsize=(9, 4))\n# plt.axis('off')\n# plt.subplot(1,3,1)\n# plt.imshow(img)\n# plt.subplot(1,3,2)\n# plt.imshow(msk)\n# plt.subplot(1,3,3)\n# plt.imshow(img, cmap='bone')\n# plt.imshow(msk, alpha=0.5)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T20:56:28.175489Z","iopub.execute_input":"2023-12-07T20:56:28.175870Z","iopub.status.idle":"2023-12-07T20:56:28.185295Z","shell.execute_reply.started":"2023-12-07T20:56:28.175839Z","shell.execute_reply":"2023-12-07T20:56:28.184448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for image: h,w\n\n# def save_patches(img, patch_s,directory, name, over_lap = 0.1):\n#     l = []\n#     width = img.shape[1]\n#     height = img.shape[0]\n#     max_stride = patch_s * (1-over_lap)\n#     num_patches = np.ceil(np.array([height, width]) / max_stride).astype(np.int64)\n#     starts = [np.int64(np.linspace(0, width - patch_s, num_patches[1])),\n#                           np.int64(np.linspace(0, height - patch_s, num_patches[0]))]\n#     stops = [starts[0] + patch_s, starts[1] +patch_s]\n#     for y1, y2 in zip(starts[1], stops[1]):\n#         for x1, x2 in zip(starts[0], stops[0]):\n#             this_region = img[y1:y2, x1:x2]\n#             l.append(this_region)\n    \n#     # save the images: \n#     if not os.path.exists(directory):\n#         os.makedirs(directory)\n#     for i in range(len(l)):\n#         file_path = os.path.join(directory,f'{name}_{i}.tif')\n#         cv2.imwrite(file_path,l[i])\n        \n\n","metadata":{"execution":{"iopub.status.busy":"2023-12-07T20:56:28.186572Z","iopub.execute_input":"2023-12-07T20:56:28.187272Z","iopub.status.idle":"2023-12-07T20:56:28.196166Z","shell.execute_reply.started":"2023-12-07T20:56:28.187201Z","shell.execute_reply":"2023-12-07T20:56:28.195423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# create patched dataset","metadata":{}},{"cell_type":"code","source":"# for i, path in enumerate(image_train_files):\n#     img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n#     save_patches(img,512,'train5_img',name = f\"kidney_1_img_{i}\")\n# for i, path in enumerate(label_train_files):\n#     msk = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n#     save_patches(msk,512,'train5_msk',name = f\"kidney_1_msk_{i}\")\n","metadata":{"execution":{"iopub.status.busy":"2023-12-07T20:56:28.197324Z","iopub.execute_input":"2023-12-07T20:56:28.197779Z","iopub.status.idle":"2023-12-07T20:56:28.209208Z","shell.execute_reply.started":"2023-12-07T20:56:28.197753Z","shell.execute_reply":"2023-12-07T20:56:28.208478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## warning\ndataset may has the chance of loading empty images, which is weird, but need way to handle that \n1. create a checking loop, to check which file cause error then remove the file. ","metadata":{}},{"cell_type":"code","source":"# error_file_pairs = []\n# for i in range(len(image_train_files)):\n#     img_path = image_train_files[i]\n#     msk_path = label_train_files[i]\n#     try: \n#         load_img(img_path)\n#         load_msk(msk_path)\n#     except Exception as e:\n#         print(f\"error image and msk path is {img_path}, and {msk_path}\")\n#         error_file_pairs.append((img_path,msk_path))\n\n        ","metadata":{"execution":{"iopub.status.busy":"2023-12-07T20:56:28.212458Z","iopub.execute_input":"2023-12-07T20:56:28.212728Z","iopub.status.idle":"2023-12-07T20:56:28.221096Z","shell.execute_reply.started":"2023-12-07T20:56:28.212706Z","shell.execute_reply":"2023-12-07T20:56:28.220253Z"},"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-07T20:56:28.222184Z","iopub.execute_input":"2023-12-07T20:56:28.222537Z","iopub.status.idle":"2023-12-07T20:56:28.234973Z","shell.execute_reply.started":"2023-12-07T20:56:28.222503Z","shell.execute_reply":"2023-12-07T20:56:28.234099Z"},"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-07T20:56:28.236057Z","iopub.execute_input":"2023-12-07T20:56:28.236386Z","iopub.status.idle":"2023-12-07T20:56:28.251050Z","shell.execute_reply.started":"2023-12-07T20:56:28.236356Z","shell.execute_reply":"2023-12-07T20:56:28.250348Z"},"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-07T20:56:28.252089Z","iopub.execute_input":"2023-12-07T20:56:28.252536Z","iopub.status.idle":"2023-12-07T20:56:37.447186Z","shell.execute_reply.started":"2023-12-07T20:56:28.252503Z","shell.execute_reply":"2023-12-07T20:56:37.446268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"resnet: https://arxiv.org/abs/1611.05431","metadata":{}},{"cell_type":"code","source":"# sample_ids = [random.randint(0, len(X_train)) for _ in range(5)]\n# sample_psudotid = [random.randint(0, len(image_val_files)) for _ in range(5)]\n# for i in range(len(sample_ids)):\n#     img, msk = train_dataset[sample_ids[i]]\n#     img = img.permute((1, 2, 0)).numpy()*255.0\n#     img = img.astype('uint8')\n#     msk = (msk*255).numpy().astype('uint8')\n#     plt.figure(figsize=(9, 4))\n    \n#     p_img,p_msk = psudo_test_dataset[sample_psudotid[i]]\n#     p_img = p_img.permute((1, 2, 0)).numpy()*255.0\n#     p_img = p_img.astype('uint8')\n#     p_msk = (p_msk*255).numpy().astype('uint8')\n    \n#     plt.axis('off')\n#     plt.subplot(1,5,1)\n#     plt.imshow(img)\n#     plt.subplot(1,5,2)\n#     plt.imshow(msk)\n#     plt.subplot(1,5,3)\n#     plt.imshow(img, cmap='bone')\n#     plt.imshow(msk, alpha=0.5)\n    \n    \n#     plt.subplot(1,5,4)\n#     plt.imshow(p_img)\n#     plt.subplot(1,5,5)\n#     plt.imshow(p_msk)\n#     plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-07T20:56:37.448600Z","iopub.execute_input":"2023-12-07T20:56:37.449322Z","iopub.status.idle":"2023-12-07T20:56:37.454333Z","shell.execute_reply.started":"2023-12-07T20:56:37.449285Z","shell.execute_reply":"2023-12-07T20:56:37.453450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_model(backbone, num_classes, device):\n    model = smp.Unet(\n        encoder_name=backbone,      # choose encoder, e.g. mobilenet_v2 or efficientnet-b7\n        encoder_weights=None,     # use `imagenet` pre-trained weights for encoder initialization\n        in_channels=3,                  # model input channels (1 for gray-scale images, 3 for RGB, etc.)\n        classes=num_classes,        # model output channels (number of classes in your dataset)\n        activation=None,\n    )\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 = load_model(CFG.backbone, \n                   1, \n                   CFG.device, \n                   CFG.bin_path)\n# model = build_model(\n#     CFG.backbone,\n#     num_classes = 1,\n#     device = CFG.device\n#     )","metadata":{"execution":{"iopub.status.busy":"2023-12-07T20:56:37.455663Z","iopub.execute_input":"2023-12-07T20:56:37.456003Z","iopub.status.idle":"2023-12-07T20:56:38.082680Z","shell.execute_reply.started":"2023-12-07T20:56:37.455971Z","shell.execute_reply":"2023-12-07T20:56:38.081642Z"},"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-07T20:56:38.084104Z","iopub.execute_input":"2023-12-07T20:56:38.084507Z","iopub.status.idle":"2023-12-07T20:56:38.093701Z","shell.execute_reply.started":"2023-12-07T20:56:38.084471Z","shell.execute_reply":"2023-12-07T20:56:38.092851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = smp.losses.DiceLoss(mode='binary')\n","metadata":{"execution":{"iopub.status.busy":"2023-12-07T20:56:38.094765Z","iopub.execute_input":"2023-12-07T20:56:38.095077Z","iopub.status.idle":"2023-12-07T20:56:38.104455Z","shell.execute_reply.started":"2023-12-07T20:56:38.095052Z","shell.execute_reply":"2023-12-07T20:56:38.103698Z"},"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-07T20:56:38.105620Z","iopub.execute_input":"2023-12-07T20:56:38.105955Z","iopub.status.idle":"2023-12-07T20:56:38.116260Z","shell.execute_reply.started":"2023-12-07T20:56:38.105925Z","shell.execute_reply":"2023-12-07T20:56:38.115549Z"},"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-07T20:56:38.117358Z","iopub.execute_input":"2023-12-07T20:56:38.117681Z","iopub.status.idle":"2023-12-07T20:56:38.130631Z","shell.execute_reply.started":"2023-12-07T20:56:38.117656Z","shell.execute_reply":"2023-12-07T20:56:38.129864Z"},"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-07T20:56:38.131888Z","iopub.execute_input":"2023-12-07T20:56:38.132227Z","iopub.status.idle":"2023-12-07T20:56:38.147225Z","shell.execute_reply.started":"2023-12-07T20:56:38.132196Z","shell.execute_reply":"2023-12-07T20:56:38.146299Z"},"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-07T20:56:38.148411Z","iopub.execute_input":"2023-12-07T20:56:38.149471Z","iopub.status.idle":"2023-12-07T20:56:38.162943Z","shell.execute_reply.started":"2023-12-07T20:56:38.149445Z","shell.execute_reply":"2023-12-07T20:56:38.162205Z"},"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-07T20:56:38.163845Z","iopub.execute_input":"2023-12-07T20:56:38.164118Z","iopub.status.idle":"2023-12-07T20:56:38.176466Z","shell.execute_reply.started":"2023-12-07T20:56:38.164095Z","shell.execute_reply":"2023-12-07T20:56:38.175681Z"},"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-07T20:56:38.177745Z","iopub.execute_input":"2023-12-07T20:56:38.178100Z","iopub.status.idle":"2023-12-07T20:56:38.190121Z","shell.execute_reply.started":"2023-12-07T20:56:38.178065Z","shell.execute_reply":"2023-12-07T20:56:38.189390Z"},"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-07T20:56:38.191436Z","iopub.execute_input":"2023-12-07T20:56:38.191719Z","iopub.status.idle":"2023-12-07T20:57:14.769743Z","shell.execute_reply.started":"2023-12-07T20:56:38.191696Z","shell.execute_reply":"2023-12-07T20:57:14.768642Z"},"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-07T20:57:17.613745Z","iopub.execute_input":"2023-12-07T20:57:17.614555Z","iopub.status.idle":"2023-12-07T20:57:30.430715Z","shell.execute_reply.started":"2023-12-07T20:57:17.614521Z","shell.execute_reply":"2023-12-07T20:57:30.429419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}