Skip to content

Commit 83e8467

Browse files
committed
pass lint check
1 parent 8f87ed7 commit 83e8467

7 files changed

Lines changed: 661 additions & 142 deletions

File tree

‎examples/dsa_hisa/README.md‎

100755100644
Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -111,9 +111,9 @@ clean_and_maintain_logits_interface(
111111
```
112112

113113
**What it does**: for each row `m`,
114-
* positions outside `[cu_seqlen_ks[m], cu_seqlen_ke[m])` → set to `-inf`
114+
- positions outside `[cu_seqlen_ks[m], cu_seqlen_ke[m])` → set to `-inf`
115115
(so `torch.topk` ignores them),
116-
* positions `cu_seqlen_ks[m]` and `cu_seqlen_ke[m] - 1` → set to `+inf`
116+
- positions `cu_seqlen_ks[m]` and `cu_seqlen_ke[m] - 1` → set to `+inf`
117117
(force-maintain the boundary blocks: they are always picked by the
118118
subsequent top-block selection — a standard hisa trick to preserve
119119
sink and local blocks).
@@ -125,8 +125,8 @@ clean_and_maintain_logits_interface(
125125
**Meaning**: fine-grained fp8 MQA over only the **raw K tokens** inside the
126126
top-`block_topk` pool blocks selected per query. Two kernel variants are
127127
auto-dispatched by the factory:
128-
* general (`kv_block_size > block_N`): pipelined sub-block inner loop
129-
* small-pooling-size (`kv_block_size == block_N`): single pass, no pipeline
128+
- general (`kv_block_size > block_N`): pipelined sub-block inner loop
129+
- small-pooling-size (`kv_block_size == block_N`): single pass, no pipeline
130130

131131
**Interface**:
132132
```python

‎examples/dsa_hisa/block_sparse_mqa_fp8.py‎

100755100644
Lines changed: 103 additions & 46 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,8 @@
33
from tilelang.profiler import do_bench
44
import 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():
273321
if __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()

‎examples/dsa_hisa/clean_and_maintain_logits.py‎

100755100644
Lines changed: 29 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,8 @@
33
from tilelang.profiler import do_bench
44
import torch
55

6+
from tilelang_utils import prepare_ks_ke_from_cu_seqlens
7+
68

79
@tilelang.jit
810
def clean_and_maintain_logits_(
@@ -17,9 +19,9 @@ def clean_and_maintain_logits_(
1719

1820
@T.prim_func
1921
def clean_and_maintain_logits_kernel(
20-
Logits: T.Tensor([seq_len, seq_len_kv], dtype), # type: ignore
21-
CuSeqLenKS: T.Tensor([seq_len], indices_dtype), # type: ignore
22-
CuSeqLenKE: T.Tensor([seq_len], indices_dtype), # type: ignore
22+
Logits: T.Tensor([seq_len, seq_len_kv], dtype), # type: ignore
23+
CuSeqLenKS: T.Tensor([seq_len], indices_dtype), # type: ignore
24+
CuSeqLenKE: T.Tensor([seq_len], indices_dtype), # type: ignore
2325
):
2426
with T.Kernel(seq_len, threads=threads) as bx:
2527
tx = T.thread_binding(0, threads, thread="threadIdx.x")
@@ -65,12 +67,22 @@ def ref_clean_and_maintain_logits(
6567
return out
6668

6769

68-
def test_clean_and_maintain_logits(M: int = 4096, N: int = 4096):
70+
def test_clean_and_maintain_logits(M: int = 4096, N: int = 4096, num_seqs: int = 1):
71+
"""Correctness + speed test where `M` query rows are packed from
72+
`num_seqs` equal-length causal sequences. Per-row ``cu_ks / cu_ke``
73+
is derived from ``prepare_ks_ke_from_cu_seqlens`` so each row sees
74+
only the prefix of its own sequence (causal self-attention)."""
6975
torch.manual_seed(0)
70-
# Build causal prefill ranges: cu_ks[m] = 0, cu_ke[m] = m + 1.
76+
assert M % num_seqs == 0, f"M ({M}) must be divisible by num_seqs ({num_seqs})"
77+
assert (M // num_seqs) <= N, "N must accommodate the longest sequence"
78+
79+
per_seq = M // num_seqs
80+
cu_seqlens = torch.arange(num_seqs + 1, device="cuda", dtype=torch.long) * per_seq
81+
ks_long, ke_long = prepare_ks_ke_from_cu_seqlens(cu_seqlens)
82+
cu_ks = ks_long.to(torch.int32).contiguous()
83+
cu_ke = ke_long.to(torch.int32).clamp(max=N).contiguous()
84+
7185
logits_init = torch.randn(M, N, device="cuda", dtype=torch.float32)
72-
cu_ks = torch.zeros(M, device="cuda", dtype=torch.int32)
73-
cu_ke = (torch.arange(M, device="cuda") + 1).to(torch.int32).clamp(max=N)
7486

7587
# Run kernel in place on a copy.
7688
got = logits_init.clone()
@@ -85,13 +97,14 @@ def test_clean_and_maintain_logits(M: int = 4096, N: int = 4096):
8597
assert torch.equal(torch.isneginf(got), torch.isneginf(ref)), "neg-inf mask differs"
8698
finite = torch.isfinite(got) & torch.isfinite(ref)
8799
torch.testing.assert_close(got[finite], ref[finite], rtol=0.0, atol=0.0)
88-
print(f" correctness: PASS (M={M}, N={N})")
100+
print(f" correctness: PASS (M={M}, N={N}, num_seqs={num_seqs}, per_seq={per_seq})")
89101

90102
# Speed.
91103
def fn():
92104
logits = torch.randn(M, N, device="cuda", dtype=torch.float32) # fresh copy each iter
93105
clean_and_maintain_logits_interface(logits, cu_ks, cu_ke)
94106
return logits
107+
95108
ms = do_bench(fn, warmup=50, rep=200)
96109
# ~2 reads + 1 write of [M, N] f32, but mostly no-op except at mask boundaries.
97110
bytes_moved = 2 * M * N * 4
@@ -100,5 +113,12 @@ def fn():
100113

101114

102115
if __name__ == "__main__":
103-
for cfg in [(4096, 4096), (16384, 16384), (65536, 65536)]:
116+
# (M, N, num_seqs)
117+
for cfg in [
118+
(4096, 4096, 1),
119+
(4096, 4096, 4),
120+
(16384, 16384, 1),
121+
(16384, 16384, 8),
122+
(65536, 65536, 16),
123+
]:
104124
test_clean_and_maintain_logits(*cfg)

0 commit comments

Comments
 (0)