{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":24800,"datasetId":1042002,"databundleVersionId":1831594,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":4800870,"datasetId":2727590,"databundleVersionId":4864291},{"sourceType":"datasetVersion","sourceId":5297273,"datasetId":3073790,"databundleVersionId":5370374},{"sourceType":"datasetVersion","sourceId":11465621,"datasetId":7184979,"databundleVersionId":11908499},{"sourceType":"datasetVersion","sourceId":15860099,"datasetId":10167932,"databundleVersionId":16811994},{"sourceType":"datasetVersion","sourceId":5605381,"datasetId":2642145,"databundleVersionId":5680458},{"sourceType":"datasetVersion","sourceId":954111,"datasetId":518432,"databundleVersionId":982056},{"sourceType":"datasetVersion","sourceId":954197,"datasetId":518486,"databundleVersionId":982144},{"sourceType":"datasetVersion","sourceId":9864328,"datasetId":6054580,"databundleVersionId":10116567},{"sourceType":"datasetVersion","sourceId":2253105,"datasetId":1353821,"databundleVersionId":2294052},{"sourceType":"datasetVersion","sourceId":13497428,"datasetId":8569815,"databundleVersionId":14218569},{"sourceType":"datasetVersion","sourceId":2332556,"datasetId":1407957,"databundleVersionId":2374115},{"sourceType":"datasetVersion","sourceId":13541937,"datasetId":8599837,"databundleVersionId":14267433},{"sourceType":"datasetVersion","sourceId":15880601,"datasetId":10182053,"databundleVersionId":16834092}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# AI-Powered Lung Disease Detection — YOLOX-s @ 640 (v11)\n\nEnd-to-end detection of **pulmonary tumors (X-ray), tuberculosis, and pneumonia** with YOLOX-s at 640×640 resolution.\n\n## What's new in v11 (cumulative v5 → v11)\n\n- **SIIM-FISABIO-RSNA COVID-19 dataset** (v11) — ~6k radiologist-drawn opacity bboxes added to pneumonia. Same RSNA annotation philosophy as the existing RSNA Pneumonia source.\n- **JSRT dataset** (v11) — small but pristine 154 confirmed-nodule bboxes added to tumor_xray. Boxes constructed from center+size in mm using JSRT's 0.175mm/pixel scale.\n- **Pneumonia cleanup** (v11) — Lung Opacity dropped from VinDr/VinBigData/cxray14 mappings (label noise was capping pneumonia AP at ~0.19). Consolidation + Infiltration retained.\n- **cxray14 tumor_xray dropped** (v11) — NIH-derived nodule labels are ~25-30% noisy; superseded by Node21 + JSRT clean sources.\n- **Aggressive multi-radiologist merge** (v11) — VinDr/VinBigData dedup IoU lowered from 0.7 to 0.5. Close-by boxes from different radiologists now merge to a single consensus box.\n- **Excessive-bbox filter** (v11) — drops entire images with too many boxes per class (tumor>8, pneumonia>12, TB>15) — typically annotation errors or pathological metastasis cases that don't generalize.\n\n## What's new in v9 (cumulative v5 → v9)\n\n- **Expanded pneumonia mapping** (v9) — Consolidation, Infiltration, Lung Opacity now map to pneumonia across VinDr, VinBigData, ChestX-Det, cxray14. Adds ~5–8k pneumonia bboxes; some label noise tradeoff.\n- **CXR Lung (Node21) added** (v9) — ~1,476 high-quality nodule bboxes added to tumor_xray.\n- **LIDC re-enabled** (v9) — discovered + parsed when TARGET_MODALITY in ('ct', 'all'); skipped silently for xray runs.\n- **Auxiliary image-level classification head** (v9) — multi-label BCE on global-pooled deepest backbone feature, weight 0.3. Provides calibrated 'Pneumonia: 87%' image-level probability for the UI.\n\n## What's new in v8 (cumulative v5 → v8)\n\n- **640×640 input, bs=16, grad_accum=1** (v6) — 2× throughput on T4 vs the prior 768@bs=8 config, real BN statistics, no gradient accumulation needed\n- **Differential LR** (v6) — backbone trains at 0.1× head LR via AdamW param groups. Replaces the freeze/unfreeze scheme that was collapsing mAP at epoch 4 (0.092 → 0.015, never recovered).\n- **No backbone freeze** (v6) — `BACKBONE_FREEZE_EPOCHS = 0`. Differential LR is safer than freeze+unfreeze on a small medical-imaging dataset.\n- **Mosaic / MixUp disabled** (v6) — designed for COCO-style large multi-class data. For 3 structurally-similar X-ray classes on ~10k images they add noise faster than signal. Copy-paste stays on to boost rare TB.\n- **Tightened pneumonia/tumor labels** (v7) — dropped \"Consolidation\", \"Lung Opacity\", \"Infiltration\" from VinDr / VinBigData / ChestX-Det mappings. These are non-specific radiological findings, not diagnoses, and were corrupting the class signal. Pneumonia is now RSNA-dominated (radiologist-confirmed boxes); tumor_xray is Nodule/Mass only.\n- **LIDC-IDRI dropped** (v8) — it's CT-only (produces `tumor_ct`) and was already filtered out by the xray-only modality split. Skipping the multi-minute DICOM+XML parse on every run just saves startup time.\n- **Resume-memory fix** (v4) — free checkpoint dict + `gc.collect()` + `torch.cuda.empty_cache()` before the DataLoader forks workers. Prevents the RAM-inheritance OOM that only triggered on session resume.\n- **DataLoader hardening** (v3) — `num_workers=2`, `persistent_workers=True`, `prefetch_factor=2`. Stable on Kaggle's 30 GB RAM at 640² with bs=16.\n- **Aggressive inverse-frequency sampling** (v3) — rarest-label-per-image picker instead of `.first()`, plus clamp at 8× the min weight. Lifts TB from ~3 % to ~30 % expected sampled class frequency. A diagnostic line prints the post-sampling class mix.\n\n## Expected T4 training time\n\nAt 640² with YOLOX-s and bs=16 on a T4: **~12–15 min per training epoch** + ~2 min validation + ~1 min mAP. Call it ~15 min/epoch total. Kaggle's 12-hour GPU limit gives you ~45 epochs per session — a full 60-epoch run fits with checkpoint resume.\n\n## Realistic mAP expectations (graduation-report context)\n\nPublished benchmarks on similar datasets cap around: VinBigData Kaggle winner ~0.314, RSNA ~0.26, TBX11K ~0.40. A 3-class unified dataset with mixed annotation styles sits in the **0.20–0.30 mAP** band when executed well. **Pair detection mAP with image-level classification AUC** (~0.85–0.90 is realistic) — two complementary metrics tell a stronger story than either alone.\n\n## What v8 keeps from earlier versions\n\n- **Lung-field masking** via `lungmask` U-Net — kills shortcut learning from text overlays / EKG leads / scanner borders\n- **Percentile-based intensity normalization** (0.5–99.5 pct clip) + single shared CLAHE\n- **Shared `preprocess_image()`** across train / val / inference / TTA — no distribution drift\n- **Train/test perceptual-hash leak removal** (was 116 leaks silently kept in v1)\n- **Albumentations** with medical-aware transforms (ElasticTransform, GaussNoise, RandomGamma, GridDropout)\n- **Copy-paste augmentation** for rare TB and tumor_xray crops\n- **EMA of model weights** (decay 0.9998) — used for validation + mAP\n- **Checkpoint resume** every 5 epochs (fixed in v4 so it doesn't OOM on reload)\n- **TTA inference** with weighted box fusion\n","metadata":{}},{"cell_type":"markdown","source":"## 1. Environment Setup","metadata":{}},{"cell_type":"code","source":"%%capture\n# PyTorch with CUDA\n!pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121\n\n# Core viz + data\n!pip install matplotlib pandas pillow torchtnt==0.2.0 tqdm opencv-python seaborn\n\n# Data formats\n!pip install tabulate pyarrow fastparquet\n\n# Visualization utilities\n!pip install distinctipy\n\n# COCO tools\n!pip install pycocotools\n\n# Perceptual hashing for duplicate detection\n!pip install imagehash\n\n# Statistics & ML\n!pip install scipy scikit-learn\n\n# YOLOX utilities (kept as per requirement)\n!pip install cjm_pandas_utils cjm_psl_utils cjm_pil_utils cjm_pytorch_utils cjm_yolox_pytorch cjm_torchvision_tfms\n!pip install torchmetrics\n\n# NEW: Albumentations\n!pip install albumentations==1.4.0\n\n# NEW: lungmask (pretrained U-Net for lung segmentation)\n!pip install lungmask\n\n# SimpleITK for proper lungmask input format\n!pip install SimpleITK\n\n# NEW: weighted-boxes-fusion for TTA\n!pip install ensemble-boxes\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T16:55:33.678633Z","iopub.execute_input":"2026-04-24T16:55:33.678903Z","iopub.status.idle":"2026-04-24T16:56:20.149375Z","shell.execute_reply.started":"2026-04-24T16:55:33.678868Z","shell.execute_reply":"2026-04-24T16:56:20.148537Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Standard library\nimport datetime, hashlib, json, math, multiprocessing, os, random, warnings\nfrom functools import partial\nfrom glob import glob\nfrom pathlib import Path\nfrom collections import Counter, defaultdict\nwarnings.filterwarnings(\"ignore\")\n\n# CJM utilities\nfrom cjm_psl_utils.core import download_file, file_extract\nfrom cjm_pil_utils.core import resize_img, get_img_files, stack_imgs\nfrom cjm_pytorch_utils.core import tensor_to_pil, get_torch_device, set_seed, denorm_img_tensor\nfrom cjm_pandas_utils.core import markdown_to_pandas, convert_to_numeric, convert_to_string\nfrom cjm_torchvision_tfms.core import ResizeMax, PadSquare\n\n# YOLOX\nfrom cjm_yolox_pytorch.model import build_model, MODEL_CFGS, NORM_STATS\nfrom cjm_yolox_pytorch.utils import generate_output_grids\nfrom cjm_yolox_pytorch.loss import YOLOXLoss\nfrom cjm_yolox_pytorch.inference import YOLOXInferenceWrapper\n\n# Visualization\nfrom distinctipy import distinctipy\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nimport seaborn as sns\n\n# Scientific computing\nimport numpy as np\nimport cv2\nfrom scipy import stats\nfrom scipy.stats import chi2_contingency\nimport imagehash\n\n# Data handling\nimport pandas as pd\npd.set_option(\"max_colwidth\", None, \"display.max_rows\", None, \"display.max_columns\", None)\n\n# PIL\nfrom PIL import Image, ImageFilter\n\n# PyTorch\nimport torch\nfrom torch.amp import autocast\nfrom torch.cuda.amp import GradScaler\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchtnt.utils import get_module_summary\n\n# Torchvision\nimport torchvision\ntorchvision.disable_beta_transforms_warning()\nfrom torchvision.tv_tensors import BoundingBoxes\nfrom torchvision.utils import draw_bounding_boxes\nimport torchvision.transforms.v2 as transforms\n\n# Progress & ML\nfrom tqdm.auto import tqdm\n# Quiet tqdm for Kaggle logs: ANSI escapes become spam in their log viewer.\n# Refresh every 10s instead of every iteration, ASCII-only, no trailing line.\nimport sys as _sys\nfrom functools import partialmethod as _partialmethod\n_tqdm_defaults = dict(mininterval=10.0, maxinterval=30.0, ascii=True,\n                      leave=False, dynamic_ncols=False, ncols=80)\n# Disable entirely when stdout is not a TTY (i.e. captured by Kaggle)\nif not getattr(_sys.stdout, \"isatty\", lambda: False)():\n    _tqdm_defaults[\"disable\"] = True\ntqdm.__init__ = _partialmethod(tqdm.__init__, **_tqdm_defaults)\nfrom sklearn.model_selection import train_test_split\nfrom torchmetrics.detection import MeanAveragePrecision\n\n# NEW: Albumentations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# NEW: lungmask\nfrom lungmask import LMInferer\n\n# NEW: weighted box fusion\nfrom ensemble_boxes import weighted_boxes_fusion\n\nprint(\"All dependencies loaded ✓\")\n\n\n# ── DICOM → PNG conversion helper ──────────────────────────────────────────\n_dim_cache = {}\nDICOM_TARGET_SIZE = 1024  # resize DICOM-derived PNGs to this longest side to save disk\n\ndef dicom_to_png(dcm_path, cache_dir=\"/tmp/dicom_png_cache\", target_size=DICOM_TARGET_SIZE):\n    \"\"\"Read a DICOM, convert to 8-bit PNG resized to target_size (longest side), cache it.\n    Returns (png_path, new_w, new_h, orig_w, orig_h) or None.\n    Caller can use orig_w/orig_h to scale bbox coords from native-DICOM space to PNG space.\n    \"\"\"\n    import pydicom\n    os.makedirs(cache_dir, exist_ok=True)\n    stem = os.path.splitext(os.path.basename(dcm_path))[0]\n    png_path = os.path.join(cache_dir, stem + \".png\")\n\n    # Cache hit — need to recover orig dims from sidecar JSON or re-read header\n    if os.path.exists(png_path):\n        sidecar = png_path + \".meta.json\"\n        if png_path in _dim_cache:\n            new_w, new_h, orig_w, orig_h = _dim_cache[png_path]\n            return png_path, new_w, new_h, orig_w, orig_h\n        if os.path.exists(sidecar):\n            try:\n                with open(sidecar) as f:\n                    m = json.load(f)\n                new_w, new_h = m[\"new_w\"], m[\"new_h\"]\n                orig_w, orig_h = m[\"orig_w\"], m[\"orig_h\"]\n                _dim_cache[png_path] = (new_w, new_h, orig_w, orig_h)\n                return png_path, new_w, new_h, orig_w, orig_h\n            except Exception:\n                pass\n        # fallback: read png dims and original DICOM header for native dims\n        img = cv2.imread(png_path)\n        if img is None:\n            return None\n        new_h, new_w = img.shape[:2]\n        try:\n            ds = pydicom.dcmread(dcm_path, stop_before_pixels=True)\n            orig_w = int(getattr(ds, \"Columns\", new_w))\n            orig_h = int(getattr(ds, \"Rows\", new_h))\n        except Exception:\n            orig_w, orig_h = new_w, new_h\n        _dim_cache[png_path] = (new_w, new_h, orig_w, orig_h)\n        return png_path, new_w, new_h, orig_w, orig_h\n\n    try:\n        ds = pydicom.dcmread(dcm_path)\n        pixel = ds.pixel_array.astype(np.float64)\n    except Exception:\n        return None\n\n    orig_h, orig_w = pixel.shape[:2]\n\n    # Apply CT windowing if available\n    if hasattr(ds, \"WindowCenter\"):\n        wc = ds.WindowCenter\n        ww = ds.WindowWidth\n        if isinstance(wc, pydicom.multival.MultiValue):\n            wc, ww = float(wc[0]), float(ww[0])\n        else:\n            wc, ww = float(wc), float(ww)\n        lo, hi = wc - ww / 2, wc + ww / 2\n        pixel = np.clip(pixel, lo, hi)\n\n    # Normalize to 0-255\n    pmin, pmax = pixel.min(), pixel.max()\n    if pmax - pmin > 0:\n        pixel = (pixel - pmin) / (pmax - pmin) * 255.0\n    pixel = pixel.astype(np.uint8)\n\n    # Resize longest side to target_size (preserves aspect ratio, saves disk)\n    longest = max(orig_w, orig_h)\n    if longest > target_size:\n        scale = target_size / longest\n        new_w = int(round(orig_w * scale))\n        new_h = int(round(orig_h * scale))\n        pixel = cv2.resize(pixel, (new_w, new_h), interpolation=cv2.INTER_AREA)\n    else:\n        new_w, new_h = orig_w, orig_h\n\n    cv2.imwrite(png_path, pixel, [cv2.IMWRITE_PNG_COMPRESSION, 6])\n\n    # Sidecar so we can recover native dims across process restarts\n    try:\n        with open(png_path + \".meta.json\", \"w\") as f:\n            json.dump({\"new_w\": new_w, \"new_h\": new_h,\n                       \"orig_w\": int(orig_w), \"orig_h\": int(orig_h)}, f)\n    except Exception:\n        pass\n\n    _dim_cache[png_path] = (new_w, new_h, int(orig_w), int(orig_h))\n    return png_path, new_w, new_h, int(orig_w), int(orig_h)\n\ndef get_image_dims(path):\n    \"\"\"Get (width, height) for any image. Uses cache.\"\"\"\n    if path in _dim_cache:\n        v = _dim_cache[path]\n        return (v[0], v[1])\n    if path.lower().endswith((\".dcm\", \".dicom\")):\n        result = dicom_to_png(path)\n        return (result[1], result[2]) if result else (None, None)\n    try:\n        with Image.open(path) as img:\n            w, h = img.size\n        _dim_cache[path] = (w, h)\n        return w, h\n    except Exception:\n        return None, None","metadata":{"execution":{"iopub.status.busy":"2026-04-24T16:56:20.151531Z","iopub.execute_input":"2026-04-24T16:56:20.151717Z","iopub.status.idle":"2026-04-24T16:56:20.178925Z","shell.execute_reply.started":"2026-04-24T16:56:20.151692Z","shell.execute_reply":"2026-04-24T16:56:20.178348Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Configuration & Dataset Discovery\n\n**Classes (3):** `tumor_xray`, `tuberculosis`, `pneumonia`. (X-ray modality only. `tumor_ct` / LIDC dropped in v8.)\n\n**v8 config (tuned for T4 16 GB + YOLOX-s @ 640):**\n- `train_sz = 640`, `bs = 16`, `grad_accum_steps = 1` — real batch 16, real BN statistics\n- `model_type = \"yolox_s\"`\n- `lr = 5e-4` head / `5e-5` backbone — differential LR via AdamW param groups\n- `mosaic_prob = 0.0`, `mixup_prob = 0.0`, `copy_paste_prob = 0.5`\n- `multi_scale = False` (variance with no data gain at this resolution)\n- `BACKBONE_FREEZE_EPOCHS = 0` — no freeze scheme; differential LR handles the transfer-learning adaptation safely\n- `epochs = 60`, `es_patience = 15`\n- **EMA decay = 0.9998** — slow enough to average over the whole run\n\n**Memory budget on T4 16 GB at 640² + YOLOX-s with bs=16:** ~7–9 GB peak, steady throughout training (no mosaic spikes).\n\n**Dataset coverage after v7 label tightening + v8 LIDC drop:**\n- `pneumonia` ~10 k bboxes, RSNA-dominated (clean, radiologist-confirmed)\n- `tumor_xray` ~4.9 k bboxes, Nodule/Mass only across VinDr / VinBigData / ChestX-Det\n- `tuberculosis` ~1.15 k bboxes, TBX11K (rare class, copy-paste boosted)\n","metadata":{}},{"cell_type":"code","source":"seed = 42\nset_seed(seed)\n\n# Auxiliary image-level classification head (multi-label per class).\n# Stabilizes pneumonia and gives the UI a calibrated 'Pneumonia: 87%' number.\nUSE_AUX_CLS_HEAD = True\nAUX_CLS_LOSS_WEIGHT = 0.3\ndevice = get_torch_device()\ndtype = torch.float32\n\n# Training config — tuned for T4 16GB + YOLOX-m @ 1024\ntrain_sz          = 640          # 640 runs 2x faster than 768, still plenty for X-ray\nbs                = 16           # bs=16 at 640 fits T4; real BN stats, no grad accum needed\ngrad_accum_steps  = 1            # effective batch = 16, no accumulation\nepochs            = 60\nwarmup_epochs     = 3\nlr                = 5e-4         # head LR; backbone gets 0.1x this via param groups\nweight_decay      = 5e-4\n\n# Augmentation schedule\nmosaic_epochs     = epochs - 15\n# Mosaic/MixUp OFF for medical imaging — they were designed for COCO (80 classes,\n# 100k+ images). For 3 structurally-similar X-ray classes on 10k images, they add\n# noise faster than signal. Copy-paste stays on (boosts rare TB class).\nmosaic_prob       = 0.0\ncopy_paste_prob   = 0.5\nmixup_prob        = 0.0\nmulti_scale       = False        # variance with no data gain at this scale\nmulti_scale_range = (576, 704)   # unused now; kept so downstream refs don't crash\n\n# EMA (Exponential Moving Average of weights)\nuse_ema           = True\nema_decay         = 0.9998\n\n# Checkpoint every N epochs (for session resumption)\ncheckpoint_every  = 5\n\n# Early stopping (on mAP, not val_loss)\nes_patience  = 15\nes_min_delta = 1e-3\n\n# Inference / TTA\nuse_tta      = True\nconf_thresh  = 0.25\niou_thresh   = 0.5\n\n# ── Two-model split: train one modality at a time ─────────────────────────\n# Set this to \"xray\" or \"ct\"; run the whole notebook once per value.\n# Each run writes to its own ckpt dir so nothing collides.\nTARGET_MODALITY = \"xray\"   # \"xray\" | \"ct\" | \"all\" (legacy joint model)\nassert TARGET_MODALITY in (\"xray\", \"ct\", \"all\")\n\n# Dirs — re-use lung-crop cache from v2 if it exists\nproject_dir    = Path(\"/kaggle/working/yolox_lung_disease_v2\")\nproject_dir.mkdir(parents=True, exist_ok=True)\nlung_cache_dir = project_dir / \"lung_crops\"\nlung_cache_dir.mkdir(parents=True, exist_ok=True)\n\n# Optional: re-use lung_crops/ from a previous run, uploaded as a Kaggle Dataset.\n# Cell 39 checks this path first; if crop_manifest.json exists there, cropping is skipped.\nexternal_lung_crop_dir = Path(\"/kaggle/input/datasets/mickgt/dataset-crops/yolox_lung_disease_v2/lung_crops\")\n\n# Separate ckpt dir per modality so xray/ct runs don't collide.\nckpt_dir = project_dir / f\"ckpts_v6_yoloxs_768_{TARGET_MODALITY}\"\nckpt_dir.mkdir(parents=True, exist_ok=True)\n\nSEARCH_ROOTS = [\"/kaggle/input\", \"/kaggle/working\"]\n\n\n# ── Dataset cache: skip data loading on subsequent runs ─────────────────────\n# Saves the final processed `df` and `norm_stats` after first build, then\n# short-circuits all data loading + DICOM conversion + lung cropping + norm\n# stats computation on every subsequent kernel restart.\n#\n# Workflow:\n#   First run:  builds cache, saves to /kaggle/working/yolox_lung_disease_v2/cache/\n#   To reuse:   commit notebook → attach the output as input dataset on next run\n#               OR upload the cache files as their own Kaggle dataset.\nDATASET_VERSION = \"v11\"          # bump if data pipeline semantics change\nUSE_DATASET_CACHE = True         # toggle False to force a full rebuild\n\nCACHE_DIR = project_dir / \"cache\"\nCACHE_DIR.mkdir(parents=True, exist_ok=True)\nDF_CACHE_NAME   = f\"df_{DATASET_VERSION}_{TARGET_MODALITY}.parquet\"\nNORM_CACHE_NAME = f\"norm_{DATASET_VERSION}_{TARGET_MODALITY}.json\"\n\ndef _find_cache_file(name):\n    \"\"\"Search working dir first, then /kaggle/input/* for the cache file.\"\"\"\n    p = CACHE_DIR / name\n    if p.exists():\n        return p\n    if os.path.exists(\"/kaggle/input\"):\n        for r, _, files in os.walk(\"/kaggle/input\"):\n            if name in files:\n                return Path(r) / name\n    return None\n\nDF_CACHE_PATH   = _find_cache_file(DF_CACHE_NAME) if USE_DATASET_CACHE else None\nNORM_CACHE_PATH = _find_cache_file(NORM_CACHE_NAME) if USE_DATASET_CACHE else None\nDATASET_FROM_CACHE = (DF_CACHE_PATH is not None and NORM_CACHE_PATH is not None)\n\nif DATASET_FROM_CACHE:\n    print(f\"\\u26a1 Dataset cache HIT\")\n    print(f\"   df:   {DF_CACHE_PATH}\")\n    print(f\"   norm: {NORM_CACHE_PATH}\")\n    print(f\"   Skipping data loading, lung cropping, norm stats.\")\nelse:\n    print(f\"Dataset cache MISS — full rebuild ({DATASET_VERSION} {TARGET_MODALITY})\")\n    print(f\"   Will save to {CACHE_DIR}/ after build\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T16:56:20.18001Z","iopub.execute_input":"2026-04-24T16:56:20.180219Z","iopub.status.idle":"2026-04-24T16:56:20.462016Z","shell.execute_reply.started":"2026-04-24T16:56:20.180197Z","shell.execute_reply":"2026-04-24T16:56:20.461448Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def find_dir_with_marker(markers, search_roots=None, max_depth=6, prefer_keywords=None):\n    search_roots = search_roots or SEARCH_ROOTS\n    candidates = []\n    for root in search_roots:\n        if not os.path.exists(root):\n            continue\n        for dirpath, dirnames, _ in os.walk(root, followlinks=False):\n            depth = dirpath[len(root):].count(os.sep)\n            if depth > max_depth:\n                dirnames[:] = []\n                continue\n            if all(os.path.exists(os.path.join(dirpath, m)) for m in markers):\n                candidates.append(dirpath)\n    if not candidates:\n        return None\n    if prefer_keywords:\n        scored = [(sum(1 for kw in prefer_keywords if kw.lower() in c.lower()), -len(c), c)\n                  for c in candidates]\n        scored.sort(reverse=True)\n        return scored[0][2]\n    return sorted(candidates, key=len)[0]\n\ndef find_file(filename_patterns, search_roots=None, max_depth=8, prefer_keywords=None):\n    import fnmatch\n    search_roots = search_roots or SEARCH_ROOTS\n    candidates = []\n    for root in search_roots:\n        if not os.path.exists(root):\n            continue\n        for dirpath, _, filenames in os.walk(root, followlinks=False):\n            depth = dirpath[len(root):].count(os.sep)\n            if depth > max_depth:\n                continue\n            for fn in filenames:\n                if any(fnmatch.fnmatch(fn, pat) for pat in filename_patterns):\n                    candidates.append(os.path.join(dirpath, fn))\n    if not candidates:\n        return None\n    if prefer_keywords:\n        scored = [(sum(1 for kw in prefer_keywords if kw.lower() in c.lower()), -len(c), c)\n                  for c in candidates]\n        scored.sort(reverse=True)\n        return scored[0][2]\n    return sorted(candidates, key=len)[0]\n\n# Dataset 1 — Lung Tumor CT\n_tumor_root = (find_dir_with_marker([\"Images\", \"Annotations\"], prefer_keywords=[\"pidata\",\"tumor\",\"lung\"])\n               or find_dir_with_marker([\"Images\", \"Masks\"], prefer_keywords=[\"pidata\",\"tumor\",\"lung\"]))\nif _tumor_root:\n    tumor_image_dir = os.path.join(_tumor_root, \"Images\")\n    _ann = os.path.join(_tumor_root, \"Annotations\")\n    tumor_mask_dir = _ann if os.path.exists(_ann) else os.path.join(_tumor_root, \"Masks\")\n    print(f\"Lung Tumor CT: {_tumor_root}\")\nelse:\n    tumor_image_dir = tumor_mask_dir = None\n    print(\"⚠ Lung Tumor CT: NOT FOUND\")\n\n# Dataset 2 — TBX11K\n_tbx_anno = find_file([\"TBX11K_trainval_only_tb.json\", \"*trainval_only_tb*.json\"],\n                      prefer_keywords=[\"TBX11K\",\"tbx\"])\nif _tbx_anno:\n    tbx_anno_file = _tbx_anno\n    p = Path(_tbx_anno).parent\n    while p != p.parent:\n        if (p / \"imgs\").exists():\n            tbx_image_dir = str(p / \"imgs\"); break\n        p = p.parent\n    else:\n        tbx_image_dir = None\n    print(f\"TBX11K: {Path(_tbx_anno).parent.parent}\")\nelse:\n    tbx_anno_file = tbx_image_dir = None\n    print(\"⚠ TBX11K: NOT FOUND\")\n\n# Dataset 3 — Shenzhen TB (skipped)\nshenzhen_img_dir = shenzhen_ann_dir = shenzhen_mask_dir = None\nprint(\"Shenzhen TB: skipped\")\n\n# Dataset 4 — VinDr-CXR\n_vindr_root = (find_dir_with_marker([os.path.join(\"annotations\",\"instances_train.json\"), \"images\"],\n                                    prefer_keywords=[\"vindr\",\"vinbig\",\"cxr\"])\n               or find_dir_with_marker([os.path.join(\"annotations\",\"instances_val.json\"), \"images\"],\n                                       prefer_keywords=[\"vindr\",\"vinbig\",\"cxr\"]))\nif _vindr_root:\n    vindr_root     = _vindr_root\n    vindr_anno_dir = os.path.join(_vindr_root, \"annotations\")\n    vindr_img_dir  = os.path.join(_vindr_root, \"images\")\n    print(f\"VinDr-CXR: {_vindr_root}\")\nelse:\n    vindr_root = vindr_anno_dir = vindr_img_dir = None\n    print(\"⚠ VinDr-CXR: NOT FOUND\")\n\n# Dataset 5 — RSNA Pneumonia\n_rsna_meta = find_file([\"stage2_train_metadata.csv\",\"stage_2_train_metadata.csv\",\n                        \"stage2_train_labels.csv\",\"stage_2_train_labels.csv\"],\n                       prefer_keywords=[\"rsna\",\"pneumonia\"])\nif _rsna_meta:\n    rsna_train_meta = _rsna_meta\n    p = Path(_rsna_meta).parent\n    rsna_train_img = None\n    for c in [p/\"Training\"/\"Images\", p/\"train\"/\"images\", p/\"Images\",\n              p/\"images\", p/\"stage_2_train_images\"]:\n        if c.exists():\n            rsna_train_img = str(c); break\n    if rsna_train_img is None:\n        for root, _, files in os.walk(p):\n            if any(f.lower().endswith(\".png\") for f in files[:50]):\n                rsna_train_img = root; break\n    print(f\"RSNA Pneumonia: {p} (images: {rsna_train_img})\")\nelse:\n    rsna_train_meta = rsna_train_img = None\n    print(\"⚠ RSNA Pneumonia: NOT FOUND\")\n\n# Dataset 6 — VinBigData Chest X-ray Abnormalities\n_vinbigdata_csv = find_file([\"train.csv\"],\n                            prefer_keywords=[\"vinbig\", \"chest\", \"x-ray\", \"abnormalities\"])\nif _vinbigdata_csv:\n    vinbigdata_csv = _vinbigdata_csv\n    _vbd_root = Path(_vinbigdata_csv).parent\n    vinbigdata_img_dir = None\n    for _cand in [_vbd_root / \"train\", _vbd_root / \"images\" / \"train\", _vbd_root / \"train\" / \"images\"]:\n        if _cand.exists():\n            vinbigdata_img_dir = str(_cand); break\n    print(f\"VinBigData CXR: {_vbd_root} (images: {vinbigdata_img_dir})\")\nelse:\n    vinbigdata_csv = vinbigdata_img_dir = None\n    print(\"⚠ VinBigData CXR: NOT FOUND\")\n\n# VinDr class mapping (explicit xray tumor label)\n# Tightened: drop consolidation/opacity/infiltration — those are non-specific\n# radiological findings, not diagnoses. Mixing them into \"pneumonia\" with\n# tightly-annotated explicit pneumonia boxes was corrupting the class signal.\nVINDR_KEEP = {\n    \"nodule/mass\":   \"tumor_xray\",\n    \"lung tumor\":    \"tumor_xray\",\n    \"pneumonia\":     \"pneumonia\",\n    \"consolidation\": \"pneumonia\",\n    \"infiltration\":  \"pneumonia\",\n    # \"lung opacity\":  \"pneumonia\",  # v11: dropped — too noisy, capped pneumonia AP at ~0.19\n}\n\n# VinBigData: only Nodule/Mass survives — this dataset has no explicit\n# \"Pneumonia\" label, so nothing remains for pneumonia class from here.\nVINBIGDATA_KEEP = {\n    \"Nodule/Mass\":   \"tumor_xray\",\n    \"Consolidation\": \"pneumonia\",\n    \"Infiltration\":  \"pneumonia\",\n    # \"Lung Opacity\":  \"pneumonia\",  # v11: dropped — too noisy\n}\n\n# Dataset 7 — LIDC-IDRI (CT-only, produces tumor_ct).\n# Only useful when TARGET_MODALITY in (\"ct\", \"all\"). For \"xray\" runs, skip\n# discovery to avoid the multi-minute XML/DICOM scan.\nif TARGET_MODALITY in (\"ct\", \"all\"):\n    _lidc_root = (find_dir_with_marker([\"LIDC-IDRI\"], prefer_keywords=[\"lidc\"])\n                  or find_dir_with_marker([\"DOI\"], prefer_keywords=[\"lidc\",\"idri\"]))\n    if _lidc_root is None:\n        # Fallback: any directory with both .xml and .dcm files\n        for _r in SEARCH_ROOTS:\n            if not os.path.exists(_r): continue\n            for _dp, _, _files in os.walk(_r):\n                if any(_f.endswith(\".xml\") for _f in _files) and \\\n                   any(_f.lower().endswith((\".dcm\", \".dicom\")) for _f in _files):\n                    _lidc_root = _dp\n                    break\n            if _lidc_root: break\n    if _lidc_root and os.path.exists(_lidc_root):\n        lidc_img_dir = _lidc_root\n        lidc_xml_dir = _lidc_root\n        print(f\"LIDC-IDRI: {_lidc_root}\")\n    else:\n        lidc_xml_dir = lidc_img_dir = None\n        print(\"\\u26a0 LIDC-IDRI: NOT FOUND (attach the LIDC dataset on Kaggle)\")\nelse:\n    lidc_xml_dir = lidc_img_dir = None\n    print(f\"LIDC-IDRI: SKIPPED (TARGET_MODALITY={TARGET_MODALITY!r}, CT not needed)\")\n\n# Dataset 8 — ChestX-Det\n_chexdet_json = find_file([\"ChestX_Det_train.json\"],\n                          prefer_keywords=[\"chexdet\", \"chestx\", \"det\", \"annotations\"])\nif _chexdet_json:\n    chexdet_train_json = _chexdet_json\n    _cd_root = Path(_chexdet_json)\n    # Walk up to find train_data/train\n    chexdet_img_dir = None\n    for _p in [_cd_root.parent, _cd_root.parent.parent, _cd_root.parent.parent.parent]:\n        for _cand in [_p / \"train_data\" / \"train\", _p / \"train\" / \"images\",\n                      _p / \"images\" / \"train\", _p / \"train\"]:\n            if _cand.exists():\n                chexdet_img_dir = str(_cand); break\n        if chexdet_img_dir:\n            break\n    print(f\"ChestX-Det: json={_chexdet_json}, imgs={chexdet_img_dir}\")\nelse:\n    chexdet_train_json = chexdet_img_dir = None\n    print(\"\\u26a0 ChestX-Det: NOT FOUND\")\n\n# ChestX-Det: only Nodule + Mass kept. \"Consolidation\" is a non-specific\n# finding (was dominating pneumonia class with visually-different boxes than RSNA).\n# ChestX-Det has no actual \"Pneumonia\" annotation in the data itself.\nCHEXDET_KEEP = {\n    \"Nodule\":        \"tumor_xray\",\n    \"Mass\":          \"tumor_xray\",\n    \"Consolidation\": \"pneumonia\",\n    # ChestX-Det has no Infiltration/Lung Opacity classes; Effusion is not pneumonia.\n}\n\n# Modality-aware class scheme (driven by TARGET_MODALITY from cell 5)\n_ALL_CLASSES = [\"tumor_ct\", \"tumor_xray\", \"tuberculosis\", \"pneumonia\"]\nif TARGET_MODALITY == \"ct\":\n    CLASS_NAMES = [\"tumor_ct\"]\nelif TARGET_MODALITY == \"xray\":\n    CLASS_NAMES = [\"tumor_xray\", \"tuberculosis\", \"pneumonia\"]\nelse:  # \"all\" — legacy joint model\n    CLASS_NAMES = _ALL_CLASSES\nCLASS_TO_IDX = {n: i for i, n in enumerate(CLASS_NAMES)}\nNUM_CLASSES  = len(CLASS_NAMES)\nprint(f\"[{TARGET_MODALITY}] Classes ({NUM_CLASSES}): {CLASS_NAMES}\")\n\nprint(f\"\\nDevice: {device}\")\nprint(f\"Resolution: {train_sz}x{train_sz}\")\nprint(f\"Classes ({NUM_CLASSES}): {CLASS_NAMES}\")\n\n# Dataset 9 — cxray14 (VinBig + NH-Xray merged, YOLO-format labels, 640x640)\ncxray14_root = \"/kaggle/input/datasets/sandipacharya10/cxray14/cxray14\"\nif not os.path.exists(cxray14_root):\n    # Fallback: search for it (structure: <root>/train/labels/*.txt)\n    _cand = find_dir_with_marker([os.path.join(\"train\", \"labels\")],\n                                 prefer_keywords=[\"cxray14\", \"cxray\"])\n    cxray14_root = _cand\nif cxray14_root and os.path.exists(cxray14_root):\n    print(f\"cxray14: {cxray14_root}\")\nelse:\n    cxray14_root = None\n    print(\"⚠ cxray14: NOT FOUND\")\n\n\n# Dataset 10 — CXR Lung (Node21-style nodule detection)\ncxrlung_root = \"/kaggle/input/datasets/hryan007/cxr-lung-dataset-png\"\ncxrlung_meta = os.path.join(cxrlung_root, \"metadata.csv\")\nif os.path.exists(cxrlung_meta):\n    print(f\"CXR Lung (Node21): {cxrlung_root}\")\nelse:\n    cxrlung_root = None\n    cxrlung_meta = None\n    print(\"⚠ CXR Lung (Node21): NOT FOUND\")\n\n\n# Dataset 11 — SIIM-FISABIO-RSNA COVID-19 (radiologist-drawn opacity bboxes)\nsiim_root = Path(\"/kaggle/input/datasets/surajghuwalewala/siim-covid19-detection-data\")\nsiim_csv  = siim_root / \"img_relative_train_image_level.csv\"\nif siim_csv.exists():\n    print(f\"SIIM-COVID19: {siim_root}\")\nelse:\n    siim_root = None\n    siim_csv = None\n    print(\"\\u26a0 SIIM-COVID19: NOT FOUND\")\n\n# Dataset 12 — JSRT (Japanese Society of Radiological Technology) — pristine nodules\njsrt_root = Path(\"/kaggle/input/datasets/raddar/nodules-in-chest-xrays-jsrt\")\njsrt_meta = jsrt_root / \"jsrt_metadata.csv\"\njsrt_img_dir = jsrt_root / \"images\" / \"images\"\nif jsrt_meta.exists() and jsrt_img_dir.exists():\n    print(f\"JSRT: {jsrt_root}\")\nelse:\n    jsrt_root = None\n    jsrt_meta = None\n    jsrt_img_dir = None\n    print(\"\\u26a0 JSRT: NOT FOUND\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T19:02:18.11523Z","iopub.execute_input":"2026-04-24T19:02:18.115496Z","iopub.status.idle":"2026-04-24T19:13:59.399688Z","shell.execute_reply.started":"2026-04-24T19:02:18.115461Z","shell.execute_reply":"2026-04-24T19:13:59.399084Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Multi-Source Data Loading\n\nUnified DataFrame schema across all five annotation formats. The `max_rel_area=0.40` safety filter from v1 is kept — no box covering more than 40% of the image can reach training.\n","metadata":{}},{"cell_type":"markdown","source":"### 3.1 Lung Tumor CT — Mask-to-BBox with safety filter","metadata":{}},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    tumor_data = []\n    print(\"\\u23e9 Lung Tumor CT: skipped (cached)\")\nelse:\n    def mask_to_bbox_advanced(mask_path, img_path, min_area=50, max_rel_area=0.40):\n        \"\"\"Extract bboxes from a binary lesion mask with whole-lung safety filter.\"\"\"\n        mask = cv2.imread(str(mask_path), cv2.IMREAD_GRAYSCALE)\n        img  = cv2.imread(str(img_path))\n        if mask is None or img is None:\n            return []\n        img_h, img_w = img.shape[:2]\n        _, binary = cv2.threshold(mask, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)\n        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3))\n        cleaned = cv2.morphologyEx(binary, cv2.MORPH_OPEN,  kernel, iterations=2)\n        cleaned = cv2.morphologyEx(cleaned, cv2.MORPH_CLOSE, kernel, iterations=1)\n        contours, _ = cv2.findContours(cleaned, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n        bboxes = []\n        for c in contours:\n            a = cv2.contourArea(c)\n            if a <= min_area:\n                continue\n            x, y, w, h = cv2.boundingRect(c)\n            ba = w * h\n            if ba / (img_w * img_h) > max_rel_area:\n                continue\n            bboxes.append({\n                \"xmin\": x, \"ymin\": y, \"xmax\": x + w, \"ymax\": y + h,\n                \"width\": w, \"height\": h, \"area\": ba,\n                \"relative_area_pct\": round(ba / (img_w * img_h) * 100, 2),\n                \"aspect_ratio\": round(w / max(h, 1), 3),\n                \"center_x\": round((x + w / 2) / img_w, 3),\n                \"center_y\": round((y + h / 2) / img_h, 3),\n            })\n        return bboxes\n\n    tumor_data = []\n    if tumor_image_dir and tumor_mask_dir and os.path.exists(tumor_image_dir) and os.path.exists(tumor_mask_dir):\n        img_f = {os.path.splitext(f)[0]: f for f in os.listdir(tumor_image_dir)\n                 if f.lower().endswith((\".png\", \".jpg\", \".jpeg\"))}\n        msk_f = {os.path.splitext(f)[0]: f for f in os.listdir(tumor_mask_dir)\n                 if f.lower().endswith((\".png\", \".jpg\", \".jpeg\"))}\n        common = sorted(set(img_f) & set(msk_f))\n        print(f\"Lung Tumor CT: {len(common)} matched\")\n        no_bb = 0\n        for k in tqdm(common, desc=\"Tumor CT bboxes\"):\n            ip = os.path.join(tumor_image_dir, img_f[k])\n            mp = os.path.join(tumor_mask_dir,  msk_f[k])\n            bb = mask_to_bbox_advanced(mp, ip)\n            if not bb:\n                no_bb += 1; continue\n            for b in bb:\n                tumor_data.append({\n                    \"image\": img_f[k], \"image_path\": ip,\n                    \"label\": \"tumor_ct\",  # CT-specific label\n                    \"modality\": \"ct\", \"source_dataset\": \"lung_tumor_ct\", **b\n                })\n        print(f\"  → {len(tumor_data)} tumor_ct bboxes ({no_bb} images skipped)\")\n    else:\n        print(\"⚠ Lung Tumor CT not found.\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T17:43:04.855285Z","iopub.execute_input":"2026-04-24T17:43:04.855468Z","iopub.status.idle":"2026-04-24T17:44:42.77538Z","shell.execute_reply.started":"2026-04-24T17:43:04.855428Z","shell.execute_reply":"2026-04-24T17:44:42.774633Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3.2 TBX11K — COCO JSON","metadata":{}},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    tb_data = []\n    print(\"\\u23e9 TBX11K: skipped (cached)\")\nelse:\n    tb_data = []\n    if tbx_anno_file and os.path.exists(tbx_anno_file):\n        with open(tbx_anno_file, \"r\") as f:\n            coco = json.load(f)\n        id2img = {img[\"id\"]: img for img in coco[\"images\"]}\n        id2cat = {cat[\"id\"]: cat[\"name\"] for cat in coco[\"categories\"]}\n        print(f\"TBX11K: {len(coco['images'])} images, {len(coco['annotations'])} annotations\")\n        tb_ids = {cid for cid, n in id2cat.items()\n                  if any(k in n.lower() for k in [\"tb\", \"tuberculosis\"])}\n        for ann in tqdm(coco[\"annotations\"], desc=\"TBX11K\"):\n            if ann[\"category_id\"] not in tb_ids:\n                continue\n            ii = id2img.get(ann[\"image_id\"])\n            if ii is None:\n                continue\n            x, y, w, h = ann[\"bbox\"]\n            iw, ih, fn = ii[\"width\"], ii[\"height\"], ii[\"file_name\"]\n            ip = None\n            for sd in [\"\", \"tb\", \"sick\", \"health\"]:\n                c = os.path.join(tbx_image_dir, sd, fn) if sd else os.path.join(tbx_image_dir, fn)\n                if os.path.exists(c):\n                    ip = c; break\n            if ip is None:\n                for root, _, files in os.walk(tbx_image_dir):\n                    if fn in files:\n                        ip = os.path.join(root, fn); break\n            if ip is None:\n                continue\n            ba = int(w * h)\n            tb_data.append({\n                \"image\": fn, \"image_path\": ip,\n                \"label\": \"tuberculosis\", \"modality\": \"xray\", \"source_dataset\": \"TBX11K\",\n                \"xmin\": int(x), \"ymin\": int(y), \"xmax\": int(x+w), \"ymax\": int(y+h),\n                \"width\": int(w), \"height\": int(h), \"area\": ba,\n                \"relative_area_pct\": round(ba/(iw*ih)*100, 2),\n                \"aspect_ratio\": round(w/max(h,1), 3),\n                \"center_x\": round((x+w/2)/iw, 3),\n                \"center_y\": round((y+h/2)/ih, 3),\n            })\n        print(f\"  → {len(tb_data)} TB bboxes from TBX11K\")\n    else:\n        print(\"⚠ TBX11K not found.\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T17:44:42.776444Z","iopub.execute_input":"2026-04-24T17:44:42.776635Z","iopub.status.idle":"2026-04-24T17:44:44.462908Z","shell.execute_reply.started":"2026-04-24T17:44:42.776611Z","shell.execute_reply":"2026-04-24T17:44:44.462215Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3.3 VinDr-CXR — Radiologist-drawn boxes","metadata":{}},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    vindr_data = []\n    print(\"\\u23e9 VinDr-CXR: skipped (cached)\")\nelse:\n    def parse_vindr_coco(split):\n        anno_path = os.path.join(vindr_anno_dir, f\"instances_{split}.json\")\n        img_dir   = os.path.join(vindr_img_dir,  split)\n        if not os.path.exists(anno_path):\n            print(f\"  VinDr-CXR {split}: annotation file not found\")\n            return []\n        with open(anno_path, \"r\") as f:\n            coco = json.load(f)\n        id2img = {img[\"id\"]: img for img in coco[\"images\"]}\n        id2cat = {cat[\"id\"]: cat[\"name\"] for cat in coco[\"categories\"]}\n        cat_id_to_label = {cid: VINDR_KEEP[cname.strip().lower()]\n                           for cid, cname in id2cat.items()\n                           if cname.strip().lower() in VINDR_KEEP}\n        print(f\"  VinDr-CXR/{split}: {len(coco['images'])} imgs, {len(coco['annotations'])} anns, \"\n              f\"{len(cat_id_to_label)}/{len(id2cat)} cats kept\")\n        records = []\n        for ann in tqdm(coco[\"annotations\"], desc=f\"VinDr-CXR {split}\", leave=False):\n            label = cat_id_to_label.get(ann[\"category_id\"])\n            if label is None:\n                continue\n            ii = id2img.get(ann[\"image_id\"])\n            if ii is None:\n                continue\n            fn = os.path.basename(ii[\"file_name\"])\n            ip = os.path.join(img_dir, fn)\n            if not os.path.exists(ip):\n                continue\n\n            x, y, w, h = ann[\"bbox\"]\n\n            # FIX ①: Bbox coords in this dataset are ALREADY in the downloaded image's\n            # coordinate space (max ~1024). The COCO JSON's width/height fields contain\n            # the ORIGINAL DICOM dims (~3000px), which must NOT be used for normalization.\n            # Read the REAL image dimensions instead.\n            real_w, real_h = get_image_dims(ip)\n            if real_w is None:\n                continue\n            iw, ih = real_w, real_h  # Use real dims (likely 1024x1024)\n\n            if w <= 1 or h <= 1:\n                continue\n            ba = int(w * h)\n            if ba / (iw * ih) > 0.40:\n                continue\n            records.append({\n                \"image\": fn, \"image_path\": ip,\n                \"label\": label, \"modality\": \"xray\", \"source_dataset\": f\"VinDr-CXR_{split}\",\n                \"xmin\": int(x), \"ymin\": int(y), \"xmax\": int(x+w), \"ymax\": int(y+h),\n                \"width\": int(w), \"height\": int(h), \"area\": ba,\n                \"relative_area_pct\": round(ba/(iw*ih)*100, 2),\n                \"aspect_ratio\": round(w/max(h,1), 3),\n                \"center_x\": round((x+w/2)/iw, 3),\n                \"center_y\": round((y+h/2)/ih, 3),\n            })\n        return records\n\n    def _dedup_vindr(records, iou_threshold=0.7):\n        \"\"\"VinDr has multi-radiologist annotations → the same lesion appears 2-3×\n        with slightly different coords. Collapse near-duplicates per (image, label).\n        Keeps the MEDIAN box (consensus) rather than the arithmetic mean — more\n        robust to a single outlier annotator.\"\"\"\n        if not records:\n            return records\n        from collections import defaultdict\n        def _iou(a, b):\n            ax1, ay1, ax2, ay2 = a; bx1, by1, bx2, by2 = b\n            ix1, iy1 = max(ax1, bx1), max(ay1, by1)\n            ix2, iy2 = min(ax2, bx2), min(ay2, by2)\n            iw, ih = max(0, ix2-ix1), max(0, iy2-iy1)\n            inter = iw * ih\n            if inter == 0:\n                return 0.0\n            union = (ax2-ax1)*(ay2-ay1) + (bx2-bx1)*(by2-by1) - inter\n            return inter / max(union, 1)\n        groups = defaultdict(list)\n        for r in records:\n            groups[(r[\"image\"], r[\"label\"])].append(r)\n        out = []\n        for key, rs in groups.items():\n            boxes = [(r[\"xmin\"], r[\"ymin\"], r[\"xmax\"], r[\"ymax\"]) for r in rs]\n            used = [False] * len(rs)\n            for i in range(len(rs)):\n                if used[i]:\n                    continue\n                cluster = [i]\n                used[i] = True\n                for j in range(i+1, len(rs)):\n                    if not used[j] and _iou(boxes[i], boxes[j]) >= iou_threshold:\n                        cluster.append(j); used[j] = True\n                if len(cluster) == 1:\n                    out.append(rs[i])\n                else:\n                    import numpy as _np\n                    arr = _np.array([boxes[k] for k in cluster], dtype=_np.float32)\n                    mx = _np.median(arr, axis=0)\n                    merged = dict(rs[cluster[0]])\n                    x1, y1, x2, y2 = int(mx[0]), int(mx[1]), int(mx[2]), int(mx[3])\n                    merged[\"xmin\"], merged[\"ymin\"] = x1, y1\n                    merged[\"xmax\"], merged[\"ymax\"] = x2, y2\n                    merged[\"width\"]  = x2 - x1\n                    merged[\"height\"] = y2 - y1\n                    merged[\"area\"]   = int((x2-x1) * (y2-y1))\n                    out.append(merged)\n        return out\n\n    vindr_data = []\n    if vindr_root and os.path.exists(vindr_root):\n        for split in [\"train\", \"val\"]:\n            raw = parse_vindr_coco(split)\n            before = len(raw)\n            raw = _dedup_vindr(raw, iou_threshold=0.5)  # v11: more aggressive merge\n            print(f\"  VinDr-CXR/{split} dedup: {before} → {len(raw)} bboxes \"\n                  f\"({before - len(raw)} multi-radiologist duplicates merged)\")\n            vindr_data.extend(raw)\n        n_tumor = sum(1 for r in vindr_data if r[\"label\"] == \"tumor_xray\")\n        n_pneu  = sum(1 for r in vindr_data if r[\"label\"] == \"pneumonia\")\n        print(f\"  → VinDr-CXR total: {len(vindr_data)} bboxes (tumor_xray: {n_tumor}, pneumonia: {n_pneu})\")\n    else:\n        print(\"⚠ VinDr-CXR not found.\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T17:44:44.464465Z","iopub.execute_input":"2026-04-24T17:44:44.464651Z","iopub.status.idle":"2026-04-24T17:44:55.992872Z","shell.execute_reply.started":"2026-04-24T17:44:44.46463Z","shell.execute_reply":"2026-04-24T17:44:55.99219Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3.4 RSNA Pneumonia","metadata":{}},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    rsna_data = []\n    print(\"\\u23e9 RSNA Pneumonia: skipped (cached)\")\nelse:\n    rsna_data = []\n    if (rsna_train_meta and rsna_train_img\n        and os.path.exists(rsna_train_meta) and os.path.exists(rsna_train_img)):\n        rsna_df = pd.read_csv(rsna_train_meta)\n        n_pos = int((rsna_df.get(\"Target\", 0) == 1).sum())\n        print(f\"RSNA Pneumonia: {len(rsna_df)} rows, {n_pos} positive\")\n        pos = rsna_df[rsna_df.get(\"Target\", 0) == 1].dropna(subset=[\"x\",\"y\",\"width\",\"height\"])\n        print(f\"  {len(pos)} positive with valid boxes\")\n        RSNA_W = RSNA_H = 1024\n        for _, r in tqdm(pos.iterrows(), total=len(pos), desc=\"RSNA Pneumonia\"):\n            fn = f\"{r['patientId']}.png\"\n            ip = os.path.join(rsna_train_img, fn)\n            if not os.path.exists(ip):\n                continue\n            x, y, w, h = float(r[\"x\"]), float(r[\"y\"]), float(r[\"width\"]), float(r[\"height\"])\n            if w <= 1 or h <= 1:\n                continue\n            ba = int(w * h)\n            if ba / (RSNA_W * RSNA_H) > 0.40:\n                continue\n            rsna_data.append({\n                \"image\": fn, \"image_path\": ip,\n                \"label\": \"pneumonia\", \"modality\": \"xray\", \"source_dataset\": \"RSNA_Pneumonia\",\n                \"xmin\": int(x), \"ymin\": int(y), \"xmax\": int(x+w), \"ymax\": int(y+h),\n                \"width\": int(w), \"height\": int(h), \"area\": ba,\n                \"relative_area_pct\": round(ba/(RSNA_W*RSNA_H)*100, 2),\n                \"aspect_ratio\": round(w/max(h,1), 3),\n                \"center_x\": round((x+w/2)/RSNA_W, 3),\n                \"center_y\": round((y+h/2)/RSNA_H, 3),\n            })\n        print(f\"  → {len(rsna_data)} pneumonia bboxes from RSNA\")\n    else:\n        print(\"⚠ RSNA not found.\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T17:44:55.993744Z","iopub.execute_input":"2026-04-24T17:44:55.993959Z","iopub.status.idle":"2026-04-24T17:45:10.941203Z","shell.execute_reply.started":"2026-04-24T17:44:55.993934Z","shell.execute_reply":"2026-04-24T17:45:10.940481Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    vinbigdata_data = []\n    print(\"\\u23e9 VinBigData CXR (DICOM-heavy): skipped (cached)\")\nelse:\n    # ── 3.5 VinBigData Chest X-ray Abnormalities ──────────────────────────────────\n    vinbigdata_data = []\n    if (vinbigdata_csv and vinbigdata_img_dir\n        and os.path.exists(vinbigdata_csv) and os.path.exists(vinbigdata_img_dir)):\n        vbd_df = pd.read_csv(vinbigdata_csv)\n        print(f\"VinBigData CXR: {len(vbd_df)} rows, {vbd_df['class_name'].nunique()} classes\")\n        print(f\"  Classes: {dict(vbd_df['class_name'].value_counts())}\")\n\n        # Filter to classes we care about\n        vbd_df[\"unified_label\"] = vbd_df[\"class_name\"].map(VINBIGDATA_KEEP)\n        vbd_pos = vbd_df.dropna(subset=[\"unified_label\"]).copy()\n        vbd_pos = vbd_pos.dropna(subset=[\"x_min\", \"y_min\", \"x_max\", \"y_max\"])\n        print(f\"  {len(vbd_pos)} annotations with valid boxes in kept classes\")\n\n        _vbd_converted = 0\n        _vbd_raw_records = []  # collect pre-dedup; we'll dedup across radiologists\n        for _, r in tqdm(vbd_pos.iterrows(), total=len(vbd_pos), desc=\"VinBigData CXR\"):\n            img_id = str(r[\"image_id\"])\n            # Find the raw file\n            raw_path = None\n            for ext in [\".dicom\", \".dcm\", \".png\", \".jpg\", \"\"]:\n                cand = os.path.join(vinbigdata_img_dir, img_id + ext)\n                if os.path.exists(cand):\n                    raw_path = cand; break\n            if raw_path is None:\n                continue\n\n            # Convert DICOM→PNG (now resized to 1024) and get both new + original dims\n            if raw_path.lower().endswith((\".dcm\", \".dicom\")):\n                result = dicom_to_png(raw_path)\n                if result is None:\n                    continue\n                ip, iw, ih, orig_w, orig_h = result\n                _vbd_converted += 1\n            else:\n                ip = raw_path\n                iw, ih = get_image_dims(ip)\n                if iw is None:\n                    continue\n                orig_w, orig_h = iw, ih  # not a DICOM, coords already in PNG space\n\n            x1, y1, x2, y2 = float(r[\"x_min\"]), float(r[\"y_min\"]), float(r[\"x_max\"]), float(r[\"y_max\"])\n\n            # FIX: VinBigData CSV coords are in ORIGINAL DICOM space.\n            # After resizing DICOM → 1024, scale bbox coords to match.\n            if (orig_w, orig_h) != (iw, ih):\n                sx = iw / max(orig_w, 1)\n                sy = ih / max(orig_h, 1)\n                x1 *= sx; x2 *= sx\n                y1 *= sy; y2 *= sy\n\n            w = x2 - x1\n            h = y2 - y1\n            if w <= 1 or h <= 1:\n                continue\n\n            ba = int(w * h)\n            if ba / (iw * ih) > 0.40:\n                continue\n\n            fn = os.path.basename(ip)\n            _vbd_raw_records.append({\n                \"image\": fn, \"image_path\": ip,\n                \"label\": r[\"unified_label\"], \"modality\": \"xray\",\n                \"source_dataset\": \"VinBigData_CXR\",\n                \"xmin\": int(x1), \"ymin\": int(y1), \"xmax\": int(x2), \"ymax\": int(y2),\n                \"width\": int(w), \"height\": int(h), \"area\": ba,\n                \"relative_area_pct\": round(ba / (iw * ih) * 100, 2),\n                \"aspect_ratio\": round(w / max(h, 1), 3),\n                \"center_x\": round((x1 + w / 2) / iw, 3),\n                \"center_y\": round((y1 + h / 2) / ih, 3),\n            })\n\n        # Dedup multi-radiologist annotations (same as VinDr; shares R1/R2/R3 style labels)\n        try:\n            _before = len(_vbd_raw_records)\n            vinbigdata_data = _dedup_vindr(_vbd_raw_records, iou_threshold=0.5)  # v11: more aggressive merge\n            print(f\"  Multi-radiologist dedup: {_before} → {len(vinbigdata_data)} bboxes\")\n        except NameError:\n            vinbigdata_data = _vbd_raw_records\n\n        n_tumor = sum(1 for r in vinbigdata_data if r[\"label\"] == \"tumor_xray\")\n        n_pneu  = sum(1 for r in vinbigdata_data if r[\"label\"] == \"pneumonia\")\n        print(f\"  ✓ Converted {_vbd_converted} DICOMs to PNG (resized to max {DICOM_TARGET_SIZE}px)\")\n        print(f\"  → VinBigData CXR total: {len(vinbigdata_data)} bboxes \"\n              f\"(tumor_xray: {n_tumor}, pneumonia: {n_pneu})\")\n    else:\n        print(\"⚠ VinBigData CXR not found.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T17:45:10.942344Z","iopub.execute_input":"2026-04-24T17:45:10.942788Z","iopub.status.idle":"2026-04-24T18:02:26.920901Z","shell.execute_reply.started":"2026-04-24T17:45:10.942747Z","shell.execute_reply":"2026-04-24T18:02:26.920234Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3.5 Unified DataFrame","metadata":{}},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    lidc_data = []\n    print(\"\\u23e9 LIDC-IDRI: skipped (cached)\")\nelse:\n    # ── 3.6 LIDC-IDRI — Lung Nodules (CT) ─────────────────────────────────────────\n    import xml.etree.ElementTree as ET  # [LIDC DROPPED — see Cell 6]\n\n    lidc_data = []\n    if (TARGET_MODALITY in (\"ct\", \"all\")\n        and lidc_xml_dir is not None and lidc_img_dir is not None\n        and os.path.exists(lidc_img_dir)):\n\n        # Build SOP UID → filepath map from .dcm files\n        print(\"LIDC-IDRI: scanning DICOM headers for SOP UID mapping...\")\n        sop_to_path = {}\n        try:\n            import pydicom\n            dcm_files = [os.path.join(lidc_img_dir, f)\n                         for f in os.listdir(lidc_img_dir) if f.lower().endswith((\".dcm\", \".dicom\"))]\n            for fp in tqdm(dcm_files, desc=\"LIDC SOP scan\", leave=False):\n                try:\n                    ds = pydicom.dcmread(fp, stop_before_pixels=True)\n                    sop_to_path[str(ds.SOPInstanceUID).strip()] = fp\n                except Exception:\n                    continue\n            print(f\"  Mapped {len(sop_to_path)} SOP UIDs to files\")\n        except ImportError:\n            print(\"  ⚠ pydicom not available, LIDC-IDRI will be skipped\")\n\n        # Helper: find element with or without namespace\n        def nsfind(el, tag, ns=\"{http://www.nih.gov}\"):\n            \"\"\"Find element trying namespaced first, then bare tag.\"\"\"\n            r = el.find(ns + tag)\n            if r is None:\n                r = el.find(tag)\n            return r\n\n        def nsfindall(el, tag, ns=\"{http://www.nih.gov}\"):\n            \"\"\"Findall trying namespaced first, then bare tag.\"\"\"\n            r = el.findall(ns + tag)\n            if not r:\n                r = el.findall(tag)\n            return r\n\n        # Parse XML annotations\n        xml_count = 0\n        _raw_rois = 0\n        _skipped_no_chars = 0\n        _skipped_inclusion = 0\n        _skipped_small = 0\n        _skipped_no_sop = 0\n        _lidc_converted = 0\n        _seen_centroids = set()  # for dedup: (sop_uid, cx_bin, cy_bin)\n\n        # Collect all XML files (may be in subdirs or flat)\n        all_xml_files = []\n        for root_dir, dirs, files in os.walk(lidc_xml_dir):\n            for f in files:\n                if f.lower().endswith(\".xml\"):\n                    all_xml_files.append(os.path.join(root_dir, f))\n        print(f\"  Found {len(all_xml_files)} XML files\")\n\n        for xml_path in tqdm(all_xml_files, desc=\"LIDC XML parse\", leave=False):\n            try:\n                tree = ET.parse(xml_path)\n            except Exception:\n                continue\n            xroot = tree.getroot()\n            xml_count += 1\n\n            for session in nsfindall(xroot, \"readingSession\"):\n                # Process BOTH unblindedReadNodule and blindedReadNodule\n                for tag in [\"unblindedReadNodule\", \"blindedReadNodule\"]:\n                    for nodule in nsfindall(session, tag):\n                        # Keep nodules WITH characteristics (≥3mm significant)\n                        chars = nsfind(nodule, \"characteristics\")\n                        if chars is None:\n                            _skipped_no_chars += 1\n                            continue\n\n                        for roi in nsfindall(nodule, \"roi\"):\n                            _raw_rois += 1\n\n                            # FIX: Filter inclusion==FALSE (excluded regions)\n                            incl_el = nsfind(roi, \"inclusion\")\n                            if incl_el is not None and incl_el.text and incl_el.text.strip().upper() == \"FALSE\":\n                                _skipped_inclusion += 1\n                                continue\n\n                            edges = nsfindall(roi, \"edgeMap\")\n                            if len(edges) < 3:\n                                _skipped_small += 1\n                                continue\n\n                            points = []\n                            for e in edges:\n                                xc = nsfind(e, \"xCoord\")\n                                yc = nsfind(e, \"yCoord\")\n                                if xc is None or yc is None:\n                                    continue\n                                points.append((int(xc.text), int(yc.text)))\n\n                            if len(points) < 3:\n                                _skipped_small += 1\n                                continue\n\n                            xs, ys = zip(*points)\n                            xmin, xmax = min(xs), max(xs)\n                            ymin, ymax = min(ys), max(ys)\n                            w = xmax - xmin\n                            h = ymax - ymin\n                            if w <= 1 or h <= 1:\n                                _skipped_small += 1\n                                continue\n\n                            # Find the image file via SOP UID\n                            sop_el = nsfind(roi, \"imageSOP_UID\")\n                            if sop_el is None or sop_el.text is None:\n                                _skipped_no_sop += 1\n                                continue\n                            sop_uid = sop_el.text.strip()\n                            dcm_path = sop_to_path.get(sop_uid)\n                            if dcm_path is None:\n                                _skipped_no_sop += 1\n                                continue\n\n                            # Dedup: same SOP + similar centroid = same nodule from different reader\n                            cx_bin = round((xmin + xmax) / 2 / 10) * 10\n                            cy_bin = round((ymin + ymax) / 2 / 10) * 10\n                            dedup_key = (sop_uid, cx_bin, cy_bin)\n                            if dedup_key in _seen_centroids:\n                                continue\n                            _seen_centroids.add(dedup_key)\n\n                            # Convert DICOM → PNG and get real dims\n                            result = dicom_to_png(dcm_path)\n                            if result is None:\n                                continue\n                            ip, iw, ih = result\n                            _lidc_converted += 1\n\n                            fn = os.path.basename(ip)\n                            ba = int(w * h)\n\n                            lidc_data.append({\n                                \"image\": fn, \"image_path\": ip,\n                                \"label\": \"tumor_ct\", \"modality\": \"ct\",\n                                \"source_dataset\": \"LIDC-IDRI\",\n                                \"xmin\": xmin, \"ymin\": ymin,\n                                \"xmax\": xmax, \"ymax\": ymax,\n                                \"width\": w, \"height\": h, \"area\": ba,\n                                \"relative_area_pct\": round(ba / (iw * ih) * 100, 2),\n                                \"aspect_ratio\": round(w / max(h, 1), 3),\n                                \"center_x\": round((xmin + w / 2) / iw, 3),\n                                \"center_y\": round((ymin + h / 2) / ih, 3),\n                            })\n\n        print(f\"  Parsed {xml_count} XMLs, {_raw_rois} raw ROIs\")\n        print(f\"  Skipped: {_skipped_no_chars} no-chars, {_skipped_inclusion} inclusion=FALSE, \"\n              f\"{_skipped_small} too-small, {_skipped_no_sop} SOP-not-found\")\n        print(f\"  ✓ Converted {_lidc_converted} DICOM slices to PNG\")\n        print(f\"  → {len(lidc_data)} nodule bboxes (after dedup)\")\n    else:\n        if TARGET_MODALITY == \"xray\":\n            print(\"LIDC-IDRI: skipped (xray-only run)\")\n        else:\n            print(\"⚠ LIDC-IDRI: enabled but path not set — check Cell 6\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T18:02:26.921906Z","iopub.execute_input":"2026-04-24T18:02:26.922127Z","iopub.status.idle":"2026-04-24T18:02:26.939649Z","shell.execute_reply.started":"2026-04-24T18:02:26.922101Z","shell.execute_reply":"2026-04-24T18:02:26.938948Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    chexdet_data = []\n    print(\"\\u23e9 ChestX-Det: skipped (cached)\")\nelse:\n    # ── 3.7 ChestX-Det — Multi-class Chest X-ray Detection ─────────────────────────\n    chexdet_data = []\n    if (chexdet_train_json and chexdet_img_dir\n        and os.path.exists(chexdet_train_json) and os.path.exists(chexdet_img_dir)):\n        with open(chexdet_train_json, \"r\") as f:\n            chexdet_anno = json.load(f)\n        print(f\"ChestX-Det: {len(chexdet_anno)} images in train JSON\")\n\n        all_syms = {}\n        for entry in chexdet_anno:\n            for s in entry.get(\"syms\", []):\n                all_syms[s] = all_syms.get(s, 0) + 1\n        print(f\"  All classes: {all_syms}\")\n\n        kept = skipped = 0\n        for entry in tqdm(chexdet_anno, desc=\"ChestX-Det\", leave=False):\n            fn = entry[\"file_name\"]\n            ip = os.path.join(chexdet_img_dir, fn)\n            if not os.path.exists(ip):\n                continue\n            syms = entry.get(\"syms\", [])\n            boxes = entry.get(\"boxes\", [])\n            if len(syms) != len(boxes):\n                continue\n\n            # FIX ③: Read REAL image dimensions\n            iw, ih = get_image_dims(ip)\n            if iw is None:\n                continue\n\n            for sym, box in zip(syms, boxes):\n                label = CHEXDET_KEEP.get(sym)\n                if label is None:\n                    skipped += 1\n                    continue\n\n                x1, y1, x2, y2 = box\n                w = x2 - x1\n                h = y2 - y1\n                if w <= 1 or h <= 1:\n                    continue\n\n                ba = int(w * h)\n                if ba / (iw * ih) > 0.40:\n                    continue\n\n                kept += 1\n                chexdet_data.append({\n                    \"image\": fn, \"image_path\": ip,\n                    \"label\": label, \"modality\": \"xray\",\n                    \"source_dataset\": \"ChestX-Det\",\n                    \"xmin\": int(x1), \"ymin\": int(y1),\n                    \"xmax\": int(x2), \"ymax\": int(y2),\n                    \"width\": int(w), \"height\": int(h), \"area\": ba,\n                    \"relative_area_pct\": round(ba / (iw * ih) * 100, 2),\n                    \"aspect_ratio\": round(w / max(h, 1), 3),\n                    \"center_x\": round((x1 + w / 2) / iw, 3),\n                    \"center_y\": round((y1 + h / 2) / ih, 3),\n                })\n\n        n_tumor = sum(1 for r in chexdet_data if r[\"label\"] == \"tumor_xray\")\n        n_pneu  = sum(1 for r in chexdet_data if r[\"label\"] == \"pneumonia\")\n        print(f\"  → ChestX-Det: {len(chexdet_data)} kept ({n_tumor} tumor_xray, {n_pneu} pneumonia), \"\n              f\"{skipped} skipped (non-target classes)\")\n    else:\n        print(\"⚠ ChestX-Det not found.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T18:02:26.940759Z","iopub.execute_input":"2026-04-24T18:02:26.940974Z","iopub.status.idle":"2026-04-24T18:03:19.3238Z","shell.execute_reply.started":"2026-04-24T18:02:26.940946Z","shell.execute_reply":"2026-04-24T18:03:19.32319Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    cxray14_data = []\n    print(\"\\u23e9 cxray14: skipped (cached)\")\nelse:\n    # -- 3.8 cxray14 (VinBig + NH-Xray merged, YOLO format, strict: Nodule/Mass only) ----\n    cxray14_data = []\n    if cxray14_root and os.path.exists(cxray14_root):\n        # cxray14 has 15 classes (0-14). We strictly keep class 8 (Nodule/Mass) only,\n        # per the v7 decision to drop non-specific findings (Consolidation, Opacity,\n        # Infiltration). The dropped classes are also not in our target set.\n        CXRAY14_KEEP = {\n            # 8: \"tumor_xray\",     # v11: dropped — NIH-derived labels are ~25-30% noisy; Node21+JSRT supersede\n            4: \"pneumonia\",        # Consolidation\n            6: \"pneumonia\",        # Infiltration\n            # 7: \"pneumonia\",      # v11: dropped — Lung Opacity too noisy\n        }\n        CXRAY14_ALL_CLASSES = {\n            0: \"Aortic enlargement\", 1: \"Atelectasis\", 2: \"Calcification\",\n            3: \"Cardiomegaly\", 4: \"Consolidation\", 5: \"ILD\", 6: \"Infiltration\",\n            7: \"Lung Opacity\", 8: \"Nodule/Mass\", 9: \"Other lesion\",\n            10: \"Pleural effusion\", 11: \"Pleural thickening\", 12: \"Pneumothorax\",\n            13: \"Pulmonary fibrosis\", 14: \"No finding\",\n        }\n        IMG_SZ = 640  # cxray14 images are pre-resized to 640x640\n        raw_class_hist = {i: 0 for i in range(15)}\n        n_images_seen = 0\n        n_boxes_kept  = 0\n\n        for split in [\"train\", \"val\", \"test\"]:\n            img_dir = os.path.join(cxray14_root, split, \"images\")\n            lbl_dir = os.path.join(cxray14_root, split, \"labels\")\n            if not (os.path.exists(img_dir) and os.path.exists(lbl_dir)):\n                continue\n            lbl_files = [f for f in os.listdir(lbl_dir) if f.endswith(\".txt\")]\n            for lbl_file in tqdm(lbl_files, desc=f\"cxray14 {split}\", leave=False):\n                stem = os.path.splitext(lbl_file)[0]\n                img_path = None\n                for ext in (\".png\", \".jpg\", \".jpeg\"):\n                    cand = os.path.join(img_dir, stem + ext)\n                    if os.path.exists(cand):\n                        img_path = cand\n                        break\n                if img_path is None:\n                    continue\n                n_images_seen += 1\n                image_key = f\"cxray14_{split}_{stem}\"\n                with open(os.path.join(lbl_dir, lbl_file), \"r\") as f:\n                    for line in f:\n                        parts = line.strip().split()\n                        if len(parts) < 5:\n                            continue\n                        try:\n                            cid = int(parts[0])\n                            xc, yc, w_n, h_n = map(float, parts[1:5])\n                        except ValueError:\n                            continue\n                        if 0 <= cid < 15:\n                            raw_class_hist[cid] += 1\n                        if cid not in CXRAY14_KEEP:\n                            continue\n                        x0 = (xc - w_n / 2) * IMG_SZ\n                        y0 = (yc - h_n / 2) * IMG_SZ\n                        bw = w_n * IMG_SZ\n                        bh = h_n * IMG_SZ\n                        if bw <= 1 or bh <= 1:\n                            continue\n                        # Safety: drop boxes that cover >40% of image\n                        rel_area = (bw * bh) / (IMG_SZ * IMG_SZ)\n                        if rel_area > 0.40:\n                            continue\n                        cxray14_data.append({\n                            \"image\": image_key,\n                            \"image_path\": img_path,\n                            \"label\": CXRAY14_KEEP[cid],\n                            \"modality\": \"xray\",\n                            \"source_dataset\": \"cxray14\",\n                            \"xmin\": int(x0), \"ymin\": int(y0),\n                            \"xmax\": int(x0 + bw), \"ymax\": int(y0 + bh),\n                            \"width\": int(bw), \"height\": int(bh),\n                            \"area\": int(bw * bh),\n                            \"relative_area_pct\": round(rel_area * 100, 2),\n                            \"aspect_ratio\": round(bw / max(bh, 1), 3),\n                            \"center_x\": round(xc, 3),\n                            \"center_y\": round(yc, 3),\n                        })\n                        n_boxes_kept += 1\n\n        print(f\"cxray14: {n_images_seen} images scanned, {n_boxes_kept} boxes kept (multi-class mapping)\")\n        print(f\"  Raw class histogram (all 15 classes, total boxes found):\")\n        for cid in sorted(raw_class_hist):\n            cnt = raw_class_hist[cid]\n            kept_sym = \"✓\" if cid in CXRAY14_KEEP else \"✗\"\n            tgt = CXRAY14_KEEP.get(cid, \"dropped\")\n            print(f\"    {kept_sym}  {cid:>2}: {CXRAY14_ALL_CLASSES[cid]:<22}  n={cnt:>6}  -> {tgt}\")\n    else:\n        print(\"cxray14: skipped (path not set)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T19:14:54.879Z","iopub.execute_input":"2026-04-24T19:14:54.879248Z","iopub.status.idle":"2026-04-24T19:15:58.841566Z","shell.execute_reply.started":"2026-04-24T19:14:54.879222Z","shell.execute_reply":"2026-04-24T19:15:58.840701Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    cxrlung_data = []\n    print(\"\\u23e9 CXR Lung Node21: skipped (cached)\")\nelse:\n    # ── 3.9 CXR Lung (Node21-style) — Nodule Detection ─────────────────────────────\n    cxrlung_data = []\n    if cxrlung_root and cxrlung_meta and os.path.exists(cxrlung_meta):\n        cm = pd.read_csv(cxrlung_meta)\n        # label==1 = nodule positive; label==0 = control image (skip — no bbox to learn)\n        cm = cm[cm[\"label\"] == 1].copy()\n        cm = cm[(cm[\"width\"] > 0) & (cm[\"height\"] > 0)]\n\n        skipped = 0\n        for _, r in tqdm(cm.iterrows(), total=len(cm), desc=\"CXR Lung Node21\", leave=False):\n            img_path = os.path.join(cxrlung_root, r[\"img_name\"])\n            if not os.path.exists(img_path):\n                skipped += 1\n                continue\n\n            iw, ih = get_image_dims(img_path)\n            if iw is None:\n                continue\n\n            x1 = float(r[\"x\"]); y1 = float(r[\"y\"])\n            w  = float(r[\"width\"]); h = float(r[\"height\"])\n            x2 = x1 + w; y2 = y1 + h\n            if w < 4 or h < 4:\n                continue\n            ba = int(w * h)\n            if ba / (iw * ih) > 0.40:\n                continue\n\n            cxrlung_data.append({\n                \"image\":          f\"cxrlung_{r['img_name']}\",\n                \"image_path\":     img_path,\n                \"label\":          \"tumor_xray\",\n                \"modality\":       \"xray\",\n                \"source_dataset\": \"CXR_Lung_Node21\",\n                \"xmin\": int(x1), \"ymin\": int(y1), \"xmax\": int(x2), \"ymax\": int(y2),\n                \"width\": int(w), \"height\": int(h), \"area\": ba,\n                \"relative_area_pct\": round(ba/(iw*ih)*100, 2),\n                \"aspect_ratio\": round(w/max(h,1), 3),\n                \"center_x\": round((x1 + w/2)/iw, 3),\n                \"center_y\": round((y1 + h/2)/ih, 3),\n            })\n        print(f\"CXR Lung (Node21): {len(cxrlung_data)} nodule bboxes \"\n              f\"from {cm['img_name'].nunique()} images ({skipped} skipped)\")\n    else:\n        cxrlung_data = []\n        print(\"⚠ CXR Lung (Node21) not found.\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    siim_data = []\n    print(\"\\u23e9 SIIM-COVID19: skipped (cached)\")\nelse:\n    # ── 3.10 SIIM-FISABIO-RSNA COVID-19 — Opacity Detection ──────────────────────\n    import ast as _ast\n    siim_data = []\n    if siim_root is not None and siim_csv is not None and siim_csv.exists():\n        sdf = pd.read_csv(siim_csv)\n        print(f\"SIIM-COVID19: {len(sdf)} rows in CSV, {sdf['ImageID'].nunique()} unique images\")\n\n        # Build image_id -> file path index by globbing JPEG/train/<study>/<series>/<image>.jpg\n        siim_img_root = siim_root / \"JPEG\" / \"train\"\n        print(\"  Indexing SIIM image files...\")\n        siim_id2path = {}\n        for jpg in tqdm(list(siim_img_root.rglob(\"*.jpg\")), desc=\"SIIM glob\", leave=False):\n            siim_id2path[jpg.stem] = str(jpg)\n        print(f\"  Indexed {len(siim_id2path)} JPEG files\")\n\n        skipped_no_box = skipped_missing = skipped_parse = 0\n        for _, r in tqdm(sdf.iterrows(), total=len(sdf), desc=\"SIIM-COVID19\", leave=False):\n            # Filter: only positives (label starts with 'opacity'); skip 'none' rows\n            boxes_raw = r.get(\"boxes\")\n            if pd.isna(boxes_raw) or boxes_raw in (\"\", \"[]\"):\n                skipped_no_box += 1\n                continue\n            try:\n                boxes = _ast.literal_eval(boxes_raw)\n            except Exception:\n                skipped_parse += 1\n                continue\n            img_id = r[\"ImageID\"]\n            img_path = siim_id2path.get(img_id)\n            if img_path is None:\n                skipped_missing += 1\n                continue\n            iw, ih = get_image_dims(img_path)\n            if iw is None:\n                continue\n\n            for b in boxes:\n                # Coords in img_relative_*.csv are normalized [0,1]\n                x1 = float(b[\"x\"]) * iw\n                y1 = float(b[\"y\"]) * ih\n                x2 = (float(b[\"x\"]) + float(b[\"width\"])) * iw\n                y2 = (float(b[\"y\"]) + float(b[\"height\"])) * ih\n                \n                # Clamp coordinates to image boundaries\n                x1 = max(0.0, min(x1, float(iw)))\n                y1 = max(0.0, min(y1, float(ih)))\n                x2 = max(0.0, min(x2, float(iw)))\n                y2 = max(0.0, min(y2, float(ih)))\n                \n                w  = x2 - x1\n                h  = y2 - y1\n                if w < 4 or h < 4:\n                    continue\n                ba = int(w * h)\n                if ba / (iw * ih) > 0.40:\n                    continue\n\n                siim_data.append({\n                    \"image\":          f\"siim_{img_id}\",\n                    \"image_path\":     img_path,\n                    \"label\":          \"pneumonia\",\n                    \"modality\":       \"xray\",\n                    \"source_dataset\": \"SIIM_COVID19\",\n                    \"xmin\": int(x1), \"ymin\": int(y1), \"xmax\": int(x2), \"ymax\": int(y2),\n                    \"width\": int(w), \"height\": int(h), \"area\": ba,\n                    \"relative_area_pct\": round(ba / (iw * ih) * 100, 2),\n                    \"aspect_ratio\": round(w / max(h, 1), 3),\n                    \"center_x\": round((x1 + w / 2) / iw, 3),\n                    \"center_y\": round((y1 + h / 2) / ih, 3),\n                })\n\n        print(f\"SIIM-COVID19: {len(siim_data)} opacity bboxes\")\n        print(f\"  Skipped: {skipped_no_box} 'none' rows, {skipped_missing} missing files, {skipped_parse} parse errors\")\n    else:\n        siim_data = []\n        print(\"\\u26a0 SIIM-COVID19 not found.\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    jsrt_data = []\n    print(\"\\u23e9 JSRT: skipped (cached)\")\nelse:\n    # ── 3.11 JSRT — Japanese SRT Nodule Database (pristine) ──────────────────────\n    jsrt_data = []\n    if jsrt_root is not None and jsrt_meta is not None and jsrt_meta.exists():\n        jdf = pd.read_csv(jsrt_meta)\n        print(f\"JSRT: {len(jdf)} rows in metadata, {jdf['study_id'].nunique()} unique images\")\n\n        # JSRT geometry: original images are 2048x2048 with 0.175 mm/pixel.\n        # The 'size' column is nodule diameter in MM; 'x','y' are nodule CENTER in pixel\n        # coords of the 2048-reference. If the stored image is a different size we scale.\n        JSRT_ORIG_SIZE = 2048\n        JSRT_MM_PER_PX = 0.175\n\n        skipped_missing = 0\n        for _, r in tqdm(jdf.iterrows(), total=len(jdf), desc=\"JSRT\", leave=False):\n            img_path = jsrt_img_dir / r[\"study_id\"]\n            if not img_path.exists():\n                skipped_missing += 1\n                continue\n            iw, ih = get_image_dims(str(img_path))\n            if iw is None:\n                continue\n\n            # Scale factor from JSRT 2048-reference to actual stored image size\n            scale = iw / JSRT_ORIG_SIZE\n\n            try:\n                x_center = float(r[\"x\"]) * scale\n                y_center = float(r[\"y\"]) * scale\n                size_mm  = float(r[\"size\"])\n            except (TypeError, ValueError):\n                continue\n\n            size_px_native = size_mm / JSRT_MM_PER_PX  # pixels at 2048 reference\n            size_px = size_px_native * scale            # pixels at actual stored res\n\n            half = size_px / 2.0\n            x1 = max(0.0, x_center - half)\n            y1 = max(0.0, y_center - half)\n            x2 = min(float(iw), x_center + half)\n            y2 = min(float(ih), y_center + half)\n\n            w = x2 - x1\n            h = y2 - y1\n            if w < 4 or h < 4:\n                continue\n            ba = int(w * h)\n            if ba / (iw * ih) > 0.40:\n                continue\n\n            jsrt_data.append({\n                \"image\":          f\"jsrt_{r['study_id']}\",\n                \"image_path\":     str(img_path),\n                \"label\":          \"tumor_xray\",\n                \"modality\":       \"xray\",\n                \"source_dataset\": \"JSRT\",\n                \"xmin\": int(x1), \"ymin\": int(y1), \"xmax\": int(x2), \"ymax\": int(y2),\n                \"width\": int(w), \"height\": int(h), \"area\": ba,\n                \"relative_area_pct\": round(ba / (iw * ih) * 100, 2),\n                \"aspect_ratio\": 1.0,  # square bbox by construction (center+size only)\n                \"center_x\": round(x_center / iw, 3),\n                \"center_y\": round(y_center / ih, 3),\n            })\n\n        print(f\"JSRT: {len(jsrt_data)} nodule bboxes ({skipped_missing} skipped — file not found)\")\n    else:\n        jsrt_data = []\n        print(\"\\u26a0 JSRT not found.\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    print(f\"\\u26a1 Loading cached df from {DF_CACHE_PATH}\")\n    df = pd.read_parquet(DF_CACHE_PATH)\n    print(f\"   loaded {len(df):,} rows, {df['image'].nunique():,} unique images\")\n    print(f\"   class counts: {dict(df['label'].value_counts())}\")\nelse:\n    df = pd.DataFrame(tumor_data + tb_data + vindr_data + rsna_data + vinbigdata_data + lidc_data + chexdet_data + cxray14_data + cxrlung_data + siim_data + jsrt_data)\n\n    # ── Two-model split: keep only rows for the current modality ───────────\n    if len(df) > 0 and TARGET_MODALITY != \"all\":\n        _before = len(df)\n        if TARGET_MODALITY == \"xray\":\n            df = df[df[\"modality\"] == \"xray\"].reset_index(drop=True)\n        elif TARGET_MODALITY == \"ct\":\n            df = df[df[\"modality\"] == \"ct\"].reset_index(drop=True)\n        # Also drop labels that no longer belong in this subset (defensive)\n        df = df[df[\"label\"].isin(CLASS_NAMES)].reset_index(drop=True)\n        print(f\"[{TARGET_MODALITY}] df filtered: {_before} → {len(df)} rows\")\n\n    if len(df) > 0:\n        # ── v11: Excessive-bbox filter — drop entire images that have too many boxes\n        # for a class. Most legitimate cases have 1-5 boxes; outliers tend to be\n        # annotation errors, extreme metastasis, or images so noisy the radiologist\n        # boxed everything. Keeping these wastes capacity learning bad patterns.\n        MAX_BBOXES_PER_IMAGE = {\n            \"tumor_xray\":   8,    # typical nodule case = 1-3 boxes\n            \"tuberculosis\": 15,   # TB can have multiple foci\n            \"pneumonia\":    12,   # pneumonia regions can be split into a few\n        }\n        _before_excess = len(df)\n        _boxes_per_img = df.groupby([\"image\", \"label\"]).size().reset_index(name=\"n_boxes\")\n        _bad_images = set()\n        for _cls, _cap in MAX_BBOXES_PER_IMAGE.items():\n            _excess = _boxes_per_img[(_boxes_per_img[\"label\"] == _cls) & (_boxes_per_img[\"n_boxes\"] > _cap)]\n            if len(_excess) > 0:\n                _bad_images.update(_excess[\"image\"].tolist())\n                print(f\"  Excessive-bbox filter ({_cls} > {_cap}): {len(_excess)} images flagged\")\n        if _bad_images:\n            df = df[~df[\"image\"].isin(_bad_images)].reset_index(drop=True)\n            print(f\"  Dropped {len(_bad_images)} images with excessive bboxes \"\n                  f\"({_before_excess} -> {len(df)} rows)\")\n\n        # FIX ⑩: Tighter per-class area caps\n        AREA_CAPS = {\"tumor_ct\": 15.0, \"tumor_xray\": 15.0, \"tuberculosis\": 15.0, \"pneumonia\": 30.0}\n        _before_area = len(df)\n        for cls, cap in AREA_CAPS.items():\n            mask = (df[\"label\"] == cls) & (df[\"relative_area_pct\"] > cap)\n            n_drop = mask.sum()\n            if n_drop > 0:\n                df = df[~mask]\n                print(f\"  Area filter: dropped {n_drop} {cls} bboxes > {cap}% area\")\n        df = df.reset_index(drop=True)\n        if _before_area != len(df):\n            print(f\"  Area filter total: {_before_area} → {len(df)} ({_before_area - len(df)} removed)\")\n\n        # Post-resize tiny-bbox filter: drop bboxes that would be <MIN_DIM px after resize to train_sz.\n        # Uses the image's native dims (looked up once per unique image).\n        MIN_DIM_AT_RESIZE = 8\n        _uniq_paths = df[\"image_path\"].drop_duplicates().tolist()\n        _dim_map = {}\n        for _p in tqdm(_uniq_paths, desc=\"Image dim cache\"):\n            _dim_map[_p] = get_image_dims(_p)\n\n        _imw = df[\"image_path\"].map(lambda p: _dim_map.get(p, (None, None))[0])\n        _imh = df[\"image_path\"].map(lambda p: _dim_map.get(p, (None, None))[1])\n        df[\"img_w\"] = _imw\n        df[\"img_h\"] = _imh\n\n        import numpy as _np\n        _valid = df[\"img_w\"].notna() & df[\"img_h\"].notna()\n        _scale = _np.where(_valid,\n                           train_sz / _np.maximum(df[\"img_w\"].fillna(1).astype(float),\n                                                  df[\"img_h\"].fillna(1).astype(float)),\n                           1.0)\n        df[\"resized_w\"] = df[\"width\"].astype(float) * _scale\n        df[\"resized_h\"] = df[\"height\"].astype(float) * _scale\n\n        _before_tiny = len(df)\n        _tiny_mask = (df[\"resized_w\"] < MIN_DIM_AT_RESIZE) | (df[\"resized_h\"] < MIN_DIM_AT_RESIZE)\n        if _tiny_mask.any():\n            print(f\"  Tiny bbox filter (< {MIN_DIM_AT_RESIZE}px at {train_sz}px) by class:\")\n            for cls in CLASS_NAMES:\n                _c = ((df[\"label\"] == cls) & _tiny_mask).sum()\n                if _c > 0:\n                    print(f\"    {cls:<14}: {_c} dropped\")\n            df = df[~_tiny_mask].reset_index(drop=True)\n            print(f\"  Tiny bbox total: {_before_tiny} → {len(df)} ({_before_tiny - len(df)} removed)\")\n\n        print(\"=\" * 62)\n        print(\"UNIFIED DATASET SUMMARY\")\n        print(\"=\" * 62)\n        print(f\"Total annotations: {len(df)}, Unique images: {df['image'].nunique()}\")\n        print()\n        print(f\"  {'Label':<14} {'Modality':<10} {'Bboxes':>7} {'Images':>7}\")\n        print(\"-\" * 45)\n        for l in CLASS_NAMES:\n            s = df[df[\"label\"] == l]\n            if len(s) == 0:\n                continue\n            mods = \",\".join(sorted(s[\"modality\"].unique()))\n            print(f\"  {l:<14} {mods:<10} {len(s):>7} {s['image'].nunique():>7}\")\n        print(\"\\nBy source dataset:\")\n        for src, g in df.groupby(\"source_dataset\"):\n            print(f\"  {src:<24s}: {g['image'].nunique():5d} images, {len(g):6d} bboxes\")\n        display(df.head())\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T19:17:58.50681Z","iopub.execute_input":"2026-04-24T19:17:58.507094Z","iopub.status.idle":"2026-04-24T19:18:08.037406Z","shell.execute_reply.started":"2026-04-24T19:17:58.507062Z","shell.execute_reply":"2026-04-24T19:18:08.036758Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Dataset Validation & Analysis\nIn-depth audit of raw classes, bounding box coordinates, and spatial distributions\nto ensure correct class mapping and bbox placement before training.","metadata":{}},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# 4.0  RAW CLASS AUDIT — Verify that we selected the correct classes\n# ═══════════════════════════════════════════════════════════════════════════════\nprint(\"=\" * 70)\nprint(\"RAW CLASS AUDIT — Every class available in each dataset\")\nprint(\"=\" * 70)\n\n# ── VinBigData ─────────────────────────────────────────────────────────────\nif vinbigdata_csv and os.path.exists(vinbigdata_csv):\n    _vbd = pd.read_csv(vinbigdata_csv)\n    print(\"\\n┌─ VinBigData Chest X-ray (train.csv) ────────────────────────────\")\n    print(f\"│  Total rows: {len(_vbd)}\")\n    for cn, cnt in _vbd[\"class_name\"].value_counts().items():\n        mapped = VINBIGDATA_KEEP.get(cn, \"── SKIPPED ──\")\n        sym = \"✓\" if cn in VINBIGDATA_KEEP else \"✗\"\n        print(f\"│  {sym}  {cn:30s}  (n={cnt:6d})  →  {mapped}\")\n    print(f\"└──────────────────────────────────────────────────────────────────\")\n\n# ── LIDC-IDRI ──────────────────────────────────────────────────────────────\nif lidc_xml_dir and os.path.exists(lidc_xml_dir):\n    print(\"\\n┌─ LIDC-IDRI (XML annotations) ─────────────────────────────────\")\n    print(\"│  This dataset only has lung nodules on CT slices.\")\n    print(\"│  Nodules WITH <characteristics> (≥3mm): → tumor_ct  ✓\")\n    print(\"│  Nodules WITHOUT <characteristics>:     → SKIPPED   ✗\")\n    print(\"│  <nonNodule> entries:                   → SKIPPED   ✗\")\n    n_with = sum(1 for r in lidc_data if r[\"label\"] == \"tumor_ct\")\n    print(f\"│  Kept: {n_with} bbox annotations (tumor_ct)\")\n    print(f\"└──────────────────────────────────────────────────────────────────\")\n\n# ── ChestX-Det ─────────────────────────────────────────────────────────────\nif chexdet_train_json and os.path.exists(chexdet_train_json):\n    with open(chexdet_train_json, \"r\") as f:\n        _cd = json.load(f)\n    all_syms = {}\n    for entry in _cd:\n        for s in entry.get(\"syms\", []):\n            all_syms[s] = all_syms.get(s, 0) + 1\n    print(\"\\n┌─ ChestX-Det (ChestX_Det_train.json) ──────────────────────────\")\n    print(f\"│  Total images: {len(_cd)}, Total annotations: {sum(all_syms.values())}\")\n    for cn in sorted(all_syms.keys(), key=lambda x: -all_syms[x]):\n        cnt = all_syms[cn]\n        mapped = CHEXDET_KEEP.get(cn, \"── SKIPPED ──\")\n        sym = \"✓\" if cn in CHEXDET_KEEP else \"✗\"\n        print(f\"│  {sym}  {cn:30s}  (n={cnt:6d})  →  {mapped}\")\n    print(f\"└──────────────────────────────────────────────────────────────────\")\n\n# ── Summary ────────────────────────────────────────────────────────────────\nprint(\"\\n\" + \"=\" * 70)\nprint(\"FINAL CLASS SCHEME\")\nprint(\"=\" * 70)\nfor cls in CLASS_NAMES:\n    n = len(df[df[\"label\"] == cls]) if len(df) > 0 else 0\n    sources = \", \".join(sorted(df[df[\"label\"] == cls][\"source_dataset\"].unique())) if n > 0 else \"none\"\n    print(f\"  {cls:14s}  →  {n:6d} bboxes  from: {sources}\")\n\n# -- cxray14 ------------------------------------------------------------\nif cxray14_root and os.path.exists(cxray14_root):\n    _cxr_rows = df[df[\"source_dataset\"] == \"cxray14\"] if len(df) > 0 else None\n    print(\"\\n+- cxray14 (VinBig + NH-Xray, YOLO-format) ---------------------\")\n    if _cxr_rows is not None and len(_cxr_rows) > 0:\n        print(f\"|  Kept boxes: {len(_cxr_rows)}  (all mapped to tumor_xray)\")\n        print(f\"|  Unique images with >=1 kept box: {_cxr_rows['image'].nunique()}\")\n        print(\"|  Strict mapping: only class 8 Nodule/Mass -> tumor_xray\")\n        print(\"|  (Consolidation/Opacity/Infiltration intentionally dropped per v7)\")\n    else:\n        print(\"|  No kept boxes.\")\n    print(\"+--------------------------------------------------------------\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T19:18:52.312578Z","iopub.execute_input":"2026-04-24T19:18:52.313296Z","iopub.status.idle":"2026-04-24T19:18:52.536935Z","shell.execute_reply.started":"2026-04-24T19:18:52.313262Z","shell.execute_reply":"2026-04-24T19:18:52.536108Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# 4.1  BBOX COORDINATE SANITY — Check pixel coords vs actual image sizes\n# ═══════════════════════════════════════════════════════════════════════════════\nfrom PIL import Image\n\nif len(df) > 0:\n    print(\"=\" * 70)\n    print(\"BBOX COORDINATE SANITY CHECK\")\n    print(\"=\" * 70)\n    print(\"Sampling images from each source dataset to verify bbox coords...\\n\")\n\n    for src_ds in sorted(df[\"source_dataset\"].unique()):\n        sub = df[df[\"source_dataset\"] == src_ds]\n        print(f\"┌─ {src_ds} ({len(sub)} bboxes, {sub['image'].nunique()} images) ─────\")\n\n        # Sample up to 5 unique images that exist\n        sampled = 0\n        issues = 0\n        coords_ok = 0\n        max_x_ratio = 0\n        max_y_ratio = 0\n\n        sample_imgs = sub.drop_duplicates(\"image\").head(20)\n        for _, row in sample_imgs.iterrows():\n            ip = row[\"image_path\"]\n            if not os.path.exists(ip):\n                continue\n\n            # Get real image dimensions\n            try:\n                if ip.lower().endswith((\".dcm\", \".dicom\")):\n                    try:\n                        import pydicom\n                        ds = pydicom.dcmread(ip, stop_before_pixels=True)\n                        real_w, real_h = int(ds.Columns), int(ds.Rows)\n                    except Exception:\n                        real_w, real_h = None, None\n                else:\n                    with Image.open(ip) as img:\n                        real_w, real_h = img.size\n            except Exception:\n                continue\n\n            if real_w is None:\n                continue\n\n            sampled += 1\n            # Check all bboxes for this image\n            img_rows = sub[sub[\"image\"] == row[\"image\"]]\n            for _, br in img_rows.iterrows():\n                x1, y1, x2, y2 = br[\"xmin\"], br[\"ymin\"], br[\"xmax\"], br[\"ymax\"]\n\n                # Check if coords are within image bounds\n                x_ratio = max(x2 / real_w, 0)\n                y_ratio = max(y2 / real_h, 0)\n                max_x_ratio = max(max_x_ratio, x_ratio)\n                max_y_ratio = max(max_y_ratio, y_ratio)\n\n                if x2 > real_w or y2 > real_h or x1 < 0 or y1 < 0:\n                    issues += 1\n                    if issues <= 3:\n                        print(f\"│  ⚠ OUT OF BOUNDS: {row['image']}  \"\n                              f\"bbox=({x1},{y1},{x2},{y2})  img=({real_w}×{real_h})\")\n                else:\n                    coords_ok += 1\n\n            if sampled >= 5:\n                break\n\n        if sampled == 0:\n            print(f\"│  ⚠ Could not read any sample images\")\n        else:\n            print(f\"│  Sampled {sampled} images:\")\n            print(f\"│  Image dimensions: {real_w}×{real_h}\")\n            print(f\"│  Coords OK: {coords_ok},  Out of bounds: {issues}\")\n            print(f\"│  Max x_ratio: {max_x_ratio:.3f},  Max y_ratio: {max_y_ratio:.3f}\")\n            if max_x_ratio > 1.0 or max_y_ratio > 1.0:\n                print(f\"│  ⚠⚠ BBOX COORDINATES EXCEED IMAGE SIZE! Possible coord system mismatch!\")\n            else:\n                print(f\"│  ✓ All bbox coords within image bounds\")\n        print(f\"└──────────────────────────────────────────────────────────────────\\n\")\n\n    # ── Distribution analysis ──────────────────────────────────────────────\n    print(\"\\n\" + \"=\" * 70)\n    print(\"BBOX SPATIAL DISTRIBUTION (per source dataset)\")\n    print(\"=\" * 70)\n    for src_ds in sorted(df[\"source_dataset\"].unique()):\n        sub = df[df[\"source_dataset\"] == src_ds]\n        if len(sub) == 0:\n            continue\n        cx = sub[\"center_x\"]\n        cy = sub[\"center_y\"]\n        print(f\"\\n  {src_ds}:\")\n        print(f\"    center_x  mean={cx.mean():.3f}  std={cx.std():.3f}  \"\n              f\"min={cx.min():.3f}  max={cx.max():.3f}\")\n        print(f\"    center_y  mean={cy.mean():.3f}  std={cy.std():.3f}  \"\n              f\"min={cy.min():.3f}  max={cy.max():.3f}\")\n        # Check for bias (center_x should be ~0.5 for balanced left/right)\n        if cx.mean() > 0.65:\n            print(f\"    ⚠ BIAS: center_x mean={cx.mean():.3f} > 0.65 — bboxes skewed RIGHT!\")\n        elif cx.mean() < 0.35:\n            print(f\"    ⚠ BIAS: center_x mean={cx.mean():.3f} < 0.35 — bboxes skewed LEFT!\")\n        else:\n            print(f\"    ✓ No obvious left/right bias\")\n\n        # Quadrant analysis\n        q_tl = ((cx < 0.5) & (cy < 0.5)).sum()\n        q_tr = ((cx >= 0.5) & (cy < 0.5)).sum()\n        q_bl = ((cx < 0.5) & (cy >= 0.5)).sum()\n        q_br = ((cx >= 0.5) & (cy >= 0.5)).sum()\n        total = len(sub)\n        print(f\"    Quadrants: TL={q_tl}({q_tl/total*100:.0f}%) \"\n              f\"TR={q_tr}({q_tr/total*100:.0f}%) \"\n              f\"BL={q_bl}({q_bl/total*100:.0f}%) \"\n              f\"BR={q_br}({q_br/total*100:.0f}%)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T19:19:39.199778Z","iopub.execute_input":"2026-04-24T19:19:39.200281Z","iopub.status.idle":"2026-04-24T19:19:39.449032Z","shell.execute_reply.started":"2026-04-24T19:19:39.200252Z","shell.execute_reply":"2026-04-24T19:19:39.44844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ═══════════════════════════════════════════════════════════════════════════════\n# 4.3  SPATIAL DISTRIBUTION PLOTS — Bbox center heatmaps per dataset\n# ═══════════════════════════════════════════════════════════════════════════════\nimport matplotlib.pyplot as plt\nimport numpy as np\n\nif len(df) > 0:\n    source_datasets = sorted(df[\"source_dataset\"].unique())\n    n = len(source_datasets)\n    fig, axes = plt.subplots(2, n, figsize=(5 * n, 10))\n    if n == 1:\n        axes = axes.reshape(-1, 1)\n    fig.suptitle(\"BBOX Spatial Distribution — Center positions & Size distributions\",\n                 fontsize=14, fontweight=\"bold\")\n\n    for col_idx, src_ds in enumerate(source_datasets):\n        sub = df[df[\"source_dataset\"] == src_ds]\n\n        # Row 1: Center position heatmap\n        ax = axes[0, col_idx]\n        cx = sub[\"center_x\"].clip(0, 1)\n        cy = sub[\"center_y\"].clip(0, 1)\n        ax.hist2d(cx, cy, bins=20, cmap=\"YlOrRd\", range=[[0, 1], [0, 1]])\n        ax.set_xlabel(\"center_x (normalized)\")\n        ax.set_ylabel(\"center_y (normalized)\")\n        ax.set_title(f\"{src_ds}\\n({len(sub)} bboxes)\", fontsize=10)\n        ax.set_aspect(\"equal\")\n        ax.invert_yaxis()\n        # Draw crosshairs at center\n        ax.axhline(0.5, color=\"cyan\", linestyle=\"--\", alpha=0.5)\n        ax.axvline(0.5, color=\"cyan\", linestyle=\"--\", alpha=0.5)\n\n        # Row 2: Width/Height scatter\n        ax2 = axes[1, col_idx]\n        w = sub[\"width\"]\n        h = sub[\"height\"]\n        ax2.scatter(w, h, alpha=0.3, s=5, c=\"dodgerblue\")\n        ax2.set_xlabel(\"bbox width (px)\")\n        ax2.set_ylabel(\"bbox height (px)\")\n        ax2.set_title(f\"Size: w={w.mean():.0f}±{w.std():.0f}, \"\n                       f\"h={h.mean():.0f}±{h.std():.0f}\", fontsize=9)\n        ax2.set_aspect(\"equal\")\n\n    plt.tight_layout()\n    plt.savefig(\"/kaggle/working/spatial_distribution_plot1.png\", dpi=300, bbox_inches=\"tight\")\n    plt.show()\n\n    # ── Relative area distribution ─────────────────────────────────────────\n    fig, axes = plt.subplots(1, n, figsize=(5 * n, 4))\n    if n == 1:\n        axes = [axes]\n    fig.suptitle(\"Relative Area Distribution (bbox_area / image_area × 100%)\",\n                 fontsize=14, fontweight=\"bold\")\n    for i, src_ds in enumerate(source_datasets):\n        sub = df[df[\"source_dataset\"] == src_ds]\n        ax = axes[i]\n        ax.hist(sub[\"relative_area_pct\"], bins=50, color=\"coral\", edgecolor=\"black\",\n                alpha=0.7)\n        ax.axvline(40.0, color=\"red\", linestyle=\"--\", linewidth=2, label=\"40% cutoff\")\n        ax.set_xlabel(\"Relative area %\")\n        ax.set_ylabel(\"Count\")\n        ax.set_title(f\"{src_ds}\\nmedian={sub['relative_area_pct'].median():.1f}%\",\n                     fontsize=10)\n        ax.legend()\n    plt.tight_layout()\n    plt.savefig(\"/kaggle/working/spatial_distribution_plot2.png\", dpi=300, bbox_inches=\"tight\")\n    plt.show()\n\n\n# ═══════════════════════════════════════════════════════════════════════════════\n# 4.4  PER-DATASET DEEP DIVE — Individual analysis for every source dataset\n# ═══════════════════════════════════════════════════════════════════════════════\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nfrom PIL import Image\nimport numpy as np\n\nif len(df) == 0:\n    print(\"⚠ df is empty — no data to analyze\")\nelse:\n    source_datasets = sorted(df[\"source_dataset\"].unique())\n    print(f\"Total source datasets in df: {len(source_datasets)}\")\n    print(f\"  {source_datasets}\\n\")\n\n    pal_map = {\"tumor_ct\": \"#c0392b\", \"tumor_xray\": \"#e67e22\",\n               \"tuberculosis\": \"#3498db\", \"pneumonia\": \"#27ae60\"}\n\n    for ds_idx, src_ds in enumerate(source_datasets):\n        sub = df[df[\"source_dataset\"] == src_ds]\n        print(\"=\" * 78)\n        print(f\"  DATASET {ds_idx + 1}/{len(source_datasets)}: {src_ds}\")\n        print(\"=\" * 78)\n\n        # ── Basic stats ────────────────────────────────────────────────────\n        n_imgs = sub[\"image\"].nunique()\n        n_boxes = len(sub)\n        modalities = sub[\"modality\"].unique().tolist()\n        labels = sub[\"label\"].value_counts().to_dict()\n\n        print(f\"  Images:     {n_imgs}\")\n        print(f\"  Bboxes:     {n_boxes}\")\n        print(f\"  Modality:   {modalities}\")\n        print(f\"  Labels:     {labels}\")\n        print(f\"  Bboxes/img: {n_boxes / max(n_imgs, 1):.2f}\")\n\n        # ── Raw data sample ────────────────────────────────────────────────\n        print(f\"\\n  ── First 3 entries ──────────────────────────────────────────\")\n        for j, (_, row) in enumerate(sub.head(3).iterrows()):\n            print(f\"  [{j}] image={row['image']}, label={row['label']}, \"\n                  f\"bbox=({row['xmin']},{row['ymin']},{row['xmax']},{row['ymax']}), \"\n                  f\"w={row['width']}, h={row['height']}, \"\n                  f\"area_pct={row['relative_area_pct']}%\")\n            print(f\"       center=({row['center_x']},{row['center_y']}), \"\n                  f\"ar={row['aspect_ratio']}, path_exists={os.path.exists(row['image_path'])}\")\n\n        # ── Image dimension validation ─────────────────────────────────────\n        print(f\"\\n  ── Image dimension check (up to 10 images) ─────────────────\")\n        dim_issues = 0\n        dim_checked = 0\n        real_dims = []\n        sample_imgs = sub.drop_duplicates(\"image\").head(10)\n        for _, row in sample_imgs.iterrows():\n            ip = row[\"image_path\"]\n            if not os.path.exists(ip):\n                print(f\"  ⚠ FILE MISSING: {ip}\")\n                continue\n\n            try:\n                if ip.lower().endswith((\".dcm\", \".dicom\")):\n                    try:\n                        import pydicom\n                        ds = pydicom.dcmread(ip, stop_before_pixels=True)\n                        rw, rh = int(ds.Columns), int(ds.Rows)\n                    except Exception:\n                        rw, rh = None, None\n                else:\n                    with Image.open(ip) as img:\n                        rw, rh = img.size\n            except Exception:\n                rw, rh = None, None\n\n            if rw is None:\n                continue\n\n            dim_checked += 1\n            real_dims.append((rw, rh))\n\n            # Check ALL bboxes for this image\n            img_boxes = sub[sub[\"image\"] == row[\"image\"]]\n            for _, br in img_boxes.iterrows():\n                x1, y1, x2, y2 = br[\"xmin\"], br[\"ymin\"], br[\"xmax\"], br[\"ymax\"]\n                oob = \"\"\n                if x1 < 0: oob += \" x1<0\"\n                if y1 < 0: oob += \" y1<0\"\n                if x2 > rw: oob += f\" x2({x2})>w({rw})\"\n                if y2 > rh: oob += f\" y2({y2})>h({rh})\"\n                if x1 >= x2: oob += \" x1>=x2!\"\n                if y1 >= y2: oob += \" y1>=y2!\"\n                if oob:\n                    dim_issues += 1\n                    if dim_issues <= 5:\n                        print(f\"  ⚠ {row['image']}: bbox=({x1},{y1},{x2},{y2}) \"\n                              f\"img=({rw}×{rh}) ISSUES:{oob}\")\n\n        if dim_checked > 0:\n            dims_str = \", \".join([f\"{w}×{h}\" for w, h in set(real_dims)])\n            print(f\"  Image sizes found: {dims_str}\")\n            print(f\"  Checked {dim_checked} images, {dim_issues} bbox issues\")\n            if dim_issues == 0:\n                print(f\"  ✓ All bboxes within image bounds\")\n        else:\n            print(f\"  ⚠ Could not check any image dimensions\")\n\n        # ── Bbox statistics ────────────────────────────────────────────────\n        print(f\"\\n  ── BBox statistics ──────────────────────────────────────────\")\n        for stat_col in [\"width\", \"height\", \"relative_area_pct\", \"aspect_ratio\",\n                         \"center_x\", \"center_y\"]:\n            s = sub[stat_col]\n            print(f\"  {stat_col:20s}  mean={s.mean():8.2f}  std={s.std():8.2f}  \"\n                  f\"min={s.min():8.2f}  max={s.max():8.2f}\")\n\n        # ── Left/Right bias check ──────────────────────────────────────────\n        cx = sub[\"center_x\"]\n        left_pct = (cx < 0.5).sum() / len(cx) * 100\n        right_pct = (cx >= 0.5).sum() / len(cx) * 100\n        print(f\"\\n  Left vs Right: {left_pct:.1f}% left, {right_pct:.1f}% right\")\n        if abs(left_pct - 50) > 20:\n            print(f\"  ⚠⚠ STRONG SPATIAL BIAS DETECTED!\")\n        elif abs(left_pct - 50) > 10:\n            print(f\"  ⚠ Moderate spatial bias\")\n        else:\n            print(f\"  ✓ Balanced left/right\")\n\n        # ── Visualization: 3 sample images ─────────────────────────────────\n        print(f\"\\n  ── Visualizing 3 sample images with bboxes ─────────────────\")\n        fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n        fig.suptitle(f\"Dataset: {src_ds} — Sample Images with BBoxes\",\n                     fontsize=14, fontweight=\"bold\")\n\n        shown = 0\n        for _, img_row in sample_imgs.iterrows():\n            if shown >= 3:\n                break\n            ip = img_row[\"image_path\"]\n            if not os.path.exists(ip):\n                continue\n\n            ax = axes[shown]\n            try:\n                if ip.lower().endswith((\".dcm\", \".dicom\")):\n                    try:\n                        import pydicom\n                        ds = pydicom.dcmread(ip)\n                        pixel = ds.pixel_array.astype(float)\n                        if hasattr(ds, \"WindowCenter\"):\n                            wc = ds.WindowCenter\n                            ww = ds.WindowWidth\n                            if isinstance(wc, pydicom.multival.MultiValue):\n                                wc, ww = float(wc[0]), float(ww[0])\n                            else:\n                                wc, ww = float(wc), float(ww)\n                            lo, hi = wc - ww / 2, wc + ww / 2\n                            pixel = np.clip(pixel, lo, hi)\n                        pixel = ((pixel - pixel.min()) /\n                                 (pixel.max() - pixel.min() + 1e-8) * 255).astype(np.uint8)\n                        ax.imshow(pixel, cmap=\"gray\")\n                        img_h, img_w = pixel.shape[:2]\n                    except Exception as e:\n                        ax.text(0.5, 0.5, f\"DICOM error:\\n{e}\", ha=\"center\",\n                                va=\"center\", transform=ax.transAxes, fontsize=8)\n                        shown += 1\n                        continue\n                else:\n                    pil_img = Image.open(ip).convert(\"RGB\")\n                    arr = np.array(pil_img)\n                    ax.imshow(arr)\n                    img_h, img_w = arr.shape[:2]\n            except Exception as e:\n                ax.text(0.5, 0.5, f\"Error: {e}\", ha=\"center\",\n                        va=\"center\", transform=ax.transAxes, fontsize=8)\n                shown += 1\n                continue\n\n            # Draw bboxes\n            img_bboxes = sub[sub[\"image\"] == img_row[\"image\"]]\n            for _, br in img_bboxes.iterrows():\n                x1, y1, x2, y2 = br[\"xmin\"], br[\"ymin\"], br[\"xmax\"], br[\"ymax\"]\n                color = pal_map.get(br[\"label\"], \"#ffffff\")\n                rect = mpatches.FancyBboxPatch(\n                    (x1, y1), x2 - x1, y2 - y1,\n                    linewidth=2.5, edgecolor=color, facecolor=\"none\",\n                    boxstyle=\"round,pad=0\")\n                ax.add_patch(rect)\n                ax.text(x1 + 2, y1 - 5, f\"{br['label']}\",\n                        fontsize=8, color=\"white\", fontweight=\"bold\",\n                        bbox=dict(boxstyle=\"round,pad=0.2\", fc=color, alpha=0.85))\n\n            ax.set_title(f\"{img_row['image']}\\n\"\n                         f\"({img_w}×{img_h}px, {len(img_bboxes)} boxes)\",\n                         fontsize=9)\n            ax.axis(\"off\")\n            shown += 1\n\n        for j in range(shown, 3):\n            axes[j].text(0.5, 0.5, \"No image available\",\n                         ha=\"center\", va=\"center\",\n                         transform=axes[j].transAxes)\n            axes[j].axis(\"off\")\n\n        plt.tight_layout()\n        plt.savefig(\"/kaggle/working/spatial_distribution_plot3.png\", dpi=300, bbox_inches=\"tight\")\n        plt.show()\n        print()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T19:40:59.631967Z","iopub.execute_input":"2026-04-24T19:40:59.632205Z","iopub.status.idle":"2026-04-24T19:41:31.371889Z","shell.execute_reply.started":"2026-04-24T19:40:59.632177Z","shell.execute_reply":"2026-04-24T19:41:31.371162Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4.5 Ground-Truth Visual Audit (per dataset)\n\nBefore training, render random images from each source dataset with their ground-truth boxes overlaid. Use this to spot:\n\n- Boxes in the wrong place (cropping / coordinate-system bugs)\n- Wrong class labels (mapping bugs in Cell 6 KEEP dicts)\n- Multiple overlapping boxes on the same lesion (multi-radiologist annotation — addressed in 4.6)\n\nRun this cell, flip through the figures (one per dataset), and verify visually.\n","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport matplotlib.patches as _mpatches\nimport numpy as np\nimport cv2\n\n_CLASS_COLORS = {\n    \"tumor_xray\":   \"#ff3030\",\n    \"tumor_ct\":     \"#ff6030\",\n    \"tuberculosis\": \"#ff9900\",\n    \"pneumonia\":    \"#00e0ff\",\n}\n\ndef viz_gt_per_dataset(df, samples_per_ds=9, target_size=512, seed=42):\n    \"\"\"One figure per source dataset, with GT boxes + labels. Interactive scroll.\"\"\"\n    datasets = sorted(df[\"source_dataset\"].unique())\n    for ds in datasets:\n        ds_df = df[df[\"source_dataset\"] == ds]\n        imgs = ds_df[\"image\"].unique()\n        if len(imgs) == 0:\n            continue\n        n = min(samples_per_ds, len(imgs))\n        rng = np.random.RandomState(seed + hash(ds) % 1000)\n        picks = rng.choice(imgs, size=n, replace=False)\n        cols = 3\n        rows = (n + cols - 1) // cols\n        fig, axes = plt.subplots(rows, cols, figsize=(cols * 4.2, rows * 4.2))\n        axes = np.array(axes).reshape(rows, cols)\n\n        n_imgs_tot = ds_df[\"image\"].nunique()\n        n_boxes_tot = len(ds_df)\n        cls_counts = ds_df[\"label\"].value_counts().to_dict()\n        fig.suptitle(\n            f\"{ds.upper()}   |   {n_imgs_tot} imgs, {n_boxes_tot} boxes   |   classes: {cls_counts}\",\n            fontsize=12, y=1.0)\n\n        for idx in range(rows * cols):\n            ax = axes[idx // cols, idx % cols]\n            if idx >= n:\n                ax.axis(\"off\"); continue\n            img_key = picks[idx]\n            img_rows = df[(df[\"source_dataset\"] == ds) & (df[\"image\"] == img_key)]\n            img_path = img_rows[\"image_path\"].iloc[0]\n            img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                ax.text(0.5, 0.5, f\"READ FAIL:\\n{img_path[-40:]}\",\n                        ha=\"center\", va=\"center\", fontsize=7)\n                ax.axis(\"off\"); continue\n            H0, W0 = img.shape[:2]\n            img_r = cv2.resize(img, (target_size, target_size))\n            sx, sy = target_size / W0, target_size / H0\n            ax.imshow(img_r, cmap=\"gray\")\n            for _, r in img_rows.iterrows():\n                col = _CLASS_COLORS.get(r[\"label\"], \"#ffff00\")\n                rect = _mpatches.Rectangle(\n                    (r[\"xmin\"] * sx, r[\"ymin\"] * sy),\n                    r[\"width\"] * sx, r[\"height\"] * sy,\n                    linewidth=1.7, edgecolor=col, facecolor=\"none\")\n                ax.add_patch(rect)\n                ax.text(r[\"xmin\"] * sx, max(r[\"ymin\"] * sy - 2, 10), r[\"label\"],\n                        color=col, fontsize=8,\n                        bbox=dict(facecolor=\"black\", alpha=0.55,\n                                  edgecolor=\"none\", pad=0.6))\n            ax.set_title(f\"{str(img_key)[:38]}  |  n={len(img_rows)}\", fontsize=8)\n            ax.axis(\"off\")\n        plt.tight_layout()\n        plt.savefig(\"/kaggle/{img}.png\", dpi=300, bbox_inches=\"tight\")\n        plt.show()\n\nviz_gt_per_dataset(df, samples_per_ds=9)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T19:51:54.562597Z","iopub.execute_input":"2026-04-24T19:51:54.562862Z","iopub.status.idle":"2026-04-24T19:52:47.295584Z","shell.execute_reply.started":"2026-04-24T19:51:54.562836Z","shell.execute_reply":"2026-04-24T19:52:47.294476Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4.6 Multi-Radiologist Consensus Merge\n\n**VinDr-CXR and VinBigData** each have **3 radiologists independently annotating the same image**. One real lesion therefore appears as 2–3 overlapping GT boxes with the same label. This over-counts ground truth:\n\n- SimOTA (YOLOX's anchor assigner) gets ambiguous targets\n- Loss is inflated — the model is penalized 3× for missing one lesion\n- mAP evaluation against redundant GT is noisy\n\n**Fix:** group boxes by `(image, label)` and merge any pair with IoU > 0.3 via weighted-box-fusion (coordinate averaging). 3 rad boxes on one nodule → 1 consensus box.\n\nOther datasets (RSNA, TBX11K, ChestX-Det, Lung Tumor CT) are single-annotator and are passed through unchanged.\n\nToggle `APPLY_CONSENSUS_MERGE = False` to skip (ablation only).\n","metadata":{}},{"cell_type":"code","source":"APPLY_CONSENSUS_MERGE = True\nCONSENSUS_IOU = 0.3\n\n# Which source_datasets are multi-radiologist (3+ annotators per image)\ndef _is_multirad_source(ds_name):\n    if isinstance(ds_name, str):\n        if ds_name.startswith(\"VinDr-CXR\"):  # VinDr-CXR_train / _val / _test\n            return True\n        if ds_name == \"VinBigData_CXR\":\n            return True\n        if ds_name == \"cxray14\":  # VinBig-derived; may carry unmerged rad duplicates\n            return True\n    return False\n\ndef _iou_xyxy(a, b):\n    xA, yA = max(a[0], b[0]), max(a[1], b[1])\n    xB, yB = min(a[2], b[2]), min(a[3], b[3])\n    if xB <= xA or yB <= yA:\n        return 0.0\n    inter = (xB - xA) * (yB - yA)\n    ua = (a[2] - a[0]) * (a[3] - a[1])\n    ub = (b[2] - b[0]) * (b[3] - b[1])\n    return inter / max(ua + ub - inter, 1e-9)\n\ndef _merge_cluster_xyxy(boxes_xyxy, iou_thresh):\n    \"\"\"Greedy: cluster boxes with pairwise IoU > thresh, average each cluster.\"\"\"\n    if len(boxes_xyxy) <= 1:\n        return boxes_xyxy\n    used = [False] * len(boxes_xyxy)\n    merged = []\n    for i in range(len(boxes_xyxy)):\n        if used[i]:\n            continue\n        cluster = [boxes_xyxy[i]]\n        used[i] = True\n        for j in range(i + 1, len(boxes_xyxy)):\n            if used[j]:\n                continue\n            if _iou_xyxy(boxes_xyxy[i], boxes_xyxy[j]) > iou_thresh:\n                cluster.append(boxes_xyxy[j])\n                used[j] = True\n        avg = np.mean(cluster, axis=0)  # (4,)\n        merged.append((float(avg[0]), float(avg[1]), float(avg[2]), float(avg[3])))\n    return merged\n\ndef consensus_merge_multirad(df, iou_thresh=0.3):\n    if not APPLY_CONSENSUS_MERGE:\n        print(\"Consensus merge SKIPPED (APPLY_CONSENSUS_MERGE=False)\")\n        return df\n    before = len(df)\n    ds_mask = df[\"source_dataset\"].apply(_is_multirad_source)\n    other_df  = df[~ds_mask].copy()\n    target_df = df[ds_mask]\n    new_rows = []\n    stats = {}  # source_dataset -> [orig, merged]\n    for (img, lbl), grp in target_df.groupby([\"image\", \"label\"], sort=False):\n        orig_xyxy = [(r.xmin, r.ymin, r.xmax, r.ymax) for _, r in grp.iterrows()]\n        merged_xyxy = _merge_cluster_xyxy(orig_xyxy, iou_thresh)\n        ds = grp[\"source_dataset\"].iloc[0]\n        stats.setdefault(ds, [0, 0])\n        stats[ds][0] += len(orig_xyxy)\n        stats[ds][1] += len(merged_xyxy)\n        template = grp.iloc[0]\n        for (x1, y1, x2, y2) in merged_xyxy:\n            nr = template.copy()\n            w, h = x2 - x1, y2 - y1\n            nr[\"xmin\"] = int(x1); nr[\"ymin\"] = int(y1)\n            nr[\"xmax\"] = int(x2); nr[\"ymax\"] = int(y2)\n            nr[\"width\"] = int(w); nr[\"height\"] = int(h)\n            nr[\"area\"]  = int(w * h)\n            # Optional fields: update if present in schema\n            if \"img_w\" in nr.index and \"img_h\" in nr.index and nr[\"img_w\"] and nr[\"img_h\"]:\n                nr[\"relative_area_pct\"] = round((w * h) / (nr[\"img_w\"] * nr[\"img_h\"]) * 100, 2)\n            if \"aspect_ratio\" in nr.index:\n                nr[\"aspect_ratio\"] = round(w / max(h, 1), 3)\n            new_rows.append(nr)\n    merged_df = pd.DataFrame(new_rows) if new_rows else target_df.iloc[0:0]\n    out = pd.concat([other_df, merged_df], ignore_index=True).reset_index(drop=True)\n    print(\"=\" * 62)\n    print(f\"Multi-radiologist consensus merge  (IoU > {iou_thresh})\")\n    print(\"=\" * 62)\n    print(f\"  Total rows: {before}  ->  {len(out)}\")\n    for ds, (orig_n, new_n) in sorted(stats.items()):\n        dropped = orig_n - new_n\n        pct = 100 * dropped / max(orig_n, 1)\n        print(f\"  {ds:<20}: {orig_n:>6} boxes -> {new_n:>6}  \"\n              f\"({dropped:>5} duplicates removed, {pct:>5.1f}%)\")\n    print()\n    print(\"Per-class counts after merge:\")\n    for cls, cnt in out[\"label\"].value_counts().items():\n        print(f\"  {cls:<14}: {cnt}\")\n    return out\n\ndf = consensus_merge_multirad(df, iou_thresh=CONSENSUS_IOU)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T19:58:02.119317Z","iopub.execute_input":"2026-04-24T19:58:02.119553Z","iopub.status.idle":"2026-04-24T19:58:03.671885Z","shell.execute_reply.started":"2026-04-24T19:58:02.119528Z","shell.execute_reply":"2026-04-24T19:58:03.671096Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4.7 Data Integrity Audit (programmatic)\n\nStrict bbox and file-level checks. Flags anything that would silently mis-train: invalid box coords, boxes escaping image bounds, degenerate boxes, missing image files, duplicate boxes, labels outside `CLASS_NAMES`. Runs in seconds.\n","metadata":{}},{"cell_type":"code","source":"def data_integrity_audit(df, sample_path_check=500):\n    \"\"\"Strict programmatic checks. Prints a PASS/FAIL matrix + per-dataset breakdown.\"\"\"\n    print(\"=\" * 72)\n    print(\"DATA INTEGRITY AUDIT\")\n    print(\"=\" * 72)\n    print(f\"Total rows: {len(df):,}  |  Unique images: {df['image'].nunique():,}\")\n    print(f\"Source datasets: {sorted(df['source_dataset'].unique())}\\n\")\n\n    checks = {}\n\n    # 1. Invalid box coords (xmin >= xmax or ymin >= ymax)\n    checks[\"invalid_box_coords\"]       = int(((df[\"xmin\"] >= df[\"xmax\"]) | (df[\"ymin\"] >= df[\"ymax\"])).sum())\n    # 2. Negative coords\n    checks[\"negative_coords\"]          = int(((df[\"xmin\"] < 0) | (df[\"ymin\"] < 0)).sum())\n    # 3. Out-of-bounds (box extends past image dims)\n    if \"img_w\" in df.columns and \"img_h\" in df.columns:\n        oob = df[\"img_w\"].notna() & df[\"img_h\"].notna() & (\n            (df[\"xmax\"] > df[\"img_w\"]) | (df[\"ymax\"] > df[\"img_h\"])\n        )\n        checks[\"boxes_outside_img_bounds\"] = int(oob.sum())\n    # 4. Degenerate boxes (<4 px on a side)\n    checks[\"boxes_<4px_side\"]          = int(((df[\"width\"] < 4) | (df[\"height\"] < 4)).sum())\n    # 5. Whole-image-ish boxes (>40% relative area)\n    checks[\"boxes_>40pct_area\"]        = int((df[\"relative_area_pct\"] > 40).sum())\n    # 6. Suspicious 0,0 anchor (box starts exactly at origin — common parse bug)\n    zero_anchor = (df[\"xmin\"] == 0) & (df[\"ymin\"] == 0) & (df[\"relative_area_pct\"] > 10)\n    checks[\"zero_anchor_large_box\"]    = int(zero_anchor.sum())\n    # 7. Duplicate boxes within same (image, label)\n    dup_mask = df.duplicated(subset=[\"image\", \"label\", \"xmin\", \"ymin\", \"xmax\", \"ymax\"], keep=False)\n    checks[\"exact_duplicate_boxes\"]    = int(dup_mask.sum())\n    # 8. Labels outside CLASS_NAMES\n    checks[\"labels_outside_CLASS_NAMES\"] = int((~df[\"label\"].isin(CLASS_NAMES)).sum())\n    # 9. Rows missing image_path file (sample)\n    uniq_paths = df[\"image_path\"].drop_duplicates()\n    if len(uniq_paths) > sample_path_check:\n        uniq_paths = uniq_paths.sample(sample_path_check, random_state=42)\n    missing = [p for p in uniq_paths if not os.path.exists(p)]\n    checks[f\"missing_image_files_(sample_{min(sample_path_check, df['image_path'].nunique())})\"] = len(missing)\n\n    print(\"Integrity checks:\")\n    print(\"-\" * 72)\n    for k, v in checks.items():\n        status = \"PASS\" if v == 0 else \"FAIL\"\n        print(f\"  [{status}]  {k:<40}: {v}\")\n\n    # Per-dataset breakdown\n    print(\"\\nPer-dataset issue counts:\")\n    print(\"-\" * 72)\n    print(f\"  {'dataset':<22} {'rows':>8} {'invalid':>10} {'oob':>7} {'<4px':>7} {'>40%':>7} {'dup':>7}\")\n    for ds in sorted(df[\"source_dataset\"].unique()):\n        sub = df[df[\"source_dataset\"] == ds]\n        n_inv = int(((sub[\"xmin\"] >= sub[\"xmax\"]) | (sub[\"ymin\"] >= sub[\"ymax\"])).sum())\n        n_oob = 0\n        if \"img_w\" in sub.columns:\n            _m = sub[\"img_w\"].notna() & sub[\"img_h\"].notna()\n            n_oob = int((_m & ((sub[\"xmax\"] > sub[\"img_w\"]) | (sub[\"ymax\"] > sub[\"img_h\"]))).sum())\n        n_tiny = int(((sub[\"width\"] < 4) | (sub[\"height\"] < 4)).sum())\n        n_huge = int((sub[\"relative_area_pct\"] > 40).sum())\n        n_dup  = int(sub.duplicated(subset=[\"image\", \"label\", \"xmin\", \"ymin\", \"xmax\", \"ymax\"], keep=False).sum())\n        print(f\"  {ds:<22} {len(sub):>8,} {n_inv:>10} {n_oob:>7} {n_tiny:>7} {n_huge:>7} {n_dup:>7}\")\n\n    # Details when something fails — show 5 example offenders\n    if checks[\"labels_outside_CLASS_NAMES\"] > 0:\n        print(f\"\\nSTRAY LABELS (not in CLASS_NAMES):\")\n        print(df[~df[\"label\"].isin(CLASS_NAMES)][\"label\"].value_counts().head(10))\n    if checks[\"boxes_>40pct_area\"] > 0:\n        ex = df[df[\"relative_area_pct\"] > 40].head(5)\n        print(f\"\\nExamples of huge (>40%) boxes:\")\n        print(ex[[\"source_dataset\", \"image\", \"label\", \"relative_area_pct\", \"width\", \"height\"]].to_string())\n    if missing:\n        print(f\"\\nExamples of missing image files:\")\n        for p in missing[:5]:\n            print(f\"  {p}\")\n\n    print()\n    return checks\n\n_audit_results = data_integrity_audit(df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T19:58:36.492087Z","iopub.execute_input":"2026-04-24T19:58:36.492861Z","iopub.status.idle":"2026-04-24T19:58:37.647308Z","shell.execute_reply.started":"2026-04-24T19:58:36.492826Z","shell.execute_reply":"2026-04-24T19:58:37.646633Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4.8 Statistical Audit (plots)\n\nThree views to spot annotation-style drift between datasets:\n\n- **Box size distribution** per class per dataset (box plots) — same class labeled with radically different box sizes across datasets is a red flag\n- **Spatial heatmap** per class — where do annotators tend to place boxes for each class?\n- **Boxes-per-image distribution** per dataset — over-annotation or under-annotation\n- **Aspect-ratio distribution** per class per dataset — style mismatch detector\n","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\ndef statistical_audit(df):\n    classes = [c for c in CLASS_NAMES if (df[\"label\"] == c).any()]\n    datasets = sorted(df[\"source_dataset\"].unique())\n\n    # --- 1) Box area distribution per class x dataset\n    fig, axes = plt.subplots(1, len(classes), figsize=(5.5 * len(classes), 5))\n    if len(classes) == 1: axes = [axes]\n    for i, cls in enumerate(classes):\n        data, labels = [], []\n        for ds in datasets:\n            vals = df[(df[\"label\"] == cls) & (df[\"source_dataset\"] == ds)][\"relative_area_pct\"].values\n            if len(vals) >= 5:\n                data.append(vals); labels.append(f\"{ds}\\n(n={len(vals)})\")\n        if data:\n            axes[i].boxplot(data, showfliers=False)\n            axes[i].set_xticks(range(1, len(labels) + 1))\n            axes[i].set_xticklabels(labels, rotation=40, ha=\"right\", fontsize=8)\n            axes[i].set_title(f\"{cls}: box area %  (per dataset)\")\n            axes[i].set_ylabel(\"relative_area_pct\")\n            axes[i].grid(axis=\"y\", alpha=0.3)\n    plt.tight_layout(); plt.show()\n\n    # --- 2) Spatial heatmap per class\n    fig, axes = plt.subplots(1, len(classes), figsize=(5 * len(classes), 5))\n    if len(classes) == 1: axes = [axes]\n    for i, cls in enumerate(classes):\n        sub = df[df[\"label\"] == cls]\n        if len(sub) < 10: axes[i].axis(\"off\"); continue\n        hb = axes[i].hexbin(sub[\"center_x\"], sub[\"center_y\"], gridsize=22,\n                             cmap=\"viridis\", extent=(0, 1, 0, 1))\n        axes[i].set_xlim(0, 1); axes[i].set_ylim(1, 0)  # flip y (top=0)\n        axes[i].set_title(f\"{cls}: spatial heatmap  n={len(sub):,}\")\n        axes[i].set_xlabel(\"center_x (normalized)\"); axes[i].set_ylabel(\"center_y (top=0)\")\n        plt.colorbar(hb, ax=axes[i], label=\"box count\")\n    plt.tight_layout(); plt.show()\n\n    # --- 3) Boxes-per-image distribution\n    bpi = df.groupby([\"source_dataset\", \"image\"]).size().reset_index(name=\"n_boxes\")\n    fig, ax = plt.subplots(figsize=(12, 5))\n    data_by_ds = [bpi[bpi[\"source_dataset\"] == ds][\"n_boxes\"].values for ds in datasets]\n    box = ax.boxplot(data_by_ds, showfliers=True)\n    ax.set_xticks(range(1, len(datasets) + 1))\n    ax.set_xticklabels([f\"{ds}\\n(imgs={len(bpi[bpi.source_dataset==ds])})\" for ds in datasets],\n                       rotation=30, ha=\"right\", fontsize=9)\n    ax.set_title(\"Boxes per image, by dataset\")\n    ax.set_ylabel(\"# boxes per image\"); ax.grid(axis=\"y\", alpha=0.3)\n    plt.tight_layout(); plt.show()\n\n    # --- 4) Aspect ratio distribution per class\n    fig, axes = plt.subplots(1, len(classes), figsize=(5.5 * len(classes), 4.5))\n    if len(classes) == 1: axes = [axes]\n    for i, cls in enumerate(classes):\n        data, labels = [], []\n        for ds in datasets:\n            vals = df[(df[\"label\"] == cls) & (df[\"source_dataset\"] == ds)][\"aspect_ratio\"].values\n            vals = vals[(vals > 0.1) & (vals < 10)]  # sanity clamp\n            if len(vals) >= 5:\n                data.append(vals); labels.append(ds)\n        if data:\n            axes[i].boxplot(data, showfliers=False)\n            axes[i].set_xticks(range(1, len(labels) + 1))\n            axes[i].set_xticklabels(labels, rotation=40, ha=\"right\", fontsize=8)\n            axes[i].axhline(1.0, color=\"grey\", linestyle=\"--\", linewidth=0.7, alpha=0.6)\n            axes[i].set_title(f\"{cls}: aspect ratio (w/h)\")\n            axes[i].set_ylabel(\"aspect\"); axes[i].grid(axis=\"y\", alpha=0.3)\n    plt.tight_layout(); plt.show()\n\n    # --- Print summary table\n    print(\"\\nPer-class × per-dataset box-count + median area:\")\n    print(\"-\" * 80)\n    print(f\"  {'class':<14} {'dataset':<22} {'n_boxes':>10} {'median_area%':>14} {'median_aspect':>15}\")\n    for cls in classes:\n        for ds in datasets:\n            sub = df[(df[\"label\"] == cls) & (df[\"source_dataset\"] == ds)]\n            if len(sub) == 0: continue\n            print(f\"  {cls:<14} {ds:<22} {len(sub):>10,} \"\n                  f\"{sub['relative_area_pct'].median():>14.2f} \"\n                  f\"{sub['aspect_ratio'].median():>15.2f}\")\n\nstatistical_audit(df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T19:48:37.275712Z","iopub.execute_input":"2026-04-24T19:48:37.276455Z","iopub.status.idle":"2026-04-24T19:48:39.647145Z","shell.execute_reply.started":"2026-04-24T19:48:37.276417Z","shell.execute_reply":"2026-04-24T19:48:39.646431Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4.9 Edge-Case Image Viewer\n\nShows the most extreme examples in the dataset. Look for bugs / bad annotations:\n\n- **Largest boxes by relative area** — often whole-image artifacts or mis-labeled full lungs\n- **Smallest boxes** — verify they're real lesions, not coordinate-parsing noise\n- **Images with the most boxes** — over-annotated images can drag training\n- **Biggest multi-dataset overlaps** — images that appear in multiple source datasets (potential duplicates)\n","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport matplotlib.patches as _mpatches\nimport cv2\n\n_EC_COLORS = {\"tumor_xray\": \"#ff3030\", \"tumor_ct\": \"#ff6030\",\n              \"tuberculosis\": \"#ff9900\", \"pneumonia\": \"#00e0ff\"}\n\ndef _render_box_grid(rows_df, title, cols=3, target=512):\n    n = len(rows_df)\n    if n == 0: return\n    rows = (n + cols - 1) // cols\n    fig, axes = plt.subplots(rows, cols, figsize=(cols * 4.2, rows * 4.2))\n    axes = np.array(axes).reshape(rows, cols)\n    fig.suptitle(title, fontsize=13, y=1.0)\n    rows_list = list(rows_df.itertuples(index=False))\n    for idx in range(rows * cols):\n        ax = axes[idx // cols, idx % cols]\n        if idx >= n:\n            ax.axis(\"off\"); continue\n        r = rows_list[idx]\n        img = cv2.imread(r.image_path, cv2.IMREAD_GRAYSCALE)\n        if img is None:\n            ax.text(0.5, 0.5, \"READ FAIL\", ha=\"center\", va=\"center\")\n            ax.axis(\"off\"); continue\n        H0, W0 = img.shape[:2]\n        img_r = cv2.resize(img, (target, target))\n        sx, sy = target / W0, target / H0\n        ax.imshow(img_r, cmap=\"gray\")\n        # Draw all boxes for this image (not just the \"offender\")\n        all_boxes = df[(df[\"image\"] == r.image) & (df[\"source_dataset\"] == r.source_dataset)]\n        for _, rb in all_boxes.iterrows():\n            col = _EC_COLORS.get(rb[\"label\"], \"#ffff00\")\n            rect = _mpatches.Rectangle((rb[\"xmin\"] * sx, rb[\"ymin\"] * sy),\n                                         rb[\"width\"] * sx, rb[\"height\"] * sy,\n                                         linewidth=1.7, edgecolor=col, facecolor=\"none\")\n            ax.add_patch(rect)\n            ax.text(rb[\"xmin\"] * sx, max(rb[\"ymin\"] * sy - 2, 8), rb[\"label\"],\n                    color=col, fontsize=7,\n                    bbox=dict(facecolor=\"black\", alpha=0.55, edgecolor=\"none\", pad=0.6))\n        ax.set_title(f\"{r.source_dataset}\\n{str(r.image)[:30]}  |  area={r.relative_area_pct:.2f}%\",\n                     fontsize=8)\n        ax.axis(\"off\")\n    plt.tight_layout(); plt.show()\n\ndef edge_case_viewer(df, k=6):\n    # Largest boxes\n    largest = df.nlargest(k, \"relative_area_pct\")\n    _render_box_grid(largest, f\"LARGEST boxes by relative area (top {k})\")\n\n    # Smallest non-trivial boxes\n    small = df[df[\"relative_area_pct\"] > 0.01].nsmallest(k, \"relative_area_pct\")\n    _render_box_grid(small, f\"SMALLEST boxes >0.01% area (top {k})\")\n\n    # Most boxes per image\n    bpi = df.groupby([\"source_dataset\", \"image\"]).size().reset_index(name=\"n_boxes\")\n    top_busy = bpi.nlargest(k, \"n_boxes\")\n    rows = []\n    for _, tb in top_busy.iterrows():\n        r = df[(df[\"source_dataset\"] == tb[\"source_dataset\"]) & (df[\"image\"] == tb[\"image\"])].iloc[0]\n        rows.append(r)\n    import pandas as _pd\n    busy_df = _pd.DataFrame(rows)\n    busy_df[\"_n_boxes_on_image\"] = top_busy[\"n_boxes\"].values\n    _render_box_grid(busy_df, f\"IMAGES WITH MOST BOXES (top {k})\")\n\n    # Images with both tumor_xray AND pneumonia (multi-class overlap)\n    both_classes = (df.groupby(\"image\")[\"label\"]\n                      .apply(lambda s: set(s))\n                      .reset_index())\n    multi_class_imgs = both_classes[both_classes[\"label\"].apply(lambda s: len(s) > 1)]\n    if len(multi_class_imgs) > 0:\n        picks = multi_class_imgs.sample(min(k, len(multi_class_imgs)), random_state=0)\n        rows = [df[df[\"image\"] == ik].iloc[0] for ik in picks[\"image\"]]\n        import pandas as _pd\n        mc_df = _pd.DataFrame(rows)\n        _render_box_grid(mc_df, f\"IMAGES WITH MULTIPLE CLASSES (n total={len(multi_class_imgs)}, showing {len(picks)})\")\n    else:\n        print(\"No images with multiple classes (every image is single-class)\")\n\nedge_case_viewer(df, k=6)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T18:05:22.707601Z","iopub.execute_input":"2026-04-24T18:05:22.707807Z","iopub.status.idle":"2026-04-24T18:05:26.163651Z","shell.execute_reply.started":"2026-04-24T18:05:22.707783Z","shell.execute_reply":"2026-04-24T18:05:26.163073Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4.10 Same-Class Cross-Dataset Comparison\n\nFor each class, shows examples from **every dataset that contains that class**, side by side. The critical sanity check: does \"tumor_xray\" labeled by VinDr look like the same visual concept as \"tumor_xray\" labeled by VinBigData or cxray14?\n\nIf annotation styles differ visibly across datasets (e.g., one dataset draws tight boxes around nodules, another draws large surrounding regions), the model is being asked to learn a compromise — which is exactly where mAP ceilings come from.\n","metadata":{}},{"cell_type":"code","source":"def same_class_cross_dataset(df, per_cell=3, target=448):\n    \"\"\"For each class, one row = one dataset, `per_cell` example images per cell.\"\"\"\n    classes = [c for c in CLASS_NAMES if (df[\"label\"] == c).any()]\n    for cls in classes:\n        cls_df = df[df[\"label\"] == cls]\n        datasets_with_cls = sorted(cls_df[\"source_dataset\"].unique())\n        if len(datasets_with_cls) == 0: continue\n        rows = len(datasets_with_cls)\n        cols = per_cell\n        fig, axes = plt.subplots(rows, cols, figsize=(cols * 4, rows * 4))\n        axes = np.array(axes).reshape(rows, cols)\n        fig.suptitle(f'class = \"{cls}\"  — same label across {rows} datasets', fontsize=14, y=1.01)\n        for r_idx, ds in enumerate(datasets_with_cls):\n            ds_cls_imgs = cls_df[cls_df[\"source_dataset\"] == ds][\"image\"].unique()\n            rng = np.random.RandomState(hash((cls, ds)) % 10000)\n            picks = rng.choice(ds_cls_imgs, size=min(cols, len(ds_cls_imgs)), replace=False)\n            for c_idx in range(cols):\n                ax = axes[r_idx, c_idx]\n                if c_idx >= len(picks):\n                    ax.axis(\"off\"); continue\n                img_key = picks[c_idx]\n                img_rows = df[(df[\"source_dataset\"] == ds) & (df[\"image\"] == img_key)]\n                img_path = img_rows[\"image_path\"].iloc[0]\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                if img is None:\n                    ax.text(0.5, 0.5, \"READ FAIL\", ha=\"center\", va=\"center\")\n                    ax.axis(\"off\"); continue\n                H0, W0 = img.shape[:2]\n                img_r = cv2.resize(img, (target, target))\n                sx, sy = target / W0, target / H0\n                ax.imshow(img_r, cmap=\"gray\")\n                for _, rb in img_rows.iterrows():\n                    col = _EC_COLORS.get(rb[\"label\"], \"#ffff00\")\n                    if rb[\"label\"] != cls:\n                        # Show other-class boxes in a muted color\n                        col = \"#888888\"\n                    rect = _mpatches.Rectangle((rb[\"xmin\"] * sx, rb[\"ymin\"] * sy),\n                                                 rb[\"width\"] * sx, rb[\"height\"] * sy,\n                                                 linewidth=1.8 if rb[\"label\"] == cls else 1.0,\n                                                 edgecolor=col, facecolor=\"none\")\n                    ax.add_patch(rect)\n                if c_idx == 0:\n                    ax.set_ylabel(ds, fontsize=11, rotation=0, ha=\"right\", va=\"center\", labelpad=80)\n                ax.set_xticks([]); ax.set_yticks([])\n                cnt = (img_rows[\"label\"] == cls).sum()\n                ax.set_title(f\"n_{cls}_boxes={cnt}\", fontsize=8)\n        plt.tight_layout(); plt.show()\n\nsame_class_cross_dataset(df, per_cell=3)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T18:05:26.164649Z","iopub.execute_input":"2026-04-24T18:05:26.164904Z","iopub.status.idle":"2026-04-24T18:05:28.16779Z","shell.execute_reply.started":"2026-04-24T18:05:26.164877Z","shell.execute_reply":"2026-04-24T18:05:28.167072Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Raw Quality Audit","metadata":{}},{"cell_type":"code","source":"def compute_qm(path):\n    img = cv2.imread(path)\n    if img is None:\n        return None\n    g = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    h, w = g.shape\n    br, co = np.mean(g), np.std(g)\n    sh = cv2.Laplacian(g, cv2.CV_64F).var()\n    snr = br / max(co, 1e-6)\n    return {\"width\": w, \"height\": h, \"aspect_ratio\": round(w/max(h,1), 3),\n            \"brightness\": round(br, 2), \"contrast\": round(co, 2),\n            \"sharpness\": round(sh, 2), \"snr\": round(snr, 2)}\n\nif len(df) > 0:\n    uq = df.drop_duplicates(\"image\")[[\"image\",\"image_path\",\"source_dataset\"]].reset_index(drop=True)\n    qr, fl = [], []\n    for _, r in tqdm(uq.iterrows(), total=len(uq), desc=\"Quality profiling\"):\n        m = compute_qm(r[\"image_path\"])\n        (qr if m else fl).append(\n            {\"image\": r[\"image\"], \"source_dataset\": r[\"source_dataset\"], **m} if m else r[\"image\"]\n        )\n    quality_df = pd.DataFrame(qr)\n    if fl:\n        print(f\"⚠ {len(fl)} failed\")\n    print(f\"Profiled {len(quality_df)} images\")\n    display(quality_df.describe().round(2))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T20:57:11.845233Z","iopub.execute_input":"2026-04-24T20:57:11.845415Z","iopub.status.idle":"2026-04-24T21:01:39.159299Z","shell.execute_reply.started":"2026-04-24T20:57:11.845386Z","shell.execute_reply":"2026-04-24T21:01:39.158584Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Data Preparation & Integrity\n\n**Critical fix vs v1:** this section now *removes* duplicate/leaky images instead of just warning. The v1 notebook detected 116 hash overlaps between train and test and ignored them.\n","metadata":{}},{"cell_type":"code","source":"# FIX ⑦: DICOM-aware integrity check\nif len(df) > 0:\n    corrupt = []\n    uq = df.drop_duplicates(\"image\")[[\"image\",\"image_path\",\"source_dataset\"]].reset_index(drop=True)\n    for _, r in tqdm(uq.iterrows(), total=len(uq), desc=\"Integrity check\"):\n        p = r[\"image_path\"]\n        try:\n            if p.lower().endswith((\".dcm\", \".dicom\")):\n                # DICOMs were already converted to PNG by the loaders;\n                # if the path is still .dcm it means conversion failed → skip\n                import pydicom\n                pydicom.dcmread(p, stop_before_pixels=True)\n            else:\n                img = Image.open(p); img.verify()\n        except Exception:\n            corrupt.append(r[\"image\"])\n    print(f\"✓ Checked: {len(uq)}, Corrupt: {len(corrupt)}\")\n    if corrupt:\n        df = df[~df[\"image\"].isin(corrupt)].reset_index(drop=True)\n        print(f\"  Removed {len(corrupt)} corrupt\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:01:56.070536Z","iopub.execute_input":"2026-04-24T21:01:56.071012Z","iopub.status.idle":"2026-04-24T21:02:22.365385Z","shell.execute_reply.started":"2026-04-24T21:01:56.07098Z","shell.execute_reply":"2026-04-24T21:02:22.364628Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Perceptual hashing for duplicate detection\nif len(df) > 0:\n    hrecs = []\n    uq = df.drop_duplicates(\"image\")[[\"image\",\"image_path\",\"source_dataset\"]].reset_index(drop=True)\n    for _, r in tqdm(uq.iterrows(), total=len(uq), desc=\"pHash\"):\n        try:\n            img = Image.open(r[\"image_path\"]).convert(\"RGB\")\n            hrecs.append({\"image\": r[\"image\"], \"source_dataset\": r[\"source_dataset\"],\n                          \"phash\": str(imagehash.phash(img))})\n        except Exception:\n            pass\n    hash_df = pd.DataFrame(hrecs)\n    dups = hash_df[\"phash\"].value_counts()\n    dups = dups[dups > 1]\n    print(f\"Unique hashes: {hash_df['phash'].nunique()}/{len(hash_df)}\")\n    print(f\"Duplicate groups: {len(dups)}\")\n    # We handle the leakage at split time (Section 11) using hash_df.\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:02:28.47627Z","iopub.execute_input":"2026-04-24T21:02:28.476482Z","iopub.status.idle":"2026-04-24T21:06:25.98905Z","shell.execute_reply.started":"2026-04-24T21:02:28.476448Z","shell.execute_reply":"2026-04-24T21:06:25.988396Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Exploratory Data Analysis","metadata":{}},{"cell_type":"code","source":"pal = {\"tumor_ct\": \"#c0392b\", \"tumor_xray\": \"#e67e22\",\n       \"tuberculosis\": \"#3498db\", \"pneumonia\": \"#27ae60\"}\n\nif len(df) > 0:\n    fig, axes = plt.subplots(1, len(CLASS_NAMES)+1, figsize=(6*(len(CLASS_NAMES)+1), 6))\n    fig.suptitle(\"Spatial Density Heatmap\", fontsize=14, fontweight=\"bold\")\n    cs = 100\n    for idx, label in enumerate(CLASS_NAMES):\n        hm = np.zeros((cs, cs))\n        sub = df[df[\"label\"] == label]\n        for _, r in sub.iterrows():\n            cx = np.clip(int(r[\"center_x\"]*(cs-1)), 0, cs-1)\n            cy = np.clip(int(r[\"center_y\"]*(cs-1)), 0, cs-1)\n            for dx in range(-3, 4):\n                for dy in range(-3, 4):\n                    nx, ny = cx+dx, cy+dy\n                    if 0 <= nx < cs and 0 <= ny < cs:\n                        hm[ny, nx] += np.exp(-(dx*dx+dy*dy)/4.0)\n        im = axes[idx].imshow(hm, cmap=\"YlOrRd\", interpolation=\"gaussian\")\n        axes[idx].set_title(f\"{label.upper()} ({len(sub)})\", fontweight=\"bold\")\n        plt.colorbar(im, ax=axes[idx], fraction=0.046)\n    hm_all = np.zeros((cs, cs))\n    for _, r in df.iterrows():\n        cx = np.clip(int(r[\"center_x\"]*(cs-1)), 0, cs-1)\n        cy = np.clip(int(r[\"center_y\"]*(cs-1)), 0, cs-1)\n        hm_all[cy, cx] += 1\n    im = axes[-1].imshow(hm_all, cmap=\"inferno\", interpolation=\"gaussian\")\n    axes[-1].set_title(\"COMBINED\", fontweight=\"bold\")\n    plt.colorbar(im, ax=axes[-1], fraction=0.046)\n    plt.tight_layout(); plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:06:34.442254Z","iopub.execute_input":"2026-04-24T21:06:34.442797Z","iopub.status.idle":"2026-04-24T21:06:39.613143Z","shell.execute_reply.started":"2026-04-24T21:06:34.442769Z","shell.execute_reply":"2026-04-24T21:06:39.612401Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if len(df) > 0:\n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n    for l in CLASS_NAMES:\n        s = df[df[\"label\"] == l]\n        if len(s) == 0:\n            continue\n        bpi = s.groupby(\"image\").size()\n        axes[0].hist(bpi, bins=range(1, max(bpi.max(), 2)+2), alpha=0.6,\n                     label=f\"{l} (μ={bpi.mean():.1f})\", color=pal[l])\n    axes[0].set_title(\"BBoxes per Image\", fontweight=\"bold\"); axes[0].legend()\n    cc = df[\"label\"].value_counts()\n    bars = axes[1].bar(cc.index, cc.values, color=[pal.get(l, \"gray\") for l in cc.index])\n    for b, v in zip(bars, cc.values):\n        axes[1].text(b.get_x()+b.get_width()/2, v+5, str(v), ha=\"center\", fontweight=\"bold\")\n    if len(cc) > 1:\n        axes[1].set_xlabel(f\"Imbalance: {cc.max()/cc.min():.1f}:1\")\n    axes[1].set_title(\"Class Distribution\", fontweight=\"bold\")\n    plt.xticks(rotation=20)\n    plt.tight_layout(); plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:06:59.111064Z","iopub.execute_input":"2026-04-24T21:06:59.111306Z","iopub.status.idle":"2026-04-24T21:06:59.472916Z","shell.execute_reply.started":"2026-04-24T21:06:59.11128Z","shell.execute_reply":"2026-04-24T21:06:59.472159Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Bias Analysis","metadata":{}},{"cell_type":"code","source":"if len(df) > 0:\n    print(\"SPATIAL BIAS — Quadrant Distribution\")\n    print(\"=\" * 55)\n    for l in CLASS_NAMES:\n        sub = df[df[\"label\"] == l]\n        if len(sub) == 0:\n            continue\n        qs = [(\"Top\" if r[\"center_y\"] < 0.5 else \"Bottom\") + \"-\" +\n              (\"Left\" if r[\"center_x\"] < 0.5 else \"Right\")\n              for _, r in sub.iterrows()]\n        ct = Counter(qs); tot = sum(ct.values())\n        print(f\"\\n  {l.upper()}:\")\n        for q in [\"Top-Left\",\"Top-Right\",\"Bottom-Left\",\"Bottom-Right\"]:\n            c = ct.get(q, 0)\n            print(f\"    {q:15s}: {c:5d} ({c/tot*100:.1f}%)\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:07:06.917035Z","iopub.execute_input":"2026-04-24T21:07:06.917302Z","iopub.status.idle":"2026-04-24T21:07:07.569251Z","shell.execute_reply.started":"2026-04-24T21:07:06.917274Z","shell.execute_reply":"2026-04-24T21:07:07.568465Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Lung Field Segmentation & Cropping (NEW)\n\nThis is the biggest single improvement. Without lung-field masking, the model shortcut-learns from **text overlays** (\"PORTABLE\", \"AP\", \"L\"), **EKG leads**, **scanner borders**, and **patient positioning markers** — all of which correlate with disease labels but have nothing to do with the lungs.\n\nWe run a pretrained U-Net (`lungmask`, R231 variant, trained on LUNA16) per image, crop to the lung bounding box + small margin, and adjust all annotation bboxes accordingly. Results are **cached to disk** so this only runs once per dataset.\n","metadata":{}},{"cell_type":"code","source":"%%capture\n# Initialize lungmask inferer — downloads weights on first use\nimport SimpleITK as sitk\n\nprint(\"Initializing lung segmentation model (downloads weights on first run)...\")\nlung_inferer = LMInferer(modelname=\"R231\", force_cpu=False)\nprint(\"✓ Lung segmentation ready\")\n\ndef get_lung_crop(img_bgr, margin_frac=0.05, min_area_frac=0.10):\n    \"\"\"Return (x1, y1, x2, y2) crop bbox for the lung field, or None on failure.\n    The lungmask model expects a SimpleITK image formatted as a 3D volume.\"\"\"\n    h, w = img_bgr.shape[:2]\n    gray = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2GRAY)\n    \n    try:\n        # lungmask expects a SimpleITK Image (3D volume). For a single 2D slice,\n        # we wrap it as a 1-slice volume with proper spacing.\n        vol = gray[None, ...].astype(np.float32)  # shape (1, H, W)\n        sitk_img = sitk.GetImageFromArray(vol)\n        sitk_img.SetSpacing([1.0, 1.0, 1.0])\n        mask = lung_inferer.apply(sitk_img)  # returns ndarray (1, H, W)\n    except Exception as e:\n        return None\n    \n    lung = (mask[0] > 0).astype(np.uint8)\n    if lung.sum() < min_area_frac * h * w:\n        return None\n    ys, xs = np.where(lung > 0)\n    x1, x2 = int(xs.min()), int(xs.max())\n    y1, y2 = int(ys.min()), int(ys.max())\n    mx = int((x2 - x1) * margin_frac)\n    my = int((y2 - y1) * margin_frac)\n    x1 = max(0, x1 - mx); x2 = min(w, x2 + mx)\n    y1 = max(0, y1 - my); y2 = min(h, y2 + my)\n    return (x1, y1, x2, y2)\n\n# Fallback for X-rays where lungmask (CT-focused) might fail\ndef get_lung_crop_simple(img_bgr, margin_frac=0.05, min_area_frac=0.05):\n    \"\"\"Simple lung field extraction using thresholding + morphology.\n    More reliable than lungmask for 2D chest X-rays.\"\"\"\n    h, w = img_bgr.shape[:2]\n    gray = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2GRAY)\n    blurred = cv2.GaussianBlur(gray, (5, 5), 0)\n    _, binary = cv2.threshold(blurred, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)\n    if np.mean(binary[:h//4, :]) > 127:\n        binary = 255 - binary\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (15, 15))\n    cleaned = cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel, iterations=3)\n    cleaned = cv2.morphologyEx(cleaned, cv2.MORPH_OPEN, kernel, iterations=2)\n    contours, _ = cv2.findContours(cleaned, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    if not contours:\n        return None\n    large_contours = [c for c in contours if cv2.contourArea(c) > min_area_frac * h * w]\n    if not large_contours:\n        return None\n    all_points = np.vstack(large_contours)\n    x1, y1, cw, ch = cv2.boundingRect(all_points)\n    x2, y2 = x1 + cw, y1 + ch\n    mx = int(cw * margin_frac)\n    my = int(ch * margin_frac)\n    x1 = max(0, x1 - mx); x2 = min(w, x2 + mx)\n    y1 = max(0, y1 - my); y2 = min(h, y2 + my)\n    crop_area = (x2 - x1) * (y2 - y1)\n    if crop_area < 0.3 * h * w:\n        return None\n    return (x1, y1, x2, y2)","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:07:16.777599Z","iopub.execute_input":"2026-04-24T21:07:16.778115Z","iopub.status.idle":"2026-04-24T21:07:17.109017Z","shell.execute_reply.started":"2026-04-24T21:07:16.778083Z","shell.execute_reply":"2026-04-24T21:07:17.10841Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%capture\nif DATASET_FROM_CACHE:\n    crop_manifest = {}\n    print(\"\\u23e9 Lung cropping: skipped (cached)\")\nelse:\n    # ── Lung crop cache: prefer external Kaggle dataset, fall back to local cropping ──\n    import shutil, json, cv2\n\n    # 1. Try loading manifest from external dataset (attached via Kaggle Datasets)\n    _ext_dir = Path(globals().get(\"external_lung_crop_dir\", \"\"))\n    _ext_manifest = _ext_dir / \"crop_manifest.json\" if _ext_dir.exists() else Path(\"__nonexistent__\")\n\n    if _ext_manifest.exists():\n        with open(_ext_manifest, \"r\") as f:\n            crop_manifest = json.load(f)\n        # Re-anchor cache_path from old working-dir paths to the external dataset location\n        for k, v in crop_manifest.items():\n            fname = Path(v[\"cache_path\"]).name\n            v[\"cache_path\"] = str(_ext_dir / fname)\n        print(f\"\\u2713 Using external lung-crop cache: {len(crop_manifest)} entries\")\n        print(f\"  Path: {_ext_dir}\")\n    else:\n        # No external cache — use working-dir cache (rebuild if needed)\n        print(\"No external lung-crop cache found; will crop locally.\")\n        crop_manifest_path = lung_cache_dir / \"crop_manifest.json\"\n        if crop_manifest_path.exists():\n            with open(crop_manifest_path, \"r\") as f:\n                crop_manifest = json.load(f)\n            print(f\"Loaded local crop manifest: {len(crop_manifest)} entries\")\n        else:\n            crop_manifest = {}\n\n    # 2. Identify images that still need cropping\n    if len(df) > 0:\n        unique_imgs = df.drop_duplicates(\"image\")[[\"image\",\"image_path\",\"modality\"]].to_dict(\"records\")\n        to_process = [u for u in unique_imgs if u[\"image\"] not in crop_manifest]\n        print(f\"Images to process: {len(to_process)}  (already cached: {len(crop_manifest)})\")\n\n        if to_process:\n            failed = []\n            lung_cache_dir.mkdir(parents=True, exist_ok=True)\n            from tqdm.notebook import tqdm\n            for u in tqdm(to_process, desc=\"Lung crops\"):\n                img = cv2.imread(u[\"image_path\"])\n                if img is None:\n                    failed.append(u[\"image\"]); continue\n\n                if u[\"modality\"] == \"ct\":\n                    h, w = img.shape[:2]\n                    crop = (0, 0, w, h)\n                else:\n                    crop = get_lung_crop(img)\n                    if crop is None:\n                        crop = get_lung_crop_simple(img)\n                    if crop is None:\n                        h, w = img.shape[:2]\n                        crop = (0, 0, w, h)\n                        failed.append(u[\"image\"])\n                x1, y1, x2, y2 = crop\n                cropped = img[y1:y2, x1:x2]\n\n                safe_name = u[\"image\"].replace(\"/\", \"_\").replace(\"\\\\\", \"_\")\n                cache_path = lung_cache_dir / safe_name\n                # v11 fix: JPEG quality=92 is 5-10x smaller than PNG for grayscale CXRs,\n                # visually lossless. Loader resolves by full path from manifest, not by\n                # extension, so mixing .png/.jpg works.\n                cache_path = cache_path.with_suffix(\".jpg\")\n\n                # v11 fix: cap crop to 1024 long-side. We train at 640, so storing\n                # >1024 source is wasted disk. Most CXRs are 1024-2048 native.\n                ch, cw = cropped.shape[:2]\n                if max(ch, cw) > 1024:\n                    _scale = 1024.0 / max(ch, cw)\n                    cropped = cv2.resize(cropped,\n                                         (int(cw * _scale), int(ch * _scale)),\n                                         interpolation=cv2.INTER_AREA)\n\n                _ok = cv2.imwrite(str(cache_path), cropped,\n                                  [cv2.IMWRITE_JPEG_QUALITY, 92])\n                if not _ok:\n                    failed.append(u[\"image\"])\n                    continue\n\n                crop_manifest[u[\"image\"]] = {\n                    \"cache_path\": str(cache_path),\n                    \"crop_box\": [x1, y1, x2, y2],\n                    \"orig_w\": img.shape[1], \"orig_h\": img.shape[0],\n                }\n\n                # v11 fix: persist manifest every 500 crops so a crash mid-loop\n                # doesn't lose all progress\n                if len(crop_manifest) % 500 == 0:\n                    _interim_path = lung_cache_dir / \"crop_manifest.json\"\n                    with open(_interim_path, \"w\") as _f:\n                        json.dump(crop_manifest, _f)\n\n            # Save updated manifest locally\n            crop_manifest_path = lung_cache_dir / \"crop_manifest.json\"\n            with open(crop_manifest_path, \"w\") as f:\n                json.dump(crop_manifest, f)\n            succeeded = len(to_process) - len(failed)\n            print(f\"\\u2713 Cropped {succeeded} new images, {len(failed)} fallbacks.\")\n        else:\n            print(\"\\u2713 All images already in cache \\u2014 skipping lung cropping entirely.\")\n\n    def adjust_bbox(row):\n        \"\"\"Shift bboxes from original image coords to lung-crop coords.\n        Clips to the crop boundary so no bbox exceeds the cropped image.\"\"\"\n        m = crop_manifest.get(row[\"image\"])\n        if m is None:\n            return row\n\n        x1c, y1c, x2c, y2c = m[\"crop_box\"]\n        crop_w = x2c - x1c\n        crop_h = y2c - y1c\n\n        new_xmin = max(0, row[\"xmin\"] - x1c)\n        new_ymin = max(0, row[\"ymin\"] - y1c)\n        new_xmax = min(crop_w, row[\"xmax\"] - x1c)\n        new_ymax = min(crop_h, row[\"ymax\"] - y1c)\n\n        row[\"xmin\"]   = new_xmin\n        row[\"ymin\"]   = new_ymin\n        row[\"xmax\"]   = new_xmax\n        row[\"ymax\"]   = new_ymax\n        row[\"width\"]  = max(0, new_xmax - new_xmin)\n        row[\"height\"] = max(0, new_ymax - new_ymin)\n        row[\"area\"]   = row[\"width\"] * row[\"height\"]\n\n        if crop_w > 0 and crop_h > 0:\n            row[\"relative_area_pct\"] = round(row[\"area\"] / (crop_w * crop_h) * 100, 2)\n            row[\"center_x\"] = round((new_xmin + row[\"width\"] / 2) / crop_w, 3)\n            row[\"center_y\"] = round((new_ymin + row[\"height\"] / 2) / crop_h, 3)\n\n        row[\"image_path\"] = m[\"cache_path\"]\n        return row\n\n    # ── Apply adjustments + validate against actual images ──\n    if len(df) > 0:\n        print(\"Applying crop adjustments to annotations...\")\n        df = df.apply(adjust_bbox, axis=1)\n\n        before_area = len(df)\n        df = df[df[\"relative_area_pct\"] <= 40.0].reset_index(drop=True)\n        print(f\"  Dropped {before_area - len(df)} bboxes exceeding 40% relative area after cropping\")\n\n        before = len(df)\n        df = df[(df[\"width\"] >= 10) & (df[\"height\"] >= 10)].reset_index(drop=True)\n        print(f\"  Dropped {before - len(df)} bboxes that became too small after cropping\")\n\n        print(\"Validating bboxes against actual crop dimensions...\")\n        bad_images = set()\n        checked = 0\n        fixed = 0\n        unique_imgs = df.drop_duplicates(\"image\")[[\"image\", \"image_path\"]].to_dict(\"records\")\n\n        from tqdm.notebook import tqdm\n        for rec in tqdm(unique_imgs, desc=\"Bbox validation\"):\n            img = cv2.imread(rec[\"image_path\"])\n            if img is None:\n                bad_images.add(rec[\"image\"])\n                continue\n            actual_h, actual_w = img.shape[:2]\n            checked += 1\n\n            mask = df[\"image\"] == rec[\"image\"]\n            sub = df.loc[mask]\n\n            oob = ((sub[\"xmax\"] > actual_w + 1) | (sub[\"ymax\"] > actual_h + 1) |\n                   (sub[\"xmin\"] < -1) | (sub[\"ymin\"] < -1))\n\n            if oob.any():\n                df.loc[mask, \"xmin\"]   = df.loc[mask, \"xmin\"].clip(lower=0)\n                df.loc[mask, \"ymin\"]   = df.loc[mask, \"ymin\"].clip(lower=0)\n                df.loc[mask, \"xmax\"]   = df.loc[mask, \"xmax\"].clip(upper=actual_w)\n                df.loc[mask, \"ymax\"]   = df.loc[mask, \"ymax\"].clip(upper=actual_h)\n                df.loc[mask, \"width\"]  = df.loc[mask, \"xmax\"] - df.loc[mask, \"xmin\"]\n                df.loc[mask, \"height\"] = df.loc[mask, \"ymax\"] - df.loc[mask, \"ymin\"]\n                df.loc[mask, \"area\"]   = df.loc[mask, \"width\"] * df.loc[mask, \"height\"]\n                fixed += oob.sum()\n\n        if bad_images:\n            df = df[~df[\"image\"].isin(bad_images)].reset_index(drop=True)\n            print(f\"  Removed {len(bad_images)} images that could not be loaded\")\n\n        before2 = len(df)\n        df = df[(df[\"width\"] >= 10) & (df[\"height\"] >= 10)].reset_index(drop=True)\n        print(f\"  Validated {checked} images, clipped {fixed} OOB bboxes\")\n        print(f\"  Dropped {before2 - len(df)} additional degenerate bboxes\")\n        n_imgs = df['image'].nunique()\n        print(f\"  Final: {len(df)} bboxes on {n_imgs} images\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:07:21.707701Z","iopub.execute_input":"2026-04-24T21:07:21.707951Z","iopub.status.idle":"2026-04-24T21:12:38.70738Z","shell.execute_reply.started":"2026-04-24T21:07:21.707925Z","shell.execute_reply":"2026-04-24T21:12:38.706785Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Save processed df to cache for fast resume ─────────────────────────────\nif not DATASET_FROM_CACHE and len(df) > 0:\n    _cache_path = CACHE_DIR / DF_CACHE_NAME\n    df.to_parquet(_cache_path, index=False)\n    print(f\"\\U0001f4be Saved df cache: {_cache_path}\")\n    print(f\"   {len(df):,} rows, {df['image'].nunique():,} unique images\")\n    print(f\"   On next run, data loading + cropping will be SKIPPED.\")\n    print(f\"   To use across sessions: commit + attach output as input dataset.\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9. Shared Preprocessing Function\n\n**This is the critical fix** for the normalization/CLAHE inconsistency bug. We define `preprocess_image()` once and call it from:\n- The training `Dataset.__getitem__`\n- The validation `Dataset.__getitem__`\n- The `predict()` function at inference\n- The TTA pipeline\n\nAll four paths now see identical input distributions.\n\n**Changes vs v1:**\n- Bilateral filter removed (was smoothing out small lesions)\n- Percentile-based intensity clipping (robust to exposure variation)\n- Single set of CLAHE parameters (no per-modality CLAHE) — since images are already lung-cropped, per-modality tuning is no longer needed\n","metadata":{}},{"cell_type":"code","source":"# Single, consistent CLAHE setup (no per-modality split — images are lung-cropped now)\nCLAHE_CLIP = 2.5\nCLAHE_TILE = (8, 8)\n_clahe = cv2.createCLAHE(clipLimit=CLAHE_CLIP, tileGridSize=CLAHE_TILE)\n\n# CT slices have much wider dynamic range (HU-like already normalized to 8-bit)\n# and benefit from tighter percentile clipping + stronger local contrast so\n# small lesions don't get crushed by lung parenchyma variation.\nCT_CLAHE_CLIP = 3.5\nCT_CLAHE_TILE = (10, 10)\n_clahe_ct = cv2.createCLAHE(clipLimit=CT_CLAHE_CLIP, tileGridSize=CT_CLAHE_TILE)\n\ndef percentile_normalize(gray, lo=0.5, hi=99.5):\n    \"\"\"Clip to percentile range then rescale to 0-255. Robust to exposure variance.\"\"\"\n    p_lo, p_hi = np.percentile(gray, [lo, hi])\n    if p_hi <= p_lo:\n        return gray\n    out = np.clip(gray, p_lo, p_hi)\n    out = ((out - p_lo) / (p_hi - p_lo) * 255).astype(np.uint8)\n    return out\n\ndef preprocess_image(img_bgr, modality=\"xray\"):\n    \"\"\"Canonical preprocessing applied everywhere (train / val / inference).\n\n    modality: \"ct\" routes through a tighter percentile + stronger CLAHE path\n    because CT lesions are small and low-contrast inside lung parenchyma.\n    Anything else (xray, default) keeps the original baseline.\n    \"\"\"\n    if img_bgr is None:\n        return None\n    if img_bgr.ndim == 3:\n        gray = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2GRAY)\n    else:\n        gray = img_bgr\n    if modality == \"ct\":\n        gray = percentile_normalize(gray, lo=1.0, hi=98.0)\n        gray = _clahe_ct.apply(gray)\n    else:\n        gray = percentile_normalize(gray)\n        gray = _clahe.apply(gray)\n    rgb = cv2.cvtColor(gray, cv2.COLOR_GRAY2RGB)\n    return rgb\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T21:12:55.023454Z","iopub.execute_input":"2026-04-24T21:12:55.023718Z","iopub.status.idle":"2026-04-24T21:12:55.031609Z","shell.execute_reply.started":"2026-04-24T21:12:55.023692Z","shell.execute_reply":"2026-04-24T21:12:55.030935Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 9.1 Preprocessing Visualization\n\nVisualize the CLAHE preprocessing pipeline: Raw → Percentile Normalization → CLAHE enhancement, with intensity histograms.","metadata":{}},{"cell_type":"code","source":"# ─── CLAHE / Preprocessing Visualization ───────────────────────────────────────\n# Shows before/after preprocessing for sample images from different sources.\n\nif len(df) > 0:\n    # Pick one sample from each source dataset\n    sources = df[\"source_dataset\"].unique()\n    n_show = min(len(sources), 4)\n    fig, axes = plt.subplots(n_show, 3, figsize=(15, 4 * n_show))\n    if n_show == 1:\n        axes = [axes]\n    fig.suptitle(\"Preprocessing Pipeline: Raw → Percentile Norm → CLAHE\", \n                 fontsize=14, fontweight=\"bold\", y=1.02)\n\n    for i, src in enumerate(sources[:n_show]):\n        row = df[df[\"source_dataset\"] == src].iloc[0]\n        img_bgr = cv2.imread(row[\"image_path\"])\n        if img_bgr is None:\n            continue\n        \n        # Step 1: Raw grayscale\n        gray_raw = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2GRAY)\n        \n        # Step 2: After percentile normalization\n        gray_pnorm = percentile_normalize(gray_raw)\n        \n        # Step 3: After CLAHE\n        gray_clahe = _clahe.apply(gray_pnorm)\n        \n        axes[i][0].imshow(gray_raw, cmap=\"gray\")\n        axes[i][0].set_title(f\"Raw ({src})\", fontsize=10)\n        axes[i][0].axis(\"off\")\n        \n        axes[i][1].imshow(gray_pnorm, cmap=\"gray\")\n        axes[i][1].set_title(\"Percentile Normalized\", fontsize=10)\n        axes[i][1].axis(\"off\")\n        \n        axes[i][2].imshow(gray_clahe, cmap=\"gray\")\n        axes[i][2].set_title(\"+ CLAHE Enhanced\", fontsize=10)\n        axes[i][2].axis(\"off\")\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Histogram comparison for the last sample\n    fig, axes2 = plt.subplots(1, 3, figsize=(15, 3))\n    fig.suptitle(\"Intensity Histograms (before vs after)\", fontsize=12, fontweight=\"bold\")\n    for ax, img_data, title in zip(axes2, [gray_raw, gray_pnorm, gray_clahe],\n                               [\"Raw\", \"Percentile Norm\", \"+ CLAHE\"]):\n        ax.hist(img_data.ravel(), bins=128, range=(0, 255), color=\"steelblue\", alpha=0.7)\n        ax.set_title(title); ax.set_xlim(0, 255)\n        ax.set_xlabel(\"Pixel Intensity\"); ax.set_ylabel(\"Count\")\n    plt.tight_layout(); plt.show()\n    print(\"✓ Preprocessing visualization complete\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T21:13:00.396242Z","iopub.execute_input":"2026-04-24T21:13:00.39645Z","iopub.status.idle":"2026-04-24T21:13:03.72469Z","shell.execute_reply.started":"2026-04-24T21:13:00.396426Z","shell.execute_reply":"2026-04-24T21:13:03.72397Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 9.1 Preprocessing Visualization\n\nVisualize the CLAHE preprocessing pipeline: Raw → Percentile Normalization → CLAHE enhancement, with intensity histograms.","metadata":{}},{"cell_type":"code","source":"# ─── CLAHE / Preprocessing Visualization ───────────────────────────────────────\n# Shows before/after preprocessing for sample images from different sources.\n\nif len(df) > 0:\n    # Pick one sample from each source dataset\n    sources = df[\"source_dataset\"].unique()\n    n_show = min(len(sources), 4)\n    fig, axes = plt.subplots(n_show, 3, figsize=(15, 4 * n_show))\n    if n_show == 1:\n        axes = [axes]\n    fig.suptitle(\"Preprocessing Pipeline: Raw → Percentile Norm → CLAHE\", \n                 fontsize=14, fontweight=\"bold\", y=1.02)\n\n    for i, src in enumerate(sources[:n_show]):\n        row = df[df[\"source_dataset\"] == src].iloc[0]\n        img_bgr = cv2.imread(row[\"image_path\"])\n        if img_bgr is None:\n            continue\n        \n        # Step 1: Raw grayscale\n        gray_raw = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2GRAY)\n        \n        # Step 2: After percentile normalization\n        gray_pnorm = percentile_normalize(gray_raw)\n        \n        # Step 3: After CLAHE\n        gray_clahe = _clahe.apply(gray_pnorm)\n        \n        axes[i][0].imshow(gray_raw, cmap=\"gray\")\n        axes[i][0].set_title(f\"Raw ({src})\", fontsize=10)\n        axes[i][0].axis(\"off\")\n        \n        axes[i][1].imshow(gray_pnorm, cmap=\"gray\")\n        axes[i][1].set_title(\"Percentile Normalized\", fontsize=10)\n        axes[i][1].axis(\"off\")\n        \n        axes[i][2].imshow(gray_clahe, cmap=\"gray\")\n        axes[i][2].set_title(\"+ CLAHE Enhanced\", fontsize=10)\n        axes[i][2].axis(\"off\")\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Histogram comparison for the last sample\n    fig, axes2 = plt.subplots(1, 3, figsize=(15, 3))\n    fig.suptitle(\"Intensity Histograms (before vs after)\", fontsize=12, fontweight=\"bold\")\n    for ax, img_data, title in zip(axes2, [gray_raw, gray_pnorm, gray_clahe],\n                               [\"Raw\", \"Percentile Norm\", \"+ CLAHE\"]):\n        ax.hist(img_data.ravel(), bins=128, range=(0, 255), color=\"steelblue\", alpha=0.7)\n        ax.set_title(title); ax.set_xlim(0, 255)\n        ax.set_xlabel(\"Pixel Intensity\"); ax.set_ylabel(\"Count\")\n    plt.tight_layout(); plt.show()\n    print(\"✓ Preprocessing visualization complete\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T21:13:55.844842Z","iopub.execute_input":"2026-04-24T21:13:55.845577Z","iopub.status.idle":"2026-04-24T21:13:59.183901Z","shell.execute_reply.started":"2026-04-24T21:13:55.845545Z","shell.execute_reply":"2026-04-24T21:13:59.183127Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Unified Normalization Statistics\n\nSingle set of channel stats computed on a sample of lung-cropped, preprocessed images. Used identically for training, validation, and inference.\n","metadata":{}},{"cell_type":"code","source":"if DATASET_FROM_CACHE:\n    with open(NORM_CACHE_PATH) as f:\n        _ns = json.load(f)\n    norm_stats = (_ns[\"mean\"], _ns[\"std\"])\n    print(f\"\\u26a1 norm_stats from cache: mean={norm_stats[0]}, std={norm_stats[1]}\")\nelse:\n    if len(df) > 0:\n        print(\"Computing unified channel statistics on preprocessed images...\")\n        sample_paths = df.drop_duplicates(\"image\")[\"image_path\"].tolist()\n        sp = random.sample(sample_paths, min(1000, len(sample_paths)))\n        csums, csq, px = np.zeros(3), np.zeros(3), 0\n        for p in tqdm(sp, desc=\"Norm stats\"):\n            img = cv2.imread(p)\n            if img is None:\n                continue\n            rgb = preprocess_image(img).astype(np.float64) / 255.0\n            csums += rgb.sum(axis=(0, 1))\n            csq   += (rgb ** 2).sum(axis=(0, 1))\n            px    += rgb.shape[0] * rgb.shape[1]\n        mean = (csums / px).tolist()\n        std  = np.sqrt(np.maximum(csq/px - (csums/px)**2, 0)).tolist()\n        norm_stats = (mean, std)\n        print(f\"Mean: {[round(x,4) for x in mean]}\")\n        print(f\"Std:  {[round(x,4) for x in std]}\")\n    # Save norm stats cache\n    if 'norm_stats' in dir() and norm_stats is not None:\n        with open(CACHE_DIR / NORM_CACHE_NAME, 'w') as f:\n            json.dump({\"mean\": list(norm_stats[0]), \"std\": list(norm_stats[1])}, f)\n        print(f\"\\U0001f4be Saved norm cache: {CACHE_DIR / NORM_CACHE_NAME}\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:14:15.445409Z","iopub.execute_input":"2026-04-24T21:14:15.445637Z","iopub.status.idle":"2026-04-24T21:15:25.851311Z","shell.execute_reply.started":"2026-04-24T21:14:15.445611Z","shell.execute_reply":"2026-04-24T21:15:25.85062Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Stratified Splitting With Leak Removal\n\nStratify on `label + modality`. After splitting, scan for perceptual-hash collisions between train and val/test and **drop** overlapping images from val/test (the fix to v1's 116 leaks).\n","metadata":{}},{"cell_type":"code","source":"if len(df) > 0:\n    # T1.4: preserve ALL labels per image in strat_key so multi-label images\n    # (e.g. VinDr images with both tumor_xray + pneumonia boxes) don't get\n    # collapsed to just their first class for stratification.\n    _lbl_set = (df.groupby(\"image\")[\"label\"]\n                  .agg(lambda s: \"+\".join(sorted(set(s)))))\n    il = df.groupby(\"image\").agg({\n        \"modality\": \"first\",\n        \"image_path\": \"first\", \"source_dataset\": \"first\"\n    }).reset_index()\n    il[\"label\"] = il[\"image\"].map(_lbl_set)\n    il[\"strat_key\"] = il[\"label\"] + \"_\" + il[\"modality\"]\n\n    vc = il[\"strat_key\"].value_counts()\n    small_strata = vc[vc < 2].index.tolist()\n    if small_strata:\n        il.loc[il[\"strat_key\"].isin(small_strata), \"strat_key\"] = \"other\"\n\n    tr, tmp = train_test_split(il, test_size=0.20, random_state=seed, stratify=il[\"strat_key\"])\n    vl, ts  = train_test_split(tmp, test_size=0.50, random_state=seed, stratify=tmp[\"strat_key\"])\n\n    train_keys = tr[\"image\"].tolist()\n    val_keys   = vl[\"image\"].tolist()\n    test_keys  = ts[\"image\"].tolist()\n\n    assert not (set(train_keys) & set(val_keys))\n    assert not (set(train_keys) & set(test_keys))\n    assert not (set(val_keys) & set(test_keys))\n\n    # REMOVE hash-leaks (v1 only detected them, never removed)\n    if \"hash_df\" in dir() and len(hash_df) > 0:\n        train_hashes = set(hash_df[hash_df[\"image\"].isin(train_keys)][\"phash\"])\n        leaky = set(hash_df[\n            (hash_df[\"image\"].isin(val_keys + test_keys)) &\n            (hash_df[\"phash\"].isin(train_hashes))\n        ][\"image\"])\n        if leaky:\n            before_v, before_t = len(val_keys), len(test_keys)\n            val_keys  = [k for k in val_keys  if k not in leaky]\n            test_keys = [k for k in test_keys if k not in leaky]\n            print(f\"✓ Removed {len(leaky)} leaky images: \"\n                  f\"val {before_v}→{len(val_keys)}, test {before_t}→{len(test_keys)}\")\n        else:\n            print(\"✓ No hash leaks\")\n\n    print(f\"\\nSplit: Train={len(train_keys)}, Val={len(val_keys)}, Test={len(test_keys)}\")\n    for sn, sk in [(\"Train\", train_keys), (\"Val\", val_keys), (\"Test\", test_keys)]:\n        sd = df[df[\"image\"].isin(sk)]\n        parts = [f\"{l}: {sd[sd['label']==l]['image'].nunique()}\" for l in CLASS_NAMES]\n        print(f\"  {sn:6s}: {' | '.join(parts)}\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:15:38.504903Z","iopub.execute_input":"2026-04-24T21:15:38.505377Z","iopub.status.idle":"2026-04-24T21:15:38.698533Z","shell.execute_reply.started":"2026-04-24T21:15:38.505347Z","shell.execute_reply":"2026-04-24T21:15:38.69787Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# WeightedRandomSampler — aggressive inverse-frequency so rare classes (TB) get real gradient\nif len(df) > 0:\n    ccounts = df[df[\"image\"].isin(train_keys)][\"label\"].value_counts().to_dict()\n    total   = sum(ccounts.values())\n    # Pure inverse-frequency (dropped sqrt — TB was collapsing to majority class at sqrt scaling)\n    cweights = {l: total / (len(ccounts) * c) for l, c in ccounts.items()}\n    # Clamp extreme ratios so a class that has ~100 boxes doesn't dwarf everything else\n    _max_ratio = 8.0\n    _w_min = min(cweights.values())\n    cweights = {l: min(w, _max_ratio * _w_min) for l, w in cweights.items()}\n    print(\"Class sampling weights:\", {l: round(w, 3) for l, w in cweights.items()})\n    print(\"Class bbox counts:    \", ccounts)\n\n    # FIX: pick the RAREST label per image (not .first()), so multi-label images\n    # like VinDr (pneumonia + tumor_xray) help the rarer class pull up its frequency.\n    def _rarest(labels):\n        return min(labels, key=lambda l: ccounts.get(l, 1e9))\n    img_labels = (df[df[\"image\"].isin(train_keys)]\n                    .groupby(\"image\")[\"label\"].apply(list)\n                    .apply(_rarest))\n    image_weights = {img: cweights.get(l, 1.0) for img, l in img_labels.items()}\n\n    # Size-aware boost for small lesions — stronger multiplier than before (1.3 -> 1.5)\n    img_max_area = df[df[\"image\"].isin(train_keys)].groupby(\"image\")[\"area\"].max()\n    area_med = img_max_area.median()\n    boosted = 0\n    for img in image_weights:\n        if img in img_max_area.index and img_max_area[img] < area_med * 0.5:\n            image_weights[img] *= 1.5\n            boosted += 1\n    print(f\"Size boost applied to {boosted} small-lesion images\")\n\n    tsw = [image_weights.get(i, 1.0) for i in train_keys]\n    sampler = WeightedRandomSampler(tsw, num_samples=len(train_keys), replacement=True)\n    # Sanity check — expected class frequency AFTER sampling\n    _img_lbl = dict(zip(img_labels.index, img_labels.values))\n    _sum_w = sum(tsw)\n    _expected = {}\n    for k, w in zip(train_keys, tsw):\n        l = _img_lbl.get(k)\n        if l is not None:\n            _expected[l] = _expected.get(l, 0.0) + w / _sum_w\n    print(\"Expected class frequency after sampling:\",\n          {k: round(v, 3) for k, v in _expected.items()})\n    print(\"✓ WeightedRandomSampler configured\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:16:09.83758Z","iopub.execute_input":"2026-04-24T21:16:09.838347Z","iopub.status.idle":"2026-04-24T21:16:10.031978Z","shell.execute_reply.started":"2026-04-24T21:16:09.83831Z","shell.execute_reply":"2026-04-24T21:16:10.031314Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 12. Dataset & Albumentations Pipeline\n\nReplaces v1's torchvision pipeline with Albumentations, adding medical-imaging-aware augmentations: ElasticTransform (simulates anatomical variation), RandomGamma (X-ray exposure variation), GaussNoise (low-dose simulation), GridDropout (forces non-shortcut learning).\n","metadata":{}},{"cell_type":"code","source":"def build_train_transform(img_size, norm_stats):\n    \"\"\"Full training augmentation. Bbox-aware.\"\"\"\n    mean, std = norm_stats\n    return A.Compose([\n        # Geometric\n        A.LongestMaxSize(max_size=img_size),\n        A.PadIfNeeded(min_height=img_size, min_width=img_size,\n                      border_mode=cv2.BORDER_CONSTANT, value=0),\n        A.HorizontalFlip(p=0.5),\n        A.Affine(scale=(0.85, 1.15), translate_percent=(-0.05, 0.05),\n                 rotate=(-12, 12), shear=(-4, 4),\n                 interpolation=cv2.INTER_LINEAR, cval=0, p=0.7),\n        A.ElasticTransform(alpha=20, sigma=5, alpha_affine=5,\n                           border_mode=cv2.BORDER_CONSTANT, value=0, p=0.2),\n\n        # Photometric (medical-aware)\n        A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.5),\n        A.RandomGamma(gamma_limit=(85, 115), p=0.3),\n        A.CLAHE(clip_limit=(1.0, 3.0), tile_grid_size=(8, 8), p=0.3),\n        A.GaussNoise(var_limit=(5.0, 20.0), p=0.2),\n        A.GaussianBlur(blur_limit=(3, 5), p=0.1),\n\n        # Final\n        A.Normalize(mean=mean, std=std),\n        ToTensorV2(),\n    ], bbox_params=A.BboxParams(\n        format=\"pascal_voc\", label_fields=[\"labels\"],\n        min_visibility=0.3, min_area=16\n    ))\n\ndef build_tb_train_transform(img_size, norm_stats):\n    \"\"\"TB-specific aggressive augmentation.\n\n    TBX11K has two problems this addresses:\n      (a) 17× underrepresented vs pneumonia → each TB sample sees more diverse views\n      (b) Strong upper-lung spatial prior (93% of TB boxes in top half) → RandomSizedBBoxSafeCrop\n          and stronger translation break that prior so the model cannot cheat on position.\n    \"\"\"\n    mean, std = norm_stats\n    return A.Compose([\n        # Break spatial prior via safe random crops + pad\n        A.LongestMaxSize(max_size=img_size),\n        A.PadIfNeeded(min_height=img_size, min_width=img_size,\n                      border_mode=cv2.BORDER_CONSTANT, value=0),\n        A.RandomSizedBBoxSafeCrop(height=img_size, width=img_size,\n                                  erosion_rate=0.1, p=0.5),\n\n        # Geometric — MORE aggressive than baseline\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.15),  # Anatomy-incorrect but forces lesion-feature learning\n        A.Affine(scale=(0.70, 1.30), translate_percent=(-0.12, 0.12),\n                 rotate=(-20, 20), shear=(-8, 8),\n                 interpolation=cv2.INTER_LINEAR, cval=0, p=0.85),\n        A.ElasticTransform(alpha=40, sigma=7, alpha_affine=8,\n                           border_mode=cv2.BORDER_CONSTANT, value=0, p=0.35),\n        A.GridDistortion(num_steps=5, distort_limit=0.15,\n                         border_mode=cv2.BORDER_CONSTANT, value=0, p=0.25),\n\n        # Photometric — wider variation\n        A.RandomBrightnessContrast(brightness_limit=0.25, contrast_limit=0.25, p=0.7),\n        A.RandomGamma(gamma_limit=(75, 130), p=0.5),\n        A.CLAHE(clip_limit=(1.0, 4.0), tile_grid_size=(8, 8), p=0.5),\n        A.GaussNoise(var_limit=(5.0, 30.0), p=0.35),\n        A.GaussianBlur(blur_limit=(3, 7), p=0.2),\n        A.Sharpen(alpha=(0.1, 0.3), lightness=(0.8, 1.2), p=0.2),\n\n        # Final\n        A.Normalize(mean=mean, std=std),\n        ToTensorV2(),\n    ], bbox_params=A.BboxParams(\n        format=\"pascal_voc\", label_fields=[\"labels\"],\n        min_visibility=0.3, min_area=16\n    ))\n\ndef build_small_lesion_transform(img_size, norm_stats):\n    \"\"\"Transform tuned for small lesions (tumor_ct, tumor_xray).\n\n    Small-lesion classes have two problems:\n      (a) Lesions are small (often <40 px at 1024) so aggressive scale-down or\n          strong blur erases the signal entirely.\n      (b) Lesion borders matter more than texture elsewhere, so sharpen/unsharp\n          help edge features; we deliberately drop vertical flip (anatomy prior).\n    \"\"\"\n    mean, std = norm_stats\n    return A.Compose([\n        A.LongestMaxSize(max_size=img_size),\n        A.PadIfNeeded(min_height=img_size, min_width=img_size,\n                      border_mode=cv2.BORDER_CONSTANT, value=0),\n\n        # Geometric — narrower scale range so small lesions stay resolvable\n        A.HorizontalFlip(p=0.5),\n        A.Affine(scale=(0.92, 1.20), translate_percent=(-0.06, 0.06),\n                 rotate=(-10, 10), shear=(-3, 3),\n                 interpolation=cv2.INTER_LINEAR, cval=0, p=0.7),\n        A.ElasticTransform(alpha=15, sigma=5, alpha_affine=4,\n                           border_mode=cv2.BORDER_CONSTANT, value=0, p=0.2),\n\n        # Photometric — sharpen/unsharp to keep lesion edges crisp\n        A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.20, p=0.55),\n        A.RandomGamma(gamma_limit=(85, 115), p=0.3),\n        A.CLAHE(clip_limit=(1.5, 3.5), tile_grid_size=(8, 8), p=0.4),\n        A.Sharpen(alpha=(0.2, 0.4), lightness=(0.9, 1.1), p=0.35),\n        A.UnsharpMask(blur_limit=(3, 5), sigma_limit=0.0,\n                      alpha=(0.2, 0.4), threshold=5, p=0.25),\n        # Keep noise/blur mild so we don't destroy sub-40px lesion signal\n        A.GaussNoise(var_limit=(3.0, 12.0), p=0.15),\n\n        A.Normalize(mean=mean, std=std),\n        ToTensorV2(),\n    ], bbox_params=A.BboxParams(\n        format=\"pascal_voc\", label_fields=[\"labels\"],\n        min_visibility=0.4, min_area=12\n    ))\n\ndef build_valid_transform(img_size, norm_stats):\n    \"\"\"Validation: deterministic resize + pad + normalize.\"\"\"\n    mean, std = norm_stats\n    return A.Compose([\n        A.LongestMaxSize(max_size=img_size),\n        A.PadIfNeeded(min_height=img_size, min_width=img_size,\n                      border_mode=cv2.BORDER_CONSTANT, value=0),\n        A.Normalize(mean=mean, std=std),\n        ToTensorV2(),\n    ], bbox_params=A.BboxParams(format=\"pascal_voc\", label_fields=[\"labels\"]))\n\ndef build_map_transform(img_size):\n    \"\"\"For mAP evaluation: resize+pad, NO normalization (YOLOXInferenceWrapper handles that).\"\"\"\n    return A.Compose([\n        A.LongestMaxSize(max_size=img_size),\n        A.PadIfNeeded(min_height=img_size, min_width=img_size,\n                      border_mode=cv2.BORDER_CONSTANT, value=0),\n        ToTensorV2(),\n    ], bbox_params=A.BboxParams(format=\"pascal_voc\", label_fields=[\"labels\"]))\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:16:15.097314Z","iopub.execute_input":"2026-04-24T21:16:15.097949Z","iopub.status.idle":"2026-04-24T21:16:15.117128Z","shell.execute_reply.started":"2026-04-24T21:16:15.097913Z","shell.execute_reply":"2026-04-24T21:16:15.116181Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 12.1 Augmentation Visualization\n\nShowcase of the training augmentation pipeline with bounding boxes, demonstrating rotations, flips, elastic transforms, brightness/contrast changes, CLAHE, noise, and GridDropout.","metadata":{}},{"cell_type":"code","source":"# ─── Augmentation Showcase ─────────────────────────────────────────────────────\n# Visualizes the training augmentation pipeline with actual bounding boxes.\n\nif len(df) > 0:\n    # Build a visualization transform (same as training, but NO normalize, NO ToTensorV2)\n    viz_transform = A.Compose([\n        A.LongestMaxSize(max_size=train_sz),\n        A.PadIfNeeded(min_height=train_sz, min_width=train_sz,\n                      border_mode=cv2.BORDER_CONSTANT, value=0),\n        A.HorizontalFlip(p=0.5),\n        A.Affine(scale=(0.85, 1.15), translate_percent=(-0.05, 0.05),\n                 rotate=(-12, 12), shear=(-4, 4),\n                 interpolation=cv2.INTER_LINEAR, cval=0, p=0.7),\n        A.ElasticTransform(alpha=20, sigma=5, alpha_affine=5,\n                           border_mode=cv2.BORDER_CONSTANT, value=0, p=0.3),\n        A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.5),\n        A.RandomGamma(gamma_limit=(85, 115), p=0.3),\n        A.CLAHE(clip_limit=(1.0, 3.0), tile_grid_size=(8, 8), p=0.3),\n        A.GaussNoise(var_limit=(5.0, 20.0), p=0.3),\n    ], bbox_params=A.BboxParams(\n        format=\"pascal_voc\", label_fields=[\"labels\"],\n        min_visibility=0.3, min_area=16\n    ))\n    \n    # Pick sample images — one from each class\n    class_colors = {\n        \"tumor_ct\": (255, 0, 0), \"tumor_xray\": (255, 128, 0),\n        \"tuberculosis\": (0, 128, 255), \"pneumonia\": (0, 200, 0)\n    }\n    \n    sample_img_name = None\n    for cls in CLASS_NAMES:\n        candidates = df[df[\"label\"] == cls].drop_duplicates(\"image\")\n        if len(candidates) > 0:\n            sample_img_name = candidates.iloc[0][\"image\"]\n            break\n    \n    if sample_img_name:\n        rows = df[df[\"image\"] == sample_img_name]\n        ip = rows[\"image_path\"].iloc[0]\n        img = cv2.imread(str(ip))\n        \n        if img is not None:\n            img_pp = preprocess_image(img)\n            bboxes = rows[[\"xmin\",\"ymin\",\"xmax\",\"ymax\"]].values.astype(float).tolist()\n            labels = rows[\"label\"].tolist()\n            \n            clean_b, clean_l = [], []\n            for (x1, y1, x2, y2), l in zip(bboxes, labels):\n                if x2 > x1 + 2 and y2 > y1 + 2:\n                    clean_b.append([x1, y1, x2, y2])\n                    clean_l.append(l)\n            \n            n_aug = 8\n            fig, axes = plt.subplots(2, 4, figsize=(20, 10))\n            fig.suptitle(f\"Training Augmentation Showcase — {sample_img_name} ({clean_l[0] if clean_l else 'N/A'})\",\n                         fontsize=14, fontweight=\"bold\")\n            \n            for i in range(n_aug):\n                ax = axes[i // 4][i % 4]\n                try:\n                    aug = viz_transform(image=img_pp.copy(), bboxes=clean_b, labels=clean_l)\n                    aug_img = aug[\"image\"]\n                    aug_bboxes = aug[\"bboxes\"]\n                    aug_labels = aug[\"labels\"]\n                    \n                    display_img = aug_img.copy()\n                    for (bx1, by1, bx2, by2), bl in zip(aug_bboxes, aug_labels):\n                        col = class_colors.get(bl, (255, 255, 0))\n                        cv2.rectangle(display_img, (int(bx1), int(by1)), (int(bx2), int(by2)), col, 2)\n                        cv2.putText(display_img, bl, (int(bx1), int(by1)-5),\n                                    cv2.FONT_HERSHEY_SIMPLEX, 0.5, col, 1)\n                    \n                    if display_img.ndim == 3 and display_img.shape[2] == 3:\n                        ax.imshow(cv2.cvtColor(display_img, cv2.COLOR_BGR2RGB))\n                    else:\n                        ax.imshow(display_img, cmap=\"gray\")\n                    ax.set_title(f\"Aug #{i+1} ({len(aug_bboxes)} boxes)\", fontsize=9)\n                except Exception as e:\n                    ax.set_title(f\"Aug #{i+1}: {type(e).__name__}\", fontsize=9, color=\"red\")\n                ax.axis(\"off\")\n            \n            plt.tight_layout()\n            plt.show()\n            print(f\"✓ Augmentation showcase: {n_aug} augmented versions\")\n            print(f\"  Transforms: HorizontalFlip, Affine(rotate/scale/shear),\")\n            print(f\"  ElasticTransform, RandomBrightnessContrast, RandomGamma,\")\n            print(f\"  CLAHE, GaussNoise, GridDropout\")\n        else:\n            print(f\"⚠ Could not load sample image: {ip}\")\n    else:\n        print(\"⚠ No sample images available\")","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:16:22.338768Z","iopub.execute_input":"2026-04-24T21:16:22.339328Z","iopub.status.idle":"2026-04-24T21:16:23.934855Z","shell.execute_reply.started":"2026-04-24T21:16:22.339299Z","shell.execute_reply":"2026-04-24T21:16:23.933425Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 12.1 Augmentation Visualization\n\nShowcase of the training augmentation pipeline with bounding boxes, demonstrating rotations, flips, elastic transforms, brightness/contrast changes, CLAHE, noise, and GridDropout.","metadata":{}},{"cell_type":"code","source":"# ─── Augmentation Showcase ─────────────────────────────────────────────────────\n# Visualizes augmentations across ALL 4 classes.\n\nif len(df) > 0:\n    viz_transform = A.Compose([\n        A.LongestMaxSize(max_size=train_sz),\n        A.PadIfNeeded(min_height=train_sz, min_width=train_sz,\n                      border_mode=cv2.BORDER_CONSTANT, value=0),\n        A.HorizontalFlip(p=0.5),\n        A.Affine(scale=(0.85, 1.15), translate_percent=(-0.05, 0.05),\n                 rotate=(-12, 12), shear=(-4, 4),\n                 interpolation=cv2.INTER_LINEAR, cval=0, p=0.7),\n        A.ElasticTransform(alpha=20, sigma=5, alpha_affine=5,\n                           border_mode=cv2.BORDER_CONSTANT, value=0, p=0.3),\n        A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.5),\n        A.RandomGamma(gamma_limit=(85, 115), p=0.3),\n        A.CLAHE(clip_limit=(1.0, 3.0), tile_grid_size=(8, 8), p=0.3),\n        A.GaussNoise(var_limit=(5.0, 20.0), p=0.3),\n    ], bbox_params=A.BboxParams(\n        format=\"pascal_voc\", label_fields=[\"labels\"],\n        min_visibility=0.3, min_area=16\n    ))\n\n    class_colors = {\n        \"tumor_ct\": (255, 0, 0), \"tumor_xray\": (255, 128, 0),\n        \"tuberculosis\": (0, 128, 255), \"pneumonia\": (0, 200, 0)\n    }\n\n    # Show 2 augmentations per class = 8 total (one per class across 2 rows)\n    fig, axes = plt.subplots(len(CLASS_NAMES), 2, figsize=(12, 5 * len(CLASS_NAMES)))\n    fig.suptitle(\"Training Augmentation Showcase — All Classes\", fontsize=16, fontweight=\"bold\")\n\n    for row_idx, cls in enumerate(CLASS_NAMES):\n        candidates = df[df[\"label\"] == cls].drop_duplicates(\"image\")\n        if len(candidates) == 0:\n            axes[row_idx][0].set_title(f\"{cls}: no data\"); axes[row_idx][0].axis(\"off\")\n            axes[row_idx][1].axis(\"off\")\n            continue\n\n        sample = candidates.iloc[0]\n        sample_name = sample[\"image\"] if \"image\" in sample.index else candidates.index[0]\n        rows = df[df[\"image\"] == sample_name]\n        ip = rows[\"image_path\"].iloc[0]\n        img = cv2.imread(str(ip))\n        if img is None:\n            axes[row_idx][0].set_title(f\"{cls}: load failed\"); axes[row_idx][0].axis(\"off\")\n            axes[row_idx][1].axis(\"off\")\n            continue\n\n        img_pp = preprocess_image(img)\n        bboxes = rows[[\"xmin\",\"ymin\",\"xmax\",\"ymax\"]].values.astype(float).tolist()\n        labels = rows[\"label\"].tolist()\n        clean_b, clean_l = [], []\n        for (x1, y1, x2, y2), l in zip(bboxes, labels):\n            if x2 > x1 + 2 and y2 > y1 + 2:\n                clean_b.append([x1, y1, x2, y2]); clean_l.append(l)\n\n        for col_idx in range(2):\n            ax = axes[row_idx][col_idx]\n            try:\n                aug = viz_transform(image=img_pp.copy(), bboxes=clean_b, labels=clean_l)\n                display_img = aug[\"image\"].copy()\n                for (bx1, by1, bx2, by2), bl in zip(aug[\"bboxes\"], aug[\"labels\"]):\n                    col = class_colors.get(bl, (255, 255, 0))\n                    cv2.rectangle(display_img, (int(bx1), int(by1)), (int(bx2), int(by2)), col, 2)\n                    cv2.putText(display_img, bl, (int(bx1), int(by1)-5),\n                                cv2.FONT_HERSHEY_SIMPLEX, 0.5, col, 1)\n                if display_img.ndim == 3 and display_img.shape[2] == 3:\n                    ax.imshow(cv2.cvtColor(display_img, cv2.COLOR_BGR2RGB))\n                else:\n                    ax.imshow(display_img, cmap=\"gray\")\n                ax.set_title(f\"{cls} — Aug #{col_idx+1} ({len(aug['bboxes'])} boxes)\", fontsize=10)\n            except Exception as e:\n                ax.set_title(f\"{cls} — {type(e).__name__}\", fontsize=10, color=\"red\")\n            ax.axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n    print(\"✓ Augmentation showcase: 2 augmented versions × 4 classes\")\n    print(\"  Transforms: HorizontalFlip, Affine, ElasticTransform,\")\n    print(\"  RandomBrightnessContrast, RandomGamma, CLAHE, GaussNoise\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:16:55.565919Z","iopub.execute_input":"2026-04-24T21:16:55.566505Z","iopub.status.idle":"2026-04-24T21:16:57.279507Z","shell.execute_reply.started":"2026-04-24T21:16:55.566476Z","shell.execute_reply":"2026-04-24T21:16:57.278785Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 13. Copy-Paste Bank for Rare Classes\n\nHarvest lesion crops from the training set, pre-indexed by class. The training dataset will randomly paste 0-2 of these into each training image (outside existing bboxes). Massively increases effective TB / tumor_xray sample count.\n","metadata":{}},{"cell_type":"code","source":"def build_copy_paste_bank(df, train_keys, classes_to_augment, max_per_class=500):\n    \"\"\"Pre-extract lesion crops for rare classes.\"\"\"\n    bank = defaultdict(list)\n    tdf = df[df[\"image\"].isin(train_keys) & df[\"label\"].isin(classes_to_augment)]\n    # Group by image to minimize re-reads\n    for img_name, grp in tqdm(tdf.groupby(\"image\"), desc=\"Copy-paste bank\"):\n        ip = grp[\"image_path\"].iloc[0]\n        img = cv2.imread(ip)\n        if img is None:\n            continue\n        for _, r in grp.iterrows():\n            x1, y1, x2, y2 = int(r[\"xmin\"]), int(r[\"ymin\"]), int(r[\"xmax\"]), int(r[\"ymax\"])\n            if x2 <= x1+4 or y2 <= y1+4:\n                continue\n            crop = img[y1:y2, x1:x2].copy()\n            if crop.size == 0:\n                continue\n            bank[r[\"label\"]].append(crop)\n    # Cap\n    for k in bank:\n        if len(bank[k]) > max_per_class:\n            bank[k] = random.sample(bank[k], max_per_class)\n        print(f\"  {k}: {len(bank[k])} crops in bank\")\n    return bank\n\n# Build bank for rare classes only\nif len(df) > 0:\n    rare_classes = [\"tuberculosis\", \"tumor_xray\"]\n    copy_paste_bank = build_copy_paste_bank(df, train_keys, rare_classes)\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:17:23.571026Z","iopub.execute_input":"2026-04-24T21:17:23.571697Z","iopub.status.idle":"2026-04-24T21:18:09.909072Z","shell.execute_reply.started":"2026-04-24T21:17:23.571664Z","shell.execute_reply":"2026-04-24T21:18:09.908307Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def apply_copy_paste(img, bboxes, labels, bank, max_pastes=2, min_size=40, max_size=160):\n    \"\"\"Paste 0-max_pastes crops from the bank into the image, avoiding existing bboxes.\n\n    Args:\n      img: HxWx3 uint8\n      bboxes: list of (x1,y1,x2,y2) in pixel coords (pascal_voc)\n      labels: list of str (class names)\n    Returns:\n      modified img, updated bboxes list, updated labels list\n    \"\"\"\n    if not bank or len(bank) == 0:\n        return img, bboxes, labels\n    h, w = img.shape[:2]\n    n = random.randint(0, max_pastes)\n    img = img.copy()\n    bboxes = list(bboxes); labels = list(labels)\n\n    for _ in range(n):\n        cls = random.choice(list(bank.keys()))\n        if not bank[cls]:\n            continue\n        crop = random.choice(bank[cls]).copy()\n        ch, cw = crop.shape[:2]\n        # Random resize to a plausible lesion size\n        target = random.randint(min_size, max_size)\n        if ch >= cw:\n            new_h = target; new_w = max(8, int(cw * target / ch))\n        else:\n            new_w = target; new_h = max(8, int(ch * target / cw))\n        if new_w >= w or new_h >= h:\n            continue\n        crop = cv2.resize(crop, (new_w, new_h))\n        # v6: Removed GaussianBlur — was destroying small lesion features\n        # that the model needs to learn to detect\n\n        # Find a non-overlapping position (try up to 10 times)\n        placed = False\n        for _ in range(10):\n            x1 = random.randint(0, w - new_w)\n            y1 = random.randint(0, h - new_h)\n            x2, y2 = x1 + new_w, y1 + new_h\n            # Avoid overlapping existing bboxes\n            overlap = False\n            for (bx1, by1, bx2, by2) in bboxes:\n                if not (x2 < bx1 or x1 > bx2 or y2 < by1 or y1 > by2):\n                    overlap = True; break\n            if overlap:\n                continue\n            # Paste with a soft alpha-blend at the edges\n            alpha = np.ones((new_h, new_w, 1), dtype=np.float32)\n            feather = 4\n            for i in range(feather):\n                alpha[i, :, 0] *= (i + 1) / feather\n                alpha[-(i+1), :, 0] *= (i + 1) / feather\n                alpha[:, i, 0] *= (i + 1) / feather\n                alpha[:, -(i+1), 0] *= (i + 1) / feather\n            region = img[y1:y2, x1:x2].astype(np.float32)\n            blended = (alpha * crop.astype(np.float32) + (1 - alpha) * region).astype(np.uint8)\n            img[y1:y2, x1:x2] = blended\n            bboxes.append((x1, y1, x2, y2))\n            labels.append(cls)\n            placed = True; break\n        # If couldn't place, just skip\n\n    return img, bboxes, labels\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:18:39.046774Z","iopub.execute_input":"2026-04-24T21:18:39.047432Z","iopub.status.idle":"2026-04-24T21:18:39.058601Z","shell.execute_reply.started":"2026-04-24T21:18:39.0474Z","shell.execute_reply":"2026-04-24T21:18:39.057694Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Debug: check tumor_xray sample images\nfor cls in CLASS_NAMES:\n    candidates = df[df[\"label\"] == cls].drop_duplicates(\"image\")\n    if len(candidates) == 0:\n        print(f\"{cls}: NO DATA\")\n        continue\n    sample = candidates.iloc[0]\n    ip = sample[\"image_path\"]\n    img = cv2.imread(str(ip))\n    loaded = \"✓\" if img is not None else \"✗ FAILED\"\n    \n    # Check bbox sizes\n    cls_df = df[df[\"label\"] == cls]\n    avg_rel = cls_df[\"relative_area_pct\"].mean()\n    max_rel = cls_df[\"relative_area_pct\"].max()\n    \n    print(f\"{cls:14s}: {loaded} | path: {ip}\")\n    print(f\"{'':14s}  bboxes: {len(cls_df)} | avg area: {avg_rel:.1f}% | max area: {max_rel:.1f}%\")\n    print(f\"{'':14s}  sample bbox: x={sample['xmin']:.0f},{sample['ymin']:.0f} → {sample['xmax']:.0f},{sample['ymax']:.0f}\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:18:56.310324Z","iopub.execute_input":"2026-04-24T21:18:56.310842Z","iopub.status.idle":"2026-04-24T21:18:56.394982Z","shell.execute_reply.started":"2026-04-24T21:18:56.310812Z","shell.execute_reply":"2026-04-24T21:18:56.394216Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 14. Dataset Class","metadata":{}},{"cell_type":"code","source":"class LungDiseaseDataset(Dataset):\n    \"\"\"Applies preprocessing + Albumentations + optional Copy-Paste.\"\"\"\n    # Classes that benefit from the small-lesion transform (sharpen + narrower scale)\n    SMALL_LESION_CLASSES = {\"tumor_ct\", \"tumor_xray\"}\n\n    def __init__(self, img_keys, df, class_to_idx, transform,\n                 preprocess_fn=None, copy_paste_bank=None, copy_paste_prob=0.0,\n                 tb_transform=None, small_lesion_transform=None):\n        self.img_keys = img_keys\n        self.df = df.set_index(\"image\")\n        self.class_to_idx = class_to_idx\n        self.transform = transform\n        self.tb_transform = tb_transform  # stronger aug applied when TB box is present\n        self.small_lesion_transform = small_lesion_transform  # for tumor_ct / tumor_xray\n        self.preprocess_fn = preprocess_fn\n        self.cp_bank = copy_paste_bank\n        self.cp_prob = copy_paste_prob\n\n    def __len__(self):\n        return len(self.img_keys)\n\n    def _get_rows(self, nm):\n        rows = self.df.loc[[nm]] if nm in self.df.index else self.df.loc[nm:nm]\n        if isinstance(rows, pd.Series):\n            rows = rows.to_frame().T\n        return rows\n\n    def __getitem__(self, idx):\n        nm   = self.img_keys[idx]\n        rows = self._get_rows(nm)\n        ip   = rows[\"image_path\"].iloc[0]\n\n        img = cv2.imread(str(ip))\n        if img is None:\n            img = np.zeros((train_sz, train_sz, 3), dtype=np.uint8)\n\n        # Apply the SHARED preprocessing (same function used at inference)\n        # Modality routes CT slices through tighter percentile + stronger CLAHE.\n        _modality = rows[\"modality\"].iloc[0] if \"modality\" in rows.columns else \"xray\"\n        if self.preprocess_fn is not None:\n            try:\n                img = self.preprocess_fn(img, modality=_modality)\n            except TypeError:\n                img = self.preprocess_fn(img)\n\n        bboxes = rows[[\"xmin\",\"ymin\",\"xmax\",\"ymax\"]].values.astype(float).tolist()\n        labels = rows[\"label\"].tolist()\n\n        # Copy-paste for rare classes\n        # T1.2: CP bank holds only xray crops (TB / tumor_xray); skip on CT images\n        # so we don't composite xray lesions onto CT slices (cross-modality nonsense).\n        if self.cp_bank and _modality != \"ct\" and random.random() < self.cp_prob:\n            img, bboxes, labels = apply_copy_paste(img, bboxes, labels, self.cp_bank)\n\n        # Filter any degenerate bboxes & clamp to image bounds before Albumentations\n        h_img, w_img = img.shape[:2]\n        clean_b, clean_l = [], []\n        for (x1, y1, x2, y2), l in zip(bboxes, labels):\n            # Clamp to image dimensions (fixes 1-2px overshoot from rounding)\n            x1 = max(0.0, min(float(x1), w_img))\n            y1 = max(0.0, min(float(y1), h_img))\n            x2 = max(0.0, min(float(x2), w_img))\n            y2 = max(0.0, min(float(y2), h_img))\n            if x2 > x1 + 2 and y2 > y1 + 2:\n                clean_b.append([x1, y1, x2, y2]); clean_l.append(l)\n\n        # Transform routing priority:\n        #   TB present         → tb_transform (strongest, breaks spatial prior)\n        #   small-lesion class → small_lesion_transform (sharpen, narrower scale)\n        #   otherwise          → baseline transform\n        tfm = self.transform\n        if self.tb_transform is not None and \"tuberculosis\" in clean_l:\n            tfm = self.tb_transform\n        elif self.small_lesion_transform is not None and any(\n            l in self.SMALL_LESION_CLASSES for l in clean_l\n        ):\n            tfm = self.small_lesion_transform\n\n        augmented = tfm(image=img, bboxes=clean_b, labels=clean_l)\n        img_t = augmented[\"image\"].float()\n        if img_t.dtype == torch.uint8:\n            img_t = img_t.float() / 255.0\n\n        final_bboxes = augmented[\"bboxes\"]\n        final_labels = augmented[\"labels\"]\n\n        if len(final_bboxes) == 0:\n            boxes_t  = torch.zeros((0, 4), dtype=torch.float32)\n            labels_t = torch.zeros((0,), dtype=torch.long)\n        else:\n            boxes_t  = torch.tensor(final_bboxes, dtype=torch.float32)\n            labels_t = torch.tensor([self.class_to_idx[l] for l in final_labels], dtype=torch.long)\n\n        # Multi-hot image-level label vector for the auxiliary classification head.\n        img_label = torch.zeros(len(self.class_to_idx), dtype=torch.float32)\n        for cls_name in set(final_labels):\n            if cls_name in self.class_to_idx:\n                img_label[self.class_to_idx[cls_name]] = 1.0\n\n        target = {\n            \"boxes\": BoundingBoxes(boxes_t, format=\"xyxy\",\n                                   canvas_size=(img_t.shape[1], img_t.shape[2])),\n            \"labels\": labels_t,\n            \"img_label\": img_label,\n        }\n        return img_t, target\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:18:59.85676Z","iopub.execute_input":"2026-04-24T21:18:59.857301Z","iopub.status.idle":"2026-04-24T21:18:59.871511Z","shell.execute_reply.started":"2026-04-24T21:18:59.857275Z","shell.execute_reply":"2026-04-24T21:18:59.870681Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 15. Mosaic + MixUp Wrapper\n\nStandard YOLOX augmentations implemented as a dataset wrapper. **In v8 both Mosaic and MixUp are disabled** (`mosaic_prob = 0.0`, `mixup_prob = 0.0`) — they hurt more than they help on 3 structurally-similar X-ray classes with ~10 k images. The wrapper still exists so downstream code doesn't break and so you can toggle them back on for experiments; copy-paste (driven by `copy_paste_prob = 0.5`) remains active from epoch 1.\n","metadata":{}},{"cell_type":"code","source":"class MosaicMixUpDataset(Dataset):\n    \"\"\"Wraps a base dataset and returns either:\n       - A mosaic of 4 images (with prob mosaic_prob), or\n       - A mixup of 2 mosaics (with small prob), or\n       - A single image (otherwise, passes through to base dataset).\n\n    Call .disable_mosaic() near end of training.\n    \"\"\"\n    def __init__(self, base_dataset, mosaic_prob=1.0, mixup_prob=0.15, img_size=768,\n                 sample_weights=None):\n        self.base = base_dataset\n        self.mosaic_prob = mosaic_prob\n        self.mixup_prob  = mixup_prob\n        self.img_size = img_size\n        self.enabled = True\n        # T2.5: per-index weights so the other 3 mosaic tiles are class-balanced,\n        # not uniformly sampled (otherwise pneumonia dominates 3/4 of every mosaic).\n        self.sample_weights = list(sample_weights) if sample_weights is not None else None\n\n    def disable_mosaic(self):\n        self.enabled = False\n\n    def __len__(self):\n        return len(self.base)\n\n    def _load_one(self, idx):\n        img_t, target = self.base[idx]\n        return img_t, target[\"boxes\"], target[\"labels\"]\n\n    def _mosaic(self, idx):\n        \"\"\"Stitch 4 images into a 2*S x 2*S mosaic, then crop back to S.\"\"\"\n        S = self.img_size\n        # Random center\n        cx = int(random.uniform(S * 0.5, S * 1.5))\n        cy = int(random.uniform(S * 0.5, S * 1.5))\n        # T2.5: weighted selection for the 3 extra tiles\n        if self.sample_weights is not None:\n            extras = random.choices(range(len(self.base)), weights=self.sample_weights, k=3)\n        else:\n            extras = random.sample(range(len(self.base)), 3)\n        indices = [idx] + extras\n\n        mosaic_img = torch.zeros(3, S*2, S*2, dtype=torch.float32)\n        all_boxes, all_labels = [], []\n\n        for i, ix in enumerate(indices):\n            img_t, boxes, labels = self._load_one(ix)\n            _, h, w = img_t.shape\n            # Determine placement region within mosaic\n            if i == 0:    # top-left\n                x1a, y1a, x2a, y2a = max(cx-w, 0), max(cy-h, 0), cx, cy\n                x1b, y1b, x2b, y2b = w-(x2a-x1a), h-(y2a-y1a), w, h\n            elif i == 1:  # top-right\n                x1a, y1a, x2a, y2a = cx, max(cy-h, 0), min(cx+w, S*2), cy\n                x1b, y1b, x2b, y2b = 0, h-(y2a-y1a), min(w, x2a-x1a), h\n            elif i == 2:  # bottom-left\n                x1a, y1a, x2a, y2a = max(cx-w, 0), cy, cx, min(cy+h, S*2)\n                x1b, y1b, x2b, y2b = w-(x2a-x1a), 0, w, min(h, y2a-y1a)\n            else:         # bottom-right\n                x1a, y1a, x2a, y2a = cx, cy, min(cx+w, S*2), min(cy+h, S*2)\n                x1b, y1b, x2b, y2b = 0, 0, min(w, x2a-x1a), min(h, y2a-y1a)\n\n            mosaic_img[:, y1a:y2a, x1a:x2a] = img_t[:, y1b:y2b, x1b:x2b]\n            padw, padh = x1a - x1b, y1a - y1b\n            if len(boxes) > 0:\n                shifted = boxes.clone().float()\n                shifted[:, [0, 2]] += padw\n                shifted[:, [1, 3]] += padh\n                all_boxes.append(shifted); all_labels.append(labels)\n\n        if all_boxes:\n            all_boxes  = torch.cat(all_boxes, dim=0)\n            all_labels = torch.cat(all_labels, dim=0)\n            # Clip\n            all_boxes[:, [0, 2]] = all_boxes[:, [0, 2]].clamp(0, S*2)\n            all_boxes[:, [1, 3]] = all_boxes[:, [1, 3]].clamp(0, S*2)\n            # Drop degenerate\n            valid = (all_boxes[:, 2] - all_boxes[:, 0] >= 4) & (all_boxes[:, 3] - all_boxes[:, 1] >= 4)\n            all_boxes  = all_boxes[valid]\n            all_labels = all_labels[valid]\n        else:\n            all_boxes  = torch.zeros((0, 4), dtype=torch.float32)\n            all_labels = torch.zeros((0,),   dtype=torch.long)\n\n        # v6 FIX: Random-crop SxS from the 2Sx2S canvas instead of resizing.\n        # Resizing was halving all box dimensions, making small lesions undetectable.\n        crop_x = random.randint(0, S)\n        crop_y = random.randint(0, S)\n        mosaic_img = mosaic_img[:, crop_y:crop_y+S, crop_x:crop_x+S]\n        if len(all_boxes) > 0:\n            all_boxes[:, [0, 2]] -= crop_x\n            all_boxes[:, [1, 3]] -= crop_y\n            all_boxes[:, [0, 2]] = all_boxes[:, [0, 2]].clamp(0, S)\n            all_boxes[:, [1, 3]] = all_boxes[:, [1, 3]].clamp(0, S)\n            # Drop boxes that are now too small after cropping\n            valid = (all_boxes[:, 2] - all_boxes[:, 0] >= 4) & (all_boxes[:, 3] - all_boxes[:, 1] >= 4)\n            all_boxes = all_boxes[valid]\n            all_labels = all_labels[valid]\n\n        return mosaic_img, all_boxes, all_labels\n\n    def __getitem__(self, idx):\n        if not self.enabled or random.random() > self.mosaic_prob:\n            img_t, target = self.base[idx]\n            return img_t, target\n\n        img_t, boxes, labels = self._mosaic(idx)\n\n        # MixUp with another mosaic (small probability)\n        if random.random() < self.mixup_prob:\n            idx2 = random.randint(0, len(self.base) - 1)\n            img2_t, boxes2, labels2 = self._mosaic(idx2)\n            lam = np.random.beta(32.0, 32.0)  # narrow around 0.5\n            img_t = lam * img_t + (1 - lam) * img2_t\n            if len(boxes2) > 0:\n                boxes  = torch.cat([boxes, boxes2], dim=0) if len(boxes) > 0 else boxes2\n                labels = torch.cat([labels, labels2], dim=0) if len(labels) > 0 else labels2\n\n        S = self.img_size\n        target = {\n            \"boxes\": BoundingBoxes(boxes, format=\"xyxy\", canvas_size=(S, S)),\n            \"labels\": labels,\n        }\n        return img_t, target\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:19:08.655966Z","iopub.execute_input":"2026-04-24T21:19:08.656538Z","iopub.status.idle":"2026-04-24T21:19:08.677945Z","shell.execute_reply.started":"2026-04-24T21:19:08.656507Z","shell.execute_reply":"2026-04-24T21:19:08.677157Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\ngc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T21:19:43.012422Z","iopub.execute_input":"2026-04-24T21:19:43.012685Z","iopub.status.idle":"2026-04-24T21:19:43.454509Z","shell.execute_reply.started":"2026-04-24T21:19:43.012658Z","shell.execute_reply":"2026-04-24T21:19:43.453609Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 16. Build Datasets & DataLoaders","metadata":{}},{"cell_type":"code","source":"if len(df) > 0:\n    train_tfm        = build_train_transform(train_sz, norm_stats)\n    tb_train_tfm     = build_tb_train_transform(train_sz, norm_stats)\n    small_lesion_tfm = build_small_lesion_transform(train_sz, norm_stats)\n    valid_tfm        = build_valid_transform(train_sz, norm_stats)\n    map_tfm          = build_map_transform(train_sz)\n\n    base_train_ds = LungDiseaseDataset(\n        train_keys, df, CLASS_TO_IDX, transform=train_tfm,\n        preprocess_fn=preprocess_image,\n        copy_paste_bank=copy_paste_bank, copy_paste_prob=copy_paste_prob,\n        tb_transform=tb_train_tfm,\n        small_lesion_transform=small_lesion_tfm,\n    )\n    train_dataset = MosaicMixUpDataset(base_train_ds, mosaic_prob=mosaic_prob,\n                                       mixup_prob=mixup_prob, img_size=train_sz,\n                                       sample_weights=tsw)\n\n    valid_dataset = LungDiseaseDataset(\n        val_keys, df, CLASS_TO_IDX, transform=valid_tfm,\n        preprocess_fn=preprocess_image,\n    )\n\n    map_dataset = LungDiseaseDataset(\n        val_keys, df, CLASS_TO_IDX, transform=map_tfm,\n        preprocess_fn=preprocess_image,\n    )\n\n    def collate_fn(batch):\n        return tuple(zip(*batch))\n\n    train_loader = DataLoader(\n        train_dataset, batch_size=bs, sampler=sampler,\n        num_workers=0,\n        collate_fn=collate_fn, pin_memory=(\"cuda\" in device),\n        drop_last=True,\n    )\n    valid_loader = DataLoader(\n        valid_dataset, batch_size=bs,\n        num_workers=0,\n        collate_fn=collate_fn, pin_memory=(\"cuda\" in device),\n        drop_last=False,\n    )\n    map_loader = DataLoader(\n        map_dataset, batch_size=bs,\n        num_workers=0,\n        collate_fn=collate_fn,\n        drop_last=False,\n    )\n    print(f\"Train: {len(train_loader)} batches (bs={bs}, grad_accum={grad_accum_steps}, \"\n          f\"effective={bs*grad_accum_steps}) | Val: {len(valid_loader)} | mAP: {len(map_loader)}\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:25:53.58969Z","iopub.execute_input":"2026-04-24T21:25:53.590564Z","iopub.status.idle":"2026-04-24T21:25:53.605745Z","shell.execute_reply.started":"2026-04-24T21:25:53.590536Z","shell.execute_reply":"2026-04-24T21:25:53.605004Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if len(df) > 0:\n    print(\"=\" * 60)\n    print(\"DIAGNOSTIC: Checking training DataLoader output\")\n    print(\"=\" * 60)\n \n    empty_count = 0\n    total_boxes = 0\n    total_images = 0\n \n    for i, (imgs, targets) in enumerate(train_loader):\n        if i >= 50:\n            break\n        for t in targets:\n            n = len(t[\"labels\"])\n            total_boxes += n\n            total_images += 1\n            if n == 0:\n                empty_count += 1\n \n    pct_empty = empty_count / max(1, total_images) * 100\n    avg_boxes = total_boxes / max(1, total_images)\n \n    print(f\"  Images sampled:  {total_images}\")\n    print(f\"  Total boxes:     {total_boxes}\")\n    print(f\"  Avg boxes/image: {avg_boxes:.1f}\")\n    print(f\"  Empty images:    {empty_count} ({pct_empty:.1f}%)\")\n    print()\n \n    if pct_empty > 50:\n        print(\"  ⚠ CRITICAL: >50% of training images have ZERO bounding boxes!\")\n        print(\"    → The model cannot learn from empty targets.\")\n        print(\"    → Bbox coordinates are likely misaligned with images.\")\n        print(\"    → Make sure PATCH CELL 1 was applied before this point.\")\n    elif pct_empty > 20:\n        print(\"  ⚠ WARNING: >20% empty images. Augmentation may be dropping\")\n        print(\"    too many boxes. Try reducing augmentation strength.\")\n    else:\n        print(\"  ✓ Bbox pipeline looks healthy.\")\n \n    # Raw DataFrame sanity check\n    print()\n    print(\"Per-class bbox stats in DataFrame:\")\n    for cls in CLASS_NAMES:\n        sub = df[df[\"label\"] == cls]\n        if len(sub) == 0:\n            print(f\"  {cls}: NO DATA\"); continue\n        neg = ((sub[\"xmin\"] < 0) | (sub[\"ymin\"] < 0)).sum()\n        zero = (sub[\"area\"] <= 0).sum()\n        tiny = (sub[\"area\"] < 100).sum()\n        print(f\"  {cls:14s}: {len(sub):5d} boxes | \"\n              f\"negative: {neg} | zero-area: {zero} | tiny(<100px²): {tiny} | \"\n              f\"mean area: {sub['area'].mean():.0f}px²\")\n ","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:25:58.902627Z","iopub.execute_input":"2026-04-24T21:25:58.903158Z","iopub.status.idle":"2026-04-24T21:27:46.016397Z","shell.execute_reply.started":"2026-04-24T21:25:58.903129Z","shell.execute_reply":"2026-04-24T21:27:46.015562Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 17. Final Data Integrity Validation","metadata":{}},{"cell_type":"code","source":"if len(df) > 0:\n    print(\"=\" * 60)\n    print(\"FINAL DATA INTEGRITY VALIDATION\")\n    print(\"=\" * 60)\n    ok = 0\n\n    # 1. No leakage\n    if (not (set(train_keys) & set(val_keys))\n        and not (set(train_keys) & set(test_keys))\n        and not (set(val_keys)   & set(test_keys))):\n        print(\"✓ [1/5] No train/val/test leakage\"); ok += 1\n\n    # 2. Images load\n    errs = sum(1 for i in random.sample(train_keys, min(50, len(train_keys)))\n               if cv2.imread(df[df[\"image\"] == i][\"image_path\"].iloc[0]) is None)\n    if errs == 0:\n        print(\"✓ [2/5] All sampled images load\"); ok += 1\n    else:\n        print(f\"✗ [2/5] {errs} load failures\")\n\n    # 3. BBox bounds (post-crop)\n    be = 0\n    for _, r in df.sample(min(200, len(df))).iterrows():\n        img = cv2.imread(r[\"image_path\"])\n        if img is not None:\n            h, w = img.shape[:2]\n            if r[\"xmin\"] < 0 or r[\"ymin\"] < 0 or r[\"xmax\"] > w + 2 or r[\"ymax\"] > h + 2:\n                be += 1\n    print(f\"{'✓' if be == 0 else '✗'} [3/5] BBoxes within bounds (OOB: {be})\")\n    if be == 0:\n        ok += 1\n\n    # 4. DataLoader\n    try:\n        for bi, bt in train_loader:\n            assert len(bi) == bs\n            break\n        print(\"✓ [4/5] DataLoader OK\"); ok += 1\n    except Exception as e:\n        print(f\"✗ [4/5] {e}\")\n\n    # 5. Class coverage\n    tl = df[df[\"image\"].isin(train_keys)][\"label\"].unique()\n    if set(tl) == set(CLASS_NAMES):\n        print(\"✓ [5/5] All classes in train\"); ok += 1\n    else:\n        missing = set(CLASS_NAMES) - set(tl)\n        print(f\"✗ [5/5] Missing from train: {missing}\")\n\n    print(f\"\\nRESULT: {ok}/5 passed {'✓ READY' if ok == 5 else ''}\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:28:47.950928Z","iopub.execute_input":"2026-04-24T21:28:47.951162Z","iopub.status.idle":"2026-04-24T21:28:55.052355Z","shell.execute_reply.started":"2026-04-24T21:28:47.951135Z","shell.execute_reply":"2026-04-24T21:28:55.051593Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 18. Model & Loss Setup\n\n- **YOLOX-s** (~9 M params). First-run weights download is ~35 MB.\n- `bbox_loss_weight = 5.0` — Megvii default. Earlier 10.0 was starving the classifier.\n- **Differential LR via AdamW param groups** — backbone @ `lr * 0.1`, head/neck @ `lr`. Set up in Cell 71.\n- **EMA of weights** — maintained alongside training weights, used for validation and inference.\n\nThe v6+ config is already memory-safe on T4 (~7–9 GB peak), so no OOM fallback table is needed. If you do need to free headroom, drop `bs` to 8 first — everything else is already at the conservative setting.\n","metadata":{}},{"cell_type":"code","source":"from copy import deepcopy\n\n\nclass ModelEMA:\n    \"\"\"Exponential Moving Average of model weights.\n    Used at validation/inference; gives a consistent +1-2 mAP for ~20 LoC.\n    \"\"\"\n    def __init__(self, model, decay=0.9998):\n        self.ema = deepcopy(model).eval()\n        for p in self.ema.parameters():\n            p.requires_grad_(False)\n        self.decay = decay\n        self.updates = 0\n\n    @torch.no_grad()\n    def update(self, model):\n        self.updates += 1\n        # Warmup decay: ramp up from 0 to `decay` over first ~1000 updates\n        d = self.decay * (1 - math.exp(-self.updates / 1000))\n        msd = model.state_dict()\n        for k, v in self.ema.state_dict().items():\n            if v.dtype.is_floating_point:\n                v.mul_(d).add_(msd[k].detach(), alpha=1 - d)\n\n    def state_dict(self):\n        return {\"ema\": self.ema.state_dict(), \"updates\": self.updates, \"decay\": self.decay}\n\n    def load_state_dict(self, sd):\n        self.ema.load_state_dict(sd[\"ema\"])\n        self.updates = sd[\"updates\"]\n        self.decay   = sd[\"decay\"]\n\n\nif len(df) > 0:\n    model_type = \"yolox_s\"     # v4: yolox_s (yolox_tiny was under-powered at 1024)\n    print(f\"Building {model_type} ({NUM_CLASSES} classes) — first-run weights download ~100MB\")\n    # 1. Build model WITHOUT pretrained (avoids Broken HuggingFace download)\n    model = build_model(model_type, NUM_CLASSES, pretrained=False).to(device)\n    \n    # 2. Load Megvii official YOLOX-s weights manually\n    # IMPORTANT: Adjust this path to wherever your yolox_s.pth is located on Kaggle\n    _ckpt_path = \"/kaggle/input/datasets/mickgt/yolo-s-weights/yolox-s-weights/yolox_s.pth\"  \n    # /kaggle/input/private-data-source/yolox-s-weights/yolox_s.pth\n    import os\n    if os.path.exists(_ckpt_path):\n        _ckpt = torch.load(_ckpt_path, map_location=device)\n        _sd = _ckpt.get(\"model\", _ckpt)\n        \n        _model_sd = model.state_dict()\n        _loaded, _skipped = 0, 0\n        for k, v in _sd.items():\n            if k in _model_sd and _model_sd[k].shape == v.shape:\n                _model_sd[k] = v\n                _loaded += 1\n            else:\n                _skipped += 1\n        model.load_state_dict(_model_sd)\n        print(f\"\\n\\u2713 Loaded pretrained backbone: {_loaded} params, {_skipped} skipped (head shape mismatch)\\n\")\n    else:\n        print(f\"\\n\\u26A0 WARNING: Could not find {_ckpt_path}. Training from scratch!\\n\")\n\n    loss_func = YOLOXLoss(num_classes=NUM_CLASSES, bbox_loss_weight=5.0)\n\n    # Differential LR: backbone fine-tunes at 0.1x head LR. Standard practice for\n    # transfer learning — lets pretrained features slowly adapt to medical imaging\n    # without shocking them, AND avoids the freeze/unfreeze collapse we hit before.\n    _BACKBONE_KEY_MATCH = (\"backbone\", \"stem\", \"dark\")\n    _bb_params, _head_params = [], []\n    for _n, _p in model.named_parameters():\n        if not _p.requires_grad:\n            continue\n        if any(k in _n.lower() for k in _BACKBONE_KEY_MATCH):\n            _bb_params.append(_p)\n        else:\n            _head_params.append(_p)\n    print(f\"Optimizer param groups: backbone={len(_bb_params)} @ {lr*0.1:.1e}, \"\n          f\"head/neck={len(_head_params)} @ {lr:.1e}\")\n    optimizer = torch.optim.AdamW([\n        {\"params\": _bb_params,   \"lr\": lr * 0.1},\n        {\"params\": _head_params, \"lr\": lr},\n    ], weight_decay=weight_decay)\n\n    ema = ModelEMA(model, decay=ema_decay) if use_ema else None\n\n    # ── Image-level auxiliary classification head ──────────────────────────\n    # Forward hook captures the deepest backbone feature map without\n    # touching the model definition. Triggered on every model.forward(x).\n    if USE_AUX_CLS_HEAD:\n        class _FeatureCapture:\n            def __init__(self):\n                self.feat = None\n            def __call__(self, module, inp, out):\n                if isinstance(out, (list, tuple)) and len(out) > 0:\n                    self.feat = out[-1]\n                else:\n                    self.feat = out\n        feat_capture = _FeatureCapture()\n        model.backbone.register_forward_hook(feat_capture)\n\n        # Probe channel count with a dummy forward\n        model.eval()\n        with torch.no_grad():\n            _dummy = torch.zeros(1, 3, train_sz, train_sz, device=device)\n            _ = model(_dummy)\n        model.train()\n        _feat_c = feat_capture.feat.shape[1]\n        print(f\"Aux cls head: feature channels = {_feat_c}\")\n\n        cls_head = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.Dropout(0.2),\n            nn.Linear(_feat_c, NUM_CLASSES),\n        ).to(device)\n\n        optimizer.add_param_group({\"params\": cls_head.parameters(),\n                                   \"lr\": lr, \"weight_decay\": weight_decay})\n        print(f\"Aux cls head added: {sum(p.numel() for p in cls_head.parameters())} params, weight={AUX_CLS_LOSS_WEIGHT}\")\n    else:\n        feat_capture = None\n        cls_head = None\n\n\n    # Cosine schedule with linear warmup. IMPORTANT: scheduler steps count OPTIMIZER updates,\n    # not dataloader iterations, since we use gradient accumulation.\n    from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR\n    opt_steps_per_epoch = max(1, len(train_loader) // grad_accum_steps)\n    warmup_steps = warmup_epochs * opt_steps_per_epoch\n    total_steps  = epochs * opt_steps_per_epoch\n    warmup  = LinearLR(optimizer, start_factor=0.01, end_factor=1.0, total_iters=warmup_steps)\n    cosine  = CosineAnnealingLR(optimizer, T_max=total_steps - warmup_steps, eta_min=lr * 0.05)  # T3.9\n    scheduler = SequentialLR(optimizer, [warmup, cosine], milestones=[warmup_steps])\n\n\n    print(f\"Model:      {model_type} @ {train_sz}x{train_sz}\")\n    print(f\"Optimizer:  AdamW, lr={lr}, wd={weight_decay}\")\n    print(f\"Schedule:   {warmup_epochs}-epoch warmup → Cosine, {total_steps} optimizer steps\")\n    print(f\"Grad accum: {grad_accum_steps} (effective bs = {bs * grad_accum_steps})\")\n    print(f\"EMA:        {'on (decay=' + str(ema_decay) + ')' if use_ema else 'off'}\")\n    print(f\"Mosaic:     on for first {mosaic_epochs} epochs (prob={mosaic_prob})\")\n    print(f\"ES:         mAP@0.5, patience={es_patience}\")\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T21:29:25.3895Z","iopub.execute_input":"2026-04-24T21:29:25.389838Z","iopub.status.idle":"2026-04-24T21:29:25.758536Z","shell.execute_reply.started":"2026-04-24T21:29:25.389813Z","shell.execute_reply":"2026-04-24T21:29:25.757854Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 19. Training Loop\n\n**v8 training-loop behavior:**\n- **Single optimizer step per batch** (`grad_accum_steps = 1`) at `bs=16` — scheduler advances per step\n- **Differential LR param groups** — backbone at 0.1× head LR throughout; no freeze, no unfreeze ramp\n- **EMA of weights** updated after every optimizer step; used for validation and mAP evaluation\n- **Checkpoint resumption** every 5 epochs — saves model/optimizer/scheduler/scaler/EMA state. v4 fixed the resume-OOM by freeing the checkpoint dict + `gc.collect()` before the DataLoader forks workers.\n- **mAP + early stopping** computed every epoch on the EMA weights\n- **No multi-scale, no mosaic, no mixup** — copy-paste is the only heavy augmentation left\n","metadata":{}},{"cell_type":"code","source":"CKPT_LATEST = ckpt_dir / \"latest.pth\"\nCKPT_BEST   = ckpt_dir / \"best.pth\"\n\n\ndef save_checkpoint(path, model, ema, optimizer, scheduler, scaler, epoch,\n                    best_map, es_counter, history):\n    state = {\n        \"model\": model.state_dict(),\n        \"ema\": ema.state_dict() if ema is not None else None,\n        \"cls_head\": cls_head.state_dict() if cls_head is not None else None,\n        \"optimizer\": optimizer.state_dict(),\n        \"scheduler\": scheduler.state_dict(),\n        \"scaler\": scaler.state_dict() if scaler is not None else None,\n        \"epoch\": epoch,\n        \"best_map\": best_map,\n        \"es_counter\": es_counter,\n        \"history\": history,\n        \"config\": {\n            \"model_type\": model_type, \"train_sz\": train_sz,\n            \"bs\": bs, \"grad_accum_steps\": grad_accum_steps,\n            \"class_names\": CLASS_NAMES,\n        },\n    }\n    torch.save(state, path)\n\n\ndef try_resume(path, model, ema, optimizer, scheduler, scaler):\n    if not path.exists():\n        return 0, -1.0, 0, []\n    print(f\"Resuming from {path}\")\n    # Load to CPU first so GPU isn't holding a duplicate while we dissect ck.\n    ck = torch.load(path, map_location=\"cpu\")\n    model.load_state_dict(ck[\"model\"]); del ck[\"model\"]\n    if ema is not None and ck.get(\"ema\") is not None:\n        ema.load_state_dict(ck[\"ema\"]); ck[\"ema\"] = None\n    if cls_head is not None and ck.get(\"cls_head\") is not None:\n        cls_head.load_state_dict(ck[\"cls_head\"]); ck[\"cls_head\"] = None\n    optimizer.load_state_dict(ck[\"optimizer\"]); del ck[\"optimizer\"]\n    scheduler.load_state_dict(ck[\"scheduler\"]); del ck[\"scheduler\"]\n    if scaler is not None and ck.get(\"scaler\") is not None:\n        scaler.load_state_dict(ck[\"scaler\"]); ck[\"scaler\"] = None\n    start_epoch = ck[\"epoch\"] + 1\n    best_map    = ck.get(\"best_map\", -1.0)\n    es_counter  = ck.get(\"es_counter\", 0)\n    history     = ck.get(\"history\", [])\n    # Release checkpoint dict before first DataLoader iteration so worker fork\n    # (when num_workers>0) doesn't inherit bloated parent RAM.\n    del ck\n    import gc\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    # Move optimizer state to GPU (map_location='cpu' leaves momentum on CPU).\n    for state in optimizer.state.values():\n        for k, v in state.items():\n            if isinstance(v, torch.Tensor):\n                state[k] = v.to(device)\n    print(f\"  Resumed: epoch {start_epoch}, best_map={best_map:.4f}, history={len(history)} entries\")\n    return start_epoch, best_map, es_counter, history\n\n\n# Per-class NMS IoU thresholds. YOLOX-standard default is 0.65. We use slightly\n# tighter values for small-object classes (tumor_*) so nearby but distinct\n# lesions aren't suppressed into one box.\nPER_CLASS_IOU_DEFAULT = {\n    \"tumor_ct\":     0.50,\n    \"tumor_xray\":   0.50,\n    \"tuberculosis\": 0.60,\n    \"pneumonia\":    0.60,\n}\n\n\ndef compute_map(model_or_ema, map_loader, device, norm_stats,\n                conf_thresh=0.01,         # Canonical YOLOX mAP eval uses 0.01\n                iou_thresh=0.5,            # mAP IoU (fixed — matched the metric init)\n                nms_iou=0.65,              # fallback NMS IoU if class not in per_class_iou\n                per_class_iou=None):\n    \"\"\"mAP@0.5 with per-class NMS. Replaces WBF (which is for multi-model\n    ensembling, not single-model dedup). Preds are capped at 300/image.\n\n    Input normalization contract:\n      - map_loader yields float32 tensors in [0, 255] (build_map_transform has no\n        A.Normalize; LungDiseaseDataset.__getitem__ does .float() without /255).\n      - YOLOXInferenceWrapper(scale_inp=True) divides by 255 and applies mean/std,\n        which matches training's A.Normalize(max_pixel_value=255) step.\n    \"\"\"\n    if per_class_iou is None:\n        per_class_iou = PER_CLASS_IOU_DEFAULT\n    eval_model = model_or_ema.ema if isinstance(model_or_ema, ModelEMA) else model_or_ema\n    mean_t = torch.tensor(norm_stats[0]).view(1, 3, 1, 1).to(device)\n    std_t  = torch.tensor(norm_stats[1]).view(1, 3, 1, 1).to(device)\n\n    wrapper = YOLOXInferenceWrapper(eval_model, mean_t, std_t, scale_inp=True)\n    metric = MeanAveragePrecision(iou_thresholds=[0.5], class_metrics=True)\n\n    eval_model.eval()\n    total_preds = 0\n    total_gts   = 0\n\n    with torch.no_grad():\n        for inputs, targets in tqdm(map_loader, desc=\"mAP@0.5\", leave=False):\n            inputs_t = torch.stack(inputs).to(device)\n            # inputs_t is float32 [0, 255] by dataset contract; no /255 here —\n            # the wrapper does it with scale_inp=True.\n            if inputs_t.dtype == torch.uint8:\n                inputs_t = inputs_t.float()\n\n            with autocast(device_type=torch.device(device).type):\n                out = wrapper(inputs_t)\n\n            img_h, img_w = inputs_t.shape[-2], inputs_t.shape[-1]\n\n            for b_idx, tgt in enumerate(targets):\n                preds = out[b_idx].cpu()\n                mask  = preds[:, -1] > conf_thresh\n                p = preds[mask]\n\n                total_gts += len(tgt[\"labels\"])\n\n                if len(p) > 0:\n                    # Wrapper output is (x0, y0, w, h, label, prob) — xywh top-left\n                    boxes_xyxy = torchvision.ops.box_convert(p[:, :4], \"xywh\", \"xyxy\").float()\n                    # Clamp to image bounds (dense head can emit slightly-OOB boxes)\n                    boxes_xyxy[:, [0, 2]] = boxes_xyxy[:, [0, 2]].clamp(0, img_w)\n                    boxes_xyxy[:, [1, 3]] = boxes_xyxy[:, [1, 3]].clamp(0, img_h)\n                    # Drop degenerate\n                    valid = (boxes_xyxy[:, 2] - boxes_xyxy[:, 0] >= 1) & \\\n                            (boxes_xyxy[:, 3] - boxes_xyxy[:, 1] >= 1)\n                    boxes_xyxy = boxes_xyxy[valid]\n                    scores     = p[:, 5].float()[valid]\n                    labels     = p[:, 4].long()[valid]\n\n                    # Per-class NMS — standard YOLOX eval path\n                    kept_b, kept_s, kept_l = [], [], []\n                    for cid in labels.unique().tolist():\n                        m = (labels == cid)\n                        cls_name = CLASS_NAMES[cid] if 0 <= cid < len(CLASS_NAMES) else \"\"\n                        iou_thr_c = per_class_iou.get(cls_name, nms_iou)\n                        keep = torchvision.ops.nms(\n                            boxes_xyxy[m], scores[m], iou_threshold=iou_thr_c\n                        )\n                        kept_b.append(boxes_xyxy[m][keep])\n                        kept_s.append(scores[m][keep])\n                        kept_l.append(labels[m][keep])\n\n                    if kept_b:\n                        all_b = torch.cat(kept_b, dim=0)\n                        all_s = torch.cat(kept_s, dim=0)\n                        all_l = torch.cat(kept_l, dim=0)\n                        # Sort by score desc, cap at 300 per image\n                        order = all_s.argsort(descending=True)[:300]\n                        pred_dict = dict(\n                            boxes  = all_b[order],\n                            scores = all_s[order],\n                            labels = all_l[order].long(),\n                        )\n                    else:\n                        pred_dict = dict(\n                            boxes  = torch.zeros(0, 4),\n                            scores = torch.zeros(0),\n                            labels = torch.zeros(0, dtype=torch.long),\n                        )\n                    total_preds += len(pred_dict[\"boxes\"])\n                else:\n                    pred_dict = dict(\n                        boxes  = torch.zeros(0, 4),\n                        scores = torch.zeros(0),\n                        labels = torch.zeros(0, dtype=torch.long),\n                    )\n\n                metric.update(\n                    [pred_dict],\n                    [dict(boxes=tgt[\"boxes\"].cpu().float(),\n                          labels=tgt[\"labels\"].cpu().long())],\n                )\n\n    print(f\"  [mAP] Predictions kept (conf>{conf_thresh}, after NMS): {total_preds} | \"\n          f\"GT boxes: {total_gts}\")\n    return metric.compute()\n\n\ndef train_one_epoch(model, ema, loader, optimizer, scheduler, loss_func, device,\n                    scaler, epoch_idx, grad_accum_steps=1,\n                    multi_scale=False, scale_range=(640, 896),\n                    unfreeze_ramp_steps=0, unfreeze_ramp_min_factor=0.1,\n                    freeze_bn_fn=None):\n    \"\"\"Training loop with gradient accumulation. `scheduler` advances per OPTIMIZER step.\n\n    If `unfreeze_ramp_steps > 0`, the first N optimizer steps of this epoch have\n    their LR multiplied by a factor that ramps linearly from `min_factor` -> 1.0.\n    Used at the freeze->unfreeze boundary so the newly-unfrozen backbone doesn't\n    take a single full-LR step the instant it becomes trainable. The scheduler's\n    own values are not modified — next epoch resumes the cosine uninterrupted.\n    \"\"\"\n    model.train()\n    if freeze_bn_fn is not None:\n        freeze_bn_fn(model)\n    total_loss = 0.0\n    n_batches = len(loader)\n    if n_batches == 0:\n        return 0.0\n\n    pbar = tqdm(loader, desc=f\"Train ep{epoch_idx+1}\")\n    current_scale = train_sz\n    optimizer.zero_grad(set_to_none=True)\n    opt_step_in_epoch = 0\n\n    for bid, (inputs, targets) in enumerate(pbar):\n        if multi_scale and bid % 10 == 0:\n            lo, hi = scale_range\n            current_scale = random.choice(list(range(lo, hi + 1, 32)))\n\n        inputs = torch.stack(inputs).to(device, non_blocking=True)\n        if inputs.dtype == torch.uint8:\n            inputs = inputs.float() / 255.0\n\n        if current_scale != inputs.shape[-1]:\n            scale_factor = current_scale / inputs.shape[-1]\n            inputs = F.interpolate(inputs, size=(current_scale, current_scale),\n                                    mode=\"bilinear\", align_corners=False)\n            new_targets = []\n            for t in targets:\n                b = t[\"boxes\"].float() * scale_factor\n                new_targets.append({\"boxes\": b, \"labels\": t[\"labels\"]})\n            targets = new_targets\n\n        gb = [t[\"boxes\"].to(device).float()  for t in targets]\n        gl = [t[\"labels\"].to(device).long()  for t in targets]\n\n        with autocast(device_type=torch.device(device).type):\n            cs, bp, os2 = model(inputs)\n            losses = loss_func(cs, bp, os2, gb, gl)\n            det_loss = sum(losses.values())\n\n            # Auxiliary image-level classification loss\n            if cls_head is not None and feat_capture is not None and feat_capture.feat is not None:\n                img_lbls = torch.stack([t[\"img_label\"] for t in targets]).to(device)\n                cls_logits = cls_head(feat_capture.feat)\n                cls_loss = F.binary_cross_entropy_with_logits(cls_logits, img_lbls)\n            else:\n                cls_loss = torch.zeros((), device=device)\n\n            loss = (det_loss + AUX_CLS_LOSS_WEIGHT * cls_loss) / grad_accum_steps\n\n        if scaler:\n            scaler.scale(loss).backward()\n        else:\n            loss.backward()\n\n        is_last = (bid + 1 == n_batches)\n        if (bid + 1) % grad_accum_steps == 0 or is_last:\n            if scaler:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 10.0)\n                scaler.step(optimizer); scaler.update()\n            else:\n                torch.nn.utils.clip_grad_norm_(model.parameters(), 10.0)\n                optimizer.step()\n            scheduler.step()\n\n            if unfreeze_ramp_steps > 0 and opt_step_in_epoch < unfreeze_ramp_steps:\n                t = opt_step_in_epoch / max(1, unfreeze_ramp_steps - 1)\n                ramp_factor = unfreeze_ramp_min_factor + (1.0 - unfreeze_ramp_min_factor) * t\n                for pg in optimizer.param_groups:\n                    pg[\"lr\"] *= ramp_factor\n            opt_step_in_epoch += 1\n\n            optimizer.zero_grad(set_to_none=True)\n\n            if ema is not None:\n                ema.update(model)\n\n        unscaled_loss = loss.item() * grad_accum_steps\n        total_loss += unscaled_loss\n        pbar.set_postfix(loss=f\"{unscaled_loss:.3f}\",\n                         avg=f\"{total_loss/(bid+1):.3f}\",\n                         sz=current_scale,\n                         lr=f\"{optimizer.param_groups[0]['lr']:.2e}\")\n\n        if not math.isfinite(unscaled_loss):\n            print(f\"Non-finite loss; stopping epoch {epoch_idx+1} early\")\n            break\n\n    return total_loss / max(1, n_batches)\n\n\ndef validate_one_epoch(model, loader, loss_func, device):\n    model.eval()\n    total_loss = 0.0\n    if len(loader) == 0:\n        return 0.0\n    with torch.no_grad():\n        for inputs, targets in tqdm(loader, desc=\"Val\", leave=False):\n            inputs = torch.stack(inputs).to(device)\n            if inputs.dtype == torch.uint8:\n                inputs = inputs.float() / 255.0\n            gb = [t[\"boxes\"].to(device).float() for t in targets]\n            gl = [t[\"labels\"].to(device).long() for t in targets]\n            with autocast(device_type=torch.device(device).type):\n                cs, bp, os2 = model(inputs)\n                losses = loss_func(cs, bp, os2, gb, gl)\n                # NOTE: cls_loss intentionally NOT computed in validation.\n                # The forward hook is on `model.backbone` (training model);\n                # in val we pass `ema.ema` so feat_capture.feat is stale and\n                # the batch size won't match the val batch (bs may be < bs\n                # on the last batch since val_loader has drop_last=False).\n                # val_loss is only logged, never drives decisions.\n                total_loss += sum(losses.values()).item()\n    return total_loss / max(1, len(loader))\n","metadata":{"execution":{"iopub.status.busy":"2026-04-24T22:03:01.477422Z","iopub.execute_input":"2026-04-24T22:03:01.477661Z","iopub.status.idle":"2026-04-24T22:03:01.511672Z","shell.execute_reply.started":"2026-04-24T22:03:01.477634Z","shell.execute_reply":"2026-04-24T22:03:01.511093Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if len(df) > 0:\n    scaler = GradScaler() if \"cuda\" in device else None\n\n    # Attempt to resume from latest checkpoint (if this is a continued session)\n    start_epoch, best_map, es_counter, history = try_resume(\n        CKPT_LATEST, model, ema, optimizer, scheduler, scaler\n    )\n    if start_epoch > 0:\n        # If resuming mid-training, restore mosaic state for this epoch\n        if start_epoch >= mosaic_epochs:\n            train_dataset.disable_mosaic()\n            print(f\"  Mosaic already disabled (past epoch {mosaic_epochs})\")\n\n    # Mosaic/MixUp are OFF globally (prob=0), so no warmup phase is needed.\n    # Copy-paste stays on from epoch 1 to boost rare TB class.\n    mosaic_warmup = 0\n\n    # No freeze — differential LR (backbone 0.1x) lets backbone adapt safely from\n    # epoch 1. Freeze+unfreeze was causing collapse; differential LR avoids it.\n    BACKBONE_FREEZE_EPOCHS = 0\n    _BACKBONE_KEYS = (\"backbone\", \"stem\", \"dark\")  # covers CSPDarknet naming\n    def _freeze_backbone(m, freeze):\n        n_touched = 0\n        for name, p in m.named_parameters():\n            if any(k in name.lower() for k in _BACKBONE_KEYS):\n                p.requires_grad = not freeze; n_touched += 1\n        return n_touched\n\n    def _freeze_backbone_bn(m):\n        \"\"\"Pin backbone BatchNorm layers to eval() so bs=2 does not corrupt\n        the pretrained running_mean/running_var. Called every epoch after\n        model.train(), since train() resets BN to training mode.\n        \"\"\"\n        n = 0\n        for name, mod in m.named_modules():\n            if isinstance(mod, nn.BatchNorm2d) and any(k in name.lower() for k in _BACKBONE_KEYS):\n                mod.eval()\n                n += 1\n        return n\n\n    for epoch in range(start_epoch, epochs):\n        # T3.10: freeze/unfreeze backbone at boundary\n        if epoch == start_epoch and epoch < BACKBONE_FREEZE_EPOCHS:\n            n = _freeze_backbone(model, True)\n            print(f\"  ✂  Froze {n} backbone params for first {BACKBONE_FREEZE_EPOCHS} epochs\")\n        if epoch == BACKBONE_FREEZE_EPOCHS:\n            n = _freeze_backbone(model, False)\n            print(f\"  ▶  Unfroze {n} backbone params\")\n\n        # v4: Warmup phase — disable heavy augmentation\n        if epoch < mosaic_warmup:\n            train_dataset.enabled = False       # disable mosaic\n            base_train_ds.cp_prob = 0.0         # disable copy-paste\n        elif epoch == mosaic_warmup:\n            train_dataset.enabled = True\n            base_train_ds.cp_prob = copy_paste_prob\n            print(f\"\\n*** Epoch {epoch+1}: ENABLING Mosaic/MixUp/CopyPaste (warmup done) ***\\n\")\n            es_counter = 0  # v6 FIX: reset ES counter — mosaic causes a natural dip\n\n        if epoch == mosaic_epochs:\n            print(f\"\\n*** Epoch {epoch+1}: DISABLING Mosaic/MixUp ***\\n\")\n            train_dataset.disable_mosaic()\n            base_train_ds.cp_prob = 0.0\n\n        # Slow ramp: 3 epochs of 0.01x -> 1.0x (was 1 epoch at 0.1x -> 1.0x which\n        # from 0.1 × scheduler-LR back to full over ~1 epoch of optimizer steps.\n        # Uses math.ceil so short epochs still get the full ramp.\n        ramp_steps = 0\n        if epoch == BACKBONE_FREEZE_EPOCHS and BACKBONE_FREEZE_EPOCHS > 0:\n            # 3 full epochs of ramp, not 1\n            ramp_steps = max(1, 3 * math.ceil(len(train_loader) / max(1, grad_accum_steps)))\n            print(f\"  ↗  Post-unfreeze LR ramp: {ramp_steps} optimizer steps \"\n                  f\"(0.1× → 1.0× of scheduler LR)\")\n\n        print(f\"\\nEpoch {epoch+1}/{epochs}\")\n        # Disable multi-scale while we are still in the warmup phase — bs=2 can\n        # not afford the extra variance on top of BN noise.\n        _use_ms = multi_scale and epoch >= mosaic_warmup\n        tl = train_one_epoch(model, ema, train_loader, optimizer, scheduler, loss_func,\n                             device, scaler, epoch,\n                             grad_accum_steps=grad_accum_steps,\n                             multi_scale=_use_ms, scale_range=multi_scale_range,\n                             unfreeze_ramp_steps=ramp_steps,\n                             unfreeze_ramp_min_factor=0.01,\n                             freeze_bn_fn=_freeze_backbone_bn)\n        # T1.3: val_loss uses EMA weights to match the weights used for mAP + best-ckpt\n        _val_model = ema.ema if (use_ema and ema is not None) else model\n        vl = validate_one_epoch(_val_model, valid_loader, loss_func, device)\n\n        # Evaluate mAP using EMA weights (standard YOLOX practice)\n        eval_model = ema if use_ema else model\n        result = compute_map(eval_model, map_loader, device, norm_stats,\n                              conf_thresh=0.05, iou_thresh=iou_thresh)  # T2.7\n        map50 = result[\"map_50\"].item()\n        per_class = {}\n        # T1.1: torchmetrics returns AP in the order of classes PRESENT in the batch,\n        # not CLASS_NAMES order. Use result[\"classes\"] for the correct id→name mapping.\n        if \"map_per_class\" in result:\n            aps = result[\"map_per_class\"]\n            aps = aps.tolist() if hasattr(aps, \"tolist\") else list(aps)\n            ids = result.get(\"classes\", None)\n            if ids is not None:\n                ids = ids.tolist() if hasattr(ids, \"tolist\") else list(ids)\n                for cid, ap in zip(ids, aps):\n                    if 0 <= int(cid) < len(CLASS_NAMES):\n                        per_class[CLASS_NAMES[int(cid)]] = ap\n            else:  # fallback (shouldn't happen with class_metrics=True)\n                for i, ap in enumerate(aps):\n                    if i < len(CLASS_NAMES):\n                        per_class[CLASS_NAMES[i]] = ap\n\n        print(f\"  Train loss: {tl:.4f} | Val loss: {vl:.4f} | mAP@0.5: {map50:.4f}\")\n        for cn, ap in per_class.items():\n            print(f\"    {cn:<14s}: AP={ap:.4f}\")\n\n        history.append({\"epoch\": epoch+1, \"train_loss\": tl, \"val_loss\": vl,\n                        \"map50\": map50, **{f\"ap_{k}\": v for k, v in per_class.items()}})\n\n        improved = map50 > best_map + es_min_delta\n        if improved:\n            best_map = map50\n            es_counter = 0\n            # Save BEST checkpoint (EMA weights go into a separate \"model_ema\" field for\n            # clean loading at inference)\n            torch.save({\n                \"model\": model.state_dict(),\n                \"ema\": ema.ema.state_dict() if ema is not None else None,\n                \"cls_head\": cls_head.state_dict() if cls_head is not None else None,\n                \"epoch\": epoch, \"map50\": best_map, \"history\": history,\n                \"config\": {\"model_type\": model_type, \"train_sz\": train_sz,\n                           \"class_names\": CLASS_NAMES},\n            }, CKPT_BEST)\n            print(\"    ✓ Saved best checkpoint (improved mAP)\")\n        else:\n            es_counter += 1\n            print(f\"    No improvement ({es_counter}/{es_patience})\")\n\n        # Save LATEST checkpoint every `checkpoint_every` epochs for resumption\n        if (epoch + 1) % checkpoint_every == 0 or improved or (epoch + 1 == epochs):\n            save_checkpoint(CKPT_LATEST, model, ema, optimizer, scheduler, scaler,\n                             epoch, best_map, es_counter, history)\n            print(f\"    ✓ Saved resumable checkpoint → {CKPT_LATEST.name}\")\n\n        if es_counter >= es_patience:\n            print(f\"\\nEarly stopping after {epoch+1} epochs\")\n            break\n\n    print(f\"\\nTraining complete. Best mAP@0.5: {best_map:.4f}\")\n\n    # Plot curves\n    hist_df = pd.DataFrame(history)\n    fig, axes = plt.subplots(1, 2, figsize=(14, 4))\n    axes[0].plot(hist_df[\"epoch\"], hist_df[\"train_loss\"], label=\"Train\")\n    axes[0].plot(hist_df[\"epoch\"], hist_df[\"val_loss\"],   label=\"Val\")\n    axes[0].set_xlabel(\"Epoch\"); axes[0].set_ylabel(\"Loss\")\n    axes[0].set_title(\"Loss Curve\"); axes[0].legend(); axes[0].grid(alpha=0.3)\n\n    axes[1].plot(hist_df[\"epoch\"], hist_df[\"map50\"], marker=\"o\", color=\"green\",\n                 label=\"mAP@0.5 (EMA)\", linewidth=2)\n    for cn in CLASS_NAMES:\n        col = f\"ap_{cn}\"\n        if col in hist_df.columns:\n            axes[1].plot(hist_df[\"epoch\"], hist_df[col], alpha=0.6, label=cn)\n    axes[1].set_xlabel(\"Epoch\"); axes[1].set_ylabel(\"mAP\")\n    axes[1].set_title(\"Validation mAP (per class)\"); axes[1].legend(fontsize=8)\n    axes[1].grid(alpha=0.3)\n    plt.tight_layout(); plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 19.5 Visual Prediction Inspection — 30 val images\n\nRun this cell *before* starting a new long training run to see what the current best checkpoint is actually predicting. Green = ground truth, red dashed = top-5 model predictions with confidence. Three failure modes to look for:\n\n- **Localization** — red boxes are nowhere near green. Fix: bbox loss / anchor / scale.\n- **Classification** — red boxes overlap green boxes but the class label is wrong. Fix: class-balanced sampling / label noise audit.\n- **Confidence** — red boxes overlap green with correct class but confidences <0.1. Fix: longer training / `conf_thresh` lowered at inference.\n","metadata":{}},{"cell_type":"code","source":"# ── Visual inspection: 30 val images, GT vs predictions side-by-side ──\nif len(df) > 0:\n    from matplotlib import patches as _mpatch\n\n    # Load best checkpoint's EMA weights if available; else current model state\n    _inspect_model = ema.ema if (use_ema and ema is not None) else model\n    if CKPT_BEST.exists():\n        _ck = torch.load(CKPT_BEST, map_location=device)\n        _sd = _ck.get(\"ema\") or _ck.get(\"model\")\n        if _sd is not None:\n            _inspect_model.load_state_dict(_sd)\n            print(f\"Loaded weights from {CKPT_BEST.name} (epoch {_ck.get('epoch','?')}, \"\n                  f\"mAP={_ck.get('map50', float('nan')):.4f})\")\n    else:\n        print(\"No checkpoint yet — showing predictions from current in-memory model.\")\n    _inspect_model.eval()\n\n    _mean_t = torch.tensor(norm_stats[0]).view(1, 3, 1, 1).to(device)\n    _std_t  = torch.tensor(norm_stats[1]).view(1, 3, 1, 1).to(device)\n    _wrap = YOLOXInferenceWrapper(_inspect_model, _mean_t, _std_t, scale_inp=True)\n\n    # Balanced pick across classes\n    _rng = random.Random(0)\n    _by = defaultdict(list)\n    _val_df = df[df[\"image\"].isin(val_keys)]\n    for _k in val_keys:\n        for _l in _val_df[_val_df[\"image\"] == _k][\"label\"].unique():\n            _by[_l].append(_k)\n    _per_cls = max(1, 30 // max(1, len(CLASS_NAMES)))\n    _picked = []\n    for _ks in _by.values():\n        _picked += _rng.sample(_ks, min(_per_cls, len(_ks)))\n    _picked = list(dict.fromkeys(_picked))[:30]\n    while len(_picked) < 30 and len(val_keys) > len(_picked):\n        _c = _rng.choice(val_keys)\n        if _c not in _picked: _picked.append(_c)\n\n    # Build a map-style transform (no normalize — wrapper handles it)\n    _tfm = A.Compose([\n        A.LongestMaxSize(max_size=train_sz),\n        A.PadIfNeeded(min_height=train_sz, min_width=train_sz,\n                      border_mode=cv2.BORDER_CONSTANT, value=0),\n        ToTensorV2(),\n    ], bbox_params=A.BboxParams(format=\"pascal_voc\",\n                                label_fields=[\"labels\"], min_visibility=0.0))\n\n    _cols = 5\n    _rows = (len(_picked) + _cols - 1) // _cols\n    fig, axes = plt.subplots(_rows, _cols, figsize=(4 * _cols, 4 * _rows))\n    axes = axes.flat if hasattr(axes, \"flat\") else [axes]\n\n    for _ax, _key in zip(axes, _picked):\n        _rows_df = df[df[\"image\"] == _key]\n        _rec = _rows_df.iloc[0]\n        _img = cv2.imread(_rec[\"image_path\"])\n        if _img is None:\n            _ax.set_visible(False); continue\n        _img = preprocess_image(_img, modality=_rec.get(\"modality\", \"xray\"))\n        _gt_b = _rows_df[[\"xmin\", \"ymin\", \"xmax\", \"ymax\"]].values.astype(float).tolist()\n        _gt_l = _rows_df[\"label\"].tolist()\n        # clamp GT to image size to survive Albumentations\n        _h, _w = _img.shape[:2]\n        _gt_b = [[max(0,min(_w,x1)), max(0,min(_h,y1)),\n                  max(0,min(_w,x2)), max(0,min(_h,y2))] for x1,y1,x2,y2 in _gt_b]\n        _out = _tfm(image=_img, bboxes=_gt_b, labels=_gt_l)\n        _img_t = _out[\"image\"].float().unsqueeze(0).to(device)\n\n        with torch.no_grad():\n            _p = _wrap(_img_t)[0].cpu()\n\n        _vis = _out[\"image\"].permute(1, 2, 0).cpu().numpy().astype(\"uint8\")\n        _ax.imshow(_vis)\n\n        # GT in green\n        for (_x1, _y1, _x2, _y2), _lbl in zip(_out[\"bboxes\"], _out[\"labels\"]):\n            _ax.add_patch(_mpatch.Rectangle((_x1, _y1), _x2-_x1, _y2-_y1,\n                          linewidth=2, edgecolor=\"lime\", facecolor=\"none\"))\n            _ax.text(_x1, max(0, _y1-3), f\"GT:{_lbl}\", color=\"lime\",\n                     fontsize=7, weight=\"bold\")\n\n        # Predictions in red dashed — top 5 by confidence\n        if len(_p) > 0:\n            _bx = torchvision.ops.box_convert(_p[:, :4], \"xywh\", \"xyxy\")\n            _sc = _p[:, -1]\n            _ci = _p[:, 4].long()\n            _ord = _sc.argsort(descending=True)[:5]\n            for _i in _ord.tolist():\n                _x1, _y1, _x2, _y2 = _bx[_i].tolist()\n                _cid = int(_ci[_i]); _cf = float(_sc[_i])\n                _cn = CLASS_NAMES[_cid] if 0 <= _cid < len(CLASS_NAMES) else str(_cid)\n                _ax.add_patch(_mpatch.Rectangle((_x1, _y1), _x2-_x1, _y2-_y1,\n                              linewidth=1.3, edgecolor=\"red\",\n                              facecolor=\"none\", linestyle=\"--\"))\n                _ax.text(_x1, min(_vis.shape[0]-1, _y2+10),\n                         f\"P:{_cn} {_cf:.2f}\", color=\"red\",\n                         fontsize=7, weight=\"bold\")\n\n        _ax.set_title(f\"{_key[:24]} [{_rec.get('modality','?')}]\", fontsize=8)\n        _ax.axis(\"off\")\n\n    # hide any leftover axes\n    for _ax in list(axes)[len(_picked):]:\n        _ax.set_visible(False)\n\n    plt.tight_layout()\n    _save_path = ckpt_dir / \"val_predictions_preview.png\"\n    plt.savefig(str(_save_path), dpi=90, bbox_inches=\"tight\")\n    plt.show()\n    print(f\"\\u2713 Saved preview \\u2192 {_save_path}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T18:59:24.44676Z","iopub.status.idle":"2026-04-24T18:59:24.44713Z","shell.execute_reply.started":"2026-04-24T18:59:24.446916Z","shell.execute_reply":"2026-04-24T18:59:24.446938Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 20. Inference with Test-Time Augmentation\n\nTTA strategy:\n1. Predict on original image\n2. Predict on horizontal flip, flip boxes back\n3. Merge both prediction sets with **weighted box fusion** (superior to NMS for overlapping lesions)\n","metadata":{}},{"cell_type":"code","source":"if len(df) > 0 and CKPT_BEST.exists():\n    # Load from new v3 checkpoint format — prefer EMA weights, fall back to main model\n    ck = torch.load(CKPT_BEST, map_location=device)\n\n    # ── v11: load auxiliary classification head from checkpoint ──\n    \n    # ── v11: image-level classification head probabilities ──\n    # After running inference on an image (variable `img` or `inputs_t`),\n    # call this helper to print/render image-level disease probabilities.\n    # Useful for the radiologist UI: 'Pneumonia: 87% confidence' overlay.\n    def get_image_level_probs(img_tensor_640):\n        \"\"\"Returns dict mapping class_name -> probability [0,1].\n        Pass a tensor of shape (1, 3, 640, 640) — float32, [0, 255] range\n        (same contract as map_loader). Uses the wrapper + cls_head hook.\"\"\"\n        if cls_head is None:\n            return {cn: None for cn in CLASS_NAMES}\n        with torch.no_grad():\n            _ = wrapper(img_tensor_640.to(device))  # triggers feat_capture hook\n            if feat_capture is not None and feat_capture.feat is not None:\n                logits = cls_head(feat_capture.feat)\n                probs = torch.sigmoid(logits)[0].cpu().tolist()\n                return {cn: float(p) for cn, p in zip(CLASS_NAMES, probs)}\n        return {cn: None for cn in CLASS_NAMES}\n\n    if cls_head is not None and ck.get(\"cls_head\") is not None:\n        cls_head.load_state_dict(ck[\"cls_head\"])\n        cls_head.eval()\n        print(\"\\u2713 Loaded auxiliary cls_head from checkpoint\")\n    if ck.get(\"ema\") is not None:\n        print(\"Loading EMA weights for inference (preferred)\")\n        model.load_state_dict(ck[\"ema\"])\n    else:\n        print(\"Loading main model weights\")\n        model.load_state_dict(ck[\"model\"])\n    model.eval()\n    best_map = ck.get(\"map50\", best_map if \"best_map\" in dir() else -1.0)\n\n    mean_t = torch.tensor(norm_stats[0]).view(1, 3, 1, 1).to(device)\n    std_t  = torch.tensor(norm_stats[1]).view(1, 3, 1, 1).to(device)\n    wrapped_model = YOLOXInferenceWrapper(model, mean_t, std_t, scale_inp=True)\n\n    def _preprocess_for_inference(img_bgr, size=train_sz):\n        \"\"\"Mirror training preprocessing for inference.\"\"\"\n        rgb = preprocess_image(img_bgr)\n        h, w = rgb.shape[:2]\n        scale = size / max(h, w)\n        new_w, new_h = int(w * scale), int(h * scale)\n        resized = cv2.resize(rgb, (new_w, new_h))\n        # Pad to square\n        pad_h = size - new_h; pad_w = size - new_w\n        padded = cv2.copyMakeBorder(resized, 0, pad_h, 0, pad_w,\n                                     cv2.BORDER_CONSTANT, value=0)\n        return padded, scale, (pad_w, pad_h)\n\n    def _run_inference(img_tensor):\n        with torch.no_grad():\n            inp = img_tensor.unsqueeze(0).to(device)\n            if inp.dtype == torch.uint8:\n                inp = inp.float() / 255.0\n            out = wrapped_model(inp).cpu()\n        return out\n\n    def predict_single(img_bgr, conf=0.3):\n        \"\"\"Single-scale prediction (no TTA).\"\"\"\n        padded, scale, _ = _preprocess_for_inference(img_bgr)\n        t = torch.from_numpy(padded).permute(2, 0, 1).float() / 255.0\n        out = _run_inference(t)\n        mask = out[0, :, -1] > conf\n        props = out[0, mask]\n        if len(props) == 0:\n            return np.zeros((0, 4)), np.zeros(0), np.zeros(0, dtype=int)\n        boxes = torchvision.ops.box_convert(props[:, :4], \"xywh\", \"xyxy\").numpy()\n        boxes /= scale\n        scores = props[:, 5].numpy()\n        labels = props[:, 4].long().numpy()\n        return boxes, scores, labels\n\n    def predict_tta(img_bgr, conf=0.3):\n        \"\"\"TTA: original + horizontal flip, merged with weighted box fusion.\"\"\"\n        h, w = img_bgr.shape[:2]\n\n        # 1. Original\n        b1, s1, l1 = predict_single(img_bgr, conf=conf)\n\n        # 2. Flipped\n        flipped = cv2.flip(img_bgr, 1)\n        b2, s2, l2 = predict_single(flipped, conf=conf)\n        if len(b2) > 0:\n            b2[:, [0, 2]] = w - b2[:, [2, 0]]\n\n        if len(b1) == 0 and len(b2) == 0:\n            return np.zeros((0, 4)), np.zeros(0), np.zeros(0, dtype=int)\n\n        all_boxes_norm, all_scores, all_labels = [], [], []\n        for bs_set, ss_set, ls_set in [(b1, s1, l1), (b2, s2, l2)]:\n            if len(bs_set) == 0:\n                all_boxes_norm.append([]); all_scores.append([]); all_labels.append([])\n                continue\n            norm = bs_set.copy().astype(np.float32)\n            norm[:, [0, 2]] /= w; norm[:, [1, 3]] /= h\n            norm = np.clip(norm, 0, 1)\n            all_boxes_norm.append(norm.tolist())\n            all_scores.append(ss_set.tolist())\n            all_labels.append(ls_set.tolist())\n\n        fused_boxes, fused_scores, fused_labels = weighted_boxes_fusion(\n            all_boxes_norm, all_scores, all_labels,\n            iou_thr=0.5, skip_box_thr=0.001,\n        )\n        fused_boxes = np.array(fused_boxes)\n        if len(fused_boxes) > 0:\n            fused_boxes[:, [0, 2]] *= w; fused_boxes[:, [1, 3]] *= h\n        return fused_boxes, np.array(fused_scores), np.array(fused_labels, dtype=int)\n\n    def predict_and_draw(img_path, conf=0.3, use_tta_flag=True):\n        img = cv2.imread(str(img_path))\n        if img is None:\n            return None\n        predictor = predict_tta if use_tta_flag else predict_single\n        boxes, scores, labels = predictor(img, conf=conf)\n\n        img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB).copy()\n        cols = distinctipy.get_colors(NUM_CLASSES, rng=0)\n        for (x1, y1, x2, y2), sc, lb in zip(boxes, scores, labels):\n            col = tuple(int(c*255) for c in cols[int(lb)])\n            cv2.rectangle(img_rgb, (int(x1), int(y1)), (int(x2), int(y2)), col, 2)\n            tag = f\"{CLASS_NAMES[int(lb)]} {sc:.2f}\"\n            cv2.putText(img_rgb, tag, (int(x1), int(y1) - 6),\n                        cv2.FONT_HERSHEY_SIMPLEX, 0.5, col, 2)\n        return img_rgb\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T18:59:24.448014Z","iopub.status.idle":"2026-04-24T18:59:24.448307Z","shell.execute_reply.started":"2026-04-24T18:59:24.448136Z","shell.execute_reply":"2026-04-24T18:59:24.448154Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 21. Test Set Predictions","metadata":{}},{"cell_type":"code","source":"if len(df) > 0 and CKPT_BEST.exists():\n    print(f\"Best Val mAP@0.5: {best_map:.4f}\")\n    nt = min(6, len(test_keys))\n    if nt > 0:\n        fig, axes = plt.subplots(nt, 2, figsize=(12, 5*nt))\n        if nt == 1:\n            axes = [axes]\n        for i in range(nt):\n            tn = test_keys[i]\n            tp = df[df[\"image\"] == tn][\"image_path\"].iloc[0]\n            img_orig = cv2.cvtColor(cv2.imread(tp), cv2.COLOR_BGR2RGB)\n            axes[i][0].imshow(img_orig)\n            axes[i][0].set_title(f\"Original: {tn}\"); axes[i][0].axis(\"off\")\n\n            img_pred = predict_and_draw(tp, conf=0.3, use_tta_flag=use_tta)\n            if img_pred is not None:\n                axes[i][1].imshow(img_pred)\n                axes[i][1].set_title(\"Prediction (TTA)\" if use_tta else \"Prediction\")\n            axes[i][1].axis(\"off\")\n        plt.tight_layout(); plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T18:59:24.449641Z","iopub.status.idle":"2026-04-24T18:59:24.454145Z","shell.execute_reply.started":"2026-04-24T18:59:24.453914Z","shell.execute_reply":"2026-04-24T18:59:24.453942Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 22. Final Test Set Evaluation\n\nCompute per-class AP on the held-out test set using TTA. This is the number to report.\n","metadata":{}},{"cell_type":"code","source":"if len(df) > 0 and CKPT_BEST.exists():\n    test_dataset = LungDiseaseDataset(\n        test_keys, df, CLASS_TO_IDX, transform=build_map_transform(train_sz),\n        preprocess_fn=preprocess_image,\n    )\n    test_loader = DataLoader(test_dataset, batch_size=bs, num_workers=2,\n                              collate_fn=lambda b: tuple(zip(*b)), drop_last=False)\n\n    # `model` was already loaded with EMA weights in the inference cell above\n    result = compute_map(model, test_loader, device, norm_stats,\n                          conf_thresh=conf_thresh, iou_thresh=iou_thresh)\n    print(\"=\" * 50)\n    print(\"TEST SET RESULTS (EMA weights, no TTA)\")\n    print(\"=\" * 50)\n    print(f\"mAP@0.5:           {result['map_50'].item():.4f}\")\n    if \"map\" in result:\n        print(f\"mAP@0.5:0.95:      {result['map'].item():.4f}\")\n    if \"map_small\" in result:\n        print(f\"mAP (small):       {result['map_small'].item():.4f}\")\n        print(f\"mAP (medium):      {result['map_medium'].item():.4f}\")\n        print(f\"mAP (large):       {result['map_large'].item():.4f}\")\n    print()\n    print(\"Per-class AP:\")\n    if \"map_per_class\" in result:\n        for i, ap in enumerate(result[\"map_per_class\"].tolist()):\n            if i < len(CLASS_NAMES):\n                print(f\"  {CLASS_NAMES[i]:<14s}: AP={ap:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T18:59:24.456386Z","iopub.status.idle":"2026-04-24T18:59:24.45678Z","shell.execute_reply.started":"2026-04-24T18:59:24.45656Z","shell.execute_reply":"2026-04-24T18:59:24.456583Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Summary\n\n**v8 changes on top of the original v4 baseline (still YOLOX):**\n\n| Thing                      | v4 baseline                | v8                                                  |\n|----------------------------|----------------------------|-----------------------------------------------------|\n| Classes                    | 4 (incl. `tumor_ct`)       | **3 (x-ray only)**                                  |\n| Model                      | YOLOX-m / YOLOX-s @ 1024   | **YOLOX-s @ 640**                                   |\n| Resolution                 | 1024 × 1024                | **640 × 640**                                       |\n| Real batch size            | 2                          | **16**                                              |\n| Effective batch size       | 8 (grad accum 4)           | **16 (no accumulation)**                            |\n| LR strategy                | single LR, freeze/unfreeze | **differential LR (bb 0.1× head); no freeze**       |\n| Mosaic / MixUp             | 0.5 / 0.1                  | **0.0 / 0.0**                                       |\n| Multi-scale                | 832–1088                   | **off**                                             |\n| Sampler                    | sqrt inverse-freq          | **pure inverse-freq + rarest-label picker (v3)**    |\n| `pneumonia` sources        | RSNA + VinDr + VinBigData + ChestX-Det | **RSNA + VinDr (explicit pneumonia only)** |\n| `tumor_xray` sources       | mixed                      | **Nodule/Mass only (v7)**                           |\n| LIDC-IDRI parse            | runs every start           | **skipped (v8)**                                    |\n| Resume memory              | ck dict lingered → OOM     | **freed pre-fork (v4)**                             |\n| DataLoader workers         | 4                          | **2 + `persistent_workers` + `prefetch_factor=2`**  |\n| mAP eval model             | main weights               | EMA weights                                         |\n\n**Realistic per-class AP targets (graduation-project honest):**\n- `tumor_xray`:    0.20–0.30\n- `tuberculosis`:  0.25–0.40 (small class, boosted by inverse-freq sampling + copy-paste)\n- `pneumonia`:     0.20–0.30 (RSNA-clean after v7 label tightening)\n- **Overall mAP:** ~0.20–0.30 — in line with published VinBigData / RSNA benchmarks\n\n**Pair with image-level AUC** (~0.85–0.90 achievable) for a stronger report narrative. Detection mAP and classification AUC measure different things; reporting both shows you understand the task.\n\n**If you want to push further post-submission:**\n1. Swap to a **RadImageNet-pretrained backbone** (biggest remaining win; requires surgery on the YOLOX stem)\n2. Train **3 seeds** and ensemble with WBF (+3 mAP)\n3. **Pseudo-label** unlabeled CheXpert / MIMIC-CXR for pneumonia (a lot of work for modest gain)\n","metadata":{}},{"cell_type":"code","source":"# ═════════════════════════════════════════════════════════════════════════\n#  Per-class confidence threshold tuning  (run AFTER training finishes)\n# ═════════════════════════════════════════════════════════════════════════\n#\n# Sweeps confidence thresholds per class on the val set, computes\n# (precision, recall, FP/image) at each. Use the result to set\n# PER_CLASS_DEPLOY_THRESH for the inference / serving path.\n#\n# Pick the threshold per class that gives recall ≥ 0.85 with FP/image ≤ 2.\n# Bake into your inference cell as e.g.:\n#     PER_CLASS_DEPLOY_THRESH = {\"tumor_xray\": 0.18, \"tuberculosis\": 0.30, \"pneumonia\": 0.12}\n\nif CKPT_BEST.exists():\n    from collections import defaultdict\n    import numpy as np\n\n    print(\"Loading best checkpoint for threshold tuning...\")\n    _ck = torch.load(CKPT_BEST, map_location=device)\n    if \"ema\" in _ck and _ck[\"ema\"] is not None:\n        ema.ema.load_state_dict(_ck[\"ema\"])\n    eval_model = ema.ema if (use_ema and ema is not None) else model\n    eval_model.eval()\n\n    _mean_t = torch.tensor(norm_stats[0]).view(1, 3, 1, 1).to(device)\n    _std_t  = torch.tensor(norm_stats[1]).view(1, 3, 1, 1).to(device)\n    _wrapper = YOLOXInferenceWrapper(eval_model, _mean_t, _std_t, scale_inp=True)\n\n    # Collect per-class (score, is_TP) pairs across the val set\n    all_preds = defaultdict(list)\n    all_gt    = defaultdict(int)\n    n_images  = 0\n\n    with torch.no_grad():\n        for inputs, targets in tqdm(map_loader, desc=\"PR scan\", leave=False):\n            x = torch.stack(inputs).to(device).float()\n            with autocast(device_type=torch.device(device).type):\n                out = _wrapper(x)\n            for b, tgt in enumerate(targets):\n                n_images += 1\n                preds = out[b].cpu()\n                preds = preds[preds[:, -1] > 0.005]   # very permissive prefilter\n                for cid in range(NUM_CLASSES):\n                    cls_preds = preds[preds[:, 4].long() == cid]\n                    cls_gts   = tgt[\"boxes\"][tgt[\"labels\"] == cid]\n                    all_gt[cid] += len(cls_gts)\n\n                    if len(cls_preds) == 0:\n                        continue\n                    if len(cls_gts) == 0:\n                        for s in cls_preds[:, 5].tolist():\n                            all_preds[cid].append((s, 0))\n                        continue\n\n                    # Match preds to GT by IoU >= 0.5, descending score\n                    pb = torchvision.ops.box_convert(cls_preds[:, :4], \"xywh\", \"xyxy\")\n                    iou = torchvision.ops.box_iou(pb, cls_gts.float())\n                    matched_gt = set()\n                    order = cls_preds[:, 5].argsort(descending=True)\n                    for i in order.tolist():\n                        s = cls_preds[i, 5].item()\n                        best_gt = int(iou[i].argmax().item())\n                        if iou[i, best_gt].item() >= 0.5 and best_gt not in matched_gt:\n                            all_preds[cid].append((s, 1))\n                            matched_gt.add(best_gt)\n                        else:\n                            all_preds[cid].append((s, 0))\n\n    # Per-class PR sweep\n    print()\n    print(f\"PR scan over {n_images} val images\")\n    print(f\"{'Class':<14} {'thresh':>7} {'recall':>8} {'precision':>11} {'FP/img':>8}\")\n    print(\"-\" * 55)\n    suggested = {}\n    for cid, cn in enumerate(CLASS_NAMES):\n        preds = sorted(all_preds[cid], key=lambda x: -x[0])\n        gt = all_gt[cid]\n        rows = []\n        for thr in [0.05, 0.08, 0.10, 0.12, 0.15, 0.18, 0.20, 0.25, 0.30, 0.40, 0.50]:\n            kept = [p for p in preds if p[0] >= thr]\n            tp = sum(1 for s, m in kept if m == 1)\n            fp = sum(1 for s, m in kept if m == 0)\n            recall    = tp / max(1, gt)\n            precision = tp / max(1, tp + fp)\n            fppi      = fp / max(1, n_images)\n            rows.append((thr, recall, precision, fppi))\n            print(f\"{cn:<14} {thr:>7.2f} {recall:>8.2%} {precision:>11.2%} {fppi:>8.2f}\")\n        # Pick threshold: highest recall with FP/img <= 2.0\n        viable = [r for r in rows if r[3] <= 2.0]\n        if viable:\n            best = max(viable, key=lambda r: r[1])  # max recall among viable\n            suggested[cn] = best[0]\n        else:\n            suggested[cn] = 0.25\n        print()\n\n    print(\"\\u2728 Suggested PER_CLASS_DEPLOY_THRESH (recall-maximizing at FP/img ≤ 2):\")\n    for cn, thr in suggested.items():\n        print(f'    \"{cn}\": {thr:.2f},')\n    print()\n    print(\"Paste into your inference cell:\")\n    print(f\"PER_CLASS_DEPLOY_THRESH = {suggested}\")\nelse:\n    print(\"No best.pth — train first, then run this cell.\")\n","metadata":{},"outputs":[],"execution_count":null}]}