longitudinal/scripts/05_build_splits.py
Furen Xiao 8fe1818c7d 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.
2026-09-27 07:14:17 +08:00

128 lines
No EOL
5.3 KiB
Python

"""Build patient-level train/val/test splits + unlabeled pool + tumor volume stats.
Inputs: data/manifests/{ntuh,m6_labeled,m6_unlabeled,lee_t1c_selected}.jsonl + data/proc/*
Outputs:
data/manifests/split_train.jsonl split_val.jsonl split_test.jsonl (labeled rows, w=1.0)
data/manifests/unlabeled_pool.jsonl (m6 + lee, labeled subjects removed)
data/vols.json (labeled tumor volume stats, mm3)
Each row carries both processed fields (pimg/plabel, data/proc — 1mm cropped/
normalized, used by Pipelines A and C) and native fields (img/label, the raw
source NIfTIs — used by Pipeline B, which feeds nnU-Net its own
plan_and_preprocess). Row membership still requires the processed volume to
exist, so the patient-level splits stay identical across pipelines.
"""
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import json
import argparse
import random
import numpy as np
from src.common import ROOT, d, load_jsonl, save_jsonl, read_nii_arr
def proc_row(r, prefix):
key = r["key"]
p = os.path.join(d("data/proc"), key + ".nii.gz")
if not os.path.exists(p):
return None
row = {"key": key, "subject": f"{prefix}_{r['subject']}", "pimg": p,
"img": r["img"], "label": r.get("label"), "date": r.get("date")}
lab = os.path.join(d("data/proc"), key + "_label.nii.gz")
if r.get("label") and os.path.exists(lab):
row["plabel"] = lab
return row
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--frac-val", type=float, default=0.05)
ap.add_argument("--frac-test", type=float, default=0.10)
ap.add_argument("--seed", type=int, default=0)
args = ap.parse_args()
random.seed(args.seed)
labeled_rows = []
for name, prefix in (("ntuh", "ntuh"), ("m6_labeled", "m6")):
for r in load_jsonl(os.path.join(ROOT, "data", "manifests", name + ".jsonl")):
if not r.get("label"):
continue
p = proc_row(r, prefix)
if p:
labeled_rows.append(p)
print(f"labeled processed rows: {len(labeled_rows)}")
subjects = sorted({r["subject"] for r in labeled_rows})
random.shuffle(subjects)
n_test = max(1, int(len(subjects) * args.frac_test))
n_val = max(1, int(len(subjects) * args.frac_val))
test_subj = set(subjects[:n_test])
val_subj = set(subjects[n_test:n_test + n_val])
train_subj = set(subjects[n_test + n_val:])
print(f"subjects: {len(subjects)} train={len(train_subj)} val={len(val_subj)} test={len(test_subj)}")
splits = {"train": [], "val": [], "test": []}
for r in labeled_rows:
r["w"] = 1.0
if r["subject"] in test_subj:
splits["test"].append(r)
elif r["subject"] in val_subj:
splits["val"].append(r)
else:
splits["train"].append(r)
for k, v in splits.items():
save_jsonl(v, os.path.join(ROOT, "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")):
if f"m6_{r['subject']}" in labeled_subjects:
continue
p = proc_row(r, "m6")
if p:
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_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["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,
"img": nii, "label": None, "date": r["date"], "source": "lee"}
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"))
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")
# tumor volume stats (1mm voxels == mm3)
vols = []
for r in labeled_rows:
try:
l = read_nii_arr(r["plabel"])
n = int((l > 0.5).sum())
vols.append({"key": r["key"], "subject": r["subject"], "vol_mm3": n})
except Exception:
pass
v = np.array([x["vol_mm3"] for x in vols])
stats = {"n": int(v.size),
"p2": float(np.percentile(v, 2)), "p5": float(np.percentile(v, 5)),
"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:
json.dump(stats, f, indent=1)
print("tumor volumes (mm3):", stats)
if __name__ == "__main__":
main()