From 86e59c16c5c5c571a3fa25195f0f3702288d8a03 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Wed, 3 Jun 2026 06:25:35 +0000 Subject: [PATCH] PART A: add KV gather diagnostics at blowup layer --- tests/unit/test_part_a_decode_diagnostics.py | 29 ++++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/tests/unit/test_part_a_decode_diagnostics.py b/tests/unit/test_part_a_decode_diagnostics.py index 4c91f461..99e2d6e3 100644 --- a/tests/unit/test_part_a_decode_diagnostics.py +++ b/tests/unit/test_part_a_decode_diagnostics.py @@ -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),