Add .gitignore, AGENTS.md, scripts directory, and src directory to initialize the repository.
147 lines
No EOL
5.6 KiB
Python
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() |