"""Build M6-2025 manifests: T1c volumes, labeled when GTV exists. GTV lives in CT space; CT is registered to T1c (similarity->affine MI + BSpline) and GTV warped (nearest) into T1c native space. Outputs: data/manifests/m6_labeled.jsonl {key, subject, date, img, label(warped GTV), source:'m6'} data/manifests/m6_unlabeled.jsonl {key, subject, date, img, label:None, source:'m6'} Usage: python scripts/02_build_m6_dataset.py [--workers 4] [--skip-reg] """ import os import sys sys.path.insert(0, os.path.dirname(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 src.common import ROOT, d, save_jsonl, read_nii_img BASE = "/mnt/pve/WORKSPACE/M6-2025/nii" 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) def pick_t1c(mrd): cands = [] for f in sorted(os.listdir(mrd)): if not f.endswith(".nii.gz"): continue b = f[: -len(".nii.gz")] if T1C_RE.search(b) and not EXCL_RE.search(b): pref = 0 if "MPR_Tra" in b: pref = -2 if "_c" in b.lower() or "+C" in b: pref -= 1 cands.append((pref, f)) cands.sort(key=lambda x: (x[0], x[1])) return os.path.join(mrd, cands[0][1]) if cands else None def discover(): labeled, unlabeled = [], [] for pid in sorted(os.listdir(BASE)): fp = os.path.join(BASE, pid) if not os.path.isdir(fp): continue for date in sorted(os.listdir(fp)): dp = os.path.join(fp, date) if not os.path.isdir(dp) or not re.match(r"\d{8}$", date): continue mrd = os.path.join(dp, "MR") if not os.path.isdir(mrd): continue t1c = pick_t1c(mrd) if not t1c: continue gtv = os.path.join(dp, "RT", "TV", "Struct_GTV.nii.gz") ct = os.path.join(dp, "RT", "ct_image.nii.gz") key = f"m6_{pid}_{date}" row = {"key": key, "subject": pid, "date": date, "img": t1c, "source": "m6"} if os.path.exists(gtv) and os.path.exists(ct): labeled.append({**row, "label": gtv, "ct": ct, "gtv_registered": False}) else: unlabeled.append({**row, "label": None}) return labeled, unlabeled def register_ct_to_t1c(ct, t1c, threads=12): import time t_start = time.time() try: # noqa sitk.CommonProperties.SetGlobalDefaultNumberOfThreads(threads) except AttributeError: pass ct = sitk.Cast(ct, sitk.sitkFloat32) t1c = sitk.Cast(t1c, sitk.sitkFloat32) r = sitk.ImageRegistrationMethod() r.SetMetricAsMattesMutualInformation(numberOfHistogramBins=32) r.SetOptimizerAsRegularStepGradientDescent( learningRate=10.0, minStep=0.5, numberOfIterations=600, relaxationFactor=0.5, gradientMagnitudeTolerance=1e-5, estimateLearningRate=sitk.ImageRegistrationMethod.EachIteration) r.SetOptimizerScalesFromPhysicalShift() r.SetInitialTransform(sitk.CenteredTransformInitializer(ct, t1c, sitk.Transform(3, sitk.sitkSimilarity))) r.SetShrinkFactorsPerLevel([8, 4, 2, 1]) r.SetSmoothingSigmasPerLevel([4, 2, 1, 0]) r.SetInterpolator(sitk.sitkLinear) aff = r.Execute(ct, t1c) print("affine done %.0fs" % (time.time() - t_start), flush=True) return aff def reg_worker(row): reg_dir = d("data/m6reg") key = row["key"] t1c = read_nii_img(row["img"]) ct = read_nii_img(row["ct"]) gtv = read_nii_img(row["label"]) tr = register_ct_to_t1c(ct, t1c) out = os.path.join(reg_dir, key + "_gtv.nii.gz") warped = sitk.Resample(gtv, t1c, tr, sitk.sitkNearestNeighbor, 0.0) sitk.WriteImage(warped, out, True) z = sitk.GetArrayFromImage(warped) return key, os.path.exists(out) and (z > 0.5).sum() > 0, int((z > 0.5).sum()) def main(): ap = argparse.ArgumentParser() ap.add_argument("--workers", type=int, default=2) ap.add_argument("--skip-reg", action="store_true") args = ap.parse_args() labeled, unlabeled = discover() print(f"m6 labeled candidates: {len(labeled)}, unlabeled: {len(unlabeled)}") reg_dir = d("data/m6reg") if not args.skip_reg: todo = [r for r in labeled if not os.path.exists(os.path.join(reg_dir, r["key"] + "_gtv.nii.gz"))] print(f"registering {len(todo)} GTVs with {args.workers} workers") with ProcessPoolExecutor(max_workers=args.workers) as ex: futs = {ex.submit(reg_worker, r): r for r in todo} for fu in as_completed(futs): k = futs[fu]["key"] try: key, ok, nvox = fu.result() print(f" {key}: ok={ok} nvox={nvox}", flush=True) except Exception as e: # noqa print(f" {k}: ERR {e}", flush=True) # refresh manifest state labeled, unlabeled = discover() for r in labeled: p = os.path.join(reg_dir, r["key"] + "_gtv.nii.gz") if os.path.exists(p): r["label"] = p else: r["label"] = None keep = [r for r in labeled if r.get("label")] 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")) print(f"final labeled={len(keep)} unlabeled={len(unlabeled)}") if __name__ == "__main__": main()