longitudinal/scripts_monai/05_monai_run_iterative.py
Furen Xiao 77adc2b3af feat(monai): add MONAI pipeline C implementation
Introduce the third segmentation pipeline using MONAI (1.6.x) to allow
direct comparison with Pipelines A and B. This includes the implementation
of the iterative pseudo-labeling workflow, training scripts, and
inference protocols.

- Add `scripts_monai/` directory containing the MONAI pipeline scripts.
- Update documentation in `README.md` and `AGENTS.md` to include MONAI
  package requirements and pipeline details.
- Configure `.gitignore` to exclude MONAI-specific run directories.
- Update data directory descriptions to include MONAI pseudo-labels.
2026-09-26 11:04:20 +08:00

180 lines
No EOL
8.1 KiB
Python

"""Orchestrates the MONAI iterative pseudo-labeling study (parallel to
scripts/09 and scripts_nnu/06).
Round 0: train the MONAI UNet on the labeled patient-level train split (internal
validation = held-out patient-level val split, subject-disjoint).
Round k: dataset grows with accepted pseudo-labels (pos + neg) from rounds 1..k
(de-duplicated by key, pseudo rows weighted); warm-start training
(full weights) from round k-1's best checkpoint at a lower LR; then
pseudo-label the remaining unlabeled pool (gates identical to
Pipelines A/B) and evaluate the model on the held-out test split.
Usage: python scripts_monai/05_monai_run_iterative.py [--rounds 4] [--gpus 3]
"""
import os
import sys
import json
import shutil
import argparse
import subprocess
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from src.common import ROOT, d, load_jsonl, save_jsonl
S = os.path.dirname(os.path.abspath(__file__)) # code lives next to this file, not under ROOT
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--rounds", type=int, default=4)
ap.add_argument("--gpus", type=int, default=3)
ap.add_argument("--base-epochs", type=int, default=100)
ap.add_argument("--base-lr", type=float, default=3e-4)
ap.add_argument("--round-epochs", type=int, default=30)
ap.add_argument("--round-lr", type=float, default=1e-3)
ap.add_argument("--pseudo-weight", type=float, default=0.3)
ap.add_argument("--batch", type=int, default=3)
ap.add_argument("--val-every", type=int, default=10)
ap.add_argument("--val-limit", type=int, default=60)
ap.add_argument("--no-neg-pseudo", action="store_true", help="do not add negative pseudo cases to training")
ap.add_argument("--no-tta", action="store_true", help="disable flip TTA at inference")
ap.add_argument("--tau-pos", type=float, default=0.95)
ap.add_argument("--tau-neg", type=float, default=0.98)
ap.add_argument("--neg-frac", type=float, default=0.90)
ap.add_argument("--vol-qp", type=float, nargs=2, default=[2, 98])
ap.add_argument("--min-cc-frac", type=float, default=0.2)
ap.add_argument("--max-rel-dist", type=float, default=40.0)
ap.add_argument("--vol-ratio", type=float, default=10.0)
args = ap.parse_args()
man = os.path.join(ROOT, "data/manifests")
train_f = os.path.join(man, "split_train.jsonl")
val_f = os.path.join(man, "split_val.jsonl")
test_f = os.path.join(man, "split_test.jsonl")
pool_f = os.path.join(man, "unlabeled_pool.jsonl")
results_dir = d("results")
logdir = d("logs/monai")
main_log = os.path.join(logdir, "iterative.log")
pdir_root = d("data/pseudo_monai")
added_f = os.path.join(pdir_root, "added_keys.jsonl")
def run(cmd, log, retries=2):
for attempt in range(retries + 1):
try:
with open(log, "a") as f:
f.write(f"$ (attempt {attempt + 1}) " + cmd + "\n")
f.flush()
print(f"$ (attempt {attempt + 1}) " + cmd, flush=True)
subprocess.run(cmd, shell=True, cwd=ROOT, stdout=f, stderr=subprocess.STDOUT, check=True)
return
except subprocess.CalledProcessError:
if attempt == retries:
raise
print(f"[monai-orch] command failed, retrying in 60s: {cmd}", flush=True)
import time
time.sleep(60)
def tr(script, extra):
t = shutil.which("torchrun") or f"{sys.executable} -m torch.distributed.run"
return f"{t} --standalone --nproc_per_node {args.gpus} {os.path.join(S, script)} {extra}"
def py(script, extra, log=main_log):
run(f"{sys.executable} {os.path.join(S, script)} {extra}", log)
def trrun(script, extra, log):
run(tr(script, extra), log)
def build_rows(out, accepted, no_neg=False):
extra = f"--train {train_f} --pseudo-weight {args.pseudo_weight} --out {out}"
if no_neg:
extra += " --no-neg"
for f in accepted:
extra += f" --accepted {f}"
py("01_monai_build_rows.py", extra)
def train(round, rows_f, epochs, lr, warmstart, val_every):
extra = (f"--rows {rows_f} --val {val_f} --epochs {epochs} --lr {lr} "
f"--batch {args.batch} --ckpt-dir runs_monai/round{round} "
f"--val-every {val_every} --val-limit {args.val_limit}")
if warmstart:
extra += f" --pretrained {warmstart}"
trrun("02_monai_train.py", extra, os.path.join(logdir, f"train_r{round}.log"))
def predict_pool(round, ckpt, added=None):
out = os.path.join(pdir_root, f"round{round}")
extra = (f"--ckpt {ckpt} --pool {pool_f} --out {out} "
f"--tau-pos {args.tau_pos} --tau-neg {args.tau_neg} --neg-frac {args.neg_frac} "
f"--vol-qp {args.vol_qp[0]} {args.vol_qp[1]} --min-cc-frac {args.min_cc_frac} "
f"--max-rel-dist {args.max_rel_dist} --vol-ratio {args.vol_ratio}")
if added and os.path.exists(added):
extra += f" --already {added}"
if args.no_tta:
extra += " --no-tta"
trrun("03_monai_pseudo_label.py", extra, os.path.join(logdir, f"pseudo_r{round}.log"))
return os.path.join(out, "accepted.jsonl")
def eval_test(round, ckpt):
out = os.path.join(results_dir, f"round{round}_test_monai.json")
extra = f"--rows {test_f} --ckpt {ckpt} --out {out}"
if args.no_tta:
extra += " --no-tta"
trrun("04_monai_eval_test.py", extra, os.path.join(logdir, f"eval_r{round}.log"))
return out
def best(round):
return os.path.join(ROOT, "runs_monai", f"round{round}", "best.pt")
# ---- round 0: baseline ----
rows0 = os.path.join(pdir_root, "round0_rows.jsonl")
build_rows(rows0, [])
train(0, rows0, args.base_epochs, args.base_lr, None, args.val_every)
accepted_f = [predict_pool(1, best(0))]
save_jsonl(load_jsonl(accepted_f[0]), added_f)
eval_test(0, best(0))
for k in range(1, args.rounds + 1):
rows_f = os.path.join(pdir_root, f"round{k}_rows.jsonl")
build_rows(rows_f, accepted_f, no_neg=args.no_neg_pseudo)
train(k, rows_f, args.round_epochs, args.round_lr, best(k - 1), args.val_every)
accepted_f.append(predict_pool(k + 1, best(k), added_f))
save_jsonl([r for f in accepted_f for r in load_jsonl(f)], added_f)
eval_test(k, best(k))
# ---- report (same layout as scripts/09 and scripts_nnu/06) ----
table = []
for k in range(args.rounds + 1):
resf = os.path.join(results_dir, f"round{k}_test_monai.json")
if not os.path.exists(resf):
continue
r = json.load(open(resf))
row = {"round": k, "test_dice": round(r["dice"], 4), "n_test": r["n"], "ckpt": best(k)}
pf = os.path.join(pdir_root, f"round{k}", "summary.json")
if k > 0 and os.path.exists(pf):
s = json.load(open(pf))
row.update({"n_pos": s["n_pos"], "n_neg": s["n_neg"],
"n_rej_cons": s["n_rejected_consistency"], "n_error": s["n_error"],
"pos_vol_med_mm3": s["pos_vol_mm3"]["med"]})
table.append(row)
save_jsonl(table, os.path.join(results_dir, "iterative_table_monai.jsonl"))
print(json.dumps(table, indent=1))
try:
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
rs = [t["round"] for t in table]
ds = [t["test_dice"] for t in table]
plt.figure(figsize=(6, 4))
plt.plot(rs, ds, "o-")
plt.xlabel("pseudo-labeling round")
plt.ylabel("holdout tumor Dice (MONAI)")
for x, y in zip(rs, ds):
plt.annotate(f"{y:.3f}", (x, y), textcoords="offset points", xytext=(0, 8), fontsize=8)
plt.grid(alpha=0.3)
plt.tight_layout()
plt.savefig(os.path.join(results_dir, "iterative_dice_monai.png"), dpi=150)
except Exception as e: # noqa
print("plot failed:", repr(e))
if __name__ == "__main__":
main()