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.
156 lines
No EOL
7.2 KiB
Python
156 lines
No EOL
7.2 KiB
Python
"""MONAI DDP training entrypoint for one pseudo-labeling round (torchrun).
|
|
|
|
torchrun --standalone --nproc_per_node 3 scripts_monai/02_monai_train.py \
|
|
--rows data/pseudo_monai/round0_rows.jsonl --val data/manifests/split_val.jsonl \
|
|
--epochs 100 --lr 3e-4 --batch 3 --ckpt-dir runs_monai/round0 \
|
|
[--pretrained runs_monai/round0/best.pt]
|
|
|
|
- backbone: monai.networks.nets.UNet (monai_common.build_model)
|
|
- data: MONAI Dataset + transform chain (load → pad → flip/rotate90/intensity →
|
|
foreground-aware 96³ crop → typed tensors); row "w" carries the per-sample
|
|
loss weight (pseudo rows down-weighted)
|
|
- loss: per-sample Dice+CE, weighted batch mean (Pipeline A's convention)
|
|
- schedule: AdamW + linear warmup + cosine (Pipeline A's schedule)
|
|
- checkpoints in --ckpt-dir: best.pt (max val dice), final.pt, state.pt
|
|
(auto-resume after a crash, as in Pipeline A)
|
|
"""
|
|
import argparse
|
|
import os
|
|
import sys
|
|
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__))))
|
|
|
|
import numpy as np
|
|
import torch
|
|
from src.common import load_jsonl, read_nii_arr
|
|
from monai_common import (build_model, WeightedDiceCELoss, make_dataloader,
|
|
predict_probs, prob_dice, rank_info, init_dist, barrier, destroy_dist)
|
|
|
|
|
|
def main():
|
|
ap = argparse.ArgumentParser()
|
|
ap.add_argument("--rows", required=True, help="jsonl of {key,pimg,plabel,w}")
|
|
ap.add_argument("--val", required=True, help="held-out val split jsonl ({pimg,plabel})")
|
|
ap.add_argument("--epochs", type=int, default=100)
|
|
ap.add_argument("--lr", type=float, default=3e-4)
|
|
ap.add_argument("--batch", type=int, default=3, help="per-GPU batch")
|
|
ap.add_argument("--patch", type=int, default=96)
|
|
ap.add_argument("--workers", type=int, default=4)
|
|
ap.add_argument("--ckpt-dir", required=True)
|
|
ap.add_argument("--pretrained", default=None, help="full-weight warm-start checkpoint")
|
|
ap.add_argument("--val-every", type=int, default=10)
|
|
ap.add_argument("--val-limit", type=int, default=60)
|
|
ap.add_argument("--sw-batch", type=int, default=8, help="sliding-window batch at inference")
|
|
args = ap.parse_args()
|
|
|
|
rank, world, local_rank = rank_info()
|
|
init_dist()
|
|
torch.manual_seed(0)
|
|
np.random.seed(0)
|
|
torch.cuda.set_device(local_rank)
|
|
device = f"cuda:{local_rank}"
|
|
|
|
rows = load_jsonl(args.rows)
|
|
val_rows = load_jsonl(args.val)[:args.val_limit]
|
|
|
|
dl = make_dataloader(rows, win=args.patch, batch=args.batch, workers=args.workers)
|
|
steps_per_epoch = max(len(dl), 1)
|
|
|
|
model = build_model(device=device)
|
|
if args.pretrained:
|
|
sd0 = torch.load(args.pretrained, map_location=device, weights_only=True)
|
|
model.load_state_dict(sd0.get("model", sd0))
|
|
if rank == 0:
|
|
print(f"[monai:train:rank0] warm start (all weights) from {args.pretrained}", flush=True)
|
|
ddp = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank]) if world > 1 else model
|
|
total_steps = steps_per_epoch * args.epochs
|
|
|
|
opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=1e-4)
|
|
warmup = min(300, max(10, total_steps // 10))
|
|
base_sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=max(total_steps - warmup, 1),
|
|
eta_min=args.lr * 0.05)
|
|
sched = torch.optim.lr_scheduler.SequentialLR(
|
|
opt, [torch.optim.lr_scheduler.LinearLR(opt, start_factor=0.1, total_iters=warmup), base_sched],
|
|
milestones=[warmup])
|
|
loss_f = WeightedDiceCELoss()
|
|
|
|
best, best_epoch, start_epoch = -1.0, -1, 0
|
|
if rank == 0:
|
|
os.makedirs(args.ckpt_dir, exist_ok=True)
|
|
state_f = os.path.join(args.ckpt_dir, "state.pt")
|
|
if os.path.exists(state_f):
|
|
st = torch.load(state_f, map_location="cpu", weights_only=True)
|
|
model.load_state_dict(st["model"])
|
|
opt.load_state_dict(st["opt"])
|
|
sched.load_state_dict(st["sched"])
|
|
best, best_epoch, start_epoch = st["best"], st["best_epoch"], st["epoch"]
|
|
if rank == 0:
|
|
print(f"[monai:train] auto-resuming from state.pt @epoch {start_epoch} (best={best:.4f})", flush=True)
|
|
if rank == 0:
|
|
print(f"[monai:train] rows={len(rows)} val={len(val_rows)} epochs={args.epochs} "
|
|
f"steps/epoch={steps_per_epoch} world={world}", flush=True)
|
|
barrier()
|
|
|
|
def save_state():
|
|
tmp = state_f + ".tmp"
|
|
torch.save({"model": model.state_dict(), "opt": opt.state_dict(), "sched": sched.state_dict(),
|
|
"best": best, "best_epoch": best_epoch, "epoch": epoch + 1}, tmp)
|
|
os.replace(tmp, state_f)
|
|
|
|
for epoch in range(start_epoch, args.epochs):
|
|
if hasattr(dl.sampler, "set_epoch"):
|
|
dl.sampler.set_epoch(epoch)
|
|
model.train()
|
|
run_loss, run_n = 0.0, 0
|
|
for batch in dl:
|
|
img = batch["pimg"].to(device, non_blocking=True)
|
|
lab = batch["plabel"].squeeze(1).to(device, non_blocking=True)
|
|
wts = batch["w"].to(device, non_blocking=True)
|
|
with torch.autocast("cuda", dtype=torch.bfloat16):
|
|
logits = ddp(img)
|
|
loss = loss_f(logits, lab, wts)
|
|
opt.zero_grad(set_to_none=True)
|
|
loss.backward()
|
|
torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)
|
|
opt.step()
|
|
sched.step()
|
|
run_loss += float(loss.detach())
|
|
run_n += 1
|
|
if rank == 0:
|
|
print(f" epoch {epoch+1}/{args.epochs} loss={run_loss/max(run_n,1):.4f} "
|
|
f"lr={opt.param_groups[0]['lr']:.2e}", flush=True)
|
|
if rank == 0 and (epoch + 1) % max(1, args.val_every // 2) == 0:
|
|
save_state()
|
|
if val_rows and rank == 0 and (epoch + 1) % args.val_every == 0:
|
|
dv, n_ok = 0.0, 0
|
|
for r in val_rows:
|
|
try:
|
|
vol = read_nii_arr(r["pimg"]).astype("float32")
|
|
labv = read_nii_arr(r["plabel"]).astype("uint8")
|
|
p = predict_probs(model, vol, device, args.patch, tta=True, sw_batch=args.sw_batch)
|
|
if p.shape == labv.shape:
|
|
dv += prob_dice(p, labv)
|
|
n_ok += 1
|
|
except Exception as e: # noqa
|
|
print(" val err", r.get("key"), repr(e))
|
|
dv = dv / max(n_ok, 1)
|
|
print(f" [val] epoch {epoch+1} dice={dv:.4f}", flush=True)
|
|
if dv > best:
|
|
best, best_epoch = dv, epoch + 1
|
|
torch.save({"model": model.state_dict(), "epoch": epoch + 1, "val_dice": best},
|
|
os.path.join(args.ckpt_dir, "best.pt"))
|
|
save_state()
|
|
barrier()
|
|
if rank == 0:
|
|
torch.save({"model": model.state_dict(), "epoch": args.epochs, "val_dice": best},
|
|
os.path.join(args.ckpt_dir, "final.pt"))
|
|
best_f = os.path.join(args.ckpt_dir, "best.pt")
|
|
if not os.path.exists(best_f):
|
|
torch.save({"model": model.state_dict(), "epoch": args.epochs, "val_dice": 0.0}, best_f)
|
|
print("[monai:train] no val run; best.pt = final weights", flush=True)
|
|
print(f"[monai:train] done best_val_dice={best:.4f}@{best_epoch}", flush=True)
|
|
destroy_dist()
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main() |