diff --git a/.gitignore b/.gitignore index 26e8bc1..3d000e0 100644 --- a/.gitignore +++ b/.gitignore @@ -1,4 +1,4 @@ -data/ +data runs/ results/ logs/ diff --git a/.vscode/settings.json b/.vscode/settings.json new file mode 100644 index 0000000..4b5a294 --- /dev/null +++ b/.vscode/settings.json @@ -0,0 +1,4 @@ +{ + "python-envs.defaultEnvManager": "ms-python.python:conda", + "python-envs.defaultPackageManager": "ms-python.python:conda" +} \ No newline at end of file diff --git a/AGENTS.md b/AGENTS.md index 565bd7e..3d0e3b0 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -16,6 +16,8 @@ Key packages: torch 2.14 (+cu126), torchvision, monai 1.6 (pip, Pipeline C), num Longitudinal (repeated-measures) analysis of medical imaging data. Repo is at an early stage — see README.md for any project notes. +**All path config lives in `config/paths.json`** — the external data roots the code reads or writes (`data`, `lee`, `m6`, `ntuh_register_inv`). Code resolves them through `src.common` (`PATHS` dict, `path(name)` lookup, `DATA`/`d()` for `data/...` paths). To move any dataset, only edit `config/paths.json`; the repo root is derived from the code location (override with `LONGITUDINAL_ROOT` if needed). + ## Conventions - All Python commands must run inside the `longitudinal` conda environment. diff --git a/config/paths.json b/config/paths.json new file mode 100644 index 0000000..b2f2c57 --- /dev/null +++ b/config/paths.json @@ -0,0 +1,6 @@ +{ + "data": "/mnt/t24/Public/xfr/longitudinal", + "lee": "/mnt/t24/Public/lee", + "m6": "/mnt/pve/WORKSPACE/M6-2025/nii", + "ntuh_register_inv": "/mnt/pve/SRS/NTUH2022G4/register_inv" +} \ No newline at end of file diff --git a/scripts/01_build_ntuh_manifest.py b/scripts/01_build_ntuh_manifest.py index ed2852b..162f65c 100644 --- a/scripts/01_build_ntuh_manifest.py +++ b/scripts/01_build_ntuh_manifest.py @@ -10,7 +10,7 @@ import sys sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import re from collections import defaultdict -from src.common import ROOT, save_jsonl +from src.common import ROOT, DATA, save_jsonl, path T1C_RE = re.compile(r"T1.*\+C|\+C.*T1") EXCL_RE = re.compile(r"MRA|FLAIR|TOF|T2|SWI|DWI|_ROI|ROI1", re.I) @@ -18,8 +18,8 @@ SPINE_RE = re.compile(r"[TLC]\s?\d+\s?[-–]\s?[TLC]\s?\d+|[TLC]\d{2}", re.I) def main(limit=None): - out = os.path.join(ROOT, "data", "manifests", "ntuh.jsonl") - reg_inv = "/mnt/pve/SRS/NTUH2022G4/register_inv" + out = os.path.join(DATA, "manifests", "ntuh.jsonl") + reg_inv = path("ntuh_register_inv") cands = {} # (subj, case, ts) -> list of (pref, ser, fname) n_subj = 0 for s in sorted(os.listdir(reg_inv)): diff --git a/scripts/02_build_m6_dataset.py b/scripts/02_build_m6_dataset.py index ba7a62b..1f46519 100644 --- a/scripts/02_build_m6_dataset.py +++ b/scripts/02_build_m6_dataset.py @@ -16,9 +16,9 @@ import argparse import numpy as np import SimpleITK as sitk from concurrent.futures import ProcessPoolExecutor, as_completed -from src.common import ROOT, d, save_jsonl, read_nii_img +from src.common import ROOT, DATA, d, save_jsonl, read_nii_img, path -BASE = "/mnt/pve/WORKSPACE/M6-2025/nii" +BASE = path("m6") T1C_RE = re.compile(r"T1.*\+C|fl3d.*\+.*c", re.I) EXCL_RE = re.compile(r"FLAIR|DTI|vibe|dixon|t2|SWI|T2|MPR_Cor", re.I) @@ -138,8 +138,8 @@ def main(): unlabeled = [r for r in unlabeled if r["key"] not in {k["key"] for k in keep}] for r in keep + unlabeled: r.pop("ct", None) - save_jsonl(keep, os.path.join(ROOT, "data/manifests/m6_labeled.jsonl")) - save_jsonl(unlabeled, os.path.join(ROOT, "data/manifests/m6_unlabeled.jsonl")) + save_jsonl(keep, os.path.join(DATA, "manifests/m6_labeled.jsonl")) + save_jsonl(unlabeled, os.path.join(DATA, "manifests/m6_unlabeled.jsonl")) print(f"final labeled={len(keep)} unlabeled={len(unlabeled)}") diff --git a/scripts/03_scan_lee_t1c.py b/scripts/03_scan_lee_t1c.py index 66391af..b7829ba 100644 --- a/scripts/03_scan_lee_t1c.py +++ b/scripts/03_scan_lee_t1c.py @@ -15,9 +15,9 @@ import re import json import argparse import random -from src.common import ROOT, save_jsonl, load_jsonl, is_head_series +from src.common import ROOT, DATA, save_jsonl, load_jsonl, is_head_series, path -BASE = "/mnt/t24/Public/lee" +BASE = path("lee") T1_NAME_RE = re.compile(r"t1|tfl|spgr|mp2rage|tse3d|vfl|mpage", re.I) EXCL_RE = re.compile(r"\bt2\b|dwi|dti|mra|mrv|angi|swi|bold|\bpp2d|\bpp3d|perf|t2\*|t2star", re.I) @@ -128,15 +128,17 @@ def scan_timepoint(sid, date, tpd): "jpg_dir": jpg_dir, "max_sp": max_spacing(t), "acq": "3" if "3D" in t.get((24, 35), "") else "2"}) - # series thicker than MAX_SPACING are kept only when this timepoint has - # no thinner valid T1c candidate; then keep the best thick one + # series thicker than MAX_SPACING are the primary candidates only when this + # timepoint has no thinner valid T1c; the others are flagged + # thick_dropped: out of the selected manifest, but kept in raw so 04 can + # fall back to them when every other T1c of the timepoint is lost thin = [c for c in cand if c["max_sp"] <= MAX_SPACING] - if thin: - cand = thin - elif cand: - cand = [max(cand, key=lambda c: (c["acq"], c["n_slices"]))] - cand.sort(key=lambda c: (c["acq"], c["n_slices"]), reverse=True) - return cand + keep = thin if thin else [max(cand, key=lambda c: (c["acq"], c["n_slices"]))] + dropped = [c for c in cand if c not in keep] + for c in dropped: + c["thick_dropped"] = True + keep.sort(key=lambda c: (c["acq"], c["n_slices"]), reverse=True) + return keep + dropped def main(): @@ -151,7 +153,7 @@ def main(): help="select every head T1c candidate (no subject/timepoint caps)") args = ap.parse_args() random.seed(args.seed) - raw_path = os.path.join(ROOT, "data/manifests/lee_t1c_raw.jsonl") + raw_path = os.path.join(DATA, "manifests/lee_t1c_raw.jsonl") if args.select_only: rows = load_jsonl(raw_path) print(f"select-only: {len(rows)} raw rows") @@ -162,6 +164,7 @@ def main(): with ThreadPoolExecutor(max_workers=32) as ex: head = list(ex.map(is_head_row, rows)) rows = [r for r, ok in zip(rows, head) if ok] + rows = [r for r in rows if not r.get("thick_dropped")] print(f"head/brain rows: {len(rows)}") else: rows = [] @@ -179,22 +182,28 @@ def main(): continue date = tp[:8] for c in scan_timepoint(sid, date, tpd): - rows.append({"sid": sid, "date": date, "tp": tp, - "jpg_dir": c["jpg_dir"], "ser": c["ser"], "n_slices": c["n_slices"], - "txt_first": c["txt_first"], "why": c["why"], - "key": f"lee_{sid}_{date}_s{c['ser']}"}) + row = {"sid": sid, "date": date, "tp": tp, + "jpg_dir": c["jpg_dir"], "ser": c["ser"], "n_slices": c["n_slices"], + "txt_first": c["txt_first"], "why": c["why"], + "key": f"lee_{sid}_{date}_s{c['ser']}"} + if c.get("thick_dropped"): + row["thick_dropped"] = True + rows.append(row) if i % 100 == 0 and i: print(f"scanned {i}/{len(sids)} subjects, {len(rows)} T1c candidates", flush=True) save_jsonl(rows, raw_path) print(f"total T1c candidates: {len(rows)}") + # thick_dropped rows stay out of the selection; they remain in raw for 04's + # last-resort pass on timepoints that otherwise yield no T1c volume + base = [r for r in rows if not r.get("thick_dropped")] if args.all: # no caps: every head T1c candidate enters the dataset - sel = list(rows) + sel = list(base) else: # subset selection: prefer subjects with more timepoints (longitudinal consistency) by_subj = {} - for r in rows: + for r in base: by_subj.setdefault(r["sid"], []).append(r) for v in by_subj.values(): v.sort(key=lambda x: x["date"]) @@ -219,7 +228,7 @@ def main(): for k in list(single)[: args.max_single_subj]: sel.extend(single[k]) sel.sort(key=lambda x: (x["sid"], x["date"], int(x["ser"]))) - save_jsonl(sel, os.path.join(ROOT, "data/manifests/lee_t1c_selected.jsonl")) + save_jsonl(sel, os.path.join(DATA, "manifests/lee_t1c_selected.jsonl")) print(f"selected: {len(sel)} timepoints from {len(set(r['sid'] for r in sel))} subjects") diff --git a/scripts/04_reconstruct_lee.py b/scripts/04_reconstruct_lee.py index 80f88a9..ffd175c 100644 --- a/scripts/04_reconstruct_lee.py +++ b/scripts/04_reconstruct_lee.py @@ -5,7 +5,7 @@ slice positions are reconstructed by fitting a linear IPP(s) model, validated against all available samples. Usage: python scripts/04_reconstruct_lee.py [--manifest ...selected.jsonl] [--workers 48] Output: data/lee_nii//_s.nii.gz (uint8, native grid) - + data/screenshots//_s.png QA screenshot (1x3 axial/coronal/sagittal) + + data/qa//_s.png QA screenshot (1x3 axial/coronal/sagittal) Non-head (spine/abdomen/breast/...) series are rejected via is_head_series. Fallback: if a timepoint otherwise yields no series, its rows are retried with a relaxed min head extent (--fallback-extent, default 40 mm, vs 60 mm), so a @@ -14,6 +14,12 @@ candidate for that exam. Dynamic-frame series (descriptions like "(exam/frame/phase)-(exam/frame/phase)", contrast-dynamics exports) are excluded whenever the same timepoint has other T1c candidates, and any artifacts of theirs are pruned. +Last resort: a timepoint whose T1c series all fail is reconstructed from ALL +of its T1c candidates, including thick series the scan filtered out of the +selected manifest (thick_dropped rows of the raw manifest). +Last ditch: a timepoint that still has no T1c volume keeps its candidates +with no min-extent floor (all other quality gates still apply) — a thin slab +is preferred over an empty timepoint. """ import os import sys @@ -26,7 +32,7 @@ import numpy as np import SimpleITK as sitk from concurrent.futures import ProcessPoolExecutor, as_completed from PIL import Image -from src.common import ROOT, d, load_jsonl, is_head_series +from src.common import ROOT, DATA, d, load_jsonl, is_head_series from make_screenshots import screenshot_from_volume @@ -94,10 +100,10 @@ def prune_artifacts(keys): for k in sorted(keys): parts = k.split("_") sid, date, sser = parts[1], parts[2], parts[3] - for p in (os.path.join(ROOT, "data", "lee_nii", sid, f"{date}_{sser}.nii.gz"), - os.path.join(ROOT, "data", "proc", k + ".nii.gz"), - os.path.join(ROOT, "data", "procmeta", k + ".json"), - os.path.join(ROOT, "data", "screenshots", sid, f"{date}_{sser}.png")): + for p in (os.path.join(DATA, "lee_nii", sid, f"{date}_{sser}.nii.gz"), + os.path.join(DATA, "proc", k + ".nii.gz"), + os.path.join(DATA, "procmeta", k + ".json"), + os.path.join(DATA, "qa", sid, f"{date}_{sser}.png")): if os.path.exists(p): os.remove(p) n += 1 @@ -236,7 +242,7 @@ def reconstruct(row, extent_min=60.0): img.SetDirection(direction) os.makedirs(os.path.dirname(out), exist_ok=True) sitk.WriteImage(img, out, True) - shot = screenshot_from_volume(vol, direction, row["key"], d("data/screenshots"), + shot = screenshot_from_volume(vol, direction, row["key"], d("data/qa"), spacing=(float(ps[1]), float(ps[0]), float(dmed)), series_desc=series, protocol=proto) msg = f"{vol.shape} n={nread}" + (f" shot={os.path.basename(shot)}" if shot @@ -244,7 +250,7 @@ def reconstruct(row, extent_min=60.0): return row["key"], True, msg -def run_pool(rows, extent_min, workers): +def run_pool(rows, extent_min, workers, reasons=None): ok = err = 0 if rows: with ProcessPoolExecutor(max_workers=workers) as ex: @@ -258,6 +264,8 @@ def run_pool(rows, extent_min, workers): ok += 1 else: err += 1 + if reasons is not None: + reasons[k] = msg if err <= 40: print(" ERR", k, msg, flush=True) if i % 100 == 0: @@ -278,37 +286,120 @@ def zero_tp_rows(uniq): and r["sid"] in sub_ok and (r["sid"], r["date"]) not in tp_ok] +def write_excluded_notes(zero_rows, all_rows, thick_rows, dyn_excl, reasons): + """Write data/qa//_no_t1c.md for every timepoint with no + surviving T1c volume, listing each candidate series and its rejection reason.""" + tps = {(r["sid"], r["date"]) for r in zero_rows} + written = [] + for sid, date in sorted(tps): + lst = [r for r in all_rows if (r["sid"], r["date"]) == (sid, date)] + lst += [r for r in thick_rows if (r["sid"], r["date"]) == (sid, date)] + lst.sort(key=lambda r: int(r["ser"])) + lines = [f"# {sid} {date}: no T1c volume", "", + "This timepoint has no reconstructed T1c volume: every T1c " + "candidate series was rejected.", "", + "| series | slices | description | reason |", + "|---|---|---|---|"] + for r in lst: + if r["key"] in dyn_excl: + why = "excluded: dynamic-frame series (other T1c in this timepoint)" + else: + why = reasons.get(r["key"], "no output") + desc = (r.get("why") or "").strip().replace("|", "/") + lines.append(f"| s{r['ser']} | {r.get('n_slices', '?')} | {desc} | {why} |") + out = os.path.join(DATA, "qa", sid, f"{date}_no_t1c.md") + os.makedirs(os.path.dirname(out), exist_ok=True) + with open(out, "w") as f: + f.write("\n".join(lines) + "\n") + written.append(out) + return written + + +def remove_stale_notes(): + """Drop no_t1c notes for timepoints that now have at least one volume.""" + import glob + n = 0 + for f in glob.glob(os.path.join(DATA, "qa", "*", "*_no_t1c.md")): + sid = os.path.basename(os.path.dirname(f)) + date = os.path.basename(f)[:8] + if glob.glob(os.path.join(DATA, "lee_nii", sid, f"{date}_s*.nii.gz")): + os.remove(f) + n += 1 + return n + + def main(): ap = argparse.ArgumentParser() - ap.add_argument("--manifest", default=os.path.join(ROOT, "data/manifests/lee_t1c_selected.jsonl")) + ap.add_argument("--manifest", default=os.path.join(DATA, "manifests/lee_t1c_selected.jsonl")) ap.add_argument("--workers", type=int, default=48) ap.add_argument("--fallback-extent", type=float, default=40.0, help="min head extent (mm) when retrying timepoints that otherwise " "yielded no series (default 40 vs the usual 60; 0 disables)") ap.add_argument("--fallback-only", action="store_true", help="skip the main pass; only run the zero-timepoint fallback") + ap.add_argument("--raw-manifest", + default=os.path.join(DATA, "manifests/lee_t1c_raw.jsonl"), + help="raw candidate manifest holding thick_dropped rows") args = ap.parse_args() rows = load_jsonl(args.manifest) - seen, uniq = set(), [] + seen, uniq_all = set(), [] for r in rows: # same key can appear twice (double-exported timepoints) if r["key"] not in seen: seen.add(r["key"]) - uniq.append(r) - dyn_excl = excluded_dynamic_keys(uniq) + uniq_all.append(r) + dyn_excl = excluded_dynamic_keys(uniq_all) + uniq = uniq_all if dyn_excl: n = prune_artifacts(dyn_excl) - uniq = [r for r in uniq if r["key"] not in dyn_excl] + uniq = [r for r in uniq_all if r["key"] not in dyn_excl] print(f"excluded {len(dyn_excl)} dynamic-frame series (other T1c in the same " f"timepoint), pruned {n} artifacts") + reasons = {} if not args.fallback_only: todo = [r for r in uniq if not os.path.exists(out_path_for(r))] print(f"todo={len(todo)} workers={args.workers}") - run_pool(todo, extent_min=60.0, workers=args.workers) + run_pool(todo, extent_min=60.0, workers=args.workers, reasons=reasons) if args.fallback_extent > 0: fb = zero_tp_rows(uniq) print(f"fallback todo={len(fb)} extent>={args.fallback_extent:g}mm " f"(timepoints with no other T1c candidate)") - run_pool(fb, extent_min=args.fallback_extent, workers=args.workers) + run_pool(fb, extent_min=args.fallback_extent, workers=args.workers, reasons=reasons) + runiq = [] + if args.fallback_extent > 0 and os.path.exists(args.raw_manifest): + # last resort: a timepoint whose T1c series all failed keeps them ALL + # (e.g. thick series the scan filtered out of the selected manifest) + rseen = set() + for r in load_jsonl(args.raw_manifest): + if r.get("thick_dropped") and r["key"] not in rseen: + rseen.add(r["key"]) + runiq.append(r) + if runiq: + zero = {(r["sid"], r["date"]) for r in zero_tp_rows(uniq)} + lr = [r for r in runiq + if (r["sid"], r["date"]) in zero and not os.path.exists(out_path_for(r))] + print(f"last-resort todo={len(lr)} extent>={args.fallback_extent:g}mm " + f"(keep all T1c of an otherwise-empty timepoint)") + run_pool(lr, extent_min=args.fallback_extent, workers=args.workers, + reasons=reasons) + # last ditch: a timepoint that still has no T1c volume keeps its + # candidates without the min-extent floor (thin slab > empty timepoint) + zero = zero_tp_rows(uniq + runiq) + if zero: + zt = {(r["sid"], r["date"]) for r in zero} + ld = [r for r in uniq + runiq + if (r["sid"], r["date"]) in zt and not os.path.exists(out_path_for(r))] + print(f"last-ditch todo={len(ld)} (empty timepoints keep all T1c, " + f"no extent floor)") + run_pool(ld, extent_min=0.0, workers=args.workers, reasons=reasons) + zero = zero_tp_rows(uniq + runiq) # thick rescues count as timepoint output + if zero: + notes = write_excluded_notes(zero, uniq_all, runiq, dyn_excl, reasons) + print(f"exclusion notes: {len(notes)}") + for p in notes: + print(" ", os.path.relpath(p, ROOT)) + nstale = remove_stale_notes() + if nstale: + print(f"removed {nstale} stale notes") if __name__ == "__main__": diff --git a/scripts/05_build_splits.py b/scripts/05_build_splits.py index 1b48322..4f4dce5 100644 --- a/scripts/05_build_splits.py +++ b/scripts/05_build_splits.py @@ -19,7 +19,7 @@ import json import argparse import random import numpy as np -from src.common import ROOT, d, load_jsonl, save_jsonl, read_nii_arr +from src.common import ROOT, DATA, d, load_jsonl, save_jsonl, read_nii_arr def proc_row(r, prefix): @@ -45,7 +45,7 @@ def main(): labeled_rows = [] for name, prefix in (("ntuh", "ntuh"), ("m6_labeled", "m6")): - for r in load_jsonl(os.path.join(ROOT, "data", "manifests", name + ".jsonl")): + for r in load_jsonl(os.path.join(DATA, "manifests", name + ".jsonl")): if not r.get("label"): continue p = proc_row(r, prefix) @@ -72,14 +72,14 @@ def main(): else: splits["train"].append(r) for k, v in splits.items(): - save_jsonl(v, os.path.join(ROOT, "data/manifests", f"split_{k}.jsonl")) + save_jsonl(v, os.path.join(DATA, "manifests", f"split_{k}.jsonl")) print(f" split_{k}: {len(v)} volumes / {len(set(r['subject'] for r in v))} subjects") labeled_subjects = train_subj | val_subj | test_subj # unlabeled pool pool = [] - for r in load_jsonl(os.path.join(ROOT, "data/manifests/m6_unlabeled.jsonl")): + for r in load_jsonl(os.path.join(DATA, "manifests/m6_unlabeled.jsonl")): if f"m6_{r['subject']}" in labeled_subjects: continue p = proc_row(r, "m6") @@ -87,10 +87,21 @@ def main(): p["date"] = r.get("date") p["source"] = "m6" pool.append(p) - lee_sel_path = os.path.join(ROOT, "data/manifests/lee_t1c_selected.jsonl") + lee_sel_path = os.path.join(DATA, "manifests/lee_t1c_selected.jsonl") lee_rows = {} + lee_manifest = [] if os.path.exists(lee_sel_path): - for r in load_jsonl(lee_sel_path): + lee_manifest = load_jsonl(lee_sel_path) + raw_path = os.path.join(DATA, "manifests", "lee_t1c_raw.jsonl") + if os.path.exists(raw_path): + # last-resort rescues: thick series the scan filtered out, kept + # when their timepoint otherwise yielded no T1c volume + have = {r["key"] for r in lee_manifest} + for r in load_jsonl(raw_path): + if r.get("thick_dropped") and r["key"] not in have: + lee_manifest.append(r) + have.add(r["key"]) + for r in lee_manifest: nii = os.path.join(d("data/lee_nii"), r["sid"], f"{r['date']}_s{r['ser']}.nii.gz") proc = os.path.join(d("data/proc"), r["key"] + ".nii.gz") if os.path.exists(proc): @@ -99,7 +110,7 @@ def main(): lee_rows[r["key"]] = p pool.extend(lee_rows.values()) pool.sort(key=lambda r: (r["subject"], r.get("date", ""))) - save_jsonl(pool, os.path.join(ROOT, "data/manifests/unlabeled_pool.jsonl")) + save_jsonl(pool, os.path.join(DATA, "manifests/unlabeled_pool.jsonl")) print(f"unlabeled pool: {len(pool)} volumes ({sum(1 for r in pool if r['source']=='m6')} m6, " f"{sum(1 for r in pool if r['source']=='lee')} lee) from {len(set(r['subject'] for r in pool))} subjects") @@ -118,8 +129,8 @@ def main(): "p50": float(np.percentile(v, 50)), "p95": float(np.percentile(v, 95)), "p98": float(np.percentile(v, 98)), "max": float(v.max()), "zero_frac": float((v == 0).mean())} - save_jsonl(vols, os.path.join(ROOT, "data/vols.jsonl")) - with open(os.path.join(ROOT, "data/vols.json"), "w") as f: + save_jsonl(vols, os.path.join(DATA, "vols.jsonl")) + with open(os.path.join(DATA, "vols.json"), "w") as f: json.dump(stats, f, indent=1) print("tumor volumes (mm3):", stats) diff --git a/scripts/06_pseudo_label.py b/scripts/06_pseudo_label.py index 019ef8e..9db934e 100644 --- a/scripts/06_pseudo_label.py +++ b/scripts/06_pseudo_label.py @@ -17,7 +17,7 @@ import argparse import numpy as np import SimpleITK as sitk from scipy import ndimage -from src.common import ROOT, d, load_jsonl, save_jsonl, read_nii_arr +from src.common import ROOT, DATA, d, load_jsonl, save_jsonl, read_nii_arr from src import training @@ -90,8 +90,8 @@ def main(): shard = rows[rank::world] vstats = {} - if os.path.exists(os.path.join(ROOT, "data/vols.json")): - vstats = json.load(open(os.path.join(ROOT, "data/vols.json"))) + if os.path.exists(os.path.join(DATA, "vols.json")): + vstats = json.load(open(os.path.join(DATA, "vols.json"))) vol_lo = vstats.get(f"p{args.vol_qp[0]:.0f}", 1.0) vol_hi = vstats.get(f"p{args.vol_qp[1]:.0f}", 50000.0) if rank == 0: diff --git a/scripts/09_run_iterative.py b/scripts/09_run_iterative.py index 56a519d..156e6de 100644 --- a/scripts/09_run_iterative.py +++ b/scripts/09_run_iterative.py @@ -15,7 +15,7 @@ import argparse import json import shutil import subprocess -from src.common import ROOT, load_jsonl, save_jsonl +from src.common import ROOT, DATA, load_jsonl, save_jsonl def run(cmd, log, retries=2): @@ -51,7 +51,7 @@ def main(): ap.add_argument("--batch", type=int, default=3) args = ap.parse_args() - man = os.path.join(ROOT, "data/manifests") + man = os.path.join(DATA, "manifests") train_f = os.path.join(man, "split_train.jsonl") val_f = os.path.join(man, "split_val.jsonl") test_f = os.path.join(man, "split_test.jsonl") @@ -69,7 +69,7 @@ def main(): run(f"python scripts/08_eval.py --rows {test_f} --ckpt runs/round0/best.pt", f"{main_log}") for k in range(1, args.rounds + 1): - pdir = os.path.join(ROOT, "data/pseudo/round" + str(k)) + pdir = os.path.join(DATA, "pseudo/round" + str(k)) run(torchrun(args.gpus, "scripts/06_pseudo_label.py", f"--ckpt runs/round{k-1}/best.pt --unlabeled {pool_f} --out {pdir}"), f"{main_log}") # build round-k training manifest: labeled + accepted pseudo rows (weighted) @@ -100,7 +100,7 @@ def main(): r = json.load(open(f)) table.append({"round": k, "test_dice": round(r["dice"], 4), "n_test": r["n"], "ckpt_epoch": r.get("epoch")}) - pf = os.path.join(ROOT, f"data/pseudo/round{k}/summary.json") if k > 0 else None + pf = os.path.join(DATA, f"pseudo/round{k}/summary.json") if k > 0 else None if pf and os.path.exists(pf): s = json.load(open(pf)) table[-1].update({"n_pos": s["n_pos"], "n_neg": s["n_neg"], diff --git a/scripts/benchmark_pipelines.py b/scripts/benchmark_pipelines.py index e71b264..87b37df 100644 --- a/scripts/benchmark_pipelines.py +++ b/scripts/benchmark_pipelines.py @@ -52,7 +52,7 @@ sys.path.insert(0, os.path.join(ROOT, "scripts_nnu")) import numpy as np import torch -from src.common import load_jsonl, read_nii_arr +from src.common import DATA, load_jsonl, read_nii_arr from src.training import dice_np @@ -95,7 +95,7 @@ def parse_args(): def native_manifests(): """key -> {img, label} from the labeled source manifests (ntuh + m6_labeled).""" - man = os.path.join(ROOT, "data", "manifests") + man = os.path.join(DATA, "manifests") out = {} for name in ("ntuh.jsonl", "m6_labeled.jsonl"): p = os.path.join(man, name) @@ -119,7 +119,7 @@ def native_rows_for(keys): def select_rows(args): - man = os.path.join(ROOT, "data", "manifests") + man = os.path.join(DATA, "manifests") train = [r for r in load_jsonl(os.path.join(man, "split_train.jsonl")) if r.get("plabel")] test = [r for r in load_jsonl(os.path.join(man, "split_test.jsonl")) if r.get("plabel")] if not train: diff --git a/scripts/make_screenshots.py b/scripts/make_screenshots.py index cfcaf8d..d1f5667 100644 --- a/scripts/make_screenshots.py +++ b/scripts/make_screenshots.py @@ -13,7 +13,7 @@ centered vertically, widths follow the physical aspect of each cut). Usage: python scripts/make_screenshots.py --keys k1 k2 ... python scripts/make_screenshots.py --sample 20 [--keys ...] -Output: data/screenshots//.png (per-patient, mirrors data/lee_nii) +Output: data/qa//.png (per-patient, mirrors data/lee_nii) """ import os import sys @@ -163,7 +163,7 @@ def main(): ap.add_argument("--keys", nargs="*", default=[]) ap.add_argument("--sample", type=int, default=0) ap.add_argument("--seed", type=int, default=0) - ap.add_argument("--out", default=d("data/screenshots")) + ap.add_argument("--out", default=d("data/qa")) args = ap.parse_args() os.makedirs(args.out, exist_ok=True) allkeys = [os.path.basename(f)[:-7] for f in glob.glob(os.path.join(d("data/proc"), "*.nii.gz")) diff --git a/scripts/scan_procs.py b/scripts/scan_procs.py index c7db157..939853e 100644 --- a/scripts/scan_procs.py +++ b/scripts/scan_procs.py @@ -4,6 +4,7 @@ import sys sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import glob import json +from src.common import DATA import numpy as np import SimpleITK as sitk from concurrent.futures import ProcessPoolExecutor, as_completed @@ -18,7 +19,7 @@ def check(f): def main(): - files = sorted(glob.glob(os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "data", "proc", "*.nii.gz"))) + files = sorted(glob.glob(os.path.join(DATA, "proc", "*.nii.gz"))) print("checking", len(files), flush=True) bad = [] with ProcessPoolExecutor(max_workers=48) as ex: @@ -32,7 +33,7 @@ def main(): print("BAD volumes:", len(bad)) for b in sorted(bad): print(" ", b) - with open(os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "data", "bad_procs.json"), "w") as f: + with open(os.path.join(DATA, "bad_procs.json"), "w") as f: json.dump(sorted(bad), f, indent=1) diff --git a/scripts/test_dataloader.py b/scripts/test_dataloader.py index 9fec62b..a548b15 100644 --- a/scripts/test_dataloader.py +++ b/scripts/test_dataloader.py @@ -6,8 +6,8 @@ sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) def main(): from src import training - from src.common import load_jsonl - rows = load_jsonl(os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "data", "manifests", "split_train.jsonl")) + from src.common import DATA, load_jsonl + rows = load_jsonl(os.path.join(DATA, "manifests", "split_train.jsonl")) for seed in (0, 1, 2): ds, dl = training.make_dataloader(rows, (96, 96, 96), 3, True, num_workers=4, seed=seed) n_none, n = 0, 0 diff --git a/scripts_monai/monai_common.py b/scripts_monai/monai_common.py index eb3beb3..77f8614 100644 --- a/scripts_monai/monai_common.py +++ b/scripts_monai/monai_common.py @@ -20,8 +20,9 @@ import os import sys # Data/run dirs: overridable for scratch runs. Code (src/, scripts_nnu/): next to this file. -ROOT = os.environ.get("LONGITUDINAL_ROOT", "/mnt/b4/xfr/git26/longitudinal") -_REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +ROOT = os.environ.get( + "LONGITUDINAL_ROOT", os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +_REPO = ROOT sys.path.insert(0, _REPO) sys.path.insert(0, os.path.join(_REPO, "scripts_nnu")) diff --git a/scripts_nnu/06_nnu_run_iterative.py b/scripts_nnu/06_nnu_run_iterative.py index 8575449..91f3e31 100644 --- a/scripts_nnu/06_nnu_run_iterative.py +++ b/scripts_nnu/06_nnu_run_iterative.py @@ -19,7 +19,7 @@ import subprocess sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from src.common import load_jsonl, save_jsonl -from nnu_common import ROOT, d, best_ckpt +from nnu_common import ROOT, DATA, d, best_ckpt def main(): @@ -41,7 +41,7 @@ def main(): ap.add_argument("--vol-ratio", type=float, default=10.0) args = ap.parse_args() - man = os.path.join(ROOT, "data/manifests") + man = os.path.join(DATA, "manifests") train_f = os.path.join(man, "split_train.jsonl") val_f = os.path.join(man, "split_val.jsonl") test_f = os.path.join(man, "split_test.jsonl") diff --git a/scripts_nnu/nnu_common.py b/scripts_nnu/nnu_common.py index e81d9e9..17faa23 100644 --- a/scripts_nnu/nnu_common.py +++ b/scripts_nnu/nnu_common.py @@ -10,9 +10,10 @@ import json import shutil import subprocess -ROOT = os.environ.get("LONGITUDINAL_ROOT", "/mnt/b4/xfr/git26/longitudinal") +ROOT = os.environ.get( + "LONGITUDINAL_ROOT", os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) sys.path.insert(0, ROOT) -from src.common import load_jsonl, save_jsonl, read_nii_arr, head_mask_from_image, largest_cc # noqa: E402 +from src.common import DATA, d, load_jsonl, save_jsonl, read_nii_arr, head_mask_from_image, largest_cc # noqa: E402 import numpy as np import SimpleITK as sitk @@ -61,12 +62,6 @@ def final_ckpt(): return os.path.join(fold_dir(), "checkpoint_final.pth") -def d(name): - p = os.path.join(ROOT, name) - os.makedirs(p, exist_ok=True) - return p - - def nnu_env(epoch=None, lr=None, warmstart=None): env = dict(os.environ) env["nnUNet_raw"] = os.path.join(nnu_root(), "raw") @@ -366,5 +361,5 @@ def consistency_filter(rows, out_dir, max_rel_dist_mm=40.0, vol_ratio_max=10.0, def load_voxel_stats(): - p = os.path.join(ROOT, "data/vols.json") + p = os.path.join(DATA, "vols.json") return json.load(open(p)) if os.path.exists(p) else {} \ No newline at end of file diff --git a/src/common.py b/src/common.py index 58c8e34..f2a4890 100644 --- a/src/common.py +++ b/src/common.py @@ -6,7 +6,38 @@ import numpy as np import SimpleITK as sitk from scipy import ndimage -ROOT = os.environ.get("LONGITUDINAL_ROOT", "/mnt/b4/xfr/git26/longitudinal") +# Repo root: next to this file's parent; overridable for scratch checkouts. +ROOT = os.environ.get( + "LONGITUDINAL_ROOT", os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + + +def _load_paths(): + """All path config from config/paths.json ({} when absent/unreadable).""" + p = os.path.join(ROOT, "config", "paths.json") + if os.path.exists(p): + try: + with open(p) as f: + v = json.load(f) + if isinstance(v, dict): + return v + except (OSError, ValueError): + pass + return {} + + +PATHS = _load_paths() + + +def path(name): + """Named external path from config/paths.json (e.g. 'data', 'lee', 'm6').""" + try: + return PATHS[name] + except KeyError: + raise KeyError( + f"missing path {name!r} in config/paths.json (have: {sorted(PATHS)})") + + +DATA = PATHS.get("data") or os.path.join(ROOT, "data") # series from non-head (non-brain/non-skull) exams must not enter the dataset # (includes "head and neck" exams — neck studies are excluded wholesale) @@ -22,7 +53,12 @@ def is_head_series(study_desc, series_desc): def d(name): - p = os.path.join(ROOT, name) + if name == "data": + p = DATA + elif name.startswith("data/"): + p = os.path.join(DATA, name[len("data/"):]) + else: + p = os.path.join(ROOT, name) os.makedirs(p, exist_ok=True) return p