{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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"}],"dockerImageVersionId":30840,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import seaborn as sns\n\nimport matplotlib.pyplot as plt\nimport os\nimport time\nimport numpy as np\nimport glob\nimport json\nimport collections\nimport torch\nimport torch.nn as nn\n\nimport pydicom as dicom\nimport matplotlib.patches as patches\n\nfrom matplotlib import animation, rc\nimport pandas as pd\n\nimport pydicom as dicom # dicom\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:28:18.668435Z","iopub.execute_input":"2025-06-02T15:28:18.668771Z","iopub.status.idle":"2025-06-02T15:28:23.462345Z","shell.execute_reply.started":"2025-06-02T15:28:18.668744Z","shell.execute_reply":"2025-06-02T15:28:23.461238Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# read data\ntrain_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\n\ntrain  = pd.read_csv(train_path + 'train.csv')\nlabel = pd.read_csv(train_path + 'train_label_coordinates.csv')\ntrain_desc  = pd.read_csv(train_path + 'train_series_descriptions.csv')\ntest_desc   = pd.read_csv(train_path + 'test_series_descriptions.csv')\nsub         = pd.read_csv(train_path + 'sample_submission.csv')\nlen(test_desc) #number of test_description.csv rows ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:28:23.463641Z","iopub.execute_input":"2025-06-02T15:28:23.464171Z","iopub.status.idle":"2025-06-02T15:28:23.645630Z","shell.execute_reply.started":"2025-06-02T15:28:23.464137Z","shell.execute_reply":"2025-06-02T15:28:23.644731Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_desc.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:28:23.647382Z","iopub.execute_input":"2025-06-02T15:28:23.647698Z","iopub.status.idle":"2025-06-02T15:28:23.667186Z","shell.execute_reply.started":"2025-06-02T15:28:23.647668Z","shell.execute_reply":"2025-06-02T15:28:23.666333Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_desc.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:28:23.668556Z","iopub.execute_input":"2025-06-02T15:28:23.668796Z","iopub.status.idle":"2025-06-02T15:28:23.678520Z","shell.execute_reply.started":"2025-06-02T15:28:23.668763Z","shell.execute_reply":"2025-06-02T15:28:23.677641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:28:23.679404Z","iopub.execute_input":"2025-06-02T15:28:23.679756Z","iopub.status.idle":"2025-06-02T15:28:23.709451Z","shell.execute_reply.started":"2025-06-02T15:28:23.679726Z","shell.execute_reply":"2025-06-02T15:28:23.708513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to generate image paths based on directory structure\ndef generate_image_paths(df, data_dir):\n    image_paths = []\n    for study_id, series_id in zip(df['study_id'], df['series_id']):\n        study_dir = os.path.join(data_dir, str(study_id))\n        series_dir = os.path.join(study_dir, str(series_id))\n        images = os.listdir(series_dir)\n        image_paths.extend([os.path.join(series_dir, img) for img in images])\n    return image_paths\n\n# Generate image paths for train and test data\ntrain_image_paths = generate_image_paths(train_desc, f'{train_path}/train_images')\ntest_image_paths = generate_image_paths(test_desc, f'{train_path}/test_images')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:28:23.710370Z","iopub.execute_input":"2025-06-02T15:28:23.710740Z","iopub.status.idle":"2025-06-02T15:29:41.464470Z","shell.execute_reply.started":"2025-06-02T15:28:23.710707Z","shell.execute_reply":"2025-06-02T15:29:41.463733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(train_desc)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:41.465300Z","iopub.execute_input":"2025-06-02T15:29:41.465597Z","iopub.status.idle":"2025-06-02T15:29:41.470401Z","shell.execute_reply.started":"2025-06-02T15:29:41.465568Z","shell.execute_reply":"2025-06-02T15:29:41.469720Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(train_image_paths)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:41.472535Z","iopub.execute_input":"2025-06-02T15:29:41.472743Z","iopub.status.idle":"2025-06-02T15:29:41.484883Z","shell.execute_reply.started":"2025-06-02T15:29:41.472725Z","shell.execute_reply":"2025-06-02T15:29:41.484265Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define function to reshape a single row of the DataFrame\ndef reshape_row(row):\n    data = {'study_id': [], 'condition': [], 'level': [], 'severity': []}\n    \n    for column, value in row.items():\n        if column not in ['study_id', 'series_id', 'instance_number', 'x', 'y', 'series_description']:\n            parts = column.split('_')\n            condition = ' '.join([word.capitalize() for word in parts[:-2]])\n            level = parts[-2].capitalize() + '/' + parts[-1].capitalize()\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# Reshape the DataFrame for all rows\nnew_train_df = pd.concat([reshape_row(row) for _, row in train.iterrows()], ignore_index=True)\n\n# Display the first few rows of the reshaped dataframe\nnew_train_df.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:41.486189Z","iopub.execute_input":"2025-06-02T15:29:41.486463Z","iopub.status.idle":"2025-06-02T15:29:42.492557Z","shell.execute_reply.started":"2025-06-02T15:29:41.486442Z","shell.execute_reply":"2025-06-02T15:29:42.491689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Print columns in a neat way\nprint(\"\\nColumns in new_train_df:\")\nprint(\",\".join(new_train_df.columns))\n\nprint(\"\\nColumns in label:\")\nprint(\",\".join(label.columns))\n\nprint(\"\\nColumns in test_desc:\")\nprint(\",\".join(test_desc.columns))\n\nprint(\"\\nColumns in sub:\")\nprint(\",\".join(sub.columns))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:42.493472Z","iopub.execute_input":"2025-06-02T15:29:42.493803Z","iopub.status.idle":"2025-06-02T15:29:42.500976Z","shell.execute_reply.started":"2025-06-02T15:29:42.493771Z","shell.execute_reply":"2025-06-02T15:29:42.500114Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Merge the dataframes on the common columns\nmerged_df = pd.merge(new_train_df, label, on=['study_id', 'condition', 'level'], how='inner')\n# Merge the dataframes on the common column 'series_id'\nfinal_merged_df = pd.merge(merged_df, train_desc, on='series_id', how='inner')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:42.501848Z","iopub.execute_input":"2025-06-02T15:29:42.502158Z","iopub.status.idle":"2025-06-02T15:29:42.570436Z","shell.execute_reply.started":"2025-06-02T15:29:42.502130Z","shell.execute_reply":"2025-06-02T15:29:42.569590Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Merge the dataframes on the common column 'series_id'\nfinal_merged_df = pd.merge(merged_df, train_desc, on=['series_id','study_id'], how='inner')\n# Display the first few rows of the final merged dataframe\nfinal_merged_df.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:42.571336Z","iopub.execute_input":"2025-06-02T15:29:42.571677Z","iopub.status.idle":"2025-06-02T15:29:42.598384Z","shell.execute_reply.started":"2025-06-02T15:29:42.571638Z","shell.execute_reply":"2025-06-02T15:29:42.597420Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# Create the row_id column\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# Create the image_path column\nfinal_merged_df['image_path'] = (\n    f'{train_path}/train_images/' + \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\n# Note: Check image path, since there's 1 instance id, for 1 image, but there's many more images other than the ones labelled in the instance ID. \n\n# Display the updated dataframe\nfinal_merged_df.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:42.599315Z","iopub.execute_input":"2025-06-02T15:29:42.599583Z","iopub.status.idle":"2025-06-02T15:29:42.752176Z","shell.execute_reply.started":"2025-06-02T15:29:42.599551Z","shell.execute_reply":"2025-06-02T15:29:42.751224Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Normal/Mild\"].value_counts().sum()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:42.753166Z","iopub.execute_input":"2025-06-02T15:29:42.753505Z","iopub.status.idle":"2025-06-02T15:29:42.865584Z","shell.execute_reply.started":"2025-06-02T15:29:42.753473Z","shell.execute_reply":"2025-06-02T15:29:42.864696Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Moderate\"].value_counts().sum()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:42.866525Z","iopub.execute_input":"2025-06-02T15:29:42.866879Z","iopub.status.idle":"2025-06-02T15:29:42.902623Z","shell.execute_reply.started":"2025-06-02T15:29:42.866819Z","shell.execute_reply":"2025-06-02T15:29:42.901734Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Severe\"].value_counts().sum()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:42.903510Z","iopub.execute_input":"2025-06-02T15:29:42.903755Z","iopub.status.idle":"2025-06-02T15:29:42.926305Z","shell.execute_reply.started":"2025-06-02T15:29:42.903732Z","shell.execute_reply":"2025-06-02T15:29:42.925593Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the base path for test images\nbase_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/'\n\n# Function to get image paths for a series\ndef get_image_paths(row):\n    series_path = os.path.join(base_path, str(row['study_id']), str(row['series_id']))\n    if os.path.exists(series_path):\n        return [os.path.join(series_path, f) for f in os.listdir(series_path) if os.path.isfile(os.path.join(series_path, f))]\n    return []\n\n# Mapping of series_description to conditions\ncondition_mapping = {\n    'Sagittal T1': {'left': 'left_neural_foraminal_narrowing', 'right': 'right_neural_foraminal_narrowing'},\n    'Axial T2': {'left': 'left_subarticular_stenosis', 'right': 'right_subarticular_stenosis'},\n    'Sagittal T2/STIR': 'spinal_canal_stenosis'\n}\n\n# Create a list to store the expanded rows\nexpanded_rows = []\n\n# Expand the dataframe by adding new rows for each file path\nfor index, row in test_desc.iterrows():\n    image_paths = get_image_paths(row)\n    conditions = condition_mapping.get(row['series_description'], {})\n    if isinstance(conditions, str):  # Single condition\n        conditions = {'left': conditions, 'right': conditions}\n    for side, condition in conditions.items():\n        for image_path in image_paths:\n            expanded_rows.append({\n                'study_id': row['study_id'],\n                'series_id': row['series_id'],\n                'series_description': row['series_description'],\n                'image_path': image_path,\n                'condition': condition,\n                'row_id': f\"{row['study_id']}_{condition}\"\n            })\n\n# Create a new dataframe from the expanded rows\nexpanded_test_desc = pd.DataFrame(expanded_rows)\n\n# Display the resulting dataframe\nexpanded_test_desc.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:42.927183Z","iopub.execute_input":"2025-06-02T15:29:42.927502Z","iopub.status.idle":"2025-06-02T15:29:43.090642Z","shell.execute_reply.started":"2025-06-02T15:29:42.927472Z","shell.execute_reply":"2025-06-02T15:29:43.089934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# change severity column labels\n#Normal/Mild': 'normal_mild', 'Moderate': 'moderate', 'Severe': 'severe'}\nfinal_merged_df['severity'] = final_merged_df['severity'].map({'Normal/Mild': 'normal_mild', 'Moderate': 'moderate', 'Severe': 'severe'})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:43.091384Z","iopub.execute_input":"2025-06-02T15:29:43.091616Z","iopub.status.idle":"2025-06-02T15:29:43.099418Z","shell.execute_reply.started":"2025-06-02T15:29:43.091597Z","shell.execute_reply":"2025-06-02T15:29:43.098598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_data = expanded_test_desc\ntrain_data = final_merged_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:43.100162Z","iopub.execute_input":"2025-06-02T15:29:43.100466Z","iopub.status.idle":"2025-06-02T15:29:43.112974Z","shell.execute_reply.started":"2025-06-02T15:29:43.100436Z","shell.execute_reply":"2025-06-02T15:29:43.112219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:43.113614Z","iopub.execute_input":"2025-06-02T15:29:43.113822Z","iopub.status.idle":"2025-06-02T15:29:43.137859Z","shell.execute_reply.started":"2025-06-02T15:29:43.113804Z","shell.execute_reply":"2025-06-02T15:29:43.137050Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data['series_description'].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:43.138625Z","iopub.execute_input":"2025-06-02T15:29:43.138868Z","iopub.status.idle":"2025-06-02T15:29:43.153999Z","shell.execute_reply.started":"2025-06-02T15:29:43.138822Z","shell.execute_reply":"2025-06-02T15:29:43.153300Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_dicom(path):\n    dicom = pydicom.dcmread(path)\n    data = dicom.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:43.154688Z","iopub.execute_input":"2025-06-02T15:29:43.154925Z","iopub.status.idle":"2025-06-02T15:29:43.167671Z","shell.execute_reply.started":"2025-06-02T15:29:43.154906Z","shell.execute_reply":"2025-06-02T15:29:43.166907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nimport matplotlib.pyplot as plt\n\n# Yeni sıfırlanmış indekslerle rastgele seçim yapalım\nfinal_merged_df_reset = final_merged_df.reset_index(drop=True)\n\n# Rastgele iki indeks seçelim\nselected_indices = random.sample(range(len(final_merged_df_reset)), 2)\n\nimages = []\nrow_ids = []\n\n# Seçilen indekslerle görselleri yükleyelim\nfor i in selected_indices:\n    image = load_dicom(final_merged_df_reset['image_path'][i])  # Yeni sıfırlanmış indeksi kullan\n    images.append(image)\n    row_ids.append(final_merged_df_reset['row_id'][i])  # Yeni sıfırlanmış indeksi kullan\n\n# Görselleri çizdirelim\nfig, ax = plt.subplots(1, 2, figsize=(8, 4))\nfor i in range(2):\n    ax[i].imshow(images[i], cmap='gray')\n    ax[i].set_title(f'Row ID: {row_ids[i]}', fontsize=8)\n    ax[i].axis('off')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:43.171171Z","iopub.execute_input":"2025-06-02T15:29:43.171441Z","iopub.status.idle":"2025-06-02T15:29:43.652914Z","shell.execute_reply.started":"2025-06-02T15:29:43.171409Z","shell.execute_reply":"2025-06-02T15:29:43.651933Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:43.654318Z","iopub.execute_input":"2025-06-02T15:29:43.654625Z","iopub.status.idle":"2025-06-02T15:29:43.668607Z","shell.execute_reply.started":"2025-06-02T15:29:43.654598Z","shell.execute_reply":"2025-06-02T15:29:43.667898Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = train_data.dropna()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:43.669524Z","iopub.execute_input":"2025-06-02T15:29:43.669923Z","iopub.status.idle":"2025-06-02T15:29:43.705340Z","shell.execute_reply.started":"2025-06-02T15:29:43.669890Z","shell.execute_reply":"2025-06-02T15:29:43.704416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torchvision.models as models\nfrom torchvision.models import ResNet50_Weights\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    confusion_matrix,\n    classification_report,\n    roc_curve,\n    auc,\n    accuracy_score\n)\n\nimport matplotlib.pyplot as plt\nimport itertools\nfrom tqdm import tqdm\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:43.706355Z","iopub.execute_input":"2025-06-02T15:29:43.706585Z","iopub.status.idle":"2025-06-02T15:29:46.263154Z","shell.execute_reply.started":"2025-06-02T15:29:43.706566Z","shell.execute_reply":"2025-06-02T15:29:46.262475Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torchvision.models as models\nfrom torchvision.models import ResNet50_Weights\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    confusion_matrix,\n    classification_report,\n    roc_curve,\n    auc,\n    accuracy_score\n)\n\nimport matplotlib.pyplot as plt\nimport itertools\nfrom tqdm import tqdm\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:46.264033Z","iopub.execute_input":"2025-06-02T15:29:46.264554Z","iopub.status.idle":"2025-06-02T15:29:46.269221Z","shell.execute_reply.started":"2025-06-02T15:29:46.264530Z","shell.execute_reply":"2025-06-02T15:29:46.268335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nfrom PIL import Image\n\ndef load_dicom(path):\n    \"\"\"\n    Reads a single DICOM file and returns a 2D NumPy array normalized to [0,255].\n    \"\"\"\n    d = pydicom.dcmread(path)\n    arr = d.pixel_array.astype(np.float32)\n    arr -= arr.min()\n    arr /= arr.max()\n    arr = (arr * 255).astype(np.uint8)\n    return arr\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:46.270097Z","iopub.execute_input":"2025-06-02T15:29:46.270327Z","iopub.status.idle":"2025-06-02T15:29:46.292733Z","shell.execute_reply.started":"2025-06-02T15:29:46.270307Z","shell.execute_reply":"2025-06-02T15:29:46.291904Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, dataframe, transform=None, augment_severe=False):\n        \"\"\"\n        dataframe: pandas DataFrame with columns ['image_path', 'severity', 'series_description', …]\n        transform: torchvision.transforms pipeline (ToPILImage → Resize → Grayscale(3) → ToTensor)\n        augment_severe: if True, apply one random augmentation when severity == 'Severe'\n        \"\"\"\n        self.dataframe = dataframe.reset_index(drop=True)\n        self.transform = transform\n        self.augment_severe = augment_severe\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, index):\n        # 1) Load the DICOM slice\n        image_path = self.dataframe.loc[index, 'image_path']\n        image = load_dicom(image_path)  # 2D NumPy array (H, W)\n\n        # 2) Severity label (string)\n        label = self.dataframe.loc[index, 'severity']\n\n        # 3) If it’s Severe and augment_severe=True, apply one random augmentation\n        if self.augment_severe and label == 'Severe':\n            image = self.apply_augmentation(image)\n\n        # 4) Apply the torchvision transform (converts to 3‐channel tensor)\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label\n\n    def apply_augmentation(self, image):\n        \"\"\"\n        Randomly pick one augmentation from the list and apply it to the 2D NumPy image.\n        \"\"\"\n        augmentations = [\n            lambda img: np.rot90(img, k=np.random.randint(1, 4)),  # Random 90° rotations\n            lambda img: np.fliplr(img),                            # Horizontal flip\n            lambda img: np.flipud(img),                            # Vertical flip\n            lambda img: self.random_crop(img),                      # Random crop to 224×224\n            lambda img: self.random_zoom(img),                      # Random zoom (returns original)\n        ]\n        augmentation = np.random.choice(augmentations)\n        return augmentation(image)\n\n    def random_crop(self, image, crop_size=(224, 224)):\n        \"\"\"\n        Randomly crop a (224, 224) patch from the input 2D NumPy image.\n        If image is smaller, return it unchanged.\n        \"\"\"\n        h, w = image.shape\n        new_h, new_w = crop_size\n        if h <= new_h or w <= new_w:\n            return image\n        top = np.random.randint(0, h - new_h + 1)\n        left = np.random.randint(0, w - new_w + 1)\n        return image[top : top + new_h, left : left + new_w]\n\n    def random_zoom(self, image, zoom_range=(0.8, 1.2)):\n        \"\"\"\n        Stub for random zoom. Returns original image unchanged.\n        \"\"\"\n        return image\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:46.293704Z","iopub.execute_input":"2025-06-02T15:29:46.294032Z","iopub.status.idle":"2025-06-02T15:29:46.309874Z","shell.execute_reply.started":"2025-06-02T15:29:46.294002Z","shell.execute_reply":"2025-06-02T15:29:46.309043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Clean out any NaNs in 'severity' or 'series_description'\nfinal_merged_df = final_merged_df.dropna(subset=['severity', 'series_description']).reset_index(drop=True)\n\n# 1) Hold out 15% as test (stratify on severity)\ntrain_val_df, test_df = train_test_split(\n    final_merged_df,\n    test_size=0.15,\n    random_state=42,\n    stratify=final_merged_df['severity']\n)\n\n# 2) Split remaining 85% into train (80%) and val (20%), stratified\ntrain_df, val_df = train_test_split(\n    train_val_df,\n    test_size=0.20,\n    random_state=42,\n    stratify=train_val_df['severity']\n)\n\n# Reset indices\ntrain_df = train_df.reset_index(drop=True)\nval_df   = val_df.reset_index(drop=True)\ntest_df  = test_df.reset_index(drop=True)\n\nprint(\"Sizes → Train:\", len(train_df), \"Val:\", len(val_df), \"Test:\", len(test_df))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:46.310759Z","iopub.execute_input":"2025-06-02T15:29:46.311014Z","iopub.status.idle":"2025-06-02T15:29:46.455134Z","shell.execute_reply.started":"2025-06-02T15:29:46.310993Z","shell.execute_reply":"2025-06-02T15:29:46.454322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_datasets_and_loaders(\n    train_df, val_df,\n    series_description,\n    transform,\n    batch_size=8,\n    augment_severe=False\n):\n    \"\"\"\n    Build train + val DataLoaders for a given 'series_description'.\n    train_df, val_df: pre‐split DataFrames.\n    \"\"\"\n    # Filter each by series_description\n    train_view_df = train_df[ train_df['series_description'] == series_description ].reset_index(drop=True)\n    val_view_df   = val_df[   val_df['series_description']   == series_description ].reset_index(drop=True)\n\n    # Create CustomDataset instances\n    train_dataset = CustomDataset(train_view_df, transform=transform, augment_severe=augment_severe)\n    val_dataset   = CustomDataset(val_view_df,   transform=transform, augment_severe=False)\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    return trainloader, valloader, len(train_view_df), len(val_view_df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:46.455882Z","iopub.execute_input":"2025-06-02T15:29:46.456094Z","iopub.status.idle":"2025-06-02T15:29:46.461276Z","shell.execute_reply.started":"2025-06-02T15:29:46.456076Z","shell.execute_reply":"2025-06-02T15:29:46.460280Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --------------------------------\n# Define transforms: DICOM(2D) → 3-channel Tensor(3×224×224)\n# --------------------------------\ntransform = transforms.Compose([\n    transforms.Lambda(lambda x: (x * 255).astype(np.uint8)),  # Convert float->uint8\n    transforms.ToPILImage(),           # PILImage from NumPy array\n    transforms.Resize((224, 224)),     # Resize to 224×224\n    transforms.Grayscale(num_output_channels=3),  # Make 3 channels\n    transforms.ToTensor(),             # Convert to Tensor, range [0,1]\n    # You can add Normalize(...) here if desired\n])\n\n# --------------------------------\n# Build Train/Val Loaders for each view (augment only Severe in train)\n# --------------------------------\ntrainloader_t1, valloader_t1, len_train_t1, len_val_t1 = create_datasets_and_loaders(\n    train_df, val_df, 'Sagittal T1', transform, batch_size=8, augment_severe=True\n)\ntrainloader_t2, valloader_t2, len_train_t2, len_val_t2 = create_datasets_and_loaders(\n    train_df, val_df, 'Axial T2', transform, batch_size=8, augment_severe=True\n)\ntrainloader_t2stir, valloader_t2stir, len_train_t2stir, len_val_t2stir = create_datasets_and_loaders(\n    train_df, val_df, 'Sagittal T2/STIR', transform, batch_size=8, augment_severe=True\n)\n\n# --------------------------------\n# Build Test Loaders for each view (no augmentation)\n# --------------------------------\ntest_dataset_t1 = CustomDataset(\n    test_df[test_df['series_description'] == 'Sagittal T1'].reset_index(drop=True),\n    transform=transform,\n    augment_severe=False\n)\ntest_loader_t1 = DataLoader(test_dataset_t1, batch_size=8, shuffle=False)\n\ntest_dataset_t2 = CustomDataset(\n    test_df[test_df['series_description'] == 'Axial T2'].reset_index(drop=True),\n    transform=transform,\n    augment_severe=False\n)\ntest_loader_t2 = DataLoader(test_dataset_t2, batch_size=8, shuffle=False)\n\ntest_dataset_t2stir = CustomDataset(\n    test_df[test_df['series_description'] == 'Sagittal T2/STIR'].reset_index(drop=True),\n    transform=transform,\n    augment_severe=False\n)\ntest_loader_t2stir = DataLoader(test_dataset_t2stir, batch_size=8, shuffle=False)\n\n# Map severity string → integer index\nlabel_map = {'normal_mild': 0, 'moderate': 1, 'severe': 2}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:46.462042Z","iopub.execute_input":"2025-06-02T15:29:46.462257Z","iopub.status.idle":"2025-06-02T15:29:46.509632Z","shell.execute_reply.started":"2025-06-02T15:29:46.462238Z","shell.execute_reply":"2025-06-02T15:29:46.508923Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CustomResNet50(nn.Module):\n    def __init__(self, num_classes=3, pretrained_weights=None):\n        super(CustomResNet50, self).__init__()\n        self.model = models.resnet50(weights=ResNet50_Weights.IMAGENET1K_V1)\n        if pretrained_weights:\n            self.model.load_state_dict(torch.load(pretrained_weights))\n        num_ftrs = self.model.fc.in_features\n        self.model.fc = nn.Linear(num_ftrs, num_classes)\n\n    def forward(self, x):\n        return self.model(x)\n\n    def unfreeze_middle_layers(self):\n        \"\"\"\n        Unfreeze ResNet layers in layer3 and layer4; freeze the rest.\n        \"\"\"\n        for name, param in self.model.named_parameters():\n            if ('layer3' in name) or ('layer4' in name):\n                param.requires_grad = True\n            else:\n                param.requires_grad = False\n\n# Device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Instantiate three models\nsagittal_t1_model   = CustomResNet50(num_classes=3).to(device)\naxial_t2_model      = CustomResNet50(num_classes=3).to(device)\nsagittal_t2stir_model = CustomResNet50(num_classes=3).to(device)\n\n# Unfreeze middle layers\nfor m in [sagittal_t1_model, axial_t2_model, sagittal_t2stir_model]:\n    m.unfreeze_middle_layers()\n\n# Loss (weighted cross‐entropy)\nweights = torch.tensor([1.0, 2.0, 4.0]).to(device)  # heavier for Severe\ncriterion = nn.CrossEntropyLoss(weight=weights)\n\n# Optimizers\noptimizer_sagittal_t1   = optim.Adam(sagittal_t1_model.parameters(),   lr=1e-3)\noptimizer_axial_t2      = optim.Adam(axial_t2_model.parameters(),      lr=1e-3)\noptimizer_sagittal_t2stir = optim.Adam(sagittal_t2stir_model.parameters(), lr=1e-3)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:46.510364Z","iopub.execute_input":"2025-06-02T15:29:46.510571Z","iopub.status.idle":"2025-06-02T15:29:48.910795Z","shell.execute_reply.started":"2025-06-02T15:29:46.510552Z","shell.execute_reply":"2025-06-02T15:29:48.910117Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(\n    model, trainloader, valloader, \n    len_train, len_val, \n    optimizer, \n    num_epochs=10, \n    patience=3, \n    model_desc=\"model\"\n):\n    \"\"\"\n    Trains the model for up to num_epochs (with early stopping).\n    Returns: (best_model, best_val_acc, train_losses, train_accs, val_losses, val_accs)\n    \"\"\"\n    # Learning‐rate scheduler\n    scheduler = lr_scheduler.StepLR(optimizer, step_size=2, gamma=0.1)\n\n    best_val_acc  = 0.0\n    best_val_loss = float('inf')\n    best_model_wts = model.state_dict().copy()\n    counter = 0\n\n    # Lists to record per‐epoch metrics\n    train_losses = []\n    train_accs   = []\n    val_losses   = []\n    val_accs     = []\n\n    for epoch in range(num_epochs):\n        # ------------ TRAIN ------------\n        model.train()\n        running_loss = 0.0\n        running_corrects = 0\n        total_samples = 0\n\n        with tqdm(trainloader, unit=\"batch\", desc=f\"[{model_desc}] Epoch {epoch+1}/{num_epochs} (Train)\") as tepoch:\n            for images, labels in tepoch:\n                images = images.to(device)\n                labels = torch.tensor([label_map[lbl] for lbl in labels]).to(device)\n\n                optimizer.zero_grad()\n                outputs = model(images)         # logits (B, 3)\n                loss = criterion(outputs, labels)\n                loss.backward()\n                optimizer.step()\n\n                batch_size = images.size(0)\n                running_loss += loss.item() * batch_size\n                probs = torch.softmax(outputs, dim=1)\n                _, preds = torch.max(probs, 1)\n                running_corrects += torch.sum(preds == labels.data).item()\n                total_samples += batch_size\n\n                tepoch.set_postfix(train_loss=running_loss/total_samples)\n\n        scheduler.step()\n        epoch_train_loss = running_loss / total_samples\n        epoch_train_acc  = running_corrects / total_samples\n        train_losses.append(epoch_train_loss)\n        train_accs.append(epoch_train_acc)\n\n        # --------- VALIDATE ----------\n        model.eval()\n        val_running_loss = 0.0\n        val_running_corrects = 0\n        val_samples = 0\n\n        with torch.no_grad():\n            with tqdm(valloader, unit=\"batch\", desc=f\"[{model_desc}] Epoch {epoch+1}/{num_epochs} (Val)\") as vepoch:\n                for images, labels in vepoch:\n                    images = images.to(device)\n                    labels = torch.tensor([label_map[lbl] for lbl in labels]).to(device)\n\n                    outputs = model(images)\n                    loss = criterion(outputs, labels)\n\n                    batch_size = images.size(0)\n                    val_running_loss += loss.item() * batch_size\n                    probs = torch.softmax(outputs, dim=1)\n                    if probs.dim() == 1:\n                        _, preds = torch.max(probs, 0)\n                    else:\n                        _, preds = torch.max(probs, 1)\n                    val_running_corrects += torch.sum(preds == labels.data).item()\n                    val_samples += batch_size\n\n                    vepoch.set_postfix(val_loss=val_running_loss/val_samples)\n\n        epoch_val_loss = val_running_loss / val_samples\n        epoch_val_acc  = val_running_corrects / val_samples\n        val_losses.append(epoch_val_loss)\n        val_accs.append(epoch_val_acc)\n\n        print(\n            f\"Epoch {epoch+1}/{num_epochs} → \"\n            f\"Train Loss: {epoch_train_loss:.4f}, Train Acc: {epoch_train_acc*100:.2f}%  |  \"\n            f\"Val Loss: {epoch_val_loss:.4f}, Val Acc: {epoch_val_acc*100:.2f}%\"\n        )\n\n        # Save best model\n        if (epoch_val_acc > best_val_acc) or (\n            epoch_val_acc == best_val_acc and epoch_val_loss < best_val_loss\n        ):\n            best_val_acc  = epoch_val_acc\n            best_val_loss = epoch_val_loss\n            best_model_wts = model.state_dict().copy()\n            counter = 0\n\n            # Save to disk\n            safe_name = model_desc.replace(\" \", \"_\").replace(\"/\", \"_\")\n            model_path = f\"best_model_{safe_name}.pth\"\n            if os.path.exists(model_path):\n                os.remove(model_path)\n            torch.save(best_model_wts, model_path)\n        else:\n            counter += 1\n\n        if counter >= patience:\n            print(f\"Early stopping after epoch {epoch+1}\")\n            break\n\n    # Load best weights\n    model.load_state_dict(best_model_wts)\n    return model, best_val_acc, train_losses, train_accs, val_losses, val_accs\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:48.911646Z","iopub.execute_input":"2025-06-02T15:29:48.911933Z","iopub.status.idle":"2025-06-02T15:29:48.924134Z","shell.execute_reply.started":"2025-06-02T15:29:48.911909Z","shell.execute_reply":"2025-06-02T15:29:48.923142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_epochs = 10\npatience   = 3\n\n# --- Sagittal T1 ---\nprint(\"\\n=== Training Sagittal T1 Model ===\")\nsagittal_t1_model, best_acc_t1, \\\n train_losses_t1, train_accs_t1, \\\n val_losses_t1,   val_accs_t1   = train_model(\n       sagittal_t1_model,\n       trainloader_t1,\n       valloader_t1,\n       len_train_t1,\n       len_val_t1,\n       optimizer_sagittal_t1,\n       num_epochs=num_epochs,\n       patience=patience,\n       model_desc=\"Sagittal_T1\"\n)\n\nplt.figure(figsize=(12,5))\nplt.subplot(1, 2, 1)\nplt.plot(train_losses_t1, marker='o', label=\"Train Loss\")\nplt.plot(val_losses_t1,   marker='o', label=\"Val Loss\")\nplt.title(\"Sagittal T1: Loss vs. Epoch\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()\n\nplt.subplot(1, 2, 2)\nplt.plot([x * 100 for x in train_accs_t1], marker='o', label=\"Train Acc\")\nplt.plot([x * 100 for x in val_accs_t1],   marker='o', label=\"Val Acc\")\nplt.title(\"Sagittal T1: Accuracy vs. Epoch\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy (%)\")\nplt.legend()\n\nplt.tight_layout()\nplt.show()\n\n# --- Axial T2 ---\nprint(\"\\n=== Training Axial T2 Model ===\")\naxial_t2_model, best_acc_t2, \\\n train_losses_t2, train_accs_t2, \\\n val_losses_t2,   val_accs_t2   = train_model(\n       axial_t2_model,\n       trainloader_t2,\n       valloader_t2,\n       len_train_t2,\n       len_val_t2,\n       optimizer_axial_t2,\n       num_epochs=num_epochs,\n       patience=patience,\n       model_desc=\"Axial_T2\"\n)\n\nplt.figure(figsize=(12,5))\nplt.subplot(1, 2, 1)\nplt.plot(train_losses_t2, marker='o', label=\"Train Loss\")\nplt.plot(val_losses_t2,   marker='o', label=\"Val Loss\")\nplt.title(\"Axial T2: Loss vs. Epoch\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()\n\nplt.subplot(1, 2, 2)\nplt.plot([x * 100 for x in train_accs_t2], marker='o', label=\"Train Acc\")\nplt.plot([x * 100 for x in val_accs_t2],   marker='o', label=\"Val Acc\")\nplt.title(\"Axial T2: Accuracy vs. Epoch\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy (%)\")\nplt.legend()\n\nplt.tight_layout()\nplt.show()\n\n# --- Sagittal T2/STIR ---\nprint(\"\\n=== Training Sagittal T2/STIR Model ===\")\nsagittal_t2stir_model, best_acc_t2stir, \\\n train_losses_t2stir, train_accs_t2stir, \\\n val_losses_t2stir,   val_accs_t2stir = train_model(\n       sagittal_t2stir_model,\n       trainloader_t2stir,\n       valloader_t2stir,\n       len_train_t2stir,\n       len_val_t2stir,\n       optimizer_sagittal_t2stir,\n       num_epochs=num_epochs,\n       patience=patience,\n       model_desc=\"Sagittal_T2_STIR\"\n)\n\nplt.figure(figsize=(12,5))\nplt.subplot(1, 2, 1)\nplt.plot(train_losses_t2stir, marker='o', label=\"Train Loss\")\nplt.plot(val_losses_t2stir,   marker='o', label=\"Val Loss\")\nplt.title(\"Sagittal T2/STIR: Loss vs. Epoch\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()\n\nplt.subplot(1, 2, 2)\nplt.plot([x * 100 for x in train_accs_t2stir], marker='o', label=\"Train Acc\")\nplt.plot([x * 100 for x in val_accs_t2stir],   marker='o', label=\"Val Acc\")\nplt.title(\"Sagittal T2/STIR: Accuracy vs. Epoch\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy (%)\")\nplt.legend()\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T15:29:48.925048Z","iopub.execute_input":"2025-06-02T15:29:48.925301Z","iopub.status.idle":"2025-06-02T16:41:04.624209Z","shell.execute_reply.started":"2025-06-02T15:29:48.925280Z","shell.execute_reply":"2025-06-02T16:41:04.623293Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def _plot_confusion_matrix(cm, classes, normalize=False, title=\"Confusion matrix\", cmap=plt.cm.Blues):\n    if normalize:\n        cm = cm.astype(\"float\") / cm.sum(axis=1)[:, np.newaxis]\n\n    plt.imshow(cm, interpolation=\"nearest\", cmap=cmap)\n    plt.title(title)\n    plt.colorbar()\n    tick_marks = np.arange(len(classes))\n    plt.xticks(tick_marks, classes, rotation=45)\n    plt.yticks(tick_marks, classes)\n\n    fmt = \".2f\" if normalize else \"d\"\n    thresh = cm.max() / 2.0\n    for i, j in itertools.product(range(cm.shape[0]), range(cm.shape[1])):\n        plt.text(j, i, format(cm[i, j], fmt),\n                 ha=\"center\",\n                 color=\"white\" if cm[i, j] > thresh else \"black\")\n\n    plt.ylabel(\"True label\")\n    plt.xlabel(\"Predicted label\")\n    plt.tight_layout()\n\ndef evaluate_model(model, test_loader, view_name=\"Model\"):\n    model.eval()\n    y_true_list  = []\n    y_pred_list  = []\n    y_proba_list = []\n\n    with torch.no_grad():\n        for images, labels in test_loader:\n            images = images.to(device)\n            labels = torch.tensor([label_map[lbl] for lbl in labels]).to(device)\n\n            outputs = model(images)\n            probs   = torch.softmax(outputs, dim=1)\n            _, preds = torch.max(probs, dim=1)\n\n            y_true_list.extend(labels.cpu().numpy())\n            y_pred_list.extend(preds.cpu().numpy())\n            y_proba_list.extend(probs.cpu().numpy())\n\n    y_true  = np.array(y_true_list)\n    y_pred  = np.array(y_pred_list)\n    y_proba = np.array(y_proba_list)  # shape (N_test, 3)\n\n    # 1) Overall accuracy\n    acc = accuracy_score(y_true, y_pred)\n    print(f\"\\n***** {view_name} Test Accuracy: {acc*100:.2f}% *****\\n\")\n\n    # 2) Classification report\n    print(\"Classification Report:\")\n    print(classification_report(\n        y_true, y_pred, \n        target_names=[\"Normal/Mild\", \"Moderate\", \"Severe\"]\n    ))\n\n    # 3) Confusion matrix\n    cm = confusion_matrix(y_true, y_pred)\n    plt.figure(figsize=(5,5))\n    _plot_confusion_matrix(\n        cm,\n        classes=[\"Normal/Mild\", \"Moderate\", \"Severe\"],\n        normalize=False,\n        title=f\"{view_name} Confusion Matrix\"\n    )\n    plt.show()\n\n    plt.figure(figsize=(5,5))\n    _plot_confusion_matrix(\n        cm,\n        classes=[\"Normal/Mild\", \"Moderate\", \"Severe\"],\n        normalize=True,\n        title=f\"{view_name} Normalized Confusion Matrix\"\n    )\n    plt.show()\n\n    # 4) ROC curves + AUC\n    plt.figure(figsize=(6,6))\n    fpr    = {}\n    tpr    = {}\n    roc_auc = {}\n\n    for i, class_name in enumerate([\"Normal/Mild\", \"Moderate\", \"Severe\"]):\n        y_true_bin = (y_true == i).astype(int)\n        y_score    = y_proba[:, i]\n        fpr[i], tpr[i], _ = roc_curve(y_true_bin, y_score)\n        roc_auc[i] = auc(fpr[i], tpr[i])\n        plt.plot(fpr[i], tpr[i], lw=2, label=f\"{class_name} (AUC = {roc_auc[i]:.2f})\")\n\n    plt.plot([0,1], [0,1], color=\"grey\", lw=1, linestyle=\"--\")\n    plt.xlim([0.0, 1.0])\n    plt.ylim([0.0, 1.0])\n    plt.xlabel(\"False Positive Rate\")\n    plt.ylabel(\"True Positive Rate\")\n    plt.title(f\"{view_name} ROC Curves\")\n    plt.legend(loc=\"lower right\")\n    plt.grid(True)\n    plt.show()\n\n    return acc, cm, roc_auc\n\n# Finally, call evaluate_model for each view’s test loader:\nprint(\"\\n=== Evaluating on Test Set ===\\n\")\n\nacc_t1, cm_t1, roc_auc_t1 = evaluate_model(\n    sagittal_t1_model, test_loader_t1, view_name=\"Sagittal T1\"\n)\n\nacc_t2, cm_t2, roc_auc_t2 = evaluate_model(\n    axial_t2_model, test_loader_t2, view_name=\"Axial T2\"\n)\n\nacc_t2stir, cm_t2stir, roc_auc_t2stir = evaluate_model(\n    sagittal_t2stir_model, test_loader_t2stir, view_name=\"Sagittal T2/STIR\"\n)\n\n# Average test accuracy across all three views\navg_test_acc = (acc_t1 + acc_t2 + acc_t2stir) / 3.0\nprint(f\"\\nAverage Test Accuracy (all three views): {avg_test_acc*100:.2f}%\\n\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-02T16:41:04.625245Z","iopub.execute_input":"2025-06-02T16:41:04.625477Z","iopub.status.idle":"2025-06-02T16:43:03.030792Z","shell.execute_reply.started":"2025-06-02T16:41:04.625455Z","shell.execute_reply":"2025-06-02T16:43:03.030109Z"}},"outputs":[],"execution_count":null}]}