PART A: add KV gather diagnostics at blowup layer

This commit is contained in:
2026-06-03 06:25:35 +00:00
parent 262f844e2e
commit 86e59c16c5

View File

@@ -269,6 +269,35 @@ def main():
X_diag, A_l_a, attn_norms.get(li).to(dev, torch.float32))
x_normed = dequantize_nvfp4(x_quant_attn.x_fp4, x_quant_attn.x_sf, x_quant_attn.gsa)
print(f" |x_normed|={x_normed.abs().max().item():.2f} gsa={x_quant_attn.gsa}", flush=True)
# Print KV cache state
kc_diag = kv_caches[li]
swa_kv_d, swa_pos_d = kc_diag.get_swa()
print(f" KV: n_comp={kc_diag.n_comp} swa_len={swa_kv_d.shape[0]}", flush=True)
# Gather KV and print
ratio_diag = cr[li] if li < len(cr) else 128
if kc_diag.n_comp > 0:
if ratio_diag == 4:
topk_idx_d = None
if indexers.get(li) is not None:
topk_idx_d = indexers[li].forward(q_a_d, x_normed, kc_diag, pos, layer_idx=li)
if topk_idx_d is not None:
tk_d = topk_idx_d[0].clamp(0, kc_diag.n_comp - 1).int()
kv_nope_fp8_d, kv_nope_scale_d, kv_rope_bf16_d = kc_diag.gather_mixed_selective(tk_d)
else:
kv_nope_fp8_d, kv_nope_scale_d, kv_rope_bf16_d = kc_diag.gather_mixed_swa_only()
elif ratio_diag > 4:
kv_nope_fp8_d, kv_nope_scale_d, kv_rope_bf16_d = kc_diag.gather_mixed_all()
else:
kv_nope_fp8_d, kv_nope_scale_d, kv_rope_bf16_d = kc_diag.gather_mixed_swa_only()
else:
kv_nope_fp8_d, kv_nope_scale_d, kv_rope_bf16_d = kc_diag.gather_mixed_swa_only()
seq_len_d = kv_nope_scale_d.shape[0]
nope_max = kv_nope_fp8_d.view(torch.float8_e4m3fn).float().abs().max().item()
scale_max = kv_nope_scale_d.abs().max().item()
rope_max = kv_rope_bf16_d.float().abs().max().item()
print(f" Gathered KV: seq_len={seq_len_d} |nope_fp8|={nope_max:.2f} |nope_scale|={scale_max:.6f} |rope_bf16|={rope_max:.2f}", flush=True)
nope_dequant_max = (kv_nope_fp8_d.view(torch.float8_e4m3fn).float() * kv_nope_scale_d.unsqueeze(-1).float()).abs().max().item()
print(f" |nope_dequant_max|={nope_dequant_max:.4f}", flush=True)
F_attn_d, q_a_d = forward_attention(
x_normed, layer_w[li], li, cfg, *rope_caches[gpu],
kv_caches[li], pos, compressors.get(li), indexers.get(li), prod_lins.get(li),