{"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":36363,"databundleVersionId":4050810,"sourceType":"competition"},{"sourceId":7659622,"sourceType":"datasetVersion","datasetId":3607309}],"dockerImageVersionId":30648,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"!pip install nibabel torch torchvision torchaudio --no-index --find-links=file:///kaggle/input/rsna-spinal/packages","metadata":{"execution":{"iopub.status.busy":"2024-02-20T02:10:43.979325Z","iopub.execute_input":"2024-02-20T02:10:43.980272Z","iopub.status.idle":"2024-02-20T02:10:57.862689Z","shell.execute_reply.started":"2024-02-20T02:10:43.980228Z","shell.execute_reply":"2024-02-20T02:10:57.861528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!pip install nibabel torch torchvision torchaudio ","metadata":{"execution":{"iopub.status.busy":"2024-02-20T02:10:57.864867Z","iopub.execute_input":"2024-02-20T02:10:57.865202Z","iopub.status.idle":"2024-02-20T02:10:57.869883Z","shell.execute_reply.started":"2024-02-20T02:10:57.86517Z","shell.execute_reply":"2024-02-20T02:10:57.868838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport nibabel as nib\nimport tensorflow as tf\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom sklearn.metrics import accuracy_score, precision_recall_fscore_support, roc_auc_score\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import models\nfrom scipy.ndimage import zoom\nimport torch.nn.functional as F\nimport torchvision.models as models\nfrom torchvision import models\nfrom PIL import Image\nfrom tqdm import tqdm\nimport pickle ","metadata":{"execution":{"iopub.status.busy":"2024-02-20T02:10:57.878424Z","iopub.execute_input":"2024-02-20T02:10:57.87869Z","iopub.status.idle":"2024-02-20T02:11:16.486899Z","shell.execute_reply.started":"2024-02-20T02:10:57.878666Z","shell.execute_reply":"2024-02-20T02:11:16.486003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#  Precomputing","metadata":{}},{"cell_type":"code","source":"# Constants and configuration settings\nsegmentation_dir = '/kaggle/input/rsna-2022-cervical-spine-fracture-detection/segmentations'\ncsv_file = '/kaggle/input/rsna-spinal/train_file_mask_path.csv'\nbatch_size = 4\nnum_workers = 4\ndesired_shape = (128, 128, 128)\nnum_classes_classification = 7 \nnum_classes_segmentation = 1\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2024-02-20T02:11:16.488033Z","iopub.execute_input":"2024-02-20T02:11:16.488596Z","iopub.status.idle":"2024-02-20T02:11:16.540503Z","shell.execute_reply.started":"2024-02-20T02:11:16.488568Z","shell.execute_reply":"2024-02-20T02:11:16.53936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":" ## Dataset","metadata":{}},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, image_paths, mask_paths, labels, transform=None):\n        \"\"\"\n        Custom dataset for handling 3D NIfTI images and segmentation masks.\n\n        Args:\n            image_paths (list): List of paths to 3D NIfTI images.\n            mask_paths (list): List of paths to segmentation masks.\n            labels (list): List of corresponding labels.\n            transform (callable, optional): Transformations to be applied to the images and masks.\n        \"\"\"\n        self.image_paths = image_paths\n        self.mask_paths = mask_paths\n        self.labels = labels\n        self.transform = transform\n\n    def __len__(self):\n        \"\"\"\n        Returns the total number of samples in the dataset.\n        \"\"\"\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        \"\"\"\n        Retrieves a sample from the dataset.\n\n        Args:\n            idx (int): Index of the sample to retrieve.\n\n        Returns:\n            tuple: A tuple containing the 3D NIfTI image, segmentation mask, and label.\n        \"\"\"\n        # Retrieve file paths for the image and mask\n        image_path, mask_path = self.image_paths[idx], self.mask_paths[idx]\n\n        # Load the 3D NIfTI image using nibabel\n        image = nib.load(image_path).get_fdata()\n\n        # Load and process the segmentation mask if available\n        segmentation_mask = self.load_and_process_mask(mask_path)\n\n        # Apply transformations if provided to the image and segmentation mask\n        if self.transform:\n            image = self.transform(image)\n            segmentation_mask = self.transform(segmentation_mask) if segmentation_mask is not None else torch.zeros_like(image)\n\n        # Convert label to torch tensor\n        label = torch.tensor(self.labels[idx], dtype=torch.float32)\n\n        return image, segmentation_mask, label\n\n    def load_and_process_mask(self, mask_path):\n        \"\"\"\n        Load and process the segmentation mask.\n\n        Args:\n            mask_path (str): Path to the segmentation mask.\n\n        Returns:\n            torch.Tensor or None: Processed segmentation mask as a torch tensor or None if not available.\n        \"\"\"\n        segmentation_mask = None\n        if pd.notna(mask_path):\n            segmentation_mask = nib.load(mask_path)\n            segmentation_mask_data = segmentation_mask.get_fdata()\n            resized_data = resize_nifti(segmentation_mask_data, desired_shape)\n            segmentation_mask_data_affine = segmentation_mask.affine\n            resized_affine = segmentation_mask_data_affine\n            segmentation_mask = nib.Nifti1Image(resized_data, affine=resized_affine).get_fdata()\n\n        return segmentation_mask\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T02:11:16.541928Z","iopub.execute_input":"2024-02-20T02:11:16.542258Z","iopub.status.idle":"2024-02-20T02:11:16.555611Z","shell.execute_reply.started":"2024-02-20T02:11:16.54223Z","shell.execute_reply":"2024-02-20T02:11:16.554735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Resizing","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\n# Function to resize NIfTI data\ndef resize_nifti(nifti_data, target_shape):\n    factors = (target_shape[0] / nifti_data.shape[0],\n               target_shape[1] / nifti_data.shape[1],\n               target_shape[2] / nifti_data.shape[2])\n    resized_data = zoom(nifti_data, factors, order=3)  # Cubic interpolation (higher quality)\n    return resized_data","metadata":{"execution":{"iopub.status.busy":"2024-02-20T02:11:16.557021Z","iopub.execute_input":"2024-02-20T02:11:16.55772Z","iopub.status.idle":"2024-02-20T02:11:16.568212Z","shell.execute_reply.started":"2024-02-20T02:11:16.557684Z","shell.execute_reply":"2024-02-20T02:11:16.56727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Assignment","metadata":{}},{"cell_type":"code","source":"# Transformations\ntransform = transforms.Compose([\n    transforms.ToTensor(),  # Convert to tensor\n    # Add more transformations if necessary\n])\n\n# Load the CSV file\ndata = pd.read_csv(csv_file)\ndata_length = len(data)\nprint(\"Length of DataFrame:\", data_length)\n\n# Filter rows where the 'mask_path' column is not empty\n#data = data.dropna(subset=['mask_path']).reset_index(drop=True)\n\n# Split the data into training, validation, and test sets\ntrain_data, temp_data = train_test_split(data, test_size=0.3, random_state=42)\nval_data, test_data = train_test_split(temp_data, test_size=0.5, random_state=42)\n\n# Define a function to extract file paths and labels\ndef extract_paths_and_labels(dataset):\n    paths = dataset['file_path'].values\n    mask_paths = dataset['mask_path'].values\n    labels = dataset[['C1', 'C2', 'C3', 'C4', 'C5', 'C6', 'C7']].values\n    return paths, mask_paths, labels\n\n# Extract file paths and labels from the data\ntrain_paths, train_mask_paths, train_labels = extract_paths_and_labels(train_data)\nval_paths, val_mask_paths, val_labels = extract_paths_and_labels(val_data)\ntest_paths, test_mask_paths, test_labels = extract_paths_and_labels(test_data)\n\n# Instantiate the datasets\ntrain_dataset = CustomDataset(train_paths, train_mask_paths, train_labels, transform=transform)\nval_dataset = CustomDataset(val_paths, val_mask_paths, val_labels, transform=transform)\ntest_dataset = CustomDataset(test_paths, test_mask_paths, test_labels, transform=transform)\n\n# Instantiate the data loaders\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers, drop_last=True)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\n\nprint('Length of train_loader:', len(train_loader))\nprint('Length of val_loader:', len(val_loader))\nprint('Length of test_loader:', len(test_loader))\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T02:11:16.569869Z","iopub.execute_input":"2024-02-20T02:11:16.570437Z","iopub.status.idle":"2024-02-20T02:11:16.617649Z","shell.execute_reply.started":"2024-02-20T02:11:16.5704Z","shell.execute_reply":"2024-02-20T02:11:16.616591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class MultiLabel3DAttentionModel(nn.Module):\n    def __init__(self, num_classes_classification, num_classes_segmentation):\n        super(MultiLabel3DAttentionModel, self).__init__()\n\n        # Load a pre-trained ResNet3D backbone\n        self.backbone = models.video.r3d_18(pretrained=True)\n        \n        # Modify the stem to accept the correct input channels (128)\n        self.backbone.stem[0] = nn.Sequential(\n            nn.Conv3d(1, 64, kernel_size=(3, 7, 7), stride=(1, 2, 2), padding=(1, 3, 3)),\n            nn.BatchNorm3d(64),\n            nn.ReLU(inplace=True))\n\n        # Attention block\n        self.attention = nn.Sequential(\n            nn.Conv3d(1, 128, kernel_size=1),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(128, 1, kernel_size=1),\n            nn.Sigmoid()\n        )\n        \n        # Classification head\n        self.classification_head = nn.Sequential(\n            nn.AdaptiveAvgPool3d(1),\n            nn.Flatten(),\n            nn.Linear(128, num_classes_classification),  # Fix input size here\n            nn.ReLU(inplace=True),\n            nn.Linear(num_classes_classification, num_classes_classification),  # Adjust output size\n            nn.Sigmoid()\n        )\n        \n        # Segmentation head\n        self.segmentation_head = nn.Sequential(\n            nn.Conv3d(1, 128, kernel_size=1),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(128, num_classes_segmentation, kernel_size=1),\n            nn.Sigmoid()\n        )\n        \n    def forward(self, x, segmentation_mask):\n        # Feature extraction with the backbone\n        features = self.backbone(x)\n\n        # Apply attention to features\n        features = features.view(features.size(0), 1, 1, 1, features.size(1))\n        attention_weights = self.attention(features)\n        attended_features = features * attention_weights\n        \n        # Classification branch\n        classification_output = self.classification_head(attended_features)\n        \n        # Segmentation branch\n        segmentation_output = self.segmentation_head(attended_features)\n        segmentation_output = F.interpolate(segmentation_output, size=segmentation_mask.shape[2:], mode='trilinear')\n        segmentation_output = segmentation_output * segmentation_mask\n        \n        return classification_output, segmentation_output\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T02:11:16.620971Z","iopub.execute_input":"2024-02-20T02:11:16.621251Z","iopub.status.idle":"2024-02-20T02:11:16.63584Z","shell.execute_reply.started":"2024-02-20T02:11:16.621228Z","shell.execute_reply":"2024-02-20T02:11:16.634727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_classes_classification = 7 \nnum_classes = 7 \nnum_classes_segmentation = 1\nmodel = MultiLabel3DAttentionModel(num_classes_classification, num_classes_segmentation)","metadata":{"execution":{"iopub.status.busy":"2024-02-20T02:11:16.637207Z","iopub.execute_input":"2024-02-20T02:11:16.637636Z","iopub.status.idle":"2024-02-20T02:11:21.314354Z","shell.execute_reply.started":"2024-02-20T02:11:16.637609Z","shell.execute_reply":"2024-02-20T02:11:21.313347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training and Validation","metadata":{}},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.optim as optim\nimport torch.nn as nn\nfrom tqdm import tqdm\nfrom sklearn.metrics import accuracy_score\n\ndef train_epoch(model, train_loader, optimizer, criterion, device):\n    model.train()\n    running_loss = 0.0\n    correct_train = 0\n    total_train = 0\n\n    for batch_images, batch_segmentation_masks, batch_labels in tqdm(train_loader, desc=\"Training\"):\n        optimizer.zero_grad()\n        \n        # Move data to the GPU if available\n        batch_images = batch_images.to(torch.float32).to(device)\n        batch_segmentation_masks = batch_segmentation_masks.to(torch.float32).to(device)\n        batch_labels = batch_labels.to(torch.float32).to(device)\n\n        # Assuming batch_images has shape (batch_size, num_frames, num_channels, height, width)\n        batch_images = batch_images.unsqueeze(1)  # Add a singleton dimension for channels\n        batch_segmentation_masks = batch_segmentation_masks.unsqueeze(1)\n\n        # Forward pass\n        classification_outputs, segmentation_outputs = model(batch_images, batch_segmentation_masks)\n\n        # Calculate binary cross-entropy loss for each class separately\n        total_loss = criterion(classification_outputs, batch_labels)\n\n        # Check if segmentation mask is available\n        if batch_segmentation_masks is not None:\n            # Apply sigmoid activation to segmentation_outputs\n            segmentation_outputs = torch.sigmoid(segmentation_outputs)\n            # Calculate segmentation loss\n            segmentation_loss = criterion(segmentation_outputs, batch_segmentation_masks)\n            total_loss += segmentation_loss\n\n        running_loss += total_loss.item()\n\n        # Calculate accuracy for each class separately\n        accuracies = []\n        for class_index in range(7):\n            class_labels = batch_labels[:, class_index]  # Select labels for the current class\n            class_outputs = classification_outputs[:, class_index]  # Select model outputs for the current class\n\n            # Calculate binary predictions based on a threshold (e.g., 0.5)\n            predicted = (class_outputs >= 0.5).float()\n\n            class_accuracy = accuracy_score(class_labels.cpu(), predicted.cpu())\n            accuracies.append(class_accuracy)\n\n        # Calculate overall accuracy\n        batch_accuracy = sum(accuracies) / 7\n        correct_train += batch_accuracy\n        total_train += 1\n\n        # Backpropagation and optimization\n        total_loss.backward()\n        \n        optimizer.step()\n\n    # Calculate and return average training accuracy and loss\n    avg_train_accuracy = correct_train / total_train\n    avg_train_loss = running_loss / len(train_loader)\n    return avg_train_accuracy, avg_train_loss\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T01:01:13.932669Z","iopub.execute_input":"2024-02-20T01:01:13.932952Z","iopub.status.idle":"2024-02-20T01:01:13.945378Z","shell.execute_reply.started":"2024-02-20T01:01:13.932928Z","shell.execute_reply":"2024-02-20T01:01:13.944362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Validation","metadata":{}},{"cell_type":"code","source":"def validate_epoch(model, val_loader, criterion, device):\n    model.eval()\n    total_val_loss = 0.0\n    correct_val = 0\n    total_val = 0\n    \n    with torch.no_grad():\n        for batch_images, batch_segmentation_masks, batch_labels in tqdm(val_loader, desc=\"Validation\"):\n            batch_images = batch_images.to(torch.float32).to(device)\n            batch_segmentation_masks = batch_segmentation_masks.to(torch.float32).to(device)\n            batch_labels = batch_labels.to(torch.float32).to(device)\n\n            batch_images = batch_images.unsqueeze(1)\n            batch_segmentation_masks = batch_segmentation_masks.unsqueeze(1)\n\n            # Forward pass\n            classification_outputs, segmentation_outputs = model(batch_images, batch_segmentation_masks)\n\n            # Calculate binary cross-entropy loss for each class separately\n            total_loss = criterion(classification_outputs, batch_labels)\n\n            # Check if segmentation mask is available\n            if batch_segmentation_masks is not None:\n                # Apply sigmoid activation to segmentation_outputs\n                segmentation_outputs = torch.sigmoid(segmentation_outputs)\n                # Calculate segmentation loss\n                segmentation_loss = criterion(segmentation_outputs, batch_segmentation_masks)\n                total_loss += segmentation_loss\n\n            total_val_loss += total_loss.item()\n\n            # Calculate accuracy for each class separately\n            accuracies = []\n            for class_index in range(7):\n                class_labels = batch_labels[:, class_index]\n                class_outputs = classification_outputs[:, class_index]\n\n                # Calculate binary predictions based on a threshold (e.g., 0.5)\n                predicted = (class_outputs >= 0.5).float()\n\n                class_accuracy = accuracy_score(class_labels.cpu(), predicted.cpu())\n                accuracies.append(class_accuracy)\n\n            # Calculate overall accuracy\n            batch_accuracy = sum(accuracies) / 7\n            correct_val += batch_accuracy\n            total_val += batch_labels.size(0)\n\n    # Calculate and return average validation accuracy and loss\n    val_accuracy = correct_val / total_val\n    avg_val_loss = total_val_loss / len(val_loader)\n    return val_accuracy, avg_val_loss\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T01:01:13.946651Z","iopub.execute_input":"2024-02-20T01:01:13.946944Z","iopub.status.idle":"2024-02-20T01:01:13.96116Z","shell.execute_reply.started":"2024-02-20T01:01:13.946906Z","shell.execute_reply":"2024-02-20T01:01:13.960247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train - Eval","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nimport pickle\n\ndef main_train_and_validate(num_epochs=10, patience=10, model_path='sample.pth', history_path='samplehistory.pkl'):\n    # Set device\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    model.to(device)\n    \n    # Define criterion and optimizer\n    criterion = nn.BCEWithLogitsLoss()\n    optimizer = optim.Adam(model.parameters(), lr=0.01)\n\n    # Initialize metrics history\n    val_loss_history = []\n    val_acc_history = []\n    train_loss_history = []\n    train_acc_history = []\n    best_val_loss = float('inf')  # Initialize with a large value\n    no_improvement_count = 0  # Initialize the count\n\n    for epoch in range(num_epochs):\n        # Train the model\n        avg_train_accuracy, avg_train_loss = train_epoch(model, train_loader, optimizer, criterion, device)\n        print(f\"Epoch [{epoch + 1}/{num_epochs}]\")\n        print(f\"Train Accuracy: {avg_train_accuracy:.4f} | Train Loss: {avg_train_loss:.4f}\")\n\n        # Validate the model\n        val_accuracy, avg_val_loss = validate_epoch(model, val_loader, criterion, device)\n        print(f\"Validation Accuracy: {val_accuracy:.4f} | Validation Loss: {avg_val_loss:.4f}\")\n\n        # Save metrics\n        val_acc_history.append(val_accuracy)\n        val_loss_history.append(avg_val_loss)\n        train_acc_history.append(avg_train_accuracy)\n        train_loss_history.append(avg_train_loss)\n\n        # Check if validation loss has improved\n        if avg_val_loss < best_val_loss:\n            best_val_loss = avg_val_loss\n            no_improvement_count = 0  # Reset the count since there was an improvement\n\n            # Save the best model checkpoint\n            torch.save(model.state_dict(), model_path)\n        else:\n            no_improvement_count += 1\n\n        # If no improvement for 'early_stopping_patience' epochs, stop training\n        if no_improvement_count >= patience:\n            print(f'Early stopping after {epoch + 1} epochs without improvement.')\n            break\n\n    # Save the lists of metrics to a file for later plotting\n    metrics_history = {\n        'val_loss_history': val_loss_history,\n        'val_acc_history': val_acc_history,\n        'train_loss_history': train_loss_history,\n        'train_acc_history': train_acc_history,\n    }\n\n    with open(history_path, 'wb') as f:\n        pickle.dump(metrics_history, f)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T01:01:13.962406Z","iopub.execute_input":"2024-02-20T01:01:13.962686Z","iopub.status.idle":"2024-02-20T01:01:13.978871Z","shell.execute_reply.started":"2024-02-20T01:01:13.962663Z","shell.execute_reply":"2024-02-20T01:01:13.977905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Looped Runs","metadata":{}},{"cell_type":"code","source":"# For first time run\nmain_train_and_validate(num_epochs=5, patience=3)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T01:34:30.91784Z","iopub.execute_input":"2024-02-20T01:34:30.918155Z","iopub.status.idle":"2024-02-20T01:54:52.889198Z","shell.execute_reply.started":"2024-02-20T01:34:30.91812Z","shell.execute_reply":"2024-02-20T01:54:52.887944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # For recursive run\n# model.load_state_dict(torch.load('/kaggle/input/sample/sample.pth'))  # Load the pre-trained weights\n# main_train_and_validate(num_epochs=10, patience=1, model_path='/kaggle/input/sample/sample.pth', history_path='/kaggle/input/sample/metrics_history.pkl')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Testing","metadata":{}},{"cell_type":"code","source":"# Test loop\nmodel.eval()\ntotal_correct = 0\ntotal_samples = 0\n\n# Initialize lists to store per-class metrics\nprecision_list = []\nrecall_list = []\nf1_list = []\n\nwith torch.no_grad():\n    for batch_images, batch_segmentation_masks, batch_labels in test_loader:\n        batch_images = batch_images.to(torch.float32).to(device)\n        batch_segmentation_masks = batch_segmentation_masks.to(torch.float32).to(device)\n        batch_labels = batch_labels.to(torch.float32).to(device)\n\n        batch_images = batch_images.unsqueeze(1)\n        batch_segmentation_masks = batch_segmentation_masks.unsqueeze(1)\n\n        # Forward pass\n        classification_outputs, segmentation_outputs = model(batch_images, batch_segmentation_masks)\n\n        # Apply sigmoid activation to the classification outputs\n        classification_outputs = torch.sigmoid(classification_outputs)\n\n        # Initialize batch-level variables for accuracy calculation\n        batch_true_positives = 0\n        batch_false_positives = 0\n        batch_false_negatives = 0\n        batch_samples = batch_labels.size(0)\n\n        for class_index in range(7):\n            class_labels = batch_labels[:, class_index]\n            class_outputs = classification_outputs[:, class_index]\n\n            # Calculate binary predictions based on a threshold\n            best_threshold = 0.5  # Default threshold\n            best_f1 = 0.0\n\n            for threshold in torch.linspace(0.1, 0.9, 9):  # Adjust the range as needed\n                predicted = (class_outputs > threshold).float()\n\n                # Calculate true positives, false positives, and false negatives\n                true_positives = torch.sum((class_labels == 1) & (predicted == 1))\n                false_positives = torch.sum((class_labels == 0) & (predicted == 1))\n                false_negatives = torch.sum((class_labels == 1) & (predicted == 0))\n\n                # Calculate precision, recall, and F1-score for the batch\n                batch_precision = true_positives / (true_positives + false_positives + 1e-12)\n                batch_recall = true_positives / (true_positives + false_negatives + 1e-12)\n                batch_f1 = 2 * (batch_precision * batch_recall) / (batch_precision + batch_recall + 1e-12)\n\n                # Update best threshold if F1-score improves\n                if batch_f1 > best_f1:\n                    best_threshold = threshold.item()\n                    best_f1 = batch_f1.item()\n\n            # Use the best threshold for predictions\n            predicted = (class_outputs > best_threshold).float()\n\n            # Append metrics for the batch\n            precision_list.append(batch_precision.item())\n            recall_list.append(batch_recall.item())\n            f1_list.append(batch_f1.item())\n\n            # Accumulate totals\n            batch_true_positives += torch.sum(class_labels == predicted).item()\n            batch_false_positives += torch.sum((class_labels == 0) & (predicted == 1)).item()\n            batch_false_negatives += torch.sum((class_labels == 1) & (predicted == 0)).item()\n\n        # Accumulate batch-level accuracy\n        total_correct += batch_true_positives\n        total_samples += batch_samples\n\n# Calculate global accuracy\ntest_accuracy = total_correct / total_samples\nprint(f\"Test Accuracy: {test_accuracy:.4f}\")\n\n# Calculate average precision, recall, and F1-score across all batches\navg_precision = sum(precision_list) / len(precision_list)\navg_recall = sum(recall_list) / len(recall_list)\navg_f1 = sum(f1_list) / len(f1_list)\n\nprint(f\"Average Precision: {avg_precision:.4f}\")\nprint(f\"Average Recall: {avg_recall:.4f}\")\nprint(f\"Average F1 Score: {avg_f1:.4f}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-02-20T02:00:05.537708Z","iopub.execute_input":"2024-02-20T02:00:05.538764Z","iopub.status.idle":"2024-02-20T02:01:43.73577Z","shell.execute_reply.started":"2024-02-20T02:00:05.538722Z","shell.execute_reply":"2024-02-20T02:01:43.734095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Plots","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}