{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":52254,"databundleVersionId":9674523,"sourceType":"competition"},{"sourceId":12242890,"sourceType":"datasetVersion","datasetId":7713895},{"sourceId":12245569,"sourceType":"datasetVersion","datasetId":7715459},{"sourceId":12247149,"sourceType":"datasetVersion","datasetId":7715595},{"sourceId":12247192,"sourceType":"datasetVersion","datasetId":7715723},{"sourceId":12247822,"sourceType":"datasetVersion","datasetId":7716931},{"sourceId":12248026,"sourceType":"datasetVersion","datasetId":7717019},{"sourceId":12248216,"sourceType":"datasetVersion","datasetId":7717073},{"sourceId":12248344,"sourceType":"datasetVersion","datasetId":7716417},{"sourceId":12248429,"sourceType":"datasetVersion","datasetId":7717311},{"sourceId":12248555,"sourceType":"datasetVersion","datasetId":7717450},{"sourceId":12248570,"sourceType":"datasetVersion","datasetId":7717422},{"sourceId":12248642,"sourceType":"datasetVersion","datasetId":7717550},{"sourceId":12248646,"sourceType":"datasetVersion","datasetId":7717525},{"sourceId":12248817,"sourceType":"datasetVersion","datasetId":7717528},{"sourceId":12248822,"sourceType":"datasetVersion","datasetId":7717451},{"sourceId":12248899,"sourceType":"datasetVersion","datasetId":7717562},{"sourceId":12248924,"sourceType":"datasetVersion","datasetId":7717559},{"sourceId":12248941,"sourceType":"datasetVersion","datasetId":7717513},{"sourceId":12248986,"sourceType":"datasetVersion","datasetId":7717737},{"sourceId":12249017,"sourceType":"datasetVersion","datasetId":7717764},{"sourceId":12249031,"sourceType":"datasetVersion","datasetId":7717797},{"sourceId":12249087,"sourceType":"datasetVersion","datasetId":7717809},{"sourceId":12249259,"sourceType":"datasetVersion","datasetId":7717907},{"sourceId":12249263,"sourceType":"datasetVersion","datasetId":7717914},{"sourceId":12249332,"sourceType":"datasetVersion","datasetId":7717947},{"sourceId":12249364,"sourceType":"datasetVersion","datasetId":7717984},{"sourceId":12249398,"sourceType":"datasetVersion","datasetId":7717979},{"sourceId":12249401,"sourceType":"datasetVersion","datasetId":7717995},{"sourceId":12249507,"sourceType":"datasetVersion","datasetId":7718311},{"sourceId":12249579,"sourceType":"datasetVersion","datasetId":7718001},{"sourceId":12249711,"sourceType":"datasetVersion","datasetId":7717756},{"sourceId":12249730,"sourceType":"datasetVersion","datasetId":7718247},{"sourceId":12249734,"sourceType":"datasetVersion","datasetId":7718243},{"sourceId":12250069,"sourceType":"datasetVersion","datasetId":7717555},{"sourceId":12250080,"sourceType":"datasetVersion","datasetId":7717565},{"sourceId":12250358,"sourceType":"datasetVersion","datasetId":7718213},{"sourceId":12250425,"sourceType":"datasetVersion","datasetId":7718208},{"sourceId":12250577,"sourceType":"datasetVersion","datasetId":7718838},{"sourceId":12250614,"sourceType":"datasetVersion","datasetId":7718879},{"sourceId":12251815,"sourceType":"datasetVersion","datasetId":7719516},{"sourceId":12252028,"sourceType":"datasetVersion","datasetId":7717569},{"sourceId":12275973,"sourceType":"datasetVersion","datasetId":7718258},{"sourceId":12276921,"sourceType":"datasetVersion","datasetId":7736108},{"sourceId":12276982,"sourceType":"datasetVersion","datasetId":7736025},{"sourceId":12277537,"sourceType":"datasetVersion","datasetId":7736720},{"sourceId":12277584,"sourceType":"datasetVersion","datasetId":7736693},{"sourceId":12277591,"sourceType":"datasetVersion","datasetId":7736109},{"sourceId":12277608,"sourceType":"datasetVersion","datasetId":7736643},{"sourceId":12277881,"sourceType":"datasetVersion","datasetId":7737297},{"sourceId":12278371,"sourceType":"datasetVersion","datasetId":7736402},{"sourceId":12324434,"sourceType":"datasetVersion","datasetId":7762794},{"sourceId":12325463,"sourceType":"datasetVersion","datasetId":7769216},{"sourceId":12335941,"sourceType":"datasetVersion","datasetId":7772212},{"sourceId":12338543,"sourceType":"datasetVersion","datasetId":7775974}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install kaggle --quiet","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:30:01.522194Z","iopub.execute_input":"2025-07-01T09:30:01.522818Z","iopub.status.idle":"2025-07-01T09:30:06.61872Z","shell.execute_reply.started":"2025-07-01T09:30:01.522793Z","shell.execute_reply":"2025-07-01T09:30:06.617885Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Install necessary libraries that may not be available by default in Kaggle\n\n# pydicom: For reading DICOM files\n# nibabel: For working with NIfTI files (reading, writing)\n# monai: Medical imaging AI framework (includes augmentations for 3D data)\n# SimpleITK: Useful for medical image processing including NIfTI and DICOM support\n# torchio: Alternative 3D medical image processing and augmentation toolkit\n\n# Only install if not already available\ntry:\n    import pydicom, nibabel, monai, SimpleITK, torchio\nexcept ImportError:\n    !pip install -q pydicom nibabel monai SimpleITK torchio\n\n# Import standard libraries\nimport os\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom glob import glob\nimport pandas as pd\nimport json\nimport random\nfrom pathlib import Path\n\n# Import DICOM and NIfTI handling libraries\nimport pydicom\nimport nibabel as nib\nimport SimpleITK as sitk\n\n# Import PyTorch and MONAI for model training and augmentations\nimport torch\nimport monai\n\n# For progress bars\nfrom tqdm.notebook import tqdm\n\n# For warnings suppression (optional)\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Ensure Reproducibility\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n\nset_seed()\n\n# Automatically use GPU if available, else CPU\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)\nprint(\"Environment setup complete.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:30:06.620333Z","iopub.execute_input":"2025-07-01T09:30:06.620844Z","iopub.status.idle":"2025-07-01T09:32:05.25001Z","shell.execute_reply.started":"2025-07-01T09:30:06.62082Z","shell.execute_reply":"2025-07-01T09:32:05.249342Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport shutil\n\n# Make sure the .config/kaggle directory exists\nos.makedirs(\"/root/.config/kaggle\", exist_ok=True)\n\n# Move kaggle.json to expected directory\nshutil.copy(\"/kaggle/input/kaggle-json/kaggle.json\", \"/root/.config/kaggle/kaggle.json\")\n\n# Set permissions (optional but recommended)\nos.chmod(\"/root/.config/kaggle/kaggle.json\", 0o600)\n\n# Now import and authenticate\nfrom kaggle.api.kaggle_api_extended import KaggleApi\n\napi = KaggleApi()\napi.authenticate()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:05.250783Z","iopub.execute_input":"2025-07-01T09:32:05.251453Z","iopub.status.idle":"2025-07-01T09:32:05.619484Z","shell.execute_reply.started":"2025-07-01T09:32:05.251425Z","shell.execute_reply":"2025-07-01T09:32:05.618699Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model_from_dataset_path = \"/kaggle/input/densenet121-model/model_best.pth\"\n\n# if os.path.exists(model_from_dataset_path):\n#     print(f\"✅ Found model in dataset: {model_from_dataset_path}\")\n#     model.load_state_dict(torch.load(model_from_dataset_path, map_location=device))\n# else:\n#     print(\"⚠️ No model found in Kaggle dataset input.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:05.621165Z","iopub.execute_input":"2025-07-01T09:32:05.621401Z","iopub.status.idle":"2025-07-01T09:32:05.624962Z","shell.execute_reply.started":"2025-07-01T09:32:05.621383Z","shell.execute_reply":"2025-07-01T09:32:05.624235Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Paths\nDATASET_ROOTS = [\n    '/kaggle/input/abdominal-nifti-0-100',\n    '/kaggle/input/abdominal-nifti-100-200',\n    '/kaggle/input/abdominal-trauma-nifti-200-300',\n    '/kaggle/input/abdominal-trauma-nifti-300-400',\n    '/kaggle/input/abdominal-trauma-nifti-400-500',\n    '/kaggle/input/abdominal-trauma-nifti-500-600',\n    '/kaggle/input/abdominal-trauma-nifti-600-700',\n    '/kaggle/input/abdominal-trauma-nifti-700-800',\n    '/kaggle/input/abdominal-trauma-nifti-800-900',\n    '/kaggle/input/abdominal-trauma-nifti-900-1000',\n    \n    '/kaggle/input/abdominal-nifti-1000-1100',\n    '/kaggle/input/abdominal-nifti-1100-1200',\n    '/kaggle/input/abdominal-nifti-1200-1300',\n    '/kaggle/input/abdominal-nifti-1300-1400',\n    '/kaggle/input/abdominal-nifti-1400-1500',\n    '/kaggle/input/abdominal-nifti-1500-1600',\n    '/kaggle/input/abdominal-nifti-1600-1700',\n    '/kaggle/input/abdominal-nifti-1700-1800',\n    '/kaggle/input/abdominal-nifti-1800-1900',\n    '/kaggle/input/abdominal-nifti-1900-2000',\n    \n    '/kaggle/input/abdominal-trauma-nifti-2000-above',\n    '/kaggle/input/abdominal-nifti-2100-2150',\n    '/kaggle/input/abdominal-trauma-nifti-2230-2360',\n    '/kaggle/input/abdominal-trauma-nifti-2300-2400',\n    '/kaggle/input/abdominal-trauma-nifti-2400-2500',\n    '/kaggle/input/abdominal-trauma-nifti-2500-2600',\n    '/kaggle/input/abdominal-trauma-nifti-2600-2700',\n    '/kaggle/input/abdominal-trauma-nifti-2700-2800',\n    '/kaggle/input/abdominal-trauma-nifti-2800-2900',\n    '/kaggle/input/abdominal-trauma-nifti-2900-3000',\n\n    '/kaggle/input/abdominal-trauma-nifti-3000-3100',\n    '/kaggle/input/abdominal-trauma-nifti-3100-3200',\n    '/kaggle/input/abdominal-trauma-nifti-3200-3300',\n    '/kaggle/input/abdominal-trauma-nifti-3300-3400',\n    '/kaggle/input/abdominal-trauma-nifti-3400-3500',\n    '/kaggle/input/abdominal-trauma-nifti-3500-3600',\n    '/kaggle/input/abdominal-trauma-nifti-3600-3700',\n    '/kaggle/input/abdominal-trauma-nifti-3700-3800',\n    '/kaggle/input/abdominal-trauma-nifti-3800-3900',\n    '/kaggle/input/abdominal-trauma-nifti-3900-4000',\n\n    '/kaggle/input/abdominal-trauma-nifti-4000-4100',\n    '/kaggle/input/abdominal-trauma-nifti-4100-4200',\n    '/kaggle/input/abdominal-trauma-nifti-4200-4300',\n    '/kaggle/input/abdominal-nifti-4290-4400',\n    '/kaggle/input/abdominal-nifti-4380-4470',\n    '/kaggle/input/abdominal-nifti-4470-4560',\n    '/kaggle/input/abdominal-nifti-4560-4650',\n    '/kaggle/input/abdominal-nifti-4650-4710',\n    \n]\n\nLABELS_CSV_PATH = '/kaggle/input/rsna-2023-abdominal-trauma-detection/train_2024.csv'\nOUTPUT_JSON_PATH = '/kaggle/working/train_metadata.json'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:05.625998Z","iopub.execute_input":"2025-07-01T09:32:05.626283Z","iopub.status.idle":"2025-07-01T09:32:05.650161Z","shell.execute_reply.started":"2025-07-01T09:32:05.626256Z","shell.execute_reply":"2025-07-01T09:32:05.649588Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load labels CSV\nlabels_df = pd.read_csv(LABELS_CSV_PATH)\nlabels_df['patient_id'] = labels_df['patient_id'].astype(str)\nlabels_dict_map = labels_df.set_index('patient_id').to_dict(orient='index')\nlabel_cols = [col for col in labels_df.columns if col != 'patient_id']\n\nmetadata_list = []\n\nfor dataset_root in DATASET_ROOTS:\n    nifti_files = sorted(Path(dataset_root).rglob(\"*.nii*\"))  # .nii or .nii.gz both\n\n    print(f\"🔍 Found {len(nifti_files)} NIfTI files in {dataset_root}\")\n\n    for nii_path in nifti_files:\n        stem = nii_path.stem  # e.g. \"12345_67890\"\n        try:\n            patient_id, study_id = stem.split(\"_\")\n        except ValueError:\n            print(f\"⚠️ Skipping malformed filename: {stem}\")\n            continue\n\n        if patient_id not in labels_dict_map:\n            print(f\"⚠️ No label for patient {patient_id}, skipping...\")\n            continue\n\n        labels = {col: int(labels_dict_map[patient_id][col]) for col in label_cols}\n\n        metadata_list.append({\n            \"patient_id\": patient_id,\n            \"study_id\": study_id,\n            \"nifti_path\": str(nii_path),\n            \"labels\": labels\n        })\n\nprint(f\"✅ Total metadata entries: {len(metadata_list)}\")\n\n# Save JSON\nwith open(OUTPUT_JSON_PATH, 'w') as f:\n    json.dump(metadata_list, f, indent=2)\n\nprint(f\"📁 Metadata saved to {OUTPUT_JSON_PATH}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:05.650742Z","iopub.execute_input":"2025-07-01T09:32:05.650981Z","iopub.status.idle":"2025-07-01T09:32:16.554129Z","shell.execute_reply.started":"2025-07-01T09:32:05.650966Z","shell.execute_reply":"2025-07-01T09:32:16.553283Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import json\n# from collections import Counter\n\n# # Load JSON file (assuming it's a list of entries)\n# with open(\"/kaggle/working/train_metadata.json\", \"r\") as f:\n#     data = json.load(f)\n\n# # Collect (patient_id, study_id) pairs\n# id_pairs = [(entry[\"patient_id\"], entry[\"study_id\"]) for entry in data]\n\n# # Count how many times each pair appears\n# pair_counts = Counter(id_pairs)\n\n# # Find duplicates\n# duplicates = [pair for pair, count in pair_counts.items() if count > 1]\n\n# # Print results\n# if duplicates:\n#     print(f\"Found {len(duplicates)} duplicate entries:\")\n#     for pair in duplicates:\n#         print(f\" - patient_id: {pair[0]}, study_id: {pair[1]}\")\n# else:\n#     print(\"✅ No duplicate (patient_id, study_id) entries found.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:16.555001Z","iopub.execute_input":"2025-07-01T09:32:16.555208Z","iopub.status.idle":"2025-07-01T09:32:16.559305Z","shell.execute_reply.started":"2025-07-01T09:32:16.555193Z","shell.execute_reply":"2025-07-01T09:32:16.558458Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Extract and flatten labels into a DataFrame ---\nlabel_rows = []\nfor entry in metadata_list:\n    row = entry[\"labels\"]\n    label_rows.append(row)\n\n\nlabels_df = pd.DataFrame(label_rows)\n\n# --- Aggregate injury counts by organ ---\nlabels_agg = pd.DataFrame({\n    'bowel_healthy': [labels_df['bowel_healthy'].sum()],\n    'bowel_injury': [labels_df['bowel_injury'].sum()],\n    'extravasation_healthy': [labels_df['extravasation_healthy'].sum()],\n    'extravasation_injury': [labels_df['extravasation_injury'].sum()],\n    'kidney_healthy': [labels_df['kidney_healthy'].sum()],\n    'kidney_injury': [labels_df['kidney_low'].sum() + labels_df['kidney_high'].sum()],\n    'liver_healthy': [labels_df['liver_healthy'].sum()],\n    'liver_injury': [labels_df['liver_low'].sum() + labels_df['liver_high'].sum()],\n    'spleen_healthy': [labels_df['spleen_healthy'].sum()],\n    'spleen_injury': [labels_df['spleen_low'].sum() + labels_df['spleen_high'].sum()]\n})\n\n# --- Prepare for plotting ---\nlabels_agg = labels_agg.T.reset_index()\nlabels_agg.columns = ['label', 'count']\nlabels_agg[['organ', 'status']] = labels_agg['label'].str.rsplit('_', n=1, expand=True)\npivot_df = labels_agg.pivot(index='organ', columns='status', values='count').fillna(0)\n\n# --- Print counts in console ---\nprint(\"Injury counts per class:\")\nprint(labels_agg[['label', 'count']].to_string(index=False))\n\n\n# --- Plot ---\npivot_df.plot(kind='bar', figsize=(10, 6), color=['skyblue', 'salmon'])\nplt.title(\"Healthy vs Injury Distribution per Organ\")\nplt.ylabel(\"Number of Samples\")\nplt.xlabel(\"Organ\")\nplt.xticks(rotation=0)\nplt.grid(axis='y', linestyle='--', alpha=0.7)\nplt.legend(title='Status')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:16.56024Z","iopub.execute_input":"2025-07-01T09:32:16.560524Z","iopub.status.idle":"2025-07-01T09:32:16.953579Z","shell.execute_reply.started":"2025-07-01T09:32:16.560499Z","shell.execute_reply":"2025-07-01T09:32:16.952916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Create a new column 'injury_status'\n# labels_df['injury_status'] = labels_df['any_injury'].apply(lambda x: 'injured' if x == 1 else 'healthy')\n\n# # Count samples by injury_status\n# status_counts = labels_df['injury_status'].value_counts()\n\n# print(\"Counts by injury status:\")\n# print(status_counts)\n\n# # Plot pie chart\n# status_counts.plot(\n#     kind='pie',\n#     colors=['skyblue', 'salmon'],\n#     autopct='%1.1f%%',\n#     startangle=90,\n#     ylabel='',  # Hide ylabel for cleaner plot\n#     title='Proportion of Healthy vs Injured Samples'\n# )\n# plt.tight_layout()\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:16.954345Z","iopub.execute_input":"2025-07-01T09:32:16.95464Z","iopub.status.idle":"2025-07-01T09:32:16.958366Z","shell.execute_reply.started":"2025-07-01T09:32:16.954616Z","shell.execute_reply":"2025-07-01T09:32:16.957852Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Multi-Label Co-Occurrence Matrix (Heatmap)","metadata":{}},{"cell_type":"code","source":"# import seaborn as sns\n\n# # Select injury columns only\n# injury_cols = ['bowel_injury', 'extravasation_injury', 'kidney_low', 'kidney_high',\n#                'liver_low', 'liver_high', 'spleen_low', 'spleen_high']\n\n# injury_only = labels_df[injury_cols].copy()\n# injury_only['kidney_injury'] = injury_only['kidney_low'] + injury_only['kidney_high']\n# injury_only['liver_injury'] = injury_only['liver_low'] + injury_only['liver_high']\n# injury_only['spleen_injury'] = injury_only['spleen_low'] + injury_only['spleen_high']\n\n# # Keep only binary (0/1)\n# injury_matrix = injury_only[['bowel_injury', 'extravasation_injury', \n#                              'kidney_injury', 'liver_injury', 'spleen_injury']]\n\n# # Compute correlation/co-occurrence matrix\n# co_occurrence = injury_matrix.T @ injury_matrix\n# sns.heatmap(co_occurrence, annot=True, fmt='d', cmap=\"Reds\")\n# plt.title(\"Injury Co-occurrence Heatmap\")\n# plt.tight_layout()\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:16.961244Z","iopub.execute_input":"2025-07-01T09:32:16.961672Z","iopub.status.idle":"2025-07-01T09:32:16.983422Z","shell.execute_reply.started":"2025-07-01T09:32:16.961655Z","shell.execute_reply":"2025-07-01T09:32:16.982791Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Class Imbalance Ratio per Organ","metadata":{}},{"cell_type":"code","source":"# for organ in ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']:\n#     if organ == 'kidney' or organ == 'liver' or organ == 'spleen':\n#         injury = labels_df[f'{organ}_low'] + labels_df[f'{organ}_high']\n#     else:\n#         injury = labels_df[f'{organ}_injury']\n#     healthy = labels_df[f'{organ}_healthy']\n#     ratio = injury.sum() / (injury.sum() + healthy.sum())\n#     print(f\"{organ.capitalize()} Injury Ratio: {ratio:.2%}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:16.984009Z","iopub.execute_input":"2025-07-01T09:32:16.984173Z","iopub.status.idle":"2025-07-01T09:32:17.000419Z","shell.execute_reply.started":"2025-07-01T09:32:16.98416Z","shell.execute_reply":"2025-07-01T09:32:16.999839Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Injury Count Histogram per Sample","metadata":{}},{"cell_type":"code","source":"# # Count how many injury labels each sample has\n# injury_per_sample = injury_matrix.sum(axis=1)\n\n# # Plot histogram\n# plt.figure(figsize=(6, 4))\n# injury_per_sample.hist(bins=range(0, 7), color='teal', rwidth=0.8)\n# plt.xlabel(\"Number of Injured Organs\")\n# plt.ylabel(\"Number of Samples\")\n# plt.title(\"Histogram of Injury Counts per Sample\")\n# plt.grid(axis='y', linestyle='--', alpha=0.7)\n# plt.tight_layout()\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.001159Z","iopub.execute_input":"2025-07-01T09:32:17.001382Z","iopub.status.idle":"2025-07-01T09:32:17.018129Z","shell.execute_reply.started":"2025-07-01T09:32:17.001363Z","shell.execute_reply":"2025-07-01T09:32:17.017536Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Patient-wise Injury Distribution","metadata":{}},{"cell_type":"code","source":"# # Optional: Visualize first 10 patients' labels as a heatmap\n# subset = injury_matrix.iloc[:10]  # First 10 rows\n# sns.heatmap(subset, annot=True, cbar=False, cmap=\"YlGnBu\")\n# plt.title(\"First 10 Patients - Injury Pattern Heatmap\")\n# plt.xlabel(\"Organ\")\n# plt.ylabel(\"Patient Index\")\n# plt.tight_layout()\n# plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.018796Z","iopub.execute_input":"2025-07-01T09:32:17.01897Z","iopub.status.idle":"2025-07-01T09:32:17.03962Z","shell.execute_reply.started":"2025-07-01T09:32:17.018956Z","shell.execute_reply":"2025-07-01T09:32:17.03907Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Check Missing or Anomalous Labels","metadata":{}},{"cell_type":"code","source":"# for organ in ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']:\n#     if organ in ['kidney', 'liver', 'spleen']:\n#         total = labels_df[f'{organ}_low'] + labels_df[f'{organ}_high'] + labels_df[f'{organ}_healthy']\n#     else:\n#         total = labels_df[f'{organ}_injury'] + labels_df[f'{organ}_healthy']\n    \n#     if not all(total == 1):\n#         print(f\"Inconsistency found in {organ} labels\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.040314Z","iopub.execute_input":"2025-07-01T09:32:17.040504Z","iopub.status.idle":"2025-07-01T09:32:17.055013Z","shell.execute_reply.started":"2025-07-01T09:32:17.040491Z","shell.execute_reply":"2025-07-01T09:32:17.054448Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Import MONAI Transform Classes\n\n# MONAI is a deep learning framework for medical imaging.\n# These are various image transformation tools used for preprocessing and augmentation.\n\nfrom monai.transforms import (\n    Compose, EnsureChannelFirst, EnsureType,\n    Orientation, Spacing, RandAffine, RandFlip,\n    NormalizeIntensity, RandScaleIntensity, RandShiftIntensity,\n    RandGaussianNoise, RandGaussianSmooth, RandAdjustContrast,\n    Resize, RandBiasField,\n    ToTensor\n)\n\nfrom monai.data import MetaTensor\nfrom monai.transforms import OneOf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.055645Z","iopub.execute_input":"2025-07-01T09:32:17.055834Z","iopub.status.idle":"2025-07-01T09:32:17.07254Z","shell.execute_reply.started":"2025-07-01T09:32:17.05582Z","shell.execute_reply":"2025-07-01T09:32:17.071926Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Configuration Class\n# Stores all important constants and settings in one place (like a dictionary).\nclass Config:\n    SEED = 42\n    IMAGE_SIZE = (128, 128, 128)  \n    BATCH_SIZE = 4\n    EPOCHS = 20\n    LR =1e-4\n    \n    # Target columns (i.e., all the labels you want to predict)\n    TARGET_COLS = [\n        \"bowel_healthy\", \"extravasation_healthy\",\n        \"bowel_injury\", \"extravasation_injury\",\n        \"kidney_healthy\", \"kidney_low\", \"kidney_high\",\n        \"liver_healthy\", \"liver_low\", \"liver_high\",\n        \"spleen_healthy\", \"spleen_low\", \"spleen_high\",\n    ]\n\n    NUM_CLASSES = len(TARGET_COLS)  # Assumes LABELS is a predefined list of column names\n\n    VOXEL_SPACING = (1.0, 1.0, 1.0)  # Used to normalize spacing in 3D CT scans\n\n    SPLIT_MODE = \"group\"  # Use 'group' split for stratified grouping, or 'random' for simple random split\n\n# Create an instance of the Config class to use in other parts of the code\nconfig = Config()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.073332Z","iopub.execute_input":"2025-07-01T09:32:17.07414Z","iopub.status.idle":"2025-07-01T09:32:17.089135Z","shell.execute_reply.started":"2025-07-01T09:32:17.07412Z","shell.execute_reply":"2025-07-01T09:32:17.088497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_transforms = Compose([\n    EnsureChannelFirst(),\n    EnsureType(),\n    Orientation(axcodes=\"RAS\"),\n    Spacing(pixdim=(1.0, 1.0, 1.0), mode=\"bilinear\"),\n    Resize(spatial_size=config.IMAGE_SIZE),\n\n    # Original transforms you had:\n    OneOf([\n        RandAffine(\n            rotate_range=(0, 0, np.pi/12),  # Limited Z-axis rotation only\n            shear_range=(0.1, 0.1, 0.1),\n            translate_range=(10, 10, 5),\n            scale_range=(0.1, 0.1, 0.1),\n            prob=0.5,\n            mode=\"bilinear\"\n        ),\n        RandFlip(prob=0.5, spatial_axis=0),\n        RandFlip(prob=0.5, spatial_axis=1),\n        RandFlip(prob=0.5, spatial_axis=2),\n    ]),\n    NormalizeIntensity(nonzero=True, channel_wise=True),\n    RandScaleIntensity(factors=0.1, prob=1.0),\n    RandShiftIntensity(offsets=0.1, prob=1.0),\n    RandGaussianNoise(prob=0.3, mean=0.0, std=0.1),\n    RandAdjustContrast(prob=0.3, gamma=(0.7, 1.5)),\n\n    # New safe additions:\n    RandGaussianSmooth(\n        prob=0.2, \n        sigma_x=(0.25, 0.5),  # Very mild smoothing\n        sigma_y=(0.25, 0.5),\n        sigma_z=(0.25, 0.5)\n    ),\n    RandBiasField(\n        prob=0.2, \n        coeff_range=(0.1, 0.3)  # Subtle intensity variations\n    ),\n\n    ToTensor()\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.089852Z","iopub.execute_input":"2025-07-01T09:32:17.090042Z","iopub.status.idle":"2025-07-01T09:32:17.1186Z","shell.execute_reply.started":"2025-07-01T09:32:17.090028Z","shell.execute_reply":"2025-07-01T09:32:17.117902Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Validation Transforms\n\n# These are simpler than training transforms and only do standard preprocessing.\n# No random changes here, just make sure data is consistent and normalized.\n\nval_transforms = Compose([\n    EnsureChannelFirst(),  # (Z, H, W) → (1, Z, H, W)\n    EnsureType(),  # Convert to MetaTensor\n    Orientation(axcodes=\"RAS\"),  # Set standard orientation\n    Spacing(pixdim=(1.0, 1.0, 1.0), mode=\"bilinear\"),  # Make voxel spacing uniform\n    Resize(spatial_size=config.IMAGE_SIZE),\n    NormalizeIntensity(nonzero=True, channel_wise=True),  # Normalize image intensities\n    ToTensor()  # Convert to tensor\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.119372Z","iopub.execute_input":"2025-07-01T09:32:17.119607Z","iopub.status.idle":"2025-07-01T09:32:17.128853Z","shell.execute_reply.started":"2025-07-01T09:32:17.119581Z","shell.execute_reply":"2025-07-01T09:32:17.12829Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Test Transforms\n# Same as validation transforms—no randomness. Used during testing and inference.\n\ntest_transforms = Compose([\n    EnsureChannelFirst(),\n    EnsureType(),\n    Orientation(axcodes=\"RAS\"),\n    Spacing(pixdim=(1.0, 1.0, 1.0), mode=\"bilinear\"),\n    Resize(spatial_size=config.IMAGE_SIZE),\n    NormalizeIntensity(nonzero=True, channel_wise=True),\n    ToTensor()\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.129634Z","iopub.execute_input":"2025-07-01T09:32:17.129833Z","iopub.status.idle":"2025-07-01T09:32:17.146073Z","shell.execute_reply.started":"2025-07-01T09:32:17.129818Z","shell.execute_reply":"2025-07-01T09:32:17.145379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport nibabel as nib\nfrom torch.utils.data import Dataset\nfrom monai.data import MetaTensor  \n\nclass RSNADataset(Dataset):\n    def __init__(self, metadata_list, transforms=None, has_labels=True):\n        \"\"\"\n        Args:\n            metadata_list: List of dictionaries with 'nifti_path' and 'labels'.\n            transforms: MONAI or array-style transforms (non-dict style).\n            has_labels: Whether to return labels (True during training/val).\n        \"\"\"\n        self.metadata_list = metadata_list\n        self.transforms = transforms\n        self.has_labels = has_labels\n\n    def __len__(self):\n        return len(self.metadata_list)\n\n    def __getitem__(self, idx):\n        entry = self.metadata_list[idx]\n\n        # --- Load the NIfTI file ---\n        nifti_path = entry[\"nifti_path\"]  # Full path to .nii.gz\n        nifti_img = nib.load(nifti_path)\n        volume = nifti_img.get_fdata().astype(np.float32)\n\n        # --- Rearrange dimensions (X, Y, Z) → (Z, Y, X) ---\n        volume = np.transpose(volume, (2, 1, 0))\n\n        # --- Add channel dimension: (1, Z, Y, X) ---\n        volume = np.expand_dims(volume, axis=0)\n\n        # --- Wrap in MetaTensor (optional, for MONAI compatibility) ---\n        meta = {\"original_channel_dim\": 0}\n        sample = MetaTensor(volume, meta=meta)\n\n        # --- Apply transforms if provided ---\n        if self.transforms:\n            sample = self.transforms(sample)\n\n        # --- Package label if available ---\n        if self.has_labels:\n            sample = {\n                \"image\": sample,\n                \"label\": np.array([entry[\"labels\"][key] for key in config.TARGET_COLS], dtype=np.float32),\n            }\n        else:\n            sample = {\"image\": sample}\n\n        return sample\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.146954Z","iopub.execute_input":"2025-07-01T09:32:17.147681Z","iopub.status.idle":"2025-07-01T09:32:17.167608Z","shell.execute_reply.started":"2025-07-01T09:32:17.147664Z","shell.execute_reply":"2025-07-01T09:32:17.166624Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Prepare Dataloaders for Training, Validation, and Testing\n# This function handles data splitting and creates PyTorch DataLoader objects.\n\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader\n\ndef prepare_dataloaders(metadata_list, train_transforms, val_transforms, test_transforms, config, subset_fraction=1.0):\n    \"\"\"\n    Prepare DataLoaders for training, validation, and testing with an option to use a subset.\n    \"\"\"\n    # Optional: Create a stratified subset\n    if subset_fraction < 1.0:\n        # Stratify by target labels (e.g., bowel_healthy, kidney_healthy)\n        df = pd.DataFrame(metadata_list)\n        df[\"combined_label\"] = df[\"labels\"].apply(lambda x: ''.join([str(x[k]) for k in sorted(x.keys())]))\n        \n        # Sample a fraction of the dataset with stratification\n        subset_df = df.groupby(\"combined_label\", group_keys=False).apply(\n            lambda x: x.sample(frac=subset_fraction, random_state=42)\n        )\n        metadata_list = subset_df.to_dict(orient=\"records\")\n        print(f\"⚡ Sampled {len(metadata_list)} samples for hyperparameter tuning.\")\n    \n    # Now prepare the data splits (random or group-based)\n    if config.SPLIT_MODE == \"random\":\n        train_meta, temp_meta = train_test_split(\n            metadata_list, test_size=0.3, random_state=42, shuffle=True\n        )\n        val_meta, test_meta = train_test_split(\n            temp_meta, test_size=0.5, random_state=42, shuffle=True\n        )\n\n    elif config.SPLIT_MODE == \"group\":\n        # Custom stratified split to maintain label distribution\n        train_meta, val_meta, test_meta = split_metadata_train_val_test(\n            metadata_list,\n            target_cols=config.TARGET_COLS,\n            val_size=0.15,\n            test_size=0.15,\n            seed=42\n        )\n    else:\n        raise ValueError(f\"Unknown SPLIT_MODE: {config.SPLIT_MODE}\")\n\n    # Create dataset objects\n    train_ds = RSNADataset(train_meta, transforms=train_transforms, has_labels=True)\n    val_ds   = RSNADataset(val_meta, transforms=val_transforms, has_labels=True)\n    test_ds  = RSNADataset(test_meta, transforms=test_transforms, has_labels=False)\n\n    # Create DataLoader objects (for batching and shuffling)\n    train_loader = DataLoader(train_ds, batch_size=config.BATCH_SIZE, shuffle=True)\n    val_loader   = DataLoader(val_ds, batch_size=config.BATCH_SIZE, shuffle=False)\n    test_loader  = DataLoader(test_ds, batch_size=1, shuffle=False)\n    print(f\"Training data size: {len(train_loader.dataset)}, Validation data size: {len(val_loader.dataset)}, Test data size: {len(test_loader.dataset)}\")\n    \n    return train_ds, val_ds, test_ds, train_loader, val_loader, test_loader\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.168387Z","iopub.execute_input":"2025-07-01T09:32:17.168593Z","iopub.status.idle":"2025-07-01T09:32:17.204599Z","shell.execute_reply.started":"2025-07-01T09:32:17.168552Z","shell.execute_reply":"2025-07-01T09:32:17.203826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def split_metadata_train_val_test(metadata_list, target_cols, val_size=0.1, test_size=0.1, seed=42):\n    \"\"\"\n    Custom stratified split to maintain the label distribution for train, validation, and test sets.\n    \"\"\"\n    # Convert metadata into a DataFrame\n    df = pd.DataFrame(metadata_list)\n\n    # Extract individual labels into separate columns\n    label_df = pd.json_normalize(df['labels'])\n    df = pd.concat([df.drop(columns='labels'), label_df], axis=1)\n\n    # Group rows by all target label combinations\n    grouped = df.groupby(target_cols)\n\n    # Empty splits\n    train_df, val_df, test_df = pd.DataFrame(), pd.DataFrame(), pd.DataFrame()\n\n    val_test_size = val_size + test_size\n\n    for _, group in grouped:\n        n = len(group)\n        if n == 1:\n            r = np.random.rand()\n            if r < test_size:\n                test_df = pd.concat([test_df, group], ignore_index=True)\n            elif r < val_test_size:\n                val_df = pd.concat([val_df, group], ignore_index=True)\n            else:\n                train_df = pd.concat([train_df, group], ignore_index=True)\n        else:\n            train_split, val_test_split = train_test_split(group, test_size=val_test_size, random_state=seed)\n\n            if len(val_test_split) < 2:\n                val_split = val_test_split\n                test_split = pd.DataFrame()\n            else:\n                relative_test_size = test_size / val_test_size if val_test_size > 0 else 0\n                val_split, test_split = train_test_split(val_test_split, test_size=relative_test_size, random_state=seed)\n\n            train_df = pd.concat([train_df, train_split], ignore_index=True)\n            val_df = pd.concat([val_df, val_split], ignore_index=True)\n            test_df = pd.concat([test_df, test_split], ignore_index=True)\n\n    # Convert DataFrame rows back to metadata format\n    def row_to_metadata(row):\n        return {\n            \"nifti_path\": row[\"nifti_path\"],\n            \"labels\": {col: row[col] for col in target_cols}\n        }\n\n    train_list = [row_to_metadata(row) for _, row in train_df.iterrows()]\n    val_list   = [row_to_metadata(row) for _, row in val_df.iterrows()]\n    test_list  = [row_to_metadata(row) for _, row in test_df.iterrows()]\n\n    return train_list, val_list, test_list\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.205651Z","iopub.execute_input":"2025-07-01T09:32:17.205935Z","iopub.status.idle":"2025-07-01T09:32:17.214625Z","shell.execute_reply.started":"2025-07-01T09:32:17.205911Z","shell.execute_reply":"2025-07-01T09:32:17.214048Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Prepare Data Using the Loader Function (with subset for tuning)\n# It splits the metadata into train/val/test and creates DataLoader objects for batching and shuffling.\n\ntrain_ds, val_ds, test_ds, train_loader, val_loader, test_loader = prepare_dataloaders(\n    metadata_list=metadata_list,\n    train_transforms=train_transforms,\n    val_transforms=val_transforms,\n    test_transforms=test_transforms,\n    config=config,\n    subset_fraction=0.3  # 👈 Add this line to use 20% of the data for tuning\n)\n\nprint(f\"Train size: {len(train_loader.dataset)}\")\nprint(f\"Val size:   {len(val_loader.dataset)}\")\nprint(f\"Test size:  {len(test_loader.dataset)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.215482Z","iopub.execute_input":"2025-07-01T09:32:17.215743Z","iopub.status.idle":"2025-07-01T09:32:17.46234Z","shell.execute_reply.started":"2025-07-01T09:32:17.215728Z","shell.execute_reply":"2025-07-01T09:32:17.461625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import Counter, defaultdict\nimport matplotlib.pyplot as plt\n\ndef print_class_distribution(metadata_list, target_cols):\n    \"\"\"\n    Prints and optionally plots the class distribution of each target in the metadata list.\n    \"\"\"\n    label_counts = defaultdict(Counter)\n\n    for entry in metadata_list:\n        labels = entry[\"labels\"]\n        for col in target_cols:\n            label_counts[col][labels[col]] += 1\n\n    print(\"📊 Class distribution in the current subset:\")\n    for col in target_cols:\n        print(f\"\\n🔸 {col}\")\n        for label, count in sorted(label_counts[col].items()):\n            print(f\"   Label {label}: {count} samples\")\n\n        # Optional: plot a bar chart\n        plt.figure(figsize=(4, 2))\n        plt.bar(label_counts[col].keys(), label_counts[col].values(), color=\"#4a6cf7\")\n        plt.title(f\"{col} class distribution\")\n        plt.xlabel(\"Label\")\n        plt.ylabel(\"Count\")\n        plt.grid(axis='y', linestyle='--', alpha=0.6)\n        plt.tight_layout()\n        plt.show()\n\n# 🔍 Call this after dataloader prep:\nprint_class_distribution(train_ds.metadata_list, config.TARGET_COLS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:17.46332Z","iopub.execute_input":"2025-07-01T09:32:17.464179Z","iopub.status.idle":"2025-07-01T09:32:19.091377Z","shell.execute_reply.started":"2025-07-01T09:32:17.464145Z","shell.execute_reply":"2025-07-01T09:32:19.090776Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Updated model architecture","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom monai.networks.nets import DenseNet121\n\nclass DenseNet121model(nn.Module):\n    def __init__(self, in_channels=1, pretrained=False):\n        super().__init__()\n        \n        # Grad-CAM hooks\n        self.activations = None\n        self.gradients = None\n        \n        # Backbone - using MONAI's DenseNet121 which already includes GAP\n        self.backbone = DenseNet121(\n            spatial_dims=3,\n            in_channels=in_channels,\n            out_channels=512,  # This is the output after GAP\n            pretrained=pretrained\n        )\n        \n        # Register hook for Grad-CAM on the last conv layer\n        self.backbone.features[-1].register_forward_hook(self.save_activation)\n        \n        # Classification heads\n        self.bowel_head = self._create_binary_head()\n        self.extra_head = self._create_binary_head()\n        self.liver_head = self._create_multiclass_head()\n        self.kidney_head = self._create_multiclass_head()\n        self.spleen_head = self._create_multiclass_head()\n        \n    def _create_binary_head(self):\n        return nn.Sequential(\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.SiLU(),\n            nn.Dropout(0.3),\n            nn.Linear(256, 1)\n        )\n    \n    def _create_multiclass_head(self):\n        return nn.Sequential(\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.SiLU(),\n            nn.Dropout(0.3),\n            nn.Linear(256, 3)\n        )\n    \n    def save_activation(self, module, input, output):\n        \"\"\"Save activations for Grad-CAM\"\"\"\n        self.activations = output\n        if output.requires_grad:\n            output.register_hook(self.save_gradient)\n    \n    def save_gradient(self, grad):\n        \"\"\"Save gradients for Grad-CAM\"\"\"\n        self.gradients = grad\n    \n    def forward(self, x):\n        # Forward pass through backbone\n        if x.requires_grad:\n            x.register_hook(self.save_gradient)\n        self.activations = x\n        \n        # Get features (shape: [B, 512] after GAP)\n        features = self.backbone(x)\n        \n        return {\n            \"bowel\": self.bowel_head(features),\n            \"extra\": self.extra_head(features),\n            \"liver\": self.liver_head(features),\n            \"kidney\": self.kidney_head(features),\n            \"spleen\": self.spleen_head(features)\n        }\n    \n    def get_activations_gradient(self):\n        return self.gradients\n    \n    def get_activations(self):\n        return self.activations","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.092232Z","iopub.execute_input":"2025-07-01T09:32:19.092464Z","iopub.status.idle":"2025-07-01T09:32:19.101697Z","shell.execute_reply.started":"2025-07-01T09:32:19.092447Z","shell.execute_reply":"2025-07-01T09:32:19.101095Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Initialize Model and Move to Device (GPU or CPU)\nmodel = DenseNet121model().to(device)\nprint(f\"Model moved to {next(model.parameters()).device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.10277Z","iopub.execute_input":"2025-07-01T09:32:19.102994Z","iopub.status.idle":"2025-07-01T09:32:19.591897Z","shell.execute_reply.started":"2025-07-01T09:32:19.102974Z","shell.execute_reply":"2025-07-01T09:32:19.591253Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\nfrom scipy.ndimage import zoom\n\ndef compute_gradcam(model, input_tensor, target_head=\"bowel\", class_index=None, target_shape=None):\n    model.eval()\n    model.zero_grad()\n\n    # Forward pass\n    output = model(input_tensor)\n\n    # Select target class output\n    if class_index is None:\n        class_index = 0\n    loss = output[target_head][0, class_index]\n    loss.backward()\n\n    # Grab activations and gradients from the model\n    activations = model.activations  # Shape: (B, C, D, H, W)\n    grads = model.gradients          # Same shape\n\n    # Global average pooling of gradients over spatial dims\n    pooled_grads = torch.mean(grads, dim=(2, 3, 4), keepdim=True)  # Shape: (B, C, 1, 1, 1)\n\n    # Weighted sum of activations\n    weighted_activations = activations * pooled_grads  # Shape: (B, C, D, H, W)\n    cam = weighted_activations.sum(dim=1).squeeze()    # Shape: (D, H, W)\n\n    # ReLU and normalize\n    cam = torch.relu(cam)\n    cam = cam / (cam.max() + 1e-5)\n\n    # Convert to NumPy\n    cam_np = cam.detach().cpu().numpy()  # Shape: (D, H, W)\n\n    # Resize to original volume shape if provided\n    if target_shape:\n        if len(target_shape) != 3:\n            target_shape = target_shape[-3:]\n        zoom_factors = [t / c for t, c in zip(target_shape, cam_np.shape)]\n        cam_np = zoom(cam_np, zoom=zoom_factors, order=1)  # Linear interpolation\n\n    return cam_np\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.59264Z","iopub.execute_input":"2025-07-01T09:32:19.592892Z","iopub.status.idle":"2025-07-01T09:32:19.599798Z","shell.execute_reply.started":"2025-07-01T09:32:19.592856Z","shell.execute_reply":"2025-07-01T09:32:19.599062Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass BinaryFocalLoss(nn.Module):\n    def __init__(self, alpha=0.8, gamma=2.0, reduction=\"mean\", label_smoothing=0.1):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n        self.label_smoothing = label_smoothing\n        \n    def forward(self, inputs, targets):\n        # Apply label smoothing\n        targets = targets * (1 - self.label_smoothing) + 0.5 * self.label_smoothing\n        \n        # Numerically stable implementation\n        bce_loss = F.binary_cross_entropy_with_logits(\n            inputs, targets, \n            reduction='none'\n        )\n        \n        # Compute pt\n        pt = torch.exp(-bce_loss)\n        focal_loss = self.alpha * (1-pt)**self.gamma * bce_loss\n        \n        if self.reduction == \"mean\":\n            return focal_loss.mean()\n        elif self.reduction == \"sum\":\n            return focal_loss.sum()\n        return focal_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.603903Z","iopub.execute_input":"2025-07-01T09:32:19.604158Z","iopub.status.idle":"2025-07-01T09:32:19.620846Z","shell.execute_reply.started":"2025-07-01T09:32:19.604143Z","shell.execute_reply":"2025-07-01T09:32:19.620258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultiClassFocalLoss(nn.Module):\n    def __init__(self, weight=None, gamma=2.0, reduction='mean'):\n        super().__init__()\n        self.weight = weight\n        self.gamma = gamma\n        self.reduction = reduction\n        \n    def forward(self, inputs, targets):\n        ce_loss = F.cross_entropy(inputs, targets, reduction='none', weight=self.weight)\n        pt = torch.exp(-ce_loss)\n        focal_loss = ((1 - pt) ** self.gamma) * ce_loss\n        \n        if self.reduction == \"mean\":\n            return focal_loss.mean()\n        elif self.reduction == \"sum\":\n            return focal_loss.sum()\n        return focal_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.621708Z","iopub.execute_input":"2025-07-01T09:32:19.622237Z","iopub.status.idle":"2025-07-01T09:32:19.639859Z","shell.execute_reply.started":"2025-07-01T09:32:19.622212Z","shell.execute_reply":"2025-07-01T09:32:19.63924Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Calculate class frequencies\ndef compute_class_freqs(metadata_list):\n    labels = torch.stack([torch.tensor([d[\"labels\"][col] for col in config.TARGET_COLS]) \n                         for d in metadata_list])\n    \n    bowel_pos = labels[:, 2].sum().item()  # bowel_injury\n    extra_pos = labels[:, 3].sum().item() # extravasation_injury\n    \n    liver_counts = labels[:, 7:10].sum(dim=0)\n    kidney_counts = labels[:, 4:7].sum(dim=0)\n    spleen_counts = labels[:, 10:13].sum(dim=0)\n    \n    return {\n        \"bowel\": torch.tensor([1.0, max(2.0, len(metadata_list)/(2*bowel_pos))]),\n        \"extra\": torch.tensor([1.0, max(2.0, len(metadata_list)/(2*extra_pos))]),\n        \"liver\": 1.0 / (liver_counts + 1e-6),\n        \"kidney\": 1.0 / (kidney_counts + 1e-6),\n        \"spleen\": 1.0 / (spleen_counts + 1e-6)\n    }\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.640638Z","iopub.execute_input":"2025-07-01T09:32:19.640893Z","iopub.status.idle":"2025-07-01T09:32:19.656165Z","shell.execute_reply.started":"2025-07-01T09:32:19.640871Z","shell.execute_reply":"2025-07-01T09:32:19.655506Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compute_class_weights(metadata_list):\n    labels_df = pd.DataFrame([entry[\"labels\"] for entry in metadata_list])\n    \n    # More robust alpha calculation\n    alpha_dict = {\n        \"bowel\": min(0.9, max(0.6, 1 - labels_df[\"bowel_injury\"].mean())),\n        \"extra\": min(0.85, max(0.6, 1 - labels_df[\"extravasation_injury\"].mean()))\n    }\n    \n    # Smoother class weights for multi-class\n    ce_weights = {}\n    for organ in [\"kidney\", \"liver\", \"spleen\"]:\n        counts = labels_df[[f\"{organ}_healthy\", f\"{organ}_low\", f\"{organ}_high\"]].mean(axis=0)\n        weights = 1.0 / (counts + 1e-6)\n        ce_weights[organ] = (weights / weights.sum()).values\n    return alpha_dict, ce_weights\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.656957Z","iopub.execute_input":"2025-07-01T09:32:19.657216Z","iopub.status.idle":"2025-07-01T09:32:19.676758Z","shell.execute_reply.started":"2025-07-01T09:32:19.657193Z","shell.execute_reply":"2025-07-01T09:32:19.676058Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Define Loss Functions for Each Output Head\n# - BCEWithLogitsLoss is used for binary classification (outputs NOT passed through sigmoid yet).\n# - CrossEntropyLoss is used for multi-class classification (outputs NOT passed through softmax yet).\n\n# Initialize losses\nfreqs = compute_class_freqs(train_ds.metadata_list)\n\nloss_fn_dict = {\n    \"bowel\": BinaryFocalLoss(\n        alpha=0.9, \n        gamma=2.0,\n        label_smoothing=0.05\n    ),\n    \"extra\": BinaryFocalLoss(\n        alpha=0.95,\n        gamma=3.0,\n        label_smoothing=0.05\n    ),\n    \"liver\": MultiClassFocalLoss(\n        weight=torch.tensor(freqs[\"liver\"], dtype=torch.float32).to(device),\n        gamma=1.5\n    ),\n    \"kidney\": MultiClassFocalLoss(\n        weight=torch.tensor(freqs[\"kidney\"], dtype=torch.float32).to(device),\n        gamma=1.5\n    ),\n    \"spleen\": MultiClassFocalLoss(\n        weight=torch.tensor(freqs[\"spleen\"], dtype=torch.float32).to(device),\n        gamma=1.5\n    )\n}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.677647Z","iopub.execute_input":"2025-07-01T09:32:19.67786Z","iopub.status.idle":"2025-07-01T09:32:19.727497Z","shell.execute_reply.started":"2025-07-01T09:32:19.677826Z","shell.execute_reply":"2025-07-01T09:32:19.726718Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Replace current optimizer setup with:\noptimizer = torch.optim.AdamW(model.parameters(), lr=config.LR, weight_decay=1e-5)\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer, \n    max_lr=1e-4,\n    epochs=config.EPOCHS,\n    steps_per_epoch=len(train_loader),\n    pct_start=0.3\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.728335Z","iopub.execute_input":"2025-07-01T09:32:19.72858Z","iopub.status.idle":"2025-07-01T09:32:19.734647Z","shell.execute_reply.started":"2025-07-01T09:32:19.728538Z","shell.execute_reply":"2025-07-01T09:32:19.733853Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import torch\n# import matplotlib.pyplot as plt\n# import copy\n\n# def lr_finder(model, train_loader, loss_fn_dict, optimizer, device, start_lr=1e-7, end_lr=1, num_iter=100):\n#     model.train()\n#     model = model.to(device)\n\n#     # Save initial weights so we can restore later\n#     initial_state = copy.deepcopy(model.state_dict())\n\n#     lr_lambda = lambda x: (end_lr / start_lr) ** (x / num_iter)\n#     scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lr_lambda)\n\n#     lrs = []\n#     losses = []\n\n#     iter_count = 0\n\n#     for batch in train_loader:\n#         inputs = batch[\"image\"]\n#         targets = batch[\"label\"]\n\n#         inputs = inputs.to(device)\n#         targets = targets.to(device)\n\n#         # ➤ Split targets into dict of organ-specific outputs\n#         targets_dict = {\n#             \"bowel\": targets[:, 1].unsqueeze(1),                  # ➤ Shape [B, 1]\n#             \"extra\": targets[:, 3].unsqueeze(1),                  # ➤ Shape [B, 1]\n#             \"liver\": targets[:, 4:7].argmax(dim=1),               # ➤ Shape [B]\n#             \"kidney\": targets[:, 7:10].argmax(dim=1),             # ➤ Shape [B]\n#             \"spleen\": targets[:, 10:13].argmax(dim=1),            # ➤ Shape [B]\n#         }\n\n#         optimizer.zero_grad()\n#         outputs = model(inputs)\n\n#         loss = 0\n#         for key in loss_fn_dict:\n#             loss += loss_fn_dict[key](outputs[key], targets_dict[key])\n\n#         loss.backward()\n#         optimizer.step()\n#         scheduler.step()\n\n#         lrs.append(optimizer.param_groups[0][\"lr\"])\n#         losses.append(loss.item())\n\n#         iter_count += 1\n#         if iter_count >= num_iter:\n#             break\n\n#     # Restore original weights\n#     model.load_state_dict(initial_state)\n\n#     # Plot\n#     plt.figure(figsize=(8, 6))\n#     plt.plot(lrs, losses)\n#     plt.xscale('log')\n#     plt.xlabel(\"Learning Rate\")\n#     plt.ylabel(\"Loss\")\n#     plt.title(\"Learning Rate Finder\")\n#     plt.grid(True)\n#     plt.show()\n\n# lr_finder(\n#     model=model,\n#     train_loader=train_loader,\n#     loss_fn_dict=loss_fn_dict,\n#     optimizer=optimizer,\n#     device=device,\n#     start_lr=1e-7,\n#     end_lr=1,\n#     num_iter=100\n# )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.735323Z","iopub.execute_input":"2025-07-01T09:32:19.735707Z","iopub.status.idle":"2025-07-01T09:32:19.752073Z","shell.execute_reply.started":"2025-07-01T09:32:19.735625Z","shell.execute_reply":"2025-07-01T09:32:19.751342Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nfrom collections import defaultdict\nimport torch\nimport numpy as np\nfrom sklearn.metrics import precision_recall_fscore_support, roc_auc_score, accuracy_score\n\ndef train_one_epoch(model, loader, optimizer, loss_fn_dict, scheduler=None, grad_clip=None):\n    model.train()\n    running_loss = 0.0\n    task_losses = defaultdict(float)\n    pbar = tqdm(loader, desc=\"Training\", leave=False)\n\n    for batch in pbar:\n        inputs = batch[\"image\"].to(device, dtype=torch.float32)\n        labels = batch[\"label\"].to(device, dtype=torch.float32)\n\n        # Handle potential missing dimension\n        if inputs.ndim == 4:\n            inputs = inputs.unsqueeze(1)  # Add channel dimension if missing\n\n        optimizer.zero_grad(set_to_none=True)  # More memory efficient\n        \n        # Forward pass\n        outputs = model(inputs)\n\n        # Prepare targets\n        targets = {\n            \"bowel\": labels[:, 0:2].max(dim=1)[0].float(),  # Binary: take max (one-hot)\n            \"extra\": labels[:, 2:4].max(dim=1)[0].float(),\n            \"kidney\": labels[:, 4:7].argmax(dim=1),  # Multi-class: argmax\n            \"liver\": labels[:, 7:10].argmax(dim=1),\n            \"spleen\": labels[:, 10:13].argmax(dim=1),\n        }\n\n        # Calculate loss per task\n        loss = 0.0\n        for key in outputs:\n            pred = outputs[key]\n            target = targets[key]\n            \n            # Handle different task types\n            if key in [\"bowel\", \"extra\"]:  # Binary tasks\n                pred = pred.squeeze(-1) if pred.ndim > 1 else pred\n                target = target.float()\n                task_loss = loss_fn_dict[key](pred, target)\n            else:  # Multi-class tasks\n                target = target.long()\n                task_loss = loss_fn_dict[key](pred, target)\n            \n            # Store individual task losses for monitoring\n            task_losses[key] += task_loss.item()\n            loss += task_loss\n\n        # Backpropagation with gradient clipping\n        loss.backward()\n        if grad_clip is not None:\n            torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)\n        \n        optimizer.step()\n        if scheduler is not None:\n            scheduler.step()  # For per-batch scheduling (e.g., OneCycleLR)\n\n        # Update progress bar\n        running_loss += loss.item()\n        avg_loss = running_loss / (pbar.n + 1)\n        postfix = {\"loss\": avg_loss}\n        \n        # Add task-specific losses to progress bar\n        for key in task_losses:\n            postfix[f\"{key}_loss\"] = task_losses[key] / (pbar.n + 1)\n        \n        pbar.set_postfix(postfix)\n\n    # Convert task losses to averages\n    task_losses = {k: v/len(loader) for k,v in task_losses.items()}\n    return avg_loss, task_losses","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.75286Z","iopub.execute_input":"2025-07-01T09:32:19.753113Z","iopub.status.idle":"2025-07-01T09:32:19.772542Z","shell.execute_reply.started":"2025-07-01T09:32:19.75308Z","shell.execute_reply":"2025-07-01T09:32:19.771811Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef validate(model, loader, loss_fn_dict):\n    model.eval()\n    val_loss = 0.0\n    pbar = tqdm(loader, desc=\"Validation\", leave=False)\n\n    all_preds = defaultdict(list)\n    all_targets = defaultdict(list)\n\n    for batch in pbar:\n        inputs = batch[\"image\"].to(device, dtype=torch.float32)\n        labels = batch[\"label\"].to(device, dtype=torch.float32)\n\n        if inputs.ndim == 4:\n            inputs = inputs.unsqueeze(2)\n\n        outputs = model(inputs)\n\n        targets = {\n            \"bowel\": labels[:, 0:2].max(dim=1)[0].float(),\n            \"extra\": labels[:, 2:4].max(dim=1)[0].float(),\n            \"kidney\": labels[:, 4:7].argmax(dim=1),\n            \"liver\": labels[:, 7:10].argmax(dim=1),\n            \"spleen\": labels[:, 10:13].argmax(dim=1),\n        }\n\n        loss = 0.0\n        for key in outputs:\n            pred = outputs[key]\n            target = targets[key]\n\n            if pred.shape[-1] == 1:\n                prob = torch.sigmoid(pred).view(-1).cpu().numpy()\n                bin_pred = (prob >= 0.5).astype(int)\n                target_np = target.cpu().numpy().astype(int)\n\n                all_preds[key].extend(bin_pred)\n                all_targets[key].extend(target_np)\n                loss += loss_fn_dict[key](pred.view(-1), target.float().view(-1))\n\n            else:\n                softmax_pred = torch.softmax(pred, dim=1)\n                class_pred = torch.argmax(softmax_pred, dim=1).cpu().numpy()\n                target_np = target.cpu().numpy()\n\n                all_preds[key].extend(class_pred)\n                all_targets[key].extend(target_np)\n                loss += loss_fn_dict[key](pred, target)\n\n        val_loss += loss.item()\n        pbar.set_postfix({\"val_loss\": val_loss / (pbar.n + 1)})\n\n    metrics = {}\n    print(\"\\n--- Evaluation Metrics ---\")\n    for key in all_preds:\n        y_true = np.array(all_targets[key])\n        y_pred = np.array(all_preds[key])\n\n        metrics[key] = {}\n\n        if len(np.unique(y_true)) <= 1:\n            print(f\"{key}: Not enough class diversity in ground truth to compute metrics.\")\n            continue\n\n        if set(np.unique(y_true)) <= {0, 1}:\n            precision, recall, f1, _ = precision_recall_fscore_support(y_true, y_pred, average='binary')\n            acc = accuracy_score(y_true, y_pred)\n            try:\n                roc = roc_auc_score(y_true, y_pred)\n            except ValueError:\n                roc = np.nan\n        else:\n            precision, recall, f1, _ = precision_recall_fscore_support(y_true, y_pred, average='macro')\n            acc = accuracy_score(y_true, y_pred)\n            try:\n                roc = roc_auc_score(\n                    y_true,\n                    torch.nn.functional.one_hot(torch.tensor(y_pred), num_classes=len(np.unique(y_true))),\n                    multi_class='ovo')\n            except ValueError:\n                roc = np.nan\n\n        metrics[key]['precision'] = precision\n        metrics[key]['recall'] = recall\n        metrics[key]['f1'] = f1\n        metrics[key]['accuracy'] = acc\n        metrics[key]['roc_auc'] = roc\n\n        print(f\"{key.capitalize()} | Acc: {acc:.3f} | Precision: {precision:.3f} | Recall: {recall:.3f} | F1: {f1:.3f} | ROC-AUC: {roc:.3f}\")\n\n    return val_loss / len(loader), metrics\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.773345Z","iopub.execute_input":"2025-07-01T09:32:19.773532Z","iopub.status.idle":"2025-07-01T09:32:19.797077Z","shell.execute_reply.started":"2025-07-01T09:32:19.773516Z","shell.execute_reply":"2025-07-01T09:32:19.79642Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport seaborn as sns\nimport zipfile\nimport os\nimport torch\nfrom datetime import datetime\nfrom kaggle.api.kaggle_api_extended import KaggleApi\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix\nimport seaborn as sns\n\n\n# Authenticate once globally\napi = KaggleApi()\napi.authenticate()\n\ndef upload_to_kaggle_model(dataset_owner, dataset_slug, model_path, checkpoint_path=None, version_note=\"\"):\n    import zipfile\n    import os\n    import json\n\n    zip_path = \"/kaggle/working/model_upload.zip\"\n    \n    # Zip model(s)\n    with zipfile.ZipFile(zip_path, 'w') as zipf:\n        zipf.write(model_path, arcname=os.path.basename(model_path))\n        if checkpoint_path:\n            zipf.write(checkpoint_path, arcname=os.path.basename(checkpoint_path))\n    \n    print(f\"📦 Zipped model(s) to: {zip_path}\")\n\n    # Create dataset-metadata.json file for Kaggle API\n    metadata = {\n        \"title\": f\"{dataset_slug} model\",\n        \"id\": f\"{dataset_owner}/{dataset_slug}\",\n        \"licenses\": [{\"name\": \"CC0-1.0\"}]\n    }\n    metadata_path = \"/kaggle/working/dataset-metadata.json\"\n    with open(metadata_path, \"w\") as f:\n        json.dump(metadata, f, indent=2)\n    print(f\"✅ Created metadata file at {metadata_path}\")\n\n    # Upload as new version to Kaggle dataset\n    api.dataset_create_version(\n        folder=\"/kaggle/working\",\n        version_notes=version_note,\n        delete_old_versions=False,\n        convert_to_csv=False\n    )\n    print(f\"✅ Uploaded to Kaggle Dataset: {dataset_owner}/{dataset_slug}\")\n\n\n\ndef extract_checkpoint_from_dataset(dataset_input_path=\"/kaggle/input/densenet-training-model-arch-focal-loss\",\n                                    extract_path=\"/kaggle/working\"):\n    checkpoint_src = os.path.join(dataset_input_path, \"checkpoint.pth\")\n    best_model_src = os.path.join(dataset_input_path, \"model_best.pth\")\n    \n    copied = False\n\n    if os.path.exists(checkpoint_src):\n        shutil.copy(checkpoint_src, os.path.join(extract_path, \"checkpoint.pth\"))\n        print(f\"✅ Copied checkpoint.pth to {extract_path}\")\n        copied = True\n    else:\n        print(\"⚠️ checkpoint.pth not found in dataset input.\")\n\n    if os.path.exists(best_model_src):\n        shutil.copy(best_model_src, os.path.join(extract_path, \"model_best.pth\"))\n        print(f\"✅ Copied model_best.pth to {extract_path}\")\n        copied = True\n    else:\n        print(\"⚠️ model_best.pth not found in dataset input.\")\n\n    return extract_path if copied else None\n\ndef train(model, train_loader, val_loader, optimizer, loss_fn_dict, num_epochs,\n          save_dir=\"/kaggle/working\", resume=False,\n          upload_to_kaggle=False, dataset_owner=None, dataset_slug=None, hyperparam_note=\"\"):\n\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n    os.makedirs(save_dir, exist_ok=True)\n\n    best_val_loss = float('inf')\n    start_epoch = 0\n\n    checkpoint_path = os.path.join(save_dir, \"checkpoint.pth\")\n    best_model_path = os.path.join(save_dir, \"model_best.pth\")\n\n    history = {\n        'train_loss': [],\n        'val_loss': [],\n        'metrics': {\n            'bowel': {'precision': [], 'recall': [], 'f1': [], 'roc_auc': [], 'accuracy': []},\n            'extra': {'precision': [], 'recall': [], 'f1': [], 'roc_auc': [], 'accuracy': []},\n            'kidney': {'precision': [], 'recall': [], 'f1': [], 'roc_auc': [], 'accuracy': []},\n            'liver': {'precision': [], 'recall': [], 'f1': [], 'roc_auc': [], 'accuracy': []},\n            'spleen': {'precision': [], 'recall': [], 'f1': [], 'roc_auc': [], 'accuracy': []},\n        }\n    }\n    \n    # Try to extract checkpoint if resume and dataset input zip exists\n    if resume:\n        extracted_path = extract_checkpoint_from_dataset()\n        if extracted_path is not None:\n            checkpoint_path_from_extract = os.path.join(extracted_path, \"checkpoint.pth\")\n            best_model_path_from_extract = os.path.join(extracted_path, \"model_best.pth\")\n            # Copy extracted checkpoint and best model to save_dir for training continuity\n            if os.path.exists(checkpoint_path_from_extract):\n                os.replace(checkpoint_path_from_extract, checkpoint_path)\n                print(f\"✅ Checkpoint copied to {checkpoint_path}\")\n            if os.path.exists(best_model_path_from_extract):\n                os.replace(best_model_path_from_extract, best_model_path)\n                print(f\"✅ Best model copied to {best_model_path}\")\n\n    # Resume from local checkpoint if exists\n    # Resume from local checkpoint if exists\n    if resume and os.path.exists(checkpoint_path):\n        checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)\n        model.load_state_dict(checkpoint['model_state_dict'])\n        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n        start_epoch = checkpoint['epoch'] + 1\n        best_val_loss = checkpoint.get('best_val_loss', best_val_loss)\n        history = checkpoint.get('history', history)\n        print(f\"🔁 Resumed from epoch {start_epoch}\")\n    \n        # 🔧 Patch: Add 'accuracy' if missing from metrics history\n        for organ in history['metrics']:\n            if 'accuracy' not in history['metrics'][organ]:\n                history['metrics'][organ]['accuracy'] = []\n\n    for epoch in range(start_epoch, num_epochs):\n        print(f\"\\nEpoch [{epoch+1}/{num_epochs}]\")\n        train_loss = train_one_epoch(model, train_loader, optimizer, loss_fn_dict)\n        val_loss, metrics = validate(model, val_loader, loss_fn_dict)\n\n        history['train_loss'].append(train_loss)\n        history['val_loss'].append(val_loss)\n        for organ in metrics:\n            for metric in metrics[organ]:\n                history['metrics'][organ][metric].append(metrics[organ][metric])\n\n        # Save checkpoint\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'best_val_loss': best_val_loss,\n            'history': history,\n        }, checkpoint_path)\n\n        # Save best model + upload if improved\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            torch.save(model.state_dict(), best_model_path)\n            print(f\"✅ New best model saved! Val Loss = {val_loss:.4f}\")\n\n            if upload_to_kaggle and dataset_owner and dataset_slug:\n                note = f\"{hyperparam_note} | Epoch {epoch+1}, Val Loss: {val_loss:.4f}\"\n                upload_to_kaggle_model(dataset_owner, dataset_slug, best_model_path, checkpoint_path, version_note=note)\n        else:\n            print(f\"No improvement. Val Loss = {val_loss:.4f}\")\n\n            if upload_to_kaggle and dataset_owner and dataset_slug:\n                note = f\"{hyperparam_note} | Epoch {epoch+1}, No improvement\"\n                upload_to_kaggle_model(dataset_owner, dataset_slug, model_path=checkpoint_path, version_note=note)\n                \n    return model, history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:19.797831Z","iopub.execute_input":"2025-07-01T09:32:19.798048Z","iopub.status.idle":"2025-07-01T09:32:20.003448Z","shell.execute_reply.started":"2025-07-01T09:32:19.798024Z","shell.execute_reply":"2025-07-01T09:32:20.0029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset_owner = \"anusapkota\"\ndataset_slug = \"densenet-training-model-arch-focal-loss\"\ndataset_id = f\"{dataset_owner}/{dataset_slug}\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:20.004532Z","iopub.execute_input":"2025-07-01T09:32:20.004829Z","iopub.status.idle":"2025-07-01T09:32:20.0085Z","shell.execute_reply.started":"2025-07-01T09:32:20.004806Z","shell.execute_reply":"2025-07-01T09:32:20.007613Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Subset your training and validation loader to only 1 image\n# from torch.utils.data import Subset\n\n# train_subset = Subset(train_ds, [0])\n# val_subset = Subset(val_ds, [0])\n\n# train_loader = DataLoader(train_subset, batch_size=1)\n# val_loader = DataLoader(val_subset, batch_size=1)\n\n# # Train for 1 epoch\n# model, history = train(\n#     model,\n#     train_loader,\n#     val_loader,\n#     optimizer,\n#     loss_fn_dict,\n#     num_epochs=1,\n#     save_dir=\"/kaggle/working/\",\n#     resume=False,\n#     upload_to_kaggle=True,\n#     dataset_owner=dataset_owner,\n#     dataset_slug=dataset_slug,\n#     hyperparam_note=\"🧪 Test run with 1 sample\"\n# )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:20.009471Z","iopub.execute_input":"2025-07-01T09:32:20.009759Z","iopub.status.idle":"2025-07-01T09:32:20.024933Z","shell.execute_reply.started":"2025-07-01T09:32:20.009736Z","shell.execute_reply":"2025-07-01T09:32:20.024287Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Full Training Loop Over All Epochs\nNUM_EPOCHS = config.EPOCHS\nsave_dir = '/kaggle/working/'\n\nhyperparam_note=f\"lr={config.LR}, bs={config.BATCH_SIZE}, image_size={config.IMAGE_SIZE}, focal_loss: fine tuned subset = 30% \"\n\nprint(\"Is CUDA available?\", torch.cuda.is_available())\nif torch.cuda.is_available():\n    print(\"Current device:\", torch.cuda.current_device())\n    print(\"Device name:\", torch.cuda.get_device_name(torch.cuda.current_device()))\n\nmodel, history = train(\n    model,\n    train_loader,\n    val_loader,\n    optimizer,\n    loss_fn_dict,\n    NUM_EPOCHS,\n    save_dir,\n    resume=True,\n    upload_to_kaggle=True,\n    dataset_owner=dataset_owner,\n    dataset_slug=dataset_slug,\n    hyperparam_note=hyperparam_note\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-01T09:32:20.025704Z","iopub.execute_input":"2025-07-01T09:32:20.025973Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# from torch.utils.data import Subset, DataLoader\n\n# # Use only first 5 samples for training and validation\n# small_train_ds = Subset(train_ds, range(5))\n# small_val_ds = Subset(val_ds, range(5))\n\n# # Create smaller loaders\n# train_loader = DataLoader(small_train_ds, batch_size=1, shuffle=True)\n# val_loader = DataLoader(small_val_ds, batch_size=1, shuffle=False)\n\n# # Train only 1 epoch on small data\n# NUM_EPOCHS = 1\n# model, history = train(model, train_loader, val_loader, optimizer, loss_fn_dict, NUM_EPOCHS)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_model(path=\"model_best.pth\"):\n    model = DenseNet121model()\n    model.load_state_dict(torch.load(path, map_location=device))\n    model.eval().to(device)\n    return model\n\nload_model('checkpoints/model_best.pth')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_training_history(history):\n    # Create a figure with appropriate size\n    plt.figure(figsize=(20, 15))\n    \n    # Plot losses\n    plt.subplot(3, 2, 1)  # 3 rows, 2 columns, position 1\n    plt.plot(history['train_loss'], label='Train Loss')\n    plt.plot(history['val_loss'], label='Validation Loss')\n    plt.title('Training and Validation Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    \n    # Plot metrics for each task\n    tasks = list(history['metrics'].keys())\n    metrics = ['precision', 'recall', 'f1', 'roc_auc']\n    \n    # Plot each metric in its own subplot\n    for i, metric in enumerate(metrics):\n        plt.subplot(3, 2, i+2)  # Positions 2-5\n        for task in tasks:\n            if history['metrics'][task][metric]:  # Check if metric exists\n                plt.plot(history['metrics'][task][metric], label=f'{task}')\n        plt.title(f'{metric.capitalize()} per epoch')\n        plt.xlabel('Epoch')\n        plt.ylabel(metric.capitalize())\n        plt.legend()\n    \n    plt.tight_layout()\n    plt.show()\n    \ndef plot_confusion_matrices(model, loader, loss_fn_dict):\n    model.eval()\n    all_preds = defaultdict(list)\n    all_targets = defaultdict(list)\n    \n    with torch.no_grad():\n        for batch in tqdm(loader, desc=\"Generating predictions for confusion matrices\"):\n            inputs = batch[\"image\"].to(device, dtype=torch.float32)\n            labels = batch[\"label\"].to(device, dtype=torch.float32)\n\n            if inputs.ndim == 4:\n                inputs = inputs.unsqueeze(2)\n\n            outputs = model(inputs)\n\n            targets = {\n                \"bowel\": labels[:, 0:2].max(dim=1)[0].float(),\n                \"extra\": labels[:, 2:4].max(dim=1)[0].float(),\n                \"kidney\": labels[:, 4:7].argmax(dim=1),\n                \"liver\": labels[:, 7:10].argmax(dim=1),\n                \"spleen\": labels[:, 10:13].argmax(dim=1),\n            }\n\n            for key in outputs:\n                pred = outputs[key]\n                target = targets[key]\n\n                if pred.shape[-1] == 1:  # Binary case\n                    prob = torch.sigmoid(pred).view(-1).cpu().numpy()\n                    bin_pred = (prob >= 0.5).astype(int)\n                    target_np = target.cpu().numpy().astype(int)\n                else:  # Multi-class\n                    softmax_pred = torch.softmax(pred, dim=1)\n                    class_pred = torch.argmax(softmax_pred, dim=1).cpu().numpy()\n                    target_np = target.cpu().numpy()\n\n                all_preds[key].extend(class_pred if pred.shape[-1] != 1 else bin_pred)\n                all_targets[key].extend(target_np)\n    \n    # Plot confusion matrices\n    plt.figure(figsize=(20, 15))\n    for i, key in enumerate(all_preds, 1):\n        plt.subplot(2, 3, i)\n        cm = confusion_matrix(all_targets[key], all_preds[key])\n        sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n                   xticklabels=np.unique(all_targets[key]), \n                   yticklabels=np.unique(all_targets[key]))\n        plt.title(f'Confusion Matrix - {key.capitalize()}')\n        plt.xlabel('Predicted')\n        plt.ylabel('True')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot training history\nplot_training_history(history)\n\n# Plot confusion matrices on validation set\nplot_confusion_matrices(model, val_loader, loss_fn_dict)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef test_inference(model, loader, device=\"cuda\", max_batches=None):\n    model.eval()\n    predictions = []\n    indices = []\n\n    pbar = tqdm(enumerate(loader), total=max_batches or len(loader), desc=\"Inference\", leave=False)\n\n    for i, batch in pbar:\n        if max_batches is not None and i >= max_batches:\n            break\n\n        inputs = batch[\"image\"].to(device, dtype=torch.float32)\n\n        if inputs.ndim == 4:\n            inputs = inputs.unsqueeze(2)  # Ensure 5D\n\n        outputs = model(inputs)\n\n        # Apply activations to get probabilities\n        batch_preds = {\n            \"bowel\": torch.sigmoid(outputs[\"bowel\"]).cpu().numpy(),\n            \"extra\": torch.sigmoid(outputs[\"extra\"]).cpu().numpy(),\n            \"kidney\": F.softmax(outputs[\"kidney\"], dim=1).cpu().numpy(),\n            \"liver\": F.softmax(outputs[\"liver\"], dim=1).cpu().numpy(),\n            \"spleen\": F.softmax(outputs[\"spleen\"], dim=1).cpu().numpy(),\n        }\n\n        predictions.append(batch_preds)\n        indices.append(i)  # You can append batch-level index or patient ID if available\n\n    return predictions, indices","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import Subset, DataLoader\n\n# Create small subset of test dataset\nsmall_test_ds = Subset(test_ds, list(range(10)))  # First 10 samples\nsmall_test_loader =  DataLoader(small_test_ds, batch_size=1, shuffle=False)\n\n# Run inference on this subset\npreds, ids = test_inference(model, small_test_loader, device=device, max_batches=10)\n\ntest_sample = next(iter(small_test_loader))  # Get one sample from the test DataLoader\ninput_tensor = test_sample[\"image\"] \n\n# Example output inspection\nprint(\"Predictions for first sample:\")\nprint(preds[0])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i, pred in enumerate(preds[:3]):\n    print(f\"\\nSample {i}:\")\n    print(\"Bowel Injury Probability:\", pred[\"bowel\"].squeeze())\n    print(\"Extravasation Probability:\", pred[\"extra\"].squeeze())\n    print(\"Kidney Class Probabilities:\", pred[\"kidney\"].squeeze())\n    print(\"Liver Class Probabilities:\", pred[\"liver\"].squeeze())\n    print(\"Spleen Class Probabilities:\", pred[\"spleen\"].squeeze())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def interpret_predictions(preds, class_labels=None):\n    if class_labels is None:\n        class_labels = [\"healthy\", \"low\", \"high\"]  # for multiclass labels\n\n    for i, pred in enumerate(preds):\n        print(f\"\\nSample {i} Prediction:\")\n\n        # Binary: Injury Present or Not\n        bowel_status = \"Injured\" if pred[\"bowel\"].squeeze() > 0.5 else \"Healthy\"\n        extra_status = \"Extravasation\" if pred[\"extra\"].squeeze() > 0.5 else \"No Extravasation\"\n        print(f\"  ▸ Bowel Injury: {bowel_status} \")\n        print(f\"  ▸ Extravasation: {extra_status} \")\n\n        # Multiclass: argmax for label\n        kidney_label = class_labels[np.argmax(pred[\"kidney\"])]\n        liver_label = class_labels[np.argmax(pred[\"liver\"])]\n        spleen_label = class_labels[np.argmax(pred[\"spleen\"])]\n\n        print(f\"  ▸ Kidney Condition: {kidney_label}\")\n        print(f\"  ▸ Liver Condition: {liver_label}\")\n        print(f\"  ▸ Spleen Condition: {spleen_label}\")\n\ninterpret_predictions(preds[:3])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from skimage import measure\n\ndef show_gradcam_overlay(cam, original_volume, slice_axis=0, alpha=0.5, threshold=0.5, show_contours=True, slice_range=5):\n    \"\"\"\n    Show Grad-CAM overlay on multiple slices of the CT volume.\n    \n    Parameters:\n    - cam: 3D Grad-CAM heatmap (D, H, W)\n    - original_volume: 3D CT scan (D, H, W)\n    - slice_axis: 0, 1, or 2 → axis to slice along\n    - alpha: blending factor for overlay\n    - threshold: value (0–1) for contour visualization\n    - show_contours: whether to show contour on top of heatmap\n    - slice_range: number of slices to show before and after mid-slice\n    \"\"\"\n    assert cam.shape == original_volume.shape, \"CAM and CT shape mismatch\"\n\n    mid_slice = original_volume.shape[slice_axis] // 2\n    slice_indices = range(mid_slice - slice_range, mid_slice + slice_range + 1)\n\n    n_cols = 5\n    n_rows = int(np.ceil(len(slice_indices) / n_cols))\n\n    fig, axes = plt.subplots(n_rows, n_cols, figsize=(15, 3 * n_rows))\n    axes = axes.flatten()\n\n    for i, idx in enumerate(slice_indices):\n        if idx < 0 or idx >= original_volume.shape[slice_axis]:\n            continue\n\n        # Slice selection\n        if slice_axis == 0:\n            base = original_volume[idx, :, :]\n            heat = cam[idx, :, :]\n        elif slice_axis == 1:\n            base = original_volume[:, idx, :]\n            heat = cam[:, idx, :]\n        elif slice_axis == 2:\n            base = original_volume[:, :, idx]\n            heat = cam[:, :, idx]\n\n        # Normalize base image\n        base = (base - np.min(base)) / (np.max(base) - np.min(base) + 1e-5)\n\n        ax = axes[i]\n        ax.imshow(base, cmap=\"gray\")\n        im = ax.imshow(heat, cmap=\"jet\", alpha=alpha)\n        \n        if show_contours:\n            contours = measure.find_contours(heat, threshold)\n            for contour in contours:\n                ax.plot(contour[:, 1], contour[:, 0], linewidth=1.5, color='white')\n\n        ax.set_title(f\"Slice {idx}\")\n        ax.axis(\"off\")\n\n    # Remove unused axes\n    for j in range(i + 1, len(axes)):\n        fig.delaxes(axes[j])\n\n    # Colorbar\n    cbar_ax = fig.add_axes([0.92, 0.15, 0.015, 0.7])\n    fig.colorbar(im, cax=cbar_ax, label='Grad-CAM Intensity')\n\n    plt.suptitle(\"Grad-CAM Overlay (Multiple Slices)\", fontsize=16)\n    plt.tight_layout(rect=[0, 0, 0.9, 0.95])\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import scipy.ndimage\n\ndef preprocess_nifti(nifti_path, target_shape=(64, 64, 64)):\n    nii = nib.load(nifti_path)\n    volume = nii.get_fdata().astype(np.float32)\n    volume = np.transpose(volume, (2, 1, 0))  # to Z, Y, X\n\n    # Downsample/resample volume to target_shape\n    factors = [t / s for t, s in zip(target_shape, volume.shape)]\n    volume = scipy.ndimage.zoom(volume, zoom=factors, order=1)  # linear interpolation\n\n    # Normalize\n    volume = (volume - np.min(volume)) / (np.max(volume) - np.min(volume) + 1e-5)\n\n    volume = np.expand_dims(volume, axis=0)  # Add channel\n    volume = np.expand_dims(volume, axis=0)  # Add batch\n    return torch.tensor(volume).to(device), volume[0, 0]\n\n# Load model\nmodel = DenseNet121model()\nmodel.load_state_dict(torch.load(\"checkpoints/model_best.pth\", map_location=device))\nmodel.to(device)\n    \n# Run inference + Grad-CAM\ninput_tensor, original_vol = preprocess_nifti(\"/kaggle/input/abdominal-nifti-0-100/1316_43094.nii\")\n\n# Ensure target shape is 3D only (D, H, W)\ntarget_shape_3d = original_vol.shape  # should be (D, H, W)\nif len(target_shape_3d) != 3:\n    target_shape_3d = target_shape_3d[-3:]  # slice off batch/channel if needed\n\n# Compute Grad-CAM with resizing\ncam = compute_gradcam(model, input_tensor, target_head=\"bowel\", target_shape=target_shape_3d)\n\n# Show overlay with enhanced visualization\nshow_gradcam_overlay(\n    cam,\n    original_vol,\n    slice_axis=0,       # 0 = axial; 1 = coronal; 2 = sagittal\n    alpha=0.5,          # blending of heatmap\n    threshold=0.5,      # threshold for contour detection\n    show_contours=True, # enable contour lines\n    slice_range=5       # show mid±5 slices\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}