{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9538404,"sourceType":"datasetVersion","datasetId":5810054},{"sourceId":128188,"sourceType":"modelInstanceVersion","modelInstanceId":107961,"modelId":132298}],"dockerImageVersionId":30775,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/pydicom1')  #\nimport pydicom\nprint(pydicom.__version__)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-10-08T23:52:14.596539Z","iopub.execute_input":"2024-10-08T23:52:14.596996Z","iopub.status.idle":"2024-10-08T23:53:55.012120Z","shell.execute_reply.started":"2024-10-08T23:52:14.596948Z","shell.execute_reply":"2024-10-08T23:53:55.011237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nimport numpy as np\nimport pandas as pd\nimport os\nimport cv2\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score, precision_score, recall_score, f1_score\nimport tensorflow as tf\nfrom tensorflow.keras.applications import EfficientNetB0, ResNet50\nfrom tensorflow.keras.layers import Dense, GlobalAveragePooling2D, Dropout, BatchNormalization\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.optimizers import Adam\nfrom tensorflow.keras.callbacks import ReduceLROnPlateau, EarlyStopping\n\n# Load CSV files\ntrain_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\ncoord_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv')\nseries_des_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')\n\n# Directory for the images\nimage_dir = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/'\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:53:55.014119Z","iopub.execute_input":"2024-10-08T23:53:55.014689Z","iopub.status.idle":"2024-10-08T23:54:09.477425Z","shell.execute_reply.started":"2024-10-08T23:53:55.014639Z","shell.execute_reply":"2024-10-08T23:54:09.476373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\n\ndef reshape_row(row):\n    data = {'study_id':[], 'condition':[], 'level':[], 'severity':[]} # Create the structure for tidy format\n\n    for column, value in row.items():\n        if column not in ['study_id']:  # Ignore study_id for reshaping\n            parts = column.split('_')\n            condition = ' '.join([word.capitalize() for word in parts[:-2]])  # e.g., 'Spinal Canal Stenosis'\n            level = parts[-2].capitalize() + '/' + parts[-1].capitalize()  # e.g., 'L1/L2'\n            data['study_id'].append(row['study_id'])\n            data['condition'].append(condition)\n            data['level'].append(level)\n            data['severity'].append(value)\n\n    return pd.DataFrame(data)\n\n# Apply reshape_row to each row in the training data\n\nnew_train_df = pd.concat([reshape_row(row) for _, row in train_df.iterrows()], ignore_index=True)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:09.479034Z","iopub.execute_input":"2024-10-08T23:54:09.479356Z","iopub.status.idle":"2024-10-08T23:54:11.012784Z","shell.execute_reply.started":"2024-10-08T23:54:09.479322Z","shell.execute_reply":"2024-10-08T23:54:11.011931Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n# Merge new_train_df with label coordinates based on study_id, condition, and level\nmerged_df = pd.merge(new_train_df, coord_df, on=['study_id', 'condition', 'level'], how='inner')\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:11.014760Z","iopub.execute_input":"2024-10-08T23:54:11.015062Z","iopub.status.idle":"2024-10-08T23:54:11.090633Z","shell.execute_reply.started":"2024-10-08T23:54:11.015028Z","shell.execute_reply":"2024-10-08T23:54:11.089812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n# Merge with series descriptions\nfinal_merged_df = pd.merge(merged_df, series_des_df, on=['study_id', 'series_id'], how='inner')\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:11.091729Z","iopub.execute_input":"2024-10-08T23:54:11.092027Z","iopub.status.idle":"2024-10-08T23:54:11.119925Z","shell.execute_reply.started":"2024-10-08T23:54:11.091995Z","shell.execute_reply":"2024-10-08T23:54:11.119220Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create unique row_id based on study_id, condition, and level\nfinal_merged_df['row_id'] = (\n    final_merged_df['study_id'].astype(str) + '_' +\n    final_merged_df['condition'].str.lower().str.replace(' ', '_') + '_' +\n    final_merged_df['level'].str.lower().str.replace('/', '_')\n)\n\n# Add image_path column\nfinal_merged_df['image_path'] = (\n    f'{image_dir}/' + \n    final_merged_df['study_id'].astype(str) + '/' +\n    final_merged_df['series_id'].astype(str) + '/' +\n    final_merged_df['instance_number'].astype(str) + '.dcm'\n)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:11.120932Z","iopub.execute_input":"2024-10-08T23:54:11.121201Z","iopub.status.idle":"2024-10-08T23:54:11.341374Z","shell.execute_reply.started":"2024-10-08T23:54:11.121172Z","shell.execute_reply":"2024-10-08T23:54:11.340546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df['severity'] = final_merged_df['severity'].map({'Normal/Mild': 'normal_mild', 'Moderate': 'moderate', 'Severe': 'severe'})\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:11.342556Z","iopub.execute_input":"2024-10-08T23:54:11.342865Z","iopub.status.idle":"2024-10-08T23:54:11.353974Z","shell.execute_reply.started":"2024-10-08T23:54:11.342832Z","shell.execute_reply":"2024-10-08T23:54:11.353058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check condition distribution across series descriptions\npd.crosstab(final_merged_df['condition'], final_merged_df['series_description'])\n\n# Check severity value counts\nfinal_merged_df['severity'].value_counts()\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:11.355470Z","iopub.execute_input":"2024-10-08T23:54:11.355831Z","iopub.status.idle":"2024-10-08T23:54:11.401194Z","shell.execute_reply.started":"2024-10-08T23:54:11.355789Z","shell.execute_reply":"2024-10-08T23:54:11.400311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:11.402306Z","iopub.execute_input":"2024-10-08T23:54:11.402618Z","iopub.status.idle":"2024-10-08T23:54:11.420210Z","shell.execute_reply.started":"2024-10-08T23:54:11.402586Z","shell.execute_reply":"2024-10-08T23:54:11.419437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport numpy as np\nimport os\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:11.424278Z","iopub.execute_input":"2024-10-08T23:54:11.424577Z","iopub.status.idle":"2024-10-08T23:54:16.382243Z","shell.execute_reply.started":"2024-10-08T23:54:11.424545Z","shell.execute_reply":"2024-10-08T23:54:16.381444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\n# Define the device based on availability of GPU\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:16.383321Z","iopub.execute_input":"2024-10-08T23:54:16.383841Z","iopub.status.idle":"2024-10-08T23:54:16.425610Z","shell.execute_reply.started":"2024-10-08T23:54:16.383807Z","shell.execute_reply":"2024-10-08T23:54:16.424475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Import necessary libraries\nimport torch\nfrom torch.utils.data import Dataset\nimport numpy as np\nimport pydicom\nimport cv2\nfrom torchvision import transforms\n\n# Define a custom dataset class for multi-label, multi-class classification \nclass SpineConditionDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        \"\"\"\n        Initialize the SpineConditionDataset with a dataframe and optional transforms.\n        \"\"\"\n        self.dataframe = dataframe\n        self.transform = transform\n        self.level_map = {level: idx for idx, level in enumerate(['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1'])}\n        self.severity_map = {severity: idx for idx, severity in enumerate(['normal_mild', 'moderate', 'severe'])}\n\n    def __len__(self):\n        \"\"\"\n        Returns the total number of samples in the dataset.\n        \"\"\"\n        return len(self.dataframe)\n\n    def __getitem__(self, index):\n        \"\"\"\n        Retrieves the image and label for a given index.\n        \"\"\"\n        image_path = self.dataframe['image_path'].iloc[index]\n        image = self.load_image(image_path)\n        level = self.dataframe['level'].iloc[index]\n        severity = self.dataframe['severity'].iloc[index]\n\n        label = self.create_label(level, severity)\n\n        if self.transform:\n            image = self.transform(image)\n        \n        return image, label\n\n    def load_image(self, image_path):\n        \"\"\"\n        Load and preprocess the DICOM image, converting it to uint8 for compatibility with transforms.\n        \"\"\"\n        dicom = pydicom.dcmread(image_path)\n        img = dicom.pixel_array\n    \n        # Convert uint16 to uint8 (just to ensure compatibility, no scaling)\n        img = np.clip(img, 0, 255).astype(np.uint8)\n    \n        # Convert the grayscale image to a 3-channel RGB image\n        img_rgb = np.stack([img] * 3, axis=-1)  # Stack grayscale to RGB\n\n        return img_rgb\n\n    def create_label(self, level, severity):\n        \"\"\"\n        Create a multi-class label based on the spinal level and severity.\n        \"\"\"\n        label = np.zeros(5, dtype=int)\n        if level in self.level_map and severity in self.severity_map:\n            label[self.level_map[level]] = self.severity_map[severity]\n        return torch.tensor(label, dtype=torch.long)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:16.426906Z","iopub.execute_input":"2024-10-08T23:54:16.427277Z","iopub.status.idle":"2024-10-08T23:54:16.441854Z","shell.execute_reply.started":"2024-10-08T23:54:16.427239Z","shell.execute_reply":"2024-10-08T23:54:16.441060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the transforms for ResNet-18 preprocessing\ntransform = transforms.Compose([\n    transforms.ToPILImage(),  # Convert to PIL image\n    transforms.Resize((224, 224)),  # Resize the image\n    transforms.ToTensor(),  # Convert image to tensor, range [0, 1]\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),  # Normalize using ImageNet stats\n])\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:16.442889Z","iopub.execute_input":"2024-10-08T23:54:16.443173Z","iopub.status.idle":"2024-10-08T23:54:16.455829Z","shell.execute_reply.started":"2024-10-08T23:54:16.443133Z","shell.execute_reply":"2024-10-08T23:54:16.454917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\n\n# Assuming `final_merged_df` contains the following columns:\n# study_id, series_description, image_path\n\n# Define the list of expected series descriptions\nexpected_series = ['Sagittal T1', 'Axial T2', 'Sagittal T2/STIR']\n\n# Group by study_id and collect the series descriptions\ngrouped_series = final_merged_df.groupby('study_id')['series_description'].apply(set).reset_index()\n\n# Count how many series each study_id has\ngrouped_series['num_series'] = grouped_series['series_description'].apply(lambda x: len(x.intersection(expected_series)))\n\n# Filter for studies missing one or more series\nmissing_series = grouped_series[grouped_series['num_series'] < len(expected_series)]\n\n# Print summary\nprint(f\"Total studies: {len(grouped_series)}\")\nprint(f\"Studies with all series: {len(grouped_series[grouped_series['num_series'] == len(expected_series)])}\")\nprint(f\"Studies missing one or more series: {len(missing_series)}\")\nprint(f\"Breakdown of missing series:\\n{missing_series['num_series'].value_counts()}\")\n\n# Optionally, you can inspect the missing cases\nprint(missing_series.head())\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:16.456794Z","iopub.execute_input":"2024-10-08T23:54:16.457101Z","iopub.status.idle":"2024-10-08T23:54:16.540752Z","shell.execute_reply.started":"2024-10-08T23:54:16.457067Z","shell.execute_reply":"2024-10-08T23:54:16.539826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\n    Create train and validation dataloaders for each series, excluding studies with missing series.\n    Args:\n        df (pd.DataFrame): Merged dataframe containing all series descriptions and labels.\n        series_descriptions (list): List of required series descriptions.\n        transform (torchvision.transforms): Image transformations.\n        batch_size (int): Batch size for the dataloaders.\n    Returns:\n        dict: Dataloaders for each series.\n        dict: Sizes of the train and validation datasets for each series.\n    \"\"\"\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:16.541774Z","iopub.execute_input":"2024-10-08T23:54:16.542041Z","iopub.status.idle":"2024-10-08T23:54:16.548442Z","shell.execute_reply.started":"2024-10-08T23:54:16.542010Z","shell.execute_reply":"2024-10-08T23:54:16.547225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Function to filter out studies with missing series and create datasets and dataloaders\ndef create_dataloaders(df, series_descriptions, transform, batch_size=16):\n    \n    # Step 1: Filter out studies with missing series\n    # Group by study_id to count unique series descriptions for each study\n    grouped_series = df.groupby('study_id')['series_description'].apply(set).reset_index()\n    grouped_series['num_series'] = grouped_series['series_description'].apply(len)\n\n    # Define the expected series set\n    expected_series = set(series_descriptions)\n\n    # Filter out studies that do not have all the required series descriptions\n    valid_studies = grouped_series[grouped_series['num_series'] == len(expected_series)]['study_id']\n    filtered_df = df[df['study_id'].isin(valid_studies)]\n\n    # Step 2: Create dataloaders for the filtered studies with all required series\n    dataloaders = {}\n    lengths = {}\n\n    for series in series_descriptions:\n        # Filter dataframe for the current series description\n        series_df = filtered_df[filtered_df['series_description'] == series]\n\n        # Split into training and validation sets\n        train_df, val_df = train_test_split(series_df, test_size=0.2, random_state=42)\n\n        # Create datasets using the custom SpineConditionDataset class\n        train_dataset = SpineConditionDataset(train_df, transform)\n        val_dataset = SpineConditionDataset(val_df, transform)\n\n        # Create dataloaders\n        trainloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\n        valloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)\n\n        # Store dataloaders and their sizes\n        dataloaders[series] = (trainloader, valloader)\n        lengths[series] = (len(train_df), len(val_df))\n\n    return dataloaders, lengths, filtered_df\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:16.549562Z","iopub.execute_input":"2024-10-08T23:54:16.549811Z","iopub.status.idle":"2024-10-08T23:54:16.559125Z","shell.execute_reply.started":"2024-10-08T23:54:16.549783Z","shell.execute_reply":"2024-10-08T23:54:16.558346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# List of series descriptions to create dataloaders for\nseries_descriptions = ['Sagittal T1', 'Axial T2', 'Sagittal T2/STIR']\n\n# Create dataloaders for each series description\ndataloaders, lengths,filtered_df = create_dataloaders(final_merged_df, series_descriptions, transform)\n\n# Print dataset statistics\nfor series, (train_loader, val_loader) in dataloaders.items():\n    print(f\"Series: {series} - Training Size: {lengths[series][0]}, Validation Size: {lengths[series][1]}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:16.560416Z","iopub.execute_input":"2024-10-08T23:54:16.560683Z","iopub.status.idle":"2024-10-08T23:54:16.697573Z","shell.execute_reply.started":"2024-10-08T23:54:16.560653Z","shell.execute_reply":"2024-10-08T23:54:16.696406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\nimport torch\n\n# Function to visualize images and check if transforms are applied\ndef visualize_dataloader(dataloader, series_name):\n    print(f\"\\nVisualizing samples from: {series_name}\")\n    batch = next(iter(dataloader))  # Get a batch from the dataloader\n    images, labels = batch\n\n    # Check the shape of the first image to confirm resizing (it should be [3, 224, 224] for RGB and 224x224 size)\n    print(f\"Image shape: {images[0].shape}\")\n    \n    # Check mean and standard deviation to verify normalization\n    print(f\"Image mean (first image): {images[0].mean().item()}\")\n    print(f\"Image std (first image): {images[0].std().item()}\")\n    \n    # Plot first few images\n    fig, ax = plt.subplots(1, 5, figsize=(20, 5))\n    for i in range(5):\n        img = images[i].numpy().transpose(1, 2, 0)  # Transpose to (H, W, C) format for plotting\n        img = np.clip(img * 0.229 + 0.485, 0, 1)  # De-normalize the image for visualization\n        ax[i].imshow(img)\n        ax[i].set_title(f\"Label: {labels[i].numpy()}\")\n        ax[i].axis('off')\n    plt.show()\n\n# Visualize a batch from each dataloader to check transforms\nfor series, (train_loader, val_loader) in dataloaders.items():\n    print(f\"Checking {series} - Train Loader\")\n    visualize_dataloader(train_loader, series)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:16.698803Z","iopub.execute_input":"2024-10-08T23:54:16.699104Z","iopub.status.idle":"2024-10-08T23:54:20.009475Z","shell.execute_reply.started":"2024-10-08T23:54:16.699072Z","shell.execute_reply":"2024-10-08T23:54:20.008459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create a summary to check for missing series in the filtered dataframe\nseries_summary = final_merged_df.groupby('series_description')['study_id'].nunique().reset_index()\nseries_summary.columns = ['series_description', 'num_studies']\n\n# Print the series summary\nprint(\"Series Summary Before Filtering:\")\nprint(series_summary)\n\n# Check which series are present in the filtered dataframe\nfiltered_series_summary = filtered_df.groupby('series_description')['study_id'].nunique().reset_index()\nfiltered_series_summary.columns = ['series_description', 'num_studies']\n\n# Print the filtered series summary\nprint(\"\\nSeries Summary After Filtering:\")\nprint(filtered_series_summary)\n\n# Check for any missing series\nmissing_series = set(series_descriptions) - set(filtered_series_summary['series_description'])\nif missing_series:\n    print(\"\\nMissing Series:\")\n    print(missing_series)\nelse:\n    print(\"\\nNo Series are Missing After Filtering.\")\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:20.010679Z","iopub.execute_input":"2024-10-08T23:54:20.011063Z","iopub.status.idle":"2024-10-08T23:54:20.035341Z","shell.execute_reply.started":"2024-10-08T23:54:20.011029Z","shell.execute_reply":"2024-10-08T23:54:20.034519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\n\n# Channel Attention Mechanism\nclass ChannelAttention(nn.Module):\n    def __init__(self, in_channels, reduction_ratio=16):\n        super(ChannelAttention, self).__init__()\n        self.fc1 = nn.Linear(in_channels, in_channels // reduction_ratio, bias=False)\n        self.relu = nn.ReLU()\n        self.fc2 = nn.Linear(in_channels // reduction_ratio, in_channels, bias=False)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg_out = torch.mean(x, dim=[2, 3], keepdim=True)\n        avg_out = avg_out.view(avg_out.size(0), -1)\n        avg_out = self.fc1(avg_out)\n        avg_out = self.relu(avg_out)\n        avg_out = self.fc2(avg_out)\n        avg_out = self.sigmoid(avg_out).view(x.size(0), -1, 1, 1)\n        return x * avg_out\n\n# Multi-Series Multi-Input Model with Attention Mechanism\nclass MultiInputCNNWithAttention(nn.Module):\n    def __init__(self, num_classes=3, dropout_prob=0.5, resnet_path='/kaggle/input/resnet18/pytorch/default/1/resnet18-f37072fd.pth'):\n        super(MultiInputCNNWithAttention, self).__init__()\n\n        self.num_classes = num_classes\n        self.dropout_prob = dropout_prob\n\n        # Load ResNet-18 architecture for each input series\n        self.resnet_t1 = models.resnet18()\n        self.resnet_t2 = models.resnet18()\n        self.resnet_t2stir = models.resnet18()\n\n        # Load pre-trained weights\n        resnet_weights = torch.load(resnet_path)\n        self.resnet_t1.load_state_dict(resnet_weights)\n        self.resnet_t2.load_state_dict(resnet_weights)\n        self.resnet_t2stir.load_state_dict(resnet_weights)\n\n        # Modify the fully connected layers\n        self.resnet_t1.fc = nn.Identity()\n        self.resnet_t2.fc = nn.Identity()\n        self.resnet_t2stir.fc = nn.Identity()\n\n        # Attention layer for each series before pooling\n        self.attention_t1 = ChannelAttention(in_channels=512)\n        self.attention_t2 = ChannelAttention(in_channels=512)\n        self.attention_t2stir = ChannelAttention(in_channels=512)\n\n        # Batch Normalization for the features from each series\n        self.bn_t1 = nn.BatchNorm1d(512)\n        self.bn_t2 = nn.BatchNorm1d(512)\n        self.bn_t2stir = nn.BatchNorm1d(512)\n\n        # Fusion layer to combine features from the series (adjust size dynamically)\n        # Initialize with max size for 3 series, but adjust dynamically\n        self.fc_fusion = nn.Linear(512 * 3, 1024)\n\n        # Attention mechanism for the fused features\n        self.attention_layer = nn.Sequential(\n            nn.Linear(1024, 512),\n            nn.Tanh(),\n            nn.Linear(512, 1024),\n            nn.Softmax(dim=1)\n        )\n\n        # Batch Normalization and Dropout after fusion\n        self.bn_fusion = nn.BatchNorm1d(1024)\n        self.dropout_fusion = nn.Dropout(self.dropout_prob)\n\n        # Additional fully connected layers\n        self.fc_block1 = nn.Linear(1024, 512)\n        self.bn_block1 = nn.BatchNorm1d(512)\n        self.dropout_block1 = nn.Dropout(self.dropout_prob)\n        self.residual_projection1 = nn.Linear(1024, 512)\n\n        self.fc_block2 = nn.Linear(512, 256)\n        self.bn_block2 = nn.BatchNorm1d(256)\n        self.dropout_block2 = nn.Dropout(self.dropout_prob)\n        self.residual_projection2 = nn.Linear(512, 256)\n\n        # Output layer for 3 series, 5 levels, and 3 severity classes\n        self.fc_output = nn.Linear(256, num_classes * 5 * 3)\n\n        # Activation function\n        self.leaky_relu = nn.LeakyReLU(negative_slope=0.01)\n\n    def forward(self, x_t1=None, x_t2=None, x_t2stir=None):\n        features = []\n        num_inputs = 0\n\n        # Extract features from T1 series if not None\n        if x_t1 is not None:\n            num_inputs += 1\n            features_t1 = self.resnet_t1.conv1(x_t1)\n            features_t1 = self.resnet_t1.bn1(features_t1)\n            features_t1 = self.resnet_t1.relu(features_t1)\n            features_t1 = self.resnet_t1.maxpool(features_t1)\n            features_t1 = self.resnet_t1.layer1(features_t1)\n            features_t1 = self.resnet_t1.layer2(features_t1)\n            features_t1 = self.resnet_t1.layer3(features_t1)\n            features_t1 = self.resnet_t1.layer4(features_t1)\n            \n            # Apply attention mechanism\n            features_t1 = self.attention_t1(features_t1)\n            # Global Average Pooling\n            features_t1 = torch.mean(features_t1, dim=[2, 3])\n            # Batch Normalization\n            features_t1 = self.bn_t1(features_t1)\n            \n            features.append(features_t1)\n\n        # Extract features from T2 series if not None\n        if x_t2 is not None:\n            num_inputs += 1\n            features_t2 = self.resnet_t2.conv1(x_t2)\n            features_t2 = self.resnet_t2.bn1(features_t2)\n            features_t2 = self.resnet_t2.relu(features_t2)\n            features_t2 = self.resnet_t2.maxpool(features_t2)\n            features_t2 = self.resnet_t2.layer1(features_t2)\n            features_t2 = self.resnet_t2.layer2(features_t2)\n            features_t2 = self.resnet_t2.layer3(features_t2)\n            features_t2 = self.resnet_t2.layer4(features_t2)\n            \n            # Apply attention mechanism\n            features_t2 = self.attention_t2(features_t2)\n            # Global Average Pooling\n            features_t2 = torch.mean(features_t2, dim=[2, 3])\n            # Batch Normalization\n            features_t2 = self.bn_t2(features_t2)\n            \n            features.append(features_t2)\n\n        # Extract features from T2STIR series if not None\n        if x_t2stir is not None:\n            num_inputs += 1\n            features_t2stir = self.resnet_t2stir.conv1(x_t2stir)\n            features_t2stir = self.resnet_t2stir.bn1(features_t2stir)\n            features_t2stir = self.resnet_t2stir.relu(features_t2stir)\n            features_t2stir = self.resnet_t2stir.maxpool(features_t2stir)\n            features_t2stir = self.resnet_t2stir.layer1(features_t2stir)\n            features_t2stir = self.resnet_t2stir.layer2(features_t2stir)\n            features_t2stir = self.resnet_t2stir.layer3(features_t2stir)\n            features_t2stir = self.resnet_t2stir.layer4(features_t2stir)\n            \n            # Apply attention mechanism\n            features_t2stir = self.attention_t2stir(features_t2stir)\n            # Global Average Pooling\n            features_t2stir = torch.mean(features_t2stir, dim=[2, 3])\n            # Batch Normalization\n            features_t2stir = self.bn_t2stir(features_t2stir)\n            \n            features.append(features_t2stir)\n\n        # Concatenate features if any series input is provided\n        if len(features) > 1:\n            combined_features = torch.cat(features, dim=1)\n        elif len(features) == 1:\n            combined_features = features[0]\n        else:\n            raise ValueError(\"At least one input (x_t1, x_t2, x_t2stir) must be provided.\")\n\n        # Adjust fc_fusion layer input size based on the number of inputs\n        input_size = 512 * num_inputs\n        self.fc_fusion = nn.Linear(input_size, 1024).to(combined_features.device)\n\n        # Pass through the fusion layer\n        fused = self.fc_fusion(combined_features)\n        fused = self.bn_fusion(fused)\n        fused = self.leaky_relu(fused)\n        fused = self.dropout_fusion(fused)        \n\n        # Apply attention mechanism to fused features\n        attention_weights = self.attention_layer(fused)\n        fused = fused * attention_weights\n\n        # Residual connection with projection for matching dimensions\n        residual1 = self.residual_projection1(fused)\n        x = self.fc_block1(fused)\n        x = self.bn_block1(x)\n        x = self.leaky_relu(x)\n        x = self.dropout_block1(x)\n        x = x + residual1\n\n        # Second residual block\n        residual2 = self.residual_projection2(x)\n        x = self.fc_block2(x)\n        x = self.bn_block2(x)\n        x = self.leaky_relu(x)\n        x = self.dropout_block2(x)\n        x = x + residual2\n\n        # Final output layer\n        output = self.fc_output(x)\n\n        # Reshape the flat output into (batch_size, 3, 5, num_classes)\n        output = output.view(output.size(0), 3, 5, self.num_classes)\n\n        return output\n\n# Example usage\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = MultiInputCNNWithAttention(num_classes=3, dropout_prob=0.5, resnet_path='/kaggle/input/resnet18/pytorch/default/1/resnet18-f37072fd.pth').to(device)\n\n\n        ","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:20.036603Z","iopub.execute_input":"2024-10-08T23:54:20.036893Z","iopub.status.idle":"2024-10-08T23:54:21.434657Z","shell.execute_reply.started":"2024-10-08T23:54:20.036862Z","shell.execute_reply":"2024-10-08T23:54:21.433641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Ensure all BatchNorm layers are in training mode\nmodel.resnet_t1.train()\nmodel.resnet_t2.train()\nmodel.resnet_t2stir.train()\n\n# Freeze all layers except the final block for fine-tuning\nfor param in model.resnet_t1.parameters():\n    param.requires_grad = False\nfor param in model.resnet_t2.parameters():\n    param.requires_grad = False\nfor param in model.resnet_t2stir.parameters():\n    param.requires_grad = False\n\n# Unfreeze the last few layers (layer4) for fine-tuning\nfor param in model.resnet_t1.layer4.parameters():\n    param.requires_grad = True\nfor param in model.resnet_t2.layer4.parameters():\n    param.requires_grad = True\nfor param in model.resnet_t2stir.layer4.parameters():\n    param.requires_grad = True\nfor module in [model.resnet_t1, model.resnet_t2, model.resnet_t2stir]:\n    for m in module.modules():\n        if isinstance(m, nn.BatchNorm1d):\n            m.requires_grad_(True)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:21.435898Z","iopub.execute_input":"2024-10-08T23:54:21.436222Z","iopub.status.idle":"2024-10-08T23:54:21.447123Z","shell.execute_reply.started":"2024-10-08T23:54:21.436188Z","shell.execute_reply":"2024-10-08T23:54:21.446137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = torch.optim.Adam([\n    {'params': model.resnet_t1.layer4.parameters(), 'lr': 5e-5},\n    {'params': model.resnet_t2.layer4.parameters(), 'lr': 5e-5},\n    {'params': model.resnet_t2stir.layer4.parameters(), 'lr': 5e-5},\n    {'params': model.fc_fusion.parameters(), 'lr': 1e-4}, \n    {'params': model.fc_block1.parameters(), 'lr': 1e-4},  \n    {'params': model.bn_block1.parameters(), 'lr': 1e-4},  \n    {'params': model.fc_block2.parameters(), 'lr': 1e-4},  \n    {'params': model.bn_block2.parameters(), 'lr': 1e-4},  \n    {'params': model.fc_output.parameters(), 'lr': 5e-4}  \n], weight_decay=1e-4)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:21.448427Z","iopub.execute_input":"2024-10-08T23:54:21.448742Z","iopub.status.idle":"2024-10-08T23:54:21.461069Z","shell.execute_reply.started":"2024-10-08T23:54:21.448712Z","shell.execute_reply":"2024-10-08T23:54:21.460307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"weights = torch.tensor([0.431, 5, 15]).to(device)  # Further increase minority class weights\n # Further increase minority class weights\nimport torch\nimport torch.nn as nn\n\n# Define Focal Loss\nclass FocalLoss(nn.Module):\n    def __init__(self, weights=None, gamma=4):\n        super(FocalLoss, self).__init__()\n        self.gamma = gamma\n        self.weights = weights  # Class weights to handle class imbalance\n\n    def forward(self, outputs, labels):\n        # Apply cross-entropy loss\n        ce_loss = nn.CrossEntropyLoss(weight=self.weights)(outputs, labels)\n\n        # Get the probability of the correct class\n        pt = torch.exp(-ce_loss)  # Probabilities of the true class\n        \n        # Apply Focal Loss formula\n        focal_loss = ((1 - pt) ** self.gamma) * ce_loss\n        \n        return focal_loss\n\n# Modify your weighted loss function to use Focal Loss\ndef weighted_loss_function(outputs, labels_t1, labels_t2, labels_t2stir, weights):\n    # Focal Loss with gamma and weights (adjust gamma as necessary)\n    focal_loss_fn = FocalLoss(weights=weights, gamma=4)\n\n    # Average loss calculations for each series (Sagittal T1, Axial T2, Sagittal T2/STIR)\n    loss_t1 = sum([focal_loss_fn(outputs[:, 0, level, :], labels_t1[:, level]) for level in range(5)]) / 5\n    loss_t2 = sum([focal_loss_fn(outputs[:, 1, level, :], labels_t2[:, level]) for level in range(5)]) / 5\n    loss_t2stir = sum([focal_loss_fn(outputs[:, 2, level, :], labels_t2stir[:, level]) for level in range(5)]) / 5\n\n    # Total loss is the sum of averaged losses from all series\n    total_loss = loss_t1 + loss_t2 + loss_t2stir\n    \n    return total_loss\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:21.462099Z","iopub.execute_input":"2024-10-08T23:54:21.462372Z","iopub.status.idle":"2024-10-08T23:54:21.474064Z","shell.execute_reply.started":"2024-10-08T23:54:21.462342Z","shell.execute_reply":"2024-10-08T23:54:21.473209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.optim as optim\n\n\n# Assume 'model' is already defined somewhere in your code\n# Optimizer and Learning Rate Scheduler\n#optimizer = optim.Adam(model.parameters(), lr=1e-4)\n#scheduler = lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=5)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:21.475195Z","iopub.execute_input":"2024-10-08T23:54:21.475863Z","iopub.status.idle":"2024-10-08T23:54:21.487033Z","shell.execute_reply.started":"2024-10-08T23:54:21.475813Z","shell.execute_reply":"2024-10-08T23:54:21.486094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom itertools import cycle\n\n# Assuming separate dataloaders for each series\ntrain_loaders = {\n    'Sagittal T1': dataloaders['Sagittal T1'][0],\n    'Axial T2': dataloaders['Axial T2'][0],\n    'Sagittal T2/STIR': dataloaders['Sagittal T2/STIR'][0]\n}\nval_loaders = {\n    'Sagittal T1': dataloaders['Sagittal T1'][1],\n    'Axial T2': dataloaders['Axial T2'][1],\n    'Sagittal T2/STIR': dataloaders['Sagittal T2/STIR'][1]\n}\n\n# Print the number of batches in each dataloader\nprint(\"\\n--- DataLoader Information ---\")\nfor series, loader in train_loaders.items():\n    print(f\"{series}: {len(loader)} batches for training.\")\nfor series, loader in val_loaders.items():\n    print(f\"{series}: {len(loader)} batches for validation.\")\n\n# Get the first batch from each dataloader to inspect input/output sizes and label information\nbatch_t1 = next(iter(train_loaders['Sagittal T1']))\nbatch_t2 = next(iter(train_loaders['Axial T2']))\nbatch_t2stir = next(iter(train_loaders['Sagittal T2/STIR']))\n\n# Unpack images and labels for each series\nx_t1, labels_t1 = batch_t1\nx_t2, labels_t2 = batch_t2\nx_t2stir, labels_t2stir = batch_t2stir\n\n# Print the shape of the images and labels for each series\nprint(\"\\n--- Batch Information ---\")\nprint(f\"Sagittal T1 batch size: {x_t1.size(0)}, Image shape: {x_t1.shape}, Labels shape: {labels_t1.shape}\")\nprint(f\"Axial T2 batch size: {x_t2.size(0)}, Image shape: {x_t2.shape}, Labels shape: {labels_t2.shape}\")\nprint(f\"Sagittal T2/STIR batch size: {x_t2stir.size(0)}, Image shape: {x_t2stir.shape}, Labels shape: {labels_t2stir.shape}\")\n\n# Model information: forward pass with dummy data to print input/output shapes\nx_t1, x_t2, x_t2stir = x_t1.to(device), x_t2.to(device), x_t2stir.to(device)\n\n# Forward pass to inspect input/output shapes\noutputs = model(x_t1, x_t2, x_t2stir)\n\n# Print model input and output shapes\nprint(\"\\n--- Model Input/Output Information ---\")\nprint(f\"Model input shapes: T1 {x_t1.shape}, T2 {x_t2.shape}, T2STIR {x_t2stir.shape}\")\nprint(f\"Model output shape: {outputs.shape}\")\n\n# Check the number of batches that will be used during training\nnum_batches = max(len(train_loaders['Sagittal T1']), len(train_loaders['Axial T2']), len(train_loaders['Sagittal T2/STIR']))\nprint(f\"\\nNumber of batches for training: {num_batches}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:21.488093Z","iopub.execute_input":"2024-10-08T23:54:21.488410Z","iopub.status.idle":"2024-10-08T23:54:23.244505Z","shell.execute_reply.started":"2024-10-08T23:54:21.488356Z","shell.execute_reply":"2024-10-08T23:54:23.243461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Using GPU: \", torch.cuda.is_available())\nprint(\"GPU Name: \", torch.cuda.get_device_name(0))\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:23.246069Z","iopub.execute_input":"2024-10-08T23:54:23.246494Z","iopub.status.idle":"2024-10-08T23:54:23.251760Z","shell.execute_reply.started":"2024-10-08T23:54:23.246448Z","shell.execute_reply":"2024-10-08T23:54:23.250770Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport gc  # For garbage collection\nfrom itertools import cycle\nfrom sklearn.metrics import accuracy_score, f1_score, precision_score, recall_score\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\ntrain_loaders = {\n    'Sagittal T1': dataloaders['Sagittal T1'][0],\n    'Axial T2': dataloaders['Axial T2'][0],\n    'Sagittal T2/STIR': dataloaders['Sagittal T2/STIR'][0]\n}\nval_loaders = {\n    'Sagittal T1': dataloaders['Sagittal T1'][1],\n    'Axial T2': dataloaders['Axial T2'][1],\n    'Sagittal T2/STIR': dataloaders['Sagittal T2/STIR'][1]\n}\n\n# Use cycle() for unequal batch sizes\ntrain_t1 = cycle(train_loaders['Sagittal T1'])\ntrain_t2 = cycle(train_loaders['Axial T2'])\ntrain_t2stir = cycle(train_loaders['Sagittal T2/STIR'])\n\n# Number of batches for the longest dataloader\nnum_batches = max(len(train_loaders['Sagittal T1']), len(train_loaders['Axial T2']), len(train_loaders['Sagittal T2/STIR']))\n\nnum_epochs = 15\n#optimizer = torch.optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-4)\n\n# Learning rate scheduler\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=1)\n\n# Early stopping parameters\npatience = 2  # How many epochs to wait for improvement before stopping\nbest_val_loss = float('inf')\nepochs_without_improvement = 0\n\n# Path to save the best model\nbest_model_path = 'best_model.pth'\n\n\ndef pad_batch(batch, target_size, device):\n    \"\"\" Pad a batch (images and labels) with zeros to match the target batch size. \"\"\"\n    images, labels = batch\n    current_size = images.size(0)\n    \n    # Only pad if the current batch size is smaller than the target size\n    if current_size < target_size:\n        pad_size = target_size - current_size\n        \n        # Print a message when padding occurs\n        print(f\"Padding batch: current size {current_size}, target size {target_size}, padding {pad_size} items.\")\n        \n        # Pad images and labels\n        padded_images = torch.cat([images, torch.zeros((pad_size, *images.shape[1:]), device=device)], dim=0)\n        padded_labels = torch.cat([labels, torch.zeros((pad_size, *labels.shape[1:]), device=device, dtype=labels.dtype)], dim=0)\n    else:\n        padded_images, padded_labels = images, labels\n    \n    return padded_images, padded_labels\n\n\n\ndef compute_f1_score(true_labels, predicted_labels):\n    \"\"\" Compute the F1 score for multi-class, multi-output classification \"\"\"\n    true_labels_flat = true_labels.view(-1).cpu().numpy()\n    predicted_labels_flat = predicted_labels.view(-1).cpu().numpy()\n\n    # Compute the weighted F1 score\n    return f1_score(true_labels_flat, predicted_labels_flat, average='weighted')\n\n\nfor epoch in range(num_epochs):\n    # Training Phase\n    model.train()\n    running_loss = 0.0\n    total_train = 0\n    correct_train = 0\n    train_f1_scores = []\n\n    for _ in range(num_batches):\n        # Fetching batches from the dataloaders using cycle\n        batch_t1 = next(train_t1)\n        batch_t2 = next(train_t2)\n        batch_t2stir = next(train_t2stir)\n\n        # Move data to GPU\n        batch_t1 = (batch_t1[0].to(device), batch_t1[1].to(device))\n        batch_t2 = (batch_t2[0].to(device), batch_t2[1].to(device))\n        batch_t2stir = (batch_t2stir[0].to(device), batch_t2stir[1].to(device))\n\n        # Find the maximum batch size among the three batches\n        max_batch_size = max(batch_t1[0].size(0), batch_t2[0].size(0), batch_t2stir[0].size(0))\n\n        # Pad each batch to match the maximum batch size\n        x_t1, labels_t1 = pad_batch(batch_t1, max_batch_size, device)\n        x_t2, labels_t2 = pad_batch(batch_t2, max_batch_size, device)\n        x_t2stir, labels_t2stir = pad_batch(batch_t2stir, max_batch_size, device)\n\n        # Zero gradients\n        optimizer.zero_grad()\n\n        # Forward pass through the model\n        outputs = model(x_t1, x_t2, x_t2stir)\n\n        # Compute loss using the custom weighted loss function\n        loss = weighted_loss_function(outputs, labels_t1, labels_t2, labels_t2stir, weights)\n\n        # Backward pass and optimization\n        loss.backward()\n         # Gradient clipping\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0)\n        optimizer.step()\n\n        # Track training loss\n        running_loss += loss.item()\n\n        # Track training accuracy\n        _, predicted = torch.max(outputs[:, 2, :, :], 2)  # Predictions for Sagittal T2/STIR across all 5 levels\n        correct_train += (predicted == labels_t2stir).sum().item()\n        total_train += labels_t2stir.numel()\n\n        # Calculate F1 score for the batch\n        batch_f1 = compute_f1_score(labels_t2stir, predicted)\n        train_f1_scores.append(batch_f1)\n\n    train_accuracy = correct_train / total_train\n    train_f1 = sum(train_f1_scores) / len(train_f1_scores)\n\n    # Validation Phase\n    model.eval()\n    val_loss = 0.0\n    val_correct = 0\n    val_total = 0\n    val_f1_scores = []\n    # Keep track of total validation batches\n    total_val_batches = 0\n\n    with torch.no_grad():\n        for val_batch_t1, val_batch_t2, val_batch_t2stir in zip(val_loaders['Sagittal T1'], val_loaders['Axial T2'], val_loaders['Sagittal T2/STIR']):\n            \n            total_val_batches += 1\n            # Move validation batches to GPU\n            val_batch_t1 = (val_batch_t1[0].to(device), val_batch_t1[1].to(device))\n            val_batch_t2 = (val_batch_t2[0].to(device), val_batch_t2[1].to(device))\n            val_batch_t2stir = (val_batch_t2stir[0].to(device), val_batch_t2stir[1].to(device))\n\n            # Find the maximum batch size for validation batches\n            max_val_batch_size = max(val_batch_t1[0].size(0), val_batch_t2[0].size(0), val_batch_t2stir[0].size(0))\n\n            # Pad validation batches\n            x_val_t1, val_labels_t1 = pad_batch(val_batch_t1, max_val_batch_size, device)\n            x_val_t2, val_labels_t2 = pad_batch(val_batch_t2, max_val_batch_size, device)\n            x_val_t2stir, val_labels_t2stir = pad_batch(val_batch_t2stir, max_val_batch_size, device)\n\n            # Forward pass\n            outputs_val = model(x_val_t1, x_val_t2, x_val_t2stir)\n\n            # Compute validation loss\n            val_loss += weighted_loss_function(outputs_val, val_labels_t1, val_labels_t2, val_labels_t2stir, weights).item()\n\n            # Compute validation accuracy\n            _, val_predicted = torch.max(outputs_val[:, 2, :, :], 2)\n            val_correct += (val_predicted == val_labels_t2stir).sum().item()\n            val_total += val_labels_t2stir.numel()\n\n            # Calculate F1 score for the validation batch\n            val_batch_f1 = compute_f1_score(val_labels_t2stir, val_predicted)\n            val_f1_scores.append(val_batch_f1)\n\n    # Compute average validation loss for the epoch (divide by total validation batches)\n    avg_val_loss = val_loss / total_val_batches\n\n    # Learning rate scheduling based on the average validation loss\n    scheduler.step(avg_val_loss)\n\n    val_accuracy = val_correct / val_total\n    val_f1 = sum(val_f1_scores) / len(val_f1_scores)\n    \n\n    # Check for validation loss improvement\n    if avg_val_loss < best_val_loss:\n        best_val_loss = avg_val_loss\n        epochs_without_improvement = 0\n        \n        # Save the best model\n        torch.save(model.state_dict(), best_model_path)\n        print(f'New best model saved with validation loss: {best_val_loss:.4f}')\n    else:\n        epochs_without_improvement += 1\n\n    # Early stopping check\n    if epochs_without_improvement >= patience:\n        print(f'Early stopping at epoch {epoch + 1}')\n        break\n\n    # Print statistics for the epoch\n    print(f'Epoch [{epoch + 1}/{num_epochs}], '\n          f'Training Loss: {running_loss/num_batches:.4f}, '\n          f'Training Accuracy: {train_accuracy:.4f}, '\n          f'Training F1 Score: {train_f1:.4f}, '\n          f'Validation Loss: {avg_val_loss:.4f}, '\n          f'Validation Accuracy: {val_accuracy:.4f}, '\n          f'Validation F1 Score: {val_f1:.4f}')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T23:54:23.253041Z","iopub.execute_input":"2024-10-08T23:54:23.253371Z","iopub.status.idle":"2024-10-09T00:11:30.811131Z","shell.execute_reply.started":"2024-10-08T23:54:23.253338Z","shell.execute_reply":"2024-10-09T00:11:30.808900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pandas as pd\n\n# Paths for test series descriptions and test images\ntest_desc_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv'\ntest_images_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images'\n\n# Load test series descriptions CSV\ntest_desc = pd.read_csv(test_desc_path)","metadata":{"execution":{"iopub.status.busy":"2024-10-09T00:11:40.377270Z","iopub.execute_input":"2024-10-09T00:11:40.377692Z","iopub.status.idle":"2024-10-09T00:11:40.391170Z","shell.execute_reply.started":"2024-10-09T00:11:40.377650Z","shell.execute_reply":"2024-10-09T00:11:40.390324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pandas as pd\n\n# Define the test images path\ntest_images_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images\"\n\n# Define conditions with their corresponding series and sides\ncondition_df = pd.DataFrame([\n    {'series_description': 'Sagittal T1', 'condition': 'left_neural_foraminal_narrowing', 'side': 'left'},\n    {'series_description': 'Sagittal T1', 'condition': 'right_neural_foraminal_narrowing', 'side': 'right'},\n    {'series_description': 'Axial T2', 'condition': 'left_subarticular_stenosis', 'side': 'left'},\n    {'series_description': 'Axial T2', 'condition': 'right_subarticular_stenosis', 'side': 'right'},\n    {'series_description': 'Sagittal T2/STIR', 'condition': 'spinal_canal_stenosis', 'side': 'both'}\n])\n\n# Merge test descriptions with condition mappings\nmerged_df = pd.merge(test_desc, condition_df, on='series_description', how='left')\n\n# Function to get image paths for each series\ndef get_image_paths(row):\n    series_folder_path = os.path.join(test_images_path, str(row['study_id']), str(row['series_id']))\n    if os.path.exists(series_folder_path):\n        image_paths = [os.path.join(series_folder_path, f) for f in os.listdir(series_folder_path) if f.endswith('.dcm')]\n\n        # Check if the series is associated with left/right conditions\n        if row['side'] in ['left', 'right']:\n            # Split images for left/right conditions, if applicable\n            split_idx = len(image_paths) // 2\n            if row['side'] == 'left':\n                return image_paths[:split_idx]  # Assign first half to 'left'\n            elif row['side'] == 'right':\n                return image_paths[split_idx:]  # Assign second half to 'right'\n        else:\n            # For 'both' (e.g., spinal canal stenosis), use all images\n            return image_paths\n    else:\n        print(f\"Series path does not exist: {series_folder_path}\")\n    return []\n\n# Apply get_image_paths to add image paths to each row\nmerged_df['image_paths'] = merged_df.apply(get_image_paths, axis=1)\n\n# Explode image paths so that each row corresponds to a single image\nexpanded_test_desc = merged_df.explode('image_paths').reset_index(drop=True)\n\n# Rename image_paths column to image_path for consistency\nexpanded_test_desc = expanded_test_desc.rename(columns={'image_paths': 'image_path'})\n\n# Drop unnecessary columns (like 'side') and reset index\nexpanded_test_desc = expanded_test_desc.drop(columns=['side']).reset_index(drop=True)\n\n# Display the number of rows and inspect the final DataFrame\nprint(f\"Number of rows in expanded_test_desc: {expanded_test_desc.shape[0]}\")\n\n# Inspect the DataFrame to check the distribution of images for left, right, and both conditions\nexpanded_test_desc.head()\n","metadata":{"execution":{"iopub.status.busy":"2024-10-09T00:11:46.619281Z","iopub.execute_input":"2024-10-09T00:11:46.619740Z","iopub.status.idle":"2024-10-09T00:11:46.687051Z","shell.execute_reply.started":"2024-10-09T00:11:46.619694Z","shell.execute_reply":"2024-10-09T00:11:46.686123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\n\n# Create the 'row_id' column by concatenating 'study_id' and 'condition' columns\nexpanded_test_desc['row_id'] = expanded_test_desc['study_id'].astype(str) + \"_\" + expanded_test_desc['condition']\n\n# Display the first 10 rows with the new 'row_id' column\nexpanded_test_desc.head()","metadata":{"execution":{"iopub.status.busy":"2024-10-09T00:11:56.130905Z","iopub.execute_input":"2024-10-09T00:11:56.131723Z","iopub.status.idle":"2024-10-09T00:11:56.145369Z","shell.execute_reply.started":"2024-10-09T00:11:56.131678Z","shell.execute_reply":"2024-10-09T00:11:56.144371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport numpy as np\nimport cv2\nimport pydicom\n\n# Define the transformations (same as training for consistency)\ntransform = transforms.Compose([\n    transforms.ToPILImage(),  # Convert to PIL image\n    transforms.Resize((224, 224)),  # Resize the image to the expected input size\n    transforms.ToTensor(),  # Convert image to tensor\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])\n\n# Custom Dataset for the test/series dataset\nclass SeriesDataset(Dataset):\n    def __init__(self, dataframe, series_column, transform=None):\n        \"\"\"\n        Initialize the SeriesDataset with a dataframe and the column containing image paths.\n        \"\"\"\n        self.dataframe = dataframe\n        self.series_column = series_column\n        self.transform = transform if transform is not None else self.default_transform()\n\n    def __len__(self):\n        \"\"\"\n        Return the total number of samples in the dataset.\n        \"\"\"\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        \"\"\"\n        Retrieve the image and row_id for a given index.\n        \"\"\"\n        row = self.dataframe.iloc[idx]\n        image_path = row[self.series_column]\n        \n        # Check if image path is valid\n        if pd.notnull(image_path):\n            # Load the image\n            image = self.load_image(image_path)\n            \n            # Apply transformations if image is loaded successfully\n            if image is not None:\n                image_tensor = self.transform(image)\n            else:\n                print(f\"Image not found or failed to load at path: {image_path} for row_id: {row['row_id']}\")\n                return None\n        else:\n            print(f\"Missing image path for row_id: {row['row_id']}\")\n            return None\n\n        # Return a dictionary with the image tensor and the row_id\n        return {'image_tensor': image_tensor, 'row_id': row['row_id']}\n\n    def load_image(self, image_path):\n        \"\"\"\n        Load and preprocess a single DICOM image.\n        \"\"\"\n        try:\n            dicom = pydicom.dcmread(image_path)\n            img = dicom.pixel_array\n            # Convert uint16 to uint8 (just to ensure compatibility, no scaling)\n            img = np.clip(img, 0, 255).astype(np.uint8)\n\n            img_rgb = np.stack([img] * 3, axis=-1)  # Convert to 3 channels (RGB)\n            return img_rgb\n        except Exception as e:\n            print(f\"Error loading image at {image_path}: {e}\")\n            return None\n\n    @staticmethod\n    def default_transform():\n        \"\"\"\n        Define the default transformations to be applied to the images if no custom transform is provided.\n        \"\"\"\n        return transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.Resize((224, 224)),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),  # Normalize using ImageNet stats\n        ])\n\n","metadata":{"execution":{"iopub.status.busy":"2024-10-09T00:12:02.359364Z","iopub.execute_input":"2024-10-09T00:12:02.360097Z","iopub.status.idle":"2024-10-09T00:12:02.375248Z","shell.execute_reply.started":"2024-10-09T00:12:02.360054Z","shell.execute_reply":"2024-10-09T00:12:02.374233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create separate DataFrames for each series\nt1_df = expanded_test_desc[expanded_test_desc['series_description'] == 'Sagittal T1'].reset_index(drop=True)\nt2_df = expanded_test_desc[expanded_test_desc['series_description'] == 'Axial T2'].reset_index(drop=True)\nt2stir_df = expanded_test_desc[expanded_test_desc['series_description'] == 'Sagittal T2/STIR'].reset_index(drop=True)\n\n# Initialize datasets\nt1_dataset = SeriesDataset(t1_df, 'image_path')\nt2_dataset = SeriesDataset(t2_df, 'image_path')\nt2stir_dataset = SeriesDataset(t2stir_df, 'image_path')\n\n# Create DataLoaders\nt1_dataloader = DataLoader(t1_dataset, batch_size=1, shuffle=False)\nt2_dataloader = DataLoader(t2_dataset, batch_size=1, shuffle=False)\nt2stir_dataloader = DataLoader(t2stir_dataset, batch_size=1, shuffle=False)\n\nprint(f\"Sagittal T1 DataLoader size: {len(t1_dataloader)} batches\")\nprint(f\"Axial T2 DataLoader size: {len(t2_dataloader)} batches\")\nprint(f\"Sagittal T2/STIR DataLoader size: {len(t2stir_dataloader)} batches\")","metadata":{"execution":{"iopub.status.busy":"2024-10-09T00:12:11.446906Z","iopub.execute_input":"2024-10-09T00:12:11.447315Z","iopub.status.idle":"2024-10-09T00:12:11.461213Z","shell.execute_reply.started":"2024-10-09T00:12:11.447276Z","shell.execute_reply":"2024-10-09T00:12:11.460145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load the best model saved from training\nmodel.load_state_dict(torch.load('best_model.pth'))\nmodel.eval()  # Set the model to evaluation mode","metadata":{"execution":{"iopub.status.busy":"2024-10-09T00:12:19.434724Z","iopub.execute_input":"2024-10-09T00:12:19.435111Z","iopub.status.idle":"2024-10-09T00:12:19.609008Z","shell.execute_reply.started":"2024-10-09T00:12:19.435073Z","shell.execute_reply":"2024-10-09T00:12:19.607580Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport pandas as pd\n\n# Temperature Scaler Class\nclass TemperatureScaler:\n    def __init__(self, temperature=1.0):\n        self.temperature = temperature\n\n    def apply_temperature_scaling(self, logits):\n        return logits / self.temperature\n\n\n# Load DataLoaders for each series \nt1_dataloader = DataLoader(t1_dataset, batch_size=1, shuffle=False)\nt2_dataloader = DataLoader(t2_dataset, batch_size=1, shuffle=False)\nt2stir_dataloader = DataLoader(t2stir_dataset, batch_size=1, shuffle=False)\n# Initialize temperature scalers with known values\nt1_scaler = TemperatureScaler(temperature=1.0)\nt2_scaler = TemperatureScaler(temperature=1.0)\nt2stir_scaler = TemperatureScaler(temperature=1.0)\n\n# List to store individual prediction rows\nindividual_rows = []\n\n# Helper function to get probabilities from the correct series-specific image and dataloader\ndef get_probabilities_from_dataloader(dataloader, series_label, levels, temp_scaler):\n    with torch.no_grad():\n        for batch in dataloader:\n            image = batch['image_tensor']\n            row_id_base = batch['row_id'][0]\n\n            image = image.squeeze(1).to(device)  # Squeeze singleton dimension and move to device\n\n            if image is not None:\n                if series_label == 'Sagittal T1':\n                    output = model(image, None, None).squeeze(0)  # Use only T1 image for prediction\n                    output = output[0]\n                elif series_label == 'Axial T2':\n                    output = model(None, image, None).squeeze(0)  # Use only T2 image for prediction\n                    output = output[1]\n                elif series_label == 'Sagittal T2/STIR':\n                    output = model(None, None, image).squeeze(0)  # Use only T2/STIR image for prediction\n                    output = output[2]\n\n                # Apply temperature scaling and get probabilities\n                output = temp_scaler.apply_temperature_scaling(output)\n                for level_idx, level_name in enumerate(levels):\n                    level_probabilities = torch.softmax(output[level_idx], dim=-1).cpu().numpy()\n                    level_row_id = f\"{row_id_base}_{level_name}\"\n                    individual_rows.append([level_row_id, *level_probabilities])\n            else:\n                # Handle missing images by defaulting to \"Normal/Mild\"\n                print(f\"Missing image for row_id: {row_id_base}\")\n                for level_name in levels:\n                    level_row_id = f\"{row_id_base}_{level_name}\"\n                    individual_rows.append([level_row_id, 1.0, 0.0, 0.0])\n\n\n# List of levels\nlevels = ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\n\n# Generate predictions for each series using respective dataloaders\nget_probabilities_from_dataloader(t1_dataloader, 'Sagittal T1', levels, t1_scaler)\nget_probabilities_from_dataloader(t2_dataloader, 'Axial T2', levels, t2_scaler)\nget_probabilities_from_dataloader(t2stir_dataloader, 'Sagittal T2/STIR', levels, t2stir_scaler)\n\n# Create DataFrame for individual predictions\nresults_df_with_condition = pd.DataFrame(individual_rows, columns=['row_id', 'normal_mild', 'moderate', 'severe'])\n\n# Aggregate the predictions by averaging over the same row_ids\nsubmission = results_df_with_condition.groupby('row_id').mean().reset_index()\n\n# Save submission to CSV in the required format\nsubmission_file = \"/kaggle/working/submission.csv\"\nsubmission.to_csv(submission_file, index=False)\n\nprint(f\"Submission saved to: {submission_file}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-10-09T00:12:28.164384Z","iopub.execute_input":"2024-10-09T00:12:28.165273Z","iopub.status.idle":"2024-10-09T00:12:28.835783Z","shell.execute_reply.started":"2024-10-09T00:12:28.165229Z","shell.execute_reply":"2024-10-09T00:12:28.834040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2024-10-09T00:11:30.825649Z","iopub.status.idle":"2024-10-09T00:11:30.825999Z","shell.execute_reply.started":"2024-10-09T00:11:30.825823Z","shell.execute_reply":"2024-10-09T00:11:30.825841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2024-10-09T00:11:30.827944Z","iopub.status.idle":"2024-10-09T00:11:30.828280Z","shell.execute_reply.started":"2024-10-09T00:11:30.828110Z","shell.execute_reply":"2024-10-09T00:11:30.828126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.describe()","metadata":{"execution":{"iopub.status.busy":"2024-10-09T00:11:30.829860Z","iopub.status.idle":"2024-10-09T00:11:30.830217Z","shell.execute_reply.started":"2024-10-09T00:11:30.830038Z","shell.execute_reply":"2024-10-09T00:11:30.830056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}