{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":8776139,"sourceType":"datasetVersion","datasetId":5274724},{"sourceId":96721,"sourceType":"modelInstanceVersion","modelInstanceId":81126,"modelId":105493}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 2d Segmentation of Sagittal Lumbar Spine MRI\n\n- Training data Spider dataset (https://doi.org/10.5281/zenodo.10159290)\n- Used the notebook https://www.kaggle.com/code/anoukstein/2d-segmentation-of-sagittal-lumbar-spine-mri by author Anouk Stein, MD\n- Tweaked loss function and calculated the dice score for each class\n- The trianed model will be later used for segmenting the vertebrae's to calculate intervertebrae disc centroids and crop them, so as to prepare a custom dataset for the RSNA2024 competition","metadata":{}},{"cell_type":"code","source":"!pip install segmentation_models_pytorch -q\n","metadata":{"execution":{"iopub.status.busy":"2024-11-03T16:37:00.125354Z","iopub.execute_input":"2024-11-03T16:37:00.12605Z","iopub.status.idle":"2024-11-03T16:37:12.615436Z","shell.execute_reply.started":"2024-11-03T16:37:00.12601Z","shell.execute_reply":"2024-11-03T16:37:12.614138Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport os\nfrom pathlib import Path\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import KFold\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset\n\nfrom segmentation_models_pytorch import Unet","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-11-03T16:37:12.617598Z","iopub.execute_input":"2024-11-03T16:37:12.617925Z","iopub.status.idle":"2024-11-03T16:37:12.623787Z","shell.execute_reply.started":"2024-11-03T16:37:12.617896Z","shell.execute_reply":"2024-11-03T16:37:12.622898Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#transforms\nnewsize = (256, 256)\n#dataset\nfold = 1\n#dataloader\nbatch_size = 64\nnum_workers = 4\n#model\nnum_classes = 6\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")#run\nepochs = 50\nlearning_rate = 1e-3\n\nTRAIN = True #or  False for inference only","metadata":{"execution":{"iopub.status.busy":"2024-11-03T16:37:15.649476Z","iopub.execute_input":"2024-11-03T16:37:15.650344Z","iopub.status.idle":"2024-11-03T16:37:15.711912Z","shell.execute_reply.started":"2024-11-03T16:37:15.650309Z","shell.execute_reply":"2024-11-03T16:37:15.710916Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Model","metadata":{}},{"cell_type":"markdown","source":"### Create folds","metadata":{}},{"cell_type":"code","source":"output_dir = \"/kaggle/input/spider-mri-spine-t2-png/data\"\nim_dir = os.path.join(output_dir, \"images\")\nmask_dir = os.path.join(output_dir, \"masks\")\n\n# get list of data\nitems = list(Path(im_dir).glob(\"*.png\"))\nimage_names = [o.name for o in items]\nimages = list(set([o.split('_')[0] for o in image_names]))\n\nfold_df = pd.DataFrame({\"image_name\": images})\n# Seed for reproducibility\nnp.random.seed(42)\n\n# Split the DataFrame into 5 folds\nkf = KFold(n_splits=5, shuffle=True, random_state=42)\nfor i, (_, v_ind) in enumerate(kf.split(fold_df)):\n    fold_df.loc[v_ind, 'fold'] = i+1\n\n# Create df with image_names and their respective folds\ndef get_fold(fn, df):\n    image_name = fn.name.split(\"_\")[0] \n    return df.loc[df.image_name==image_name, 'fold'].values[0]\n\nfolds = [get_fold(o, fold_df) for o in items]\ndf = pd.DataFrame({\"image\": image_names, \"fold\": folds})\n\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-11-03T16:37:23.978469Z","iopub.execute_input":"2024-11-03T16:37:23.978839Z","iopub.status.idle":"2024-11-03T16:37:25.559055Z","shell.execute_reply.started":"2024-11-03T16:37:23.978812Z","shell.execute_reply":"2024-11-03T16:37:25.557934Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_df=df[df['fold']==5]\ntrain_df=df[df['fold']!=5]\nlen(train_df),len(test_df)","metadata":{"execution":{"iopub.status.busy":"2024-11-03T16:37:28.615834Z","iopub.execute_input":"2024-11-03T16:37:28.616774Z","iopub.status.idle":"2024-11-03T16:37:28.627175Z","shell.execute_reply.started":"2024-11-03T16:37:28.616737Z","shell.execute_reply":"2024-11-03T16:37:28.626243Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Dataset class","metadata":{}},{"cell_type":"code","source":"class SEGDataset(Dataset):\n    def __init__(self, df, mode, transforms=None):\n        self.df = df.reset_index()\n        self.mode = mode\n        self.transforms = transforms\n\n    def __len__(self):\n        return self.df.shape[0]\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n\n        image_path = os.path.join(im_dir, row.image)\n        mask_path = os.path.join(mask_dir, row.image)\n\n        # Open image\n        image = Image.open(image_path).convert('L')\n        image = np.asarray(image)\n        if (image > 1).any():  # Normalize if pixel values are between 0-255\n            image = image / 255.0\n\n        # Open mask\n        mask = Image.open(mask_path)\n        mask = np.asarray(mask)\n        mask= np.where(mask<=5,mask,0) #considering only L5,L4,L3,L2,L1 vertebrae's\n        assert mask.max() < num_classes, f\"Mask value {mask.max()} exceeds number of classes {num_classes}\"\n\n        # Apply transformations\n        if self.transforms is not None:\n            transformed = self.transforms(image=image, mask=mask)\n            image = transformed[\"image\"]\n            mask = transformed[\"mask\"]\n        \n        # Create one layer for each label\n        mask = torch.as_tensor(mask).long()\n        mask = torch.nn.functional.one_hot(mask, num_classes=num_classes).permute(2,0,1).float()\n        #mask = torch.nn.functional.one_hot(mask, num_classes=num_classes).permute(0,3,1,2).squeeze(0).float()\n\n        # Convert image to tensor\n        image = torch.as_tensor(image).float()\n\n        return image, mask          ","metadata":{"execution":{"iopub.status.busy":"2024-11-03T16:37:30.880467Z","iopub.execute_input":"2024-11-03T16:37:30.880813Z","iopub.status.idle":"2024-11-03T16:37:30.891935Z","shell.execute_reply.started":"2024-11-03T16:37:30.880789Z","shell.execute_reply":"2024-11-03T16:37:30.890992Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Transforms","metadata":{}},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\ntransforms_train = A.Compose([\n    A.Resize(newsize[0], newsize[1]),\n    A.HorizontalFlip(),\n    A.Normalize(\n        mean=[0.485],\n        std=[0.229],\n    ),\n    ToTensorV2()\n])\n\ntransforms_valid = A.Compose([\n    A.Resize(newsize[0], newsize[1]),\n    A.Normalize(\n        mean=[0.485],\n        std=[0.229],\n    ),\n    ToTensorV2()\n])","metadata":{"execution":{"iopub.status.busy":"2024-11-03T16:37:40.368024Z","iopub.execute_input":"2024-11-03T16:37:40.368411Z","iopub.status.idle":"2024-11-03T16:37:40.858705Z","shell.execute_reply.started":"2024-11-03T16:37:40.368382Z","shell.execute_reply":"2024-11-03T16:37:40.857752Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Loss\nA combined loss function with cross-entropy loss and jaccard loss is created ","metadata":{}},{"cell_type":"code","source":"class CombinedLoss(nn.Module):\n    def __init__(self, weight_ce=0.4, weight_iou=0.6):\n        super(CombinedLoss, self).__init__()\n        self.weight_ce = weight_ce\n        self.weight_iou = weight_iou\n        self.cross_entropy_loss = nn.CrossEntropyLoss(reduction='mean')\n\n    def forward(self, inputs, targets):\n        # Cross-Entropy Loss\n        ce_loss = self.cross_entropy_loss(inputs, targets.argmax(dim=1))\n\n        # IoU Loss\n        probs = F.softmax(inputs, dim=1)\n        intersection = torch.sum(probs * targets, dim=(2, 3))\n        union = torch.sum(probs, dim=(2, 3)) + torch.sum(targets, dim=(2, 3)) - intersection\n        iou = (intersection + 1e-6) / (union + 1e-6)\n        iou_loss = 1 - iou.mean(dim=1)  # Average over classes, resulting in (batch_size,)\n\n        # Combine losses\n        loss = self.weight_ce * ce_loss + self.weight_iou * iou_loss.mean()\n        return loss","metadata":{"execution":{"iopub.status.busy":"2024-11-03T16:37:51.948117Z","iopub.execute_input":"2024-11-03T16:37:51.949445Z","iopub.status.idle":"2024-11-03T16:37:51.957177Z","shell.execute_reply.started":"2024-11-03T16:37:51.949409Z","shell.execute_reply":"2024-11-03T16:37:51.956306Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Create datasets and dataloaders","metadata":{}},{"cell_type":"code","source":"train_ = train_df[train_df['fold'] != fold].reset_index(drop=True)\nvalid_ = train_df[train_df['fold'] == fold].reset_index(drop=True)\nprint(len(train_),len(valid_))\ndataset_train = SEGDataset(train_, 'train',  transforms_train)\ndataset_valid = SEGDataset(valid_, 'valid',  transforms_valid)\ndataset_test  = SEGDataset(test_df,  'valid',  transforms_valid)\n\ntrain_loader = torch.utils.data.DataLoader(dataset_train, batch_size=batch_size, shuffle=True, num_workers=num_workers)\nval_loader = torch.utils.data.DataLoader(dataset_valid, batch_size=batch_size, shuffle=False, num_workers=num_workers)\ntest_loader = torch.utils.data.DataLoader(dataset_test, batch_size=batch_size, shuffle=False, num_workers=num_workers)","metadata":{"execution":{"iopub.status.busy":"2024-11-03T16:37:55.749203Z","iopub.execute_input":"2024-11-03T16:37:55.74957Z","iopub.status.idle":"2024-11-03T16:37:55.762234Z","shell.execute_reply.started":"2024-11-03T16:37:55.749545Z","shell.execute_reply":"2024-11-03T16:37:55.761244Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Run function","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\ndef run(train_loader, val_loader, model, learning_rate, criterion, epochs, device):\n    optimizer = optim.Adam(model.parameters(), lr=learning_rate)\n    scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=5, verbose=True)\n    best_val_loss=float('inf')\n    for epoch in range(epochs):\n        model.train()\n        train_loss = 0.0\n        with tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs}\", unit=\"batch\") as train_bar:\n            for images, masks in train_bar:\n                images, masks = images.to(device), masks.to(device)\n                outputs = model(images)\n                loss = criterion(outputs, masks)\n                optimizer.zero_grad()\n                loss.backward()\n                optimizer.step()\n                train_loss += loss.item()\n                #train_bar.set_postfix({\"Train Loss\": loss.item()})\n        train_loss /= len(train_loader)\n\n        model.eval()\n        val_loss = 0.0\n        with torch.no_grad():\n            with tqdm(val_loader, desc=\"Validation\", unit=\"batch\") as val_bar:\n                for images, masks in val_bar:\n                    images, masks = images.to(device), masks.to(device)\n                    outputs = model(images)\n                    val_loss += criterion(outputs, masks).item()\n                    \n        val_loss /= len(val_loader)\n        if best_val_loss>val_loss:\n            print('Saving model...')\n            best_val_loss=val_loss\n            torch.save(model.state_dict(), './simple_unet.pth')\n\n        scheduler.step(val_loss)\n        print(f\"Epoch: {epoch+1}/{epochs} | Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n\n    ","metadata":{"execution":{"iopub.status.busy":"2024-11-03T16:37:59.744239Z","iopub.execute_input":"2024-11-03T16:37:59.744912Z","iopub.status.idle":"2024-11-03T16:37:59.756108Z","shell.execute_reply.started":"2024-11-03T16:37:59.744881Z","shell.execute_reply":"2024-11-03T16:37:59.75536Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = Unet(\n  encoder_name=\"resnet34\",  # Choose encoder (e.g. resnet18, efficientnet-b0)\n  classes=num_classes,  # Number of output classes\n  in_channels=1  # Number of input channels (e.g. 3 for RGB)\n)","metadata":{"execution":{"iopub.status.busy":"2024-11-03T16:38:05.269529Z","iopub.execute_input":"2024-11-03T16:38:05.269902Z","iopub.status.idle":"2024-11-03T16:38:06.475852Z","shell.execute_reply.started":"2024-11-03T16:38:05.269874Z","shell.execute_reply":"2024-11-03T16:38:06.475012Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Train","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision import transforms\nimport random\n\nclass ContrastiveLoss(nn.Module):\n    def __init__(self, margin=1.0):\n        super(ContrastiveLoss, self).__init__()\n        self.margin = margin\n\n    def forward(self, output1, output2, label):\n        euclidean_distance = F.pairwise_distance(output1, output2)\n        loss_contrastive = torch.mean((1 - label) * torch.pow(euclidean_distance, 2) +\n                                      (label) * torch.pow(torch.clamp(self.margin - euclidean_distance, min=0.0), 2))\n        return loss_contrastive\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-03T17:02:24.010386Z","iopub.execute_input":"2024-11-03T17:02:24.010802Z","iopub.status.idle":"2024-11-03T17:02:24.018343Z","shell.execute_reply.started":"2024-11-03T17:02:24.010769Z","shell.execute_reply":"2024-11-03T17:02:24.017308Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from segmentation_models_pytorch.encoders import get_encoder\n\nencoder = get_encoder(\"resnet34\", in_channels=1, depth=5, weights=None)  # Modify as per chosen encoder\nencoder.to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-03T17:02:40.338242Z","iopub.execute_input":"2024-11-03T17:02:40.338641Z","iopub.status.idle":"2024-11-03T17:02:40.761077Z","shell.execute_reply.started":"2024-11-03T17:02:40.338611Z","shell.execute_reply":"2024-11-03T17:02:40.760123Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ContrastiveDataset(Dataset):\n    def __init__(self, image_paths, transform=None):\n        self.image_paths = image_paths\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, index):\n        img1 = Image.open(self.image_paths[index]).convert('L')\n        label = random.randint(0, 1)\n\n        if label == 1:\n            img2 = img1  # Positive pair (same image)\n        else:\n            img2 = Image.open(random.choice(self.image_paths)).convert('L')  # Negative pair (different image)\n\n        if self.transform:\n            img1 = self.transform(img1)\n            img2 = self.transform(img2)\n\n        return img1, img2, torch.tensor([label], dtype=torch.float32)\n\ntransform = transforms.Compose([\n    transforms.Resize(newsize),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485], std=[0.229])\n])\n\ncontrastive_dataset = ContrastiveDataset(items, transform=transform)\ncontrastive_loader = torch.utils.data.DataLoader(contrastive_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-03T17:03:03.945027Z","iopub.execute_input":"2024-11-03T17:03:03.945429Z","iopub.status.idle":"2024-11-03T17:03:03.95587Z","shell.execute_reply.started":"2024-11-03T17:03:03.945397Z","shell.execute_reply":"2024-11-03T17:03:03.954862Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def pretrain_contrastive(encoder, contrastive_loader, epochs, device):\n    criterion = ContrastiveLoss()\n    optimizer = torch.optim.Adam(encoder.parameters(), lr=learning_rate)\n\n    encoder.train()\n    for epoch in range(epochs):\n        total_loss = 0.0\n        for img1, img2, label in contrastive_loader:\n            img1, img2, label = img1.to(device), img2.to(device), label.to(device)\n            \n            # Get the final feature representations\n            out1 = encoder(img1)[-1]  # Select last layer's output\n            out2 = encoder(img2)[-1]\n\n            # Ensure out1 and out2 are tensors\n            if isinstance(out1, list):\n                out1 = out1[-1]  # Use the deepest feature map if still a list\n            if isinstance(out2, list):\n                out2 = out2[-1]\n\n            # Compute contrastive loss\n            loss = criterion(out1, out2, label)\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n            total_loss += loss.item()\n\n        avg_loss = total_loss / len(contrastive_loader)\n        print(f\"Epoch {epoch+1}/{epochs}, Loss: {avg_loss:.4f}\")\n\npretrain_contrastive(encoder, contrastive_loader, epochs=10, device=device)\ntorch.save(encoder.state_dict(), 'contrastive_pretrained_encoder.pth')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-03T17:05:37.245499Z","iopub.execute_input":"2024-11-03T17:05:37.24588Z","iopub.status.idle":"2024-11-03T17:05:40.157004Z","shell.execute_reply.started":"2024-11-03T17:05:37.245846Z","shell.execute_reply":"2024-11-03T17:05:40.155426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = CombinedLoss()\nmodel.to(device)\n\nif TRAIN:\n    run(train_loader, val_loader, model, learning_rate, criterion, epochs, device)\nelse:\n    model.load_state_dict(torch.load(\"/kaggle/input/lumbar-vertebrae-segmentation/pytorch/default/1/simple_unet.pth\"))\n                      ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-03T17:07:33.052833Z","iopub.execute_input":"2024-11-03T17:07:33.053189Z","iopub.status.idle":"2024-11-03T17:07:35.970544Z","shell.execute_reply.started":"2024-11-03T17:07:33.053158Z","shell.execute_reply":"2024-11-03T17:07:35.969374Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Inference","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef inference(model, dataloader, device, num_samples=16):\n    model.eval()\n    images_batch = []\n    preds_batch = []\n    \n    with torch.no_grad():\n        for images, _ in dataloader:\n            images = images.to(device)\n            outputs = model(images)\n            preds = torch.argmax(outputs, dim=1)\n            \n            images_batch.append(images.cpu())\n            preds_batch.append(preds.cpu())\n            \n            if len(images_batch) * images.size(0) >= num_samples:\n                break\n\n    images_batch = torch.cat(images_batch)[:num_samples]\n    preds_batch = torch.cat(preds_batch)[:num_samples]\n    \n    return images_batch, preds_batch\n\n\n# Define a color map with fixed colors for each label\ndef get_label_colors(num_classes):\n    colors = plt.cm.tab20(np.linspace(0, 1, num_classes))\n    return colors\n\ndef visualize_predictions(images, masks, num_classes=6, num_samples=16):\n    num_samples = min(num_samples, len(images))\n    plt.figure(figsize=(10,10))\n    \n    label_colors = get_label_colors(num_classes)\n    \n    for i in range(num_samples):\n        plt.subplot(4, 2, i * 2 + 1)\n        im = images[i].numpy()\n        im = np.transpose(im, (1, 2, 0))\n        #denormalize\n        im = ((im * [0.229]) + [0.485]) * 255\n        plt.imshow(im,cmap='gray')\n        plt.title(\"Input Image\")\n        plt.axis('off')\n        \n        plt.subplot(4, 2, i * 2 + 2)\n        mask = masks[i].numpy()\n\n        color_mask = np.zeros((mask.shape[0], mask.shape[1], 3))\n        for label in range(num_classes):\n            color_mask[mask == label] = label_colors[label][:3] * 255\n        \n        plt.imshow(color_mask.astype(np.uint8))\n        plt.title(\"Predicted Mask\")\n        plt.axis('off')\n\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nmodel.to(device)\n\nimages, masks = inference(model, val_loader, device, num_samples=4)\nvisualize_predictions(images, masks, num_samples=6)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\n\ndef dice_score(inputs, targets, sm=1e-4):\n    # Apply softmax to get probabilities\n    prob = F.softmax(inputs, dim=1)\n    \n    # Calculate intersection and union\n    intersection = torch.sum(prob * targets, dim=(2, 3))  # Sum over spatial dimensions\n    union = torch.sum(prob, dim=(2, 3)) + torch.sum(targets, dim=(2, 3))  # Sum over spatial dimensions\n    \n    # Calculate Dice score\n    score = (2 * intersection + sm) / (union + sm)\n    \n    # Average over classes\n    scores = score.mean(dim=0)\n    \n    # Convert to list\n    return scores.tolist()\n\nmodel.eval()\nmodel.to(device)\no_scores = []\n\nwith torch.no_grad():  # Disable gradient computation\n    for images, masks in test_loader:\n        images = images.to(device)\n        masks = masks.to(device)\n        outputs = model(images)\n        scores = dice_score(outputs, masks)\n        o_scores.append(scores)\n\n# Convert the list of scores to a tensor and calculate the mean\no_scores = torch.tensor(o_scores).mean(dim=0)\n\n# Display Dice scores for each vertebrae and the overall score\nfor i in range(1,6):\n    print(f'The Dice score for vertebrae {i} is {o_scores[6-i]:.4f}')\nprint(f'Overall Dice score is {o_scores.mean():.4f}')\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\n# Example data for Dice scores (these should be replaced by the actual scores you have)\nvertebrae_classes = ['L1', 'L2', 'L3', 'L4', 'L5']\ndice_scores = (o_scores.tolist())[::-1][:-1]  # Replace with your actual scores\noverall_score = np.mean(dice_scores)\n\n# Create a simple bar plot for the Dice scores\nfigure=plt.figure(figsize=(10, 6))\nbars = plt.bar(vertebrae_classes, dice_scores, color='skyblue', edgecolor='black')\n\n# Annotate the bars with the Dice scores\nfor bar in bars:\n    yval = bar.get_height()\n    plt.text(bar.get_x() + bar.get_width() / 2, yval, f'{yval:.2f}', ha='center', va='bottom')\n\n# Add overall score as a separate text\nplt.text(0.5, 0.95, f'Overall Dice Score: {overall_score:.2f}', ha='center', va='center', fontsize=14, transform=plt.gca().transAxes)\n\n# Customize plot\nplt.title('Dice Scores for Vertebrae Segmentation', fontsize=16)\nplt.ylabel('Dice Score', fontsize=12)\nplt.ylim(0, 1)\nplt.grid(axis='y', linestyle='--', alpha=0.7)\nfigure.savefig('./dice_Scores.png')\n# Display plot\nplt.tight_layout()\n\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference on DCM Files ","metadata":{}},{"cell_type":"code","source":"import pydicom","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(torch.load('/kaggle/working/simple_unet.pth'))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport pydicom\nimport matplotlib.pyplot as plt\n\nmodel.eval()\nmodel.to(device)\n\n# Define a color map with fixed colors for each label\ndef get_label_colors(num_classes):\n    colors = plt.cm.tab20(np.linspace(0, 1, num_classes))\n    return colors\n\nfigure=plt.figure(figsize=(10, 10))\n\nlabel_colors = get_label_colors(num_classes)\n\nfile_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/1028684462/2914428894/8.dcm'\nim = pydicom.dcmread(file_path).pixel_array\nim = im / 255.0\nim = transforms_valid(image=im)['image']  \nim = im.unsqueeze(0).float().to(device)  # Ensure the transformed image is sent to the device\n\nwith torch.no_grad():\n    output = model(im)\n    pred = torch.argmax(output, dim=1)\n    im = im.squeeze(0).cpu()  # Remove the batch dimension and move to CPU for plotting\n    pred = pred.squeeze(0).cpu()  # Squeeze and move to CPU\n\n    # Plotting the original image\n    plt.subplot(1, 2, 1)\n    print(im.shape)\n    im = np.transpose(im.numpy(), (1, 2, 0))  # Transpose to (H, W, C)\n    im = ((im * 0.229) + 0.485) * 255  # Adjust normalization (ensure these values match your training)\n    plt.imshow(im, cmap='gray')\n    plt.title(\"IMAGE\")\n    plt.axis('off')\n\n    # Plotting the predicted mask\n    plt.subplot(1, 2, 2)\n    mask = pred.numpy()  # Convert prediction to numpy\n\n    color_mask = np.zeros((mask.shape[0], mask.shape[1], 3))  # Initialize the color mask\n    for label in range(num_classes):\n        color_mask[mask == label] = label_colors[label][:3] * 255  # Apply the colors based on label\n\n    plt.imshow(color_mask.astype(np.uint8))\n    plt.title(\"Predicted Mask\")\n    plt.axis('off')\n\nplt.tight_layout()\nplt.show()\nfigure.savefig('./Prediction.png')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport pydicom\nimport matplotlib.pyplot as plt\n\nmodel.eval()\nmodel.to(device)\n\n# Define a color map with fixed colors for each label\ndef get_label_colors(num_classes):\n    colors = plt.cm.tab20(np.linspace(0, 1, num_classes))\n    return colors\n\nfigure=plt.figure(figsize=(10, 10))\n\nlabel_colors = get_label_colors(num_classes)\n\nfile_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/1028684462/2914428894/9.dcm'\nim = pydicom.dcmread(file_path).pixel_array\nim = im / 255.0\nim = transforms_valid(image=im)['image']  # Assuming transforms_valid is already defined\nim = im.unsqueeze(0).float().to(device)  # Ensure the transformed image is sent to the device\n\nwith torch.no_grad():\n    output = model(im)\n    pred = torch.argmax(output, dim=1)\n    im = im.squeeze(0).cpu()  # Remove the batch dimension and move to CPU for plotting\n    pred = pred.squeeze(0).cpu()  # Squeeze and move to CPU\n\n    # Plotting the original image\n    plt.subplot(1, 2, 1)\n    im_np = np.transpose(im.numpy(), (1, 2, 0))  # Transpose to (H, W, C)\n    im_np = ((im_np * 0.229) + 0.485) * 255  # Adjust normalization (ensure these values match your training)\n    im_np = im_np  # Convert to uint8 for display\n    plt.imshow(im_np, cmap='gray')\n    plt.title(\"Original Image\")\n    plt.axis('off')\n\n    # Plotting the regions based on labels 1-5\n    plt.subplot(1, 2, 2)\n    combined_mask = np.zeros_like(im_np)  # Initialize a combined mask for all labels 1-5\n    \n    for label in range(1, 6):  # Loop through labels 1-5\n        binary_mask = (pred.numpy() == label)  # Create a binary mask for the current label\n        label_region = np.zeros_like(im_np)  # Initialize an empty image for the label region\n        \n        label_region[:, :,0] = im_np[:, :,0] * binary_mask\n\n        # Add the label region to the combined mask\n        combined_mask += label_region\n\n    plt.imshow(combined_mask,cmap='gray')\n    plt.title(\"Region Highlighted (Labels 1-5)\")\n    plt.axis('off')\n\nplt.tight_layout()\nfigure.savefig('./demo.png')\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}