Feature/wandb integration #18

Merged
lars merged 10 commits from feature/wandb-integration into master 2026-07-29 10:52:52 +02:00
2 changed files with 234 additions and 37 deletions
Showing only changes of commit a986f96ba3 - Show all commits
+24
View File
@@ -566,6 +566,30 @@ class Router(nn.Module):
"""
return torch.zeros((), device=cond_cont.device)
def gate_stats(
self, cond_cont: torch.Tensor, cond_cat: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
"""Diagnostics for catching a router that fails to specialize.
Returns `(norm_entropy, importance)`:
- `norm_entropy`: scalar, the batch-mean of each row's gate entropy
divided by `log(n_experts)`, in [0, 1] and comparable across
routers with different `n_experts` (1.0 = uniform/collapsed
gating, 0.0 = fully hard routing).
- `importance`: (n_experts,) tensor, `gate(...).sum(dim=0)` — the
*unnormalized* per-expert weight mass for this batch. Callers
wanting a global utilization share across many batches must sum
this across batches first and normalize once at the end;
averaging per-batch shares instead would treat every batch as
equally important regardless of size and understate a
rarely-but-fully-used expert.
"""
gate = self.gate(cond_cont, cond_cat) # (B, n_experts)
row_entropy = -(gate * (gate + 1e-8).log()).sum(dim=-1) # (B,)
norm_entropy = row_entropy.mean() / math.log(self.n_experts)
importance = gate.sum(dim=0) # (n_experts,)
return norm_entropy, importance
ROUTER_REGISTRY: dict[str, type[Router]] = {}
+210 -37
View File
@@ -32,6 +32,7 @@ _METRICS_FIELDS = [
"train_loss_s2",
"train_loss_balance",
"train_loss_proc",
"train_nsec_acc",
"d_loss",
"g_loss",
"wasserstein_estimate",
@@ -42,9 +43,24 @@ _METRICS_FIELDS = [
"val_loss_s2",
"val_loss_balance",
"val_loss_proc",
"val_nsec_acc",
"val_marginal_kl",
"router_s1_entropy",
"router_s1_util_min",
"router_s1_util_max",
"router_s1_util_std",
"router_s2_entropy",
"router_s2_util_min",
"router_s2_util_max",
"router_s2_util_std",
"lr",
"critic_lr",
"grad_norm",
"grad_norm_d",
"grad_norm_g",
"gpu_mem_mb",
"samples_per_sec",
"is_best",
"epoch_time_s",
]
@@ -108,9 +124,15 @@ def _compute_losses(
lambda_balance: float = 0.0,
lambda_proc: float = 0.0,
) -> tuple[
torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
torch.Tensor,
]:
"""Compute (total_loss, L_s1, L_nsec, L_s2, L_balance, L_proc) for one batch."""
"""Compute (total_loss, L_s1, L_nsec, L_s2, L_balance, L_proc, nsec_acc) for one batch."""
cond_cont, cond_cat, x1_s1, n_sec, sec_cont, proc_idx = batch
cond_cont = cond_cont.to(device)
cond_cat = cond_cat.to(device)
@@ -129,6 +151,7 @@ def _compute_losses(
# n_sec classification loss
n_sec_logits = stage1_model.predict_n_sec(cond_cont, cond_cat)
l_nsec = F.cross_entropy(n_sec_logits, n_sec)
nsec_acc = (n_sec_logits.argmax(dim=-1) == n_sec).float().mean()
# Stage-2 secondary flow loss
# Use a noiseless Stage-1 target as context (detach to avoid back-prop
@@ -171,7 +194,7 @@ def _compute_losses(
total = total + lambda_balance * l_balance
if lambda_proc > 0:
total = total + lambda_proc * l_proc
return total, l_s1, l_nsec, l_s2, l_balance, l_proc
return total, l_s1, l_nsec, l_s2, l_balance, l_proc, nsec_acc
def _wgan_train_step(
@@ -263,7 +286,9 @@ def _wgan_train_step(
# --- Generator (+ n_sec) step ---
did_g_step = step_count % n_critic == 0
l_nsec = F.cross_entropy(generator.predict_n_sec(cond_cont, cond_cat), n_sec)
n_sec_logits = generator.predict_n_sec(cond_cont, cond_cat)
l_nsec = F.cross_entropy(n_sec_logits, n_sec)
nsec_acc = (n_sec_logits.argmax(dim=-1) == n_sec).float().mean()
optimizer_g.zero_grad()
if did_g_step:
g1 = generator_loss(critic_fn1, fake1)
@@ -283,8 +308,11 @@ def _wgan_train_step(
"wasserstein_estimate": wasserstein_estimate,
"gp_loss": (gp1 + lambda_s2 * gp2).detach(),
"l_nsec": l_nsec.detach(),
"nsec_acc": nsec_acc.detach(),
"did_g_step": did_g_step,
"grad_norm": grad_norm_d.item() + grad_norm_g.item(),
"grad_norm_d": grad_norm_d.item(),
"grad_norm_g": grad_norm_g.item(),
}
@@ -328,6 +356,18 @@ def train(
out_dir = Path(out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
stage1_params = sum(p.numel() for p in stage1_model.parameters())
sec_decoder_params = sum(p.numel() for p in sec_decoder.parameters())
critic_params = (
sum(p.numel() for p in critic.parameters()) if critic is not None else 0
)
sec_critic_params = (
sum(p.numel() for p in sec_critic.parameters()) if sec_critic is not None else 0
)
total_params = (
stage1_params + sec_decoder_params + critic_params + sec_critic_params
)
wandb_run = None
if use_wandb:
try:
@@ -358,6 +398,11 @@ def train(
"n_critic": n_critic,
"gp_weight": gp_weight,
"model": model_config or {},
"stage1_params": stage1_params,
"sec_decoder_params": sec_decoder_params,
"critic_params": critic_params,
"sec_critic_params": sec_critic_params,
"total_params": total_params,
},
)
@@ -370,6 +415,13 @@ def train(
critic = critic.to(device)
sec_critic = sec_critic.to(device)
# MoE routing trunk (RoutedDenoisingMLP/RoutedSecondaryDecoder) is
# optional and orthogonal to `mode` — both stages carry a `.router`
# when enabled. Each router is an independent instance (their
# `n_experts` need not match), used both for the batch-level gate
# entropy snapshot below and the val-level gate stats further down.
has_router = hasattr(stage1_model, "router") and hasattr(sec_decoder, "router")
# Flow-matching/diffusion models sample noticeably better from an EMA of
# the weights than from the raw SGD-noisy ones — buffers (e.g. the fixed
# sinusoidal-embedding freqs, or non-learned router centers) never change
@@ -487,6 +539,8 @@ def train(
with _GracefulShutdown() as shutdown:
for epoch in range(start_epoch, epochs + 1):
epoch_start = time.monotonic()
if device.type == "cuda":
torch.cuda.reset_peak_memory_stats(device)
stage1_model.train()
sec_decoder.train()
if mode == "wgan":
@@ -503,6 +557,9 @@ def train(
train_g_sum = 0.0
train_wasserstein_sum = 0.0
train_gp_sum = 0.0
train_nsec_acc_sum = 0.0
train_grad_norm_d_sum = 0.0
train_grad_norm_g_sum = 0.0
train_n = 0
train_batches = 0
grad_norm_sum = 0.0
@@ -559,18 +616,23 @@ def train(
train_g_sum += stats["g_loss"].item() * B
train_wasserstein_sum += stats["wasserstein_estimate"].item() * B
train_gp_sum += stats["gp_loss"].item() * B
train_nsec_acc_sum += stats["nsec_acc"].item() * B
train_grad_norm_d_sum += stats["grad_norm_d"] * B
train_grad_norm_g_sum += stats["grad_norm_g"] * B
else:
loss, l_s1, l_nsec, l_s2, l_balance, l_proc = _compute_losses(
stage1_model,
sec_decoder,
batch,
mode,
ddpm_schedule,
device,
lambda_nsec,
lambda_s2,
lambda_balance,
lambda_proc,
loss, l_s1, l_nsec, l_s2, l_balance, l_proc, nsec_acc = (
_compute_losses(
stage1_model,
sec_decoder,
batch,
mode,
ddpm_schedule,
device,
lambda_nsec,
lambda_s2,
lambda_balance,
lambda_proc,
)
)
optimizer.zero_grad()
loss.backward()
@@ -593,6 +655,7 @@ def train(
train_s2_sum += l_s2.item() * B
train_balance_sum += l_balance.item() * B
train_proc_sum += l_proc.item() * B
train_nsec_acc_sum += nsec_acc.item() * B
train_n += B
train_batches += 1
@@ -615,16 +678,42 @@ def train(
and wandb_log_every > 0
and global_step % wandb_log_every == 0
):
wandb_run.log(
{
"batch/epoch": epoch,
"batch/loss": batch_loss,
"batch/loss_ema": ema_loss,
"batch/grad_norm": batch_grad_norm,
"batch/lr": optimizer.param_groups[0]["lr"],
},
step=global_step,
)
log_payload = {
"batch/epoch": epoch,
"batch/loss": batch_loss,
"batch/loss_ema": ema_loss,
"batch/grad_norm": batch_grad_norm,
"batch/lr": optimizer.param_groups[0]["lr"],
"batch/critic_lr": (
optimizer_d.param_groups[0]["lr"]
if optimizer_d is not None
else 0.0
),
}
if has_router:
# Cheap re-use of the batch already in hand — no
# extra data loading, just a small forward through
# each router's own gate function. Only entropy is
# logged at this granularity (not per-expert
# utilization): a single batch's importance sum is
# too noisy as a "global share" estimate, whereas
# the val-loop aggregate (below) sums over the
# whole val set for that. Batch-level entropy alone
# is still enough to see a router collapsing in
# real time, mid-epoch, rather than only at the
# next validation pass.
with torch.no_grad():
cond_cont_b = batch[0].to(device)
cond_cat_b = batch[1].to(device)
s1_entropy, _ = stage1_model.router.gate_stats(
cond_cont_b, cond_cat_b
)
s2_entropy, _ = sec_decoder.router.gate_stats(
cond_cont_b, cond_cat_b
)
log_payload["batch/router_s1_entropy"] = s1_entropy.item()
log_payload["batch/router_s2_entropy"] = s2_entropy.item()
wandb_run.log(log_payload, step=global_step)
if shutdown.requested:
break
@@ -635,7 +724,13 @@ def train(
train_loss = train_loss_sum / max(train_n, 1)
train_grad_norm = grad_norm_sum / max(train_batches, 1)
train_nsec_acc = train_nsec_acc_sum / max(train_n, 1)
train_grad_norm_d = train_grad_norm_d_sum / max(train_n, 1)
train_grad_norm_g = train_grad_norm_g_sum / max(train_n, 1)
current_lr = optimizer.param_groups[0]["lr"]
critic_lr_value = (
optimizer_d.param_groups[0]["lr"] if optimizer_d is not None else 0.0
)
stage1_model.eval()
sec_decoder.eval()
@@ -671,8 +766,12 @@ def train(
val_loss = val_marginal_kl
val_s1_sum = val_nsec_sum = val_s2_sum = val_balance_sum = (
val_proc_sum
) = 0.0
) = val_nsec_acc_sum = 0.0
val_n = 1
val_nsec_acc = 0.0
router_s1_entropy = router_s2_entropy = 0.0
router_s1_util_min = router_s1_util_max = router_s1_util_std = 0.0
router_s2_util_min = router_s2_util_max = router_s2_util_std = 0.0
else:
val_loss_sum = 0.0
val_s1_sum = 0.0
@@ -680,22 +779,36 @@ def train(
val_s2_sum = 0.0
val_balance_sum = 0.0
val_proc_sum = 0.0
val_nsec_acc_sum = 0.0
val_n = 0
if has_router:
n_experts_s1 = stage1_model.router.n_experts
n_experts_s2 = sec_decoder.router.n_experts
val_router_s1_entropy_sum = 0.0
val_router_s2_entropy_sum = 0.0
val_router_s1_importance_sum = torch.zeros(
n_experts_s1, device=device
)
val_router_s2_importance_sum = torch.zeros(
n_experts_s2, device=device
)
with torch.no_grad():
for val_batch_idx, batch in enumerate(val_loader):
if max_val_batches > 0 and val_batch_idx >= max_val_batches:
break
loss, l_s1, l_nsec, l_s2, l_balance, l_proc = _compute_losses(
stage1_model,
sec_decoder,
batch,
mode,
ddpm_schedule,
device,
lambda_nsec,
lambda_s2,
lambda_balance,
lambda_proc,
loss, l_s1, l_nsec, l_s2, l_balance, l_proc, nsec_acc = (
_compute_losses(
stage1_model,
sec_decoder,
batch,
mode,
ddpm_schedule,
device,
lambda_nsec,
lambda_s2,
lambda_balance,
lambda_proc,
)
)
B = batch[0].size(0)
val_loss_sum += loss.item() * B
@@ -704,8 +817,47 @@ def train(
val_s2_sum += l_s2.item() * B
val_balance_sum += l_balance.item() * B
val_proc_sum += l_proc.item() * B
val_nsec_acc_sum += nsec_acc.item() * B
if has_router:
cond_cont = batch[0].to(device)
cond_cat = batch[1].to(device)
s1_entropy, s1_importance = stage1_model.router.gate_stats(
cond_cont, cond_cat
)
s2_entropy, s2_importance = sec_decoder.router.gate_stats(
cond_cont, cond_cat
)
val_router_s1_entropy_sum += s1_entropy.item() * B
val_router_s2_entropy_sum += s2_entropy.item() * B
val_router_s1_importance_sum += s1_importance
val_router_s2_importance_sum += s2_importance
val_n += B
val_loss = val_loss_sum / max(val_n, 1)
val_nsec_acc = val_nsec_acc_sum / max(val_n, 1)
if has_router:
router_s1_entropy = val_router_s1_entropy_sum / max(val_n, 1)
router_s2_entropy = val_router_s2_entropy_sum / max(val_n, 1)
s1_util = val_router_s1_importance_sum / (
val_router_s1_importance_sum.sum().clamp_min(1e-8)
)
s2_util = val_router_s2_importance_sum / (
val_router_s2_importance_sum.sum().clamp_min(1e-8)
)
router_s1_util_min = s1_util.min().item()
router_s1_util_max = s1_util.max().item()
router_s1_util_std = (
s1_util.std().item() if n_experts_s1 > 1 else 0.0
)
router_s2_util_min = s2_util.min().item()
router_s2_util_max = s2_util.max().item()
router_s2_util_std = (
s2_util.std().item() if n_experts_s2 > 1 else 0.0
)
else:
router_s1_entropy = router_s2_entropy = 0.0
router_s1_util_min = router_s1_util_max = router_s1_util_std = 0.0
router_s2_util_min = router_s2_util_max = router_s2_util_std = 0.0
if validate_every > 0 and epoch % validate_every == 0:
print(f"[epoch {epoch}] marginal validation:")
@@ -721,6 +873,11 @@ def train(
val_marginal_kl = float(np.mean(marginal_result["kl_divergence"]))
epoch_time = time.monotonic() - epoch_start
gpu_mem_mb = (
torch.cuda.max_memory_allocated(device) / (1024 * 1024)
if device.type == "cuda"
else 0.0
)
is_best = val_loss < best_val_loss
marker = " [best]" if is_best else ""
@@ -746,6 +903,7 @@ def train(
"train_loss_s2": train_s2_sum / max(train_n, 1),
"train_loss_balance": train_balance_sum / max(train_n, 1),
"train_loss_proc": train_proc_sum / max(train_n, 1),
"train_nsec_acc": train_nsec_acc,
"d_loss": train_d_sum / max(train_n, 1),
"g_loss": train_g_sum / max(train_n, 1),
"wasserstein_estimate": train_wasserstein_sum / max(train_n, 1),
@@ -756,9 +914,24 @@ def train(
"val_loss_s2": val_s2_sum / max(val_n, 1),
"val_loss_balance": val_balance_sum / max(val_n, 1),
"val_loss_proc": val_proc_sum / max(val_n, 1),
"val_nsec_acc": val_nsec_acc,
"val_marginal_kl": val_marginal_kl,
"router_s1_entropy": router_s1_entropy,
"router_s1_util_min": router_s1_util_min,
"router_s1_util_max": router_s1_util_max,
"router_s1_util_std": router_s1_util_std,
"router_s2_entropy": router_s2_entropy,
"router_s2_util_min": router_s2_util_min,
"router_s2_util_max": router_s2_util_max,
"router_s2_util_std": router_s2_util_std,
"lr": current_lr,
"critic_lr": critic_lr_value,
"grad_norm": train_grad_norm,
"grad_norm_d": train_grad_norm_d,
"grad_norm_g": train_grad_norm_g,
"gpu_mem_mb": gpu_mem_mb,
"samples_per_sec": train_n / max(epoch_time, 1e-8),
"is_best": int(is_best),
"epoch_time_s": epoch_time,
}
metrics_writer.writerow(metrics_row)