[Example] Add HISA: hierarchical sparse attention indexer - #2069
Conversation
|
👋 Hi! Thank you for contributing to the TileLang project. Please remember to run We appreciate you taking this step! Our team will review your contribution, and we look forward to your awesome work! 🚀 |
📝 WalkthroughWalkthroughImplements a complete HISA prefill pipeline: FP8 block mean-pooling and re-quantization, pooled block MQA for block logits, in-place logits masking, FP8 block-sparse token-level MQA over selected blocks, and Python orchestration/tests to produce per-query top-k K indices. Changes
Sequence Diagram(s)sequenceDiagram
participant Q as Query Tensors
participant K as Raw K Tokens + Scales
participant BPool as Block Mean Pooling
participant PoolMQA as Pool MQA (block logits)
participant BlockSel as Block Selection (torch.topk)
participant FineMQA as Block-Sparse MQA (token logits)
participant TokenSel as Token Selection (torch.topk)
participant Out as K Index Output
Q->>PoolMQA: FP8 Q
K->>BPool: FP8 K + per-token scales
BPool->>PoolMQA: FP8 pooled blocks + block scales
PoolMQA->>PoolMQA: FP8×FP8 GEMM, per-block FP32 scaling, clamp & reduce
PoolMQA->>BlockSel: Block-level logits (masked)
BlockSel->>FineMQA: Top block indices
FineMQA->>FineMQA: Block-sparse FP8×FP8 MQA on raw K tokens -> token logits (out-of-range = -inf)
FineMQA->>TokenSel: Token-level logits
TokenSel->>Out: Per-query top-k K token offsets (out-of-window -> -1)
Estimated code review effort🎯 4 (Complex) | ⏱️ ~60 minutes Possibly related PRs
Suggested reviewers
Poem
🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✏️ Tip: You can configure your own custom pre-merge checks in the settings. ✨ Finishing Touches🧪 Generate unit tests (beta)
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
There was a problem hiding this comment.
Actionable comments posted: 10
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Inline comments:
In `@examples/dsa_hisa/block_sparse_mqa_fp8.py`:
- Line 210: The unpacking of q_fp8.shape assigns an unused H (M, H, D =
q_fp8.shape) — change it to use a throwaway variable (e.g., M, _, D =
q_fp8.shape or M, _H, D = q_fp8.shape) so linters stop complaining; also replace
any Unicode multiplication signs (e.g., '×') used around line 315 with the ASCII
asterisk '*' (search for occurrences in strings/comments or expressions) to
address the RUF003 warnings.
- Around line 71-101: The code copies IndexK and reads IndexKScale for a full
block using block_s_i and block_N before applying the cu_k_s/cu_k_e mask, which
can read past the end when seq_len_kv % kv_block_size != 0; fix by performing
guarded loads or padding: when computing T.copy(IndexK[block_s_i : block_s_i +
block_N, :], index_k_shared) and the subsequent IndexKScale loads, ensure you
only copy entries where block_s_i + bn_i < seq_len_kv (e.g., copy a shortened
slice and fill the remaining index_k_shared rows with safe defaults or use
conditional copies), and similarly guard the IndexKScale access
(IndexKScale[block_s_i + bn_i]) with block_s_i + bn_i < seq_len_kv; apply the
same guarded-load/padding logic to the other kernel section mentioned (lines
132-160) so no out-of-bounds reads occur.
In `@examples/dsa_hisa/clean_and_maintain_logits.py`:
- Around line 31-37: The loop can produce idx values >= logits width
(seq_len_kv) when block_K > actual width, causing out-of-bounds writes to
Logits; wrap or extend the existing conditions to ensure any write to Logits[bx,
idx] is guarded by a bounds check (e.g., require idx < seq_len_kv) before
assigning +inf or -inf. Locate the pipelined loops using T.Pipelined/T.ceildiv
and the variables idx, block_K, threads, tx, cu_k_s, cu_k_e and add a single idx
< seq_len_kv guard (or combine it into the existing if conditions) so no write
occurs when idx is outside [0, seq_len_kv-1].
In `@examples/dsa_hisa/fp8_block_mean_pooling.py`:
- Around line 48-66: The loop currently does an unconditional T.copy and
unconditional KScale loads which can read past seq_len_k for the last pooling
tile; change the loads to be guarded by cur_tl_block_size (or pad K/KScale to
block_N) so you only copy/assign up to cur_tl_block_size entries and set the
remaining lanes to zero. Specifically, replace the unconditional
T.copy(K[tl_block_s : tl_block_s + block_N, :], index_k) and the for bn_i in
T.Parallel(block_N): scale[bn_i] = KScale[tl_block_s + bn_i] with
guarded/limited loads that use cur_tl_block_size (or conditional bn_i <
cur_tl_block_size) to avoid indexing beyond seq_len_k, and ensure any untouched
index_k/scale lanes are explicitly zeroed so the subsequent index_k scaling and
T.reduce_sum are safe.
In `@examples/dsa_hisa/hisa.py`:
- Line 82: Replace the Unicode multiplication character (×) with the ASCII
letter "x" in the comments inside hisa.py—specifically update the comment
"blocks' raw tokens (block_topk_eff blocks × k_block_size tokens" and the other
comment occurrences that use "×" (also present in the nearby comments around the
token/block size descriptions and later comment at the end of the file). Search
for comments containing the "×" glyph and change them to use "x" so they satisfy
Ruff RUF003 while keeping the original comment wording intact.
- Around line 101-132: The function currently returns topk_indices with shape
[M, topk_tokens_eff] when topk_tokens > block_sparse_logits.shape[-1], but the
docstring promises [M, topk_tokens] with invalid slots set to -1; after
computing and masking topk_indices (symbols: topk_tokens, topk_tokens_eff,
block_sparse_logits, relevant_topk_indices, topk_indices), if topk_tokens_eff <
topk_tokens create a new int32 tensor filled with -1 of shape [M, topk_tokens]
and copy the existing topk_indices into the leftmost topk_tokens_eff columns
(leave remaining columns as -1), then return that padded tensor so the function
always returns shape [M, topk_tokens].
In `@examples/dsa_hisa/pool_mqa_fp8.py`:
- Line 219: Several comments in examples/dsa_hisa/pool_mqa_fp8.py use the
Unicode multiplication sign (×) which triggers Ruff RUF003; replace the Unicode
× with ASCII x in those comments (e.g., change "fp8×fp8 GEMM" to "fp8xfp8 GEMM"
and any other occurrences on the same comment lines). Search for the '×'
character in the file and update the comment text accordingly, keeping only
ASCII characters to silence the RUF003 warnings.
- Around line 79-106: The kernel currently assumes full tiles for queries and K
(block_Q, block_N) causing out-of-bounds reads from
blocked_kv_fp8/blocked_kv_scale and writes to Logits when seq_len_blocked_kv or
the final query tile is partial; fix by adding bounds checks and guarded
loads/stores or by computing padded tile extents before the loops: clamp
cu_k_s_min/cu_k_e_max and the per-tile index offsets using seq_len_blocked_kv
and CuSeqLenBlockedKS/CuSeqLenBlockedKE, use guarded copy/load for IndexBlockedK
and IndexBlockedKScale when cu_k_s_min + nbn_i*block_N + bn_i >=
seq_len_blocked_kv, and only write to Logits when seq_len_i + bq_i and the K
index are within valid lengths (use CuSeqLenBlockedKS/KE and seq_len_blocked_kv
to decide); alternatively ensure inputs are pre-padded to block_Q/block_N
multiples before kernel launch.
In `@examples/dsa_hisa/README.md`:
- Line 87: Update the three fenced code blocks containing the snippets starting
with "block_k_score[m, n] = ..." , the one containing "k_abs = blk *
kv_block_size + i" and the line "(1.1) fp8_native_block_mean_pooling
K, k_scale → blocked_k, blocked_k_scale" to include a language identifier by
replacing the opening ``` with ```text so the fences are annotated (e.g.,
```text block_k_score..., ```text k_abs = ..., ```text (1.1)
fp8_native_block_mean_pooling ...).
In `@examples/dsa_hisa/tilelang_utils.py`:
- Around line 264-270: The current slice uses last_seq_id =
torch.where(cu_seqlens.cumsum(0) >= total_seqlen)[0][0] which raises if no
cumulative sum reaches total_seqlen and prevents the fallback; change the logic
around cu_seqlens/cumsum to first compute the indices (e.g., idxs =
torch.where(cu_seqlens.cumsum(0) >= total_seqlen)[0]) and check idxs.numel() (or
.nelement()) before indexing: if idxs is empty, leave cu_seqlens unchanged so
the subsequent if cu_seqlens.sum() < total_seqlen fallback can run, otherwise
set last_seq_id = idxs[0].item() and slice cu_seqlens =
cu_seqlens[:last_seq_id]; keep references to cu_seqlens, total_seqlen,
last_seq_id and the cumsum check when applying the fix.
🪄 Autofix (Beta)
Fix all unresolved CodeRabbit comments on this PR:
- Push a commit to this branch (recommended)
- Create a new PR with the fixes
ℹ️ Review info
⚙️ Run configuration
Configuration used: defaults
Review profile: CHILL
Plan: Pro
Run ID: 7457af49-5ac9-4c35-8186-1df669cbf1d9
📒 Files selected for processing (7)
examples/dsa_hisa/README.mdexamples/dsa_hisa/block_sparse_mqa_fp8.pyexamples/dsa_hisa/clean_and_maintain_logits.pyexamples/dsa_hisa/fp8_block_mean_pooling.pyexamples/dsa_hisa/hisa.pyexamples/dsa_hisa/pool_mqa_fp8.pyexamples/dsa_hisa/tilelang_utils.py
| for n_i in T.serial(topk): | ||
| topk_block_id = T.cast(TopKBlockIndex[seq_len_i, n_i], index_dtype) | ||
| block_s = topk_block_id * kv_block_size | ||
| for b_i in T.Pipelined(kv_block_size // block_N, num_stages=num_stages): | ||
| block_s_i = block_s + b_i * block_N | ||
|
|
||
| T.copy(IndexK[block_s_i : block_s_i + block_N, :], index_k_shared) | ||
| for bn_i in T.Parallel(block_N): | ||
| scale_shared[bn_i] = IndexKScale[block_s_i + bn_i] | ||
|
|
||
| T.gemm( | ||
| index_k_shared, | ||
| index_q_shared, | ||
| s, | ||
| transpose_B=True, | ||
| clear_accum=True, | ||
| policy=T.GemmWarpPolicy.FullRow, | ||
| ) | ||
|
|
||
| for bn_i, bq_i, h_i in T.Parallel(block_N, H_per_block // heads, heads): | ||
| 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] | ||
|
|
||
| T.reduce_sum(s_reshaped, logits, dim=-1, clear=True) | ||
|
|
||
| for i_i in T.Parallel(block_N): | ||
| k_i = block_s_i + i_i | ||
| if k_i < cu_k_s_min or k_i >= cu_k_e_max: | ||
| logits[i_i, 0] = -T.infinity(accum_dtype) | ||
|
|
||
| for bn_i in T.Parallel(block_N): | ||
| Logits[seq_len_i, n_i * kv_block_size + b_i * block_N + bn_i] = logits[bn_i, 0] |
There was a problem hiding this comment.
Guard ragged final K blocks before loading.
Both kernel variants load IndexK[block_s_i : block_s_i + block_N] and IndexKScale[block_s_i + bn_i] before applying the [cu_k_s, cu_k_e) mask. If N % kv_block_size != 0 and the last block is selected, those loads can read past k. Pad k/k_scale to a full block or use guarded loads for block_s_i + bn_i < seq_len_kv.
Also applies to: 132-160
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@examples/dsa_hisa/block_sparse_mqa_fp8.py` around lines 71 - 101, The code
copies IndexK and reads IndexKScale for a full block using block_s_i and block_N
before applying the cu_k_s/cu_k_e mask, which can read past the end when
seq_len_kv % kv_block_size != 0; fix by performing guarded loads or padding:
when computing T.copy(IndexK[block_s_i : block_s_i + block_N, :],
index_k_shared) and the subsequent IndexKScale loads, ensure you only copy
entries where block_s_i + bn_i < seq_len_kv (e.g., copy a shortened slice and
fill the remaining index_k_shared rows with safe defaults or use conditional
copies), and similarly guard the IndexKScale access (IndexKScale[block_s_i +
bn_i]) with block_s_i + bn_i < seq_len_kv; apply the same guarded-load/padding
logic to the other kernel section mentioned (lines 132-160) so no out-of-bounds
reads occur.
| cu_seqlen_ks: torch.Tensor, | ||
| cu_seqlen_ke: torch.Tensor, | ||
| ) -> torch.Tensor: | ||
| M, H, D = q_fp8.shape |
There was a problem hiding this comment.
Clean up the remaining Ruff warnings.
Line 210 unpacks an unused H, and Line 315 uses Unicode multiplication signs flagged by RUF003.
🧹 Proposed fix
- M, H, D = q_fp8.shape
+ M, _H, D = q_fp8.shape
...
- # FLOPs: M × topk × kv_block_size × H × D (fp8×fp8) × 2 (mul+add).
+ # FLOPs: M x topk x kv_block_size x H x D (fp8 x fp8) x 2 (mul+add).Also applies to: 315-315
🧰 Tools
🪛 Ruff (0.15.10)
[warning] 210-210: Unpacked variable H is never used
Prefix it with an underscore or any other dummy variable pattern
(RUF059)
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@examples/dsa_hisa/block_sparse_mqa_fp8.py` at line 210, The unpacking of
q_fp8.shape assigns an unused H (M, H, D = q_fp8.shape) — change it to use a
throwaway variable (e.g., M, _, D = q_fp8.shape or M, _H, D = q_fp8.shape) so
linters stop complaining; also replace any Unicode multiplication signs (e.g.,
'×') used around line 315 with the ASCII asterisk '*' (search for occurrences in
strings/comments or expressions) to address the RUF003 warnings.
| for b_i in T.serial(T.ceildiv(cur_pooling_block_size, block_N)): | ||
| T.fill(index_k, 0.0) | ||
|
|
||
| tl_block_s = k_start + b_i * block_N | ||
| tl_block_e = T.min(k_start + (b_i + 1) * block_N, k_end) | ||
| T.copy(K[tl_block_s : tl_block_s + block_N, :], index_k) | ||
| for bn_i in T.Parallel(block_N): | ||
| scale[bn_i] = KScale[tl_block_s + bn_i] | ||
|
|
||
| for bn_i, d_i in T.Parallel(block_N, dim): | ||
| index_k[bn_i, d_i] = index_k[bn_i, d_i] * scale[bn_i] | ||
|
|
||
| cur_tl_block_size = tl_block_e - tl_block_s | ||
| for n_i in T.parallel(block_N): | ||
| for d_i in T.parallel(dim): | ||
| if n_i >= cur_tl_block_size: | ||
| index_k[n_i, d_i] = T.cast(0, accum_dtype) | ||
|
|
||
| T.reduce_sum(index_k, acc, dim=0, clear=False) |
There was a problem hiding this comment.
Avoid reading past ragged pooling tiles.
The kernel zeroes invalid lanes after the full T.copy and KScale loads. For the final pooling block, tl_block_s + bn_i can exceed seq_len_k, so Lines 53 and 55 can read out of bounds before Lines 60-65 clear the lane. Use guarded loads or pad K/KScale to the next block_N boundary.
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@examples/dsa_hisa/fp8_block_mean_pooling.py` around lines 48 - 66, The loop
currently does an unconditional T.copy and unconditional KScale loads which can
read past seq_len_k for the last pooling tile; change the loads to be guarded by
cur_tl_block_size (or pad K/KScale to block_N) so you only copy/assign up to
cur_tl_block_size entries and set the remaining lanes to zero. Specifically,
replace the unconditional T.copy(K[tl_block_s : tl_block_s + block_N, :],
index_k) and the for bn_i in T.Parallel(block_N): scale[bn_i] =
KScale[tl_block_s + bn_i] with guarded/limited loads that use cur_tl_block_size
(or conditional bn_i < cur_tl_block_size) to avoid indexing beyond seq_len_k,
and ensure any untouched index_k/scale lanes are explicitly zeroed so the
subsequent index_k scaling and T.reduce_sum are safe.
|
|
||
| # ------------------------------------------------------------------ | ||
| # Stage 2: fp8 fine-grained Q·K MQA over only the selected | ||
| # blocks' raw tokens (block_topk_eff blocks × k_block_size tokens |
There was a problem hiding this comment.
Fix the remaining Ruff RUF003 comment warnings.
Replace Unicode multiplication signs with ASCII x in comments.
🧹 Proposed fix
- # blocks' raw tokens (block_topk_eff blocks × k_block_size tokens
+ # blocks' raw tokens (block_topk_eff blocks x k_block_size tokens
...
- # × k_block_size candidate tokens. Gives per-query slot ids.
+ # x k_block_size candidate tokens. Gives per-query slot ids.
...
- # slot = block_id_in_topk × k_block_size + offset_in_block
+ # slot = block_id_in_topk x k_block_size + offset_in_block
...
- # absolute_k = topk_block_indices[m, block_id_in_topk] × k_block_size + offset_in_block
+ # absolute_k = topk_block_indices[m, block_id_in_topk] x k_block_size + offset_in_block
...
- # topk_tokens) (clipped by K range and by block_topk × k_block_size).
+ # topk_tokens) (clipped by K range and by block_topk x k_block_size).Also applies to: 99-99, 114-116, 192-192
🧰 Tools
🪛 Ruff (0.15.10)
[warning] 82-82: Comment contains ambiguous × (MULTIPLICATION SIGN). Did you mean x (LATIN SMALL LETTER X)?
(RUF003)
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@examples/dsa_hisa/hisa.py` at line 82, Replace the Unicode multiplication
character (×) with the ASCII letter "x" in the comments inside
hisa.py—specifically update the comment "blocks' raw tokens (block_topk_eff
blocks × k_block_size tokens" and the other comment occurrences that use "×"
(also present in the nearby comments around the token/block size descriptions
and later comment at the end of the file). Search for comments containing the
"×" glyph and change them to use "x" so they satisfy Ruff RUF003 while keeping
the original comment wording intact.
| topk_tokens_eff = min(topk_tokens, block_sparse_logits.shape[-1]) | ||
| relevant_topk_indices = torch.topk( | ||
| block_sparse_logits, | ||
| k=topk_tokens_eff, | ||
| dim=-1, | ||
| ).indices # [M, topk_tokens_eff] int64 | ||
|
|
||
| # ------------------------------------------------------------------ | ||
| # Stage 3 (post, Python): translate slot ids → absolute K token | ||
| # position → per-query relative offset (matches vLLM indexer | ||
| # output buffer). Slots whose relative offset falls outside the | ||
| # query's visible range are set to -1. | ||
| # ------------------------------------------------------------------ | ||
| # slot = block_id_in_topk × k_block_size + offset_in_block | ||
| # where block_id_in_topk ∈ [0, block_topk_eff) | ||
| # absolute_k = topk_block_indices[m, block_id_in_topk] × k_block_size + offset_in_block | ||
| absolute_topk_block_indices = torch.gather( | ||
| topk_block_indices, | ||
| dim=-1, | ||
| index=(relevant_topk_indices // k_block_size), | ||
| ) | ||
| topk_indices = absolute_topk_block_indices * k_block_size + (relevant_topk_indices % k_block_size) | ||
| topk_indices = topk_indices.to(torch.int32) | ||
|
|
||
| # Relative to this query's K start. | ||
| topk_indices -= cu_seqlen_ks[:, None] | ||
| mask_lo = topk_indices >= 0 | ||
| mask_hi = topk_indices - (cu_seqlen_ke - cu_seqlen_ks)[:, None] < 0 | ||
| mask = mask_lo & mask_hi | ||
| topk_indices = topk_indices.masked_fill(~mask, -1) | ||
|
|
||
| return topk_indices |
There was a problem hiding this comment.
Preserve the documented [M, topk_tokens] output shape.
When topk_tokens > block_sparse_logits.shape[-1], Line 101 clips the selection and the function returns [M, topk_tokens_eff], despite the docstring promising [M, topk_tokens] with invalid slots set to -1. Pad the result after masking.
🐛 Proposed fix
topk_indices = topk_indices.masked_fill(~mask, -1)
+ if topk_tokens_eff < topk_tokens:
+ topk_indices = torch.cat(
+ [
+ topk_indices,
+ topk_indices.new_full((topk_indices.shape[0], topk_tokens - topk_tokens_eff), -1),
+ ],
+ dim=-1,
+ )
return topk_indices🧰 Tools
🪛 Ruff (0.15.10)
[warning] 114-114: Comment contains ambiguous × (MULTIPLICATION SIGN). Did you mean x (LATIN SMALL LETTER X)?
(RUF003)
[warning] 116-116: Comment contains ambiguous × (MULTIPLICATION SIGN). Did you mean x (LATIN SMALL LETTER X)?
(RUF003)
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@examples/dsa_hisa/hisa.py` around lines 101 - 132, The function currently
returns topk_indices with shape [M, topk_tokens_eff] when topk_tokens >
block_sparse_logits.shape[-1], but the docstring promises [M, topk_tokens] with
invalid slots set to -1; after computing and masking topk_indices (symbols:
topk_tokens, topk_tokens_eff, block_sparse_logits, relevant_topk_indices,
topk_indices), if topk_tokens_eff < topk_tokens create a new int32 tensor filled
with -1 of shape [M, topk_tokens] and copy the existing topk_indices into the
leftmost topk_tokens_eff columns (leave remaining columns as -1), then return
that padded tensor so the function always returns shape [M, topk_tokens].
| for bq_i in T.serial(block_Q): | ||
| cu_k_s_min = T.min(cu_k_s_min, T.min(CuSeqLenBlockedKS[seq_len_i + bq_i], seq_len_blocked_kv)) | ||
| for bq_i in T.serial(block_Q): | ||
| cu_k_e_max = T.max(cu_k_e_max, T.min(CuSeqLenBlockedKE[seq_len_i + bq_i], seq_len_blocked_kv)) | ||
|
|
||
| T.copy(IndexQ[seq_len_i * heads, 0], index_q_shared) | ||
| T.copy(Weights[seq_len_i, 0], weights) | ||
|
|
||
| for nbn_i in T.Pipelined(T.ceildiv(cu_k_e_max - cu_k_s_min, block_N), num_stages=num_stages): | ||
| T.copy(IndexBlockedK[cu_k_s_min + nbn_i * block_N, 0], index_k_shared) | ||
| T.copy(IndexBlockedKScale[cu_k_s_min + nbn_i * block_N], index_k_scale_fragment) | ||
|
|
||
| T.gemm( | ||
| index_k_shared, | ||
| index_q_shared, | ||
| s, | ||
| transpose_B=True, | ||
| clear_accum=True, | ||
| policy=T.GemmWarpPolicy.FullCol, | ||
| ) | ||
|
|
||
| for bn_i, bq_i, h_i in T.Parallel(block_N, block_Q, heads): | ||
| s_reshaped[bn_i, bq_i, h_i] = T.max(s_reshaped[bn_i, bq_i, h_i] * index_k_scale_fragment[bn_i], 0) * weights[bq_i, h_i] | ||
|
|
||
| T.reduce_sum(s_reshaped, logits, dim=-1, clear=True) | ||
|
|
||
| for bq_i, bn_i in T.Parallel(block_Q, block_N): | ||
| Logits[seq_len_i + bq_i, cu_k_s_min + nbn_i * block_N + bn_i] = logits[bn_i, bq_i] |
There was a problem hiding this comment.
Handle partial query and K tiles before copying/storing.
This kernel assumes full block_Q query tiles and full block_N K tiles. In the end-to-end path, seq_len_blocked_kv can be smaller than the default block_N=256, so Lines 88-89 read past blocked_kv_fp8/blocked_kv_scale, and Line 106 writes past Logits. The final partial query tile has the same issue at Lines 80-85.
Pad inputs to tile multiples before launch or add guarded loads/stores in the kernel.
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@examples/dsa_hisa/pool_mqa_fp8.py` around lines 79 - 106, The kernel
currently assumes full tiles for queries and K (block_Q, block_N) causing
out-of-bounds reads from blocked_kv_fp8/blocked_kv_scale and writes to Logits
when seq_len_blocked_kv or the final query tile is partial; fix by adding bounds
checks and guarded loads/stores or by computing padded tile extents before the
loops: clamp cu_k_s_min/cu_k_e_max and the per-tile index offsets using
seq_len_blocked_kv and CuSeqLenBlockedKS/CuSeqLenBlockedKE, use guarded
copy/load for IndexBlockedK and IndexBlockedKScale when cu_k_s_min +
nbn_i*block_N + bn_i >= seq_len_blocked_kv, and only write to Logits when
seq_len_i + bq_i and the K index are within valid lengths (use
CuSeqLenBlockedKS/KE and seq_len_blocked_kv to decide); alternatively ensure
inputs are pre-padded to block_Q/block_N multiples before kernel launch.
| ref = ref_clean_and_maintain_logits(ref, cu_blocked_ks, cu_blocked_ke) | ||
|
|
||
| # After the mask, +/-inf positions must agree exactly. Compare the | ||
| # remaining finite values under an fp8×fp8 GEMM tolerance. |
There was a problem hiding this comment.
Fix the remaining Ruff RUF003 comment warnings.
Use ASCII x instead of the Unicode multiplication sign in these comments so lint stays clean.
🧹 Proposed fix
- # remaining finite values under an fp8×fp8 GEMM tolerance.
+ # remaining finite values under an fp8 x fp8 GEMM tolerance.
...
- # FLOPs: fp8×fp8 GEMM dominates = 2 * M * H * Nb * D (mul+add).
+ # FLOPs: fp8 x fp8 GEMM dominates = 2 * M * H * Nb * D (mul+add).
...
- # M × k_block_size^-1 must be a multiple of block_N=256.
+ # M x k_block_size^-1 must be a multiple of block_N=256.Also applies to: 239-239, 246-246
🧰 Tools
🪛 Ruff (0.15.10)
[warning] 219-219: Comment contains ambiguous × (MULTIPLICATION SIGN). Did you mean x (LATIN SMALL LETTER X)?
(RUF003)
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@examples/dsa_hisa/pool_mqa_fp8.py` at line 219, Several comments in
examples/dsa_hisa/pool_mqa_fp8.py use the Unicode multiplication sign (×) which
triggers Ruff RUF003; replace the Unicode × with ASCII x in those comments
(e.g., change "fp8×fp8 GEMM" to "fp8xfp8 GEMM" and any other occurrences on the
same comment lines). Search for the '×' character in the file and update the
comment text accordingly, keeping only ASCII characters to silence the RUF003
warnings.
|
|
||
| **What it does**: for each query `m` and each pool block `n` in | ||
| `[cu_seqlen_blocked_ks[m], cu_seqlen_blocked_ke[m])`, | ||
| ``` |
There was a problem hiding this comment.
Add languages to fenced code blocks.
Markdownlint MD040 flags these fences. Use text for formula/pipeline blocks.
🧹 Proposed fix
-```
+```text
block_k_score[m, n] = sum_h ReLU(q[m, h] · blocked_k[n]) * blocked_k_scale[n] * weights[m, h]...
- +text
k_abs = blk * kv_block_size + i
if k_abs ∉ [cu_seqlen_ks[m], cu_seqlen_ke[m]) or k_abs >= N:
block_sparse_logits[m, t * kv_block_size + i] = -inf
else:
block_sparse_logits[m, t * kv_block_size + i] =
sum_h ReLU(q[m, h] · k[k_abs]) * k_scale[k_abs] * weights[m, h]
...
-```
+```text
(1.1) fp8_native_block_mean_pooling K, k_scale → blocked_k, blocked_k_scale
Also applies to: 149-149, 187-187
🧰 Tools
🪛 markdownlint-cli2 (0.22.0)
[warning] 87-87: Fenced code blocks should have a language specified
(MD040, fenced-code-language)
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@examples/dsa_hisa/README.md` at line 87, Update the three fenced code blocks
containing the snippets starting with "block_k_score[m, n] = ..." , the one
containing "k_abs = blk * kv_block_size + i" and the line "(1.1)
fp8_native_block_mean_pooling K, k_scale → blocked_k,
blocked_k_scale" to include a language identifier by replacing the opening ```
with ```text so the fences are annotated (e.g., ```text block_k_score...,
```text k_abs = ..., ```text (1.1) fp8_native_block_mean_pooling ...).
|
@xuyufei-a Thanks for your contributions! Would you mind refactoring this example into eager style? In this style, we do not require T.prim_func. cc: https://github.com/tile-ai/tilelang/blob/main/examples/eager_jit/eagerjit.zh.ipynb |
There was a problem hiding this comment.
Actionable comments posted: 1
♻️ Duplicate comments (6)
examples/dsa_hisa/pool_mqa_fp8.py (2)
219-246:⚠️ Potential issue | 🟡 MinorRuff RUF003: replace Unicode
×in comments.Lines 219, 239, 246 still use
×(U+00D7). Use ASCIIxto clear the lint.🧹 Proposed fix
- # remaining finite values under an fp8×fp8 GEMM tolerance. + # remaining finite values under an fp8 x fp8 GEMM tolerance. ... - # FLOPs: fp8×fp8 GEMM dominates = 2 * M * H * Nb * D (mul+add). + # FLOPs: fp8 x fp8 GEMM dominates = 2 * M * H * Nb * D (mul+add). ... - # M × k_block_size^-1 must be a multiple of block_N=256. + # M x k_block_size^-1 must be a multiple of block_N=256.🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@examples/dsa_hisa/pool_mqa_fp8.py` around lines 219 - 246, Summary: Replace the Unicode multiplication sign (×, U+00D7) used in comments with ASCII 'x' to satisfy Ruff RUF003. Fix: edit the comment strings and any inline comments around the code using pool_mqa_attn_return_logits_fp8_interface, do_bench, and the module-level guard so that occurrences like "M × k_block_size^-1" and any prints or comment lines use "x" instead of "×"; do not change logic or variable names, only replace the character in comments and string literals. Ensure all three occurrences noted near the benchmarking block and the module doc/guard are updated.
79-106:⚠️ Potential issue | 🔴 CriticalPartial Q / K tiles still unguarded.
- Lines 79-82: when
seq_len % block_Q != 0, the tail tile readsCuSeqLenBlockedKS/KE[seq_len_i + bq_i]pastseq_len.- Line 84-85:
T.copy(IndexQ[seq_len_i * heads, 0], index_q_shared)/Weights[seq_len_i, 0]copies a fullblock_Q*headsrows and can overrunseq_len*heads.- Lines 88-89: when
seq_len_blocked_kv < block_N(e.g., HISA passes smallNb),T.copy(IndexBlockedK[cu_k_s_min + nbn_i*block_N, 0], ...)reads past the K buffer.- Line 106:
Logits[seq_len_i + bq_i, cu_k_s_min + nbn_i*block_N + bn_i]writes pastseq_len/seq_len_blocked_kvin the same conditions.The current tests dodge this via
M % num_seqs == 0andN_blocked % block_N == 0asserts, but any non-tile-aligned call fromhisa.pywill OOB. Add guarded loads/stores keyed toseq_len/seq_len_blocked_kv, or pad the inputs before launch.🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@examples/dsa_hisa/pool_mqa_fp8.py` around lines 79 - 106, The loops and copies (using symbols CuSeqLenBlockedKS, CuSeqLenBlockedKE, IndexQ, Weights, IndexBlockedK, Logits) can read/write past the valid sequence/KV ranges when seq_len % block_Q != 0 or seq_len_blocked_kv < block_N; fix by adding explicit bounds guards or masked/padded copies: before any T.copy or array read use min/conditional logic against seq_len and seq_len_blocked_kv (e.g., clamp seq_len_i + bq_i and cu_k_s_min+nbn_i*block_N + bn_i), and for writes to Logits only emit stores when the target index < seq_len_blocked_kv and seq_len; alternatively pad IndexQ/Weights/IndexBlockedK inputs to full tile sizes before launch. Ensure guards are applied for the T.copy(IndexQ..., Weights...), T.copy(IndexBlockedK...), uses of CuSeqLenBlockedKS/KE, and the final Logits[...] write so no OOB accesses occur.examples/dsa_hisa/fp8_block_mean_pooling.py (1)
48-63:⚠️ Potential issue | 🔴 CriticalRagged pooling tile still reads past
seq_len_k.For the last pool block when
seq_len_k % pooling_block_size != 0,T.copy(K[tl_block_s : tl_block_s + block_N, :], index_k)(Line 50) andKScale[tl_block_s + bn_i](Line 52) load lanes withtl_block_s + bn_i >= seq_len_kbefore then_i >= cur_tl_block_sizezero-fill on Lines 58-61 ever runs. Guard the loads (or pad K/KScale to the nextblock_Nboundary) so no index pastseq_len_kis dereferenced; zeroing after the fact does not prevent the read.🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@examples/dsa_hisa/fp8_block_mean_pooling.py` around lines 48 - 63, The code reads past seq_len_k in the last pooling block; compute cur_tl_block_size = tl_block_e - tl_block_s before doing any loads and guard the copies and KScale loads so you only read lanes with bn_i < cur_tl_block_size (or alternatively pad K/KScale to block_N); specifically, change the T.copy(K[...], index_k) and KScale[tl_block_s + bn_i] accesses to be conditional on bn_i < cur_tl_block_size (or use a safe masked copy), ensuring index_k and scale for out-of-range bn_i are set to zero before the subsequent multiply and T.reduce_sum (referencing tl_block_s, tl_block_e, block_N, cur_tl_block_size, T.copy, K, KScale, scale, index_k).examples/dsa_hisa/clean_and_maintain_logits.py (1)
31-37:⚠️ Potential issue | 🔴 CriticalTail-index OOB write still unresolved.
With
block_K=4096(default) and HISA call sites passing much smallerseq_len_kv(e.g.,Nb),idx = n_i*block_K + k_i*threads + txcan exceedseq_len_kv - 1in the last pipelined iteration, causing OOB writes toLogits[bx, idx]. The previously suggested guard (idx < seq_len_kv) is still needed.🛡️ Proposed fix
for n_i in T.Pipelined(T.ceildiv(seq_len_kv, block_K)): - for k_i in T.serial(block_K // threads): - idx = n_i * block_K + k_i * threads + tx - if idx == cu_k_s or idx == cu_k_e - 1: - Logits[bx, idx] = T.infinity(dtype) - if idx < cu_k_s or idx >= cu_k_e: - Logits[bx, idx] = -T.infinity(dtype) + for k_i in T.serial(T.ceildiv(block_K, threads)): + idx = n_i * block_K + k_i * threads + tx + if idx < seq_len_kv: + if idx == cu_k_s or idx == cu_k_e - 1: + Logits[bx, idx] = T.infinity(dtype) + if idx < cu_k_s or idx >= cu_k_e: + Logits[bx, idx] = -T.infinity(dtype)🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@examples/dsa_hisa/clean_and_maintain_logits.py` around lines 31 - 37, The loop can compute idx beyond seq_len_kv-1, causing OOB writes to Logits; update the write conditions in the nested loops (where idx = n_i * block_K + k_i * threads + tx is computed) to only assign Logits[bx, idx] when idx < seq_len_kv (e.g., add idx < seq_len_kv to both the infinity and -infinity branches or combine checks so writes occur only when cu_k_s/cu_k_e logic AND idx < seq_len_kv are true); ensure the check uses the same dtype/variables (seq_len_kv, block_K, threads, tx, cu_k_s, cu_k_e, Logits, bx) so the tail pipelined iteration cannot write out of bounds.examples/dsa_hisa/block_sparse_mqa_fp8.py (2)
144-144:⚠️ Potential issue | 🟡 MinorResolve the remaining Ruff warnings.
His unused, and the FLOPs comment still contains Unicode multiplication signs flagged by Ruff.🧹 Proposed lint cleanup
- M, H, D = q_fp8.shape + M, _H, D = q_fp8.shape ... - # FLOPs: M × topk × kv_block_size × H × D (fp8×fp8) × 2 (mul+add). + # FLOPs: M x topk x kv_block_size x H x D (fp8 x fp8) x 2 (mul+add).Also applies to: 249-249
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@examples/dsa_hisa/block_sparse_mqa_fp8.py` at line 144, The tuple unpacking M, H, D = q_fp8.shape introduces an unused variable H and the FLOPs comment contains Unicode multiplication signs; change the unpacking to M, _, D = q_fp8.shape (or M, H_unused, D) to mark H as intentionally unused and update any FLOPs comments that use “×” to use the ASCII asterisk "*" instead so Ruff no longer flags them; search for q_fp8.shape and the FLOPs comment strings in this file (and the other occurrence around line 249) and apply these two edits.
77-101:⚠️ Potential issue | 🔴 CriticalGuard partial final KV blocks before reading K.
Line 77 and Line 79 still read a full
block_Nbefore the range mask runs. When the selected block is the ragged final block, this can read pastseq_len_kv.🐛 Proposed guard for ragged blocks
- T.copy(IndexK[block_s_i : block_s_i + block_N, :], index_k_shared) - for bn_i in T.Parallel(block_N): - scale_shared[bn_i] = IndexKScale[block_s_i + bn_i] + for bn_i, d_i in T.Parallel(block_N, index_dim): + k_i = block_s_i + bn_i + if k_i < seq_len_kv: + index_k_shared[bn_i, d_i] = IndexK[k_i, d_i] + else: + index_k_shared[bn_i, d_i] = T.cast(0, fp8_dtype) + for bn_i in T.Parallel(block_N): + k_i = block_s_i + bn_i + if k_i < seq_len_kv: + scale_shared[bn_i] = IndexKScale[k_i] + else: + scale_shared[bn_i] = T.cast(0, accum_dtype) ... - if k_i < cu_k_s_min or k_i >= cu_k_e_max: + if k_i < cu_k_s_min or k_i >= cu_k_e_max or k_i >= seq_len_kv: logits[i_i, 0] = -T.infinity(accum_dtype)🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@examples/dsa_hisa/block_sparse_mqa_fp8.py` around lines 77 - 101, The code reads a full block_N in IndexK and related loops even for the final ragged KV block; compute valid_N = max(0, min(block_N, seq_len_kv - block_s_i)) and use that to limit T.copy into index_k_shared and to bound all T.Parallel loops (the scale_shared fill from IndexKScale, the s_reshaped write loop, the per-item cu_k range check that writes logits, and the final Logits write) so no reads or writes occur past seq_len_kv; keep the existing cu_k_s_min/cu_k_e_max check but ensure iterations only go 0..valid_N-1 when referencing index_k_shared, scale_shared, s_reshaped, logits, and Logits to avoid out-of-range K reads.
🧹 Nitpick comments (1)
examples/dsa_hisa/clean_and_maintain_logits.py (1)
100-109: Bench includestorch.randnallocation.
fn()allocates a fresh[M, N]f32 tensor each iteration, somsand the reported GB/s conflate allocation + RNG + the kernel (which is mostly a light mask). Allocate once outside and clone/reset in-place, or drop the randn and just re-run the kernel on a pre-allocated buffer, so the measurement reflects the kernel.🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@examples/dsa_hisa/clean_and_maintain_logits.py` around lines 100 - 109, The benchmark currently allocates a fresh tensor inside fn() which mixes allocation/RNG overhead into the timing; move the torch.randn(M, N, device="cuda", dtype=torch.float32) allocation outside of fn() (e.g., create a reusable buffer before calling do_bench) and inside fn() either clone/reset the pre-allocated buffer in-place or simply call clean_and_maintain_logits_interface(logits, cu_ks, cu_ke) on that pre-allocated tensor so do_bench measures only the kernel work performed by clean_and_maintain_logits_interface for the [M, N] buffer.
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Inline comments:
In `@examples/dsa_hisa/block_sparse_mqa_fp8.py`:
- Around line 229-232: The current test masks out NaNs via `finite` so a kernel
that produces NaN where `ref` is finite can pass; add an explicit NaN mask
assertion before using `finite`: assert that torch.isnan(got) equals
torch.isnan(ref) (e.g., raise "NaN mask differs") so any NaN/finite mismatch
fails, then keep the existing infinity-mask checks and the finite-element
comparison using torch.testing.assert_close on got[finite] vs ref[finite].
---
Duplicate comments:
In `@examples/dsa_hisa/block_sparse_mqa_fp8.py`:
- Line 144: The tuple unpacking M, H, D = q_fp8.shape introduces an unused
variable H and the FLOPs comment contains Unicode multiplication signs; change
the unpacking to M, _, D = q_fp8.shape (or M, H_unused, D) to mark H as
intentionally unused and update any FLOPs comments that use “×” to use the ASCII
asterisk "*" instead so Ruff no longer flags them; search for q_fp8.shape and
the FLOPs comment strings in this file (and the other occurrence around line
249) and apply these two edits.
- Around line 77-101: The code reads a full block_N in IndexK and related loops
even for the final ragged KV block; compute valid_N = max(0, min(block_N,
seq_len_kv - block_s_i)) and use that to limit T.copy into index_k_shared and to
bound all T.Parallel loops (the scale_shared fill from IndexKScale, the
s_reshaped write loop, the per-item cu_k range check that writes logits, and the
final Logits write) so no reads or writes occur past seq_len_kv; keep the
existing cu_k_s_min/cu_k_e_max check but ensure iterations only go 0..valid_N-1
when referencing index_k_shared, scale_shared, s_reshaped, logits, and Logits to
avoid out-of-range K reads.
In `@examples/dsa_hisa/clean_and_maintain_logits.py`:
- Around line 31-37: The loop can compute idx beyond seq_len_kv-1, causing OOB
writes to Logits; update the write conditions in the nested loops (where idx =
n_i * block_K + k_i * threads + tx is computed) to only assign Logits[bx, idx]
when idx < seq_len_kv (e.g., add idx < seq_len_kv to both the infinity and
-infinity branches or combine checks so writes occur only when cu_k_s/cu_k_e
logic AND idx < seq_len_kv are true); ensure the check uses the same
dtype/variables (seq_len_kv, block_K, threads, tx, cu_k_s, cu_k_e, Logits, bx)
so the tail pipelined iteration cannot write out of bounds.
In `@examples/dsa_hisa/fp8_block_mean_pooling.py`:
- Around line 48-63: The code reads past seq_len_k in the last pooling block;
compute cur_tl_block_size = tl_block_e - tl_block_s before doing any loads and
guard the copies and KScale loads so you only read lanes with bn_i <
cur_tl_block_size (or alternatively pad K/KScale to block_N); specifically,
change the T.copy(K[...], index_k) and KScale[tl_block_s + bn_i] accesses to be
conditional on bn_i < cur_tl_block_size (or use a safe masked copy), ensuring
index_k and scale for out-of-range bn_i are set to zero before the subsequent
multiply and T.reduce_sum (referencing tl_block_s, tl_block_e, block_N,
cur_tl_block_size, T.copy, K, KScale, scale, index_k).
In `@examples/dsa_hisa/pool_mqa_fp8.py`:
- Around line 219-246: Summary: Replace the Unicode multiplication sign (×,
U+00D7) used in comments with ASCII 'x' to satisfy Ruff RUF003. Fix: edit the
comment strings and any inline comments around the code using
pool_mqa_attn_return_logits_fp8_interface, do_bench, and the module-level guard
so that occurrences like "M × k_block_size^-1" and any prints or comment lines
use "x" instead of "×"; do not change logic or variable names, only replace the
character in comments and string literals. Ensure all three occurrences noted
near the benchmarking block and the module doc/guard are updated.
- Around line 79-106: The loops and copies (using symbols CuSeqLenBlockedKS,
CuSeqLenBlockedKE, IndexQ, Weights, IndexBlockedK, Logits) can read/write past
the valid sequence/KV ranges when seq_len % block_Q != 0 or seq_len_blocked_kv <
block_N; fix by adding explicit bounds guards or masked/padded copies: before
any T.copy or array read use min/conditional logic against seq_len and
seq_len_blocked_kv (e.g., clamp seq_len_i + bq_i and cu_k_s_min+nbn_i*block_N +
bn_i), and for writes to Logits only emit stores when the target index <
seq_len_blocked_kv and seq_len; alternatively pad IndexQ/Weights/IndexBlockedK
inputs to full tile sizes before launch. Ensure guards are applied for the
T.copy(IndexQ..., Weights...), T.copy(IndexBlockedK...), uses of
CuSeqLenBlockedKS/KE, and the final Logits[...] write so no OOB accesses occur.
---
Nitpick comments:
In `@examples/dsa_hisa/clean_and_maintain_logits.py`:
- Around line 100-109: The benchmark currently allocates a fresh tensor inside
fn() which mixes allocation/RNG overhead into the timing; move the
torch.randn(M, N, device="cuda", dtype=torch.float32) allocation outside of fn()
(e.g., create a reusable buffer before calling do_bench) and inside fn() either
clone/reset the pre-allocated buffer in-place or simply call
clean_and_maintain_logits_interface(logits, cu_ks, cu_ke) on that pre-allocated
tensor so do_bench measures only the kernel work performed by
clean_and_maintain_logits_interface for the [M, N] buffer.
🪄 Autofix (Beta)
Fix all unresolved CodeRabbit comments on this PR:
- Push a commit to this branch (recommended)
- Create a new PR with the fixes
ℹ️ Review info
⚙️ Run configuration
Configuration used: defaults
Review profile: CHILL
Plan: Pro
Run ID: 97ac5913-80aa-4618-9d31-a955b9054041
📒 Files selected for processing (4)
examples/dsa_hisa/block_sparse_mqa_fp8.pyexamples/dsa_hisa/clean_and_maintain_logits.pyexamples/dsa_hisa/fp8_block_mean_pooling.pyexamples/dsa_hisa/pool_mqa_fp8.py
| finite = torch.isfinite(got) & torch.isfinite(ref) | ||
| assert torch.equal(torch.isposinf(got), torch.isposinf(ref)), "pos-inf mask differs" | ||
| assert torch.equal(torch.isneginf(got), torch.isneginf(ref)), "neg-inf mask differs" | ||
| torch.testing.assert_close(got[finite], ref[finite], rtol=1e-1, atol=2e-1) |
There was a problem hiding this comment.
Fail the correctness test on NaNs.
finite excludes NaN mismatches, and the infinity-mask checks do not catch them. A kernel producing NaN for a finite reference value can pass this test.
🧪 Proposed assertion fix
finite = torch.isfinite(got) & torch.isfinite(ref)
+ assert not torch.isnan(got).any().item(), "got contains NaNs"
+ assert not torch.isnan(ref).any().item(), "ref contains NaNs"
assert torch.equal(torch.isposinf(got), torch.isposinf(ref)), "pos-inf mask differs"
assert torch.equal(torch.isneginf(got), torch.isneginf(ref)), "neg-inf mask differs"
torch.testing.assert_close(got[finite], ref[finite], rtol=1e-1, atol=2e-1)📝 Committable suggestion
‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.
| finite = torch.isfinite(got) & torch.isfinite(ref) | |
| assert torch.equal(torch.isposinf(got), torch.isposinf(ref)), "pos-inf mask differs" | |
| assert torch.equal(torch.isneginf(got), torch.isneginf(ref)), "neg-inf mask differs" | |
| torch.testing.assert_close(got[finite], ref[finite], rtol=1e-1, atol=2e-1) | |
| finite = torch.isfinite(got) & torch.isfinite(ref) | |
| assert not torch.isnan(got).any().item(), "got contains NaNs" | |
| assert not torch.isnan(ref).any().item(), "ref contains NaNs" | |
| assert torch.equal(torch.isposinf(got), torch.isposinf(ref)), "pos-inf mask differs" | |
| assert torch.equal(torch.isneginf(got), torch.isneginf(ref)), "neg-inf mask differs" | |
| torch.testing.assert_close(got[finite], ref[finite], rtol=1e-1, atol=2e-1) |
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@examples/dsa_hisa/block_sparse_mqa_fp8.py` around lines 229 - 232, The
current test masks out NaNs via `finite` so a kernel that produces NaN where
`ref` is finite can pass; add an explicit NaN mask assertion before using
`finite`: assert that torch.isnan(got) equals torch.isnan(ref) (e.g., raise "NaN
mask differs") so any NaN/finite mismatch fails, then keep the existing
infinity-mask checks and the finite-element comparison using
torch.testing.assert_close on got[finite] vs ref[finite].
|
@LeiWang1999, is this failure caused by the new HISA example? What should I do next? |
|
@xuyufei-a This PR only adds a new example, so theoretically it should not affect the tests. It might be due to random fluctuations in our CI. Let me take a look. 👀 |
Summary
Adds
examples/dsa_hisa/— a Tilelang prefill implementation of HISA(Hierarchical Indexed Sparse Attention), a plug-and-play replacement for the DeepSeek sparse attention indexer that rewrites the flat-token search path into a two-stage hierarchical procedure.Paper: https://arxiv.org/pdf/2603.28458
What's in this PR
fp8_block_mean_pooling.pypool_mqa_fp8.pyclean_and_maintain_logits.pyblock_sparse_mqa_fp8.pyblock_topkselected blockshisa.pytorch.topk+ index-translation post-processing into a singlehisa_indexerentry pointtilelang_utils.pyprepare_ks_ke_from_cu_seqlens,per_custom_dims_cast_to_fp8, etc.)README.mdPipeline
Stage 1 — coarse block-level selection. Group K tokens into pool blocks of
k_block_size, mean-pool each block, then score each query against all pool blocks and pick the topblock_topkblocks per query.Stage 2 — fine-grained token-level scoring. For each query, run a full-resolution fp8 MQA over the raw tokens inside its selected blocks, then pick the top
topk_tokenstokens per query.Summary by CodeRabbit
Documentation
New Features
Tools & Utilities
Tests / Benchmarks