{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.8.17","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":52254,"databundleVersionId":6863140,"sourceType":"competition"},{"sourceId":6523471,"sourceType":"datasetVersion","datasetId":3771357},{"sourceId":6524344,"sourceType":"datasetVersion","datasetId":3771912}],"dockerImageVersionId":30529,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install nibabel","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-03T13:25:02.64237Z","iopub.execute_input":"2023-12-03T13:25:02.643051Z","iopub.status.idle":"2023-12-03T13:25:08.479178Z","shell.execute_reply.started":"2023-12-03T13:25:02.643019Z","shell.execute_reply":"2023-12-03T13:25:08.478186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.distributed.xla_multiprocessing as xmp\n\nimport torch\nimport torch.nn.functional as F\nimport torch_xla.core.xla_model as xm\nfrom torch.nn import Transformer\n\ndef xla_linear(input, weight, bias=None):\n    if isinstance(input, torch.Tensor) and input.device.type == 'xla':\n#         print(\"input\", input.shape)\n#         print(\"************************************************************************\")\n#         print(\"weight\", weight)\n#         print(\"$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$$\")\n#         print(\"bias\", bias)\n#         return torch.nn.functional.linear(input, weight, bias)\n        return torch.matmul(input.to(xm.xla_device()), weight.to(xm.xla_device()).t()) + bias.to(xm.xla_device())\n    else:\n        input_xla = input.to(xm.xla_device())\n        weight_xla = weight.to(xm.xla_device())\n        if bias is not None:\n            bias_xla = bias.to(xm.xla_device())\n        else:\n            bias_xla = None\n#         return torch.nn.functional.linear(input_xla, weight_xla, bias_xla)\n        return torch.matmul(input_xla, weight_xla.t()) + bias_xla\n    \n# Override the torch.nn.functional.linear function with the XLA version\n# F.linear = xla_linear\n\ndef xla_layer_norm(input, normalized_shape, weight=None, bias=None, eps=1e-5):\n    if input.device.type == 'xla':\n        # Calculate the mean and variance along the last dimension\n        mean = input.mean(dim=-1, keepdim=True)\n        var = input.var(dim=-1, unbiased=False, keepdim=True)\n        \n        # Reshape weight and bias to match the shape of input\n        if weight is not None:\n            weight = weight.view(*input.shape[-len(normalized_shape):])\n        if bias is not None:\n            bias = bias.view(*input.shape[-len(normalized_shape):])\n        \n        # Normalize the input\n        input = (input - mean) / torch.sqrt(var + eps)\n        \n        # Apply weight and bias\n        if weight is not None:\n            input = input * weight\n        if bias is not None:\n            input = input + bias\n        print(input.shape)\n        return input\n    else:\n        # Fall back to PyTorch's layer normalization\n        return F.layer_norm(input, normalized_shape, weight, bias, eps)\n\n# Override the torch.nn.functional.layer_norm function with the XLA version\n# F.layer_norm = xla_layer_norm\n\ninput_shape = (128, 128, 128)  # Depth x Height x Width\nnum_classes = 14  # Number of classes for classification\n\nclass Transformer3DClassifier(nn.Module):\n    def __init__(self, input_shape, num_classes, num_layers=6, d_model=16, nhead=8, dim_feedforward=2048, dropout=0.1):\n        super(Transformer3DClassifier, self).__init__()\n        \n        # Initialize d_model\n        self.d_model = d_model\n        \n        # Calculate the input size for the transformer\n        d_in = input_shape[0] * input_shape[1] * input_shape[2]  # Depth x Height x Width\n        self.embedding = nn.Linear(d_in, d_model)\n        \n        self.transformer = Transformer(\n            d_model=d_model,\n            nhead=nhead,\n            num_encoder_layers=num_layers,\n            dim_feedforward=dim_feedforward,\n            dropout=dropout\n        )\n        \n        self.fc = nn.Linear(d_model, num_classes)\n   \n\n    def forward(self, x):\n        # Flatten the input and apply linear embedding\n        x = x.view(x.size(0), -1)\n        print(\"Before embedding x.shape is \", x.shape)\n        x = self.embedding(x)\n        print(\"After embedding x.shape is \", x.shape)\n        \n        # Reshape to add a third dimension (seq_len)\n        x = x.unsqueeze(0)\n        print(\"x shape after unsqueeze\", x.shape)\n        # Create a dummy target tensor (you can adjust its size if needed)\n        tgt = torch.zeros(1, x.size(1), self.d_model).to(x.device)\n        print(\"tgt shape\", tgt.shape)\n        \n        # Transformer encoder\n        output = self.transformer(x, tgt)\n        print(\"Output shape after transformer\", output.shape)\n\n        # Remove the added dimension\n#         output = output.squeeze(0)\n#         print(\"Output shape after squeeze\", output.shape)\n        \n#         # Global average pooling\n#         output = output.mean(dim=1)\n#         print(\"Output shape after global average pooling\", output.shape)\n\n        # Classification layer\n        logits = self.fc(output)\n        \n        # Add batch dimension to logits\n        logits = logits.unsqueeze(0)\n        \n        \n        return logits\n\n# Define XLA tensors for input and hidden layer sizes\n# input_size = torch.tensor(32, device=xm.xla_device())\ninput_size = 32 # Adjust the dimensions as needed\nhidden_size = 16  # Adjust the dimensions as needed\n# hidden_size = torch.tensor(16, device=xm.xla_device())\n\nclass Custom3DViTModelTPU(nn.Module):\n    def __init__(self, in_channels, num_classes, num_classes_segmentation, batch_size):\n        super(Custom3DViTModelTPU, self).__init__()\n        self.batch_size = batch_size\n        self.num_classes = num_classes\n        \n#         self.vit_backbone = VisionTransformer3DBackboneTPU(\n#             in_channels=in_channels,\n#             embedding_dim=32,  # Adjust the embedding dimension as needed\n#             num_heads=2,       # Number of attention heads\n#             num_layers=2       # Number of transformer layers\n#         )\n        \n        self.vit_backbone = Transformer3DClassifier(\n            input_shape,\n            num_classes\n        )\n\n        self.classification_head = nn.Sequential(\n#             nn.Linear(batch_size, 16),\n#             nn.ReLU(inplace=True),\n#             nn.Linear(16, num_classes),\n#             nn.Sigmoid()\n            nn.Linear(self.vit_backbone.d_model, num_classes)\n        )\n\n        self.segmentation_head = nn.Sequential(\n            nn.Conv3d(1, num_classes_segmentation, kernel_size=1),\n            nn.Sigmoid()\n        )\n    def print_weights(self):\n        for name, param in self.named_parameters():\n            print(f\"Layer: {name}, Size: {param.size()}\")\n            print(param)\n\n    def forward(self, x, segmentation_mask):\n        print(\"x shape and segmentation_mask shape\", x.shape, segmentation_mask.shape)\n        \n        # Move input tensors to XLA devices\n        x = x.to(xm.xla_device())\n        segmentation_mask = segmentation_mask.to(xm.xla_device())\n\n        features = self.vit_backbone(x)\n        features = features.to(xm.xla_device())\n        print(\"features shape\", features.shape)\n        \n        #classification_output = self.classification_head(features)\n        # Reshape it to (32, 10)\n        classification_output = features.view(self.batch_size, self.num_classes)\n        \n        print(\"classification_output\", classification_output.shape)\n        segmentation_output = self.segmentation_head(x)\n        print(\"segmentation output\", segmentation_output.shape)\n\n#         # Resize segmentation_output to match the shape of segmentation_mask\n        segmentation_output = nn.functional.interpolate(segmentation_output, size=segmentation_mask.shape[2:], mode='trilinear')\n\n        segmentation_output = segmentation_output * segmentation_mask\n\n        return classification_output, segmentation_output\n\nbatch_size = 32\n\n# Move the entire model to XLA devices\ndef get_model():\n    return Custom3DViTModelTPU(3, 14, 5, batch_size)\n\n# Modify the run function to accept the process index\ndef run(index):\n    print(\"^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^\")\n#     l_in = torch.randn(10, device=xm.xla_device())\n#     linear = torch.nn.Linear(10, 20).to(xm.xla_device())\n#     l_out = linear(l_in)\n#     print(l_out)\n    \n    model = get_model()\n    model.print_weights()\n    model = model.to(xm.xla_device())\n    print(\">>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>\")\n    # Create sample input tensors (modify this according to your data)\n    batch_images = torch.randn(32, 1, 128, 128, 128)  # Example input shape\n    batch_segmentation_masks = torch.randn(32, 1, 128, 128, 128)  # Example mask shape\n\n    batch_images = batch_images.to(xm.xla_device())  # Move input tensors to XLA device\n    batch_segmentation_masks = batch_segmentation_masks.to(xm.xla_device())\n    print(\"*****************************************************************************\")\n\n    # Forward pass\n    classification_outputs, segmentation_outputs = model(batch_images, batch_segmentation_masks)\n    \n# # Use XLA multiprocessing to distribute across TPUs\nif __name__ == '__main__':\n     xmp.spawn(run, nprocs=1, start_method='fork')","metadata":{"execution":{"iopub.status.busy":"2023-12-03T13:25:08.48121Z","iopub.execute_input":"2023-12-03T13:25:08.481495Z","iopub.status.idle":"2023-12-03T13:26:10.299282Z","shell.execute_reply.started":"2023-12-03T13:25:08.481468Z","shell.execute_reply":"2023-12-03T13:26:10.297956Z"},"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 torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch_xla.core.xla_model as xm\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\nfrom PIL import Image\n\n# Create a function to move data to the XLA device\ndef move_data_to_xla(data):\n    return data.to(torch.float32).to(device)\n\nclass CustomDataset(Dataset):\n    def __init__(self, image_paths, mask_paths, labels, transform=None):\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        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        image_path = self.image_paths[idx]\n        mask_path = self.mask_paths[idx]\n\n        # Load the 3D NIfTI image using nibabel\n        image = nib.load(image_path).get_fdata()\n#         print(image.shape)\n#         print(\"Data shape *********:\", image.dtype)\n\n        # Load the segmentation mask if available\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        # Apply transformations if provided to the image\n        if self.transform:\n            image = self.transform(image)\n\n        # Apply transformations if provided to the segmentation mask\n        if segmentation_mask is not None and self.transform:\n            segmentation_mask = self.transform(segmentation_mask)\n        else:\n            segmentation_mask = torch.zeros_like(image)\n\n        label = torch.tensor(self.labels[idx], dtype=torch.float32)\n        \n        return image, segmentation_mask, label\n\nimport torch\nimport torch.nn as nn\nimport torchvision.models as models\n\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\n\n# Paths and settings\nsegmentation_dir = '/kaggle/input/rsna-2023-abdominal-trauma-detection/segmentations'  # Update with the correct path\ncsv_file = '/kaggle/input/abdominal-trauma-nii-csv/abdominal_trauma_nii.csv'  # Update with the correct path\nbatch_size = 32\nnum_workers = 4  # Number of CPU cores to use for data loading\nnum_classes = 14  # Number of classes\nnum_classes_segmentation = 1  # Number of classes\ndesired_shape = (128, 128, 128)\n# device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\ndevice = xm.xla_device()\nprint(device)\n\n# Define transformations if needed\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).head(4480)\n\n# Filter rows where the 'mask_path' column is not empty\n#data = data[pd.notna(data['mask_path'])]\n\n# Remove the extra space from the column name\ndata.columns = data.columns.str.strip()\n\n# Assuming 'data' is your DataFrame\ndata_length = len(data)\nprint(\"Length of DataFrame:\", data_length)\n# print(data)\n# print(data.index)\n\n# Split the data into training, validation, and test sets\ntrain_data, temp_data = train_test_split(data, test_size=0.2, random_state=42)\nval_data, test_data = train_test_split(temp_data, test_size=0.5, random_state=42)\n\n# Set the display option to show alnl rows\npd.set_option('display.max_rows', None)\n\nindex_values = train_data.index.values\n\n# chunk_size = 200  # You can adjust the chunk size\n# for i in range(0, len(index_values), chunk_size):\n#     print(index_values[i:i+chunk_size])\n\n# Reset the display option to its default value (if needed)\npd.reset_option('display.max_rows')\n\n# Extract file paths and labels from the data\ntrain_paths = train_data['file_path'].values\ntrain_mask_paths = train_data['mask_path'].values\ntrain_labels = train_data[['bowel_healthy','bowel_injury','extravasation_healthy','extravasation_injury','kidney_healthy','kidney_low','kidney_high','liver_healthy','liver_low','liver_high','spleen_healthy','spleen_low','spleen_high','any_injury']].values\n\nval_paths = val_data['file_path'].values\nval_mask_paths = val_data['mask_path'].values\nval_labels = val_data[['bowel_healthy','bowel_injury','extravasation_healthy','extravasation_injury','kidney_healthy','kidney_low','kidney_high','liver_healthy','liver_low','liver_high','spleen_healthy','spleen_low','spleen_high','any_injury']].values\n\ntest_paths = test_data['file_path'].values\ntest_mask_paths = test_data['mask_path'].values\ntest_labels = test_data[['bowel_healthy','bowel_injury','extravasation_healthy','extravasation_injury','kidney_healthy','kidney_low','kidney_high','liver_healthy','liver_low','liver_high','spleen_healthy','spleen_low','spleen_high','any_injury']].values\n\n# Instantiate the datasets\ntrain_dataset = CustomDataset(train_paths, train_mask_paths, train_labels, transform=transform)\nprint('len of train_dataset', len(train_dataset))\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)\nprint('train_loader', len(train_loader))\n# train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\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)\nprint(train_loader)\n# print(\"Indices:\", train_loader.index)  # Print the indices\n        \n# Instantiate the model with the appropriate number of classes for both classification and segmentation\nin_channels = 1  # Input channels (e.g., for grayscale images or volumes)\nnum_classes_classification = 14  # Number of classes for classification\nnum_classes_segmentation = 5   # Number of classes for segmentation (change this according to your task)\nmodel = Custom3DViTModelTPU(in_channels, num_classes_classification, \n                            num_classes_segmentation, batch_size)\n# Count the number of parameters\ntotal_params = sum(p.numel() for p in model.parameters())\nprint(f\"Total Trainable Parameters: {total_params}\")\n#model = model.to(device)\n# model = get_model()\nmodel = model.to(xm.xla_device())\n# Define loss function and optimizer\ncriterion = nn.BCELoss()  # Binary Cross-Entropy loss\noptimizer = optim.Adam(model.parameters(), lr=0.001)\n\n# Training loop\nnum_epochs = 2\nepoch_count=0\nfor epoch in range(num_epochs):\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 train_loader:\n        optimizer.zero_grad()\n        xm.mark_step()\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        print(\"batch_images shape and batch_segmentation_masks shape\", \n              batch_images.shape, batch_segmentation_masks.shape)\n        print(\"batch number: \",epoch_count)\n        epoch_count=epoch_count+1\n        \n#         model = model.to(device)\n        # Move data to the XLA device\n        \n        batch_images = batch_images.to(torch.float32)  # Convert to float32 if not already\n        batch_segmentation_masks = batch_segmentation_masks.to(torch.float32)  # Convert to float32 if not already\n\n        batch_images = move_data_to_xla(batch_images)\n        batch_segmentation_masks = move_data_to_xla(batch_segmentation_masks)\n        batch_labels = move_data_to_xla(batch_labels)\n        \n        \n\n        # Forward pass\n        classification_outputs, segmentation_outputs = model(batch_images, batch_segmentation_masks)\n        print(\"..........................................................\")\n        \n        # Apply sigmoid activation to the classification outputs\n        classification_outputs = torch.sigmoid(classification_outputs)\n        \n        print(\"classification_outputs.shape, batch_labels shape\", \n              classification_outputs.shape, batch_labels.shape)\n        \n        \n        # Calculate binary cross-entropy loss for each class separately\n        losses = []\n        for class_index in range(num_classes_classification):\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            class_loss = criterion(class_outputs, class_labels)\n#             print(class_loss)\n            losses.append(class_loss)\n\n        # Calculate the total loss as the sum of individual class losses\n        print(sum(losses),\"sum of losses\")\n        total_loss = sum(losses)/num_classes_classification\n        print(\"total classification loss\",total_loss)\n\n        # Check if segmentation mask is available\n        if batch_segmentation_masks is not None:\n            # Ensure that both input and target tensors are of type torch.float32\n            batch_segmentation_masks = batch_segmentation_masks.to(torch.float32)\n            \n            # Apply sigmoid activation to segmentation_outputs\n            segmentation_outputs = torch.sigmoid(segmentation_outputs)\n            segmentation_outputs = segmentation_outputs.to(torch.float32)\n\n            # Calculate segmentation loss\n            segmentation_loss = criterion(segmentation_outputs, batch_segmentation_masks)\n            print(segmentation_loss)\n            total_loss += segmentation_loss\n            \n\n        running_loss += total_loss.item()\n        running_loss=running_loss/2\n        print(running_loss,\"running_loss\")\n        \n        # Calculate accuracy for each class separately\n        accuracies = []\n        for class_index in range(num_classes_classification):\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) / num_classes_classification\n        correct_train += batch_accuracy\n        total_train += 1\n        print(\"batch accuracy: \",batch_accuracy)\n        print(\"batch loss: \",running_loss)\n        \n        # Backpropagation and optimization\n        total_loss.backward()\n        xm.mark_step()\n        optimizer.step()\n\n    # Calculate and print average training accuracy and loss\n    avg_train_accuracy = correct_train / len(train_loader)\n    avg_train_loss = running_loss / len(train_loader)\n    \n    print(f\"Epoch [{epoch+1}/{num_epochs}]\")\n    print(f\"Train Accuracy: {avg_train_accuracy:.4f} | Train Loss: {avg_train_loss:.4f}\")\n        # Validation loop\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 val_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)  # Add a singleton dimension for channels\n            batch_segmentation_masks = batch_segmentation_masks.unsqueeze(1)\n            print(\"batch_images shape \", \n              batch_images.shape)\n            \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            # Calculate binary cross-entropy loss for each class separately\n            losses = []\n            for class_index in range(num_classes_classification):\n                class_labels = batch_labels[:, class_index]\n                class_outputs = classification_outputs[:, class_index]\n                class_loss = criterion(class_outputs, class_labels)\n                losses.append(class_loss)\n\n            total_loss = sum(losses)/num_classes_classification\n\n            # Check if segmentation mask is available\n            if batch_segmentation_masks is not None:\n                batch_segmentation_masks = batch_segmentation_masks.to(torch.float64)\n                \n                # Apply sigmoid activation to segmentation_outputs\n                segmentation_outputs = torch.sigmoid(segmentation_outputs)\n                segmentation_outputs = segmentation_outputs.to(torch.float64)\n                \n                segmentation_loss = criterion(segmentation_outputs, batch_segmentation_masks)\n                total_loss = total_loss + segmentation_loss\n            \n            total_val_loss += total_loss.item()\n            total_val_loss=total_val_loss/2\n\n            # Calculate accuracy for each class separately\n            accuracies = []\n            for class_index in range(num_classes_classification):\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            batch_accuracy = sum(accuracies) / num_classes_classification\n            correct_val += batch_accuracy\n            total_val += batch_labels.size(0)\n\n    val_accuracy = correct_val / len(val_loader)\n    avg_val_loss = total_val_loss / len(val_loader)\n\n    print(f\"Validation Accuracy: {val_accuracy:.4f} | Validation Loss: {avg_val_loss:.4f}\")\n\n# Test loop\nmodel.eval()\ntotal_correct = 0\ntotal_samples = 0\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        # Move your model to the TPU device\n        model = model.to(device)\n\n        # Inside your training loop or forward pass\n        batch_images = batch_images.to(device)\n        batch_segmentation_masks = batch_segmentation_masks.to(device)\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_correct = 0\n        batch_samples = batch_labels.size(0)\n        \n        for class_index in range(num_classes_classification):\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            batch_correct += class_accuracy\n            \n            # Calculate precision, recall, and F1-score for the current class\n            precision, recall, f1, _ = precision_recall_fscore_support(\n                class_labels.cpu(), predicted.cpu(), average='binary')\n            \n            precision_list.append(precision)\n            recall_list.append(recall)\n            f1_list.append(f1)\n\n        # Accumulate batch-level accuracy\n        total_correct += batch_correct/num_classes_classification\n        total_samples += batch_samples\n    \n    test_accuracy = total_correct / total_samples\n    print(f\"Test Accuracy: {test_accuracy:.4f}\")\n\n    # Calculate average precision, recall, and F1-score across all classes\n    avg_precision = sum(precision_list) / total_samples\n    avg_recall = sum(recall_list) / num_classes_classification\n    avg_f1 = sum(f1_list) / num_classes_classification\n\n    print(f\"Average Precision: {avg_precision:.4f}\")\n    print(f\"Average Recall: {avg_recall:.4f}\")\n    print(f\"Average F1 Score: {avg_f1:.4f}\")\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T13:28:01.591021Z","iopub.execute_input":"2023-12-03T13:28:01.59142Z","iopub.status.idle":"2023-12-03T13:28:47.042093Z","shell.execute_reply.started":"2023-12-03T13:28:01.591391Z","shell.execute_reply":"2023-12-03T13:28:47.040734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip list\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T13:26:10.589643Z","iopub.status.idle":"2023-12-03T13:26:10.590019Z","shell.execute_reply.started":"2023-12-03T13:26:10.58984Z","shell.execute_reply":"2023-12-03T13:26:10.589857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import os\n# import pandas as pd\n# import numpy as np\n# import nibabel as nib\n# import torch\n# import torch.nn as nn\n# import torch.optim as optim\n# import torch_xla.core.xla_model as xm\n# from torch.utils.data import Dataset, DataLoader\n# from torchvision import transforms\n# from sklearn.metrics import accuracy_score, precision_recall_fscore_support, roc_auc_score\n# from sklearn.model_selection import train_test_split\n# from torchvision import models\n# from scipy.ndimage import zoom\n# import torch.nn.functional as F\n# from PIL import Image\n\n# # Create a function to move data to the XLA device\n# def move_data_to_xla(data):\n#     return data.to(torch.float32).to(device)\n\n# class CustomDataset(Dataset):\n#     def __init__(self, image_paths, mask_paths, labels, transform=None):\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#         return len(self.image_paths)\n\n#     def __getitem__(self, idx):\n#         image_path = self.image_paths[idx]\n#         mask_path = self.mask_paths[idx]\n\n#         # Load the 3D NIfTI image using nibabel\n#         image = nib.load(image_path).get_fdata()\n# #         print(image.shape)\n# #         print(\"Data shape *********:\", image.dtype)\n\n#         # Load the segmentation mask if available\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#         # Apply transformations if provided to the image\n#         if self.transform:\n#             image = self.transform(image)\n\n#         # Apply transformations if provided to the segmentation mask\n#         if segmentation_mask is not None and self.transform:\n#             segmentation_mask = self.transform(segmentation_mask)\n#         else:\n#             segmentation_mask = torch.zeros_like(image)\n\n#         label = torch.tensor(self.labels[idx], dtype=torch.float32)\n        \n#         return image, segmentation_mask, label\n\n# import torch\n# import torch.nn as nn\n# import torchvision.models as models\n\n# # Function to resize NIfTI data\n# def 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\n\n# # Paths and settings\n# segmentation_dir = '/kaggle/input/rsna-2023-abdominal-trauma-detection/segmentations'  # Update with the correct path\n# csv_file = '/kaggle/input/abdominal-trauma-nii-csv/abdominal_trauma_nii.csv'  # Update with the correct path\n# batch_size = 32\n# num_workers = 4  # Number of CPU cores to use for data loading\n# num_classes = 14  # Number of classes\n# num_classes_segmentation = 5  # Number of classes\n# desired_shape = (128, 128, 128)\n# # device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n# device = xm.xla_device()\n# print(device)\n\n# # Define transformations if needed\n# transform = transforms.Compose([\n#     transforms.ToTensor(),  # Convert to tensor\n#     # Add more transformations if necessary\n# ])\n\n# # Load the CSV file\n# data = pd.read_csv(csv_file).head(320)\n\n# # Filter rows where the 'mask_path' column is not empty\n# #data = data[pd.notna(data['mask_path'])]\n\n# # Remove the extra space from the column name\n# data.columns = data.columns.str.strip()\n\n# # Assuming 'data' is your DataFrame\n# data_length = len(data)\n# print(\"Length of DataFrame:\", data_length)\n# # print(data)\n# # print(data.index)\n\n# # Split the data into training, validation, and test sets\n# train_data, temp_data = train_test_split(data, test_size=0.2, random_state=42)\n# val_data, test_data = train_test_split(temp_data, test_size=0.5, random_state=42)\n\n# # Set the display option to show all rows\n# pd.set_option('display.max_rows', None)\n\n# index_values = train_data.index.values\n\n# # chunk_size = 200  # You can adjust the chunk size\n# # for i in range(0, len(index_values), chunk_size):\n# #     print(index_values[i:i+chunk_size])\n\n# # Reset the display option to its default value (if needed)\n# pd.reset_option('display.max_rows')\n\n# # Extract file paths and labels from the data\n# train_paths = train_data['file_path'].values\n# train_mask_paths = train_data['mask_path'].values\n# train_labels = train_data[['bowel_healthy','bowel_injury','extravasation_healthy','extravasation_injury','kidney_healthy','kidney_low','kidney_high','liver_healthy','liver_low','liver_high','spleen_healthy','spleen_low','spleen_high','any_injury']].values\n\n# val_paths = val_data['file_path'].values\n# val_mask_paths = val_data['mask_path'].values\n# val_labels = val_data[['bowel_healthy','bowel_injury','extravasation_healthy','extravasation_injury','kidney_healthy','kidney_low','kidney_high','liver_healthy','liver_low','liver_high','spleen_healthy','spleen_low','spleen_high','any_injury']].values\n\n# test_paths = test_data['file_path'].values\n# test_mask_paths = test_data['mask_path'].values\n# test_labels = test_data[['bowel_healthy','bowel_injury','extravasation_healthy','extravasation_injury','kidney_healthy','kidney_low','kidney_high','liver_healthy','liver_low','liver_high','spleen_healthy','spleen_low','spleen_high','any_injury']].values\n\n# # Instantiate the datasets\n# train_dataset = CustomDataset(train_paths, train_mask_paths, train_labels, transform=transform)\n# print('len of train_dataset', len(train_dataset))\n# val_dataset = CustomDataset(val_paths, val_mask_paths, val_labels, transform=transform)\n# test_dataset = CustomDataset(test_paths, test_mask_paths, test_labels, transform=transform)\n\n# # Instantiate the data loaders\n# train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers, drop_last=True)\n# print('train_loader', len(train_loader))\n# # train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\n# val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\n# test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\n# print(train_loader)\n# # print(\"Indices:\", train_loader.index)  # Print the indices\n        \n# # Instantiate the model with the appropriate number of classes for both classification and segmentation\n# in_channels = 1  # Input channels (e.g., for grayscale images or volumes)\n# num_classes_classification = 14  # Number of classes for classification\n# num_classes_segmentation = 1    # Number of classes for segmentation (change this according to your task)\n# model = Custom3DViTModelTPU(in_channels, num_classes_classification, \n#                             num_classes_segmentation, batch_size)\n# # Count the number of parameters\n# total_params = sum(p.numel() for p in model.parameters())\n# print(f\"Total Trainable Parameters: {total_params}\")\n# #model = model.to(device)\n# # model = get_model()\n# model = model.to(xm.xla_device())\n# # Define loss function and optimizer\n# criterion = nn.BCELoss()  # Binary Cross-Entropy loss\n# optimizer = optim.Adam(model.parameters(), lr=0.001)\n\n# # Training loop\n# num_epochs = 3\n# epoch_count=0\n# for epoch in range(num_epochs):\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 train_loader:\n#         optimizer.zero_grad()\n#         print(\"batch_labels are\",batch_labels)\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#         print(\"batch_images shape and batch_segmentation_masks shape\", \n#               batch_images.shape, batch_segmentation_masks.shape)\n#         print(\"batch number: \",epoch_count)\n#         epoch_count=epoch_count+1\n        \n# #         model = model.to(device)\n#         # Move data to the XLA device\n        \n#         batch_images = batch_images.to(torch.float32)  # Convert to float32 if not already\n#         batch_segmentation_masks = batch_segmentation_masks.to(torch.float32)  # Convert to float32 if not already\n\n#         batch_images = move_data_to_xla(batch_images)\n#         batch_segmentation_masks = move_data_to_xla(batch_segmentation_masks)\n#         batch_labels = move_data_to_xla(batch_labels)\n        \n        \n\n#         # Forward pass\n#         classification_outputs, segmentation_outputs = model(batch_images, batch_segmentation_masks)\n#         print(\"..........................................................\")\n        \n#         # Apply sigmoid activation to the classification outputs\n#         classification_outputs = torch.sigmoid(classification_outputs)\n        \n#         print(\"classification_outputs.shape, batch_labels shape\", \n#               classification_outputs.shape, batch_labels.shape)\n        \n        \n#         # Calculate binary cross-entropy loss for each class separately\n#         losses = [criterion_classification(class_output, class_label) for class_output, class_label in zip(classification_outputs.transpose(1, 0), batch_labels.transpose(1, 0))]\n#         total_loss = sum(losses)\n\n#         # Check if segmentation mask is available\n#         if batch_segmentation_masks is not None:\n#             # Ensure that both input and target tensors are of type torch.float32\n#             batch_segmentation_masks = batch_segmentation_masks.to(torch.float32)\n\n#             # Apply sigmoid activation to segmentation_outputs\n#             segmentation_outputs = torch.sigmoid(segmentation_outputs)\n#             segmentation_outputs = segmentation_outputs.to(torch.float32)\n\n#             # Calculate segmentation loss\n#             segmentation_loss = criterion_segmentation(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#         predicted_labels = (torch.sigmoid(classification_outputs) > 0.5).float()\n#         class_accuracies = [accuracy_score(class_label.cpu(), predicted_label.cpu()) for class_label, predicted_label in zip(batch_labels.transpose(1, 0), predicted_labels.transpose(1, 0))]\n\n#         # Calculate overall accuracy\n#         batch_accuracy = sum(class_accuracies) / len(class_accuracies)\n#         correct_train += batch_accuracy\n#         total_train += 1\n\n        \n#         total_loss.backward()\n#         xm.mark_step()\n#         optimizer.step()\n\n#     # Calculate and print average training accuracy and loss\n#     avg_train_accuracy = correct_train / total_train\n#     avg_train_loss = running_loss / len(train_loader)\n    \n#     print(f\"Epoch [{epoch+1}/{num_epochs}]\")\n#     print(f\"Train Accuracy: {avg_train_accuracy:.4f} | Train Loss: {avg_train_loss:.4f}\")\n#         # Validation loop\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 val_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)  # Add a singleton dimension for channels\n#             batch_segmentation_masks = batch_segmentation_masks.unsqueeze(1)\n#             print(\"batch_images shape and batch_segmentation_masks shape\", \n#               batch_images.shape, batch_segmentation_masks.shape)\n            \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#             # Calculate binary cross-entropy loss for each class separately\n#             losses = []\n#             for class_index in range(num_classes_classification):\n#                 class_labels = batch_labels[:, class_index]\n#                 class_outputs = classification_outputs[:, class_index]\n#                 class_loss = criterion(class_outputs, class_labels)\n#                 losses.append(class_loss)\n\n#             total_loss = sum(losses)\n\n#             # Check if segmentation mask is available\n#             if batch_segmentation_masks is not None:\n#                 batch_segmentation_masks = batch_segmentation_masks.to(torch.float64)\n                \n#                 # Apply sigmoid activation to segmentation_outputs\n#                 segmentation_outputs = torch.sigmoid(segmentation_outputs)\n#                 segmentation_outputs = segmentation_outputs.to(torch.float64)\n                \n#                 segmentation_loss = criterion(segmentation_outputs, batch_segmentation_masks)\n#                 total_loss = 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(num_classes_classification):\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#             batch_accuracy = sum(accuracies) / num_classes_classification\n#             correct_val += batch_accuracy\n#             total_val += batch_labels.size(0)\n\n#     val_accuracy = correct_val / total_val\n#     avg_val_loss = total_val_loss / len(val_loader)\n\n#     print(f\"Validation Accuracy: {val_accuracy:.4f} | Validation Loss: {avg_val_loss:.4f}\")\n\n# # Test loop\n# model.eval()\n# total_correct = 0\n# total_samples = 0\n# # Initialize lists to store per-class metrics\n# precision_list = []\n# recall_list = []\n# f1_list = []\n\n# with 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#         # Move your model to the TPU device\n#         model = model.to(device)\n\n#         # Inside your training loop or forward pass\n#         batch_images = batch_images.to(device)\n#         batch_segmentation_masks = batch_segmentation_masks.to(device)\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_correct = 0\n#         batch_samples = batch_labels.size(0)\n        \n#         for class_index in range(num_classes_classification):\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#             batch_correct += class_accuracy\n            \n#             # Calculate precision, recall, and F1-score for the current class\n#             precision, recall, f1, _ = precision_recall_fscore_support(\n#                 class_labels.cpu(), predicted.cpu(), average='binary')\n            \n#             precision_list.append(precision)\n#             recall_list.append(recall)\n#             f1_list.append(f1)\n\n#         # Accumulate batch-level accuracy\n#         total_correct += batch_correct\n#         total_samples += batch_samples\n    \n#     test_accuracy = total_correct / total_samples\n#     print(f\"Test Accuracy: {test_accuracy:.4f}\")\n\n#     # Calculate average precision, recall, and F1-score across all classes\n#     avg_precision = sum(precision_list) / num_classes_classification\n#     avg_recall = sum(recall_list) / num_classes_classification\n#     avg_f1 = sum(f1_list) / num_classes_classification\n\n#     print(f\"Average Precision: {avg_precision:.4f}\")\n#     print(f\"Average Recall: {avg_recall:.4f}\")\n#     print(f\"Average F1 Score: {avg_f1:.4f}\")\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T13:26:10.591718Z","iopub.status.idle":"2023-12-03T13:26:10.59207Z","shell.execute_reply.started":"2023-12-03T13:26:10.591897Z","shell.execute_reply":"2023-12-03T13:26:10.591914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}