File size: 20,584 Bytes
d65ae7d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
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)