From 8fe1818c7d5261a840262789f9aea055cef28f17 Mon Sep 17 00:00:00 2001 From: Furen Xiao Date: Sun, 27 Sep 2026 07:14:17 +0800 Subject: [PATCH] Enhance T1c series selection and reconstruction processes - Updated regex patterns for improved matching of series names and tags. - Added functionality to reject non-head series based on study and series descriptions. - Implemented max voxel spacing check to filter out series with excessive spacing. - Enhanced the reconstruction script to handle dynamic-frame series exclusions and artifact pruning. - Modified output paths for reconstructed NIfTI files and added QA screenshot generation. - Improved argument parsing in benchmark and pseudo-labeling scripts for better flexibility. - Introduced a new script for generating QA screenshots from reconstructed volumes. --- .gitignore | 1 + BENCHMARKING.md | 11 +- scripts/03_scan_lee_t1c.py | 132 ++++++++++++++------- scripts/04_reconstruct_lee.py | 209 ++++++++++++++++++++++++++------- scripts/05_build_splits.py | 2 +- scripts/06_pseudo_label.py | 6 +- scripts/benchmark_pipelines.py | 25 ++-- scripts/make_screenshots.py | 188 +++++++++++++++++++++++++++++ src/common.py | 13 ++ 9 files changed, 487 insertions(+), 100 deletions(-) create mode 100644 scripts/make_screenshots.py diff --git a/.gitignore b/.gitignore index 78e1108..26e8bc1 100644 --- a/.gitignore +++ b/.gitignore @@ -3,6 +3,7 @@ runs/ results/ logs/ nnu/ +runs_bench/ runs_nnu/ runs_monai/ __pycache__/ diff --git a/BENCHMARKING.md b/BENCHMARKING.md index 008d052..2ad0b62 100644 --- a/BENCHMARKING.md +++ b/BENCHMARKING.md @@ -57,7 +57,8 @@ conda activate longitudinal python scripts/benchmark_pipelines.py \ --train-epochs 1 --max-rows 24 --nnu-cases 32 --max-eval-rows 4 -# default full benchmark: train 2 epochs on 120 rows, eval on 12 test volumes +# default full benchmark: train 100 epochs on 80% of train cases (~1986 rows), +# eval on 10% of test volumes (~30) python scripts/benchmark_pipelines.py ``` @@ -79,10 +80,10 @@ plus the artifacts it created (safe to delete): `runs_bench/{A,C}/best.pt`, |---|---|---| | `--pipelines` | `A B C` | space-separated subset, e.g. `"A C"` | | `--gpu` | `0` | physical GPU index; all pipelines run serially on this one GPU | -| `--train-epochs` | `2` | training budget for A and C (and B via `NNU_PL_EPOCHS`) | -| `--max-rows` | `120` | train rows for A/C (first N of `split_train.jsonl` that have labels) | -| `--nnu-cases` | `120` | train cases for B (first N labeled native cases of the same split) | -| `--max-eval-rows` | `12` | held-out test volumes for inference + DICE | +| `--train-epochs` | `100` | training budget for A and C (and B via `NNU_PL_EPOCHS`) | +| `--max-rows` | `80%` of `split_train` | train rows for A/C (first N of `split_train.jsonl` that have labels; default 80% of the split, ~1986 rows) | +| `--nnu-cases` | `80%` of `split_train` | train cases for B (first N labeled native cases of the same split; default same as `--max-rows`) | +| `--max-eval-rows` | `10%` of `split_test` | held-out test volumes for inference + DICE (default 10% of the split, ~30 rows) | | `--eval-warmup` | `1` | untimed warmup cases before inference timing | | `--workers` | `4` | data-loader workers (A/C) | | `--batch` | `3` | per-GPU batch (A/C) | diff --git a/scripts/03_scan_lee_t1c.py b/scripts/03_scan_lee_t1c.py index acbdfce..66391af 100644 --- a/scripts/03_scan_lee_t1c.py +++ b/scripts/03_scan_lee_t1c.py @@ -15,11 +15,11 @@ import re import json import argparse import random -from src.common import ROOT, save_jsonl, load_jsonl +from src.common import ROOT, save_jsonl, load_jsonl, is_head_series BASE = "/mnt/t24/Public/lee" T1_NAME_RE = re.compile(r"t1|tfl|spgr|mp2rage|tse3d|vfl|mpage", re.I) -EXCL_RE = re.compile(r"flair|\bt2\b|dwi|dti|mra|angi|swi|bold|\bpp2d|\bpp3d|perf|t2\*|t2star", re.I) +EXCL_RE = re.compile(r"\bt2\b|dwi|dti|mra|mrv|angi|swi|bold|\bpp2d|\bpp3d|perf|t2\*|t2star", re.I) def parse_txt(p): @@ -28,7 +28,7 @@ def parse_txt(p): try: with open(p, errors="ignore") as f: for line in f: - m = re.match(r'"\((\d{4}),(\d{4})\)","([A-Z]+)","(\d+)",(.*),"(.*)"\s*$', line.strip()) + m = re.match(r'"\(([0-9A-Fa-f]{4}),([0-9A-Fa-f]{4})\)","([A-Z]+)","(\d+)",(.*),"(.*)"\s*$', line.strip()) if m: t[(int(m.group(1), 16), int(m.group(2), 16))] = m.group(5).strip().strip('"') except OSError: @@ -37,16 +37,24 @@ def parse_txt(p): def is_t1c(t): - bolus = any(k in t for k in ((24, 4161), (24, 4168), (8, 307))) + # newer exports tag contrast only via (0018,0010) Agent ("Y GD", "POST + # CONTRAST", ...); pre-contrast series leave it empty in all exports + bolus = any(k in t for k in ((24, 4161), (24, 4168), (8, 307))) or bool(t.get((24, 16), "").strip()) if not bolus: return False, "no bolus" - name = (t.get((24, 36), "") + " " + t.get((24, 33), "")).lower() + # newer exports leave Sequence Name empty and put the protocol in the + # series description, so match/exclude on all three + name = (t.get((24, 36), "") + " " + t.get((24, 33), "") + " " + + t.get((8, 0x103E), "")).lower() if EXCL_RE.search(name): return False, "excluded name " + name[:30] + # FLAIR is usually T2-weighted (excluded); T1-FLAIR is a valid T1c + if re.search(r"flair", name) and "t1" not in name: + return False, "excluded name (FLAIR, not T1) " + name[:30] if not T1_NAME_RE.search(name): - seq = t.get((24, 32), "") + seq = t.get((24, 32), "").strip() et = t.get((24, 129), "") - ok = seq in ("GR", "SE", "GR\\IR", "SE ", "GR ") + ok = seq in ("GR", "SE", "GR\\IR", "RM\\IR", "SE ", "GR ") try: et_ok = et.replace('"', "").split()[0].replace(" ", "")[:1] != "" and float(et.split("\\")[0].strip()) < 30 except (ValueError, IndexError): @@ -56,6 +64,30 @@ def is_t1c(t): return True, name[:40] +MAX_SPACING = 4.0 # mm; reject series with any voxel spacing above this + + +def max_spacing(t): + """Max voxel spacing (mm): (0028,0030) pixel spacing + the larger of + (0018,0050) slice thickness and (0018,0088) spacing between slices + (interleaved slices can be thinner than their center-to-center gap). + Returns 0.0 if no spacing info is available.""" + vals = [] + ps = t.get((40, 48), "") + if ps: + try: + vals += [float(x) for x in ps.split("\\")[:2]] + except ValueError: + pass + for v in (t.get((24, 80), ""), t.get((24, 136), "")): + if v: + try: + vals.append(float(v.split("\\")[0].strip())) + except ValueError: + pass + return max(vals) if vals else 0.0 + + def scan_timepoint(sid, date, tpd): """Returns list of candidate T1c series for one MR timepoint dir.""" txts = {} @@ -80,6 +112,8 @@ def scan_timepoint(sid, date, tpd): ok, why = is_t1c(t) if not ok: continue + if not is_head_series(t.get((8, 0x1030)), t.get((8, 0x103E))): + continue rows = t.get((40, 16), ""); cols = t.get((40, 17), "") # count jpgs of this series pat = re.compile(rf"{sid}_{date}_MR_{ser}_(\d+)_(\d+)\.jpg$") @@ -91,13 +125,17 @@ def scan_timepoint(sid, date, tpd): nslices_all = sum(1 for jf in jpgs if pat.match(jf)) cand.append({"ser": ser, "txt_first": first, "n_slices": nslices_all, "rows": rows, "cols": cols, "why": why, - "jpg_dir": jpg_dir}) - # prefer 3D (MRAcq) & most slices - def score(c): - t = parse_txt(c["txt_first"]) - acq = "3" if "3D" in t.get((24, 35), "") else "2" - return (acq, c["n_slices"]) - cand.sort(key=score, reverse=True) + "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 + 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 @@ -109,12 +147,22 @@ def main(): ap.add_argument("--max-single-subj", type=int, default=0) ap.add_argument("--max-subj-timepoints", type=int, default=6) ap.add_argument("--select-only", action="store_true") + ap.add_argument("--all", action="store_true", + 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") if args.select_only: rows = load_jsonl(raw_path) print(f"select-only: {len(rows)} raw rows") + from concurrent.futures import ThreadPoolExecutor + def is_head_row(r): + t = parse_txt(r["txt_first"]) + return is_head_series(t.get((8, 0x1030)), t.get((8, 0x103E))) + 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] + print(f"head/brain rows: {len(rows)}") else: rows = [] sids = [s for s in os.listdir(BASE) if os.path.isdir(os.path.join(BASE, s)) and not s.endswith(".complete")] @@ -140,32 +188,36 @@ def main(): save_jsonl(rows, raw_path) print(f"total T1c candidates: {len(rows)}") - # subset selection: prefer subjects with more timepoints (longitudinal consistency) - by_subj = {} - for r in rows: - by_subj.setdefault(r["sid"], []).append(r) - for v in by_subj.values(): - v.sort(key=lambda x: x["date"]) - multi = {k: v for k, v in by_subj.items() if len(v) >= 2} - single = {k: v for k, v in by_subj.items() if len(v) == 1} - rnd = random.Random(args.seed) - depth = {k: len(v) for k, v in multi.items()} - cand = [k for k in multi if depth[k] >= 3] - deep = sorted(cand, key=lambda k: -depth[k]) - n_deep = max(1, args.max_multi_subj * 2 // 3) - pick = deep[:n_deep] - rest = [k for k in cand if k not in pick] - rnd.shuffle(rest) - pick += rest[: args.max_multi_subj - len(pick)] - sel = [] - for k in pick: - v = multi[k] - if len(v) > args.max_subj_timepoints and args.max_subj_timepoints >= 2: - v = [v[int(j * (len(v) - 1) / (args.max_subj_timepoints - 1))] for j in range(args.max_subj_timepoints)] - sel.extend(v) - rnd.shuffle(list(single)) - for k in list(single)[: args.max_single_subj]: - sel.extend(single[k]) + if args.all: + # no caps: every head T1c candidate enters the dataset + sel = list(rows) + else: + # subset selection: prefer subjects with more timepoints (longitudinal consistency) + by_subj = {} + for r in rows: + by_subj.setdefault(r["sid"], []).append(r) + for v in by_subj.values(): + v.sort(key=lambda x: x["date"]) + multi = {k: v for k, v in by_subj.items() if len(v) >= 2} + single = {k: v for k, v in by_subj.items() if len(v) == 1} + rnd = random.Random(args.seed) + depth = {k: len(v) for k, v in multi.items()} + cand = [k for k in multi if depth[k] >= 3] + deep = sorted(cand, key=lambda k: -depth[k]) + n_deep = max(1, args.max_multi_subj * 2 // 3) + pick = deep[:n_deep] + rest = [k for k in cand if k not in pick] + rnd.shuffle(rest) + pick += rest[: args.max_multi_subj - len(pick)] + sel = [] + for k in pick: + v = multi[k] + if len(v) > args.max_subj_timepoints and args.max_subj_timepoints >= 2: + v = [v[int(j * (len(v) - 1) / (args.max_subj_timepoints - 1))] for j in range(args.max_subj_timepoints)] + sel.extend(v) + rnd.shuffle(list(single)) + 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")) 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 8b9e66a..80f88a9 100644 --- a/scripts/04_reconstruct_lee.py +++ b/scripts/04_reconstruct_lee.py @@ -4,11 +4,21 @@ The timepoint dir holds only sampled per-slice txts (typically slices 1, 2, N); 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/.nii.gz (uint8, native grid) +Output: data/lee_nii//_s.nii.gz (uint8, native grid) + + data/screenshots//_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 +thin-but-complete stack (e.g. 45 mm) is kept when there is no other T1c +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. """ import os import sys sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) import re import json import argparse @@ -16,15 +26,17 @@ 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 +from src.common import ROOT, d, load_jsonl, is_head_series +from make_screenshots import screenshot_from_volume LINE_RE = re.compile( - r'^\s*"\((\d{4}),(\d{4})\)","([A-Z]+)","(\d+)",(?:\"([^\"]*)\"|([^,]+)),"') + r'^\s*"\(([0-9A-Fa-f]{4}),([0-9A-Fa-f]{4})\)","([A-Z]+)","(\d+)",(?:\"([^\"]*)\"|([^,]+)),"') def parse_txt(p): ipp = spacing = iop = rows = cols = None + study = series = proto = "" try: with open(p, errors="ignore") as f: for raw in f: @@ -32,7 +44,8 @@ def parse_txt(p): if not m: continue tag = (int(m.group(1), 16), int(m.group(2), 16)) - if tag not in ((32, 50), (40, 48), (32, 55), (40, 16), (40, 17)): + if tag not in ((32, 50), (40, 48), (32, 55), (40, 16), (40, 17), + (8, 0x1030), (8, 0x103E), (0x18, 0x1030)): continue val = (m.group(5) if m.group(5) is not None else m.group(6)).strip() if not val: @@ -48,31 +61,67 @@ def parse_txt(p): rows = int(float(val)) elif tag == (40, 17): cols = int(float(val)) + elif tag == (8, 0x1030): + study = val + elif tag == (8, 0x103E): + series = val + elif tag == (0x18, 0x1030): + proto = val except (ValueError, IndexError): continue if ipp is None or spacing is None or iop is None or rows is None or cols is None: return None - return ipp, spacing, iop, rows, cols + return ipp, spacing, iop, rows, cols, study, series, proto except OSError: return None -def reconstruct(row): +DYN_FRAME_RE = re.compile(r"\(\d+/\d+/\d+(?:\.\.\d+)?\)\s*-\s*\(") + + +def excluded_dynamic_keys(uniq): + """Keys of dynamic-frame series when the same timepoint has other T1c + candidates (the dynamics frames are redundant re-exports of the exam).""" + has_other = {(r["sid"], r["date"]) for r in uniq + if not DYN_FRAME_RE.search(r.get("why") or "")} + return {r["key"] for r in uniq + if DYN_FRAME_RE.search(r.get("why") or "") + and (r["sid"], r["date"]) in has_other} + + +def prune_artifacts(keys): + n = 0 + 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")): + if os.path.exists(p): + os.remove(p) + n += 1 + return n + + +def out_path_for(row): + return os.path.join(d("data/lee_nii"), row["sid"], f"{row['date']}_s{row['ser']}.nii.gz") + + +def reconstruct(row, extent_min=60.0): sid, date, ser, jpg_dir = row["sid"], row["date"], row["ser"], row["jpg_dir"] - out = os.path.join(d("data/lee_nii"), row["key"] + ".nii.gz") + out = out_path_for(row) if os.path.exists(out): return row["key"], True, "cached" pat = re.compile(rf"^{re.escape(sid)}_{date}_MR_{ser}_(\d+)_(\d+)\.jpg$") tpat = re.compile(rf"^{re.escape(sid)}_{date}_MR_{ser}_(\d+)\.txt$") tp_dir = os.path.dirname(jpg_dir) - jpg_map, txt_samples = {}, [] + cands, txt_samples = [], [] try: for f in os.listdir(jpg_dir): m = pat.match(f) if m: - sl, inst = int(m.group(1)), int(m.group(2)) - if sl not in jpg_map or inst < jpg_map[sl][1]: - jpg_map[sl] = (f, inst) + cands.append((int(m.group(1)), int(m.group(2)), f)) for f in os.listdir(tp_dir): m = tpat.match(f) if m: @@ -81,9 +130,13 @@ def reconstruct(row): txt_samples.append((int(m.group(1)), p)) except OSError as e: return row["key"], False, f"listdir fail {e!r}" - if not jpg_map or not txt_samples: - return row["key"], False, f"no jpg({len(jpg_map)}) or txt({len(txt_samples)})" + if not cands or not txt_samples: + return row["key"], False, f"no jpg({len(cands)}) or txt({len(txt_samples)})" txt_samples.sort() + # reject series from non-head exams (brain/head only dataset) + desc = next((p for _, p in txt_samples if p[5] or p[6]), None) + if desc is not None and not is_head_series(desc[5], desc[6]): + return row["key"], False, "non-head series" # majority (rows, cols, pixel-spacing) as reference geometry from collections import Counter geo = Counter((p[3], p[4], round(p[1][0], 4), round(p[1][1], 4)) for _, p in txt_samples) @@ -92,7 +145,7 @@ def reconstruct(row): good_idx = [sl for sl, p in txt_samples if (p[3], p[4]) == (rows, cols)] if len(good_idx) < 2 or not ref: return row["key"], False, f"only {len(good_idx)} consistent geometry samples" - ipp0, ps, iop0, _, _ = ref[len(ref) // 2] + ipp0, ps, iop0, _, _, study, series, proto = ref[len(ref) // 2] # linear IPP model (fit on geometry-consistent samples only) good = [(sl, p) for sl, p in txt_samples if p[3] == rows and p[4] == cols and np.allclose(p[1], ps, atol=1e-3)] if len(good) < 2: @@ -107,6 +160,22 @@ def reconstruct(row): # the stack direction must follow the nominal slice normal (rejects multiplanar 2D exports) if abs(float(np.dot(step, n))) < 0.6 * float(np.linalg.norm(step)): return row["key"], False, "stack not along slice normal" + # pick one jpg per slice: a series can be exported twice (e.g. the first + # N slices re-rendered at the wrong size); prefer instances matching the + # nominal matrix, tie-break on the smaller instance + want = {(cols, rows), (rows, cols)} + jpg_map = {} + for sl, inst, f in sorted(cands): + try: + sz = Image.open(os.path.join(jpg_dir, f)).size + except Exception: + continue + good = sz in want + cur = jpg_map.get(sl) + if cur is None or inst < cur[1] or (good and not cur[2]): + jpg_map[sl] = (f, inst, good) + if not jpg_map: + return row["key"], False, "no readable jpg slices" sl_list = sorted(jpg_map) def ipp_of(s): @@ -114,11 +183,14 @@ def reconstruct(row): ipp_all = {sl: ipp_of(sl) for sl in sl_list} r0 = min(sl_list, key=lambda s: np.dot(ipp_all[s] - ipp_all[sl_list[0]], n)) p_ref = ipp_all[r0] + # In this export the IOP row/col cosines are swapped relative to the + # rendered pixel matrix (see AutoPACS uni2nii): rows (vertical) run along + # v at the row pixel spacing ps[0], cols (horizontal) along u at ps[1]. zs, pos = [], [] for sl in sl_list: off = ipp_all[sl] - p_ref zs.append(float(np.dot(off, n))) - pos.append((int(round(float(np.dot(off, u)) / ps[0])), int(round(float(np.dot(off, v)) / ps[1])))) + pos.append((int(round(float(np.dot(off, v)) / ps[0])), int(round(float(np.dot(off, u)) / ps[1])))) zs = np.array(zs) dmed = np.median(np.abs(np.diff(np.sort(zs)))) if dmed <= 0 or not np.isfinite(dmed): @@ -128,7 +200,10 @@ def reconstruct(row): ro_max = max(p[0] for p in pos) co_min = min(p[1] for p in pos) co_max = max(p[1] for p in pos) - vol = np.zeros((rows + (ro_max - ro_min), cols + (co_max - co_min), nz), dtype=np.uint8) + # vol layout (nz, rows, cols): image x = cols (jpg width), y = rows + # (jpg height), z = slices — the same axis convention as AutoPACS + # uni2nii (sitk.ReadImage of the jpg stack). + vol = np.zeros((nz, rows + (ro_max - ro_min), cols + (co_max - co_min)), dtype=np.uint8) nread = 0 for idx, (sl, jf, (ro, co)) in enumerate(zip(sl_list, [jpg_map[s][0] for s in sl_list], pos)): zi = int(round(zs[idx] / dmed)) @@ -141,49 +216,99 @@ def reconstruct(row): im = im.T.copy() else: continue - vol[ro - ro_min: ro - ro_min + rows, co - co_min: co - co_min + cols, zi] = im + vol[zi, ro - ro_min: ro - ro_min + rows, co - co_min: co - co_min + cols] = im nread += 1 if nread < 0.95 * len(sl_list) or vol.max() == 0: return row["key"], False, f"only {nread}/{len(sl_list)} slices read" - ext = (vol.shape[0] * ps[0], vol.shape[1] * ps[1], vol.shape[2] * dmed) - if min(ext) < 60 or max(ext) > 350: + ext = ((cols + (co_max - co_min)) * ps[1], (rows + (ro_max - ro_min)) * ps[0], nz * dmed) + if min(ext) < extent_min or max(ext) > 350: return row["key"], False, f"plausible extent failed {tuple(round(e, 1) for e in ext)}mm" - origin = p_ref + (ro_min * ps[0]) * u + (co_min * ps[1]) * v - direction = tuple(float(x) for x in np.concatenate([u, v, n])) + origin = p_ref + (co_min * ps[1]) * u + (ro_min * ps[0]) * v + # SetDirection takes row-major values whose COLUMNS are the image-axis + # directions: col 0 = x axis (cols, u), col 1 = y axis (rows, v), + # col 2 = z axis (slices, n). + direction = tuple(float(x) for x in np.concatenate([u, v, n]).reshape(3, 3).T.ravel()) + # GetImageFromArray maps array index 0 -> image dimension 2 (z), so the + # (nz, rows, cols) array gives image (x=cols, y=rows, z=nz). img = sitk.GetImageFromArray(vol) img.SetOrigin(tuple(float(x) for x in origin)) - img.SetSpacing((float(ps[0]), float(ps[1]), float(dmed))) + img.SetSpacing((float(ps[1]), float(ps[0]), float(dmed))) img.SetDirection(direction) os.makedirs(os.path.dirname(out), exist_ok=True) sitk.WriteImage(img, out, True) - return row["key"], True, f"{vol.shape} n={nread}" + shot = screenshot_from_volume(vol, direction, row["key"], d("data/screenshots"), + 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 + else " (no screenshot: not body-aligned)") + return row["key"], True, msg + + +def run_pool(rows, extent_min, workers): + ok = err = 0 + if rows: + with ProcessPoolExecutor(max_workers=workers) as ex: + futs = {ex.submit(reconstruct, r, extent_min): r for r in rows} + for i, fu in enumerate(as_completed(futs), 1): + try: + k, good, msg = fu.result() + except Exception as e: # noqa + k, good, msg = futs[fu]["key"], False, repr(e) + if good: + ok += 1 + else: + err += 1 + if err <= 40: + print(" ERR", k, msg, flush=True) + if i % 100 == 0: + print(f" {i}/{len(futs)} ok={ok} err={err}", flush=True) + print(f"done ok={ok} err={err}") + + +def zero_tp_rows(uniq): + """Rows without output for timepoints that yielded no series, restricted to + subjects that have at least one reconstructed volume elsewhere.""" + sub_ok, tp_ok = set(), set() + for r in uniq: + if os.path.exists(out_path_for(r)): + sub_ok.add(r["sid"]) + tp_ok.add((r["sid"], r["date"])) + return [r for r in uniq + if not os.path.exists(out_path_for(r)) + and r["sid"] in sub_ok and (r["sid"], r["date"]) not in tp_ok] def main(): ap = argparse.ArgumentParser() ap.add_argument("--manifest", default=os.path.join(ROOT, "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") args = ap.parse_args() rows = load_jsonl(args.manifest) - rows = [r for r in rows if not os.path.exists(os.path.join(d("data/lee_nii"), r["key"] + ".nii.gz"))] - print(f"todo={len(rows)} workers={args.workers}") - ok = err = 0 - with ProcessPoolExecutor(max_workers=args.workers) as ex: - futs = {ex.submit(reconstruct, r): r for r in rows} - for i, fu in enumerate(as_completed(futs), 1): - try: - k, good, msg = fu.result() - except Exception as e: # noqa - k, good, msg = futs[fu]["key"], False, repr(e) - if good: - ok += 1 - else: - err += 1 - if err <= 40: - print(" ERR", k, msg, flush=True) - if i % 100 == 0: - print(f" {i}/{len(futs)} ok={ok} err={err}", flush=True) - print(f"done ok={ok} err={err}") + seen, uniq = 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) + if dyn_excl: + n = prune_artifacts(dyn_excl) + uniq = [r for r in uniq 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") + 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) + 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) if __name__ == "__main__": diff --git a/scripts/05_build_splits.py b/scripts/05_build_splits.py index b4075ce..1b48322 100644 --- a/scripts/05_build_splits.py +++ b/scripts/05_build_splits.py @@ -91,7 +91,7 @@ def main(): lee_rows = {} if os.path.exists(lee_sel_path): for r in load_jsonl(lee_sel_path): - nii = os.path.join(d("data/lee_nii"), r["key"] + ".nii.gz") + 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): p = {"key": r["key"], "subject": f"lee_{r['sid']}", "pimg": proc, diff --git a/scripts/06_pseudo_label.py b/scripts/06_pseudo_label.py index 988d3b5..019ef8e 100644 --- a/scripts/06_pseudo_label.py +++ b/scripts/06_pseudo_label.py @@ -25,9 +25,9 @@ def grid_info(key): p = os.path.join(d("data/procmeta"), key + ".json") m = json.load(open(p)) origin = np.array(m["origin"]) - R = np.array(m["direction"]).reshape(3, 3) # row_dir, col_dir, slice_dir - cv = np.array(m["crop_vox"]) - o = origin + cv[0] * R[0] + cv[1] * R[1] + cv[2] * R[2] + R = np.array(m["direction"]).reshape(3, 3) # columns = x, y, z axis directions + cv = np.array(m["crop_vox"]) # crop starts in (z, y, x) array coords + o = origin + cv[2] * R[:, 0] + cv[1] * R[:, 1] + cv[0] * R[:, 2] return o, R diff --git a/scripts/benchmark_pipelines.py b/scripts/benchmark_pipelines.py index f3e1ba7..e71b264 100644 --- a/scripts/benchmark_pipelines.py +++ b/scripts/benchmark_pipelines.py @@ -66,10 +66,13 @@ def parse_args(): ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("--pipelines", default="A B C", help="space-separated subset of 'A B C'") ap.add_argument("--gpu", type=int, default=0, help="physical GPU index (all pipelines use one GPU)") - ap.add_argument("--train-epochs", type=int, default=2, help="training budget for A and C (and B via NNU_PL_EPOCHS)") - ap.add_argument("--max-rows", type=int, default=120, help="train rows for A/C (first N of split_train)") - ap.add_argument("--nnu-cases", type=int, default=120, help="train cases for B (first N labeled native cases)") - ap.add_argument("--max-eval-rows", type=int, default=12, help="held-out test volumes for inference + DICE") + ap.add_argument("--train-epochs", type=int, default=100, help="training budget for A and C (and B via NNU_PL_EPOCHS)") + ap.add_argument("--max-rows", type=int, default=None, + help="train rows for A/C (first N of labeled split_train; default 80%% of the split)") + ap.add_argument("--nnu-cases", type=int, default=None, + help="train cases for B (first N labeled native cases; default same as --max-rows)") + ap.add_argument("--max-eval-rows", type=int, default=None, + help="held-out test volumes for inference + DICE (default 10%% of labeled split_test)") ap.add_argument("--eval-warmup", type=int, default=1, help="untimed warmup cases before inference timing") ap.add_argument("--workers", type=int, default=4, help="dataloader workers (A/C)") ap.add_argument("--batch", type=int, default=3, help="per-GPU batch (A/C)") @@ -117,15 +120,19 @@ def native_rows_for(keys): def select_rows(args): man = os.path.join(ROOT, "data", "manifests") - train = load_jsonl(os.path.join(man, "split_train.jsonl")) - test = load_jsonl(os.path.join(man, "split_test.jsonl")) - train = [r for r in train if r.get("plabel")][:args.max_rows] - test = [r for r in test if r.get("plabel")][:args.max_eval_rows] + 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: raise SystemExit("split_train.jsonl has no labeled processed rows; run scripts/preprocess.py + 05_build_splits.py") if not test: raise SystemExit("split_test.jsonl has no labeled processed rows; run scripts/preprocess.py + 05_build_splits.py") - return train, test + if args.max_rows is None: + args.max_rows = max(1, round(0.8 * len(train))) + if args.max_eval_rows is None: + args.max_eval_rows = max(1, round(0.1 * len(test))) + if args.nnu_cases is None: + args.nnu_cases = args.max_rows + return train[:args.max_rows], test[:args.max_eval_rows] # ---------------- GPU / timing helpers ---------------- diff --git a/scripts/make_screenshots.py b/scripts/make_screenshots.py new file mode 100644 index 0000000..cfcaf8d --- /dev/null +++ b/scripts/make_screenshots.py @@ -0,0 +1,188 @@ +"""Generate 1x3 QA screenshots (axial/coronal/sagittal) per volume. + +Each proc volume is first mapped to a canonical body array (S, P, L) using the +direction matrix in procmeta (columns = image-axis directions, LPS basis), so +the panels are correct for any source orientation (LPS axial, LIP coronal, +PIR sagittal, ...). Radiological display conventions: + axial: anterior on top, patient's left on right + coronal: superior on top, patient's left on right + sagittal: superior on top, anterior on left +Layout: single row, axial | coronal | sagittal; each panel is scaled so its +maximal physical dimension (mm) is the same across the row (panels are +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) +""" +import os +import sys +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +import argparse +import json +import random +import glob + +import numpy as np +import SimpleITK as sitk +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt +from src.common import ROOT, d, read_nii_img + +B_L = np.array([1.0, 0, 0]) +B_P = np.array([0.0, 1, 0]) +B_S = np.array([0.0, 0, 1]) + + +def vol_to_body(a, M): + """a: array (z,y,x) over image axes; M: 3x3 LPS, columns = axis dirs. + Returns (body, ok, kof): body = (S, P, L)-indexed array, kof maps each + body axis to its image axis index (0=x, 1=y, 2=z); kof is None if not + body-aligned.""" + M = np.asarray(M).reshape(3, 3) + v = a.transpose(2, 1, 0) # (i0, i1, i2) image axes + kof, sgn = {}, {} + for name, b in (("L", B_L), ("P", B_P), ("S", B_S)): + dots = M.T @ b + k = int(np.argmax(np.abs(dots))) + # 0.92 ~= 23 deg: tolerates tilted clinical protocols (e.g. oblique + # sagittal MP-RAGE ~21 deg from P-S) + if abs(dots[k]) < 0.92: + return None, False, None + kof[name], sgn[name] = k, (1 if dots[k] > 0 else -1) + out = np.transpose(v, [kof["S"], kof["P"], kof["L"]]) + for ax, name in enumerate(("S", "P", "L")): + if sgn[name] < 0: + out = np.flip(out, ax) + return np.ascontiguousarray(out), True, kof + + +def panels_from_body(b): + s, p, l = b.shape + sm, pm, lm = s // 2, p // 2, l // 2 + axial = b[sm] # (P, L): top=anterior, right=left + coronal = b[:, pm, :][::-1] # (S-rev, L): top=superior, right=left + sagittal = b[:, :, lm][::-1] # (S-rev, P): top=superior, left=anterior + return axial, coronal, sagittal, (sm, pm, lm) + + +def shot_path(out_dir, key): + """Per-patient screenshot path: //.png, where the + patient id is the 2nd underscore token of the key (lee_/m6_/ntuh_ prefixes).""" + parts = key.split("_") + if len(parts) < 3: + return os.path.join(out_dir, key + ".png") + return os.path.join(out_dir, parts[1], "_".join(parts[2:]) + ".png") + + +def screenshot_from_volume(arr, direction, key, out_dir, dpi=110, spacing=None, + series_desc=None, protocol=None): + """Render the 1x3 QA screenshot (axial | coronal | sagittal) from an + image-axis array (z,y,x), a 3x3 LPS direction matrix (columns = + image-axis directions) and the image-axis spacing in mm, (x, y, z) + (None = isotropic 1mm). series_desc / protocol (optional) are shown + on a second title line. Panel cells follow the physical (mm) aspect + of each cut. Output: //.png (mirrors the + data/lee_nii layout). Returns the output path, or None if the volume + is not body-aligned.""" + if spacing is None: + spacing = (1.0, 1.0, 1.0) + body, ok, kof = vol_to_body(arr, direction) + if not ok: + return None + sp = {n: float(spacing[kof[n]]) for n in ("S", "P", "L")} + lo, hi = np.percentile(body[body > 0], [1, 99.5]) + axial, coronal, sagittal, _ = panels_from_body(body) + s, p, l = body.shape + # (image, title, aspect = mm per row / mm per col, n_rows, n_cols) + panels = [ + (axial, "axial", sp["P"] / sp["L"], p, l), + (coronal, "coronal", sp["S"] / sp["L"], s, l), + (sagittal, "sagittal", sp["S"] / sp["P"], s, p), + ] + max_in, gap_in, m_in, top_in, bot_in = 4.4, 0.25, 0.3, 1.5, 0.2 + # each panel is scaled so its maximal physical dimension spans max_in + w_in, h_in = [], [] + for _, _, a, nrow, ncol in panels: + pw, ph = ncol, nrow * a # physical extent along display x / y + m = max(pw, ph) + w_in.append(max_in * pw / m) + h_in.append(max_in * ph / m) + plot_h = max_in + fig_w = sum(w_in) + 2 * gap_in + 2 * m_in + fig_h = plot_h + top_in + bot_in + fig = plt.figure(figsize=(fig_w, fig_h)) + left = m_in / fig_w + for (im, title, a, _, _), w, h in zip(panels, w_in, h_in): + axi = fig.add_axes([left, (bot_in + (plot_h - h) / 2) / fig_h, + w / fig_w, h / fig_h]) + axi.imshow(im, cmap="gray", vmin=lo, vmax=hi, origin="upper", aspect=a) + axi.set_title(title, color="w", fontsize=11) + axi.set_xticks([]); axi.set_yticks([]) + for spine in axi.spines.values(): + spine.set_edgecolor("0.35") + left += (w + gap_in) / fig_w + mx = arr.shape[::-1] # (x, y, z) + tkey = key.split("_", 1)[1] if "_" in key else key # drop dataset prefix (lee_/m6_/ntuh_) + title = (f"{tkey} {mx[0]}×{mx[1]}×{mx[2]} @ " + f"{spacing[0]:.3f}×{spacing[1]:.3f}×{spacing[2]:.3f}mm") + def _clean(s): + # drop chars the Agg font cannot render (e.g. CJK in old study descs) + return " ".join(s.split()).encode("ascii", "ignore").decode() + + sd = _clean(series_desc) if series_desc else "" + pt = _clean(protocol) if protocol else "" + info = [] + if sd: + info.append("series: " + sd[:60]) + if pt and pt[:10] != sd[:10]: + info.append("protocol: " + pt[:60]) + if info: + title += "\n" + " ".join(info) + fig.suptitle(title, color="w", fontsize=12) + fig.patch.set_facecolor("k") + out = shot_path(out_dir, key) + os.makedirs(os.path.dirname(out), exist_ok=True) + fig.savefig(out, facecolor="k", dpi=dpi) + plt.close(fig) + return out + + +def screenshot(key, out_dir, dpi=110): + p = os.path.join(d("data/proc"), key + ".nii.gz") + if not os.path.exists(p): + return None + arr = np.asarray(sitk.GetArrayFromImage(read_nii_img(p)), dtype=np.float32) + meta = json.load(open(os.path.join(d("data/procmeta"), key + ".json"))) + return screenshot_from_volume(arr, meta["direction"], key, out_dir, dpi) + + +def main(): + ap = argparse.ArgumentParser() + 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")) + 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")) + if not f.endswith("_label.nii.gz")] + keys = list(args.keys) + if args.sample: + rnd = random.Random(args.seed) + extra = [k for k in allkeys if k not in keys] + keys += rnd.sample(extra, min(args.sample - len(keys), len(extra))) + n = 0 + for k in keys: + r = screenshot(k, args.out) + if r: + n += 1 + print("wrote", r, flush=True) + else: + print("skip", k, flush=True) + print(f"done: {n}/{len(keys)} screenshots in {args.out}") + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/src/common.py b/src/common.py index 3f2b1c6..58c8e34 100644 --- a/src/common.py +++ b/src/common.py @@ -1,4 +1,5 @@ import os +import re import json import glob import numpy as np @@ -7,6 +8,18 @@ from scipy import ndimage ROOT = os.environ.get("LONGITUDINAL_ROOT", "/mnt/b4/xfr/git26/longitudinal") +# series from non-head (non-brain/non-skull) exams must not enter the dataset +# (includes "head and neck" exams — neck studies are excluded wholesale) +HEAD_EXCL_RE = re.compile( + r"spine|breast|abdomen|pelvis|chest|extremit|brachial|urograph|prostat|" + r"\bbody\b|thigh|knee|wrist|ankle|foot\b|hand\b|elbow|shoulder|carotid|neck|" + r"(?