33from tilelang .profiler import do_bench
44import torch
55
6+ from tilelang_utils import prepare_ks_ke_from_cu_seqlens
7+
68
79@tilelang .jit (
810 pass_configs = {
@@ -37,14 +39,14 @@ def fp8_native_block_sparse_mqa_attn_return_logits(
3739
3840 @T .prim_func
3941 def fp8_native_block_sparse_mqa_attn_return_logits_kernel (
40- IndexQ : T .Tensor (index_q_shape , fp8_dtype ), # type: ignore
41- IndexK : T .Tensor (index_k_shape , fp8_dtype ), # type: ignore
42- IndexKScale : T .Tensor (index_k_scale_shape , accum_dtype ), # type: ignore
42+ IndexQ : T .Tensor (index_q_shape , fp8_dtype ), # type: ignore
43+ IndexK : T .Tensor (index_k_shape , fp8_dtype ), # type: ignore
44+ IndexKScale : T .Tensor (index_k_scale_shape , accum_dtype ), # type: ignore
4345 TopKBlockIndex : T .Tensor ([seq_len , topk ], topk_index_dtype ), # type: ignore
44- Logits : T .Tensor (logits_shape , accum_dtype ), # type: ignore
45- Weights : T .Tensor ([seq_len , heads ], accum_dtype ), # type: ignore
46- CuSeqLenKS : T .Tensor ([seq_len ], index_dtype ), # type: ignore
47- CuSeqLenKE : T .Tensor ([seq_len ], index_dtype ), # type: ignore
46+ Logits : T .Tensor (logits_shape , accum_dtype ), # type: ignore
47+ Weights : T .Tensor ([seq_len , heads ], accum_dtype ), # type: ignore
48+ CuSeqLenKS : T .Tensor ([seq_len ], index_dtype ), # type: ignore
49+ CuSeqLenKE : T .Tensor ([seq_len ], index_dtype ), # type: ignore
4850 ):
4951 with T .Kernel (seq_len , threads = threads ) as bx :
5052 index_q_shared = T .alloc_shared ([H_per_block , index_dim ], fp8_dtype )
@@ -63,7 +65,7 @@ def fp8_native_block_sparse_mqa_attn_return_logits_kernel(
6365 cu_k_s_min = CuSeqLenKS [seq_len_i ]
6466 cu_k_e_max = CuSeqLenKE [seq_len_i ]
6567
66- T .copy (IndexQ [seq_len_i * heads : seq_len_i * heads + H_per_block , :], index_q_shared )
68+ T .copy (IndexQ [seq_len_i * heads : seq_len_i * heads + H_per_block , :], index_q_shared )
6769 T .copy (Weights [seq_len_i , :], weights )
6870
6971 for n_i in T .serial (topk ):
@@ -72,7 +74,7 @@ def fp8_native_block_sparse_mqa_attn_return_logits_kernel(
7274 for b_i in T .Pipelined (kv_block_size // block_N , num_stages = num_stages ):
7375 block_s_i = block_s + b_i * block_N
7476
75- T .copy (IndexK [block_s_i : block_s_i + block_N , :], index_k_shared )
77+ T .copy (IndexK [block_s_i : block_s_i + block_N , :], index_k_shared )
7678 for bn_i in T .Parallel (block_N ):
7779 scale_shared [bn_i ] = IndexKScale [block_s_i + bn_i ]
7880
@@ -86,7 +88,7 @@ def fp8_native_block_sparse_mqa_attn_return_logits_kernel(
8688 )
8789
8890 for bn_i , bq_i , h_i in T .Parallel (block_N , H_per_block // heads , heads ):
89- s_reshaped [bn_i , bq_i , h_i ] = ( T .max (s_reshaped [bn_i , bq_i , h_i ] * scale_shared [bn_i ], 0 ) * weights [bq_i , h_i ])
91+ s_reshaped [bn_i , bq_i , h_i ] = T .max (s_reshaped [bn_i , bq_i , h_i ] * scale_shared [bn_i ], 0 ) * weights [bq_i , h_i ]
9092
9193 T .reduce_sum (s_reshaped , logits , dim = - 1 , clear = True )
9294
@@ -100,14 +102,14 @@ def fp8_native_block_sparse_mqa_attn_return_logits_kernel(
100102
101103 @T .prim_func
102104 def fp8_native_block_sparse_mqa_attn_return_logits_kernel_for_small_pooling_size (
103- IndexQ : T .Tensor (index_q_shape , fp8_dtype ), # type: ignore
104- IndexK : T .Tensor (index_k_shape , fp8_dtype ), # type: ignore
105- IndexKScale : T .Tensor (index_k_scale_shape , accum_dtype ), # type: ignore
105+ IndexQ : T .Tensor (index_q_shape , fp8_dtype ), # type: ignore
106+ IndexK : T .Tensor (index_k_shape , fp8_dtype ), # type: ignore
107+ IndexKScale : T .Tensor (index_k_scale_shape , accum_dtype ), # type: ignore
106108 TopKBlockIndex : T .Tensor ([seq_len , topk ], topk_index_dtype ), # type: ignore
107- Logits : T .Tensor (logits_shape , accum_dtype ), # type: ignore
108- Weights : T .Tensor ([seq_len , heads ], accum_dtype ), # type: ignore
109- CuSeqLenKS : T .Tensor ([seq_len ], index_dtype ), # type: ignore
110- CuSeqLenKE : T .Tensor ([seq_len ], index_dtype ), # type: ignore
109+ Logits : T .Tensor (logits_shape , accum_dtype ), # type: ignore
110+ Weights : T .Tensor ([seq_len , heads ], accum_dtype ), # type: ignore
111+ CuSeqLenKS : T .Tensor ([seq_len ], index_dtype ), # type: ignore
112+ CuSeqLenKE : T .Tensor ([seq_len ], index_dtype ), # type: ignore
111113 ):
112114 with T .Kernel (seq_len , threads = threads ) as bx :
113115 index_q_shared = T .alloc_shared ([H_per_block , index_dim ], fp8_dtype )
@@ -124,14 +126,14 @@ def fp8_native_block_sparse_mqa_attn_return_logits_kernel_for_small_pooling_size
124126 cu_k_s_min = CuSeqLenKS [seq_len_i ]
125127 cu_k_e_max = CuSeqLenKE [seq_len_i ]
126128
127- T .copy (IndexQ [seq_len_i * heads : seq_len_i * heads + H_per_block , :], index_q_shared )
129+ T .copy (IndexQ [seq_len_i * heads : seq_len_i * heads + H_per_block , :], index_q_shared )
128130 T .copy (Weights [seq_len_i , :], weights )
129131
130132 for n_i in T .serial (topk ):
131133 topk_block_id = T .cast (TopKBlockIndex [seq_len_i , n_i ], index_dtype )
132134 block_s_i = topk_block_id * kv_block_size
133135
134- T .copy (IndexK [block_s_i : block_s_i + block_N , :], index_k_shared )
136+ T .copy (IndexK [block_s_i : block_s_i + block_N , :], index_k_shared )
135137 for bn_i in T .Parallel (block_N ):
136138 scale_shared [bn_i ] = IndexKScale [block_s_i + bn_i ]
137139
@@ -145,7 +147,7 @@ def fp8_native_block_sparse_mqa_attn_return_logits_kernel_for_small_pooling_size
145147 )
146148
147149 for bn_i , bq_i , h_i in T .Parallel (block_N , H_per_block // heads , heads ):
148- s_reshaped [bn_i , bq_i , h_i ] = ( T .max (s_reshaped [bn_i , bq_i , h_i ] * scale_shared [bn_i ], 0 ) * weights [bq_i , h_i ])
150+ s_reshaped [bn_i , bq_i , h_i ] = T .max (s_reshaped [bn_i , bq_i , h_i ] * scale_shared [bn_i ], 0 ) * weights [bq_i , h_i ]
149151
150152 T .reduce_sum (s_reshaped , logits , dim = - 1 , clear = True )
151153
@@ -176,12 +178,21 @@ def fp8_native_block_sparse_mqa_attn_return_logits_interface(
176178 seq_len , heads , index_dim = q .shape
177179 topk = topk_block_index .shape [1 ]
178180 kernel = fp8_native_block_sparse_mqa_attn_return_logits (
179- heads = heads , index_dim = index_dim , kv_block_size = kv_block_size , topk = topk ,
181+ heads = heads ,
182+ index_dim = index_dim ,
183+ kv_block_size = kv_block_size ,
184+ topk = topk ,
180185 )
181186 logits = torch .empty ([seq_len , topk * kv_block_size ], device = q .device , dtype = torch .float32 )
182187 kernel (
183- q .view (seq_len * heads , index_dim ), k , k_scale ,
184- topk_block_index , logits , weights , cu_seqlen_ks , cu_seqlen_ke ,
188+ q .view (seq_len * heads , index_dim ),
189+ k ,
190+ k_scale ,
191+ topk_block_index ,
192+ logits ,
193+ weights ,
194+ cu_seqlen_ks ,
195+ cu_seqlen_ke ,
185196 )
186197 return logits
187198
@@ -200,69 +211,106 @@ def ref_fp8_block_sparse_mqa(
200211 N = k_fp8 .shape [0 ]
201212 topk = topk_block_index .shape [1 ]
202213
203- block_starts = topk_block_index .long () * kv_block_size # [M, topk]
214+ block_starts = topk_block_index .long () * kv_block_size # [M, topk]
204215 pos_in_block = torch .arange (kv_block_size , device = q_fp8 .device )
205- k_abs = block_starts [..., None ] + pos_in_block [None , None , :] # [M, topk, B]
216+ k_abs = block_starts [..., None ] + pos_in_block [None , None , :] # [M, topk, B]
206217 k_safe = k_abs .clamp (0 , N - 1 )
207218
208219 q_f = q_fp8 .float ()
209220 k_f = k_fp8 .float () * k_scale [:, None ]
210221 gathered_k = k_f [k_safe .flatten ()].reshape (M , topk , kv_block_size , D )
211222
212- s = torch .einsum ("mhd,mtid->mtih" , q_f , gathered_k ) # [M, topk, B, H]
213- logits = (s .clamp (min = 0 ) * weights [:, None , None , :]).sum (dim = - 1 ) # [M, topk, B]
223+ s = torch .einsum ("mhd,mtid->mtih" , q_f , gathered_k ) # [M, topk, B, H]
224+ logits = (s .clamp (min = 0 ) * weights [:, None , None , :]).sum (dim = - 1 ) # [M, topk, B]
214225
215- in_range = (
216- (k_abs >= cu_seqlen_ks .long ()[:, None , None ])
217- & (k_abs < cu_seqlen_ke .long ()[:, None , None ])
218- & (k_abs < N )
219- )
226+ in_range = (k_abs >= cu_seqlen_ks .long ()[:, None , None ]) & (k_abs < cu_seqlen_ke .long ()[:, None , None ]) & (k_abs < N )
220227 logits = logits .masked_fill (~ in_range , float ("-inf" ))
221228 return logits .reshape (M , topk * kv_block_size )
222229
223230
224- def test_fp8_block_sparse_mqa (M : int = 1024 , H : int = 64 , D : int = 128 , kv_block_size : int = 128 , topk : int = 64 ):
231+ def test_fp8_block_sparse_mqa (
232+ M : int = 1024 ,
233+ H : int = 64 ,
234+ D : int = 128 ,
235+ kv_block_size : int = 128 ,
236+ topk : int = 64 ,
237+ num_seqs : int = 1 ,
238+ ):
239+ """Correctness + speed test packing `num_seqs` equal-length causal
240+ sequences into the [M, H, D] Q and [M, D] K tensors. Each query sees
241+ only the prefix of its own sequence (``cu_ks = start_of_seq``,
242+ ``cu_ke = start_of_seq + position_in_seq + 1``).
243+
244+ ``topk_block_index`` is drawn at random from [0, num_k_blocks) — some
245+ picks will point to blocks outside the query's own sequence; those
246+ positions get -inf via the kernel's built-in mask, and the torch ref
247+ produces the same -inf. Comparison checks both the +/-inf mask
248+ pattern (exact) and the finite values (fp8 tolerance)."""
225249 torch .manual_seed (0 )
226- N = M # causal self-attention prefill
250+ assert M % num_seqs == 0 , f"M ({ M } ) must be divisible by num_seqs ({ num_seqs } )"
251+ N = M # causal self-attention prefill, packed
252+
253+ per_seq = M // num_seqs
254+ cu_seqlens = torch .arange (num_seqs + 1 , device = "cuda" , dtype = torch .long ) * per_seq
255+ ks_long , ke_long = prepare_ks_ke_from_cu_seqlens (cu_seqlens )
256+ cu_ks = ks_long .to (torch .int32 ).contiguous ()
257+ cu_ke = ke_long .to (torch .int32 ).contiguous ()
258+
227259 q_bf16 = torch .randn (M , H , D , device = "cuda" , dtype = torch .bfloat16 )
228260 q = q_bf16 .to (torch .float8_e4m3fn )
229261 k_bf16 = torch .randn (N , D , device = "cuda" , dtype = torch .bfloat16 )
230262 k = k_bf16 .to (torch .float8_e4m3fn )
231263 k_scale = (0.1 + 0.01 * torch .rand (N , device = "cuda" , dtype = torch .float32 )).contiguous ()
232264 weights = torch .randn (M , H , device = "cuda" , dtype = torch .float32 )
233- cu_ks = torch .zeros (M , device = "cuda" , dtype = torch .int32 )
234- cu_ke = (torch .arange (M , device = "cuda" ) + 1 ).to (torch .int32 )
235265
236266 # Random per-query top-k blocks (distinct indices drawn from [0, num_blocks)).
237267 num_k_blocks = (N + kv_block_size - 1 ) // kv_block_size
238268 topk = min (topk , num_k_blocks )
239269 g = torch .Generator (device = "cuda" ).manual_seed (42 )
240- topk_block_index = torch .stack ([
241- torch .randperm (num_k_blocks , generator = g , device = "cuda" )[:topk ]
242- for _ in range (M )
243- ]).to (torch .int64 )
270+ topk_block_index = torch .stack ([torch .randperm (num_k_blocks , generator = g , device = "cuda" )[:topk ] for _ in range (M )]).to (torch .int64 )
244271
245272 # Correctness.
246273 got = fp8_native_block_sparse_mqa_attn_return_logits_interface (
247- q , k , k_scale , topk_block_index , kv_block_size , weights , cu_ks , cu_ke ,
274+ q ,
275+ k ,
276+ k_scale ,
277+ topk_block_index ,
278+ kv_block_size ,
279+ weights ,
280+ cu_ks ,
281+ cu_ke ,
248282 )
249283 ref = ref_fp8_block_sparse_mqa (
250- q , k , k_scale , topk_block_index , kv_block_size , weights , cu_ks , cu_ke ,
284+ q ,
285+ k ,
286+ k_scale ,
287+ topk_block_index ,
288+ kv_block_size ,
289+ weights ,
290+ cu_ks ,
291+ cu_ke ,
251292 )
252293 # The kernel marks out-of-range as -inf. Compare finite positions only —
253294 # the -inf mask pattern must agree exactly, so we also check that.
254- both_inf = torch .isinf (got ) & torch .isinf (ref ) & (got .sign () == ref .sign ())
255295 finite = torch .isfinite (got ) & torch .isfinite (ref )
256296 assert torch .equal (torch .isposinf (got ), torch .isposinf (ref )), "pos-inf mask differs"
257297 assert torch .equal (torch .isneginf (got ), torch .isneginf (ref )), "neg-inf mask differs"
258298 torch .testing .assert_close (got [finite ], ref [finite ], rtol = 1e-1 , atol = 2e-1 )
259- print (f" correctness: PASS (M={ M } , H={ H } , D={ D } , kv_block_size={ kv_block_size } , topk={ topk } )" )
299+ print (f" correctness: PASS (M={ M } , H={ H } , D={ D } , kv_block_size={ kv_block_size } , topk={ topk } , num_seqs= { num_seqs } , per_seq= { per_seq } )" )
260300
261301 # Speed.
262302 def fn ():
263303 return fp8_native_block_sparse_mqa_attn_return_logits_interface (
264- q , k , k_scale , topk_block_index , kv_block_size , weights , cu_ks , cu_ke ,
304+ q ,
305+ k ,
306+ k_scale ,
307+ topk_block_index ,
308+ kv_block_size ,
309+ weights ,
310+ cu_ks ,
311+ cu_ke ,
265312 )
313+
266314 ms = do_bench (fn , warmup = 50 , rep = 200 )
267315 # FLOPs: M × topk × kv_block_size × H × D (fp8×fp8) × 2 (mul+add).
268316 total_flops = 2 * M * topk * kv_block_size * H * D
@@ -273,6 +321,15 @@ def fn():
273321if __name__ == "__main__" :
274322 # Ref path materialises [M, topk, B, D] fp32 gathered_k which is ~M GB at
275323 # topk=64, kv_block_size=128, D=128. Keep M modest to avoid OOM.
276- for cfg in [(1024 , 64 , 128 , 128 , 64 ), (4096 , 64 , 128 , 128 , 64 ), (8192 , 64 , 128 , 128 , 64 ), (8192 , 64 , 128 , 64 , 128 ), (8192 , 64 , 128 , 256 , 32 )]:
324+ # (M, H, D, kv_block_size, topk, num_seqs)
325+ for cfg in [
326+ (1024 , 64 , 128 , 128 , 64 , 1 ),
327+ (4096 , 64 , 128 , 128 , 64 , 1 ),
328+ (4096 , 64 , 128 , 128 , 64 , 4 ),
329+ (8192 , 64 , 128 , 128 , 64 , 1 ),
330+ (8192 , 64 , 128 , 128 , 64 , 8 ),
331+ (8192 , 64 , 128 , 64 , 128 , 8 ),
332+ (8192 , 64 , 128 , 256 , 32 , 8 ),
333+ ]:
277334 test_fp8_block_sparse_mqa (* cfg )
278335 torch .cuda .empty_cache ()
0 commit comments