PART A: add KV gather diagnostics at blowup layer
This commit is contained in:
@@ -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),
|
||||
|
||||
Reference in New Issue
Block a user