{"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"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install gdcm\n!pip install pylibjpeg\n!pip install pylibjpeg-libjpeg\n!pip install --upgrade pydicom\n!pip install --upgrade pylibjpeg pylibjpeg-libjpeg pylibjpeg-openjpeg\n!pip install pipdeptree\n!pip install pydicom==2.4.2\n!pip install --upgrade nibabel\n!pip install nibabel==3.2.1\n!pip install nibabel numpy scipy\n","metadata":{"execution":{"iopub.status.busy":"2023-09-03T16:41:45.68508Z","iopub.execute_input":"2023-09-03T16:41:45.68544Z","iopub.status.idle":"2023-09-03T16:42:42.127463Z","shell.execute_reply.started":"2023-09-03T16:41:45.68541Z","shell.execute_reply":"2023-09-03T16:42:42.126194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install cloud-tpu-client==0.10 https://storage.googleapis.com/tpu-pytorch/wheels/torch_xla-1.9-cp37-cp37m-linux_x86_64.whl\n","metadata":{"execution":{"iopub.status.busy":"2023-09-03T16:42:42.129478Z","iopub.execute_input":"2023-09-03T16:42:42.129799Z","iopub.status.idle":"2023-09-03T16:42:43.855501Z","shell.execute_reply.started":"2023-09-03T16:42:42.129766Z","shell.execute_reply":"2023-09-03T16:42:43.854263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install scikit-image\n","metadata":{"execution":{"iopub.status.busy":"2023-09-03T16:42:43.856952Z","iopub.execute_input":"2023-09-03T16:42:43.857274Z","iopub.status.idle":"2023-09-03T16:42:51.988974Z","shell.execute_reply.started":"2023-09-03T16:42:43.857241Z","shell.execute_reply":"2023-09-03T16:42:51.987756Z"},"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\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\nimport torch.nn.functional as F\nfrom scipy.ndimage import zoom\nfrom sklearn.model_selection import train_test_split\nimport torch_xla\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.xla_multiprocessing as xmp\nimport torch_xla.distributed.parallel_loader as pl\nimport torchvision.models as models\nimport tensorflow as tf\nfrom tensorflow.keras import layers, Model\nfrom skimage.transform import resize\nfrom tqdm import tqdm\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# Function to calculate new affine matrix after resizing\ndef calculate_new_affine(original_affine, original_shape, target_shape):\n    scale_factors = [t / o for t, o in zip(target_shape, original_shape)]\n    new_affine = np.copy(original_affine)\n    new_affine[:3, :3] = original_affine[:3, :3] * scale_factors\n    return new_affine\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        image_path = self.image_paths[idx]\n        segmentation_mask = self.labels['resized_segmentation_mask'][idx]\n\n        # Load the 3D NIfTI image using nibabel\n        image = nib.load(image_path).get_fdata()\n\n        # Load the segmentation mask if available\n        if pd.notna(mask_path):\n            segmentation_mask_nifti = nib.load(mask_path)  # Load the complete NIfTI object\n            segmentation_mask = segmentation_mask_nifti.get_fdata()\n\n            # Resize segmentation mask to desired shape\n            resized_data = resize_nifti(segmentation_mask, desired_shape)\n            resized_affine = calculate_new_affine(segmentation_mask_nifti.affine, segmentation_mask.shape, desired_shape)\n            resized_mask = nib.Nifti1Image(resized_data, affine=resized_affine)\n\n            # Convert the resized segmentation mask to a tensor\n            segmentation_mask_tensor = torch.tensor(resized_mask.get_fdata(), dtype=torch.float32)\n            segmentation_mask_tensor = F.interpolate(segmentation_mask_tensor.unsqueeze(0).unsqueeze(0), size=desired_shape,\n                                                    mode='trilinear', align_corners=False)\n            segmentation_mask_tensor = segmentation_mask_tensor.squeeze()\n        else:\n            segmentation_mask_tensor = torch.zeros(desired_shape, dtype=torch.float32)\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\n\nclass MultiLabel3DAttentionModel(nn.Module):\n    def __init__(self, num_classes, 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        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),  # Add batch normalization\n        nn.ReLU(inplace=True))\n\n        # Attention block\n        self.attention = nn.Sequential(\n            nn.Conv3d(512, 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(512, 128),\n            nn.ReLU(inplace=True),\n            nn.Linear(128, num_classes),\n            nn.Sigmoid()\n        )\n        \n        # Segmentation head\n        self.segmentation_head = nn.Sequential(\n            nn.Conv3d(512, 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        features = self.backbone(x)\n        \n        # Apply attention to features\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) * segmentation_mask\n        \n        return classification_output, segmentation_output\n     \n# Paths and settings\nsegmentation_dir = '/kaggle/input/rsna-2022-cervical-spine-fracture-detection/segmentations'\ncsv_file = '/kaggle/input/file-mask-path/train_file_mask_path_2 (1).csv'\nbatch_size = 16\nnum_workers = 4\nnum_classes_classification = 7\nnum_classes_segmentation = 1\ndesired_shape = (128, 128, 128)\n\n# Define transformations\ntransform = transforms.Compose([\n    transforms.ToTensor(),\n])\n\n# Load the CSV file\ndata = pd.read_csv(csv_file)\ndata.columns = data.columns.str.strip()\n\n\n# Initialize TPU strategy\ntpu_resolver = tf.distribute.cluster_resolver.TPUClusterResolver()\ntf.config.experimental_connect_to_cluster(tpu_resolver)\ntf.tpu.experimental.initialize_tpu_system(tpu_resolver)\nstrategy = tf.distribute.TPUStrategy(tpu_resolver)\n\n# Load and resize segmentation masks\nsegmentation_masks = []\nfor mask_path in data['mask_path']:\n    if pd.notna(mask_path):\n        segmentation_mask_nifti = nib.load(mask_path)  # Load the complete NIfTI object\n        segmentation_mask = segmentation_mask_nifti.get_fdata()\n        resized_data = resize_nifti(segmentation_mask, desired_shape)\n        resized_affine = calculate_new_affine(segmentation_mask_nifti.affine, segmentation_mask.shape, desired_shape)\n        resized_mask = nib.Nifti1Image(resized_data, affine=resized_affine)\n        segmentation_mask_tensor = torch.tensor(resized_mask.get_fdata(), dtype=torch.float32)\n        segmentation_mask_tensor = F.interpolate(segmentation_mask_tensor.unsqueeze(0).unsqueeze(0), size=desired_shape,\n                                                mode='trilinear', align_corners=False)\n        segmentation_mask_tensor = segmentation_mask_tensor.squeeze()\n    else:\n        segmentation_mask_tensor = torch.zeros(desired_shape, dtype=torch.float32)\n\n    segmentation_masks.append(segmentation_mask_tensor)\n\n\n# Add the resized segmentation masks to the data DataFrame\ndata['resized_segmentation_mask'] = segmentation_masks\n\n# Split the data\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# Instantiate the datasets\ntrain_dataset = CustomDataset(train_paths, train_labels, transform=transform)\nval_dataset = CustomDataset(val_paths, val_labels, transform=transform)\ntest_dataset = CustomDataset(test_paths, test_labels, transform=transform)\n\n# Define data loaders using TensorFlow's Dataset API\nbatch_size = 16  # You can adjust this as needed\ntrain_dataset = CustomDataset(train_data['file_path'].values, train_data['mask_path'].values,\n                              train_data[['C1', 'C2', 'C3', 'C4', 'C5', 'C6', 'C7']].values, transform=transform)\nval_dataset = CustomDataset(val_data['file_path'].values, val_data['mask_path'].values,\n                            val_data[['C1', 'C2', 'C3', 'C4', 'C5', 'C6', 'C7']].values, transform=transform)\ntest_dataset = CustomDataset(test_data['file_path'].values, test_data['mask_path'].values,\n                             test_data[['C1', 'C2', 'C3', 'C4', 'C5', 'C6', 'C7']].values, transform=transform)\n\ntrain_loader = tf.data.Dataset.from_tensor_slices((train_dataset.inputs, train_dataset.targets)).batch(batch_size)\nval_loader = tf.data.Dataset.from_tensor_slices((val_dataset.inputs, val_dataset.targets)).batch(batch_size)\ntest_loader = tf.data.Dataset.from_tensor_slices((test_dataset.inputs, test_dataset.targets)).batch(batch_size)\n\n# Define your model, loss function, and optimizer\nwith strategy.scope():\n    model = MultiLabel3DAttentionModel(num_classes_classification, num_classes_segmentation)\n    optimizer = tf.keras.optimizers.Adam(learning_rate=0.001)\n    loss_fn = tf.keras.losses.BinaryCrossentropy()\n\n# Training loop\nnum_epochs = 1\n\nfor epoch in range(num_epochs):\n    # Training\n    total_loss = 0.0\n    total_correct_train = 0\n    total_train_samples = 0\n\n    for batch_images, batch_labels in tqdm(train_loader):\n        with tf.GradientTape() as tape:\n            classification_outputs, segmentation_outputs = model(batch_images, batch_segmentation_masks)\n            classification_loss = loss_fn(batch_labels, classification_outputs)\n\n            if batch_segmentation_masks is not None:\n                segmentation_loss = loss_fn(batch_segmentation_masks, segmentation_outputs)\n                loss = classification_loss + segmentation_loss\n            else:\n                loss = classification_loss\n\n        gradients = tape.gradient(loss, model.trainable_variables)\n        optimizer.apply_gradients(zip(gradients, model.trainable_variables))\n\n        total_loss += loss\n        predicted = tf.math.argmax(classification_outputs, axis=1)\n        total_correct_train += tf.reduce_sum(tf.cast(tf.equal(predicted, batch_labels), tf.int32))\n        total_train_samples += batch_labels.shape[0]\n\n    train_accuracy = total_correct_train / total_train_samples\n    avg_train_loss = total_loss / len(train_loader)\n\n    print(f\"Epoch [{epoch+1}/{num_epochs}]\")\n    print(f\"Train Accuracy: {train_accuracy:.4f} | Train Loss: {avg_train_loss:.4f}\")\n\n    # Validation\n    total_correct_val = 0\n    total_val_samples = 0\n    total_val_loss = 0.0\n\n    for batch_images, batch_labels in tqdm(val_loader):\n        classification_outputs, _ = model(batch_images, batch_segmentation_masks)\n        val_loss = loss_fn(batch_labels, classification_outputs)\n        total_val_loss += val_loss\n        predicted = tf.math.argmax(classification_outputs, axis=1)\n        total_correct_val += tf.reduce_sum(tf.cast(tf.equal(predicted, batch_labels), tf.int32))\n        total_val_samples += batch_labels.shape[0]\n\n    val_accuracy = total_correct_val / total_val_samples\n    avg_val_loss = total_val_loss / len(val_loader)\n\n    print(f\"Epoch [{epoch+1}/{num_epochs}]\")\n    print(f\"Validation Accuracy: {val_accuracy:.4f} | Validation Loss: {avg_val_loss:.4f}\")\n\n# Test loop\ntotal_correct_test = 0\ntotal_test_samples = 0\ntotal_test_loss = 0.0\n\nfor batch_images, batch_labels in tqdm(test_loader):\n    classification_outputs, _ = model(batch_images, batch_segmentation_masks)\n    test_loss = loss_fn(batch_labels, classification_outputs)\n    total_test_loss += test_loss\n    predicted = tf.math.argmax(classification_outputs, axis=1)\n    total_correct_test += tf.reduce_sum(tf.cast(tf.equal(predicted, batch_labels), tf.int32))\n    total_test_samples += batch_labels.shape[0]\n\ntest_accuracy = total_correct_test / total_test_samples\navg_test_loss = total_test_loss / len(test_loader)\n\nprint(f\"Test Accuracy: {test_accuracy:.4f} | Test Loss: {avg_test_loss:.4f}\")\n\n\n","metadata":{"execution":{"iopub.status.busy":"2023-09-03T16:42:51.991649Z","iopub.execute_input":"2023-09-03T16:42:51.991977Z"},"trusted":true},"execution_count":null,"outputs":[]}],"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}}