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

Co-authored-by: mmangkad <mohammad.angkad@radixark.ai>
This commit is contained in:
Mohammad Miadh Angkad
2026-09-21 18:33:03 -07:00
committed by GitHub
co-authored by mmangkad
parent 35eb7cf8d6
commit e332e1b84e
2 changed files with 36 additions and 26 deletions
@@ -12,9 +12,10 @@ 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``.
State: conv_state and the SSM state are read-only. Verify is speculative, and
the commit scatter advances them from the selected intermediate window.
ReplaySSM (``cache_ring``): instead of per-step [HV, V, K] fp32 state
snapshots, stash each step's raw inputs (pre-l2norm k, pre-delta v, gate,
@@ -383,19 +384,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(
@@ -423,13 +413,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; conv_state
# and the conv-window cache stay bit-identical to the reference, the bf16
# output within one ulp (the BV=4 tile reduces K in a different order).
# The fp32 intermediate-ssm rollback cache carries that ~1 ulp/step delta
# 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 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.
# The ReplaySSM ring values are bit-exact at any num_warps: they are
# elementwise (conv FMA chain, gate, sigmoid), upstream of every tl.sum.
num_warps: int = 4,
@@ -204,7 +204,7 @@ def _compare_case(case, num_warps, use_ring=False):
rings_fus = {name: buf.clone() for name, buf in template.items()}
else:
rings_ref = rings_fus = None
o_ref, conv_ref, win_ref, ic_ref = _run_reference(
o_ref, _, win_ref, ic_ref = _run_reference(
inp, B, T, H, HV, K, V, lower_bound, rings=rings_ref
)
o_fus, conv_fus, win_fus, ic_fus = _run_fused(
@@ -213,13 +213,13 @@ def _compare_case(case, num_warps, use_ring=False):
idx_vals = inp["idx_vals"]
valid_rows = [i for i, slot in enumerate(idx_vals) if slot >= 0]
touched_slots = [slot for slot in idx_vals if slot >= 0]
o_ref_v = o_ref.reshape(B, T, HV, V)[valid_rows]
o_fus_v = o_fus.reshape(B, T, HV, V)[valid_rows]
# One bf16 ulp: the fused and reference tiles reduce K in different orders.
torch.testing.assert_close(o_fus_v, o_ref_v, rtol=2**-7, atol=1e-7)
assert torch.equal(conv_ref[touched_slots], conv_fus[touched_slots])
# conv_state is read-only in verify; the commit scatter advances it.
assert torch.equal(inp["conv_pool"], conv_fus)
assert torch.equal(win_ref[valid_rows], win_fus[valid_rows])
if use_ring:
# Full-tensor bitwise: ring values are elementwise (conv FMA chain,
@@ -239,6 +239,28 @@ def test_matches_unfused_reference(case):
_compare_case(case, num_warps=4)
def test_output_does_not_depend_on_cta_scheduling():
"""The verify output must not change with how the CTAs happen to be
scheduled. H=1 with HV=16 shares one Q/K history across 16 V tiles."""
if torch.cuda.get_device_capability()[0] < 9:
pytest.skip("green contexts need SM90 or newer")
from flashinfer.green_ctx import split_device_green_ctx_by_sm_count
case = (1, 6, 1, 16, 128, 128, 4, False, None, False, 1)
B, T, H, HV, K, V, W, has_bias, lower_bound, neg_slot, seed = case
inp = _make_inputs(B, T, H, HV, K, V, W, has_bias, neg_slot, seed)
full = _run_fused(inp, B, T, H, HV, K, V, lower_bound, num_warps=4)[0]
streams, _ = split_device_green_ctx_by_sm_count(torch.device("cuda:0"), [8])
# The green stream is non-blocking, so it must be told to wait for the
# inputs produced above; synchronize() afterwards only waits on the consumer.
streams[0].wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(streams[0]):
squeezed = _run_fused(inp, B, T, H, HV, K, V, lower_bound, num_warps=4)[0]
streams[0].synchronize()
assert torch.equal(full, squeezed)
@pytest.mark.parametrize("case", _RING_CASES)
def test_replayssm_ring_matches_unfused(case):
_compare_case(case, num_warps=4, use_ring=True)