[Fix] Don't write conv state from the fused KDA verify kernel (#39219)

Co-authored-by: mmangkad <mohammad.angkad@radixark.ai>
This commit is contained in:
Mohammad Miadh Angkad
2026-09-13 21:26:57 -07:00
committed by GitHub
co-authored by mmangkad
parent 60f6f03409
commit 2fd835b9c1
2 changed files with 40 additions and 31 deletions
@@ -12,15 +12,16 @@ two transpose copies the unfused path needs to feed the conv kernel.
Scope (v1): chain speculation only (``speculative_eagle_topk == 1``, i.e.
``retrieve_next_token is None``). The tree path keeps the unfused reference
kernels. Requires ``T >= kernel_width - 1`` (the rolled conv state is then
exactly the last ``kernel_width - 1`` input tokens, matching the reference
kernel's store).
kernels. Requires ``T >= kernel_width - 1``.
Numerics: deliberately bit-aligned with the unfused pair. The conv output is
rounded to the activation dtype (bf16) before entering the recurrence —
exactly what the unfused path does through its intermediate tensor — and all
expressions mirror the reference kernels line by line, with the same
num_warps so reduction order matches.
State: conv_state and the SSM state are read-only. Verify is speculative, and
the commit scatter advances them from the selected intermediate window.
Numerics: aligned with the unfused pair. The conv output is rounded to the
activation dtype (bf16) before entering the recurrence — exactly what the
unfused path does through its intermediate tensor — and all expressions mirror
the reference kernels line by line. Reduction order still splits differently
where many V heads share one Q/K head, worth ~1 ulp on the output.
"""
from typing import Optional
@@ -318,19 +319,8 @@ def fused_kda_conv_gating_verify_kernel(
)
tl.store(cache_ptr, b_h.to(cache_ptr.dtype.element_ty), mask=mask_h)
# Rolled conv state after consuming T >= W-1 tokens is exactly the last
# W-1 input tokens — which are the current window registers. The verify
# pass never writes the ssm state back (rollback happens at commit).
if is_qk_owner:
tl.store(cs_base + q_ch + 0 * stride_cs_tok, q_c0, mask=mask_k)
tl.store(cs_base + q_ch + 1 * stride_cs_tok, q_c1, mask=mask_k)
tl.store(cs_base + q_ch + 2 * stride_cs_tok, q_c2, mask=mask_k)
tl.store(cs_base + k_ch + 0 * stride_cs_tok, k_c0, mask=mask_k)
tl.store(cs_base + k_ch + 1 * stride_cs_tok, k_c1, mask=mask_k)
tl.store(cs_base + k_ch + 2 * stride_cs_tok, k_c2, mask=mask_k)
tl.store(cs_base + v_ch + 0 * stride_cs_tok, v_c0, mask=mask_v)
tl.store(cs_base + v_ch + 1 * stride_cs_tok, v_c1, mask=mask_v)
tl.store(cs_base + v_ch + 2 * stride_cs_tok, v_c2, mask=mask_v)
# No conv-state writeback: every V tile reads the same Q/K history, so a
# tile in a later wave would read what i_v == 0 had overwritten.
def fused_kda_conv_gating_verify(
@@ -358,14 +348,11 @@ def fused_kda_conv_gating_verify(
softplus_beta: float = 1.0,
softplus_threshold: float = 20.0,
use_qk_l2norm_in_kernel: bool = True,
# num_warps=4 is ~1.3x faster than the unfused pair in-graph; the output,
# conv_state and conv-window caches stay bit-identical to the reference.
# Only the fp32 intermediate-ssm rollback cache differs: the tl.sum
# reduction-order delta (~1 ulp/step) compounds through the delta-rule
# recurrence — measured ~6e-8 at T=4 standard gate (the production MTP
# shape), ~1.5e-5 at T=4 safe gate, ~2e-3 at T=8 safe gate. num_warps=1
# reproduces the reference reduction order exactly (all buffers
# bit-identical) but is ~2.4x slower in-graph — numerics debugging only.
# num_warps=4 is ~1.3x faster than the unfused pair in-graph; 1 restores the
# reference reduction order but is ~2.4x slower, for numerics debugging only.
# The fp32 intermediate-ssm rollback cache carries the reduction-order delta
# furthest: ~6e-8 at T=4 standard gate (the production MTP shape), ~2e-3 at
# T=8 safe gate. conv_state is not comparable to the reference at all.
num_warps: int = 4,
) -> torch.Tensor:
"""Chain-verify fast path. Returns ``o`` of shape [1, seq_len, HV, V],