SACFlow / sacflow /engine /train_loop.py
sathiiii's picture
Add SACFlow source code
d65ae7d verified
Raw
History Blame Contribute Delete
20.6 kB
from __future__ import annotations
from pathlib import Path
import time
import copy
import re
import torch
import torch.nn.functional as F
from torch.nn.parallel import DistributedDataParallel as DDP
from tqdm import tqdm
from monai.inferers import sliding_window_inference
from sacflow.data.loader import build_loader
from sacflow.models.unet3d import build_model, freeze_except_adapters
from sacflow.models.velocity_field import VelocityField3D
from sacflow.methods.sacflow_step import sacflow_forward_step, ce_loss_masked, dice_loss_masked
from sacflow.methods.source_memory import load_source_memory, class_moments
from sacflow.utils.metrics import torch_soft_dice_loss, entropy_loss, confidence_and_margin, dice_per_class, hd95_per_class
from sacflow.utils.misc import ensure_dir, count_trainable, move_to_device, unwrap_model
from sacflow.utils.distributed import is_main_process, get_world_size, get_rank, reduce_mean, barrier, is_dist_avail_and_initialized
import torch.distributed as dist
from sacflow.utils.wandb_utils import wandb_log
def build_optimizer(params, cfg):
ocfg = cfg["optim"]
params = [p for p in params if p.requires_grad]
if ocfg.get("optimizer", "adamw").lower() == "sgd":
return torch.optim.SGD(params, lr=float(ocfg["lr"]), momentum=0.9, weight_decay=float(ocfg.get("weight_decay", 0)))
return torch.optim.AdamW(params, lr=float(ocfg["lr"]), weight_decay=float(ocfg.get("weight_decay", 0)), betas=tuple(ocfg.get("betas", [0.9, 0.999])))
def update_ema(teacher, student, decay):
with torch.no_grad():
for pt, ps in zip(teacher.parameters(), student.parameters()):
pt.data.mul_(decay).add_(ps.data, alpha=1-decay)
def load_checkpoint_into(model, path, strict=False):
ckpt = torch.load(path, map_location="cpu")
state = ckpt.get("model", ckpt)
missing, unexpected = model.load_state_dict(state, strict=strict)
return missing, unexpected
def save_checkpoint(path, model, optimizer, epoch, step, best_metric=None, velocity_field=None, cfg=None, teacher=None, include_optimizer=True):
"""Save a checkpoint on rank 0 only.
Disk policy:
- best.pt is intended for evaluation/inference and is saved without optimizer by default.
- last.pt is intended for resume and includes optimizer.
This avoids filling the disk with epoch_N.pt checkpoints.
"""
if not is_main_process():
return
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
obj = {
"model": unwrap_model(model).state_dict(),
"epoch": epoch,
"step": step,
"best_metric": best_metric,
"cfg": cfg,
}
if include_optimizer and optimizer is not None:
obj["optimizer"] = optimizer.state_dict()
if velocity_field is not None:
obj["velocity_field"] = unwrap_model(velocity_field).state_dict()
if teacher is not None:
obj["teacher"] = unwrap_model(teacher).state_dict()
torch.save(obj, path)
def _epoch_number(path: Path) -> int:
m = re.search(r"epoch_(\d+)\.pt$", path.name)
return int(m.group(1)) if m else -1
def resolve_resume_checkpoint(ckpt_dir: Path, resume_value):
"""Return a usable resume checkpoint path.
resume_value can be:
- None/False: do not resume
- "auto"/True: prefer last.pt, then newest epoch_*.pt, then best.pt
- explicit checkpoint path
Corrupted/incomplete checkpoints are skipped.
"""
if not resume_value:
return None
if str(resume_value).lower() not in ("auto", "true", "1", "yes"):
return Path(resume_value)
candidates = []
last = ckpt_dir / "last.pt"
if last.exists():
candidates.append(last)
candidates.extend(sorted(ckpt_dir.glob("epoch_*.pt"), key=_epoch_number, reverse=True))
best = ckpt_dir / "best.pt"
if best.exists():
candidates.append(best)
for c in candidates:
try:
torch.load(c, map_location="cpu")
return c
except Exception as e:
if is_main_process():
print(f"Skipping unusable checkpoint {c}: {e}")
return None
def load_training_checkpoint(path, model, optimizer=None, velocity_field=None, teacher=None):
ckpt = torch.load(path, map_location="cpu")
missing, unexpected = unwrap_model(model).load_state_dict(ckpt.get("model", ckpt), strict=False)
if is_main_process():
print(f"Loaded resume model from {path} missing={len(missing)} unexpected={len(unexpected)}")
if velocity_field is not None and "velocity_field" in ckpt:
unwrap_model(velocity_field).load_state_dict(ckpt["velocity_field"], strict=False)
if teacher is not None and "teacher" in ckpt:
unwrap_model(teacher).load_state_dict(ckpt["teacher"], strict=False)
elif teacher is not None:
unwrap_model(teacher).load_state_dict(unwrap_model(model).state_dict(), strict=False)
if optimizer is not None and ckpt.get("optimizer") is not None:
optimizer.load_state_dict(ckpt["optimizer"])
start_epoch = int(ckpt.get("epoch", 0))
global_step = int(ckpt.get("step", 0))
best = float(ckpt.get("best_metric", -1e9) if ckpt.get("best_metric", None) is not None else -1e9)
return start_epoch, global_step, best
def supervised_step(model, batch, cfg):
x = batch["image"]
y = batch["label"]
logits = model(x)
ce = F.cross_entropy(logits, y.long())
dice = torch_soft_dice_loss(logits, y, cfg["data"]["num_classes"])
loss = cfg["train"].get("loss", {}).get("ce", 1.0)*ce + cfg["train"].get("loss", {}).get("dice", 1.0)*dice
return loss, {"loss_total": loss.detach(), "loss_ce": ce.detach(), "loss_dice": dice.detach()}
def proto_align_step(model, teacher, batch, memory, cfg):
x = batch["image"]
logits, feats = model(x, return_features=True)
feat = feats["prelogit"]
with torch.no_grad():
tlogits = teacher(x)
tprobs = torch.softmax(tlogits, dim=1)
conf, margin, pseudo = confidence_and_margin(tprobs)
mask = conf > float(cfg["train"].get("pseudo_conf_threshold", 0.75))
probs_f = tprobs
if probs_f.shape[-3:] != feat.shape[-3:]:
probs_f = F.interpolate(probs_f, size=feat.shape[-3:], mode="trilinear", align_corners=False)
ce = ce_loss_masked(logits, pseudo, mask)
dice = dice_loss_masked(logits, pseudo, mask, cfg["data"]["num_classes"])
proto_loss = torch.tensor(0.0, device=x.device)
if memory is not None and "feature_mu" in memory:
mu = memory["feature_mu"].to(feat.device, feat.dtype) # [C,d]
# expected source prototype at each voxel based on teacher probabilities
proto = torch.einsum("bchwz,cf->bfhwz", probs_f.detach(), mu)
proto_loss = ((feat - proto).pow(2) * probs_f.max(1, keepdim=True).values.detach()).mean()
ent = entropy_loss(logits)
loss_cfg = cfg["train"].get("loss", {})
loss = float(loss_cfg.get("ce", 1.0))*ce + float(loss_cfg.get("dice", 1.0))*dice + float(loss_cfg.get("prototype", 0.1))*proto_loss + float(loss_cfg.get("entropy", 0.01))*ent
return loss, {"loss_total": loss.detach(), "loss_pseudo_ce": ce.detach(), "loss_pseudo_dice": dice.detach(), "loss_proto_align": proto_loss.detach(), "loss_entropy": ent.detach(), "pseudo_conf_mean": conf.mean().detach(), "pseudo_accept_rate": mask.float().mean().detach()}
def pseudo_step(model, teacher, batch, cfg):
x = batch["image"]
with torch.no_grad():
tlogits = teacher(x)
tprobs = torch.softmax(tlogits, dim=1)
conf, margin, pseudo = confidence_and_margin(tprobs)
mask = conf > float(cfg["train"].get("pseudo_conf_threshold", 0.75))
logits = model(x)
ce = ce_loss_masked(logits, pseudo, mask)
dice = dice_loss_masked(logits, pseudo, mask, cfg["data"]["num_classes"])
ent = entropy_loss(logits)
loss_cfg = cfg["train"].get("loss", {})
loss = float(loss_cfg.get("ce", 1.0))*ce + float(loss_cfg.get("dice", 1.0))*dice + float(loss_cfg.get("entropy", 0.01))*ent
return loss, {"loss_total": loss.detach(), "loss_pseudo_ce": ce.detach(), "loss_pseudo_dice": dice.detach(), "loss_entropy": ent.detach(), "pseudo_conf_mean": conf.mean().detach(), "pseudo_accept_rate": mask.float().mean().detach()}
@torch.no_grad()
def evaluate(model, loader, cfg, device, max_batches=None):
"""Evaluate segmentation metrics.
In DDP this function is called on *all* ranks with a no-padding sharded
validation loader. It then all-reduces metric sums/counts so rank 0 gets
exact full-validation metrics without other ranks idling at a barrier.
"""
model.eval()
all_metrics = []
num_classes = cfg["data"]["num_classes"]
roi_size = tuple(cfg.get("eval", {}).get("roi_size", cfg["data"].get("patch_size", [96,96,96])))
sw_batch_size = int(cfg.get("eval", {}).get("sw_batch_size", 1))
overlap = float(cfg.get("eval", {}).get("overlap", 0.5))
iterator = enumerate(loader)
if is_main_process():
iterator = tqdm(iterator, total=len(loader), desc="eval", leave=False)
for i, batch in iterator:
if max_batches is not None and i >= max_batches:
break
if "label" not in batch:
continue
x = batch["image"].to(device, non_blocking=True)
y = batch["label"].numpy()
if cfg.get("eval", {}).get("sliding_window", True):
logits = sliding_window_inference(x, roi_size=roi_size, sw_batch_size=sw_batch_size, predictor=model, overlap=overlap)
else:
logits = model(x)
pred = logits.argmax(1).cpu().numpy()
for b in range(pred.shape[0]):
m = {}
m.update(dice_per_class(pred[b], y[b], num_classes))
spacing = tuple(batch.get("spacing", torch.ones(1,3))[b].cpu().numpy().tolist()) if "spacing" in batch else (1,1,1)
m.update(hd95_per_class(pred[b], y[b], num_classes, spacing=spacing))
all_metrics.append(m)
metric_keys = [f"dice_c{c}" for c in range(1, num_classes)] + ["dice_mean"] + [f"hd95_c{c}" for c in range(1, num_classes)] + ["hd95_mean"]
sums = torch.zeros(len(metric_keys), device=device, dtype=torch.float64)
counts = torch.zeros(len(metric_keys), device=device, dtype=torch.float64)
for m in all_metrics:
for j, k in enumerate(metric_keys):
v = m.get(k, float("nan"))
if v == v: # not NaN
sums[j] += float(v)
counts[j] += 1.0
if is_dist_avail_and_initialized():
dist.all_reduce(sums, op=dist.ReduceOp.SUM)
dist.all_reduce(counts, op=dist.ReduceOp.SUM)
out = {}
for j, k in enumerate(metric_keys):
if counts[j].item() > 0:
out[f"val/{k}"] = float((sums[j] / counts[j]).item())
if not out:
out["val/dice_mean"] = float("nan")
return out
def run_training(cfg, device, wandb_run=None):
mode = cfg["train"]["mode"]
out_dir = ensure_dir(cfg["output_dir"])
ckpt_dir = ensure_dir(out_dir / "checkpoints")
require_label = mode in ("source_train", "oracle_train")
split = "source_train" if mode == "source_train" else ("target_train" if mode in ("oracle_train", "self_train", "peft", "sacflow_fm", "proto_align") else "target_train")
train_loader = build_loader(cfg, split=split, training=True, require_label=require_label)
val_split = "source_val" if mode == "source_train" else "target_val"
try:
val_loader = build_loader(cfg, split=val_split, training=False, require_label=True, distributed=(get_world_size() > 1))
except Exception:
val_loader = None
model = build_model(cfg).to(device)
if cfg["train"].get("source_checkpoint"):
missing, unexpected = load_checkpoint_into(model, cfg["train"]["source_checkpoint"], strict=False)
if is_main_process():
print("Loaded source checkpoint", cfg["train"]["source_checkpoint"], "missing", len(missing), "unexpected", len(unexpected))
if mode in ("peft", "sacflow_fm", "proto_align") and cfg.get("model", {}).get("adapter", {}).get("enabled", False):
freeze_except_adapters(model, train_norm_affine=True)
teacher = copy.deepcopy(model).to(device)
for p in teacher.parameters():
p.requires_grad = False
velocity_field = None
memory = None
if mode in ("sacflow_fm", "proto_align"):
if cfg["train"].get("memory_path"):
memory = load_source_memory(cfg["train"]["memory_path"], device=device)
feat_ch = model.prelogit_channels
vcfg = cfg.get("sacflow", {}).get("velocity", {})
if cfg.get("sacflow", {}).get("use_velocity_field", True):
velocity_field = VelocityField3D(
residual_channels=feat_ch,
num_classes=cfg["data"]["num_classes"],
hidden_ratio=float(vcfg.get("hidden_ratio", 0.25)),
depth=int(vcfg.get("depth", 2)),
tau_embedding_dim=int(vcfg.get("tau_embedding_dim", 32)),
organ_embedding_dim=int(vcfg.get("organ_embedding_dim", 16)),
include_teacher_probs=bool(vcfg.get("include_teacher_probs", True)),
include_confidence=bool(vcfg.get("include_confidence", True)),
include_boundary=bool(vcfg.get("include_boundary", True)),
use_depthwise=bool(vcfg.get("use_depthwise", True)),
use_group_norm=bool(vcfg.get("use_group_norm", True)),
use_film=bool(vcfg.get("use_film", True)),
).to(device)
params = list(model.parameters()) + ([] if velocity_field is None else list(velocity_field.parameters()))
optimizer = build_optimizer(params, cfg)
if get_world_size() > 1:
# SACFlow uses a custom feature-path loss in addition to the ordinary forward pass.
# find_unused_parameters=True is safer for this mode because some classifier outputs
# from the first forward are not directly used in the loss, while PEFT adapters are
# used again for path-state classification.
find_unused = bool(cfg.get("distributed", {}).get("find_unused_parameters", False)) or mode == "sacflow_fm"
model = DDP(model, device_ids=[device.index] if device.type == "cuda" else None, find_unused_parameters=find_unused)
if velocity_field is not None:
velocity_field = DDP(velocity_field, device_ids=[device.index] if device.type == "cuda" else None, find_unused_parameters=True)
trainable, total = count_trainable(unwrap_model(model))
if velocity_field is not None:
vt, vtotal = count_trainable(unwrap_model(velocity_field))
trainable += vt
total += vtotal
if is_main_process():
print(f"Mode={mode} trainable={trainable:,} total={total:,} ({100*trainable/max(1,total):.2f}%)")
scaler = torch.cuda.amp.GradScaler(enabled=bool(cfg.get("amp", True)) and device.type == "cuda")
best = -1e9
global_step = 0
start_epoch = 0
resume_value = cfg.get("train", {}).get("resume_checkpoint")
resume_path = resolve_resume_checkpoint(ckpt_dir, resume_value)
if resume_path is not None:
start_epoch, global_step, best = load_training_checkpoint(
resume_path, model, optimizer=optimizer, velocity_field=velocity_field, teacher=teacher
)
if is_main_process():
print(f"Resuming from epoch={start_epoch}, step={global_step}, best={best:.6f}")
elif resume_value and is_main_process():
print(f"WARNING: requested resume={resume_value!r}, but no usable checkpoint was found in {ckpt_dir}")
epochs = int(cfg["train"].get("epochs", 100))
steps_per_epoch = int(cfg["train"].get("steps_per_epoch", len(train_loader)))
if start_epoch >= epochs and is_main_process():
print(f"Checkpoint epoch {start_epoch} is already >= configured epochs {epochs}; nothing to train.")
for epoch in range(start_epoch, epochs):
if hasattr(train_loader.sampler, "set_epoch"):
train_loader.sampler.set_epoch(epoch)
model.train()
if velocity_field is not None:
velocity_field.train()
iterator = iter(train_loader)
pbar = range(steps_per_epoch)
if is_main_process():
pbar = tqdm(pbar, desc=f"epoch {epoch+1}/{epochs}", dynamic_ncols=True)
for _ in pbar:
try:
batch = next(iterator)
except StopIteration:
iterator = iter(train_loader)
batch = next(iterator)
batch = move_to_device(batch, device)
optimizer.zero_grad(set_to_none=True)
with torch.cuda.amp.autocast(enabled=bool(cfg.get("amp", True)) and device.type == "cuda"):
if mode in ("source_train", "oracle_train"):
loss, logs = supervised_step(model, batch, cfg)
elif mode in ("self_train", "peft"):
loss, logs = pseudo_step(model, teacher, batch, cfg)
elif mode == "proto_align":
loss, logs = proto_align_step(model, teacher, batch, memory, cfg)
elif mode == "sacflow_fm":
loss, logs = sacflow_forward_step(model, teacher, velocity_field, batch, memory, cfg)
else:
raise ValueError(f"Unknown train mode {mode}")
scaler.scale(loss).backward()
if float(cfg["optim"].get("grad_clip_norm", 0) or 0) > 0:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_([p for p in params if p.requires_grad], float(cfg["optim"].get("grad_clip_norm")))
scaler.step(optimizer)
scaler.update()
if mode in ("self_train", "peft", "sacflow_fm", "proto_align"):
update_ema(teacher, unwrap_model(model), float(cfg["train"].get("ema_decay", 0.995)))
global_step += 1
red_logs = {}
for k, v in logs.items():
if torch.is_tensor(v):
red_logs[f"train/{k}"] = float(reduce_mean(v.float()).item())
else:
red_logs[f"train/{k}"] = v
if is_main_process() and global_step % int(cfg["train"].get("log_every", 20)) == 0:
red_logs["train/epoch"] = epoch + 1
red_logs["train/lr"] = optimizer.param_groups[0]["lr"]
wandb_log(wandb_run, red_logs, step=global_step)
if hasattr(pbar, "set_postfix"):
pbar.set_postfix({"loss": f"{red_logs.get('train/loss_total', 0):.4f}", "step": global_step})
if val_loader is not None and ((epoch + 1) % int(cfg["train"].get("val_every", 1)) == 0):
metrics = evaluate(unwrap_model(model), val_loader, cfg, device)
if is_main_process():
score = metrics.get("val/dice_mean", -1e9)
print(f"Epoch {epoch+1} validation: {metrics}")
wandb_log(wandb_run, metrics, step=global_step)
if score > best:
best = score
save_checkpoint(
ckpt_dir / "best.pt", model, optimizer, epoch+1, global_step, best,
velocity_field, cfg, teacher=teacher,
include_optimizer=bool(cfg["train"].get("save_optimizer_in_best", False)),
)
# Keep all ranks synchronized after rank0 best-checkpoint writing.
barrier()
# Always keep a resumable last.pt. Do not create epoch_N.pt files unless explicitly requested.
save_checkpoint(
ckpt_dir / "last.pt", model, optimizer, epoch+1, global_step, best,
velocity_field, cfg, teacher=teacher, include_optimizer=True,
)
# Ensure rank0 has finished writing last.pt before other ranks start the next epoch.
barrier()
if bool(cfg["train"].get("keep_epoch_checkpoints", False)) and int(cfg["train"].get("save_every", 0) or 0) > 0:
if (epoch + 1) % int(cfg["train"].get("save_every", 5)) == 0:
save_checkpoint(
ckpt_dir / f"epoch_{epoch+1}.pt", model, optimizer, epoch+1, global_step, best,
velocity_field, cfg, teacher=teacher, include_optimizer=True,
)
barrier()
# final last.pt is already saved after every epoch; save once more for completeness.
save_checkpoint(ckpt_dir / "last.pt", model, optimizer, epochs, global_step, best, velocity_field, cfg, teacher=teacher, include_optimizer=True)
return unwrap_model(model)