diff --git a/single_shot_inference.py b/single_shot_inference.py index 8bab5879..9f92303f 100644 --- a/single_shot_inference.py +++ b/single_shot_inference.py @@ -436,12 +436,14 @@ def moe_forward(x, li, moe_runner, se_runner, router, token_id): torch.cuda.synchronize(x.device) if topk_ids.max().item() >= 384 or topk_ids.min().item() < 0: print(f" L{li} BAD topk_ids: min={topk_ids.min().item()} max={topk_ids.max().item()}", flush=True) + if li >= 58: + print(f" L{li} MoE DIAG: topk_ids={topk_ids[0].tolist()} topk_w=[{','.join(f'{w:.3f}' for w in topk_w[0].tolist())}]", flush=True) if VERBOSE >= 2 and li < 3: print(f" L{li} MoE input: |x|={x.abs().max().item():.4f} has_nan={torch.isnan(x).any().item()}", flush=True) routed_out = moe_runner.run(x, topk_w, topk_ids) - if VERBOSE >= 2 and li < 3: - print(f" L{li} MoE routed: |out|={routed_out.abs().max().item():.4f} has_nan={torch.isnan(routed_out).any().item()}", flush=True) shared_out = se_runner.run(x) + if li >= 58: + print(f" L{li} MoE DIAG: |routed|={routed_out.abs().max().item():.1f} |shared|={shared_out.abs().max().item():.1f} |x|={x.abs().max().item():.1f}", flush=True) if VERBOSE >= 2 and li < 3: has_nan = torch.isnan(shared_out).any().item() out_max = shared_out.abs().max().item() if not has_nan else float('nan') @@ -475,6 +477,23 @@ def forward_layer(X_l, w, li, cfg, rope_cos, rope_sin, if VERBOSE >= 1: print(f" L{li}: |X|={X_l.abs().max().item():.1f}->{X_next.abs().max().item():.1f} " f"|Fa|={F_attn.abs().max().item():.1f} |Ff|={F_ffn.abs().max().item():.1f}", flush=True) + # Detailed diagnostics for last 3 layers or any layer with explosive growth + if li >= 58 or (li > 0 and X_next.abs().max().item() > 200): + A_a, B_a, C_a = attn_mhc._dynamic_params(X_l) + A_f, B_f, C_f = ffn_mhc._dynamic_params(X_mid) + print(f" L{li} DIAG: A_attn=[{A_a.min().item():.4f},{A_a.max().item():.4f}] " + f"C_attn=[{C_a.min().item():.4f},{C_a.max().item():.4f}] " + f"A_ffn=[{A_f.min().item():.4f},{A_f.max().item():.4f}] " + f"C_ffn=[{C_f.min().item():.4f},{C_f.max().item():.4f}]", flush=True) + print(f" L{li} DIAG: B_attn row_sum=[{B_a.sum(-1).min().item():.4f},{B_a.sum(-1).max().item():.4f}] " + f"col_sum=[{B_a.sum(-2).min().item():.4f},{B_a.sum(-2).max().item():.4f}] " + f"B_ffn row_sum=[{B_f.sum(-1).min().item():.4f},{B_f.sum(-1).max().item():.4f}] " + f"col_sum=[{B_f.sum(-2).min().item():.4f},{B_f.sum(-2).max().item():.4f}]", flush=True) + print(f" L{li} DIAG: |x_in_attn|={x_in.abs().max().item():.1f} " + f"|x_in_ffn|={x_in_f.abs().max().item():.1f} " + f"|X_l|={X_l.abs().max().item():.1f} " + f"|X_mid|={X_mid.abs().max().item():.1f} " + f"|X_next|={X_next.abs().max().item():.1f}", flush=True) return X_next # =====================================================================