From e332e1b84e643d02d6223542941030c03717f62c Mon Sep 17 00:00:00 2001 From: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com> Date: Tue, 22 Sep 2026 01:33:03 +0000 Subject: [PATCH] [Fix] Don't write conv state from the fused KDA verify kernel (#39524) Co-authored-by: mmangkad --- .../fla/fused_kda_conv_recurrent_verify.py | 34 ++++++------------- .../test_fused_kda_conv_recurrent_verify.py | 28 +++++++++++++-- 2 files changed, 36 insertions(+), 26 deletions(-) diff --git a/python/sglang/kernels/ops/attention/fla/fused_kda_conv_recurrent_verify.py b/python/sglang/kernels/ops/attention/fla/fused_kda_conv_recurrent_verify.py index de00fa56b..bb9dbf78e 100644 --- a/python/sglang/kernels/ops/attention/fla/fused_kda_conv_recurrent_verify.py +++ b/python/sglang/kernels/ops/attention/fla/fused_kda_conv_recurrent_verify.py @@ -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, diff --git a/test/registered/kernels/ops/attention/test_fused_kda_conv_recurrent_verify.py b/test/registered/kernels/ops/attention/test_fused_kda_conv_recurrent_verify.py index b3ee32a6e..4d5957404 100644 --- a/test/registered/kernels/ops/attention/test_fused_kda_conv_recurrent_verify.py +++ b/test/registered/kernels/ops/attention/test_fused_kda_conv_recurrent_verify.py @@ -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)