{"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":"# ============================================================================\n#  CERVICAL FRACTURE DETECTION — FINAL PIPELINE\n#  Task:   Binary patient-level fracture detection (patient_overall: 0/1)\n#  Method: Multiple-Instance Learning (MIL) with MAX-POOLING across slices\n#  Stage 1: Cervical-slice localization (87-patient segmentation subset)\n#  Stage 2: Fracture classification on Stage-1-cropped full 2019-patient set\n# ============================================================================","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── cell 1:Install decompressors ───────────────────────────────\n!pip install -q python-gdcm pylibjpeg pylibjpeg-libjpeg pydicom","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-08-25T13:58:20.230775Z","iopub.execute_input":"2026-08-25T13:58:20.231212Z","iopub.status.idle":"2026-08-25T13:58:24.976117Z","shell.execute_reply.started":"2026-08-25T13:58:20.23118Z","shell.execute_reply":"2026-08-25T13:58:24.975395Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":" # ───  cell 2:Imports ──────────────────────────────────────────────────────\nimport os\nimport numpy as np\nimport pandas as pd\nimport pydicom as dicom\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\n\nimport tensorflow as tf\nfrom tensorflow.keras import layers, Model\nfrom tensorflow.keras.applications import MobileNetV2\nfrom tensorflow.keras.callbacks import EarlyStopping, ModelCheckpoint\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (roc_auc_score, roc_curve, confusion_matrix,\n                             ConfusionMatrixDisplay, classification_report)\n\nprint(\">>> VERSION_CHECK_FRACTURE_FINAL <<<\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:04:45.094553Z","iopub.execute_input":"2026-08-25T09:04:45.095298Z","iopub.status.idle":"2026-08-25T09:04:59.953954Z","shell.execute_reply.started":"2026-08-25T09:04:45.095262Z","shell.execute_reply":"2026-08-25T09:04:59.953322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── cell 3:Data loading (REAL fracture labels) ──────────────────────────\nbase_dir     = r'/kaggle/input/competitions/rsna-2022-cervical-spine-fracture-detection'\ntrain_images = os.path.join(base_dir, 'train_images')\n\n# Official study-level fracture labels\ntrain_data = pd.read_csv(os.path.join(base_dir, 'train.csv'))\n\n# Patients that actually have DICOM folders on disk\navailable = set(os.listdir(train_images))\npatients_df = train_data[train_data['StudyInstanceUID'].isin(available)].reset_index(drop=True)\n\nprint(f\"Patients with images available : {len(patients_df)}\")\nprint(f\"  Fractured (patient_overall=1): {patients_df['patient_overall'].sum()}\")\nprint(f\"  Healthy   (patient_overall=0): {(patients_df['patient_overall']==0).sum()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:05:12.624634Z","iopub.execute_input":"2026-08-25T09:05:12.625555Z","iopub.status.idle":"2026-08-25T09:05:12.680303Z","shell.execute_reply.started":"2026-08-25T09:05:12.625523Z","shell.execute_reply":"2026-08-25T09:05:12.679494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 3b: Load segmentation metadata (87-patient subset) ───────────────\nseg_meta = pd.read_csv(\n    '/kaggle/input/datasets/saikiranvarma/rsna-cervical-fracture-segmentation-metadata/meta_segmentation.csv'\n)\nseg_meta['slice_num'] = seg_meta['InstanceNumber'].astype(int)\nseg_patients = seg_meta['StudyInstanceUID'].unique()\nprint(f\"Segmentation patients available: {len(seg_patients)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:05:15.908907Z","iopub.execute_input":"2026-08-25T09:05:15.909244Z","iopub.status.idle":"2026-08-25T09:05:16.062406Z","shell.execute_reply.started":"2026-08-25T09:05:15.909217Z","shell.execute_reply":"2026-08-25T09:05:16.061454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 3c: Cervical-range helper (ground-truth, 87-patient only) ────────\ndef get_cervical_range(study_id, max_gap=3, trim_pct=8):\n    \"\"\"\n    Ground-truth cervical slice range for the 87 segmentation-annotated\n    patients only. Uses longest-contiguous-run detection + edge trimming\n    to reduce skull/thoracic boundary leakage.\n    \"\"\"\n    pat = seg_meta[seg_meta['StudyInstanceUID'] == study_id].copy()\n    pat['slice_num'] = pat['InstanceNumber'].astype(int)\n    cervical = pat[pat[['C1','C2','C3','C4','C5','C6','C7']].sum(axis=1) > 0]\n    if len(cervical) == 0:\n        return None\n\n    slices = sorted(cervical['slice_num'].unique())\n    runs, start, prev = [], slices[0], slices[0]\n    for s in slices[1:]:\n        if s - prev > max_gap:\n            runs.append((start, prev)); start = s\n        prev = s\n    runs.append((start, prev))\n    lo, hi = max(runs, key=lambda r: r[1] - r[0])\n\n    span = hi - lo\n    lo_trim = lo + int(span * trim_pct / 100)\n    hi_trim = hi - int(span * trim_pct / 100)\n    return lo_trim, hi_trim","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:05:25.124666Z","iopub.execute_input":"2026-08-25T09:05:25.125374Z","iopub.status.idle":"2026-08-25T09:05:25.131953Z","shell.execute_reply.started":"2026-08-25T09:05:25.125341Z","shell.execute_reply":"2026-08-25T09:05:25.131277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 3d (optional diagnostic): Verify slice alignment on 3 patients ───\nfor pid in seg_patients[:3]:\n    csv_slices = sorted(seg_meta[seg_meta['StudyInstanceUID'] == pid]['slice_num'].tolist())\n    dcm_files  = sorted([int(f.split('.')[0]) for f in os.listdir(os.path.join(train_images, pid))])\n    rng = get_cervical_range(pid)\n\n    print(f\"Patient ...{pid[-8:]}\")\n    print(f\"  CSV slices : {len(csv_slices)} (range {min(csv_slices)}-{max(csv_slices)})\")\n    print(f\"  DICOM files: {len(dcm_files)} (range {min(dcm_files)}-{max(dcm_files)})\")\n    print(f\"  Aligned (subset): {set(csv_slices).issubset(set(dcm_files))}\")\n    print(f\"  Cervical range: {rng}\")\n    print()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:05:28.566586Z","iopub.execute_input":"2026-08-25T09:05:28.567125Z","iopub.status.idle":"2026-08-25T09:05:28.641544Z","shell.execute_reply.started":"2026-08-25T09:05:28.567083Z","shell.execute_reply":"2026-08-25T09:05:28.640642Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ───  cell 4:DICOM helpers  ─────────────\ndef get_pixels_hu(img, size=128):\n    image = cv2.resize(img.pixel_array, (size, size), interpolation=cv2.INTER_NEAREST)\n    image = image.astype(np.int16)\n    image[image <= -1000] = 0\n    intercept = np.array(img.RescaleIntercept)\n    slope     = np.array(img.RescaleSlope)\n    image = (slope * image.astype(\"float64\")) + intercept\n    return image.astype(\"int16\")\n\ndef load_slice(path, size=128):\n    if not os.path.exists(path):\n        return None\n    img  = dicom.dcmread(path)\n    data = get_pixels_hu(img, size).astype(\"float64\")\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    return data.astype(\"float32\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:05:32.245082Z","iopub.execute_input":"2026-08-25T09:05:32.245825Z","iopub.status.idle":"2026-08-25T09:05:32.251247Z","shell.execute_reply.started":"2026-08-25T09:05:32.245795Z","shell.execute_reply":"2026-08-25T09:05:32.250508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ───  cell 5:Per-patient volume loader (fixed N evenly-spaced slices) ──────\nN_SLICES = 24     # slices sampled per patient\nIMG_SIZE = 128\n\ndef load_patient_volume(study_id, n_slices=N_SLICES, size=IMG_SIZE):\n    \"\"\"Return a (n_slices, size, size, 3) float32 tensor for one patient.\"\"\"\n    folder = os.path.join(train_images, study_id)\n    files  = os.listdir(folder)\n    # sort slices by numeric index to preserve anatomical order\n    files  = sorted(files, key=lambda x: int(x.split('.')[0]))\n\n    # evenly sample n_slices across the whole scan\n    if len(files) >= n_slices:\n        idx   = np.linspace(0, len(files) - 1, n_slices).astype(int)\n        files = [files[i] for i in idx]\n    else:  # short scan: pad by repeating the last slice\n        files = files + [files[-1]] * (n_slices - len(files))\n\n    vol = []\n    for f in files:\n        s = load_slice(os.path.join(folder, f), size)\n        if s is None:\n            s = np.zeros((size, size), dtype=\"float32\")\n        vol.append(s)\n\n    vol = np.stack(vol)                 # (n_slices, size, size)\n    vol = vol[..., np.newaxis]          # (n_slices, size, size, 1)\n    vol = np.repeat(vol, 3, axis=-1)    # (n_slices, size, size, 3)  <- for MobileNetV2\n    return vol.astype(\"float32\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T07:09:58.929977Z","iopub.execute_input":"2026-08-25T07:09:58.930605Z","iopub.status.idle":"2026-08-25T07:09:58.937894Z","shell.execute_reply.started":"2026-08-25T07:09:58.930574Z","shell.execute_reply":"2026-08-25T07:09:58.937121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 5b (exploratory): Whole-scan visualization ───────────────────────\npatient_id = '1.2.826.0.1.3680043.1363'\nvol = load_patient_volume(patient_id)\nfig, axes = plt.subplots(4, 6, figsize=(18, 12))\nfor i, ax in enumerate(axes.flat):\n    ax.imshow(vol[i, :, :, 0], cmap='bone')\n    ax.set_title(f\"Slice {i+1}\")\n    ax.axis('off')\nplt.suptitle(\"Whole-scan sampling (baseline) — contains skull/chest contamination\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T07:10:29.212899Z","iopub.execute_input":"2026-08-25T07:10:29.213602Z","iopub.status.idle":"2026-08-25T07:10:31.769234Z","shell.execute_reply.started":"2026-08-25T07:10:29.213567Z","shell.execute_reply":"2026-08-25T07:10:31.768172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 5c: Ground-truth cropped loader (87-patient exploratory only) ────\n# NOTE: This function is NOT used in the final training pipeline because it\n# only works for the 87 patients with segmentation masks. It is kept here\n# to demonstrate the cropping concept was validated visually before\n# building Stage 1 (which generalizes cropping to all 2019 patients).\ndef load_patient_volume_cropped(study_id, n_slices=24, size=128):\n    folder = os.path.join(train_images, study_id)\n    files = sorted(os.listdir(folder), key=lambda x: int(x.split('.')[0]))\n\n    rng = get_cervical_range(study_id)\n    if rng is not None:\n        cmin, cmax = rng\n        cropped = [f for f in files if cmin <= int(f.split('.')[0]) <= cmax]\n        if len(cropped) >= 3:\n            files = cropped\n\n    if len(files) >= n_slices:\n        idx = np.linspace(0, len(files)-1, n_slices).astype(int)\n        files = [files[i] for i in idx]\n    else:\n        files = files + [files[-1]] * (n_slices - len(files))\n\n    vol = []\n    for f in files:\n        s = load_slice(os.path.join(folder, f), size)\n        if s is None:\n            s = np.zeros((size, size), dtype=\"float32\")\n        vol.append(s)\n    vol = np.stack(vol)[..., np.newaxis]\n    vol = np.repeat(vol, 3, axis=-1)\n    return vol.astype(\"float32\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T07:10:37.354101Z","iopub.execute_input":"2026-08-25T07:10:37.354686Z","iopub.status.idle":"2026-08-25T07:10:37.363503Z","shell.execute_reply.started":"2026-08-25T07:10:37.354653Z","shell.execute_reply":"2026-08-25T07:10:37.362741Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 5d (exploratory): Ground-truth cropped visualization ─────────────\npatient_id = seg_patients[0]\nvol = load_patient_volume_cropped(patient_id)\nfig, axes = plt.subplots(4, 6, figsize=(18, 12))\nfor i, ax in enumerate(axes.flat):\n    ax.imshow(vol[i, :, :, 0], cmap='bone')\n    ax.set_title(f\"Cropped Slice {i+1}\")\n    ax.axis('off')\nplt.suptitle(\"Ground-truth cervical crop (87-patient exploratory validation)\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T07:10:48.693754Z","iopub.execute_input":"2026-08-25T07:10:48.694041Z","iopub.status.idle":"2026-08-25T07:10:51.43705Z","shell.execute_reply.started":"2026-08-25T07:10:48.694017Z","shell.execute_reply":"2026-08-25T07:10:51.435618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 6: Patient-level split (ALL 2019 patients, stratified) ──────────\ntrain_pat, test_pat = train_test_split(\n    patients_df,\n    test_size=0.20,\n    random_state=42,\n    stratify=patients_df['patient_overall']\n)\ntrain_pat = train_pat.reset_index(drop=True)\ntest_pat  = test_pat.reset_index(drop=True)\n\nprint(f\"Train patients: {len(train_pat)}  (fractured {train_pat['patient_overall'].sum()})\")\nprint(f\"Test  patients: {len(test_pat)}  (fractured {test_pat['patient_overall'].sum()})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:06:12.732568Z","iopub.execute_input":"2026-08-25T09:06:12.733079Z","iopub.status.idle":"2026-08-25T09:06:12.743283Z","shell.execute_reply.started":"2026-08-25T09:06:12.73305Z","shell.execute_reply":"2026-08-25T09:06:12.742536Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 7: Build WHOLE-SCAN tensors (baseline, all 2019 patients) ────────\ndef build_set(df):\n    X, y = [], []\n    for _, row in tqdm(df.iterrows(), total=len(df)):\n        X.append(load_patient_volume(row['StudyInstanceUID']))\n        y.append(row['patient_overall'])\n    return np.stack(X).astype(\"float32\"), np.array(y).astype(\"float32\")\n\nprint(\"Loading training volumes (whole-scan)...\")\nX_train, y_train = build_set(train_pat)\nprint(\"Loading test volumes (whole-scan)...\")\nX_test,  y_test  = build_set(test_pat)\n\nprint(\"X_train:\", X_train.shape, \"| y_train:\", y_train.shape)\nprint(\"X_test :\", X_test.shape,  \"| y_test :\", y_test.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T07:13:03.550578Z","iopub.execute_input":"2026-08-25T07:13:03.551226Z","iopub.status.idle":"2026-08-25T07:30:52.785108Z","shell.execute_reply.started":"2026-08-25T07:13:03.551197Z","shell.execute_reply":"2026-08-25T07:30:52.784083Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 8 : MIL model (validated, frozen encoder, crash-free) ────\ndef build_mil_maxpool(n_slices=N_SLICES, size=IMG_SIZE):\n    inp = layers.Input(shape=(n_slices, size, size, 3))\n\n    backbone = MobileNetV2(include_top=False, weights=\"imagenet\",\n                           input_shape=(size, size, 3), pooling=\"avg\")\n    backbone.trainable = False   # FROZEN \n\n    x = layers.TimeDistributed(backbone)(inp)\n    x = layers.TimeDistributed(layers.Dense(256, activation=\"relu\"))(x)\n    x = layers.TimeDistributed(layers.Dropout(0.4))(x)\n    x = layers.TimeDistributed(layers.Dense(64, activation=\"relu\"))(x)\n    x = layers.TimeDistributed(layers.Dense(1, activation=\"sigmoid\"))(x)\n    out = layers.GlobalMaxPooling1D()(x)\n    return Model(inp, out, name=\"MIL-MaxPool-Fracture\")\n\nfx_model = build_mil_maxpool()\nfx_model.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T07:33:08.750858Z","iopub.execute_input":"2026-08-25T07:33:08.751678Z","iopub.status.idle":"2026-08-25T07:33:09.606263Z","shell.execute_reply.started":"2026-08-25T07:33:08.751636Z","shell.execute_reply":"2026-08-25T07:33:09.605416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 9: Compile & train — WHOLE-SCAN BASELINE ─────────────────────────\nfx_model.compile(\n    loss=\"binary_crossentropy\",\n    optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4),\n    metrics=[tf.keras.metrics.BinaryAccuracy(name=\"acc\"),\n             tf.keras.metrics.AUC(name=\"auc\"),\n             tf.keras.metrics.Recall(name=\"recall\")]\n)\n\nearly = EarlyStopping(monitor=\"val_auc\", mode=\"max\", patience=12, restore_best_weights=True)\nckpt  = ModelCheckpoint(\"fracture_pilot.keras\", monitor=\"val_auc\", mode=\"max\", save_best_only=True)\n\nprint(\">>> Training WHOLE-SCAN binary patient-level fracture detector...\")\nhistory = fx_model.fit(\n    X_train, y_train,\n    validation_data=(X_test, y_test),\n    epochs=20,                 # set to 2 for a quick smoke test if desired\n    batch_size=2,\n    callbacks=[early, ckpt],\n    class_weight={0: 1.0, 1: 1.5},\n    verbose=1\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T07:35:10.477838Z","iopub.execute_input":"2026-08-25T07:35:10.478714Z","iopub.status.idle":"2026-08-25T07:52:10.41969Z","shell.execute_reply.started":"2026-08-25T07:35:10.478679Z","shell.execute_reply":"2026-08-25T07:52:10.418636Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 10: Evaluate WHOLE-SCAN baseline ─────────────────────────────────\ny_prob = fx_model.predict(X_test, batch_size=2, verbose=0).ravel()\ny_pred = (y_prob >= 0.5).astype(int)\n\nauc = roc_auc_score(y_test, y_prob)\ncm  = confusion_matrix(y_test, y_pred, labels=[0, 1])\ntn, fp, fn, tp = cm.ravel()\nsens = tp / (tp + fn) if (tp + fn) else 0.0\nspec = tn / (tn + fp) if (tn + fp) else 0.0\n\nprint(\"\\n\" + \"=\"*55)\nprint(\"  WHOLE-SCAN — PATIENT-LEVEL FRACTURE DETECTION RESULTS\")\nprint(\"=\"*55)\nprint(f\"  AUC          : {auc:.4f}\")\nprint(f\"  Sensitivity  : {sens:.4f}\")\nprint(f\"  Specificity  : {spec:.4f}\")\nprint(f\"  Confusion    : TN={tn}  FP={fp}  FN={fn}  TP={tp}\")\nprint(\"=\"*55)\nprint(classification_report(y_test, y_pred, target_names=[\"Healthy\", \"Fractured\"], digits=4))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T07:54:25.70877Z","iopub.execute_input":"2026-08-25T07:54:25.709681Z","iopub.status.idle":"2026-08-25T07:56:01.037007Z","shell.execute_reply.started":"2026-08-25T07:54:25.70964Z","shell.execute_reply":"2026-08-25T07:56:01.036063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL 11: Figures — WHOLE-SCAN baseline ────────────────────────────────\nfpr, tpr, _ = roc_curve(y_test, y_prob)\nplt.figure(figsize=(6, 5))\nplt.plot(fpr, tpr, lw=2, label=f\"AUC = {auc:.3f}\")\nplt.plot([0, 1], [0, 1], \"--\", color=\"gray\", label=\"Random\")\nplt.xlabel(\"False Positive Rate\"); plt.ylabel(\"True Positive Rate\")\nplt.title(\"Whole-Scan — Patient-Level Fracture Detection ROC\")\nplt.legend(loc=\"lower right\"); plt.grid(alpha=0.3)\nplt.tight_layout(); plt.savefig(\"fracture_roc.png\", dpi=300); plt.show()\n\nConfusionMatrixDisplay(cm, display_labels=[\"Healthy\", \"Fractured\"]).plot(cmap=\"Blues\", values_format=\"d\")\nplt.title(\"Whole-Scan — Confusion Matrix\")\nplt.tight_layout(); plt.savefig(\"fracture_confusion.png\", dpi=300); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T07:56:24.824246Z","iopub.execute_input":"2026-08-25T07:56:24.825453Z","iopub.status.idle":"2026-08-25T07:56:25.827428Z","shell.execute_reply.started":"2026-08-25T07:56:24.825388Z","shell.execute_reply":"2026-08-25T07:56:25.826612Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nprint(os.path.exists(\"fracture_pilot.keras\"))\nprint(os.path.getsize(\"fracture_pilot.keras\") if os.path.exists(\"fracture_pilot.keras\") else \"N/A\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:00:43.091646Z","iopub.execute_input":"2026-08-25T08:00:43.092522Z","iopub.status.idle":"2026-08-25T08:00:43.097719Z","shell.execute_reply.started":"2026-08-25T08:00:43.092488Z","shell.execute_reply":"2026-08-25T08:00:43.096729Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n#  STAGE 1 — CERVICAL-SLICE LOCALIZATION CLASSIFIER\n#  (Trained on the 87-patient segmentation subset; generalizes to all 2019)\n# ============================================================================","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── STAGE1 CELL A: Build per-slice cervical/non-cervical training data ────\nseg_meta['is_cervical'] = (seg_meta[['C1','C2','C3','C4','C5','C6','C7']].sum(axis=1) > 0).astype(int)\n\nprint(f\"Total labeled slices (87 patients): {len(seg_meta)}\")\nprint(f\"Cervical slices  : {seg_meta['is_cervical'].sum()}\")\nprint(f\"Non-cervical slices: {(seg_meta['is_cervical']==0).sum()}\")\n\ns1_train_patients, s1_test_patients = train_test_split(\n    seg_patients, test_size=0.20, random_state=42\n)\nprint(f\"Stage-1 train patients: {len(s1_train_patients)}  |  test patients: {len(s1_test_patients)}\")\n\ndef build_stage1_slice_set(patient_list, max_slices_per_patient=40):\n    X, y = [], []\n    for pid in tqdm(patient_list):\n        pat = seg_meta[seg_meta['StudyInstanceUID'] == pid].copy()\n        pat['slice_num'] = pat['InstanceNumber'].astype(int)\n\n        if len(pat) > max_slices_per_patient:\n            pat = pat.sort_values('slice_num').iloc[\n                np.linspace(0, len(pat)-1, max_slices_per_patient).astype(int)\n            ]\n\n        folder = os.path.join(train_images, pid)\n        for _, row in pat.iterrows():\n            fpath = os.path.join(folder, f\"{row['slice_num']}.dcm\")\n            img = load_slice(fpath, size=128)\n            if img is None:\n                continue\n            img_rgb = np.repeat(img[..., np.newaxis], 3, axis=-1)\n            X.append(img_rgb)\n            y.append(row['is_cervical'])\n\n    return np.stack(X).astype(\"float32\"), np.array(y).astype(\"float32\")\n\nprint(\"Building Stage-1 TRAIN slices...\")\nX_s1_train, y_s1_train = build_stage1_slice_set(s1_train_patients)\nprint(\"Building Stage-1 TEST slices...\")\nX_s1_test, y_s1_test = build_stage1_slice_set(s1_test_patients)\n\nprint(f\"X_s1_train: {X_s1_train.shape}  | cervical ratio: {y_s1_train.mean():.2f}\")\nprint(f\"X_s1_test : {X_s1_test.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:07:08.408084Z","iopub.execute_input":"2026-08-25T09:07:08.408849Z","iopub.status.idle":"2026-08-25T09:08:13.385487Z","shell.execute_reply.started":"2026-08-25T09:07:08.408821Z","shell.execute_reply":"2026-08-25T09:08:13.384695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── STAGE1 CELL B: Model architecture (frozen backbone, single slice) ─────\ndef build_stage1_classifier(size=128):\n    inp = layers.Input(shape=(size, size, 3))\n    backbone = MobileNetV2(include_top=False, weights=\"imagenet\",\n                           input_shape=(size, size, 3), pooling=\"avg\")\n    backbone.trainable = False\n\n    x = backbone(inp)\n    x = layers.Dense(128, activation=\"relu\")(x)\n    x = layers.Dropout(0.3)(x)\n    out = layers.Dense(1, activation=\"sigmoid\")(x)\n    return Model(inp, out, name=\"Stage1-Cervical-Slice-Classifier\")\n\nstage1_model = build_stage1_classifier()\nstage1_model.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:09:12.624961Z","iopub.execute_input":"2026-08-25T09:09:12.62561Z","iopub.status.idle":"2026-08-25T09:09:16.525097Z","shell.execute_reply.started":"2026-08-25T09:09:12.625581Z","shell.execute_reply":"2026-08-25T09:09:16.524344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── STAGE1 CELL C: Compile & train ─────────────────────────────────────────\nstage1_model.compile(\n    loss=\"binary_crossentropy\",\n    optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4),\n    metrics=[tf.keras.metrics.BinaryAccuracy(name=\"acc\"),\n             tf.keras.metrics.AUC(name=\"auc\")]\n)\n\ns1_early = EarlyStopping(monitor=\"val_auc\", mode=\"max\", patience=8, restore_best_weights=True)\n\nprint(\">>> Training Stage-1 cervical-slice classifier...\")\ns1_history = stage1_model.fit(\n    X_s1_train, y_s1_train,\n    validation_data=(X_s1_test, y_s1_test),\n    epochs=20,\n    batch_size=32,\n    callbacks=[s1_early],\n    verbose=1\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:09:50.844872Z","iopub.execute_input":"2026-08-25T09:09:50.845578Z","iopub.status.idle":"2026-08-25T09:11:02.997444Z","shell.execute_reply.started":"2026-08-25T09:09:50.845547Z","shell.execute_reply":"2026-08-25T09:11:02.996561Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── STAGE1 CELL D: Evaluate slice-level classifier ────────────────────────\ns1_probs = stage1_model.predict(X_s1_test, batch_size=32, verbose=0).ravel()\ns1_preds = (s1_probs >= 0.5).astype(int)\n\nprint(\"Stage-1 AUC:\", roc_auc_score(y_s1_test, s1_probs))\nprint(classification_report(y_s1_test, s1_preds, target_names=[\"Non-cervical\",\"Cervical\"], digits=4))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:11:05.022362Z","iopub.execute_input":"2026-08-25T09:11:05.022625Z","iopub.status.idle":"2026-08-25T09:11:14.200215Z","shell.execute_reply.started":"2026-08-25T09:11:05.022602Z","shell.execute_reply":"2026-08-25T09:11:14.199543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── STAGE1 CELL E (exploratory): Single-patient spot-check function ──────\ndef predict_cervical_range(study_id, stage1_model, sample_every=5, size=128):\n    \"\"\"Single-patient version — used only for spot-checking Stage 1's\n    predictions against ground truth before scaling to all 2019 patients.\"\"\"\n    folder = os.path.join(train_images, study_id)\n    files = sorted(os.listdir(folder), key=lambda x: int(x.split('.')[0]))\n    sampled_files = files[::sample_every]\n\n    imgs, slice_nums = [], []\n    for f in sampled_files:\n        img = load_slice(os.path.join(folder, f), size)\n        if img is None:\n            continue\n        imgs.append(np.repeat(img[..., np.newaxis], 3, axis=-1))\n        slice_nums.append(int(f.split('.')[0]))\n\n    if len(imgs) == 0:\n        return None\n\n    imgs = np.stack(imgs).astype(\"float32\")\n    probs = stage1_model.predict(imgs, batch_size=32, verbose=0).ravel()\n\n    cervical_slices = [s for s, p in zip(slice_nums, probs) if p >= 0.5]\n    if len(cervical_slices) == 0:\n        return None\n    return min(cervical_slices), max(cervical_slices)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:05:33.441794Z","iopub.execute_input":"2026-08-25T08:05:33.442825Z","iopub.status.idle":"2026-08-25T08:05:33.450522Z","shell.execute_reply.started":"2026-08-25T08:05:33.442789Z","shell.execute_reply":"2026-08-25T08:05:33.449787Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── STAGE1 CELL F (exploratory): Spot-check against ground truth ─────────\ntest_pid = s1_test_patients[0]\nrng = predict_cervical_range(test_pid, stage1_model)\nprint(f\"Patient: {test_pid}\")\nprint(f\"Predicted cervical range: {rng}\")\ngt_rng = get_cervical_range(test_pid)\nprint(f\"Ground-truth cervical range: {gt_rng}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:05:38.401035Z","iopub.execute_input":"2026-08-25T08:05:38.401707Z","iopub.status.idle":"2026-08-25T08:05:50.932056Z","shell.execute_reply.started":"2026-08-25T08:05:38.401652Z","shell.execute_reply":"2026-08-25T08:05:50.931064Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n#  STAGE 1 APPLIED — SCALE TO ALL 2019 PATIENTS (memory-safe, chunked)\n# ============================================================================","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell V2\nfrom collections import defaultdict\n\ndef predict_cervical_ranges_chunked(patient_ids, stage1_model, sample_every=8,\n                                      size=128, batch_size=128, chunk_size=100):\n    predicted_ranges = {}\n    n_chunks = (len(patient_ids) + chunk_size - 1) // chunk_size\n\n    for chunk_idx in tqdm(range(n_chunks), desc=\"Processing chunks\"):\n        chunk_patients = patient_ids[chunk_idx*chunk_size : (chunk_idx+1)*chunk_size]\n        chunk_imgs, chunk_meta = [], []\n\n        for pid in chunk_patients:\n            folder = os.path.join(train_images, pid)\n            try:\n                files = sorted(os.listdir(folder), key=lambda x: int(x.split('.')[0]))\n            except FileNotFoundError:\n                predicted_ranges[pid] = None\n                continue\n            for f in files[::sample_every]:\n                img = load_slice(os.path.join(folder, f), size)\n                if img is None:\n                    continue\n                chunk_imgs.append(np.repeat(img[..., np.newaxis], 3, axis=-1))\n                chunk_meta.append((pid, int(f.split('.')[0])))\n\n        if len(chunk_imgs) == 0:\n            continue\n\n        chunk_imgs = np.stack(chunk_imgs).astype(\"float32\")\n        probs = stage1_model.predict(chunk_imgs, batch_size=batch_size, verbose=0).ravel()\n\n        patient_slices = defaultdict(list)\n        for (pid, slice_num), prob in zip(chunk_meta, probs):\n            if prob >= 0.5:\n                patient_slices[pid].append(slice_num)\n\n        for pid in chunk_patients:\n            slices = patient_slices.get(pid, [])\n            predicted_ranges[pid] = (min(slices), max(slices)) if slices else None\n\n        del chunk_imgs, chunk_meta\n\n    return predicted_ranges\n\n\nprint(\">>> Predicting cervical ranges for ALL 2019 patients...\")\npredicted_ranges = predict_cervical_ranges_chunked(\n    patients_df['StudyInstanceUID'].tolist(),\n    stage1_model,\n    sample_every=8,\n    batch_size=128,\n    chunk_size=100\n)\n\nimport pickle\nwith open(\"predicted_cervical_ranges.pkl\", \"wb\") as f:\n    pickle.dump(predicted_ranges, f)\nprint(f\"Saved. Coverage: {sum(1 for v in predicted_ranges.values() if v)} / {len(predicted_ranges)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:13:33.752487Z","iopub.execute_input":"2026-08-25T09:13:33.752784Z","iopub.status.idle":"2026-08-25T09:45:07.184695Z","shell.execute_reply.started":"2026-08-25T09:13:33.752759Z","shell.execute_reply":"2026-08-25T09:45:07.18383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n#  STAGE 2 — FRACTURE DETECTION ON STAGE-1-CROPPED FULL DATASET\n# ============================================================================","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL v2-A: Stage-1-guided volume loader + dataset builder ─────────────\ndef load_patient_volume_v2(study_id, n_slices=24, size=128):\n    \"\"\"Uses Stage-1's PREDICTED cervical range; falls back to whole-scan\n    if no prediction is available for a given patient.\"\"\"\n    folder = os.path.join(train_images, study_id)\n    files = sorted(os.listdir(folder), key=lambda x: int(x.split('.')[0]))\n\n    rng = predicted_ranges.get(study_id)\n    if rng is not None:\n        cmin, cmax = rng\n        cropped = [f for f in files if cmin <= int(f.split('.')[0]) <= cmax]\n        if len(cropped) >= 3:\n            files = cropped\n\n    if len(files) >= n_slices:\n        idx = np.linspace(0, len(files)-1, n_slices).astype(int)\n        files = [files[i] for i in idx]\n    else:\n        files = files + [files[-1]] * (n_slices - len(files))\n\n    vol = []\n    for f in files:\n        s = load_slice(os.path.join(folder, f), size)\n        if s is None:\n            s = np.zeros((size, size), dtype=\"float32\")\n        vol.append(s)\n    vol = np.stack(vol)[..., np.newaxis]\n    vol = np.repeat(vol, 3, axis=-1)\n    return vol.astype(\"float32\")\n\n\ndef build_set_v2(df):\n    X, y = [], []\n    for _, row in tqdm(df.iterrows(), total=len(df)):\n        X.append(load_patient_volume_v2(row['StudyInstanceUID']))\n        y.append(row['patient_overall'])\n    return np.stack(X).astype(\"float32\"), np.array(y).astype(\"float32\")\n\nprint(\"load_patient_volume_v2 and build_set_v2 defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T09:46:27.391985Z","iopub.execute_input":"2026-08-25T09:46:27.392425Z","iopub.status.idle":"2026-08-25T09:46:27.40261Z","shell.execute_reply.started":"2026-08-25T09:46:27.392397Z","shell.execute_reply":"2026-08-25T09:46:27.401955Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL v2-B: Build STAGE-1-CROPPED tensors (all 2019, same split) ───────\nprint(\"Loading STAGE-1-CROPPED volumes for ALL 2019 patients...\")\nX_train_v2, y_train_v2 = build_set_v2(train_pat)   # reuses SAME split as baseline\nX_test_v2,  y_test_v2  = build_set_v2(test_pat)\n\nprint(\"X_train_v2:\", X_train_v2.shape, \"| y_train_v2:\", y_train_v2.shape)\nprint(\"X_test_v2 :\", X_test_v2.shape,  \"| y_test_v2 :\", y_test_v2.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-25T08:46:43.292576Z","iopub.execute_input":"2026-08-25T08:46:43.293483Z","execution_failed":"2026-08-25T08:58:36.709Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL v2-C: Train + evaluate fracture model on cropped volumes ─────────\nfx_model_v2 = build_mil_maxpool()\n\nfx_model_v2.compile(\n    loss=\"binary_crossentropy\",\n    optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4),\n    metrics=[tf.keras.metrics.BinaryAccuracy(name=\"acc\"),\n             tf.keras.metrics.AUC(name=\"auc\"),\n             tf.keras.metrics.Recall(name=\"recall\")]\n)\n\nearly_v2 = EarlyStopping(monitor=\"val_auc\", mode=\"max\", patience=12, restore_best_weights=True)\nckpt_v2  = ModelCheckpoint(\"fracture_pilot_v2.keras\", monitor=\"val_auc\", mode=\"max\", save_best_only=True)\n\nprint(\">>> Training on STAGE-1-CROPPED volumes (all 2019 patients)...\")\nhistory_v2 = fx_model_v2.fit(\n    X_train_v2, y_train_v2,\n    validation_data=(X_test_v2, y_test_v2),\n    epochs=40,\n    batch_size=2,\n    callbacks=[early_v2, ckpt_v2],\n    class_weight={0: 1.0, 1: 1.5},\n    verbose=1\n)\n\ny_prob_v2 = fx_model_v2.predict(X_test_v2, batch_size=2, verbose=0).ravel()\ny_pred_v2 = (y_prob_v2 >= 0.5).astype(int)\n\nauc_v2 = roc_auc_score(y_test_v2, y_prob_v2)\ncm_v2  = confusion_matrix(y_test_v2, y_pred_v2, labels=[0, 1])\ntn, fp, fn, tp = cm_v2.ravel()\nsens_v2 = tp / (tp + fn) if (tp + fn) else 0.0\nspec_v2 = tn / (tn + fp) if (tn + fp) else 0.0\n\nprint(\"=\"*55)\nprint(\"  STAGE-1-CROPPED — PATIENT-LEVEL FRACTURE RESULTS\")\nprint(\"=\"*55)\nprint(f\"  AUC          : {auc_v2:.4f}   (whole-scan baseline was {auc:.4f})\")\nprint(f\"  Sensitivity  : {sens_v2:.4f}\")\nprint(f\"  Specificity  : {spec_v2:.4f}\")\nprint(f\"  Confusion    : TN={tn}  FP={fp}  FN={fn}  TP={tp}\")\nprint(\"=\"*55)\nprint(classification_report(y_test_v2, y_pred_v2, target_names=[\"Healthy\",\"Fractured\"], digits=4))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── CELL v2-D: Figures — Stage-1-cropped result ───────────────────────────\nfpr_v2, tpr_v2, _ = roc_curve(y_test_v2, y_prob_v2)\nplt.figure(figsize=(6, 5))\nplt.plot(fpr_v2, tpr_v2, lw=2, label=f\"AUC = {auc_v2:.3f}\")\nplt.plot([0, 1], [0, 1], \"--\", color=\"gray\", label=\"Random\")\nplt.xlabel(\"False Positive Rate\"); plt.ylabel(\"True Positive Rate\")\nplt.title(\"Stage-1-Cropped — Patient-Level Fracture Detection ROC\")\nplt.legend(loc=\"lower right\"); plt.grid(alpha=0.3)\nplt.tight_layout(); plt.savefig(\"fracture_roc_v2.png\", dpi=300); plt.show()\n\nConfusionMatrixDisplay(cm_v2, display_labels=[\"Healthy\", \"Fractured\"]).plot(cmap=\"Blues\", values_format=\"d\")\nplt.title(\"Stage-1-Cropped — Confusion Matrix\")\nplt.tight_layout(); plt.savefig(\"fracture_confusion_v2.png\", dpi=300); plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ─── FINAL: Side-by-side summary for the paper ─────────────────────────────\nprint(\"\\n\" + \"=\"*60)\nprint(\"  FINAL COMPARISON — WHOLE-SCAN vs STAGE-1-CROPPED\")\nprint(\"=\"*60)\nprint(f\"  Whole-scan baseline    : AUC = {auc:.4f}\")\nprint(f\"  Stage-1-cropped (final): AUC = {auc_v2:.4f}\")\nprint(f\"  Delta                  : {(auc_v2 - auc):+.4f}\")\nprint(\"=\"*60)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}