Add L58-60 diagnostic: mHC A/B/C, MoE routed/shared, topk
This commit is contained in:
@@ -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
|
||||
|
||||
# =====================================================================
|
||||
|
||||
Reference in New Issue
Block a user