Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
29 commits
Select commit Hold shift + click to select a range
255cd9d
[None][feat] Add CuTE DSL FP8 paged MQA logits kernel for Blackwell S…
limin2021 Apr 20, 2026
c32db8f
[None][fix] Fix docstring script filename in fp8_paged_mqa_logits.py
limin2021 Apr 20, 2026
b235eb1
[None][fix] Update copyright year in paged_mqa_logits __init__.py
limin2021 Apr 20, 2026
69d7170
[None][fix] Remove unused max_mem_gb argument from benchmark
limin2021 Apr 20, 2026
a2b9b43
[None][fix] Fix stream handling, add arch guard, and rename kernel class
limin2021 Apr 21, 2026
a2b0de7
[None][fix] Ensure CuTE DSL op registration when only logits kernel i…
limin2021 Apr 21, 2026
fe637c8
[None][refactor] Rename use_cute_dsl_logits to use_cute_dsl_paged_mqa…
limin2021 Apr 21, 2026
8f9b0e1
[None][fix] Remove redundant SM version check from DSA logits config
limin2021 Apr 21, 2026
cf418ab
[None][fix] Clean up DeepGEMM/DG-FullK references in docstrings and u…
limin2021 Apr 21, 2026
b1260fb
[None][refactor] Migrate MQA logits runner to fake tensor + TVM FFI
limin2021 Apr 21, 2026
4ef3d13
[None][fix] Remove commented-out debug code from fp8_paged_mqa_logits
limin2021 Apr 21, 2026
670cd16
[None][fix] Replace compile print with logger.debug in MQA logits
limin2021 Apr 21, 2026
0a1e93d
[None][fix] Add dtype validation to cute_dsl_fp8_paged_mqa_logits wra…
limin2021 Apr 22, 2026
fabc902
[None][test] Improve fp16 accuracy test and benchmark for MQA logits
limin2021 Apr 22, 2026
aeb32d8
[None][cleanup] Remove dead standalone test code and deduplicate dtyp…
limin2021 Apr 22, 2026
d0fa1cb
[None][fix] Validate num_heads divisible by 4 regardless of num_epi_s…
limin2021 Apr 22, 2026
38ec767
[None][cleanup] Remove unused variables and unnecessary noqa comments
limin2021 Apr 22, 2026
719b333
[None][fix] Add missing use_cute_dsl_paged_mqa_logits to test mock co…
limin2021 Apr 22, 2026
1a157f4
[None][fix] Fix TMA SMEM alignment for fp16 epilogue and support DSL …
limin2021 Apr 23, 2026
6681031
[None][fix] Fix OOB read in zero-work CTA and move test seed into helper
limin2021 Apr 23, 2026
1c68831
[None][feat] Support multi-block TMA for phys_block_kv < 128 in DSL p…
limin2021 Apr 23, 2026
862260c
[None][fix] Move shuffle before barrier acquire to match DeepGEMM sch…
limin2021 Apr 23, 2026
0331897
[None][cleanup] Remove commented-out test_deepgemm_fp8_paged_mqa_logits
limin2021 Apr 23, 2026
0ca0f27
Merge branch 'main' into add_dsl_indexer_gemm
limin2021 Apr 24, 2026
26effcc
[None][cleanup] Remove unused helpers in DSL paged MQA logits kernel …
limin2021 Apr 28, 2026
f34fb16
[None][fix] Skip DSL backend on non-SM100 archs in indexer decode test
limin2021 Apr 29, 2026
9db90f0
Merge remote-tracking branch 'github-upstream/main' into add_dsl_inde…
limin2021 May 1, 2026
0a84ecc
Merge remote-tracking branch 'github-upstream/main' into add_dsl_inde…
limin2021 May 7, 2026
4bcbd79
[None][test] Fix DSL indexer scheduler buffer aliasing and extend FP8…
limin2021 May 8, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
113 changes: 81 additions & 32 deletions tensorrt_llm/_torch/attention_backend/sparse/dsa.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,24 @@
hadamard_transform = None
HAS_FAST_HADAMARD = False

# `block_kv` arg passed to DeepGEMM's `get_paged_mqa_logits_metadata`. This is
# a SCHEDULE-granularity parameter, not the cache page size. The DG metadata
# kernel computes `SPLIT_KV = block_kv * 4` (the multiplier 4 is hardcoded;
# DeepGEMM commit 7f2a703 dropped the SM100-aware `arch == 10 ? 2 : 4` from
# nv_dev #fc97232 during the Public release 26/04 sync, leaving the formula
# uniform across SM90/SM100). Both the DG compute kernel (which hardcodes
# `split_kv = 256` at csrc/apis/attention.hpp:353) and our DSL paged-MQA-logits
# kernel (compute tile = 128 × kNumMathWarpGroups = 2 = 256) require
# `SPLIT_KV = 256`. So we must pass `block_kv = 256 / 4 = 64` here. Independent
# of the indexer K cache's physical page size (`tokens_per_block`), which only
# affects cache reads inside the compute kernel — the metadata wrapper does
# not read the cache. Passing the cache page size directly (the previous
# behavior) only works by accident when `tokens_per_block == 64`; for
# `tokens_per_block == 32` it produces a SPLIT_KV=128 schedule that the
# compute kernels misinterpret. TODO(remove once DeepGEMM restores the
# SM100-aware num_math_warpgroups in the metadata JIT impl).
_DG_SCHEDULE_BLOCK_KV = 64


def _compute_slot_mappings(
global_positions: torch.Tensor,
Expand Down Expand Up @@ -937,13 +955,15 @@ def prepare(self):
# Because the fp8_paged_mqa_logits only supports seq_len == 1/2/4 (i.e., max_draft_tokens == 0/1/3) on sm100, and
# seq_len == 1/2 (i.e., max_draft_tokens == 0/1) on sm90, for other cases, we need to flatten the q tensor and
# expand the kv_lens and block_table for MTP support.
# The CuTe DSL kernel supports arbitrary next_n natively, so it never needs expansion.
# TODO:
# - No distinction between sm90 and sm100 is needed once MTP3 is supported on sm90.
# - Remove this once fp8_paged_mqa_logits supports an arbitrary number of MTP draft tokens.
self.use_expanded_buffers_for_mtp = (
(self.max_draft_tokens > 1 and get_sm_version() == 90)
or ((self.max_draft_tokens == 2 or self.max_draft_tokens > 3)
and get_sm_version() >= 100))
_use_dsl = self.sparse_attention_config.use_cute_dsl_paged_mqa_logits
self.use_expanded_buffers_for_mtp = (not _use_dsl and (
(self.max_draft_tokens > 1 and get_sm_version() == 90) or
((self.max_draft_tokens == 2 or self.max_draft_tokens > 3)
and get_sm_version() >= 100)))
if self.use_expanded_buffers_for_mtp:
# Expand kv_lens_cuda (only generation)
num_tokens = self.num_generations * (1 + self.max_draft_tokens)
Expand Down Expand Up @@ -1055,9 +1075,10 @@ def on_update_kv_lens(self):
# column would be non-contiguous and would fail the metadata
# kernel's is_contiguous assertion.
context_lens_next_n1 = gen_kv_lens.view(-1, 1)
# `_DG_SCHEDULE_BLOCK_KV` (= 64) instead of cache `tokens_per_block`:
# see module-level constant comment for the SPLIT_KV=256 alignment.
scheduler_metadata_buffer = get_paged_mqa_logits_metadata(
context_lens_next_n1, self.kv_cache_manager.tokens_per_block,
self.num_sms)
context_lens_next_n1, _DG_SCHEDULE_BLOCK_KV, self.num_sms)
self.scheduler_metadata_buffer.copy_(scheduler_metadata_buffer,
non_blocking=True)
# When MTP is on without the expanded-tokens path, also populate
Expand All @@ -1070,8 +1091,8 @@ def on_update_kv_lens(self):
num_generations, :
next_n_cap]
scheduler_metadata_buffer_full_next_n = get_paged_mqa_logits_metadata(
context_lens_full_next_n,
self.kv_cache_manager.tokens_per_block, self.num_sms)
context_lens_full_next_n, _DG_SCHEDULE_BLOCK_KV,
self.num_sms)
self.scheduler_metadata_buffer_full_next_n.copy_(
scheduler_metadata_buffer_full_next_n, non_blocking=True)
if self.use_expanded_buffers_for_mtp:
Expand All @@ -1086,8 +1107,7 @@ def on_update_kv_lens(self):
num_tokens].view(
-1, 1)
scheduler_metadata_buffer_expanded = get_paged_mqa_logits_metadata(
kv_lens_expanded_2d, self.kv_cache_manager.tokens_per_block,
self.num_sms)
kv_lens_expanded_2d, _DG_SCHEDULE_BLOCK_KV, self.num_sms)
self.scheduler_metadata_buffer_expanded.copy_(
scheduler_metadata_buffer_expanded, non_blocking=True)
self.prepare_dense_topk_indices(self.kv_lens_cuda, device=True)
Expand Down Expand Up @@ -1206,19 +1226,24 @@ def __init__(self,
self.ln_events = [torch.cuda.Event(), torch.cuda.Event()]
self.use_cute_dsl_topk = (sparse_attention_config.use_cute_dsl_topk
and IS_CUTLASS_DSL_AVAILABLE)
self.use_cute_dsl_paged_mqa_logits = (
Comment thread
limin2021 marked this conversation as resolved.
sparse_attention_config.use_cute_dsl_paged_mqa_logits
and IS_CUTLASS_DSL_AVAILABLE)
self.weight_scale_factor = self.softmax_scale * self.n_heads**-0.5

self._enable_heuristic_topk = (
sparse_attention_config.enable_heuristic_topk
and get_sm_version() >= 100)

if self.use_cute_dsl_topk and layer_idx == 0:
if (self.use_cute_dsl_topk
or self.use_cute_dsl_paged_mqa_logits) and layer_idx == 0:
Comment thread
yuxianq marked this conversation as resolved.
from tensorrt_llm._torch.custom_ops import cute_dsl_custom_ops

# the dtype of topk input tensor, which is float32 now.
# Note, need to update it if the dtype of topk input tensor is changed.
cute_dsl_custom_ops.warmup_cute_dsl_indexer_topk(
dtype=torch.float32, top_k=self.index_topk)
if self.use_cute_dsl_topk:
# the dtype of topk input tensor, which is float32 now.
# Note, need to update it if the dtype of topk input tensor is changed.
cute_dsl_custom_ops.warmup_cute_dsl_indexer_topk(
dtype=torch.float32, top_k=self.index_topk)

if self._enable_heuristic_topk and layer_idx == 0:
# Populate static caches (sm_count, L2 cache size) inside the C++
Expand Down Expand Up @@ -1485,8 +1510,11 @@ def prepare(metadata: DSAtrtllmAttentionMetadata):
# slicing kv_lens_cuda_2d's first column would be a strided
# view that fails the metadata kernel's contiguous assertion.
context_lens_next_n1 = gen_seq_lens.view(-1, 1)
# `_DG_SCHEDULE_BLOCK_KV` (= 64) instead of cache `tokens_per_block`:
# see module-level constant comment for the SPLIT_KV=256 alignment.
scheduler_metadata_buffer = get_paged_mqa_logits_metadata(
context_lens_next_n1, tokens_per_block, metadata.num_sms)
context_lens_next_n1, _DG_SCHEDULE_BLOCK_KV,
metadata.num_sms)
metadata.scheduler_metadata_buffer.copy_(
scheduler_metadata_buffer, non_blocking=True)
# MTP main forward uses next_n = 1 + max_draft_tokens; build
Expand All @@ -1497,7 +1525,7 @@ def prepare(metadata: DSAtrtllmAttentionMetadata):
num_generations, :
next_n_cap]
scheduler_metadata_buffer_full_next_n = get_paged_mqa_logits_metadata(
context_lens_full_next_n, tokens_per_block,
context_lens_full_next_n, _DG_SCHEDULE_BLOCK_KV,
metadata.num_sms)
metadata.scheduler_metadata_buffer_full_next_n.copy_(
scheduler_metadata_buffer_full_next_n,
Expand All @@ -1512,7 +1540,8 @@ def prepare(metadata: DSAtrtllmAttentionMetadata):
num_tokens].view(
-1, 1)
scheduler_metadata_buffer_expanded = get_paged_mqa_logits_metadata(
kv_lens_expanded_2d, tokens_per_block, metadata.num_sms)
kv_lens_expanded_2d, _DG_SCHEDULE_BLOCK_KV,
metadata.num_sms)
metadata.scheduler_metadata_buffer_expanded.copy_(
scheduler_metadata_buffer_expanded, non_blocking=True)

Expand Down Expand Up @@ -1841,7 +1870,9 @@ def sparse_attn_indexer(
# schedule (via num_next_n_atoms). MTP forwards alternate
# between the full-window call (next_n == 1+max_draft_tokens)
# and per-token draft calls (next_n == 1), so we must select
# the buffer that was populated for this next_n.
# the buffer that was populated for this next_n. The DSL path
# uses its own schedule buffer (built with num_next_n_atoms=1
# via a (num_gen, 1) input shape) and overrides this below.
if next_n == 1:
scheduler_metadata_buffer = metadata.scheduler_metadata_buffer
else:
Expand All @@ -1864,19 +1895,37 @@ def sparse_attn_indexer(
k_cache = metadata.kv_cache_manager.get_indexer_k_cache_buffers(
self.layer_idx)

decode_q_scale = q_scale[num_ctx_tokens:num_ctx_tokens +
num_gen_tokens,
...] if self.use_fp4 else None
if self.use_fp4:
# q_decode shape is either (num_generations, next_n, n_heads,
# head_dim/2) [non-expanded] or (batch*next_n, 1, n_heads,
# head_dim/2) [expanded]. Match q_scale's batch/next_n dims.
decode_q_scale = decode_q_scale.view(q_decode.shape[0],
q_decode.shape[1],
self.n_heads)
logits_decode = self._call_paged_mqa_logits(
q_decode, k_cache, weights_decode, context_lens, block_table,
scheduler_metadata_buffer, max_seq_len, decode_q_scale)
if self.use_cute_dsl_paged_mqa_logits:
# DSL kernel design: 1 atom per q (atom = real next_n positions),
# kNumNextNAtoms = 1 for any real next_n. The matching schedule
# is `scheduler_metadata_buffer` — built in `Indexer.prepare()`
# with a (num_gen, 1) input shape, which makes DeepGEMM's wrapper
# compute `num_next_n_atoms = 1`. (DeepGEMM uses the same buffer
# for its own next_n=1 kernel; DSL piggy-backs on it for all
# real next_n values.) All next_n positions of a batch share
# the same KV length on this path (kv_lens_cuda_2d broadcasts),
# so passing the 1D contiguous kv_lens slice for context_lens
# avoids materializing a 2D contiguous tensor per call.
dsl_context_lens = metadata.kv_lens_cuda_runtime[
num_contexts:num_contexts + num_generations]
logits_decode = torch.ops.trtllm.cute_dsl_fp8_paged_mqa_logits(
q_decode, k_cache, weights_decode, dsl_context_lens,
block_table, metadata.scheduler_metadata_buffer,
max_seq_len)
else:
decode_q_scale = q_scale[num_ctx_tokens:num_ctx_tokens +
num_gen_tokens,
...] if self.use_fp4 else None
if self.use_fp4:
# q_decode shape is either (num_generations, next_n, n_heads,
# head_dim/2) [non-expanded] or (batch*next_n, 1, n_heads,
# head_dim/2) [expanded]. Match q_scale's batch/next_n dims.
decode_q_scale = decode_q_scale.view(
q_decode.shape[0], q_decode.shape[1], self.n_heads)
logits_decode = self._call_paged_mqa_logits(
q_decode, k_cache, weights_decode, context_lens,
block_table, scheduler_metadata_buffer, max_seq_len,
decode_q_scale)

if use_custom_topk:
# Kernel expects kv_lens (total cache length), not seq_lens (new tokens)
Expand Down
Loading
Loading