File size: 9,575 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 | from __future__ import annotations
from typing import Dict, Tuple
import torch
import torch.nn.functional as F
from sacflow.utils.metrics import torch_soft_dice_loss, entropy_loss, confidence_and_margin, finite_difference_boundary
from sacflow.methods.task_space import centered_classifier_basis, random_basis, project_task_and_residual, project_to_residual
from sacflow.methods.source_memory import class_moments, moment_transport_residual
from sacflow.utils.misc import unwrap_model
def ce_loss_masked(logits, target, mask=None):
loss = F.cross_entropy(logits, target.long(), reduction="none")
if mask is not None:
loss = loss * mask.float()
return loss.sum() / (mask.float().sum() + 1e-6)
return loss.mean()
def dice_loss_masked(logits, target, mask=None, num_classes=None, eps=1e-5):
if num_classes is None:
num_classes = logits.shape[1]
probs = torch.softmax(logits, dim=1)
target = target.clamp(0, num_classes - 1).long()
onehot = torch.nn.functional.one_hot(target, num_classes).permute(0,4,1,2,3).float()
if mask is not None:
m = mask.float().unsqueeze(1)
probs = probs * m
onehot = onehot * m
dims = tuple(range(2, logits.ndim))
inter = (probs * onehot).sum(dims)
denom = probs.sum(dims) + onehot.sum(dims)
dice = (2 * inter + eps) / (denom + eps)
return 1.0 - dice[:, 1:].mean()
def kl_masked(p, q, mask=None, eps=1e-8):
# p,q probabilities [B,C,H,W,D]
kl = (p * ((p+eps).log() - (q+eps).log())).sum(dim=1)
if mask is not None:
kl = kl * mask.float()
return kl.sum() / (mask.float().sum() + 1e-6)
return kl.mean()
def _stats_distance(mu_a, std_a, mu_b, std_b):
# simple differentiable class/channel statistic distance
return (mu_a - mu_b).abs().mean() + (std_a - std_b).abs().mean()
def build_task_basis(model, cfg, feat_dim, device, dtype):
basis = cfg.get("sacflow", {}).get("basis", "centered_svd")
W = model.final_classifier_weight().detach().to(device=device, dtype=dtype)
if basis == "random":
rank = max(1, min(W.shape[0]-1, feat_dim))
return random_basis(feat_dim, rank, device, dtype)
return centered_classifier_basis(W, foreground_only=cfg.get("sacflow", {}).get("foreground_only_basis", False))
def make_tau(B, device, dtype):
return torch.rand(B, device=device, dtype=dtype)
def sacflow_forward_step(model, teacher, velocity_field, batch, memory, cfg):
x = batch["image"]
num_classes = cfg["data"]["num_classes"]
# Use the wrapped model for the main forward when DDP is active; unwrap only for helper methods.
base_model = unwrap_model(model)
logits, feats = model(x, return_features=True)
feat = feats["prelogit"]
B, d, H, W, D = feat.shape
with torch.no_grad():
tlogits = teacher(x)
tprobs = torch.softmax(tlogits, dim=1)
if tprobs.shape[-3:] != (H,W,D):
tprobs_f = F.interpolate(tprobs, size=(H,W,D), mode="trilinear", align_corners=False)
else:
tprobs_f = tprobs
conf, margin, pseudo = confidence_and_margin(tprobs)
mask = (conf > float(cfg["train"].get("pseudo_conf_threshold", 0.75))).float()
Q = build_task_basis(base_model, cfg, d, feat.device, feat.dtype)
F_task, Rt = project_task_and_residual(feat, Q)
if cfg.get("sacflow", {}).get("flow_space", "residual") == "whole_feature":
F_task = torch.zeros_like(feat)
Rt = feat
Q = torch.empty(d, 0, device=feat.device, dtype=feat.dtype)
use_mem = cfg.get("sacflow", {}).get("use_compact_memory", True) and memory is not None
whole_feature = cfg.get("sacflow", {}).get("flow_space", "residual") == "whole_feature"
if use_mem and whole_feature and "feature_mu" in memory:
src_mu = memory["feature_mu"].to(feat.device, feat.dtype)
src_std = memory["feature_std"].to(feat.device, feat.dtype)
elif use_mem and "residual_mu" in memory:
src_mu = memory["residual_mu"].to(feat.device, feat.dtype)
src_std = memory["residual_std"].to(feat.device, feat.dtype)
else:
# strict fallback: use target stats as weak source proxy; ablate this separately.
# This makes R0 close to Rt and is intentionally weaker than compact-memory SACFlow.
src_mu, src_std, _ = class_moments(Rt.detach(), tprobs_f.detach())
R0, tgt_mu, tgt_std = moment_transport_residual(Rt.detach(), tprobs_f.detach(), src_mu, src_std)
R1 = Rt
tau = make_tau(B, feat.device, feat.dtype)
tau_view = tau.view(B,1,1,1,1)
sigma_tau = float(cfg.get("sacflow", {}).get("sigma_tau", 0.0) or 0.0)
noise = torch.randn_like(Rt) if sigma_tau > 0 else torch.zeros_like(Rt)
R_interp = (1 - tau_view) * R0 + tau_view * R1 + sigma_tau * tau_view * (1 - tau_view) * noise
u = R1 - R0 + sigma_tau * (1 - 2*tau_view) * noise
anchors = {"probs": tprobs_f.detach(), "boundary": finite_difference_boundary(tprobs_f.detach())}
use_v = bool(cfg.get("sacflow", {}).get("use_velocity_field", True))
if use_v and velocity_field is not None:
v = velocity_field(R_interp, tau, anchors)
if cfg.get("sacflow", {}).get("project_velocity_to_nullspace", True) and Q.numel() > 0:
v = project_to_residual(v, Q)
fm_loss = F.mse_loss(v, u.detach())
path_mode = cfg.get("sacflow", {}).get("path_from_velocity", "one_step")
if path_mode == "one_step":
Rtau = R0 + tau_view * v
else:
Rtau = R_interp
else:
v = torch.zeros_like(Rt)
fm_loss = torch.tensor(0.0, device=feat.device)
Rtau = R_interp
Ftau = F_task + Rtau
# Classify path states through the wrapped model so DDP can track gradients.
path_logits = model(prelogit_features=Ftau)
# resize mask/pseudo if needed
pseudo_f = pseudo
mask_f = mask
if pseudo.shape[-3:] != path_logits.shape[-3:]:
pseudo_f = F.interpolate(pseudo[:, None].float(), size=path_logits.shape[-3:], mode="nearest")[:,0].long()
mask_f = F.interpolate(mask[:, None].float(), size=path_logits.shape[-3:], mode="nearest")[:,0]
path_probs = torch.softmax(path_logits, dim=1)
task_loss = ce_loss_masked(path_logits, pseudo_f, mask_f)
dice = dice_loss_masked(path_logits, pseudo_f, mask_f, num_classes)
ent = entropy_loss(path_logits)
if tprobs.shape[-3:] != path_logits.shape[-3:]:
tprobs_path = F.interpolate(tprobs.detach(), size=path_logits.shape[-3:], mode="trilinear", align_corners=False)
else:
tprobs_path = tprobs.detach()
task_kl = kl_masked(path_probs, tprobs_path, mask_f)
b_pred = finite_difference_boundary(path_probs)
b_ref = finite_difference_boundary(tprobs_path)
boundary_loss = F.l1_loss(b_pred, b_ref)
null_leak = torch.tensor(0.0, device=feat.device)
if Q.numel() > 0:
null_leak = (Rtau - R0 - project_to_residual(Rtau - R0, Q)).pow(2).mean()
# Domain-progress proxy in residual-statistic space. A valid path state should
# move monotonically from source-like residual stats toward target residual stats.
rtau_mu, rtau_std, _ = class_moments(Rtau, tprobs_f.detach())
dist_src = _stats_distance(rtau_mu, rtau_std, src_mu.detach(), src_std.detach())
dist_tgt = _stats_distance(rtau_mu, rtau_std, tgt_mu.detach(), tgt_std.detach())
rho = dist_src / (dist_src + dist_tgt + 1e-6)
domain_progress_loss = (rho - tau.mean()).abs()
losses_cfg = cfg.get("sacflow", {}).get("losses", {})
val_cfg = cfg.get("sacflow", {}).get("validation", {})
use_path_weights = bool(val_cfg.get("use_weights", True))
if use_path_weights:
alpha = float(val_cfg.get("task_alpha", 1.0))
beta = float(val_cfg.get("anatomy_beta", 1.0))
# Use detached scalar weights so the gate selects/weights path states but
# does not create degenerate gradients that simply lower the weight.
accepted_weight = torch.exp(-alpha * task_kl.detach() - beta * boundary_loss.detach()).clamp(
min=float(val_cfg.get("min_weight", 0.05)), max=1.0
)
else:
accepted_weight = torch.tensor(1.0, device=feat.device, dtype=feat.dtype)
path_loss = accepted_weight * (
float(losses_cfg.get("path_ce", 1.0)) * task_loss +
float(losses_cfg.get("path_dice", 1.0)) * dice
)
loss = (
float(losses_cfg.get("fm", 1.0)) * fm_loss +
path_loss +
float(losses_cfg.get("task_kl", 0.25)) * task_kl +
float(losses_cfg.get("boundary", 0.05)) * boundary_loss +
float(losses_cfg.get("null_leakage", 0.1)) * null_leak +
float(losses_cfg.get("domain_progress", 0.05)) * domain_progress_loss +
float(losses_cfg.get("entropy", 0.01)) * ent
)
logs = {
"loss_total": loss.detach(),
"loss_fm": fm_loss.detach(),
"loss_path_ce": task_loss.detach(),
"loss_path_dice": dice.detach(),
"loss_task_kl": task_kl.detach(),
"loss_boundary": boundary_loss.detach(),
"loss_null_leak": null_leak.detach(),
"loss_domain_progress": domain_progress_loss.detach(),
"loss_entropy": ent.detach(),
"sacflow_rho_mean": rho.detach(),
"sacflow_tau_mean": tau.mean().detach(),
"sacflow_velocity_mag": v.detach().abs().mean(),
"sacflow_residual_gap": (R1-R0).detach().abs().mean(),
"pseudo_conf_mean": conf.detach().mean(),
"pseudo_accept_rate": mask.detach().mean(),
"path_weight": accepted_weight.detach(),
}
return loss, logs
|