Download sacflow/engine/train_loop.py from sathiiii/SACFlow: direct link, hf CLI and curl.
- Browser
- Download file 20.6 kB
-
https://huggingface.co/sathiiii/SACFlow/resolve/main/sacflow/engine/train_loop.py
- Command line
-
hf download hf://sathiiii/SACFlow/sacflow/engine/train_loop.py
-
curl -L -o train_loop.py https://huggingface.co/sathiiii/SACFlow/resolve/main/sacflow/engine/train_loop.py
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()} | |
| 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) | |