From a986f96ba3e5ef18f891d4fc11df1f433cb08b08 Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Wed, 29 Jul 2026 10:21:03 +0200 Subject: [PATCH] Log router health, WGAN grad-norm split, n_sec accuracy, GPU/throughput to W&B MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds Router.gate_stats (per-router gate entropy + per-expert utilization), logged both per-batch (entropy only, train loop) and per-epoch (full stats, over the whole val set) — the router-collapse failure mode from the roadmap's rollout postmortem is now visible during training instead of only after a full rollout+analysis run. Also splits WGAN critic/ generator grad norms instead of summing them, logs critic LR, n_sec head accuracy, GPU peak memory + samples/sec, model parameter counts (in wandb.config), and an is_best flag — all wired into both metrics.csv and W&B. Co-Authored-By: Claude Sonnet 5 --- giant/model/network.py | 24 ++++ giant/train.py | 247 +++++++++++++++++++++++++++++++++++------ 2 files changed, 234 insertions(+), 37 deletions(-) diff --git a/giant/model/network.py b/giant/model/network.py index 4f444ba..35ec968 100644 --- a/giant/model/network.py +++ b/giant/model/network.py @@ -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]] = {} diff --git a/giant/train.py b/giant/train.py index a520368..cd80217 100644 --- a/giant/train.py +++ b/giant/train.py @@ -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)