diff --git a/python/sglang/kernels/ops/attention/fla/chunk_delta_h.py b/python/sglang/kernels/ops/attention/fla/chunk_delta_h.py index 7626358e3..41c451aff 100644 --- a/python/sglang/kernels/ops/attention/fla/chunk_delta_h.py +++ b/python/sglang/kernels/ops/attention/fla/chunk_delta_h.py @@ -13,7 +13,7 @@ from sglang.kernels.ops.attention.fla.index import ( prepare_chunk_indices, prepare_chunk_offsets, ) -from sglang.kernels.ops.attention.fla.op import exp, safe_exp +from sglang.kernels.ops.attention.fla.op import exp, exp2, safe_exp from sglang.kernels.ops.attention.fla.utils import ( autotune_cache_kwargs, is_nvidia_hopper, @@ -76,6 +76,7 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64( SAVE_NEW_VALUE: tl.constexpr, IS_VARLEN: tl.constexpr, NT_BUCKET: tl.constexpr, + USE_EXP2: tl.constexpr, ): i_v, i_nh = tl.program_id(0), tl.program_id(1) i_n, i_h = i_nh // H, i_nh % H @@ -220,7 +221,10 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64( mask=(o_k1 < K), other=0.0, ) - b_h1 *= exp(b_gk_last1)[None, :] + if USE_EXP2: + b_h1 *= exp2(b_gk_last1)[None, :] + else: + b_h1 *= exp(b_gk_last1)[None, :] if K > 64: o_k2 = 64 + o_k1 b_gk_last2 = tl.load( @@ -228,7 +232,10 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64( mask=(o_k2 < K), other=0.0, ) - b_h2 *= exp(b_gk_last2)[None, :] + if USE_EXP2: + b_h2 *= exp2(b_gk_last2)[None, :] + else: + b_h2 *= exp(b_gk_last2)[None, :] if K > 128: o_k3 = 128 + o_k1 b_gk_last3 = tl.load( @@ -236,7 +243,10 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64( mask=(o_k3 < K), other=0.0, ) - b_h3 *= exp(b_gk_last3)[None, :] + if USE_EXP2: + b_h3 *= exp2(b_gk_last3)[None, :] + else: + b_h3 *= exp(b_gk_last3)[None, :] if K > 192: o_k4 = 192 + o_k1 b_gk_last4 = tl.load( @@ -244,7 +254,10 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64( mask=(o_k4 < K), other=0.0, ) - b_h4 *= exp(b_gk_last4)[None, :] + if USE_EXP2: + b_h4 *= exp2(b_gk_last4)[None, :] + else: + b_h4 *= exp(b_gk_last4)[None, :] b_v = b_v.to(k.dtype.element_ty) p_k = tl.make_block_ptr( @@ -303,7 +316,11 @@ def chunk_gated_delta_rule_fwd_h( save_new_value: bool = True, cu_seqlens: Optional[torch.LongTensor] = None, chunk_indices: Optional[torch.LongTensor] = None, + use_exp2: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + assert not ( + use_exp2 and g is not None + ), "use_exp2 covers only the per-channel gk path; scalar g stays natural-exp" B, T, Hg, K, V = *k.shape, u.shape[-1] H = u.shape[-2] BT = CHUNK_SIZE @@ -353,5 +370,6 @@ def chunk_gated_delta_rule_fwd_h( SAVE_NEW_VALUE=v_new is not None, IS_VARLEN=cu_seqlens is not None, NT_BUCKET=(0 if NT <= 32 else (1 if NT <= 128 else 2)), + USE_EXP2=use_exp2, ) return h, v_new diff --git a/python/sglang/kernels/ops/attention/fla/chunk_intra.py b/python/sglang/kernels/ops/attention/fla/chunk_intra.py index 06ed1fd6f..454eba022 100644 --- a/python/sglang/kernels/ops/attention/fla/chunk_intra.py +++ b/python/sglang/kernels/ops/attention/fla/chunk_intra.py @@ -11,7 +11,7 @@ from sglang.kernels.ops.attention.fla.chunk_intra_token_parallel import ( from sglang.kernels.ops.attention.fla.index import ( prepare_chunk_indices, ) -from sglang.kernels.ops.attention.fla.op import exp, exp2, gather +from sglang.kernels.ops.attention.fla.op import exp2, gather from sglang.kernels.ops.attention.fla.utils import ( autotune_cache_kwargs, is_gather_supported, @@ -632,7 +632,7 @@ def chunk_kda_fwd_kernel_inter_solve_fused( tl.store(p_u2, b_u2.to(p_u2.dtype.element_ty), boundary_check=(0, 1)) tl.store(p_u3, b_u3.to(p_u3.dtype.element_ty), boundary_check=(0, 1)) - # ---- w = A_inv @ (k * beta * exp(gk)), kg = k * exp(gn - gk) ---- + # ---- w = A_inv @ (k * beta * exp2(gk)), kg = k * exp2(gn - gk) ---- w_base = w_out + (bos * H + i_h) * K kg_base = kg_out + (bos * H + i_h) * K last_idx = min(i_t * BT + BT, T) - 1 @@ -680,10 +680,10 @@ def chunk_kda_fwd_kernel_inter_solve_fused( b_gk2r = tl.load(p_gk2, boundary_check=(0, 1)).to(tl.float32) b_gk3r = tl.load(p_gk3, boundary_check=(0, 1)).to(tl.float32) - b_kb0 = (b_k0r * b_b0[:, None] * exp(b_gk0r)).to(b_k0r.dtype) - b_kb1 = (b_k1r * b_b1r[:, None] * exp(b_gk1r)).to(b_k1r.dtype) - b_kb2 = (b_k2r * b_b2r[:, None] * exp(b_gk2r)).to(b_k2r.dtype) - b_kb3 = (b_k3r * b_b3r[:, None] * exp(b_gk3r)).to(b_k3r.dtype) + b_kb0 = (b_k0r * b_b0[:, None] * exp2(b_gk0r)).to(b_k0r.dtype) + b_kb1 = (b_k1r * b_b1r[:, None] * exp2(b_gk1r)).to(b_k1r.dtype) + b_kb2 = (b_k2r * b_b2r[:, None] * exp2(b_gk2r)).to(b_k2r.dtype) + b_kb3 = (b_k3r * b_b3r[:, None] * exp2(b_gk3r)).to(b_k3r.dtype) b_w0 = tl.dot(b_Ai00_h, b_kb0) b_w1 = tl.dot(b_Ai10_h, b_kb0) + tl.dot(b_Ai11_h, b_kb1) @@ -716,10 +716,10 @@ def chunk_kda_fwd_kernel_inter_solve_fused( tl.store(p_w2, b_w2.to(p_w2.dtype.element_ty), boundary_check=(0, 1)) tl.store(p_w3, b_w3.to(p_w3.dtype.element_ty), boundary_check=(0, 1)) - b_kg0 = b_k0r * exp(b_gn[None, :] - b_gk0r) - b_kg1 = b_k1r * exp(b_gn[None, :] - b_gk1r) - b_kg2 = b_k2r * exp(b_gn[None, :] - b_gk2r) - b_kg3 = b_k3r * exp(b_gn[None, :] - b_gk3r) + b_kg0 = b_k0r * exp2(b_gn[None, :] - b_gk0r) + b_kg1 = b_k1r * exp2(b_gn[None, :] - b_gk1r) + b_kg2 = b_k2r * exp2(b_gn[None, :] - b_gk2r) + b_kg3 = b_k3r * exp2(b_gn[None, :] - b_gk3r) p_kg0 = tl.make_block_ptr( kg_base, (T, K), (H * K, 1), (i_tc0, i_k * BK), (BC, BK), (1, 0) @@ -918,7 +918,6 @@ def chunk_kda_fwd_intra( chunk_size: int = 64, chunk_indices: torch.LongTensor | None = None, safe_gate: bool = False, - disable_recompute: bool = False, fuse_recompute: bool = False, fuse_diagonal: bool = False, ): @@ -1042,14 +1041,13 @@ def chunk_kda_fwd_intra( recompute_w_u_fwd as kda_recompute_w_u_fwd, ) - w, u, qg, kg = kda_recompute_w_u_fwd( + w, u, kg = kda_recompute_w_u_fwd( k=k, v=v, beta=beta, A=Akk, - q=q if disable_recompute else None, gk=gk, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, ) - return w, u, qg, kg, Aqk, Akk + return w, u, None, kg, Aqk, Akk diff --git a/python/sglang/kernels/ops/attention/fla/kda.py b/python/sglang/kernels/ops/attention/fla/kda.py index 5ffeb1ffb..091ee5c6b 100644 --- a/python/sglang/kernels/ops/attention/fla/kda.py +++ b/python/sglang/kernels/ops/attention/fla/kda.py @@ -23,7 +23,7 @@ from sglang.kernels.ops.attention.fla.index import ( prepare_chunk_indices, ) from sglang.kernels.ops.attention.fla.l2norm import l2norm_fwd -from sglang.kernels.ops.attention.fla.op import exp, log +from sglang.kernels.ops.attention.fla.op import exp, exp2, log from sglang.kernels.ops.attention.fla.utils import ( check_shared_mem, is_intel, @@ -37,6 +37,10 @@ if is_intel: BS_LIST = [32, 64] if check_shared_mem() else [16, 32] +# Convert natural-log gates to log2 space before the exp2-based chunk kernels. +# log2(e) rounded to fp32, matching flash-linear-attention. +RCP_LN2 = 1.4426950216293335 + def cdiv(a: int, b: int) -> int: """Ceiling division.""" @@ -223,6 +227,7 @@ def chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_inter( A, Aqk, scale, + gk_scale, cu_seqlens, chunk_indices, T, @@ -288,19 +293,22 @@ def chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_inter( o_k = i_k * BK + tl.arange(0, BK) m_k = o_k < K # [BK,] - b_gn = tl.load(g + (i_t * BT + i_i * BC) * H * K + o_k, mask=m_k, other=0) + b_gn = ( + tl.load(g + (i_t * BT + i_i * BC) * H * K + o_k, mask=m_k, other=0) + * gk_scale + ) # [BC, BK] - b_g = tl.load(p_g, boundary_check=(0, 1)) - b_k = tl.load(p_k, boundary_check=(0, 1)) * exp(b_g - b_gn[None, :]) + b_g = tl.load(p_g, boundary_check=(0, 1)) * gk_scale + b_k = tl.load(p_k, boundary_check=(0, 1)) * exp2(b_g - b_gn[None, :]) # [BK, BC] - b_gk = tl.load(p_gk, boundary_check=(0, 1)) + b_gk = tl.load(p_gk, boundary_check=(0, 1)) * gk_scale b_kt = tl.load(b_kt, boundary_check=(0, 1)) # [BC, BC] - b_ktg = b_kt * exp(b_gn[:, None] - b_gk) + b_ktg = b_kt * exp2(b_gn[:, None] - b_gk) b_A += tl.dot(b_k, b_ktg) b_q = tl.load(p_q, boundary_check=(0, 1)) - b_qg = b_q * exp(b_g - b_gn[None, :]) * scale + b_qg = b_q * exp2(b_g - b_gn[None, :]) * scale b_Aqk += tl.dot(b_qg, b_ktg) b_A *= b_b[:, None] @@ -328,6 +336,7 @@ def chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_intra( A, Aqk, scale, + gk_scale, cu_seqlens, chunk_indices, T, @@ -388,7 +397,7 @@ def chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_intra( ) b_q = tl.load(p_q, boundary_check=(0, 1)) b_k = tl.load(p_k, boundary_check=(0, 1)) - b_g = tl.load(p_g, boundary_check=(0, 1)) + b_g = tl.load(p_g, boundary_check=(0, 1)) * gk_scale p_b = beta + (bos + i_t * BT + i_i * BC + o_i) * H + i_h b_k = b_k * tl.load(p_b, mask=m_A, other=0)[:, None] @@ -398,8 +407,8 @@ def chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_intra( for j in range(0, min(BC, T - i_t * BT - i_i * BC)): b_kt = tl.load(p_kt, mask=m_k, other=0).to(tl.float32) - b_gk = tl.load(p_gk, mask=m_k, other=0).to(tl.float32) - b_ktg = b_kt[None, :] * exp(b_g - b_gk[None, :]) + b_gk = tl.load(p_gk, mask=m_k, other=0).to(tl.float32) * gk_scale + b_ktg = b_kt[None, :] * exp2(b_g - b_gk[None, :]) b_A = tl.sum(b_k * b_ktg, 1) b_A = tl.where(o_i > j, b_A, 0.0) b_Aqk = tl.sum(b_q * b_ktg, 1) @@ -416,6 +425,7 @@ def chunk_kda_scaled_dot_kkt_fwd( gk: torch.Tensor | None = None, beta: torch.Tensor | None = None, scale: float | None = None, + gk_scale: float = 1.0, cu_seqlens: torch.LongTensor | None = None, chunk_size: int = 64, output_dtype: torch.dtype = torch.float32, @@ -429,7 +439,13 @@ def chunk_kda_scaled_dot_kkt_fwd( beta (torch.Tensor): The beta tensor of shape `[B, T, H]`. gk (torch.Tensor): - The cumulative sum of the gate tensor of shape `[B, T, H, K]` applied to the key tensor. Default: `None`. + The cumulative sum of the gate tensor of shape `[B, T, H, K]` applied to the key tensor, + in log2 space (the kernels apply `exp2`). + Default: `None`. + gk_scale (float): + Scale multiplied onto `gk` as it is loaded in the kernels. Pass a natural-log + cumsum with `gk_scale=RCP_LN2` to convert to log2 space in-kernel without + materializing a scaled copy of `gk`. Default: `1.0`. cu_seqlens (torch.LongTensor): The cumulative sequence lengths of the input tensor. Default: None @@ -463,6 +479,7 @@ def chunk_kda_scaled_dot_kkt_fwd( A=A, Aqk=Aqk, scale=scale, + gk_scale=gk_scale, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, T=T, @@ -483,6 +500,7 @@ def chunk_kda_scaled_dot_kkt_fwd( A=A, Aqk=Aqk, scale=scale, + gk_scale=gk_scale, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, T=T, @@ -508,9 +526,7 @@ def chunk_kda_scaled_dot_kkt_fwd( ) @triton.jit(do_not_specialize=["T"]) def recompute_w_u_fwd_kernel( - q, k, - qg, kg, v, beta, @@ -527,7 +543,6 @@ def recompute_w_u_fwd_kernel( BT: tl.constexpr, BK: tl.constexpr, BV: tl.constexpr, - STORE_QG: tl.constexpr, STORE_KG: tl.constexpr, IS_VARLEN: tl.constexpr, DOT_PRECISION: tl.constexpr, @@ -605,27 +620,7 @@ def recompute_w_u_fwd_kernel( (1, 0), ) b_gk = tl.load(p_gk, boundary_check=(0, 1)) - b_kb *= exp(b_gk) - if STORE_QG: - p_q = tl.make_block_ptr( - q + (bos * H + i_h) * K, - (T, K), - (H * K, 1), - (i_t * BT, i_k * BK), - (BT, BK), - (1, 0), - ) - p_qg = tl.make_block_ptr( - qg + (bos * H + i_h) * K, - (T, K), - (H * K, 1), - (i_t * BT, i_k * BK), - (BT, BK), - (1, 0), - ) - b_q = tl.load(p_q, boundary_check=(0, 1)) - b_qg = b_q * exp(b_gk) - tl.store(p_qg, b_qg.to(p_qg.dtype.element_ty), boundary_check=(0, 1)) + b_kb *= exp2(b_gk) if STORE_KG: last_idx = min(i_t * BT + BT, T) - 1 @@ -634,7 +629,7 @@ def recompute_w_u_fwd_kernel( b_gn = tl.load( gk + ((bos + last_idx) * H + i_h) * K + o_k, mask=m_k, other=0.0 ) - b_kg = b_k * exp(b_gn - b_gk) + b_kg = b_k * exp2(b_gn - b_gk) p_kg = tl.make_block_ptr( kg + (bos * H + i_h) * K, @@ -655,11 +650,10 @@ def recompute_w_u_fwd( v: torch.Tensor, beta: torch.Tensor, A: torch.Tensor, - q: torch.Tensor | None = None, gk: torch.Tensor | None = None, cu_seqlens: torch.LongTensor | None = None, chunk_indices: torch.LongTensor | None = None, -) -> tuple[torch.Tensor, torch.Tensor]: +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]: B, T, H, K, V = *k.shape, v.shape[-1] BT = A.shape[-1] @@ -671,9 +665,7 @@ def recompute_w_u_fwd( u = torch.empty_like(v) kg = torch.empty_like(k) if gk is not None else None recompute_w_u_fwd_kernel[(NT, B * H)]( - q=q, k=k, - qg=None, kg=kg, v=v, beta=beta, @@ -688,12 +680,11 @@ def recompute_w_u_fwd( K=K, V=V, BT=BT, - STORE_QG=False, STORE_KG=kg is not None, IS_VARLEN=cu_seqlens is not None, DOT_PRECISION="tf32", ) - return w, u, None, kg + return w, u, kg @triton.autotune( @@ -780,7 +771,7 @@ def chunk_gla_fwd_kernel_o( # [BT, BK] b_g = tl.load(p_g, boundary_check=(0, 1)) # [BT, BK] - b_qg = (b_q * exp(b_g)).to(b_q.dtype) + b_qg = (b_q * exp2(b_g)).to(b_q.dtype) # [BK, BV] b_h = tl.load(p_h, boundary_check=(0, 1)) # works but dkw, owing to divine benevolence @@ -1060,6 +1051,7 @@ def chunk_kda_fwd( g, A_log=A_log, chunk_size=chunk_size, + scale=RCP_LN2, dt_bias=dt_bias, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, @@ -1070,6 +1062,7 @@ def chunk_kda_fwd( g = chunk_local_cumsum( g, chunk_size=chunk_size, + scale=RCP_LN2, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, ) @@ -1114,6 +1107,7 @@ def chunk_kda_fwd( initial_state_indices=initial_state_indices, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, + use_exp2=True, ) del w, u, kg diff --git a/python/sglang/kernels/ops/attention/linear/kda_blackwell/__init__.py b/python/sglang/kernels/ops/attention/linear/kda_blackwell/__init__.py index 12b6c5869..7ac4fbbc8 100644 --- a/python/sglang/kernels/ops/attention/linear/kda_blackwell/__init__.py +++ b/python/sglang/kernels/ops/attention/linear/kda_blackwell/__init__.py @@ -173,13 +173,20 @@ def chunk_kda_cutedsl( # injected through the cutedsl KKT/Aqk MMAs as an identity-right-operand pass: # with kL'=M (M in the first 64 K-slots) and kR'=onehot(chunk-pos), the MMA # kL'@kR'.T == M, so kkt_inv_uw/kernel_o see the correct matrix without overflow. - from sglang.kernels.ops.attention.fla.kda import chunk_kda_scaled_dot_kkt_fwd + from sglang.kernels.ops.attention.fla.kda import ( + RCP_LN2, + chunk_kda_scaled_dot_kkt_fwd, + ) ones_beta = q.new_ones(1, T, Hv, dtype=torch.float32) + # The FLA kkt kernels consume log2-space gate cumsums (exp2-based); g_cu must + # stay natural-log for the cutedsl kernels below, so let the kernels apply + # gk_scale=RCP_LN2 at load time instead of materializing a scaled copy. M_kk, M_qk = chunk_kda_scaled_dot_kkt_fwd( q.unsqueeze(0).contiguous(), k.unsqueeze(0).contiguous(), gk=g_cu.unsqueeze(0), + gk_scale=RCP_LN2, beta=ones_beta, scale=float(scale), cu_seqlens=cu_seqlens, diff --git a/python/sglang/srt/hardware_backend/xpu/kernels/fla/chunk_delta_h.py b/python/sglang/srt/hardware_backend/xpu/kernels/fla/chunk_delta_h.py index 009fff16c..607fb7d4c 100644 --- a/python/sglang/srt/hardware_backend/xpu/kernels/fla/chunk_delta_h.py +++ b/python/sglang/srt/hardware_backend/xpu/kernels/fla/chunk_delta_h.py @@ -8,7 +8,12 @@ from sglang.kernels.ops.attention.fla.index import ( prepare_chunk_indices, prepare_chunk_offsets, ) -from sglang.kernels.ops.attention.fla.op import exp, make_tensor_descriptor, safe_exp +from sglang.kernels.ops.attention.fla.op import ( + exp, + exp2, + make_tensor_descriptor, + safe_exp, +) from sglang.kernels.ops.attention.fla.utils import ( autotune_cache_kwargs, ) @@ -52,6 +57,7 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64_k_loop( SAVE_NEW_VALUE: tl.constexpr, IS_VARLEN: tl.constexpr, NT_BUCKET: tl.constexpr, # this arg is kept to align with the triton kernel for CUDA + USE_EXP2: tl.constexpr, ): i_v, i_nh = tl.program_id(0), tl.program_id(1) i_n, i_h = i_nh // H, i_nh % H @@ -199,7 +205,10 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64_k_loop( mask=(o_k1 < K), other=0.0, ) - b_h *= tl.expand_dims(exp(b_gk_last1), 0) + if USE_EXP2: + b_h *= tl.expand_dims(exp2(b_gk_last1), 0) + else: + b_h *= tl.expand_dims(exp(b_gk_last1), 0) # Delta update: h += k^T @ v b_k = tl.trans(k_desc.load([i_t * BT, k_blk])) @@ -224,7 +233,11 @@ def chunk_gated_delta_rule_fwd_h( save_new_value: bool = True, cu_seqlens: Optional[torch.LongTensor] = None, chunk_indices: Optional[torch.LongTensor] = None, + use_exp2: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + assert not ( + use_exp2 and g is not None + ), "use_exp2 covers only the per-channel gk path; scalar g stays natural-exp" B, T, Hg, K, V = *k.shape, u.shape[-1] H = u.shape[-2] BT = CHUNK_SIZE @@ -276,5 +289,6 @@ def chunk_gated_delta_rule_fwd_h( SAVE_NEW_VALUE=v_new is not None, IS_VARLEN=cu_seqlens is not None, NT_BUCKET=(0 if NT <= 32 else (1 if NT <= 128 else 2)), + USE_EXP2=use_exp2, ) return h, v_new diff --git a/test/registered/attention/test_kda_kernels.py b/test/registered/attention/test_kda_kernels.py index 41e1858d9..27cd47b38 100644 --- a/test/registered/attention/test_kda_kernels.py +++ b/test/registered/attention/test_kda_kernels.py @@ -11,11 +11,13 @@ from sglang.kernels.ops.attention.fla.fused_sigmoid_gating_recurrent import ( ) from sglang.kernels.ops.attention.fla.index import prepare_chunk_indices from sglang.kernels.ops.attention.fla.kda import ( + chunk_kda, fused_recurrent_kda, kda_gate_chunk_cumsum, ) from sglang.srt.utils.common import get_device from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.test_utils import CustomTestCase register_cuda_ci(est_time=12, stage="base-b", runner_config="1-gpu-large") register_amd_ci(est_time=12, stage="stage-b", runner_config="1-gpu-large-amd") @@ -245,6 +247,142 @@ class TestKDAGateChunkCumsum(unittest.TestCase): ) +@unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA") +class TestKDAChunkExponentDomain(CustomTestCase): + """Guard KDA prefill against mixing natural-log gates with exp2 kernels.""" + + @staticmethod + def _naive_recurrent(q, k, v, g, beta, initial_state, lengths): + q, k, v, g, beta = (tensor.float() for tensor in (q, k, v, g, beta)) + scale = q.shape[-1] ** -0.5 + output = torch.empty_like(v) + final_state = initial_state.float().clone() + + offset = 0 + for sequence_index, length in enumerate(lengths): + state = final_state[sequence_index] + for token_index in range(offset, offset + length): + state = state * g[0, token_index].exp().unsqueeze(-2) + residual = v[0, token_index] - torch.einsum( + "hvk,hk->hv", state, k[0, token_index] + ) + state = state + torch.einsum( + "hv,hk->hvk", + residual * beta[0, token_index, :, None], + k[0, token_index], + ) + output[0, token_index] = ( + torch.einsum("hvk,hk->hv", state, q[0, token_index]) * scale + ) + final_state[sequence_index] = state + offset += length + return output, final_state + + @staticmethod + def _relative_rmse(actual, expected): + error = (actual.float() - expected.float()).square().mean().sqrt() + baseline = expected.float().square().mean().sqrt().clamp_min(1e-8) + return (error / baseline).item() + + @torch.inference_mode() + def test_chunk_prefill_matches_natural_exp_recurrence(self): + device = get_device() + dtype = torch.bfloat16 + num_heads, head_dim = 2, 64 + + cases = ( + ([129], False, False), + ([15, 16, 17, 63, 65], True, True), + # 129 chunks x 2 heads = 258 CTAs > 256 -> _small_grid=False: exercises + # the standalone (non-fused) diagonal and recompute kernels. + ([2] * 129, True, False), + ) + for lengths, use_varlen, fuse_gate in cases: + with self.subTest( + lengths=lengths, use_varlen=use_varlen, fuse_gate=fuse_gate + ): + torch.manual_seed(42) + total_tokens = sum(lengths) + shape = (1, total_tokens, num_heads, head_dim) + q = torch.nn.functional.normalize( + torch.randn(shape, dtype=torch.float32, device=device), dim=-1 + ).to(dtype) + k = torch.nn.functional.normalize( + torch.randn(shape, dtype=torch.float32, device=device), dim=-1 + ).to(dtype) + v = torch.randn(shape, dtype=dtype, device=device) * 0.1 + raw_gate = ( + torch.randn(shape, dtype=torch.float32, device=device) * 0.5 - 2.0 + ).to(dtype) + A_log = torch.randn(num_heads, dtype=torch.float32, device=device) * 0.1 + dt_bias = ( + torch.randn( + num_heads * head_dim, + dtype=torch.float32, + device=device, + ) + * 0.1 + ) + activated_gate = -torch.exp( + A_log.view(1, 1, num_heads, 1) + ) * torch.nn.functional.softplus( + raw_gate.float() + dt_bias.view(1, 1, num_heads, head_dim) + ) + kernel_gate = raw_gate if fuse_gate else activated_gate.to(dtype) + reference_gate = activated_gate if fuse_gate else kernel_gate.float() + beta = torch.rand( + 1, total_tokens, num_heads, dtype=dtype, device=device + ).sigmoid() + initial_state = ( + torch.randn( + len(lengths), + num_heads, + head_dim, + head_dim, + dtype=torch.float32, + device=device, + ) + * 0.05 + ) + + expected_output, expected_state = self._naive_recurrent( + q=q, + k=k, + v=v, + g=reference_gate, + beta=beta, + initial_state=initial_state, + lengths=lengths, + ) + actual_state = initial_state.clone() + cu_seqlens = None + if use_varlen: + cu_seqlens = torch.tensor( + [0, *torch.tensor(lengths).cumsum(0).tolist()], + dtype=torch.int32, + device=device, + ) + actual_output = chunk_kda( + q=q.clone(), + k=k.clone(), + v=v.clone(), + g=kernel_gate.clone(), + beta=beta.clone(), + initial_state=actual_state, + initial_state_indices=torch.arange( + len(lengths), dtype=torch.int32, device=device + ), + cu_seqlens=cu_seqlens, + A_log=A_log if fuse_gate else None, + dt_bias=dt_bias if fuse_gate else None, + ) + + output_error = self._relative_rmse(actual_output, expected_output) + state_error = self._relative_rmse(actual_state, expected_state) + self.assertLess(output_error, 1e-2, f"output error={output_error:.3%}") + self.assertLess(state_error, 1e-2, f"state error={state_error:.3%}") + + @unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA") class TestKDAPackedDecode(unittest.TestCase): """Verify ``fused_recurrent_kda_packed_decode`` matches the existing decode