longitudinal/scripts/02_build_m6_dataset.py
Furen Xiao b6fa62a763 feat: initial project structure
Add .gitignore, AGENTS.md, scripts directory, and src directory to initialize the repository.
2026-09-25 16:00:37 +08:00

147 lines
No EOL
5.6 KiB
Python

"""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()