{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nKaggle notebook 1 - build the cached study tensors.\n\n600 GB of DICOM cannot be re-read every training run, so this pass writes one compact\nuint8 tensor per study and everything downstream reads that.\n\n  per study:  (n_slots, N_SLICES, RES, RES) uint8   + a slot presence mask\n  at RES=224, N_SLICES=16, 5 populated slots  ->  ~4 MB/study, ~17 GB for 4407 studies\n\nAll three preprocessing corrections from src/preprocess.py are applied here, so they are\nbaked into the cache and cost nothing at train time:\n\n  * laterality from FOV-centre x sign  (the Laterality tag is missing on ~41% of series)\n  * canonical medial->lateral slice order on sagittal, mirrored coronal/axial for right\n    knees, so the four side-specific targets see consistent geometry\n  * sequence type from TR/TE, giving 12 slots instead of 6 and keeping the T1/T2 pairs\n    that a plane x FS scheme has to discard\n\nRun with SHARD_ID / N_SHARDS to spread the work over several 12h Kaggle sessions.\n\nA100/H100 note: with a bigger GPU raise RES to 320 and N_SLICES to 24 (~55 GB cache) and\ndrop the uint8 quantisation in favour of float16 - on Kaggle neither fits the time budget.\n\"\"\"\nfrom __future__ import annotations\nimport os, sys, glob, json, math, warnings, threading\nfrom concurrent.futures import ThreadPoolExecutor\nimport numpy as np, pandas as pd, pydicom, cv2\nwarnings.filterwarnings('ignore')\n\n# ----------------------------------------------------------------- config\nCOMP        = os.environ.get('RSNA_COMP', '/kaggle/input/competitions/rsna-knee-abnormality-detection')\nOUT         = os.environ.get('RSNA_CACHE', '/kaggle/working/cache')\nSPLIT       = os.environ.get('RSNA_SPLIT', 'train')          # 'train' | 'test'\nRES         = int(os.environ.get('RSNA_RES', 224))\nN_SLICES    = int(os.environ.get('RSNA_NSLICES', 16))\nFOV_MM      = float(os.environ.get('RSNA_FOV', 140))\nBAND_LO     = float(os.environ.get('RSNA_BAND_LO', 0.06))\nBAND_HI     = float(os.environ.get('RSNA_BAND_HI', 0.94))\nSHARD_ID    = int(os.environ.get('RSNA_SHARD', 0))\nN_SHARDS    = int(os.environ.get('RSNA_NSHARDS', 1))\n# This pass is I/O bound, not GPU bound - run it in a CPU-only Kaggle session so it does\n# not spend the 30 h/week GPU quota. /kaggle/input is network-backed and far slower than\n# local disk, so threads matter here even though the local timing looks instant.\nWORKERS     = int(os.environ.get('RSNA_WORKERS', 8))\n\nPLANES = ('Sagittal', 'Coronal', 'Axial')\nSLOTS  = tuple(f'{p[:3]}_{s}' for p in PLANES for s in ('FS', 'T1', 'T2', 'PD'))\nSLOT_IX = {s: i for i, s in enumerate(SLOTS)}\n\n# ----------------------------------------------------------------- geometry\ndef series_geometry(ds):\n    iop = np.asarray(ds.ImageOrientationPatient, float)\n    ipp = np.asarray(ds.ImagePositionPatient, float)\n    ps  = np.asarray(ds.PixelSpacing, float)\n    row_dir, col_dir = iop[:3], iop[3:]\n    centre = ipp + col_dir * ps[1] * (ds.Columns / 2.0) + row_dir * ps[0] * (ds.Rows / 2.0)\n    return centre, np.cross(row_dir, col_dir), row_dir, col_dir\n\ndef study_laterality(headers):\n    tags = {str(h.get('Laterality', '') or h.get('ImageLaterality', '') or '') for h in headers}\n    tags.discard('')\n    if len(tags) == 1:\n        return tags.pop()\n    xs = []\n    for h in headers:\n        try: xs.append(series_geometry(h)[0][0])\n        except Exception: pass\n    return 'L' if (xs and float(np.median(xs)) > 0) else ('R' if xs else 'L')\n\ndef sequence_type(ds):\n    tr, te, ti = ds.get('RepetitionTime'), ds.get('EchoTime'), ds.get('InversionTime')\n    scanseq = str(ds.get('ScanningSequence', '') or '')\n    desc = str(ds.get('SeriesDescription', '') or '').lower()\n    try:\n        if ti not in (None, '') and float(ti) > 20: return 'STIR'\n    except Exception: pass\n    if 'GR' in scanseq and 'SE' not in scanseq: return 'GRE'\n    if tr in (None, '') or te in (None, ''):\n        for k, v in (('t1','T1'),('stir','STIR'),('t2','T2'),('pd','PD'),('dp','PD'),('int','PD')):\n            if k in desc: return v\n        return 'PD'\n    tr, te = float(tr), float(te)\n    if tr < 900 and te < 20: return 'T1'\n    if te >= 45: return 'T2'\n    return 'PD'\n\ndef slot_of(plane, fat_sup, seq):\n    if fat_sup: return f'{plane[:3]}_FS'\n    return f'{plane[:3]}_{\"T1\" if seq==\"T1\" else \"T2\" if seq==\"T2\" else \"PD\"}'\n\ndef order_slices(dss, laterality, plane):\n    try:\n        _, normal, _, _ = series_geometry(dss[0])\n        pos = np.array([float(np.dot(np.asarray(d.ImagePositionPatient, float), normal)) for d in dss])\n    except Exception:\n        pos = np.array([float(getattr(d, 'InstanceNumber', i)) for i, d in enumerate(dss)])\n        normal = np.array([1.0, 0.0, 0.0])\n    if plane == 'Sagittal':\n        key = pos * (1.0 if normal[0] >= 0 else -1.0)\n        if laterality != 'L': key = -key          # medial first for both knees\n    elif plane == 'Coronal':\n        key = pos * (1.0 if normal[1] >= 0 else -1.0)\n    else:\n        key = -pos * (1.0 if normal[2] >= 0 else -1.0)\n    return [dss[i] for i in np.argsort(key)]\n\ndef canonical_lr(img, plane, laterality):\n    return img[:, ::-1] if (plane in ('Coronal', 'Axial') and laterality == 'R') else img\n\n# ----------------------------------------------------------------- pixels\ndef to_float(ds):\n    a = ds.pixel_array.astype(np.float32)\n    return a * float(ds.get('RescaleSlope', 1) or 1) + float(ds.get('RescaleIntercept', 0) or 0)\n\ndef physical_crop(img, ps, fov_mm, out):\n    want = int(round(fov_mm / float(ps)))\n    h, w = img.shape\n    half = want // 2\n    cy, cx = h // 2, w // 2\n    y0, y1 = max(0, cy - half), min(h, cy + half)\n    x0, x1 = max(0, cx - half), min(w, cx + half)\n    crop = img[y0:y1, x0:x1]\n    if crop.size == 0: crop = img\n    s = min(crop.shape)\n    crop = crop[:s, :s]\n    return cv2.resize(crop, (out, out), interpolation=cv2.INTER_AREA)\n\ndef normalise_u8(vol):\n    lo, hi = np.percentile(vol, 1), np.percentile(vol, 99)\n    if hi <= lo: hi = lo + 1.0\n    return (np.clip((vol - lo) / (hi - lo), 0, 1) * 255).astype(np.uint8)\n\ndef band(n_have, n_want, lo=BAND_LO, hi=BAND_HI):\n    if n_have <= 0: return []\n    idx = np.linspace(lo * (n_have - 1), hi * (n_have - 1), n_want)\n    return np.clip(np.round(idx).astype(int), 0, n_have - 1).tolist()\n\n# ----------------------------------------------------------------- per study\ndef build_study(study_uid, series_rows, root):\n    \"\"\"-> (n_slots, N_SLICES, RES, RES) uint8, presence mask, meta dict\"\"\"\n    vol  = np.zeros((len(SLOTS), N_SLICES, RES, RES), np.uint8)\n    have = np.zeros(len(SLOTS), np.uint8)\n\n    heads, per_series = [], []\n    for r in series_rows:\n        sd = os.path.join(root, study_uid, r['SeriesInstanceUID'])\n        if not os.path.isdir(sd):\n            sd = os.path.join(root, r['SeriesInstanceUID'])          # flat layout fallback\n        files = glob.glob(os.path.join(sd, '*.dcm'))\n        if not files: continue\n        try: h = pydicom.dcmread(files[0], stop_before_pixels=True)\n        except Exception: continue\n        heads.append(h); per_series.append((r, sd, files, h))\n    if not per_series:\n        return vol, have, {}\n\n    lat = study_laterality(heads)\n\n    # one series per slot; on a collision keep the one with the most slices\n    chosen = {}\n    for r, sd, files, h in per_series:\n        s = slot_of(r['Anatomical_Plane'], int(r['Fat_Suppression']), sequence_type(h))\n        if s not in SLOT_IX: continue\n        if s not in chosen or len(files) > len(chosen[s][2]):\n            chosen[s] = (r, sd, files, h)\n\n    meta = {'laterality': lat, 'slots': sorted(chosen), 'n_series': len(per_series)}\n    for s, (r, sd, files, h) in chosen.items():\n        try:\n            dss = [pydicom.dcmread(f) for f in files]\n            dss = order_slices(dss, lat, r['Anatomical_Plane'])\n            idx = band(len(dss), N_SLICES)\n            ps = float(np.asarray(dss[0].PixelSpacing, float)[0])\n            stack = []\n            for i in idx:\n                a = to_float(dss[i])\n                a = canonical_lr(a, r['Anatomical_Plane'], lat)\n                stack.append(physical_crop(a, ps, FOV_MM, RES))\n            vol[SLOT_IX[s]] = normalise_u8(np.stack(stack))\n            have[SLOT_IX[s]] = 1\n        except Exception as e:\n            print(f'  ! {study_uid[-8:]} {s}: {type(e).__name__} {e}', flush=True)\n    return vol, have, meta\n\n# ----------------------------------------------------------------- driver\ndef main():\n    os.makedirs(OUT, exist_ok=True)\n    ser = pd.read_csv(f'{COMP}/{SPLIT}_series.csv', dtype=str)\n    ser['Fat_Suppression'] = ser['Fat_Suppression'].astype(int)\n    root = f'{COMP}/{SPLIT}_series'\n    studies = sorted(ser.StudyInstanceUID.unique())\n    mine = [s for i, s in enumerate(studies) if i % N_SHARDS == SHARD_ID]\n    print(f'{SPLIT}: {len(studies)} studies, shard {SHARD_ID}/{N_SHARDS} -> {len(mine)}', flush=True)\n\n    grouped = {k: v.to_dict('records') for k, v in ser.groupby('StudyInstanceUID')}\n    todo = [u for u in mine if not os.path.exists(os.path.join(OUT, f'{u}.npz'))]\n    print(f'{len(mine) - len(todo)} already cached, {len(todo)} to build, {WORKERS} workers', flush=True)\n\n    metas, lock, done = {}, threading.Lock(), [0]\n\n    def one(uid):\n        dst = os.path.join(OUT, f'{uid}.npz')\n        try:\n            vol, have, meta = build_study(uid, grouped[uid], root)\n            if have.sum() == 0:\n                print(f'  ! {uid[-8:]} no usable series, skipped', flush=True)\n                return\n            # the suffix must stay .npz: np.savez_compressed silently appends .npz to any\n            # path that lacks it, which would leave os.replace looking for a missing file\n            tmp = f'{dst[:-4]}.{os.getpid()}.{threading.get_ident()}.tmp.npz'\n            np.savez_compressed(tmp, vol=vol, have=have)\n            os.replace(tmp, dst)          # atomic, so an interrupted session resumes cleanly\n        except Exception as e:\n            print(f'  ! {uid[-8:]} FAILED {type(e).__name__}: {e}', flush=True)\n            return\n        with lock:\n            metas[uid] = meta\n            done[0] += 1\n            if done[0] % 100 == 0:\n                print(f'  [{done[0]}/{len(todo)}] {uid[-8:]} slots={len(meta.get(\"slots\", []))} '\n                      f'lat={meta.get(\"laterality\")}', flush=True)\n\n    with ThreadPoolExecutor(max_workers=WORKERS) as pool:\n        list(pool.map(one, todo))\n\n    with open(os.path.join(OUT, f'meta_{SPLIT}_{SHARD_ID}.json'), 'w') as f:\n        json.dump(metas, f)\n    keep = [f for f in os.listdir(OUT) if f.endswith('.npz') and '.tmp.' not in f]\n    for f in os.listdir(OUT):\n        if '.tmp.' in f:\n            os.remove(os.path.join(OUT, f))          # orphans from an interrupted worker\n    n_out = len(keep)\n    size = sum(os.path.getsize(os.path.join(OUT, f)) for f in keep)\n    print(f'done: {n_out} studies cached, {size/1e9:.1f} GB', flush=True)\n\nif __name__ == '__main__':\n    main()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-08-30T16:20:42.782886Z","iopub.execute_input":"2026-08-30T16:20:42.783303Z","iopub.status.idle":"2026-08-30T17:12:28.325744Z","shell.execute_reply.started":"2026-08-30T16:20:42.783264Z","shell.execute_reply":"2026-08-30T17:12:28.324601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}