longitudinal/scripts/04_reconstruct_lee.py

406 lines
No EOL
19 KiB
Python

"""Reconstruct lee T1c volumes from JPG slices + DICOM txt metadata.
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/<sid>/<date>_s<ser>.nii.gz (uint8, native grid)
+ data/qa/<sid>/<date>_s<ser>.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.
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
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
import numpy as np
import SimpleITK as sitk
from concurrent.futures import ProcessPoolExecutor, as_completed
from PIL import Image
from src.common import ROOT, DATA, d, load_jsonl, is_head_series
from make_screenshots import screenshot_from_volume
LINE_RE = re.compile(
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:
m = LINE_RE.match(raw)
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),
(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:
continue
try:
if tag == (32, 50):
ipp = np.array([float(x) for x in val.split("\\")[:3]])
elif tag == (40, 48):
spacing = np.array([float(x) for x in val.split("\\")[:2]])
elif tag == (32, 55):
iop = np.array([float(x) for x in val.split("\\")[:6]])
elif tag == (40, 16):
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, study, series, proto
except OSError:
return None
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(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
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 = 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)
cands, txt_samples = [], []
try:
for f in os.listdir(jpg_dir):
m = pat.match(f)
if m:
cands.append((int(m.group(1)), int(m.group(2)), f))
for f in os.listdir(tp_dir):
m = tpat.match(f)
if m:
p = parse_txt(os.path.join(tp_dir, f))
if p is not None:
txt_samples.append((int(m.group(1)), p))
except OSError as e:
return row["key"], False, f"listdir fail {e!r}"
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)
(rows, cols, ps0, ps1) = geo.most_common(1)[0][0]
ref = [p for _, p in txt_samples if (p[3], p[4]) == (rows, cols) and np.allclose(p[1], (ps0, ps1), atol=1e-3)]
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, _, _, 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:
return row["key"], False, "fewer than 2 consistent geometry samples"
(s_lo, p_lo), (s_hi, p_hi) = good[0], good[-1]
step = (p_hi[0] - p_lo[0]) / max(s_hi - s_lo, 1)
resid = max(np.linalg.norm(p[0] - p_lo[0] - (s - s_lo) * step) for s, p in good)
if resid > 0.35 * np.linalg.norm(step):
return row["key"], False, f"non-linear slice positions resid={resid:.3f}"
u, v = iop0[:3], iop0[3:]
n = np.cross(u, v)
# 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):
return p_lo[0] + (s - s_lo) * step
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, 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):
return row["key"], False, "bad slice spacing"
nz = int(round((zs.max() - zs.min()) / dmed)) + 1
ro_min = min(p[0] for p in pos)
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 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))
try:
im = np.asarray(Image.open(os.path.join(jpg_dir, jf)).convert("L"), dtype=np.uint8)
except Exception:
continue
if im.shape[:2] != (rows, cols):
if im.shape[:2] == (cols, rows):
im = im.T.copy()
else:
continue
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 = ((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 + (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[1]), float(ps[0]), float(dmed)))
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/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
else " (no screenshot: not body-aligned)")
return row["key"], True, msg
def run_pool(rows, extent_min, workers, reasons=None):
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 reasons is not None:
reasons[k] = msg
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 write_excluded_notes(zero_rows, all_rows, thick_rows, dyn_excl, reasons):
"""Write data/qa/<sid>/<date>_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(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_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_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_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, 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, 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__":
main()