{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"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.11.13"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":5144,"databundleVersionId":862050,"sourceType":"competition"},{"sourceId":6927,"databundleVersionId":45059,"sourceType":"competition"},{"sourceId":16306036,"sourceType":"kernelVersion"}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Necessary Imports","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset,DataLoader\nimport zipfile\nimport os","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:11.445580Z","iopub.execute_input":"2025-09-27T20:12:11.445840Z","iopub.status.idle":"2025-09-27T20:12:13.028310Z","shell.execute_reply.started":"2025-09-27T20:12:11.445816Z","shell.execute_reply":"2025-09-27T20:12:13.027729Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# To-Do\n- [x] Building a U-Net (using BatchNormalization) from scratch\n- [x] loading the dataset and visualising some images with training masks\n- [x] training the network, with logs like training-loss, validation-loss and validation metrics like IoU\n- [x] Using transfer learning on a new dataset like UltraSound Nerve Segmentation to test the generalization of the trained model\n------------\n# Model Specification \n- we made a U-Net model from the scratch\n- instead of using the archetecture discussed in the paper, we slightly modified it with non-zero paddding (this was motivated by the winning model of this competition)","metadata":{}},{"cell_type":"code","source":"import torchvision.transforms.functional as F\nclass InterConv(nn.Module):\n    def __init__(self,in_channels,out_channels,dropout_probab = 0.0):\n        super(InterConv,self).__init__()\n        layers=[\n            nn.Conv2d(in_channels,out_channels,3,1,1,bias=False), ##here \n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels,out_channels,3,1,1,bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n        ]\n        if dropout_probab > 0:\n            layers.append(nn.Dropout(p=dropout_probab))\n        self.conv = nn.Sequential(*layers)\n    def forward(self,x):\n        x = self.conv(x)\n        return x\n\n\nclass UNET(nn.Module):\n    def __init__(self,in_channels=3,out_channels=1,p=0.0,features=[64,128,256,512]):\n        super(UNET,self).__init__()\n        self.downs = nn.ModuleList()\n        self.ups = nn.ModuleList()\n        self.pool = nn.MaxPool2d(kernel_size=2,stride=2)\n\n        #downsampling - Part-01 of UNET\n        for feature in features:\n            self.downs.append(InterConv(in_channels,feature))\n            in_channels = feature\n            \n        # in between these two processes, their exists another process of \"Bottle-Necking\" used to match the dimension requirements \n        \n        #upsampling -  Part-02 of UNET\n        for feats in reversed(features):\n            self.ups.append(\n                nn.ConvTranspose2d(\n                    feats*2,\n                    feats,\n                    kernel_size=2,\n                    stride=2\n                )\n            )\n            self.ups.append(InterConv(feats*2,feats))\n\n        self.BottleNeck = InterConv(features[-1],features[-1]*2,dropout_probab=p)\n        self.FinalLayer = nn.Conv2d(features[0],out_channels,1)\n\n    def forward(self,x):\n        skip_connections = []\n        for down in self.downs:\n            x = down(x)\n            skip_connections.append(x) #these will be used in the decoder (upsampling)\n            x = self.pool(x)\n        x = self.BottleNeck(x)\n        skip_connections = skip_connections[::-1]\n\n        for idx in range(0,len(self.ups),2):\n            x = self.ups[idx](x)\n            skip_connect = skip_connections[idx//2]\n            if x.shape != skip_connect.shape: #to manage shape mismatch while forward pass, even sized dimension \n                x = F.resize(x, size=skip_connect.shape[2:])\n            concat_skip = torch.concat((skip_connect,x),dim=1)\n            x = self.ups[idx+1](concat_skip)\n        \n        return self.FinalLayer(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:13.393586Z","iopub.execute_input":"2025-09-27T20:12:13.393952Z","iopub.status.idle":"2025-09-27T20:12:14.721029Z","shell.execute_reply.started":"2025-09-27T20:12:13.393931Z","shell.execute_reply":"2025-09-27T20:12:14.720433Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Loading ","metadata":{}},{"cell_type":"code","source":"# we need train_imgs, train_mask, val_imgs, val_mask and finallly some test cases to test our models\nfor things in os.listdir('/kaggle/input/carvana-image-masking-challenge'):\n    print(things)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:15.782433Z","iopub.execute_input":"2025-09-27T20:12:15.783266Z","iopub.status.idle":"2025-09-27T20:12:15.788040Z","shell.execute_reply.started":"2025-09-27T20:12:15.783243Z","shell.execute_reply":"2025-09-27T20:12:15.787274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"folders = ['train_masks.zip','train.zip','test.zip']\ndestination = ['train_mask','train','test']\ndef UnZipFile(path,dest):\n    with zipfile.ZipFile(path,'r') as ref:\n        ref.extractall(dest)\n\nroot = r'/kaggle/input/carvana-image-masking-challenge'\n\nfor path,dest in zip(folders,destination):\n    print(f'=> Extracting {path}....')\n    path = os.path.join(root,path)\n    os.makedirs(dest,exist_ok=True)\n    UnZipFile(path,dest)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n\n#make validation set out of the training images, let's use randomized 90/10 split \n# which roughly means about 600 images for the \nimport random\nrandom.seed(42)\n#make list of img_paths and corresponding mask_paths\ntrain_dir = r'/kaggle/working/train/train'\ntrain_mask_dir = r'/kaggle/working/train_mask/train_masks'\ntest_dir = r'/kaggle/working/test/test'\n\nimages = sorted([p for p in os.listdir(train_dir)])\nmasks = sorted([p for p in os.listdir(train_mask_dir)])\ntesting = [p for p in os.listdir(test_dir)][:100]\nrandom.shuffle(images)\n\ntrain_img = images[:int(0.9*len(images))]\nval_img = images[int(0.9*len(images)):]\nprint(f'=> # of Training Images: {len(train_img)}\\n=> # of Validation Images: {len(val_img)}')\n\ntrain_masks = [m.replace(\".jpg\",\"_mask.gif\") for m in train_img]\nval_masks = [m.replace(\".jpg\",\"_mask.gif\") for m in val_img]\nprint(f'=> # of Training Masks: {len(train_masks)}\\n=> # of Validation Masks: {len(val_masks)}')\n\nprint(f'=> # of Testing Images: {len(testing)}')\n\n\n#checking if all the masks are present or not \n\nfor img,mask in zip(train_img+val_img,train_masks+val_masks):\n    img = img.split(\".\")[0]\n    mask = mask.split(\".\")[0].replace(\"_mask\",\"\")\n    if img != mask:\n        print(f'\\n{mask} not found !')\n        break\nprint(f'\\n=> All masks matched!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:18.717993Z","iopub.execute_input":"2025-09-27T20:12:18.718281Z","iopub.status.idle":"2025-09-27T20:12:18.793611Z","shell.execute_reply.started":"2025-09-27T20:12:18.718257Z","shell.execute_reply":"2025-09-27T20:12:18.792790Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from PIL import Image\nclass CarvanaDataset(Dataset):\n    \n    def __init__(self,img_root:str,mask_root:str,data:tuple,transform):\n        self.img_root = img_root\n        self.mask_root = mask_root\n        self.data = data\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data[0])\n\n    def __getitem__(self,idx:int) -> torch.tensor:\n        img_path = os.path.join(self.img_root,self.data[0][idx])\n        mask_path = os.path.join(self.mask_root,self.data[1][idx])\n        \n        img = np.array(Image.open(img_path).convert(\"RGB\"))\n        mask = np.array(Image.open(mask_path).convert(\"L\"),dtype=np.float32)\n        mask[mask == 255.0] = 1.0\n        \n        if self.transform:\n            augmentation = self.transform(image=img,mask=mask)\n            img = augmentation['image']\n            mask = augmentation['mask']\n        return img,mask\n\ndef get_loaders(directories,train_data,val_data,transforms=None,workers=2,pin=True,batch_size=64):\n    loaders = []\n    data = [train_data,val_data]\n    img_root,mask_root = directories[0],directories[1]\n    for idx in range(2):\n        dataset = CarvanaDataset(img_root,mask_root,data[idx],transform=transforms[idx])\n        loader = DataLoader(dataset,batch_size=batch_size,shuffle=True,num_workers=workers,pin_memory=pin)\n        loaders.append(loader)\n    return loaders\ndef get_loaders_med(lookup_table,data,transforms=None,workers=2,pin=True,batch_size=64):\n    loaders=[]\n    for idx in range(len(data)):\n        dataset = MedDataSet(data[idx],lookup_table,transform=(transforms[idx] if transforms else None))\n        loader= DataLoader(dataset,batch_size=batch_size,shuffle=True,num_workers=workers,pin_memory=pin)\n        loaders.append(loader)\n    return loaders\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:19.442644Z","iopub.execute_input":"2025-09-27T20:12:19.443377Z","iopub.status.idle":"2025-09-27T20:12:19.452218Z","shell.execute_reply.started":"2025-09-27T20:12:19.443347Z","shell.execute_reply":"2025-09-27T20:12:19.451539Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Procedure ","metadata":{}},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm.notebook import tqdm,trange\n\n#hyperparams \nlr = 3e-4\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nbatch_size=8\nnum_epochs = 50\n\n# CARVANA TASK\nIMAGE_HEIGTH= 320\nIMAGE_WIDTH = 480\n\n# ULTRASOUND TASK\n# IMAGE_HEIGTH= 420\n# IMAGE_WIDTH = 580\n\nPIN_MEMORY = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:21.221334Z","iopub.execute_input":"2025-09-27T20:12:21.221671Z","iopub.status.idle":"2025-09-27T20:12:22.128096Z","shell.execute_reply.started":"2025-09-27T20:12:21.221646Z","shell.execute_reply":"2025-09-27T20:12:22.127281Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#this function represents 1 Epoch of training \ndef train_this_model(loader, model, optimizer, loss_fn,scaler):\n    loop = tqdm(loader)\n    epoch_everage_loss = 0.0\n    numBatch=0\n    for idx,(data,targets) in enumerate(loop):\n        numBatch+=1\n        data = data.to(DEVICE)\n        targets = targets.float().unsqueeze(1).to(DEVICE)\n\n        #FORWARD PASS\n        with torch.amp.autocast('cuda'):\n            predictions = model(data)\n            loss = loss_fn(predictions,targets)\n            \n        #BACKWARD PASS\n        optimizer.zero_grad()\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        loop.set_postfix(loss = loss.item())\n        epoch_everage_loss += loss.item()\n    epoch_everage_loss /= numBatch\n    return epoch_everage_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:22.798237Z","iopub.execute_input":"2025-09-27T20:12:22.798671Z","iopub.status.idle":"2025-09-27T20:12:22.804095Z","shell.execute_reply.started":"2025-09-27T20:12:22.798646Z","shell.execute_reply":"2025-09-27T20:12:22.803416Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Defining Transforms","metadata":{}},{"cell_type":"code","source":"train_transform = A.Compose(\n    [\n        A.Resize(height=IMAGE_HEIGTH, width = IMAGE_WIDTH),\n        # A.ShiftScaleRotate(shift_limit=0.05,scale_limit=0.05,rotate_limit=15,p=0.5),\n        A.Rotate(limit=40,p=0.5),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomBrightnessContrast(p=0.2),\n        A.ColorJitter(p=0.2),\n        A.Normalize(\n            mean = [0.0,0.0,0.0],\n            std = [1.0,1.0,1.0],\n            max_pixel_value = 255.0\n        ),\n        ToTensorV2(),\n    ],\n)\n\nval_transform = A.Compose(\n    [\n        A.Resize(height=IMAGE_HEIGTH, width = IMAGE_WIDTH),\n        A.Normalize(\n            mean = [0.0,0.0,0.0],\n            std = [1.0,1.0,1.0],\n            max_pixel_value = 255.0\n        ),\n        ToTensorV2(),\n    ]\n)\n\n\n\nwith open('transforms.py','w') as f:\n    f.write('''\nimport albumentations as A\nimport numpy as np\nimport torch\nfrom albumentations.pytorch import ToTensorV2\nIMAGE_HEIGTH=160\nIMAGE_WIDTH = 240\ntrain_transform = A.Compose(\n    [\n        A.Resize(height=IMAGE_HEIGTH, width = IMAGE_WIDTH),\n        A.Rotate(limit=30,p=1.0),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.1),\n        A.Normalize(\n            mean = [0.0,0.0,0.0],\n            std = [1.0,1.0,1.0],\n            max_pixel_value = 255.0\n        ),\n        ToTensorV2(),\n    ],\n)\n\nval_transform = A.Compose(\n    [\n        A.Resize(height=IMAGE_HEIGTH, width = IMAGE_WIDTH),\n        A.Normalize(\n            mean = [0.0,0.0,0.0],\n            std = [1.0,1.0,1.0],\n            max_pixel_value = 255.0\n        ),\n        ToTensorV2(),\n    ]\n)\n\nif __name__ == '__main__':\n    x = np.random.randn(3,1980,2560)\n    print(f'=> Shape of original tensor: {x.shape}')\n    print(f'=> Train Transformed: {train_transform(image=x)[\"image\"].shape}')\n    print(f'=> Validation/Test Transformed: {val_transform(image=x)[\"image\"].shape}')\n    ''')\n\n# %run transforms.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:23.934889Z","iopub.execute_input":"2025-09-27T20:12:23.935251Z","iopub.status.idle":"2025-09-27T20:12:23.947134Z","shell.execute_reply.started":"2025-09-27T20:12:23.935230Z","shell.execute_reply":"2025-09-27T20:12:23.946447Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Utility Functions ","metadata":{}},{"cell_type":"code","source":"def save_model_checkpoint(state,filename='checkpoint_model.pth.tar'):\n    print(f'=> saving checkpoint')\n    torch.save(state,filename)\n    \ndef load_model_checkpoint(checkpoint,model):\n    print(f'=> Loading checkpoint')\n    model.load_state_dict(checkpoint[\"state_dict\"])\n\ndef check_accuracy_binary(loader,model,device):\n    numCorrect = 0\n    numPixels = 0\n    model.eval()\n    dice_score = 0\n    val_loss = 0.0\n    numBatch=0\n    with torch.no_grad():\n        for x,y in loader:\n            numBatch+=1\n            x = x.to(device)\n            y = y.to(device).unsqueeze(1)\n            preds = model(x)\n            loss = loss_fn(preds,y)\n            preds = torch.sigmoid(preds)\n            preds = (preds > 0.5).float()\n            numCorrect += (preds == y).sum()\n            numPixels += torch.numel(preds)\n            intersection = (preds * y).sum()\n            dice_score += (2. * intersection) / (preds.sum() + y.sum() + 1e-8)\n            val_loss += loss.item()\n    model.train()\n    acc = (numCorrect/numPixels)\n    val_loss /=numBatch\n    print(\n        f'Got {numCorrect}/{numPixels} with acc {acc*100:.2f}'\n    )\n    print(f'Dice Score: {dice_score/len(loader):.2f}')\n    return acc,val_loss\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:25.922201Z","iopub.execute_input":"2025-09-27T20:12:25.922852Z","iopub.status.idle":"2025-09-27T20:12:25.929301Z","shell.execute_reply.started":"2025-09-27T20:12:25.922826Z","shell.execute_reply":"2025-09-27T20:12:25.928488Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Formal Training ","metadata":{}},{"cell_type":"code","source":"batch_size","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:27.946399Z","iopub.execute_input":"2025-09-27T20:12:27.946956Z","iopub.status.idle":"2025-09-27T20:12:27.952107Z","shell.execute_reply.started":"2025-09-27T20:12:27.946934Z","shell.execute_reply":"2025-09-27T20:12:27.951396Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = UNET(p=0.4).to(DEVICE)\nloss_fn = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(),lr=lr)\nscaler = torch.amp.GradScaler('cuda')\n\n\ntrain_loader,val_loader = get_loaders(\n    (train_dir,train_mask_dir),\n    (train_img,train_masks),\n    (val_img,val_masks),\n    (train_transform,val_transform),\n    batch_size=batch_size,\n    pin = True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:13:41.217255Z","iopub.execute_input":"2025-09-27T20:13:41.217902Z","iopub.status.idle":"2025-09-27T20:13:41.490338Z","shell.execute_reply.started":"2025-09-27T20:13:41.217877Z","shell.execute_reply":"2025-09-27T20:13:41.489652Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_acc = float('-inf')\npatience = 5\ntrigger = 0\navg_training_loss = []\navg_val_acc=[]\navg_val_loss = []\nfor epochs in trange(num_epochs):\n    torch.cuda.empty_cache()\n    train_loss = train_this_model(train_loader,model,optimizer,loss_fn,scaler)\n    print(f'Epoch: {epochs+1}, Avg Training Loss: {train_loss}')\n    avg_training_loss.append(train_loss)\n    acc,val_loss = check_accuracy_binary(val_loader,model,device=DEVICE)\n    avg_val_acc.append(acc)\n    avg_val_loss.append(val_loss)\n    if acc > best_acc:\n        best_acc = acc\n        checkpoint = {\n            \"state_dict\":model.state_dict(),\n            \"optimzer\":optimizer.state_dict()\n        }\n        save_model_checkpoint(checkpoint)\n        trigger = 0\n    else:\n        trigger +=1\n    if trigger > patience:\n        print(f'=> EARLY STOPPING')\n        break ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_acc = [p.item() for p in avg_val_acc]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#plotting different model configuration graphs \nimport matplotlib.pyplot as plt\n\nfig,ax = plt.subplots(3,1,figsize=(10,8))\nax = ax.flatten()\ntitles = ['Avg Training Loss','Avg Val Accuracy','Avg Val Loss']\ny = [avg_training_loss,val_acc,avg_val_loss]\nx = np.arange(1,26+2,1)\nfor idx,(title,data) in enumerate(zip(titles,y)):\n    ax[idx].plot(x,data)\n    ax[idx].set_title(f'{title} v/s Epoch')\n    ax[idx].set_xlabel('Epochs')\n    ax[idx].set_ylabel(title)\n    ax[idx].grid()\nplt.tight_layout()\n\n#I forgot to run this cell !","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loading the Best Model and Testing it","metadata":{}},{"cell_type":"code","source":"my_model = UNET(p=0.4).to(DEVICE) #from the model definition \ncheckpoint_path = r'/kaggle/working/checkpoint_model.pth.tar'\ncheckpoint_dict = torch.load(checkpoint_path,map_location=DEVICE)\nload_model_checkpoint(checkpoint_dict,my_model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:31.925432Z","iopub.execute_input":"2025-09-27T20:12:31.925721Z","iopub.status.idle":"2025-09-27T20:12:32.669039Z","shell.execute_reply.started":"2025-09-27T20:12:31.925699Z","shell.execute_reply":"2025-09-27T20:12:32.668310Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#testing the model on test images \ntest_root = r'/kaggle/working/test/test'\ntest_img = [test_root + \"/\"+p  for p in testing]\n\ndef img_to_tensor(img_path,transform=None):\n    img = np.array(Image.open(img_path).convert('RGB'))\n    if transform:\n        img = transform(image=img)\n        img = img['image'].unsqueeze(0)\n    return img \n\ndef show_image(img_path,axis=None,transform=None):\n    img = np.array(Image.open(img_path).convert(\"RGB\"))\n    if transform:\n        img = transform(image=img)\n        img = img['image'].permute(1,2,0).numpy()\n    if axis:\n        axis.imshow(img)\n        axis.axis('off')\n    else:\n        plt.imshow(img)\n        plt.axis('off')\n    return img.shape\n\ndef get_mask(img_path,model,transform=val_transform,device='cpu'):\n    model.to(device)\n    x = img_to_tensor(img_path,transform).to(device)\n    y = torch.sigmoid(model(x).detach().cpu())\n    y = (y>0.5).float()\n    y = y.squeeze(0).permute(1,2,0)\n    return y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:33.401755Z","iopub.execute_input":"2025-09-27T20:12:33.402085Z","iopub.status.idle":"2025-09-27T20:12:33.409205Z","shell.execute_reply.started":"2025-09-27T20:12:33.402060Z","shell.execute_reply":"2025-09-27T20:12:33.408415Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visualising some test Images","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nn = 2\nfig,ax = plt.subplots(n,2,figsize=(12,12))\ntransformations = [None,val_transform]\ntitle = ['Original','Transformed']\nfor i in range(n):\n    idx = np.random.randint(0,100)\n    for j in range(2):\n        shape = show_image(test_img[idx],ax[i,j],transformations[j])\n        ax[i,j].set_title(f'{title[j]}, Shape:{shape}')\nplt.tight_layout()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:34.865965Z","iopub.execute_input":"2025-09-27T20:12:34.866246Z","iopub.status.idle":"2025-09-27T20:12:36.285226Z","shell.execute_reply.started":"2025-09-27T20:12:34.866222Z","shell.execute_reply":"2025-09-27T20:12:36.284450Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## MODEL Output","metadata":{}},{"cell_type":"code","source":"nimages = 5\nfig,ax = plt.subplots(nimages,3,figsize=(12,12))\nfor i in range(nimages):\n    idx = np.random.randint(0,len(test_img))\n    image = test_img[idx]\n    show_image(image,ax[i,0],None)\n    ax[i,0].set_title(f'Original Image')\n    ax[i,0].axis('off')\n    show_image(image,ax[i,1],val_transform) #shows the image\n    ax[i,1].set_title('Transformed Image(downsampled to resize)')\n    ax[i,1].axis('off')\n    mask = get_mask(image,my_model,val_transform,DEVICE)\n    ax[i,2].imshow(mask,cmap='gray')\n    ax[i,2].set_title('Predicted Mask')\n    ax[i,2].axis('off')\nplt.tight_layout()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:36.286317Z","iopub.execute_input":"2025-09-27T20:12:36.286534Z","iopub.status.idle":"2025-09-27T20:12:40.083342Z","shell.execute_reply.started":"2025-09-27T20:12:36.286517Z","shell.execute_reply":"2025-09-27T20:12:40.082418Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Transfer Learning \n- here we aim to reuse our models weights and use the same model for other Binary Segmentation task\n  - but this task can't be done directly, as their is a fundamental domain difference between two tasks\n  - so we will employ a **'TEACHER-STUDENT' Learner strategy**\n  - we'll load our Pretrained UNET weights and fine tune with extremely low lr on the train set of nerve ultrasounds, then run all the test images through it and produce some prediction masks\n  - then we'll filter the best masks and add it in the train set and fine tune another final model, this would be our teacher learned which will be more robust\n  - ofcourse we'll keep aside a small subset of images to test our final teacher (around 1% of the total train data)\n- the second task comes from another [kaggle competition](https://www.kaggle.com/c/ultrasound-nerve-segmentation)\n- we'll begin my importing and separating our data\n","metadata":{}},{"cell_type":"code","source":"import os\nSEED = 42\nultrasound_train_path = r'/kaggle/input/ultrasound-nerve-segmentation/train'\nultrasound_test_path = r'/kaggle/input/ultrasound-nerve-segmentation/test'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:46.113187Z","iopub.execute_input":"2025-09-27T20:12:46.113968Z","iopub.status.idle":"2025-09-27T20:12:46.117608Z","shell.execute_reply.started":"2025-09-27T20:12:46.113943Z","shell.execute_reply":"2025-09-27T20:12:46.116934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_images_ultrasound = []\nfor idx,things in enumerate(os.listdir(ultrasound_train_path)):\n    if things.endswith('_mask.tif'):\n        continue\n    else:\n        train_images_ultrasound.append(things)\nprint(f'# of Train Images: {len(train_images_ultrasound)}')\nprint(f'First 5 examples: {train_images_ultrasound[:5]}')\n\nlookup_ultrasound = {} ## this dictionary stores the mask path as the key \n\nfor images in train_images_ultrasound:\n    lookup_ultrasound[ultrasound_train_path+\"/\"+images] = ultrasound_train_path + \"/\"+images.replace(\".tif\",\"_mask.tif\")\nprint(f'First Five examples from the lookup dictionary')\nprint(list(lookup_ultrasound.items())[:5])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:47.274333Z","iopub.execute_input":"2025-09-27T20:12:47.275001Z","iopub.status.idle":"2025-09-27T20:12:47.293456Z","shell.execute_reply.started":"2025-09-27T20:12:47.274976Z","shell.execute_reply":"2025-09-27T20:12:47.292812Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#making splits on the data train,val and hold set \n# randomly holding out some images \nn_hold = int(0.01*len(train_images_ultrasound))\nhold_idx = np.random.randint(0,len(train_images_ultrasound),n_hold) #indices of val set \nnp.random.seed(SEED)\ntrain_ultrasound = [ultrasound_train_path+\"/\"+img for idx,img in enumerate(train_images_ultrasound) if idx not in hold_idx]\nval_ultrasound = [ultrasound_train_path+\"/\"+img for idx,img in enumerate(train_images_ultrasound) if idx in hold_idx]\n\nprint(f'=> # of Train Images: {len(train_ultrasound)}')\nprint(f'=> # of Validation Images: {len(val_ultrasound)}')\n\nfor t in train_ultrasound:\n    if t in val_ultrasound:\n        print(\"ERROR\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:12:48.858180Z","iopub.execute_input":"2025-09-27T20:12:48.858512Z","iopub.status.idle":"2025-09-27T20:12:48.913297Z","shell.execute_reply.started":"2025-09-27T20:12:48.858487Z","shell.execute_reply":"2025-09-27T20:12:48.912651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MedDataSet(Dataset):\n    def __init__(self,list_of_images,lookup_table,transform=None):\n        self.list_of_images = list_of_images\n        self.lookup_table = lookup_table\n        self.transform = transform\n    def __len__(self):\n        return len(self.list_of_images)\n    def __getitem__(self,idx):\n        img_path = self.list_of_images[idx]\n        mask_path = self.lookup_table[img_path]\n        img = np.array(Image.open(img_path).convert(\"RGB\"))\n        mask = np.array(Image.open(mask_path).convert(\"L\"))\n        mask[mask == 255.0] = 1.0\n        if self.transform:\n            transformation = self.transform(image = img,mask = mask)\n            img = transformation['image']\n            mask = transformation['mask']\n        return img,mask\n\nn_train_transform = A.Compose(\n    [\n        A.Resize(height=336, width = 464),\n        # A.ShiftScaleRotate(shift_limit=0.05,scale_limit=0.05,rotate_limit=15,p=0.5),\n        A.Rotate(limit=40,p=0.5),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomBrightnessContrast(p=0.2),\n        A.ColorJitter(p=0.2),\n        A.Normalize(\n            mean = [0.0,0.0,0.0],\n            std = [1.0,1.0,1.0],\n            max_pixel_value = 255.0\n        ),\n        ToTensorV2(),\n    ],\n)\n\nn_val_transform = A.Compose(\n    [\n        A.Resize(height=336, width = 464),\n        A.Normalize(\n            mean = [0.0,0.0,0.0],\n            std = [1.0,1.0,1.0],\n            max_pixel_value = 255.0\n        ),\n        ToTensorV2(),\n    ]\n)\n\n\ntrain_loader,val_loader = get_loaders_med(\n    lookup_ultrasound,\n    data=[train_ultrasound,val_ultrasound],\n    transforms=[n_train_transform,n_val_transform],\n    batch_size=8\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:13:21.222507Z","iopub.execute_input":"2025-09-27T20:13:21.222790Z","iopub.status.idle":"2025-09-27T20:13:21.236429Z","shell.execute_reply.started":"2025-09-27T20:13:21.222770Z","shell.execute_reply":"2025-09-27T20:13:21.235758Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_rows = 3\nfig,axis = plt.subplots(n_rows,2,figsize=(6,6)) \nfor x,y in train_loader:\n    idx = np.random.randint(0,len(x),n_rows)\n    for i in range(n_rows):\n        img = x[idx[i]]\n        mask = y[idx[i]]\n        axis[i,0].imshow(img.permute(1,2,0))\n        axis[i,0].axis('off')\n        axis[i,1].imshow(mask,cmap='gray')\n        axis[i,1].axis('off')\n        print(img.shape)\n        print(mask.shape)\n    break\nplt.tight_layout()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:13:22.615968Z","iopub.execute_input":"2025-09-27T20:13:22.616259Z","iopub.status.idle":"2025-09-27T20:13:23.434653Z","shell.execute_reply.started":"2025-09-27T20:13:22.616236Z","shell.execute_reply":"2025-09-27T20:13:23.433808Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_model = UNET(p=0.4).to(DEVICE) #from the model definition \ncheckpoint_path = r'/kaggle/working/checkpoint_model.pth.tar'\ncheckpoint_dict = torch.load(checkpoint_path,map_location=DEVICE)\nload_model_checkpoint(checkpoint_dict,my_model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:13:23.532459Z","iopub.execute_input":"2025-09-27T20:13:23.532756Z","iopub.status.idle":"2025-09-27T20:13:24.184209Z","shell.execute_reply.started":"2025-09-27T20:13:23.532729Z","shell.execute_reply":"2025-09-27T20:13:24.183558Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_LR = 1e-5\nn_optimizer = torch.optim.AdamW(n_model.parameters(),lr=n_LR)\nn_scaler = torch.amp.GradScaler('cuda')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:13:25.378990Z","iopub.execute_input":"2025-09-27T20:13:25.379749Z","iopub.status.idle":"2025-09-27T20:13:25.384314Z","shell.execute_reply.started":"2025-09-27T20:13:25.379723Z","shell.execute_reply":"2025-09-27T20:13:25.383594Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_acc = float('-inf')\npatience = 5\ntrigger = 0\nn_avg_training_loss = []\nn_avg_val_acc=[]\nn_avg_val_loss = []\nfig,ax = plt.subplots(nimages,3,figsize=(12,12))\nfor epochs in trange(5):\n    torch.cuda.empty_cache()\n    train_loss = train_this_model(train_loader,n_model,n_optimizer,loss_fn,n_scaler)\n    print(f'Epoch: {epochs+1}, Avg Training Loss: {train_loss}')\n    n_avg_training_loss.append(train_loss)\n    nimages = 2\n    for i in range(nimages):\n        idx = np.random.randint(0,len(train_ultrasound))\n        image = train_ultrasound[idx]\n        show_image(image,ax[i,0],None)\n        ax[i,0].set_title(f'Original Image')\n        ax[i,0].axis('off')\n        show_image(image,ax[i,1],n_val_transform) #shows the image\n        ax[i,1].set_title('Transformed Image(downsampled to resize)')\n        ax[i,1].axis('off')\n        mask = get_mask(image,n_model,n_val_transform,DEVICE)\n        ax[i,2].imshow(mask,cmap='gray')\n        ax[i,2].set_title('Predicted Mask')\n        ax[i,2].axis('off')\n        # break\n    plt.tight_layout()\n    plt.show()\n    \n    # acc,val_loss = check_accuracy_binary(val_loader,n_model,device=DEVICE)\n    # n_avg_val_acc.append(acc)\n    # n_avg_val_loss.append(val_loss)\n    # if acc > best_acc:\n    #     best_acc = acc\n    #     checkpoint = {\n    #         \"state_dict\":n_model.state_dict(),\n    #         \"optimzer\":n_optimizer.state_dict()\n    #     }\n    #     save_model_checkpoint(checkpoint,filename=\"fine_tuned.pth.tar\")\n    #     trigger = 0\n    # else:\n    #     trigger +=1\n    # if trigger > patience:\n    #     print(f'=> EARLY STOPPING')\n    #     break ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T20:19:18.940827Z","iopub.execute_input":"2025-09-27T20:19:18.941144Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}