Skip to content

[Example] Add HISA: hierarchical sparse attention indexer - #2069

Merged
SiriusNEO merged 3 commits into
tile-ai:mainfrom
xuyufei-a:hisa
Apr 25, 2026
Merged

SiriusNEO merged 3 commits into
tile-ai:mainfrom
xuyufei-a:hisa

Conversation

@xuyufei-a

@xuyufei-a xuyufei-a commented Apr 20, 2026 •

Copy link
Copy Markdown
Contributor

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

File Step Role
fp8_block_mean_pooling.py 1.1 Mean-pool raw K into pool blocks (fp8 + per-block f32 scale)
pool_mqa_fp8.py 1.2 fp8×fp8 score of Q against pooled K → one logit per (query, pool block)
clean_and_maintain_logits.py 1.3 In-place mask on stage-1 logits: -inf outside per-query range, +inf at first/last valid block
block_sparse_mqa_fp8.py 2.1 fp8×fp8 fine-grained score over raw tokens of the top-block_topk selected blocks
hisa.py — End-to-end orchestration: chains all four kernels + two torch.topk + index-translation post-processing into a single hisa_indexer entry point
tilelang_utils.py — Helpers (prepare_ks_ke_from_cu_seqlens, per_custom_dims_cast_to_fp8, etc.)
README.md —

Pipeline

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 top block_topk blocks 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_tokens tokens per query.

Summary by CodeRabbit

  • Documentation

    • Added a comprehensive guide for the HISA hierarchical sparse-attention prefill pipeline and usage notes.
  • New Features

    • Added FP8 examples: block mean-pooling, pooled MQA, block-sparse MQA, in-place logits masking, and a full HISA indexer pipeline with top‑K selection.
  • Tools & Utilities

    • Added tensor preprocessing, FP8 quantization helpers, and sequence-index utilities.
  • Tests / Benchmarks

    • Added per-kernel and end-to-end tests and benchmarking harnesses validating correctness and measuring latency/throughput.

@github-actions

Copy link
Copy Markdown

👋 Hi! Thank you for contributing to the TileLang project.

Please remember to run pre-commit run --all-files in the root directory of the project to ensure your changes are properly linted and formatted. This will help ensure your contribution passes the format check.

We appreciate you taking this step! Our team will review your contribution, and we look forward to your awesome work! 🚀

@coderabbitai

coderabbitai Bot commented Apr 20, 2026 •

Copy link
Copy Markdown
Contributor
📝 Walkthrough

Walkthrough

Implements 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

Cohort / File(s) Summary
Documentation
examples/dsa_hisa/README.md
New README describing the two-stage HISA pipeline, kernel responsibilities, I/O contracts for hisa_indexer, testing strategy, and end-to-end behavior.
Utilities & Helpers
examples/dsa_hisa/tilelang_utils.py
Adds tensor caching decorator, cu-seqlen/lens/position/sequence/token index builders, FP8 casting helpers (UE8M0 option), error/similarity metrics, and randomized cu-seqlen generator.
Coarse Pooling
examples/dsa_hisa/fp8_block_mean_pooling.py
FP8 block mean-pooling JIT kernel + re-quantization and reference implementation; tests and benchmarks.
Pool MQA (coarse logits)
examples/dsa_hisa/pool_mqa_fp8.py
FP8 pooled MQA kernel computing block-level logits with per-block scaling, ReLU clamp, weight reduction; Python interface, reference, tests, benchmarks.
Logits Masking
examples/dsa_hisa/clean_and_maintain_logits.py
In-place kernel to set +inf at boundaries and -inf outside per-query ranges, with reference and validation harness.
Fine-grained MQA
examples/dsa_hisa/block_sparse_mqa_fp8.py
FP8 block-sparse MQA kernel producing per-token logits within top blocks, Python wrapper, reference, and tests/benchmarks.
Pipeline Orchestration & Tests
examples/dsa_hisa/hisa.py
hisa_indexer orchestrates pooling → block top-k → block-sparse token scoring → token top-k → index translation; includes test_hisa and benchmarks.

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)
Loading

Estimated code review effort

🎯 4 (Complex) | ⏱️ ~60 minutes

Possibly related PRs

Suggested reviewers

  • LeiWang1999
  • tzj-fxz

Poem

🐰
I hopped through FP8 fields of light,
Pooled the blocks and scored each sight,
Masked the edges, picked the best,
Token hops — a top-k quest!
Hisa hums, the indices take flight.

🚥 Pre-merge checks | ✅ 4 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 31.82% which is insufficient. The required threshold is 80.00%. Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (4 passed)
Check name Status Explanation
Title check ✅ Passed The title '[Example] Add HISA: hierarchical sparse attention indexer' clearly summarizes the main change—adding a complete HISA implementation example with all associated files and supporting utilities.
Linked Issues check ✅ Passed Check skipped because no linked issues were found for this pull request.
Out of Scope Changes check ✅ Passed Check skipped because no linked issues were found for this pull request.
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.

✏️ Tip: You can configure your own custom pre-merge checks in the settings.

✨ Finishing Touches
🧪 Generate unit tests (beta)
  • Create PR with unit tests

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.

❤️ Share

Comment @coderabbitai help to get the list of available commands and usage tips.

@coderabbitai coderabbitai Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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

📥 Commits

Reviewing files that changed from the base of the PR and between 99ee74b and 83e8467.

📒 Files selected for processing (7)
  • examples/dsa_hisa/README.md
  • examples/dsa_hisa/block_sparse_mqa_fp8.py
  • examples/dsa_hisa/clean_and_maintain_logits.py
  • examples/dsa_hisa/fp8_block_mean_pooling.py
  • examples/dsa_hisa/hisa.py
  • examples/dsa_hisa/pool_mqa_fp8.py
  • examples/dsa_hisa/tilelang_utils.py

Comment on lines +71 to +101
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]

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🔴 Critical

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

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🟡 Minor

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.

Comment thread examples/dsa_hisa/clean_and_maintain_logits.py Outdated
Comment on lines +48 to +66
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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🔴 Critical

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.

Comment thread examples/dsa_hisa/hisa.py

# ------------------------------------------------------------------
# Stage 2: fp8 fine-grained Q·K MQA over only the selected
# blocks' raw tokens (block_topk_eff blocks × k_block_size tokens

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🟡 Minor

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.

Comment thread examples/dsa_hisa/hisa.py
Comment on lines +101 to +132
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

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🟠 Major

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].

Comment thread examples/dsa_hisa/pool_mqa_fp8.py Outdated
Comment on lines +79 to +106
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]

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🔴 Critical

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.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🟡 Minor

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])`,
```

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🟡 Minor

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 ...).

Comment thread examples/dsa_hisa/tilelang_utils.py
@xuyufei-a xuyufei-a changed the title [examples] Add HISA: hierarchical sparse attention indexer [Example] Add HISA: hierarchical sparse attention indexer Apr 20, 2026
@LeiWang1999

Copy link
Copy Markdown
Member

@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

@coderabbitai coderabbitai Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Actionable comments posted: 1

♻️ Duplicate comments (6)
examples/dsa_hisa/pool_mqa_fp8.py (2)

219-246: ⚠️ Potential issue | 🟡 Minor

Ruff RUF003: replace Unicode × in comments.

Lines 219, 239, 246 still use × (U+00D7). Use ASCII x to 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 | 🔴 Critical

Partial Q / K tiles still unguarded.

  • Lines 79-82: when seq_len % block_Q != 0, the tail tile reads CuSeqLenBlockedKS/KE[seq_len_i + bq_i] past seq_len.
  • Line 84-85: T.copy(IndexQ[seq_len_i * heads, 0], index_q_shared) / Weights[seq_len_i, 0] copies a full block_Q*heads rows and can overrun seq_len*heads.
  • Lines 88-89: when seq_len_blocked_kv < block_N (e.g., HISA passes small Nb), 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 past seq_len / seq_len_blocked_kv in the same conditions.

The current tests dodge this via M % num_seqs == 0 and N_blocked % block_N == 0 asserts, but any non-tile-aligned call from hisa.py will OOB. Add guarded loads/stores keyed to seq_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 | 🔴 Critical

Ragged 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) and KScale[tl_block_s + bn_i] (Line 52) load lanes with tl_block_s + bn_i >= seq_len_k before the n_i >= cur_tl_block_size zero-fill on Lines 58-61 ever runs. Guard the loads (or pad K/KScale to the next block_N boundary) so no index past seq_len_k is 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 | 🔴 Critical

Tail-index OOB write still unresolved.

With block_K=4096 (default) and HISA call sites passing much smaller seq_len_kv (e.g., Nb), idx = n_i*block_K + k_i*threads + tx can exceed seq_len_kv - 1 in the last pipelined iteration, causing OOB writes to Logits[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 | 🟡 Minor

Resolve the remaining Ruff warnings.

H is 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 | 🔴 Critical

Guard partial final KV blocks before reading K.

Line 77 and Line 79 still read a full block_N before the range mask runs. When the selected block is the ragged final block, this can read past seq_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 includes torch.randn allocation.

fn() allocates a fresh [M, N] f32 tensor each iteration, so ms and 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

📥 Commits

Reviewing files that changed from the base of the PR and between 83e8467 and 4c57379.

📒 Files selected for processing (4)
  • examples/dsa_hisa/block_sparse_mqa_fp8.py
  • examples/dsa_hisa/clean_and_maintain_logits.py
  • examples/dsa_hisa/fp8_block_mean_pooling.py
  • examples/dsa_hisa/pool_mqa_fp8.py

Comment on lines +229 to +232
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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🟠 Major

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.

Suggested change
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].

@xuyufei-a

Copy link
Copy Markdown
Contributor Author

@LeiWang1999, is this failure caused by the new HISA example? What should I do next?

@SiriusNEO

Copy link
Copy Markdown
Collaborator

@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. 👀

@SiriusNEO
SiriusNEO merged commit 0ee6345 into tile-ai:main Apr 25, 2026
11 of 12 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants