{"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":69711,"sourceType":"modelInstanceVersion","modelInstanceId":58166,"modelId":80339}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Dummy dataset and dataloader for example purposes\nfrom torch.utils.data import DataLoader, TensorDataset\n\n# Assuming input size of (N, 3, 256, 256) and number of classes is 21\ninputs = torch.randn(10, 3, 256, 256)\ntargets = torch.randint(0, 21, (10, 256, 256))\n\ndataset = TensorDataset(inputs, targets)\ndataloader = DataLoader(dataset, batch_size=2)\n\n# Initialize the model, loss function, and optimizer\nmodel = SegNet(num_classes=21)\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\n\n# Training loop\nnum_epochs = 5\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    for images, labels in dataloader:\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item()\n\n    print(f\"Epoch [{epoch+1}/{num_epochs}], Loss: {running_loss/len(dataloader)}\")\n# 2d Segmentation of Sagittal Lumbar Spine MRI\n\n- Training data Spider dataset (https://doi.org/10.5281/zenodo.10159290)\n- Very simple model using segmentation_models_pytorch\n- Used ChatGPT and Gemini to help with coding\n- Trained model attached\n- Images resized to 256x256","metadata":{}},{"cell_type":"code","source":"!pip install segmentation_models_pytorch -q","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install torch torchvision segmentation-models-pytorch\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport os\nfrom pathlib import Path\nfrom PIL import Image\n\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\n  ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#transforms\nnewsize = (256, 256)\n#dataset\nfold = 1\n#dataloader\nbatch_size = 64\nnum_workers = 4\n#model\nnum_classes = 20\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")#run\nepochs = 100\nlearning_rate = 1e-3\n\nTRAIN = True #or False for inference only","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model","metadata":{}},{"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=3  # Number of input channels (e.g. 3 for RGB)\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"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)\n        if image.mode != 'RGB':  # Ensure image is RGB\n            image = image.convert('RGB')\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        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":{"trusted":true},"execution_count":null,"outputs":[]},{"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, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225],\n    ),\n    ToTensorV2()\n])\n\ntransforms_valid = A.Compose([\n    A.Resize(newsize[0], newsize[1]),\n    A.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225],\n    ),\n    ToTensorV2()\n])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Loss","metadata":{}},{"cell_type":"code","source":"class CombinedLoss(nn.Module):\n    def __init__(self, weight_ce=1.0, weight_iou=1.0):\n        super(CombinedLoss, self).__init__()\n        self.weight_ce = weight_ce\n        self.weight_iou = weight_iou\n        self.cross_entropy_loss = nn.CrossEntropyLoss()\n\n    def forward(self, inputs, targets):\n        # Cross-Entropy Loss\n        ce_loss = self.cross_entropy_loss(inputs, targets)\n\n        # IoU Loss\n        # Apply softmax to the inputs to get probabilities\n        probs = F.softmax(inputs, dim=1)\n\n        intersection = torch.sum(probs * targets, dim=(2, 3))\n        union = torch.sum(probs + targets, dim=(2, 3)) - intersection\n        iou = (intersection + 1e-6) / (union + 1e-6)\n        iou_loss = 1 - iou.mean()\n\n        # Combine losses\n        loss = self.weight_ce * ce_loss + self.weight_iou * iou_loss\n        return loss","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Create datasets and dataloaders","metadata":{}},{"cell_type":"code","source":"train_ = df[df['fold'] != fold].reset_index(drop=True)\nvalid_ = df[df['fold'] == fold].reset_index(drop=True)\n\ndataset_train = SEGDataset(train_, 'train',  transforms_train)\ndataset_valid = SEGDataset(valid_, '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)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Run function","metadata":{}},{"cell_type":"code","source":"from torch import optim\nfrom torch.nn import BCEWithLogitsLoss\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\n\ndef run(train_loader, val_loader, model, learning_rate, criterion, epochs, device):\n  \"\"\"\n  Trains a U-net model for multi-label segmentation.\n\n  Args:\n      train_loader: DataLoader for training data.\n      val_loader: DataLoader for validation data.\n      model: U-net model instance.\n      learning_rate: Learning rate for optimizer.\n      epochs: Number of epochs to train.\n      device: Device to use for training (CPU or GPU).\n  \"\"\"\n  # Define loss function and optimizer\n  optimizer = optim.Adam(model.parameters(), lr=learning_rate)\n\n  # Define a learning rate scheduler\n  scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=5, verbose=True)\n\n\n  # Training loop\n  for epoch in range(epochs):\n    model.train()\n    train_loss = 0.0\n    for images, masks in train_loader:\n      images, masks = images.to(device), masks.to(device)\n\n      # Forward pass and calculate loss\n      outputs = model(images)\n      loss = criterion(outputs, masks)\n\n      # Backward pass and update weights\n      optimizer.zero_grad()\n      loss.backward()\n      optimizer.step()\n\n      train_loss += loss.item()\n\n    train_loss /= len(train_loader)\n\n    # Validation step (optional)\n    model.eval()\n    with torch.no_grad():\n      val_loss = 0.0\n      for images, masks in val_loader:\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    \n    # Step the scheduler\n    scheduler.step(val_loss)\n\n    # Print training and validation loss\n    print(f\"Epoch: {epoch+1}/{epochs} | Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Train","metadata":{}},{"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/simple_unet_2d_lspine/pytorch/one/1/simple_unet.pth\"))\n                      ","metadata":{"scrolled":true,"trusted":true},"execution_count":null,"outputs":[]},{"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=20, num_samples=16):\n    num_samples = min(num_samples, len(images))\n    plt.figure(figsize=(20, 20))\n    \n    label_colors = get_label_colors(num_classes)\n    \n    for i in range(num_samples):\n        plt.subplot(4, 8, i * 2 + 1)\n        im = images[i].numpy()\n        im = np.transpose(im, (1, 2, 0))\n        #denormalize\n        im = ((im * [0.229, 0.224, 0.225]) + [0.485, 0.456, 0.406]) * 255\n        plt.imshow(im)\n        plt.title(\"Input Image\")\n        plt.axis('off')\n        \n        plt.subplot(4, 8, 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},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\nmodel.to(device)\n\nimages, masks = inference(model, val_loader, device, num_samples=16)\nvisualize_predictions(images,   masks, num_samples=16)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\n# Function to calculate the pixel-wise accuracy\ndef calculate_accuracy(predictions, masks):\n    # Convert the predictions to binary format (0 or 1) based on a threshold\n    predictions = (predictions > 0.5).float()\n    \n    # Calculate the number of correct pixels\n    correct_pixels = (predictions == masks).float().sum()\n    \n    # Calculate the total number of pixels\n    total_pixels = torch.numel(predictions)\n    \n    # Calculate accuracy\n    accuracy = correct_pixels / total_pixels\n    \n    return accuracy.item()\n\n# Set the model to evaluation mode and move it to the device\nmodel.eval()\nmodel.to(device)\n\n# Loop through the validation dataset\ntotal_accuracy = 0\nnum_batches = 0\n\nwith torch.no_grad():\n    for images, masks in val_loader:\n        images = images.to(device)\n        masks = masks.to(device)\n\n        # Get the model predictions\n        outputs = model(images)\n\n        # Calculate accuracy for the batch\n        batch_accuracy = calculate_accuracy(outputs, masks)\n        total_accuracy += batch_accuracy\n        num_batches += 1\n\n# Calculate average accuracy over all batches\naverage_accuracy = total_accuracy / num_batches\n\nprint(f'Pixel-wise Accuracy: {average_accuracy:.4f}')\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Save model","metadata":{}},{"cell_type":"code","source":"torch.save(model.state_dict(), './simple_unet.pth')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**SEGNET**","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\nimport torch\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)  # SegNet usually outputs logits for each class\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\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=20, num_samples=16):\n    num_samples = min(num_samples, len(images))\n    plt.figure(figsize=(20, 20))\n    \n    label_colors = get_label_colors(num_classes)\n    \n    for i in range(num_samples):\n        plt.subplot(4, 8, i * 2 + 1)\n        im = images[i].numpy()\n        im = np.transpose(im, (1, 2, 0))\n        # Denormalize\n        im = ((im * [0.229, 0.224, 0.225]) + [0.485, 0.456, 0.406]) * 255\n        plt.imshow(im.astype(np.uint8))\n        plt.title(\"Input Image\")\n        plt.axis('off')\n        \n        plt.subplot(4, 8, i * 2 + 2)\n        mask = masks[i].numpy()\n\n        color_mask = np.zeros((mask.shape[1], mask.shape[2], 3))  # Adjust dimensions\n        for label in range(num_classes):\n            color_mask[mask[0] == label] = label_colors[label][:3] * 255  # Adjust for single-channel mask\n        \n        plt.imshow(color_mask.astype(np.uint8))\n        plt.title(\"Predicted Mask\")\n        plt.axis('off')\n\n    plt.show()\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\nmodel.to(device)\n\nimages, masks = inference(model, val_loader, device, num_samples=16)\nvisualize_predictions(images,   masks, num_samples=16)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}