{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport torch\n\n# Confirm GPU is available\nprint(\"GPU available:\", torch.cuda.is_available())\nprint(\"GPU name:\", torch.cuda.get_device_name(0) if torch.cuda.is_available() else \"None\")\n\n# Find the correct dataset path (Kaggle nests competition datasets under /kaggle/input/competitions/)\nbase_path = '/kaggle/input/competitions/rsna-2024-lumbar-spine-degenerative-classification'\n\nif not os.path.exists(base_path):\n    # fallback: search for it\n    print(\"\\nExpected path not found — searching /kaggle/input ...\")\n    for root, dirs, files in os.walk('/kaggle/input'):\n        if 'train.csv' in files:\n            base_path = root\n            break\n\nprint(\"\\nUsing base_path:\", base_path)\nprint(\"\\nFiles in dataset:\")\nfor f in os.listdir(base_path):\n    print(\" -\", f)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-08-15T12:57:48.193009Z","iopub.execute_input":"2026-08-15T12:57:48.193312Z","iopub.status.idle":"2026-08-15T12:57:54.05136Z","shell.execute_reply.started":"2026-08-15T12:57:48.193279Z","shell.execute_reply":"2026-08-15T12:57:54.050647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pd.read_csv(f'{base_path}/train.csv')\nprint(\"train.csv shape:\", train.shape)\nprint(train.columns.tolist())\ntrain.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T12:57:54.052779Z","iopub.execute_input":"2026-08-15T12:57:54.053166Z","iopub.status.idle":"2026-08-15T12:57:54.11246Z","shell.execute_reply.started":"2026-08-15T12:57:54.053085Z","shell.execute_reply":"2026-08-15T12:57:54.111732Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"series_desc = pd.read_csv(f'{base_path}/train_series_descriptions.csv')\nprint(\"\\ntrain_series_descriptions.csv shape:\", series_desc.shape)\nseries_desc.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T12:57:54.11343Z","iopub.execute_input":"2026-08-15T12:57:54.113766Z","iopub.status.idle":"2026-08-15T12:57:54.135809Z","shell.execute_reply.started":"2026-08-15T12:57:54.113743Z","shell.execute_reply":"2026-08-15T12:57:54.13519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"coords = pd.read_csv(f'{base_path}/train_label_coordinates.csv')\nprint(\"\\ntrain_label_coordinates.csv shape:\", coords.shape)\ncoords.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T12:57:54.136666Z","iopub.execute_input":"2026-08-15T12:57:54.136952Z","iopub.status.idle":"2026-08-15T12:57:54.265675Z","shell.execute_reply.started":"2026-08-15T12:57:54.136922Z","shell.execute_reply":"2026-08-15T12:57:54.264831Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check for missing labels (some studies may have incomplete annotations)\nprint(\"Missing values per column:\")\nprint(train.isnull().sum()[train.isnull().sum() > 0])\n\n# Check class balance — severity classes are usually heavily imbalanced\nlabel_cols = [c for c in train.columns if c != 'study_id']\nall_labels = train[label_cols].values.flatten()\nimport collections\nprint(\"\\nClass distribution across all labels:\")\nprint(collections.Counter(all_labels))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T12:57:54.267512Z","iopub.execute_input":"2026-08-15T12:57:54.267782Z","iopub.status.idle":"2026-08-15T12:57:54.288632Z","shell.execute_reply.started":"2026-08-15T12:57:54.26776Z","shell.execute_reply":"2026-08-15T12:57:54.287837Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\n\n# Grab one sample DICOM file to inspect available metadata\nsample_study = train['study_id'].iloc[0]\nsample_series = series_desc[series_desc['study_id'] == sample_study]['series_id'].iloc[0]\n\nseries_folder = f'{base_path}/train_images/{sample_study}/{sample_series}'\nsample_file = os.listdir(series_folder)[0]\ndcm = pydicom.dcmread(f'{series_folder}/{sample_file}')\n\nprint(dcm)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T12:57:54.289539Z","iopub.execute_input":"2026-08-15T12:57:54.289814Z","iopub.status.idle":"2026-08-15T12:57:55.034073Z","shell.execute_reply.started":"2026-08-15T12:57:54.289784Z","shell.execute_reply":"2026-08-15T12:57:55.033279Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nfrom tqdm import tqdm\n\nmetadata_records = []\n\nfor study_id in tqdm(train['study_id'].unique()):\n    study_series = series_desc[series_desc['study_id'] == study_id]\n    for _, row in study_series.iterrows():\n        series_id = row['series_id']\n        series_folder = f'{base_path}/train_images/{study_id}/{series_id}'\n        try:\n            files = os.listdir(series_folder)\n            dcm = pydicom.dcmread(f'{series_folder}/{files[0]}')\n            metadata_records.append({\n                'study_id': study_id,\n                'series_id': series_id,\n                'series_description': row['series_description'],\n                'num_slices': len(files),\n                'rows': dcm.Rows,\n                'columns': dcm.Columns,\n                'pixel_spacing_x': float(dcm.PixelSpacing[0]),\n                'pixel_spacing_y': float(dcm.PixelSpacing[1]),\n                'slice_thickness': float(dcm.SliceThickness),\n                'spacing_between_slices': float(getattr(dcm, 'SpacingBetweenSlices', np.nan)),\n                'photometric_interpretation': dcm.PhotometricInterpretation,\n                'patient_position': dcm.PatientPosition,\n            })\n        except Exception as e:\n            metadata_records.append({\n                'study_id': study_id, 'series_id': series_id,\n                'series_description': row['series_description'],\n                'error': str(e)\n            })\n\nmetadata_df = pd.DataFrame(metadata_records)\nprint(metadata_df.shape)\nmetadata_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T12:57:55.035068Z","iopub.execute_input":"2026-08-15T12:57:55.035379Z","iopub.status.idle":"2026-08-15T12:59:46.683072Z","shell.execute_reply.started":"2026-08-15T12:57:55.035348Z","shell.execute_reply":"2026-08-15T12:59:46.682319Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Confirm no extraction errors\nprint(\"Errors:\", metadata_df['error'].notna().sum() if 'error' in metadata_df.columns else 0)\n\n# Check variability across the dataset — this tells us how much standardization is needed\nprint(\"\\nRows/Columns distribution:\")\nprint(metadata_df.groupby(['rows', 'columns']).size().sort_values(ascending=False).head(10))\n\nprint(\"\\nPixel spacing range:\")\nprint(metadata_df[['pixel_spacing_x', 'pixel_spacing_y']].describe())\n\nprint(\"\\nSlice thickness distribution:\")\nprint(metadata_df['slice_thickness'].value_counts())\n\nprint(\"\\nSeries per study — should mostly be 3, but check for duplicates:\")\nseries_counts = metadata_df.groupby('study_id').size()\nprint(series_counts.value_counts().sort_index())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:00:02.244163Z","iopub.execute_input":"2026-08-15T13:00:02.244693Z","iopub.status.idle":"2026-08-15T13:00:02.272434Z","shell.execute_reply.started":"2026-08-15T13:00:02.244666Z","shell.execute_reply":"2026-08-15T13:00:02.27169Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Studies with incomplete series (only 2, missing a view)\nincomplete_studies = series_counts[series_counts == 2].index.tolist()\nprint(\"Incomplete studies (only 2 series):\", incomplete_studies)\nfor sid in incomplete_studies:\n    print(sid, metadata_df[metadata_df['study_id'] == sid]['series_description'].tolist())\n\n# Studies with duplicate series types (e.g. two Axial T2)\ndupe_check = metadata_df.groupby(['study_id', 'series_description']).size().reset_index(name='count')\ndupes = dupe_check[dupe_check['count'] > 1]\nprint(f\"\\nStudies with duplicate series of the same type: {dupes['study_id'].nunique()}\")\nprint(dupes.head(15))\n\n# Save the metadata table — you'll need this for every later stage\nmetadata_df.to_csv('/kaggle/working/series_metadata.csv', index=False)\nprint(\"\\nSaved metadata_df to /kaggle/working/series_metadata.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:00:06.630091Z","iopub.execute_input":"2026-08-15T13:00:06.630936Z","iopub.status.idle":"2026-08-15T13:00:06.710688Z","shell.execute_reply.started":"2026-08-15T13:00:06.63089Z","shell.execute_reply":"2026-08-15T13:00:06.709963Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Determine which series to keep when duplicates exist (keep the one with more slices)\ndef select_series(group):\n    if len(group) == 1:\n        return group\n    return group.sort_values('num_slices', ascending=False).iloc[[0]]\n\nkept_series = metadata_df.groupby(['study_id', 'series_description'], group_keys=False).apply(select_series)\nprint(\"Series after dedup:\", len(kept_series), \"(was\", len(metadata_df), \")\")\n\n# Assign QC status per study\nqc_status = []\nfor study_id in train['study_id'].unique():\n    views = kept_series[kept_series['study_id'] == study_id]['series_description'].tolist()\n    required = {'Sagittal T1', 'Sagittal T2/STIR', 'Axial T2'}\n    missing = required - set(views)\n    if not missing:\n        status = 'accepted'\n    elif len(missing) == 1:\n        status = 'requires_manual_review'\n    else:\n        status = 'excluded'\n    qc_status.append({'study_id': study_id, 'qc_status': status, 'missing_views': list(missing)})\n\nqc_df = pd.DataFrame(qc_status)\nprint(\"\\nQC status distribution:\")\nprint(qc_df['qc_status'].value_counts())\n\n# Save both tables\nkept_series.to_csv('/kaggle/working/series_metadata_deduped.csv', index=False)\nqc_df.to_csv('/kaggle/working/qc_status.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:00:11.560486Z","iopub.execute_input":"2026-08-15T13:00:11.560888Z","iopub.status.idle":"2026-08-15T13:00:12.823516Z","shell.execute_reply.started":"2026-08-15T13:00:11.560859Z","shell.execute_reply":"2026-08-15T13:00:12.822616Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n# Merge QC status into train, keep only accepted studies for now\ntrain_qc = train.merge(qc_df, on='study_id')\naccepted = train_qc[train_qc['qc_status'] == 'accepted'].copy()\nprint(\"Accepted studies for splitting:\", len(accepted))\n\n# Create a stratification key: worst severity across all 25 labels per study\nseverity_rank = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\nseverity_numeric = accepted[label_cols].replace(severity_rank)\naccepted['worst_severity'] = severity_numeric.max(axis=1)\n\nprint(\"\\nWorst-severity distribution (stratification target):\")\nprint(accepted['worst_severity'].value_counts().sort_index())\n\n# 70/15/15 split, patient-level (study_id = patient here), stratified\ntrain_ids, temp_ids = train_test_split(\n    accepted['study_id'], test_size=0.30, stratify=accepted['worst_severity'], random_state=42\n)\nval_ids, test_ids = train_test_split(\n    temp_ids, test_size=0.50,\n    stratify=accepted.loc[accepted['study_id'].isin(temp_ids), 'worst_severity'],\n    random_state=42\n)\n\nprint(f\"\\nTrain: {len(train_ids)}, Val: {len(val_ids)}, Test: {len(test_ids)}\")\n\n# Save splits\nsplits_df = pd.DataFrame({\n    'study_id': pd.concat([train_ids, val_ids, test_ids]),\n    'split': ['train']*len(train_ids) + ['val']*len(val_ids) + ['test']*len(test_ids)\n})\nsplits_df.to_csv('/kaggle/working/patient_splits.csv', index=False)\nprint(\"\\nSaved patient_splits.csv\")\nsplits_df['split'].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:00:15.454585Z","iopub.execute_input":"2026-08-15T13:00:15.454839Z","iopub.status.idle":"2026-08-15T13:00:16.249227Z","shell.execute_reply.started":"2026-08-15T13:00:15.454818Z","shell.execute_reply":"2026-08-15T13:00:16.248358Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Confirm severity balance is preserved across splits\nverify = accepted.merge(splits_df, on='study_id')\nprint(pd.crosstab(verify['split'], verify['worst_severity'], normalize='index').round(3))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:00:19.130424Z","iopub.execute_input":"2026-08-15T13:00:19.130682Z","iopub.status.idle":"2026-08-15T13:00:19.162601Z","shell.execute_reply.started":"2026-08-15T13:00:19.130659Z","shell.execute_reply":"2026-08-15T13:00:19.161945Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Load a sample sagittal T2/STIR series and visualize a middle slice\nsample_study = train_ids.iloc[0]\nsample_row = kept_series[(kept_series['study_id'] == sample_study) & \n                           (kept_series['series_description'] == 'Sagittal T2/STIR')].iloc[0]\n\nseries_folder = f'{base_path}/train_images/{sample_study}/{sample_row[\"series_id\"]}'\nfiles = sorted(os.listdir(series_folder), key=lambda x: int(x.replace('.dcm', '')))\nmid_file = files[len(files) // 2]\n\ndcm = pydicom.dcmread(f'{series_folder}/{mid_file}')\nimg = dcm.pixel_array\n\nprint(\"Image Orientation (Patient):\", dcm.ImageOrientationPatient)\nprint(\"Image Position (Patient):\", dcm.ImagePositionPatient)\nprint(\"Photometric Interpretation:\", dcm.PhotometricInterpretation)\nprint(\"Shape:\", img.shape, \"dtype:\", img.dtype, \"min/max:\", img.min(), img.max())\n\nplt.figure(figsize=(6,6))\nplt.imshow(img, cmap='gray')\nplt.title(f'Study {sample_study} - Sagittal T2/STIR - slice {mid_file}')\nplt.axis('off')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:00:21.102893Z","iopub.execute_input":"2026-08-15T13:00:21.103592Z","iopub.status.idle":"2026-08-15T13:00:21.369625Z","shell.execute_reply.started":"2026-08-15T13:00:21.103562Z","shell.execute_reply":"2026-08-15T13:00:21.368805Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def classify_plane(iop):\n    \"\"\"Determine imaging plane from Image Orientation (Patient) cosines.\"\"\"\n    row_cosine = np.array(iop[:3])\n    col_cosine = np.array(iop[3:])\n    normal = np.cross(row_cosine, col_cosine)\n    abs_normal = np.abs(normal)\n    axis = np.argmax(abs_normal)\n    if axis == 0:\n        return 'sagittal'\n    elif axis == 1:\n        return 'coronal'\n    else:\n        return 'axial'\n\ndef standardize_orientation(img, iop):\n    \"\"\"Flip image so anatomical directions are consistent regardless of scanner-recorded orientation.\"\"\"\n    row_cosine = np.array(iop[:3])\n    col_cosine = np.array(iop[3:])\n    \n    # Flip rows if row_cosine points in negative direction (ensures consistent L/R or A/P facing)\n    dominant_row_axis = np.argmax(np.abs(row_cosine))\n    if row_cosine[dominant_row_axis] < 0:\n        img = np.fliplr(img)\n    \n    dominant_col_axis = np.argmax(np.abs(col_cosine))\n    if col_cosine[dominant_col_axis] < 0:\n        img = np.flipud(img)\n    \n    return img\n\n# Test it on our sample\nplane = classify_plane(dcm.ImageOrientationPatient)\nprint(\"Detected plane:\", plane)\n\nstandardized_img = standardize_orientation(img, dcm.ImageOrientationPatient)\n\nfig, axes = plt.subplots(1, 2, figsize=(12, 6))\naxes[0].imshow(img, cmap='gray')\naxes[0].set_title('Original')\naxes[0].axis('off')\naxes[1].imshow(standardized_img, cmap='gray')\naxes[1].set_title('Standardized')\naxes[1].axis('off')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:00:24.735782Z","iopub.execute_input":"2026-08-15T13:00:24.736226Z","iopub.status.idle":"2026-08-15T13:00:25.056969Z","shell.execute_reply.started":"2026-08-15T13:00:24.736195Z","shell.execute_reply":"2026-08-15T13:00:25.056194Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def standardize_orientation(img, iop, verbose=True):\n    row_cosine = np.array(iop[:3])\n    col_cosine = np.array(iop[3:])\n    \n    flipped_lr = False\n    flipped_ud = False\n    \n    dominant_row_axis = np.argmax(np.abs(row_cosine))\n    if row_cosine[dominant_row_axis] < 0:\n        img = np.fliplr(img)\n        flipped_lr = True\n    \n    dominant_col_axis = np.argmax(np.abs(col_cosine))\n    if col_cosine[dominant_col_axis] < 0:\n        img = np.flipud(img)\n        flipped_ud = True\n    \n    if verbose:\n        print(f\"Row cosine: {row_cosine}, dominant axis: {dominant_row_axis}, flipped L/R: {flipped_lr}\")\n        print(f\"Col cosine: {col_cosine}, dominant axis: {dominant_col_axis}, flipped U/D: {flipped_ud}\")\n    \n    return img\n\nstandardized_img = standardize_orientation(img, dcm.ImageOrientationPatient)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:01:13.122714Z","iopub.execute_input":"2026-08-15T13:01:13.123162Z","iopub.status.idle":"2026-08-15T13:01:13.130545Z","shell.execute_reply.started":"2026-08-15T13:01:13.123087Z","shell.execute_reply":"2026-08-15T13:01:13.129686Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test on an axial series from the same study\naxial_row = kept_series[(kept_series['study_id'] == sample_study) & \n                          (kept_series['series_description'] == 'Axial T2')].iloc[0]\n\naxial_folder = f'{base_path}/train_images/{sample_study}/{axial_row[\"series_id\"]}'\naxial_files = sorted(os.listdir(axial_folder), key=lambda x: int(x.replace('.dcm', '')))\naxial_mid_file = axial_files[len(axial_files) // 2]\n\naxial_dcm = pydicom.dcmread(f'{axial_folder}/{axial_mid_file}')\naxial_img = axial_dcm.pixel_array\n\nprint(\"Axial Image Orientation (Patient):\", axial_dcm.ImageOrientationPatient)\nplane = classify_plane(axial_dcm.ImageOrientationPatient)\nprint(\"Detected plane:\", plane)\n\naxial_standardized = standardize_orientation(axial_img, axial_dcm.ImageOrientationPatient)\n\nfig, axes = plt.subplots(1, 2, figsize=(12, 6))\naxes[0].imshow(axial_img, cmap='gray')\naxes[0].set_title('Original (Axial)')\naxes[0].axis('off')\naxes[1].imshow(axial_standardized, cmap='gray')\naxes[1].set_title('Standardized (Axial)')\naxes[1].axis('off')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:01:16.493907Z","iopub.execute_input":"2026-08-15T13:01:16.494292Z","iopub.status.idle":"2026-08-15T13:01:16.800799Z","shell.execute_reply.started":"2026-08-15T13:01:16.494265Z","shell.execute_reply":"2026-08-15T13:01:16.800195Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def normalize_intensity(img, lower_pct=1, upper_pct=99):\n    \"\"\"Percentile clipping followed by z-score normalization.\"\"\"\n    img = img.astype(np.float32)\n    \n    lower = np.percentile(img, lower_pct)\n    upper = np.percentile(img, upper_pct)\n    img_clipped = np.clip(img, lower, upper)\n    \n    mean = img_clipped.mean()\n    std = img_clipped.std()\n    img_normalized = (img_clipped - mean) / (std + 1e-8)\n    \n    return img_normalized, (lower, upper, mean, std)\n\n# Test on the sagittal sample we already have\nnorm_img, stats = normalize_intensity(standardized_img)\nprint(f\"Clip range: [{stats[0]:.1f}, {stats[1]:.1f}]\")\nprint(f\"Mean: {stats[2]:.1f}, Std: {stats[3]:.1f}\")\nprint(f\"Original range: [{standardized_img.min()}, {standardized_img.max()}]\")\nprint(f\"Normalized range: [{norm_img.min():.2f}, {norm_img.max():.2f}]\")\n\nfig, axes = plt.subplots(1, 3, figsize=(18, 6))\naxes[0].imshow(standardized_img, cmap='gray')\naxes[0].set_title('Before normalization')\naxes[0].axis('off')\n\naxes[1].hist(standardized_img.flatten(), bins=100)\naxes[1].set_title('Original intensity histogram')\n\naxes[2].imshow(norm_img, cmap='gray')\naxes[2].set_title('After normalization')\naxes[2].axis('off')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:01:20.314519Z","iopub.execute_input":"2026-08-15T13:01:20.315Z","iopub.status.idle":"2026-08-15T13:01:21.040825Z","shell.execute_reply.started":"2026-08-15T13:01:20.314973Z","shell.execute_reply":"2026-08-15T13:01:21.039975Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport torch\nimport pydicom\nimport cv2\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\n# --- Paths & core data ---\nbase_path = '/kaggle/input/competitions/rsna-2024-lumbar-spine-degenerative-classification'\ntrain = pd.read_csv(f'{base_path}/train.csv')\nseries_desc = pd.read_csv(f'{base_path}/train_series_descriptions.csv')\ncoords = pd.read_csv(f'{base_path}/train_label_coordinates.csv')\nlabel_cols = [c for c in train.columns if c != 'study_id']\n\n# --- Regenerate metadata (Step 4) ---\nmetadata_records = []\nfor study_id in tqdm(train['study_id'].unique()):\n    study_series = series_desc[series_desc['study_id'] == study_id]\n    for _, row in study_series.iterrows():\n        series_id = row['series_id']\n        series_folder = f'{base_path}/train_images/{study_id}/{series_id}'\n        try:\n            files = os.listdir(series_folder)\n            dcm = pydicom.dcmread(f'{series_folder}/{files[0]}')\n            metadata_records.append({\n                'study_id': study_id, 'series_id': series_id,\n                'series_description': row['series_description'],\n                'num_slices': len(files), 'rows': dcm.Rows, 'columns': dcm.Columns,\n                'pixel_spacing_x': float(dcm.PixelSpacing[0]), 'pixel_spacing_y': float(dcm.PixelSpacing[1]),\n                'slice_thickness': float(dcm.SliceThickness),\n                'spacing_between_slices': float(getattr(dcm, 'SpacingBetweenSlices', np.nan)),\n                'photometric_interpretation': dcm.PhotometricInterpretation,\n                'patient_position': dcm.PatientPosition,\n            })\n        except Exception as e:\n            metadata_records.append({'study_id': study_id, 'series_id': series_id,\n                                      'series_description': row['series_description'], 'error': str(e)})\n\nmetadata_df = pd.DataFrame(metadata_records)\n\n# --- Dedup (Step 5/8) ---\ndef select_series(group):\n    return group if len(group) == 1 else group.sort_values('num_slices', ascending=False).iloc[[0]]\n\nkept_series = metadata_df.groupby(['study_id', 'series_description'], group_keys=False).apply(select_series, include_groups=False)\nkept_series = metadata_df.loc[kept_series.index] if 'study_id' not in kept_series.columns else kept_series\n\nqc_status = []\nfor study_id in train['study_id'].unique():\n    views = kept_series[kept_series['study_id'] == study_id]['series_description'].tolist()\n    required = {'Sagittal T1', 'Sagittal T2/STIR', 'Axial T2'}\n    missing = required - set(views)\n    status = 'accepted' if not missing else ('requires_manual_review' if len(missing) == 1 else 'excluded')\n    qc_status.append({'study_id': study_id, 'qc_status': status, 'missing_views': list(missing)})\nqc_df = pd.DataFrame(qc_status)\n\n# --- Splits (Stage 3) ---\nfrom sklearn.model_selection import train_test_split\ntrain_qc = train.merge(qc_df, on='study_id')\naccepted = train_qc[train_qc['qc_status'] == 'accepted'].copy()\nseverity_rank = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\nseverity_numeric = accepted[label_cols].replace(severity_rank)\naccepted['worst_severity'] = severity_numeric.max(axis=1)\n\ntrain_ids, temp_ids = train_test_split(accepted['study_id'], test_size=0.30, stratify=accepted['worst_severity'], random_state=42)\nval_ids, test_ids = train_test_split(temp_ids, test_size=0.50,\n    stratify=accepted.loc[accepted['study_id'].isin(temp_ids), 'worst_severity'], random_state=42)\n\nsplits_df = pd.DataFrame({\n    'study_id': pd.concat([train_ids, val_ids, test_ids]),\n    'split': ['train']*len(train_ids) + ['val']*len(val_ids) + ['test']*len(test_ids)\n})\n\n# --- SAVE IMMEDIATELY so a disconnect doesn't lose this again ---\nmetadata_df.to_csv('/kaggle/working/series_metadata.csv', index=False)\nkept_series.to_csv('/kaggle/working/series_metadata_deduped.csv', index=False)\nqc_df.to_csv('/kaggle/working/qc_status.csv', index=False)\nsplits_df.to_csv('/kaggle/working/patient_splits.csv', index=False)\n\nprint(\"Rebuilt and saved. Train:\", len(train_ids), \"Val:\", len(val_ids), \"Test:\", len(test_ids))\n\n# --- Pipeline functions ---\ndef classify_plane(iop):\n    row_cosine, col_cosine = np.array(iop[:3]), np.array(iop[3:])\n    normal = np.cross(row_cosine, col_cosine)\n    return {0: 'sagittal', 1: 'coronal', 2: 'axial'}[np.argmax(np.abs(normal))]\n\ndef standardize_orientation(img, iop, verbose=False):\n    row_cosine, col_cosine = np.array(iop[:3]), np.array(iop[3:])\n    if row_cosine[np.argmax(np.abs(row_cosine))] < 0:\n        img = np.fliplr(img)\n    if col_cosine[np.argmax(np.abs(col_cosine))] < 0:\n        img = np.flipud(img)\n    return img\n\ndef normalize_intensity(img, lower_pct=1, upper_pct=99):\n    img = img.astype(np.float32)\n    lower, upper = np.percentile(img, lower_pct), np.percentile(img, upper_pct)\n    img_clipped = np.clip(img, lower, upper)\n    mean, std = img_clipped.mean(), img_clipped.std()\n    return (img_clipped - mean) / (std + 1e-8), (lower, upper, mean, std)\n\ndef resample_and_resize(img, current_spacing, target_spacing=0.5, target_size=384):\n    scale_factor = current_spacing / target_spacing\n    new_h, new_w = int(img.shape[0] * scale_factor), int(img.shape[1] * scale_factor)\n    resampled = cv2.resize(img, (new_w, new_h), interpolation=cv2.INTER_LINEAR)\n    h, w = resampled.shape\n    if h < target_size:\n        pt = (target_size - h) // 2\n        resampled = np.pad(resampled, ((pt, target_size - h - pt), (0, 0)), constant_values=resampled.min())\n    else:\n        ct = (h - target_size) // 2\n        resampled = resampled[ct:ct + target_size, :]\n    h, w = resampled.shape\n    if w < target_size:\n        pl = (target_size - w) // 2\n        resampled = np.pad(resampled, ((0, 0), (pl, target_size - w - pl)), constant_values=resampled.min())\n    else:\n        cl = (w - target_size) // 2\n        resampled = resampled[:, cl:cl + target_size]\n    return resampled\n\nprint(\"Pipeline functions ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:01:24.369858Z","iopub.execute_input":"2026-08-15T13:01:24.370446Z","iopub.status.idle":"2026-08-15T13:01:42.89685Z","shell.execute_reply.started":"2026-08-15T13:01:24.370418Z","shell.execute_reply":"2026-08-15T13:01:42.896053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load sagittal sample fresh\nsample_study = train_ids.iloc[0]\nsample_row = kept_series[(kept_series['study_id'] == sample_study) & \n                           (kept_series['series_description'] == 'Sagittal T2/STIR')].iloc[0]\nseries_folder = f'{base_path}/train_images/{sample_study}/{sample_row[\"series_id\"]}'\nfiles = sorted(os.listdir(series_folder), key=lambda x: int(x.replace('.dcm', '')))\nmid_file = files[len(files) // 2]\n\ndcm = pydicom.dcmread(f'{series_folder}/{mid_file}')\nimg = dcm.pixel_array\n\nstandardized_img = standardize_orientation(img, dcm.ImageOrientationPatient)\nnorm_img, stats = normalize_intensity(standardized_img)\ncurrent_spacing = float(dcm.PixelSpacing[0])\nresized_img = resample_and_resize(norm_img, current_spacing=current_spacing, target_spacing=0.5, target_size=384)\n\nprint(\"Original shape:\", norm_img.shape, \"| spacing:\", current_spacing)\nprint(\"Resampled shape:\", resized_img.shape)\n\nfig, axes = plt.subplots(1, 2, figsize=(12, 6))\naxes[0].imshow(norm_img, cmap='gray')\naxes[0].set_title(f'Before resize {norm_img.shape}')\naxes[0].axis('off')\naxes[1].imshow(resized_img, cmap='gray')\naxes[1].set_title(f'After resize {resized_img.shape} @ 0.5mm spacing')\naxes[1].axis('off')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:01:48.690953Z","iopub.execute_input":"2026-08-15T13:01:48.691571Z","iopub.status.idle":"2026-08-15T13:01:49.05542Z","shell.execute_reply.started":"2026-08-15T13:01:48.691542Z","shell.execute_reply":"2026-08-15T13:01:49.054708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\n\ndef augment_image(img, \n                   rotation_range=7,       # degrees\n                   translate_range=0.05,   # fraction of image size\n                   scale_range=0.05,       # +/- 5%\n                   brightness_range=0.1,\n                   contrast_range=0.1,\n                   gamma_range=0.1,\n                   noise_std=0.02):\n    h, w = img.shape\n    \n    # Rotation + scale + translation via affine matrix\n    angle = random.uniform(-rotation_range, rotation_range)\n    scale = 1.0 + random.uniform(-scale_range, scale_range)\n    tx = random.uniform(-translate_range, translate_range) * w\n    ty = random.uniform(-translate_range, translate_range) * h\n    \n    center = (w // 2, h // 2)\n    M = cv2.getRotationMatrix2D(center, angle, scale)\n    M[0, 2] += tx\n    M[1, 2] += ty\n    \n    img_aug = cv2.warpAffine(img, M, (w, h), borderMode=cv2.BORDER_CONSTANT, borderValue=float(img.min()))\n    \n    # Brightness / contrast\n    brightness = random.uniform(-brightness_range, brightness_range)\n    contrast = 1.0 + random.uniform(-contrast_range, contrast_range)\n    img_aug = img_aug * contrast + brightness\n    \n    # Gamma (on a shifted-positive copy since our images are z-scored, can go negative)\n    gamma = 1.0 + random.uniform(-gamma_range, gamma_range)\n    img_min = img_aug.min()\n    img_shifted = img_aug - img_min + 1e-3\n    img_aug = np.power(img_shifted / img_shifted.max(), gamma) * img_shifted.max() + img_min\n    \n    # Light Gaussian noise\n    noise = np.random.normal(0, noise_std, img_aug.shape).astype(np.float32)\n    img_aug = img_aug + noise\n    \n    return img_aug\n\n# Test: apply augmentation multiple times to the same image, compare\nfig, axes = plt.subplots(1, 4, figsize=(20, 5))\naxes[0].imshow(resized_img, cmap='gray')\naxes[0].set_title('Original (post Steps 12-15)')\naxes[0].axis('off')\n\nfor i in range(1, 4):\n    aug = augment_image(resized_img)\n    axes[i].imshow(aug, cmap='gray')\n    axes[i].set_title(f'Augmented v{i}')\n    axes[i].axis('off')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:03:09.593306Z","iopub.execute_input":"2026-08-15T13:03:09.593697Z","iopub.status.idle":"2026-08-15T13:03:10.270016Z","shell.execute_reply.started":"2026-08-15T13:03:09.593668Z","shell.execute_reply":"2026-08-15T13:03:10.269228Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\n\nclass UnlabeledMRIDataset(Dataset):\n    def __init__(self, study_ids, kept_series_df, base_path, img_size=224, slices_per_series=3):\n        self.samples = []\n        for study_id in study_ids:\n            series_rows = kept_series_df[kept_series_df['study_id'] == study_id]\n            for _, row in series_rows.iterrows():\n                series_folder = f'{base_path}/train_images/{study_id}/{row[\"series_id\"]}'\n                try:\n                    files = sorted(os.listdir(series_folder), key=lambda x: int(x.replace('.dcm', '')))\n                except FileNotFoundError:\n                    continue\n                if len(files) == 0:\n                    continue\n                # Sample a few evenly-spaced slices per series rather than every slice (keeps dataset size manageable)\n                idxs = np.linspace(0, len(files) - 1, min(slices_per_series, len(files)), dtype=int)\n                for idx in idxs:\n                    self.samples.append((study_id, row['series_id'], files[idx], series_folder))\n        self.img_size = img_size\n        print(f\"Built unlabeled dataset: {len(self.samples)} slices from {len(study_ids)} studies\")\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        study_id, series_id, filename, folder = self.samples[idx]\n        dcm = pydicom.dcmread(f'{folder}/{filename}')\n        img = dcm.pixel_array.astype(np.float32)\n        img = standardize_orientation(img, dcm.ImageOrientationPatient)\n        img, _ = normalize_intensity(img)\n        current_spacing = float(dcm.PixelSpacing[0])\n        img = resample_and_resize(img, current_spacing, target_spacing=0.5, target_size=self.img_size)\n        img = augment_image(img)  # light augmentation helps SSL pretraining too\n        return torch.from_numpy(img).unsqueeze(0).float()  # (1, H, W)\n\n# Build it on train split only\nunlabeled_ds = UnlabeledMRIDataset(train_ids.tolist(), kept_series, base_path, img_size=224, slices_per_series=3)\nunlabeled_loader = DataLoader(unlabeled_ds, batch_size=32, shuffle=True, num_workers=2, drop_last=True)\n\n# Sanity check: pull one batch\nbatch = next(iter(unlabeled_loader))\nprint(\"Batch shape:\", batch.shape, \"dtype:\", batch.dtype)\nprint(\"Value range:\", batch.min().item(), batch.max().item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:05:57.781311Z","iopub.execute_input":"2026-08-15T13:05:57.781921Z","iopub.status.idle":"2026-08-15T13:06:03.953281Z","shell.execute_reply.started":"2026-08-15T13:05:57.781889Z","shell.execute_reply":"2026-08-15T13:06:03.952427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\n\nclass PatchEmbed(nn.Module):\n    def __init__(self, img_size=224, patch_size=16, in_chans=1, embed_dim=192):\n        super().__init__()\n        self.num_patches = (img_size // patch_size) ** 2\n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n\n    def forward(self, x):\n        x = self.proj(x)  # (B, embed_dim, H/patch, W/patch)\n        x = x.flatten(2).transpose(1, 2)  # (B, num_patches, embed_dim)\n        return x\n\nclass SimpleMAE(nn.Module):\n    def __init__(self, img_size=224, patch_size=16, embed_dim=192, depth=6, decoder_dim=96, decoder_depth=2, mask_ratio=0.75):\n        super().__init__()\n        self.patch_size = patch_size\n        self.mask_ratio = mask_ratio\n        self.patch_embed = PatchEmbed(img_size, patch_size, 1, embed_dim)\n        num_patches = self.patch_embed.num_patches\n\n        self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, embed_dim))\n        encoder_layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=4, dim_feedforward=embed_dim*4, batch_first=True)\n        self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=depth)\n\n        self.decoder_embed = nn.Linear(embed_dim, decoder_dim)\n        self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim))\n        self.decoder_pos_embed = nn.Parameter(torch.zeros(1, num_patches, decoder_dim))\n        decoder_layer = nn.TransformerEncoderLayer(d_model=decoder_dim, nhead=4, dim_feedforward=decoder_dim*4, batch_first=True)\n        self.decoder = nn.TransformerEncoder(decoder_layer, num_layers=decoder_depth)\n        self.decoder_pred = nn.Linear(decoder_dim, patch_size * patch_size)\n\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n        nn.init.trunc_normal_(self.decoder_pos_embed, std=0.02)\n        nn.init.trunc_normal_(self.mask_token, std=0.02)\n\n    def random_masking(self, x):\n        B, N, D = x.shape\n        len_keep = int(N * (1 - self.mask_ratio))\n        noise = torch.rand(B, N, device=x.device)\n        ids_shuffle = torch.argsort(noise, dim=1)\n        ids_restore = torch.argsort(ids_shuffle, dim=1)\n        ids_keep = ids_shuffle[:, :len_keep]\n        x_masked = torch.gather(x, dim=1, index=ids_keep.unsqueeze(-1).repeat(1, 1, D))\n        mask = torch.ones(B, N, device=x.device)\n        mask[:, :len_keep] = 0\n        mask = torch.gather(mask, dim=1, index=ids_restore)\n        return x_masked, mask, ids_restore\n\n    def forward(self, imgs):\n        x = self.patch_embed(imgs) + self.pos_embed\n        x_masked, mask, ids_restore = self.random_masking(x)\n        latent = self.encoder(x_masked)\n\n        x = self.decoder_embed(latent)\n        B, len_keep, D = x.shape\n        N = ids_restore.shape[1]\n        mask_tokens = self.mask_token.repeat(B, N - len_keep, 1)\n        x_ = torch.cat([x, mask_tokens], dim=1)\n        x_ = torch.gather(x_, dim=1, index=ids_restore.unsqueeze(-1).repeat(1, 1, D))\n        x_ = x_ + self.decoder_pos_embed\n        x_ = self.decoder(x_)\n        pred = self.decoder_pred(x_)  # (B, N, patch_size^2)\n        return pred, mask\n\n    def patchify(self, imgs):\n        p = self.patch_size\n        h = w = imgs.shape[2] // p\n        x = imgs.reshape(imgs.shape[0], 1, h, p, w, p)\n        x = x.permute(0, 2, 4, 3, 5, 1).reshape(imgs.shape[0], h * w, p * p)\n        return x\n\n# Sanity check on one batch — forward pass + loss computation only, no training yet\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = SimpleMAE(img_size=224, patch_size=16, embed_dim=192, depth=6, mask_ratio=0.75).to(device)\n\nbatch = batch.to(device)\npred, mask = model(batch)\ntarget = model.patchify(batch)\n\nloss = ((pred - target) ** 2).mean(dim=-1)\nloss = (loss * mask).sum() / mask.sum()\n\nprint(\"Pred shape:\", pred.shape)\nprint(\"Target shape:\", target.shape)\nprint(\"Mask shape:\", mask.shape, \"| masked fraction:\", mask.mean().item())\nprint(\"Initial loss (untrained, should be a moderate positive number):\", loss.item())\nprint(\"Model parameters:\", sum(p.numel() for p in model.parameters()) / 1e6, \"M\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:07:03.624524Z","iopub.execute_input":"2026-08-15T13:07:03.624994Z","iopub.status.idle":"2026-08-15T13:07:05.123008Z","shell.execute_reply.started":"2026-08-15T13:07:03.624958Z","shell.execute_reply":"2026-08-15T13:07:05.122176Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport time\n\noptimizer = AdamW(model.parameters(), lr=1.5e-4, weight_decay=0.05)\nnum_epochs = 20\nscheduler = CosineAnnealingLR(optimizer, T_max=num_epochs)\n\ntrain_losses = []\nmodel.train()\n\nstart_time = time.time()\nfor epoch in range(num_epochs):\n    epoch_loss = 0.0\n    num_batches = 0\n    for batch in unlabeled_loader:\n        batch = batch.to(device)\n        pred, mask = model(batch)\n        target = model.patchify(batch)\n        loss = ((pred - target) ** 2).mean(dim=-1)\n        loss = (loss * mask).sum() / mask.sum()\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        epoch_loss += loss.item()\n        num_batches += 1\n\n    scheduler.step()\n    avg_loss = epoch_loss / num_batches\n    train_losses.append(avg_loss)\n    elapsed = time.time() - start_time\n    print(f\"Epoch {epoch+1}/{num_epochs} | Loss: {avg_loss:.4f} | LR: {scheduler.get_last_lr()[0]:.6f} | Elapsed: {elapsed/60:.1f}min\")\n\n# Save checkpoint immediately after training\ntorch.save(model.state_dict(), '/kaggle/working/mae_pretrained.pth')\nprint(\"\\nSaved MAE weights to /kaggle/working/mae_pretrained.pth\")\n\nplt.figure(figsize=(8,5))\nplt.plot(train_losses)\nplt.xlabel('Epoch')\nplt.ylabel('Reconstruction Loss')\nplt.title('MAE Pretraining Loss')\nplt.grid(True)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:07:54.558994Z","iopub.execute_input":"2026-08-15T13:07:54.559971Z","iopub.status.idle":"2026-08-15T13:41:56.828904Z","shell.execute_reply.started":"2026-08-15T13:07:54.559938Z","shell.execute_reply":"2026-08-15T13:41:56.828005Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nwith torch.no_grad():\n    sample_batch = next(iter(unlabeled_loader)).to(device)\n    pred, mask = model(sample_batch)\n\n    # Reconstruct full image: visible patches (original) + masked patches (predicted)\n    target_patches = model.patchify(sample_batch)\n    recon_patches = target_patches * (1 - mask.unsqueeze(-1)) + pred * mask.unsqueeze(-1)\n\n    def unpatchify(patches, img_size=224, patch_size=16):\n        p = patch_size\n        h = w = img_size // p\n        x = patches.reshape(patches.shape[0], h, w, p, p, 1)\n        x = x.permute(0, 5, 1, 3, 2, 4).reshape(patches.shape[0], 1, img_size, img_size)\n        return x\n\n    recon_img = unpatchify(recon_patches)\n    masked_vis = target_patches * (1 - mask.unsqueeze(-1))  # zero out masked patches to visualize what model saw\n    masked_img = unpatchify(masked_vis)\n\n    idx = 0\n    fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n    axes[0].imshow(sample_batch[idx, 0].cpu(), cmap='gray')\n    axes[0].set_title('Original')\n    axes[0].axis('off')\n    axes[1].imshow(masked_img[idx, 0].cpu(), cmap='gray')\n    axes[1].set_title('Masked (75% hidden)')\n    axes[1].axis('off')\n    axes[2].imshow(recon_img[idx, 0].cpu(), cmap='gray')\n    axes[2].set_title('Reconstructed')\n    axes[2].axis('off')\n    plt.tight_layout()\n    plt.show()\n\nmodel.train()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:43:30.931054Z","iopub.execute_input":"2026-08-15T13:43:30.93176Z","iopub.status.idle":"2026-08-15T13:43:32.561845Z","shell.execute_reply.started":"2026-08-15T13:43:30.931724Z","shell.execute_reply":"2026-08-15T13:43:32.560757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Filter coordinates to sagittal series only (per plan: sagittal T2/STIR primary, sagittal T1 optional)\ncoords_sag = coords.merge(kept_series[['study_id', 'series_id', 'series_description']], on=['study_id', 'series_id'])\ncoords_sag = coords_sag[coords_sag['series_description'].isin(['Sagittal T2/STIR', 'Sagittal T1'])]\n\nprint(\"Coordinate rows (sagittal only):\", len(coords_sag))\nprint(\"\\nUnique levels:\", sorted(coords_sag['level'].unique()))\nprint(\"\\nSample:\")\nprint(coords_sag.head(10))\n\n# For AOL-Net++, we want ONE landmark per level per study (disc-center), averaged across conditions/series at that level\n# Prefer Sagittal T2/STIR when available, since that's the plan's primary input\nlevel_order = ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']\n\ndef get_disc_centers(study_id, coords_df, preferred_series='Sagittal T2/STIR'):\n    study_coords = coords_df[coords_df['study_id'] == study_id]\n    centers = {}\n    for level in level_order:\n        level_coords = study_coords[study_coords['level'] == level]\n        preferred = level_coords[level_coords['series_description'] == preferred_series]\n        use = preferred if len(preferred) > 0 else level_coords\n        if len(use) > 0:\n            centers[level] = (use['x'].mean(), use['y'].mean(), use['series_id'].iloc[0])\n        else:\n            centers[level] = None\n    return centers\n\n# Test on our sample study\ntest_centers = get_disc_centers(sample_study, coords_sag)\nfor level, center in test_centers.items():\n    print(level, center)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:45:21.881305Z","iopub.execute_input":"2026-08-15T13:45:21.882252Z","iopub.status.idle":"2026-08-15T13:45:21.916582Z","shell.execute_reply.started":"2026-08-15T13:45:21.882211Z","shell.execute_reply":"2026-08-15T13:45:21.915761Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check completeness across all training-split studies\nmissing_level_studies = []\nfor study_id in train_ids:\n    centers = get_disc_centers(study_id, coords_sag)\n    missing = [lvl for lvl, c in centers.items() if c is None]\n    if missing:\n        missing_level_studies.append((study_id, missing))\n\nprint(f\"Studies with at least one missing level coordinate: {len(missing_level_studies)} / {len(train_ids)}\")\nif missing_level_studies:\n    print(\"\\nSample of missing cases:\")\n    for sid, missing in missing_level_studies[:10]:\n        print(sid, missing)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:45:59.120277Z","iopub.execute_input":"2026-08-15T13:45:59.1208Z","iopub.status.idle":"2026-08-15T13:46:04.062077Z","shell.execute_reply.started":"2026-08-15T13:45:59.120771Z","shell.execute_reply":"2026-08-15T13:46:04.061416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def generate_gaussian_heatmap(center_x, center_y, img_size, sigma=4):\n    \"\"\"Generate a 2D Gaussian heatmap centered at (center_x, center_y).\"\"\"\n    x = np.arange(0, img_size, 1, np.float32)\n    y = np.arange(0, img_size, 1, np.float32)[:, np.newaxis]\n    heatmap = np.exp(-((x - center_x)**2 + (y - center_y)**2) / (2 * sigma**2))\n    return heatmap\n\n# Test: build 5 heatmaps for our sample study, using the sagittal T2/STIR image we already have\nimg_size = 384  # matches our Step 15 target size\n\n# We need to transform original pixel coords -> our standardized/resized coordinate space\n# Since we resized from original DICOM shape to img_size, coords need the same scale factor\nsample_dcm_path = f'{base_path}/train_images/{sample_study}/{test_centers[\"L1/L2\"][2]}'\nsample_files = sorted(os.listdir(sample_dcm_path), key=lambda x: int(x.replace('.dcm', '')))\nref_dcm = pydicom.dcmread(f'{sample_dcm_path}/{sample_files[0]}')\norig_h, orig_w = ref_dcm.Rows, ref_dcm.Columns\norig_spacing = float(ref_dcm.PixelSpacing[0])\n\nscale_factor = orig_spacing / 0.5  # matches target_spacing=0.5 from resample_and_resize\n# account for center-crop/pad offset (approximation: assumes resize centers content)\nresized_h, resized_w = int(orig_h * scale_factor), int(orig_w * scale_factor)\noffset_y = (img_size - resized_h) / 2\noffset_x = (img_size - resized_w) / 2\n\nheatmaps = np.zeros((5, img_size, img_size), dtype=np.float32)\nfor i, level in enumerate(level_order):\n    cx, cy, _ = test_centers[level]\n    new_x = cx * scale_factor + offset_x\n    new_y = cy * scale_factor + offset_y\n    heatmaps[i] = generate_gaussian_heatmap(new_x, new_y, img_size, sigma=6)\n\n# Visualize: combine all 5 heatmaps overlaid on the resized image for a sanity check\ncombined_heatmap = heatmaps.max(axis=0)\n\nfig, axes = plt.subplots(1, 2, figsize=(12, 6))\naxes[0].imshow(resized_img, cmap='gray')\naxes[0].set_title('Resized image')\naxes[0].axis('off')\naxes[1].imshow(resized_img, cmap='gray')\naxes[1].imshow(combined_heatmap, cmap='hot', alpha=0.5)\naxes[1].set_title('Heatmap overlay (5 levels)')\naxes[1].axis('off')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:46:58.643499Z","iopub.execute_input":"2026-08-15T13:46:58.644356Z","iopub.status.idle":"2026-08-15T13:46:59.246145Z","shell.execute_reply.started":"2026-08-15T13:46:58.644321Z","shell.execute_reply":"2026-08-15T13:46:59.245061Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q timm\n\nimport timm\n\n# Load HRNet-W32 as a feature extractor (not classifier) - outputs multi-scale feature maps\nbackbone = timm.create_model('hrnet_w32', pretrained=True, features_only=True, in_chans=1)\nbackbone = backbone.to(device)\n\nprint(\"Feature info:\", backbone.feature_info.channels())\nprint(\"Reduction factors:\", backbone.feature_info.reduction())\n\n# Test forward pass with our resized image size\ntest_input = torch.randn(2, 1, 384, 384).to(device)\nwith torch.no_grad():\n    features = backbone(test_input)\n\nfor i, f in enumerate(features):\n    print(f\"Feature level {i}: shape {f.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:47:46.677914Z","iopub.execute_input":"2026-08-15T13:47:46.678539Z","iopub.status.idle":"2026-08-15T13:48:02.397633Z","shell.execute_reply.started":"2026-08-15T13:47:46.678504Z","shell.execute_reply":"2026-08-15T13:48:02.39687Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BiFPNLayer(nn.Module):\n    def __init__(self, num_channels=128, num_levels=5, epsilon=1e-4):\n        super().__init__()\n        self.epsilon = epsilon\n        self.num_levels = num_levels\n\n        # Learnable fusion weights (one set for top-down, one for bottom-up)\n        self.td_weights = nn.Parameter(torch.ones(num_levels - 1, 2))\n        self.bu_weights = nn.Parameter(torch.ones(num_levels - 1, 3))\n\n        self.td_convs = nn.ModuleList([\n            nn.Conv2d(num_channels, num_channels, 3, padding=1) for _ in range(num_levels - 1)\n        ])\n        self.bu_convs = nn.ModuleList([\n            nn.Conv2d(num_channels, num_channels, 3, padding=1) for _ in range(num_levels - 1)\n        ])\n\n    def forward(self, feats):\n        # feats: list of [B, C, H, W] from high-res (0) to low-res (N-1)\n        # Top-down pathway\n        td_feats = [None] * self.num_levels\n        td_feats[-1] = feats[-1]\n        for i in range(self.num_levels - 2, -1, -1):\n            w = torch.relu(self.td_weights[i])\n            w = w / (w.sum() + self.epsilon)\n            upsampled = torch.nn.functional.interpolate(td_feats[i + 1], size=feats[i].shape[-2:], mode='nearest')\n            fused = w[0] * feats[i] + w[1] * upsampled\n            td_feats[i] = self.td_convs[i](fused)\n\n        # Bottom-up pathway\n        bu_feats = [None] * self.num_levels\n        bu_feats[0] = td_feats[0]\n        for i in range(1, self.num_levels):\n            w = torch.relu(self.bu_weights[i - 1])\n            w = w / (w.sum() + self.epsilon)\n            downsampled = torch.nn.functional.max_pool2d(bu_feats[i - 1], kernel_size=2)\n            if downsampled.shape[-2:] != feats[i].shape[-2:]:\n                downsampled = torch.nn.functional.interpolate(downsampled, size=feats[i].shape[-2:], mode='nearest')\n            fused = w[0] * feats[i] + w[1] * td_feats[i] + w[2] * downsampled\n            bu_feats[i] = self.bu_convs[i - 1](fused)\n\n        return bu_feats\n\nclass MultiScaleExtractor(nn.Module):\n    def __init__(self, backbone, in_channels_list, bifpn_channels=128, num_bifpn_layers=2):\n        super().__init__()\n        self.backbone = backbone\n        # 1x1 convs to project each backbone level to a common channel dim\n        self.lateral_convs = nn.ModuleList([\n            nn.Conv2d(c, bifpn_channels, 1) for c in in_channels_list\n        ])\n        self.bifpn_layers = nn.ModuleList([\n            BiFPNLayer(bifpn_channels, len(in_channels_list)) for _ in range(num_bifpn_layers)\n        ])\n\n    def forward(self, x):\n        feats = self.backbone(x)\n        feats = [lat(f) for lat, f in zip(self.lateral_convs, feats)]\n        for layer in self.bifpn_layers:\n            feats = layer(feats)\n        return feats\n\n# Build and test\nextractor = MultiScaleExtractor(backbone, in_channels_list=[64, 128, 256, 512, 1024], bifpn_channels=128, num_bifpn_layers=2).to(device)\n\nwith torch.no_grad():\n    fused_feats = extractor(test_input)\n\nfor i, f in enumerate(fused_feats):\n    print(f\"Fused level {i}: shape {f.shape}\")\n\nprint(\"\\nTotal parameters:\", sum(p.numel() for p in extractor.parameters()) / 1e6, \"M\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:48:49.466054Z","iopub.execute_input":"2026-08-15T13:48:49.466638Z","iopub.status.idle":"2026-08-15T13:48:49.596512Z","shell.execute_reply.started":"2026-08-15T13:48:49.466606Z","shell.execute_reply":"2026-08-15T13:48:49.595903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AnatomicalContextModule(nn.Module):\n    def __init__(self, channels=128, num_heads=4, depth=2):\n        super().__init__()\n        self.channels = channels\n        encoder_layer = nn.TransformerEncoderLayer(\n            d_model=channels, nhead=num_heads, dim_feedforward=channels * 4,\n            batch_first=True, dropout=0.1\n        )\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=depth)\n        # Learnable 2D positional encoding (added once, sized for the working resolution)\n        self.pos_embed = None  # built lazily based on input size\n\n    def build_pos_embed(self, h, w, device):\n        pos = torch.zeros(1, h * w, self.channels, device=device)\n        nn.init.trunc_normal_(pos, std=0.02)\n        return nn.Parameter(pos)\n\n    def forward(self, feat):\n        B, C, H, W = feat.shape\n        if self.pos_embed is None or self.pos_embed.shape[1] != H * W:\n            self.pos_embed = self.build_pos_embed(H, W, feat.device)\n\n        x = feat.flatten(2).transpose(1, 2)  # (B, H*W, C)\n        x = x + self.pos_embed\n        x = self.transformer(x)\n        x = x.transpose(1, 2).reshape(B, C, H, W)\n        return x\n\n# Test on the highest-res fused feature (level 0)\ncontext_module = AnatomicalContextModule(channels=128, num_heads=4, depth=2).to(device)\n\nwith torch.no_grad():\n    context_out = context_module(fused_feats[0])\n\nprint(\"Context module output shape:\", context_out.shape)\nprint(\"Parameters:\", sum(p.numel() for p in context_module.parameters()) / 1e6, \"M\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:49:12.333159Z","iopub.execute_input":"2026-08-15T13:49:12.333577Z","iopub.status.idle":"2026-08-15T13:49:12.728945Z","shell.execute_reply.started":"2026-08-15T13:49:12.333549Z","shell.execute_reply":"2026-08-15T13:49:12.727942Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DeformableAttention2D(nn.Module):\n    \"\"\"Single-scale deformable attention: each query samples K learned offset points from the feature map.\"\"\"\n    def __init__(self, channels=128, num_heads=4, num_points=4):\n        super().__init__()\n        self.channels = channels\n        self.num_heads = num_heads\n        self.num_points = num_points\n        self.head_dim = channels // num_heads\n        assert channels % num_heads == 0, \"channels must be divisible by num_heads\"\n\n        # Predict per-query, per-head, per-point (x,y) offsets and attention weights\n        self.offset_proj = nn.Conv2d(channels, num_heads * num_points * 2, kernel_size=3, padding=1)\n        self.attn_weight_proj = nn.Conv2d(channels, num_heads * num_points, kernel_size=3, padding=1)\n        self.value_proj = nn.Conv2d(channels, channels, kernel_size=1)\n        self.output_proj = nn.Conv2d(channels, channels, kernel_size=1)\n\n        nn.init.constant_(self.offset_proj.weight, 0)\n        nn.init.constant_(self.offset_proj.bias, 0)\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        nh, np_, hd = self.num_heads, self.num_points, self.head_dim\n\n        value = self.value_proj(x)  # (B, C, H, W)\n        value = value.view(B, nh, hd, H, W)\n\n        offsets = self.offset_proj(x)  # (B, nh*np*2, H, W)\n        offsets = offsets.view(B, nh, np_, 2, H, W)\n\n        attn_weights = self.attn_weight_proj(x)  # (B, nh*np, H, W)\n        attn_weights = attn_weights.view(B, nh, np_, H, W)\n        attn_weights = torch.softmax(attn_weights, dim=2)  # normalize over the K sample points\n\n        # Build base sampling grid (query's own location), normalized to [-1, 1] for grid_sample\n        ys, xs = torch.meshgrid(\n            torch.linspace(-1, 1, H, device=x.device),\n            torch.linspace(-1, 1, W, device=x.device),\n            indexing='ij'\n        )\n        base_grid = torch.stack([xs, ys], dim=-1)  # (H, W, 2)\n\n        out = torch.zeros(B, nh, hd, H, W, device=x.device)\n        for head in range(nh):\n            for pt in range(np_):\n                # Offset in normalized coords (scale down raw offsets so they start small/local)\n                off_x = offsets[:, head, pt, 0] / (W / 2)\n                off_y = offsets[:, head, pt, 1] / (H / 2)\n                sample_grid = base_grid.unsqueeze(0) + torch.stack([off_x, off_y], dim=-1)  # (B, H, W, 2)\n                sample_grid = sample_grid.clamp(-1, 1)\n\n                sampled = torch.nn.functional.grid_sample(\n                    value[:, head], sample_grid, mode='bilinear', padding_mode='border', align_corners=True\n                )  # (B, hd, H, W)\n\n                w = attn_weights[:, head, pt].unsqueeze(1)  # (B, 1, H, W)\n                out[:, head] += sampled * w\n\n        out = out.view(B, C, H, W)\n        out = self.output_proj(out)\n        return out + x  # residual connection\n\n# Test it on our context module output\ndeform_attn = DeformableAttention2D(channels=128, num_heads=4, num_points=4).to(device)\n\nwith torch.no_grad():\n    deform_out = deform_attn(context_out)\n\nprint(\"Deformable attention output shape:\", deform_out.shape)\nprint(\"Parameters:\", sum(p.numel() for p in deform_attn.parameters()) / 1e6, \"M\")\nprint(\"Output matches input shape:\", deform_out.shape == context_out.shape)\nprint(\"Output stats - mean:\", deform_out.mean().item(), \"std:\", deform_out.std().item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:50:07.422865Z","iopub.execute_input":"2026-08-15T13:50:07.42336Z","iopub.status.idle":"2026-08-15T13:50:07.531508Z","shell.execute_reply.started":"2026-08-15T13:50:07.42333Z","shell.execute_reply":"2026-08-15T13:50:07.530764Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class HeatmapHead(nn.Module):\n    def __init__(self, in_channels=128, num_levels=5):\n        super().__init__()\n        self.num_levels = num_levels\n        self.heads = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(in_channels, 64, 3, padding=1),\n                nn.ReLU(inplace=True),\n                nn.Conv2d(64, 1, 1)\n            ) for _ in range(num_levels)\n        ])\n\n    def forward(self, feat):\n        outputs = [head(feat) for head in self.heads]\n        return torch.cat(outputs, dim=1)  # (B, num_levels, H, W)\n\nheatmap_head = HeatmapHead(in_channels=128, num_levels=5).to(device)\n\nwith torch.no_grad():\n    pred_heatmaps = heatmap_head(deform_out)\n\nprint(\"Predicted heatmaps shape:\", pred_heatmaps.shape)\nprint(\"Parameters:\", sum(p.numel() for p in heatmap_head.parameters()) / 1e6, \"M\")\n\n# Sanity check: apply sigmoid to see what an untrained prediction looks like\npred_sigmoid = torch.sigmoid(pred_heatmaps)\nprint(\"Untrained heatmap range after sigmoid:\", pred_sigmoid.min().item(), pred_sigmoid.max().item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:50:31.649423Z","iopub.execute_input":"2026-08-15T13:50:31.649831Z","iopub.status.idle":"2026-08-15T13:50:31.769812Z","shell.execute_reply.started":"2026-08-15T13:50:31.649802Z","shell.execute_reply":"2026-08-15T13:50:31.769138Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def soft_argmax_2d(heatmaps, beta=100):\n    \"\"\"\n    heatmaps: (B, num_levels, H, W)\n    Returns: (B, num_levels, 2) coordinates in [0, W-1] x [0, H-1] pixel space\n    \"\"\"\n    B, N, H, W = heatmaps.shape\n    heatmaps_flat = heatmaps.view(B, N, -1)\n    probs = torch.softmax(heatmaps_flat * beta, dim=-1)\n    probs = probs.view(B, N, H, W)\n\n    ys = torch.linspace(0, H - 1, H, device=heatmaps.device)\n    xs = torch.linspace(0, W - 1, W, device=heatmaps.device)\n    grid_y, grid_x = torch.meshgrid(ys, xs, indexing='ij')\n\n    coord_x = (probs * grid_x.unsqueeze(0).unsqueeze(0)).sum(dim=(-2, -1))\n    coord_y = (probs * grid_y.unsqueeze(0).unsqueeze(0)).sum(dim=(-2, -1))\n\n    return torch.stack([coord_x, coord_y], dim=-1)  # (B, N, 2)\n\n# Test on our untrained heatmaps\nwith torch.no_grad():\n    pred_coords = soft_argmax_2d(pred_heatmaps, beta=100)\n\nprint(\"Predicted coords shape:\", pred_coords.shape)\nprint(\"Sample coords (untrained, will look random):\")\nprint(pred_coords[0])\n\n# Verify differentiability: run with grad enabled and check backward works\npred_heatmaps_grad = heatmap_head(deform_out.detach().requires_grad_(False))\npred_heatmaps_grad.requires_grad_(True)\ntest_coords = soft_argmax_2d(pred_heatmaps_grad, beta=100)\nloss = test_coords.sum()\nloss.backward()\nprint(\"\\nGradient flows correctly:\", pred_heatmaps_grad.grad is not None)\nprint(\"Gradient stats - mean:\", pred_heatmaps_grad.grad.mean().item(), \"std:\", pred_heatmaps_grad.grad.std().item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:51:00.111472Z","iopub.execute_input":"2026-08-15T13:51:00.112189Z","iopub.status.idle":"2026-08-15T13:51:00.334466Z","shell.execute_reply.started":"2026-08-15T13:51:00.112144Z","shell.execute_reply":"2026-08-15T13:51:00.333402Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Correct gradient test: use a genuine leaf tensor as input\ntest_heatmaps = torch.randn(2, 5, 192, 192, device=device, requires_grad=True)\ntest_coords = soft_argmax_2d(test_heatmaps, beta=100)\nloss = test_coords.sum()\nloss.backward()\n\nprint(\"Gradient flows correctly:\", test_heatmaps.grad is not None)\nprint(\"Gradient stats - mean:\", test_heatmaps.grad.mean().item(), \"std:\", test_heatmaps.grad.std().item())\n\n# Also verify gradient flows all the way through the real model (heatmap_head -> soft_argmax)\ndeform_out_leaf = deform_out.detach().clone().requires_grad_(True)\npred_heatmaps_real = heatmap_head(deform_out_leaf)\npred_coords_real = soft_argmax_2d(pred_heatmaps_real, beta=100)\nloss_real = pred_coords_real.sum()\nloss_real.backward()\n\nprint(\"\\nEnd-to-end gradient (through heatmap_head) flows:\", deform_out_leaf.grad is not None)\nprint(\"End-to-end gradient stats - mean:\", deform_out_leaf.grad.mean().item(), \"std:\", deform_out_leaf.grad.std().item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:51:29.665203Z","iopub.execute_input":"2026-08-15T13:51:29.665807Z","iopub.status.idle":"2026-08-15T13:51:29.743248Z","shell.execute_reply.started":"2026-08-15T13:51:29.665777Z","shell.execute_reply":"2026-08-15T13:51:29.742411Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def anatomical_ordering_loss(pred_coords, margin=2.0):\n    \"\"\"\n    pred_coords: (B, 5, 2) in order [L1/L2, L2/L3, L3/L4, L4/L5, L5/S1]\n    Penalize violations where y[i] is not sufficiently less than y[i+1]\n    (i.e., each level should be below the previous one in image space)\n    \"\"\"\n    y_coords = pred_coords[:, :, 1]  # (B, 5)\n    diffs = y_coords[:, 1:] - y_coords[:, :-1]  # should be positive (increasing y = moving down)\n    violation = torch.clamp(margin - diffs, min=0)  # penalize when diff < margin\n    return violation.mean()\n\n# Test on our current (untrained, likely disordered) predictions\nwith torch.no_grad():\n    order_loss = anatomical_ordering_loss(pred_coords)\n    print(\"Ordering loss on untrained predictions:\", order_loss.item())\n    print(\"\\nY-coordinates per level (should increase top to bottom if ordered correctly):\")\n    print(pred_coords[0, :, 1])\n\n# Test on a \"perfect\" synthetic example to confirm the loss goes to ~0 when correctly ordered\nperfect_coords = torch.tensor([[[100, 20], [100, 60], [100, 100], [100, 140], [100, 180]]], dtype=torch.float32, device=device)\nperfect_loss = anatomical_ordering_loss(perfect_coords)\nprint(\"\\nOrdering loss on a perfectly-ordered synthetic example (should be ~0):\", perfect_loss.item())\n\n# Test on a deliberately reversed example to confirm the loss is large\nreversed_coords = torch.tensor([[[100, 180], [100, 140], [100, 100], [100, 60], [100, 20]]], dtype=torch.float32, device=device)\nreversed_loss = anatomical_ordering_loss(reversed_coords)\nprint(\"Ordering loss on a reversed synthetic example (should be large):\", reversed_loss.item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:52:13.562976Z","iopub.execute_input":"2026-08-15T13:52:13.563655Z","iopub.status.idle":"2026-08-15T13:52:13.574685Z","shell.execute_reply.started":"2026-08-15T13:52:13.563624Z","shell.execute_reply":"2026-08-15T13:52:13.574012Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AuxiliaryHeads(nn.Module):\n    def __init__(self, in_channels=128, num_levels=5):\n        super().__init__()\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.shared = nn.Sequential(\n            nn.Linear(in_channels, 64),\n            nn.ReLU(inplace=True)\n        )\n        self.visibility_head = nn.Linear(64, num_levels)       # binary per level\n        self.confidence_head = nn.Linear(64, num_levels)       # scalar per level\n        self.roi_size_head = nn.Linear(64, num_levels)         # scalar per level (proxy target)\n        self.axial_group_head = nn.Linear(64, num_levels)      # regressed slice-index proxy\n\n    def forward(self, feat):\n        pooled = self.pool(feat).flatten(1)  # (B, C)\n        shared = self.shared(pooled)\n        return {\n            'visibility': torch.sigmoid(self.visibility_head(shared)),\n            'confidence': torch.sigmoid(self.confidence_head(shared)),\n            'roi_size': self.roi_size_head(shared),\n            'axial_group': self.axial_group_head(shared),\n        }\n\naux_heads = AuxiliaryHeads(in_channels=128, num_levels=5).to(device)\n\nwith torch.no_grad():\n    aux_out = aux_heads(deform_out)\n\nfor k, v in aux_out.items():\n    print(f\"{k}: shape {v.shape}\")\n\n# --- Derive proxy targets for ROI size (spacing between adjacent landmarks) ---\ndef derive_roi_size_targets(coords_per_level):\n    # coords_per_level: (5, 2) array of (x,y) for L1/L2...L5/S1\n    sizes = np.zeros(5)\n    for i in range(5):\n        neighbors = []\n        if i > 0: neighbors.append(np.linalg.norm(coords_per_level[i] - coords_per_level[i-1]))\n        if i < 4: neighbors.append(np.linalg.norm(coords_per_level[i] - coords_per_level[i+1]))\n        sizes[i] = np.mean(neighbors)\n    return sizes\n\ntest_coords_array = np.array([[test_centers[lvl][0], test_centers[lvl][1]] for lvl in level_order])\nroi_targets = derive_roi_size_targets(test_coords_array)\nprint(\"\\nDerived ROI size proxy targets (pixel distance to neighbor levels):\", roi_targets)\n\nprint(\"\\nTotal auxiliary head parameters:\", sum(p.numel() for p in aux_heads.parameters()) / 1e6, \"M\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-15T13:53:03.567811Z","iopub.execute_input":"2026-08-15T13:53:03.568292Z","iopub.status.idle":"2026-08-15T13:53:03.605396Z","shell.execute_reply.started":"2026-08-15T13:53:03.568262Z","shell.execute_reply":"2026-08-15T13:53:03.604608Z"}},"outputs":[],"execution_count":null}]}