Feature/wandb integration #18
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user