From 84a3c3e6d5f19679967263bbb51808ee36f4a4d9 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Wed, 5 Aug 2026 11:42:05 +0000 Subject: [PATCH 01/12] [None][test] Cover sparse MQA/GQA support matrix Document sparse attention algorithms and the internal sparse MQA/GQA kernel support boundaries. Expand unit tests across supported dtypes, head dimensions, and query-to-KV head group limits. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- docs/source/features/sparse-attention.md | 384 ++++++++++++------ .../attention/sparse/test_sparse_attention.py | 135 +++++- 2 files changed, 375 insertions(+), 144 deletions(-) diff --git a/docs/source/features/sparse-attention.md b/docs/source/features/sparse-attention.md index cb992ccd319b..7c42d94f84ff 100644 --- a/docs/source/features/sparse-attention.md +++ b/docs/source/features/sparse-attention.md @@ -1,63 +1,181 @@ # Sparse Attention - [Overview](#overview) - - [Algorithms](#algorithms) -- [RocketKV](#rocketkv) -- [DeepSeek Sparse Attention (DSA)](#deepseek-sparse-attention-dsa) -- [Skip Softmax Attention](#skip-softmax-attention) -- [Algorithm Comparison](#algorithm-comparison) +- [Supported Algorithms](#supported-algorithms) +- [Sparse MQA/GQA Kernel Support](#sparse-mqagqa-kernel-support) +- [Configure Sparse Attention](#configure-sparse-attention) +- [Algorithm Details](#algorithm-details) - [Further Reading](#further-reading) ## Overview -Sparse attention reduces the cost of long-context inference by skipping work on KV entries that contribute little to the attention output. In TensorRT LLM, sparse attention is enabled by passing a `sparse_attention_config` object to the `LLM` API, or its YAML equivalent for `trtllm-serve`, `trtllm-bench`, or `trtllm-eval`. The config object is a discriminated union: each algorithm has its own subclass of `BaseSparseAttentionConfig` selected via the `algorithm` field. - -This page focuses on the **user-facing API**: how to construct and pass the config for each supported algorithm, in both Python and YAML form. For framework design details, see [Blog 17: Sparse Attention in TensorRT-LLM](../blogs/tech_blog/blog17_Sparse_Attention_in_TensorRT-LLM.md). For developers adding a new sparse attention algorithm, see the [Sparse Attention Development Guide](../developer-guide/sparse-attention-development-guide.md). +Sparse attention reduces long-context inference cost by avoiding attention work on +KV entries that an algorithm considers unimportant. TensorRT LLM separates two +parts of that process: + +1. An algorithm selects tokens or blocks, or decides which kernel tiles can be + skipped. +2. An attention implementation consumes that sparse pattern and computes the + output. + +This distinction matters for support. A kernel that can compute sparse MQA/GQA +does not by itself define how a model selects tokens, and therefore is not a +standalone user-facing algorithm. + +The user-facing `sparse_attention_config` API is currently prototype and is +supported by the PyTorch execution backend. Each public algorithm has a config +class selected by its `algorithm` field. Model-native algorithms usually read +their geometry from the checkpoint; avoid overriding those values unless the +model-specific guide says they are tunable. + +## Supported Algorithms + +| `algorithm` | Config class | Sparse mechanism | Attention implementation | Typical use | +|---|---|---|---|---| +| `rocket` | `RocketSparseAttentionConfig` | Prompt KV eviction, then page-level Top-K selection during decode | TRTLLM or Vanilla | Training-free sparsity for standard attention models | +| `dsa` | `DeepSeekSparseAttentionConfig` | Learned token-level indexer followed by sparse MLA | TRTLLM | DeepSeek V3.2 and compatible model-native DSA architectures | +| `deepseek_v4` | `DeepSeekV4SparseAttentionConfig` | Sliding-window attention plus compressed sparse or compressed dense history | TRTLLM | DeepSeek-V4 hybrid attention | +| `minimax_m3` | `MiniMaxM3SparseAttentionConfig` | Learned block selection followed by sparse GQA | Dedicated Triton or MSA implementation | MiniMax-M3 sparse layers | +| `skip_softmax` | `SkipSoftmaxAttentionConfig` | Dynamically skips eligible softmax work inside the FMHA kernel | TRTLLM | Existing full-attention models with calibrated or direct thresholds | + +All five configs select the PyTorch execution backend. The "attention +implementation" column refers to the attention kernel/backend used inside that +execution backend. + +### Capability Comparison + +| Capability | RocketKV | DSA | DeepSeek-V4 | MiniMax-M3 | Skip Softmax | +|---|---:|---:|---:|---:|---:| +| Sparse prefill computation | No | Yes | Yes | Yes | Yes | +| Sparse decode computation | Yes | Yes | Yes | Yes | Yes | +| Reduces retained main KV history | Yes | No | Yes, through model-native compression | No | No | +| Requires a model-trained selector | No | Yes | Yes | Yes | No | +| Selection granularity | Token eviction and pages | Tokens | Compressed entries | Blocks | Kernel tiles | + +"No" for RocketKV prefill means that prompt attention is still computed +densely. RocketKV selects which prompt KV entries to retain, so it reduces cache +size and later decode work. + +## Sparse MQA/GQA Kernel Support + +TensorRT LLM contains an internal TRTLLM-Gen kernel for token-sparse +multi-query attention (MQA) and grouped-query attention (GQA). It accepts a +precomputed token-index list for each KV head and query token. Query heads in +the same KV group share the KV head's index list. + +This is a kernel capability, not a public `SparseAttentionConfig` algorithm. +There is no supported `algorithm: mqa_gqa` value for `LLM` or YAML. A sparse +algorithm must provide the selector, metadata, cache management, and attention +backend integration before applications can use this kernel through the public +API. + +### Support Matrix + +| Parameter | Supported | Not currently supported or not established | +|---|---|---| +| GPU architecture | SM100 and SM103 | Pre-Blackwell GPUs; SM120 and SM121 | +| Attention type | MQA (`num_kv_heads == 1`) and GQA | Arbitrary head mappings | +| Head relationship | `num_q_heads % num_kv_heads == 0`; at most 32 query heads per KV head | Non-divisible Q/KV head counts; MQA/GQA groups larger than 32 | +| Q/K/V dimensions | Equal QK and V head dimensions | Unequal QK/V dimensions (MLA uses a separate sparse path) | +| Head dimension | `64`, `80`, `128`, or `256` | `512` and other head dimensions | +| Q/K/V and output dtype | BF16 or FP16 | Quantized Q/output combinations are not covered by this primitive's regression tests | +| KV cache | Paged KV cache, with the KV dtype matching the input dtype; page size is a power of two and at least 8 tokens | Contiguous KV cache; non-power-of-two pages; pages smaller than 8 tokens | +| Sparse indices | `int32`, token-granular, one list per KV head and query token | A public built-in selector for generic MQA/GQA | +| Sparse Top-K | Positive multiple of 4; shorter sequences may pad unused entries with `-1` | Top-K values not divisible by 4 | +| Inference phase | Fresh context and single-token generation | Mixed context/generation batches are not covered by the regression tests | +| Beam width | `1` | Beam search | +| Attention mask/window | Causal self-attention with a fixed cache window | ALiBi, arbitrary custom masks, StreamingLLM/sink tokens, and variable cyclic windows | + +The current main branch JIT-compiles this path with NVRTC. Its support is +therefore defined by the current TRTLLM-Gen source checks, not by the set of +precompiled cubins that was present when the feature was introduced. + +The regression tests cover: + +- MQA and GQA ratios of 2:1, 4:1, and 8:1; +- the maximum supported query-head group size of 32 for both MQA and GQA; +- variable batch and sequence lengths; +- context KV compaction, context sparse computation, and decode sparse + computation; +- Top-K values `4`, `64`, and `128`, including Top-K larger than a request's + current KV length; +- backing KV-cache page sizes `32` and `64`; +- all supported equal head dimensions in both BF16 and FP16. + +The shared TRTLLM-Gen option validator also admits head dimension `512`, but +the sparse MQA/GQA path aborts before launch for that configuration on current +main. It is therefore intentionally excluded from the supported matrix and +regression tests. + +Backend developers can use +[`test_sparse_attention.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_attention.py) +as a minimal integration example. `MockSparseParams` and +`TestSparseAttention` deliberately supply fixed sparse predictions so that the +test isolates the cache/index layout and kernel computation. They are not +public application APIs. + +## Configure Sparse Attention + +Pass a config object to `LLM` in Python, or use the equivalent discriminated +YAML object with `trtllm-serve`, `trtllm-bench`, or `trtllm-eval`. -### Algorithms +```python +from tensorrt_llm import LLM +from tensorrt_llm.llmapi import RocketSparseAttentionConfig -| `algorithm` | Config class | Reference | -|---|---|---| -| `rocket` | `RocketSparseAttentionConfig` | [RocketKV paper](https://arxiv.org/pdf/2502.14051) | -| `dsa` | `DeepSeekSparseAttentionConfig` | [DeepSeek V3.2 paper](https://github.com/deepseek-ai/DeepSeek-V3.2-Exp/blob/main/DeepSeek_V3_2.pdf) | -| `skip_softmax` | `SkipSoftmaxAttentionConfig` | [BLASST paper](https://arxiv.org/pdf/2512.12087) | +llm = LLM( + model="", + sparse_attention_config=RocketSparseAttentionConfig(), +) +``` -For per-field semantics, refer to the docstring on each config class in `tensorrt_llm/llmapi/llm_args.py`. +```yaml +sparse_attention_config: + algorithm: rocket +``` -YAML configs shown below are consumed via the standard `--extra_llm_api_options` / `--config` flag: +For example: ```bash -trtllm-serve --model --config extra_config.yaml ... -trtllm-bench --model --config extra_config.yaml ... -trtllm-eval --model --config extra_config.yaml longbench_v2 --max_output_length 1024 +trtllm-serve --config config.yaml +trtllm-bench --model throughput --dataset --config config.yaml ``` -## RocketKV +The following sections list algorithm-specific settings and constraints. + +## Algorithm Details + +### RocketKV -RocketKV is a training-free, two-stage algorithm. It applies permanent KV cache eviction in the context phase, followed by dynamic Top-K token selection in the generation phase. Some framework-level algorithms, including RocketKV, currently require disabling KV cache block reuse. +[RocketKV](https://arxiv.org/pdf/2502.14051) is a training-free, two-stage +algorithm for standard attention architectures. During prefill, it computes +dense attention and permanently evicts prompt KV entries beyond a prompt +budget. During decode, it scores retained pages and attends to the selected +Top-K pages. -**Python API** +RocketKV currently requires CUDA compute capability 10.0 or newer. KV-cache +block reuse and chunked prefill must be disabled, and disaggregated serving is +not supported. ```python from tensorrt_llm import LLM, SamplingParams -from tensorrt_llm.llmapi import RocketSparseAttentionConfig, KvCacheConfig - -sparse_attention_config = RocketSparseAttentionConfig( - prompt_budget=2048, - kt_cache_dtype="float8_e5m2", -) -kv_cache_config = KvCacheConfig(enable_block_reuse=False) +from tensorrt_llm.llmapi import KvCacheConfig, RocketSparseAttentionConfig llm = LLM( model="", - sparse_attention_config=sparse_attention_config, - kv_cache_config=kv_cache_config, + sparse_attention_config=RocketSparseAttentionConfig( + prompt_budget=2048, + kt_cache_dtype="float8_e5m2", + ), + kv_cache_config=KvCacheConfig(enable_block_reuse=False), + enable_chunked_prefill=False, +) +outputs = llm.generate( + ["To be or not to be..."], + SamplingParams(max_tokens=128), ) -outputs = llm.generate(["To be or not to be..."], SamplingParams(max_tokens=128)) ``` -**YAML** - ```yaml sparse_attention_config: algorithm: rocket @@ -68,40 +186,37 @@ kv_cache_config: enable_chunked_prefill: false ``` -## DeepSeek Sparse Attention (DSA) +The TRTLLM and Vanilla attention implementations support RocketKV. The +Vanilla implementation requires a BF16 KT cache. -DSA is a model-native sparse attention mechanism introduced with DeepSeek V3.2. A lightweight learned indexer scores all KV entries, and only the top-`index_topk` entries are attended to. +### DeepSeek Sparse Attention -**Python API** +DeepSeek Sparse Attention (DSA) is a model-native mechanism introduced by +DeepSeek V3.2. A learned MQA indexer scores the KV history, Top-K selects token +indices, and sparse MLA consumes them. Checkpoint fields define the indexer +head count, index head dimension, and Top-K; the safest configuration is to let +TensorRT LLM load them from the model. ```python from tensorrt_llm import LLM from tensorrt_llm.llmapi import DeepSeekSparseAttentionConfig -sparse_attention_config = DeepSeekSparseAttentionConfig(index_topk=64) - llm = LLM( - model="", - sparse_attention_config=sparse_attention_config, + model="deepseek-ai/DeepSeek-V3.2", + sparse_attention_config=DeepSeekSparseAttentionConfig(), ) ``` -**YAML** - ```yaml sparse_attention_config: algorithm: dsa - index_topk: 64 ``` -**Optional: Guess-Verify-Refine Top-K.** On Blackwell (SM 100+), set `enable_heuristic_topk=True` to use the Guess-Verify-Refine (GVR) Top-K. GVR is currently supported only for `index_topk=2048`; other values fall back to the production insertion/radix Top-K. `TRTLLM_HEURISTIC_NMIN` overrides the small-batch lower bound, and `TRTLLM_SCHEMEX_DEBUG=1` prints the dispatcher decision. - -```python -sparse_attention_config = DeepSeekSparseAttentionConfig( - index_topk=2048, - enable_heuristic_topk=True, -) -``` +On supported Blackwell configurations, Guess-Verify-Refine (GVR) can replace +the regular decode Top-K dispatcher. The current implementation accepts +`index_topk` values `512`, `1024`, and `2048`, and indexer compression ratios +`1` and `4`. Unsupported combinations fall back to the production +insertion/radix Top-K path. ```yaml sparse_attention_config: @@ -110,91 +225,89 @@ sparse_attention_config: enable_heuristic_topk: true ``` -## Skip Softmax Attention +See the +[DeepSeek V3/V3.2 example](../../../examples/models/core/deepseek_v3/README.md) +for model precision, hardware, parallelism, MTP, chunked-prefill, cache-reuse, +and disaggregated-serving support. -Skip Softmax Attention is a kernel-level method, also known as BLASST, that dynamically skips computation in a FlashAttention-style kernel. It can accelerate existing full-attention models without changing the model architecture. +### DeepSeek-V4 Hybrid Sparse Attention -The value actually consumed by the kernel is **`threshold_scale_factor`**. The kernel combines it with the **sequence length** to compute the **threshold** at runtime. Other configuration paths resolve to that scalar before the attention backend is constructed. +DeepSeek-V4 interleaves three model-native attention modes: -### Checkpoint Config +- sliding-window attention over recent raw tokens; +- compressed sparse attention over 4x-compressed history selected by an + indexer; +- compressed dense attention over 128x-compressed history. -[NVIDIA Model Optimizer](https://github.com/NVIDIA/Model-Optimizer) (ModelOpt) can perform calibration and store metadata for Skip Softmax Attention in the model checkpoint's `config.json`. The checkpoint config provides the formula that maps `target_sparsity` to `threshold_scale_factor`. +TensorRT LLM normally constructs `DeepSeekV4SparseAttentionConfig` from the +checkpoint. An explicit config overrides matching fields; it must preserve the +model's attention layout. The current implementation requires +`window_size=128`, compression ratios from `{1, 4, 128}`, data-center Blackwell +GPUs, KV-cache blocks of `128` or `256` tokens, and beam width `1`. -This checkpoint config is **optional**. It is only required when using `target_sparsity`, which is a [0, 1] scalar that is more intuitive than directly choosing the kernel-facing `threshold_scale_factor`. But please note that `target_sparsity` only serves as a guidance, the actual **achieved** sparsity in the kernel would vary. +```yaml +sparse_attention_config: + algorithm: deepseek_v4 + window_size: 128 + index_topk: 512 +``` -Example checkpoint config: +See the +[DeepSeek-V4 example](../../../examples/models/core/deepseek_v4/README.md) for +checkpoint-derived configuration and deployment constraints. -```json -{ - "sparse_attention_config": { - "config_groups": { - "group_0": { - "algorithm": "skip_softmax", - "threshold_scale_factor": { - "formula": "a * exp(b * target_sparsity)", - "prefill": {"a": 100.0, "b": 5.0}, - "decode": {"a": 0.05, "b": 10.0} - }, - "target_sparsity": { - "prefill": 0.5, - "decode": 0.3 - }, - "ignore": [ - "model.layers.0.self_attn", - "model.layers.1.self_attn" - ] - } - } - } -} +### MiniMax-M3 Block-Sparse GQA + +MiniMax-M3 uses model-native block-sparse GQA in its sparse layers. An index +branch scores main KV-cache blocks, forces configured initial/local blocks into +the selection, and chooses the remaining Top-K blocks before sparse GQA. +Defaults such as four index heads, index dimension `128`, block size `128`, and +16 selected blocks come from the checkpoint-compatible config. + +```yaml +sparse_attention_config: + algorithm: minimax_m3 ``` -The checkpoint config may contain multiple `config_groups` for different sparse attention algorithms. At most one group may configure Skip Softmax Attention. Multiple groups whose `algorithm` is `skip_softmax` are invalid. +Two implementations are available: -- `formula` — an **arbitrary** [numexpr](https://numexpr.readthedocs.io/) expression of `threshold_scale_factor` using `target_sparsity` and one or more named coefficients. Standard math functions such as `exp`, `log`, `sqrt`, `pow`, and `**` are available. The runtime parses and evaluates it directly, so calibration is not locked to a fixed functional form. It can be configured separately for prefill and decode; otherwise both phases use the same config. -- `target_sparsity` — optional checkpoint-provided target values. It can be configured separately for prefill and decode; otherwise both phases use the same config. -- `ignore` — optional fnmatch layer patterns where the calibrated Skip Softmax Attention config should not apply. +- `triton` is the default reference implementation. +- `msa` uses `fmha_sm100` kernels and requires an SM100-family GPU (SM100 or + SM103), the `fmha_sm100` package, and `sparse_block_size=128`. -### User Configuration +```yaml +sparse_attention_config: + algorithm: minimax_m3 + implementation: msa +``` -User configuration is supplied through Python or YAML and controls how the checkpoint metadata is consumed: +The sparse path currently has no dense fallback and does not support KV-cache +reuse or MTP. See the +[MiniMax-M3 deployment guide](../deployment-guide/deployment-guide-for-minimax-m3-on-trtllm.md) +for supported checkpoints and parallel deployment settings. -- Set `threshold_scale_factor` directly to pass a concrete threshold to the kernel. This does not require checkpoint config. -- Set `target_sparsity` to request a sparsity target. The runtime resolves it to `threshold_scale_factor` using the checkpoint calibration formula. If the checkpoint does not provide the required Skip Softmax Attention metadata, the runtime raises an error. +### Skip Softmax Attention -Both `threshold_scale_factor` and `target_sparsity` take either a scalar, applied to both prefill and decode, or a `{"prefill": ..., "decode": ...}` dict. `threshold_scale_factor` and `target_sparsity` are alternatives: if both are present, `threshold_scale_factor` takes precedence and the calibration formula is not used. User-provided `target_sparsity` overrides checkpoint-default `target_sparsity`. Checkpoint `ignore` patterns always disable Skip Softmax Attention for matching layers. +Skip Softmax Attention, also known as BLASST, dynamically skips eligible work +inside a FlashAttention-style kernel. It does not select tokens, alter the +model architecture, or reduce KV-cache storage. -#### Python API +The kernel consumes `threshold_scale_factor` and combines it with sequence +length at runtime. You can provide that value directly: ```python from tensorrt_llm import LLM from tensorrt_llm.llmapi import SkipSoftmaxAttentionConfig -# Direct threshold (single value applied to both phases): -sparse_attention_config = SkipSoftmaxAttentionConfig(threshold_scale_factor=1000.0) - -# Direct threshold, per-phase: -sparse_attention_config = SkipSoftmaxAttentionConfig( - threshold_scale_factor={"prefill": 1000.0, "decode": 500.0}, -) - -# Target sparsity (requires the checkpoint to carry a calibration formula): -sparse_attention_config = SkipSoftmaxAttentionConfig(target_sparsity=0.5) - -# Target sparsity, per-phase: -sparse_attention_config = SkipSoftmaxAttentionConfig( - target_sparsity={"prefill": 0.5, "decode": 0.3}, +llm = LLM( + model="", + sparse_attention_config=SkipSoftmaxAttentionConfig( + threshold_scale_factor={"prefill": 1000.0, "decode": 500.0}, + ), ) - -llm = LLM(model="", sparse_attention_config=sparse_attention_config) ``` -Skip Softmax Attention only works with the **TRTLLM** attention backend, which is the default attention backend. Other backends silently bypass Skip Softmax Attention. - -#### YAML - ```yaml -# Direct threshold: sparse_attention_config: algorithm: skip_softmax threshold_scale_factor: @@ -202,8 +315,11 @@ sparse_attention_config: decode: 500.0 ``` +Alternatively, provide `target_sparsity`. This path requires the checkpoint to +contain a calibration formula that maps the requested target to the kernel's +threshold scale factor. + ```yaml -# Target sparsity (requires a calibrated checkpoint): sparse_attention_config: algorithm: skip_softmax target_sparsity: @@ -211,15 +327,39 @@ sparse_attention_config: decode: 0.3 ``` -## Algorithm Comparison +Both fields accept a scalar for both phases or a dictionary with `prefill` and +`decode` values. If both are present, `threshold_scale_factor` takes +precedence. User-provided `target_sparsity` overrides a checkpoint default. + +Model Optimizer can store calibration metadata in the checkpoint's +`config.json`: + +```json +{ + "sparse_attention_config": { + "config_groups": { + "group_0": { + "algorithm": "skip_softmax", + "threshold_scale_factor": { + "formula": "a * exp(b * target_sparsity)", + "prefill": {"a": 100.0, "b": 5.0}, + "decode": {"a": 0.05, "b": 10.0} + }, + "target_sparsity": {"prefill": 0.5, "decode": 0.3}, + "ignore": ["model.layers.0.self_attn"] + } + } + } +} +``` + +The formula is a [numexpr](https://numexpr.readthedocs.io/) expression over +`target_sparsity` and named coefficients. The optional `ignore` list uses +fnmatch layer patterns. At most one checkpoint config group may use the +`skip_softmax` algorithm. -| Aspect | RocketKV | DSA | Skip Softmax Attention | -|---|---|---|---| -| Prefill acceleration | No | Yes | Yes | -| Decode acceleration | Yes | Yes | Yes | -| KV cache reduction | Yes | No | No | -| Framework-level support required | Yes | Yes | No | -| Model-native | No | Yes | No | +Skip Softmax Attention requires the TRTLLM attention backend. Other attention +backends do not apply it. ## Further Reading diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_attention.py b/tests/unittest/_torch/attention/sparse/test_sparse_attention.py index e24ac9f4e060..6c4d2c025b38 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_attention.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_attention.py @@ -13,8 +13,11 @@ # See the License for the specific language governing permissions and # limitations under the License. -""" -Unit tests for sparse attention with TrtllmAttention backend. +"""Unit tests for sparse attention with the ``TrtllmAttention`` backend. + +The ``mqa_gqa`` sparse parameters below exercise the internal token-sparse +TRTLLM-Gen kernel directly. They are intentionally a test stub, not a public +``SparseAttentionConfig`` algorithm. """ import math @@ -46,20 +49,26 @@ from tensorrt_llm._torch.attention.mla import MLA from tensorrt_llm._torch.metadata import KVCacheParams from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager -from tensorrt_llm._utils import str_dtype_to_binding, torch_dtype_to_str +from tensorrt_llm._utils import is_sm_100f, str_dtype_to_binding, torch_dtype_to_str from tensorrt_llm.bindings.executor import KvCacheConfig from tensorrt_llm.mapping import Mapping ATOL = 2e-2 RTOL = 2e-2 +requires_sparse_mqa_gqa = pytest.mark.skipif( + not is_sm_100f(getSMVersion()), + reason="Sparse MQA/GQA requires an SM100-family GPU (SM100 or SM103)", +) + @dataclass(kw_only=True, frozen=False) class SparseScenario: """Base configuration for sparse attention tests. - NOTE: SparseMqaGqa trtllm-gen cubins are currently only available for BF16. - FP16 cubins have not been generated yet, so tests default to BF16. + Sparse MQA/GQA uses the TRTLLM-Gen NVRTC path on current main. BF16 is the + default for the broad scenario coverage; a focused matrix below also + covers FP16 and every supported equal QK/V head dimension. """ dtype: torch.dtype = torch.bfloat16 @@ -958,7 +967,7 @@ def _build_reference_kv_cache( return k_cache_ref, v_cache_ref -@pytest.mark.skipif(getSMVersion() < 100, reason="Sparse MQA/GQA requires SM100 (Blackwell)") +@requires_sparse_mqa_gqa @pytest.mark.parametrize( "s", [ @@ -1022,7 +1031,7 @@ def test_context_sparse_kv(s: SparseContextScenario): kv_cache_manager.shutdown() -@pytest.mark.skipif(getSMVersion() < 100, reason="Sparse MQA/GQA requires SM100 (Blackwell)") +@requires_sparse_mqa_gqa @pytest.mark.parametrize( "s", [ @@ -1058,6 +1067,28 @@ def test_context_sparse_kv(s: SparseContextScenario): past_kv_lens=(128, 256), num_pages=16, ), + SparseGenerationScenario( + num_heads=8, + num_kv_heads=4, + batch_size=2, + past_kv_lens=(128, 256), + num_pages=16, + ), + # Maximum supported query-head group size + SparseGenerationScenario( + num_heads=32, + num_kv_heads=1, + batch_size=1, + past_kv_lens=(128,), + num_pages=8, + ), + SparseGenerationScenario( + num_heads=64, + num_kv_heads=2, + batch_size=1, + past_kv_lens=(128,), + num_pages=8, + ), # topk: minimum (4), topk exceeding some past_kv_lens SparseGenerationScenario( batch_size=1, @@ -1092,6 +1123,9 @@ def test_context_sparse_kv(s: SparseContextScenario): "batch3_var_kv", "mqa_8q1kv", "gqa_16q4kv", + "gqa_8q4kv", + "mqa_group32_boundary", + "gqa_group32_boundary", "topk4_min", "topk128_exceeds_some", "batch8_varied", @@ -1100,6 +1134,11 @@ def test_context_sparse_kv(s: SparseContextScenario): ) def test_generation_sparse_attention(s: SparseGenerationScenario): """Test generation phase with sparse attention computation.""" + _run_generation_sparse_attention(s) + + +def _run_generation_sparse_attention(s: SparseGenerationScenario): + """Run a generation scenario and compare the sparse kernel with PyTorch.""" ( device, q, @@ -1111,27 +1150,79 @@ def test_generation_sparse_attention(s: SparseGenerationScenario): metadata, attention, ) = _setup_generation_test(s) + try: + k_cache_ref, v_cache_ref = _build_reference_kv_cache( + kv_cache_manager, request_ids, s, device, s.dtype + ) + ref_sparse_output = reference_generation_sparse_attention( + q, k_cache_ref, v_cache_ref, k_new, v_new, sparse_attn_indices, s + ) - k_cache_ref, v_cache_ref = _build_reference_kv_cache( - kv_cache_manager, request_ids, s, device, s.dtype - ) - ref_sparse_output = reference_generation_sparse_attention( - q, k_cache_ref, v_cache_ref, k_new, v_new, sparse_attn_indices, s - ) + qkv = torch.cat([q, k_new, v_new], dim=1) + output = attention.forward(qkv, None, None, metadata) - qkv = torch.cat([q, k_new, v_new], dim=1) - output = attention.forward(qkv, None, None, metadata) + expected_shape = (s.num_generations, s.num_heads * s.head_dim) + assert output.shape == expected_shape, f"Shape mismatch: {output.shape} vs {expected_shape}" + assert torch.isfinite(output).all(), "Output contains non-finite values" - expected_shape = (s.num_generations, s.num_heads * s.head_dim) - assert output.shape == expected_shape, f"Shape mismatch: {output.shape} vs {expected_shape}" - assert torch.isfinite(output).all(), "Output contains non-finite values" + torch.testing.assert_close(output, ref_sparse_output, atol=ATOL, rtol=RTOL) + print(f"Generation sparse attention test passed: {s}") + finally: + kv_cache_manager.shutdown() - torch.testing.assert_close(output, ref_sparse_output, atol=ATOL, rtol=RTOL) - print(f"Generation sparse attention test passed: {s}") - kv_cache_manager.shutdown() + +@requires_sparse_mqa_gqa +@pytest.mark.parametrize("num_kv_heads", [1, 4], ids=["mqa", "gqa_2to1"]) +@pytest.mark.parametrize( + ("dtype", "head_dim"), + [ + (torch.bfloat16, 64), + (torch.bfloat16, 80), + (torch.bfloat16, 128), + (torch.bfloat16, 256), + (torch.float16, 64), + (torch.float16, 80), + (torch.float16, 128), + (torch.float16, 256), + ], + ids=[ + "bf16_h64", + "bf16_h80", + "bf16_h128", + "bf16_h256", + "fp16_h64", + "fp16_h80", + "fp16_h128", + "fp16_h256", + ], +) +def test_generation_sparse_mqa_gqa_kernel_matrix( + monkeypatch: pytest.MonkeyPatch, + num_kv_heads: int, + dtype: torch.dtype, + head_dim: int, +): + """Cover every supported head dimension and both supported input dtypes. + + Each KV head owns one token-index list. All query heads in its MQA/GQA + group share that list, which is the layout expected by the kernel. + """ + monkeypatch.setenv("TLLM_FMHA_LIBS", "fallback") + scenario = SparseGenerationScenario( + dtype=dtype, + kvcache_dtype=dtype, + num_heads=8, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + batch_size=1, + past_kv_lens=(64,), + num_pages=4, + num_sparse_topk=32, + ) + _run_generation_sparse_attention(scenario) -@pytest.mark.skipif(getSMVersion() < 100, reason="Sparse MQA/GQA requires SM100 (Blackwell)") +@requires_sparse_mqa_gqa @pytest.mark.parametrize( "s", [ From 73e923ae51b679eb334186af3ae84123e4d630f8 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Mon, 24 Aug 2026 09:20:17 -0700 Subject: [PATCH 02/12] [None][test] Refine sparse MQA/GQA coverage Separate kernel-specific MQA/GQA regression tests from generic sparse attention framework tests. Cover linear draft decoding, additional head-group sizes, FP8 KV cache and output, and document the verified support matrix. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- docs/source/features/sparse-attention.md | 37 +- .../attention/sparse/test_sparse_attention.py | 1240 +-------------- .../attention/sparse/test_sparse_mqa_gqa.py | 1330 +++++++++++++++++ 3 files changed, 1373 insertions(+), 1234 deletions(-) create mode 100644 tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py diff --git a/docs/source/features/sparse-attention.md b/docs/source/features/sparse-attention.md index 7c42d94f84ff..0730d926be7a 100644 --- a/docs/source/features/sparse-attention.md +++ b/docs/source/features/sparse-attention.md @@ -75,14 +75,15 @@ API. |---|---|---| | GPU architecture | SM100 and SM103 | Pre-Blackwell GPUs; SM120 and SM121 | | Attention type | MQA (`num_kv_heads == 1`) and GQA | Arbitrary head mappings | -| Head relationship | `num_q_heads % num_kv_heads == 0`; at most 32 query heads per KV head | Non-divisible Q/KV head counts; MQA/GQA groups larger than 32 | +| Head relationship | `num_q_heads % num_kv_heads == 0`; at most 32 query heads per KV head; tests cover group sizes `2`, `3`, `4`, `8`, `16`, `24`, `31`, and `32` | Non-divisible Q/KV head counts and groups larger than 32 | | Q/K/V dimensions | Equal QK and V head dimensions | Unequal QK/V dimensions (MLA uses a separate sparse path) | | Head dimension | `64`, `80`, `128`, or `256` | `512` and other head dimensions | -| Q/K/V and output dtype | BF16 or FP16 | Quantized Q/output combinations are not covered by this primitive's regression tests | -| KV cache | Paged KV cache, with the KV dtype matching the input dtype; page size is a power of two and at least 8 tokens | Contiguous KV cache; non-power-of-two pages; pages smaller than 8 tokens | +| Model Q/K/V input dtype | BF16 or FP16 | FP8 model-input tensors are not wired through this XQA fixed-parameter path | +| Output dtype | BF16, FP16, or E4M3 FP8 | Other output dtypes | +| KV-cache dtype and layout | BF16, FP16, or E4M3 FP8 paged KV cache; page size is a power of two and at least 8 tokens | Contiguous cache, non-power-of-two pages, and pages smaller than 8 tokens | | Sparse indices | `int32`, token-granular, one list per KV head and query token | A public built-in selector for generic MQA/GQA | | Sparse Top-K | Positive multiple of 4; shorter sequences may pad unused entries with `-1` | Top-K values not divisible by 4 | -| Inference phase | Fresh context and single-token generation | Mixed context/generation batches are not covered by the regression tests | +| Inference phase | Packed prefill, single-token generation, and linear draft-token generation (`qSeqLen=4` is tested) | Tree-shaped speculative masks are ignored by this static sparse kernel; mixed context/generation batches are not covered | | Beam width | `1` | Beam search | | Attention mask/window | Causal self-attention with a fixed cache window | ALiBi, arbitrary custom masks, StreamingLLM/sink tokens, and variable cyclic windows | @@ -90,17 +91,29 @@ The current main branch JIT-compiles this path with NVRTC. Its support is therefore defined by the current TRTLLM-Gen source checks, not by the set of precompiled cubins that was present when the feature was introduced. +Linear draft-token generation is verified with four query tokens per request: +one target token plus three draft tokens. Each query has its own causal sparse +index list, including indices for K/V written earlier in the same speculative +forward. A separate branched-tree probe showed that the static sparse kernel +matches the unmasked reference rather than the tree-filtered reference, so +tree-shaped speculative masks are not supported. + +For FP8 KV cache, Q is quantized to E4M3 during QKV preprocessing and the XQA +runner selects E4M3 KV/math types while retaining BF16 or FP16 model input. +Tests cover both BF16 output with an FP8 KV cache and an E4M3 FP8-output kernel. + The regression tests cover: -- MQA and GQA ratios of 2:1, 4:1, and 8:1; -- the maximum supported query-head group size of 32 for both MQA and GQA; +- MQA and GQA group sizes `2`, `3`, `4`, `8`, `16`, `24`, `31`, and `32`; - variable batch and sequence lengths; - context KV compaction, context sparse computation, and decode sparse computation; - Top-K values `4`, `64`, and `128`, including Top-K larger than a request's current KV length; - backing KV-cache page sizes `32` and `64`; -- all supported equal head dimensions in both BF16 and FP16. +- all supported equal head dimensions in both BF16 and FP16; +- linear generation with three draft tokens; +- E4M3 FP8 KV cache and FP8 output with BF16 model input. The shared TRTLLM-Gen option validator also admits head dimension `512`, but the sparse MQA/GQA path aborts before launch for that configuration on current @@ -108,11 +121,11 @@ main. It is therefore intentionally excluded from the supported matrix and regression tests. Backend developers can use -[`test_sparse_attention.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_attention.py) -as a minimal integration example. `MockSparseParams` and -`TestSparseAttention` deliberately supply fixed sparse predictions so that the -test isolates the cache/index layout and kernel computation. They are not -public application APIs. +[`test_sparse_mqa_gqa.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py) +as a minimal integration example. `_SparseMqaGqaParams` and +`_StaticSparseMqaGqaAttention` deliberately supply fixed sparse predictions so +that the test isolates the cache/index layout and kernel computation. They are +not public application APIs. ## Configure Sparse Attention diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_attention.py b/tests/unittest/_torch/attention/sparse/test_sparse_attention.py index 6c4d2c025b38..40a151daa252 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_attention.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_attention.py @@ -13,28 +13,19 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Unit tests for sparse attention with the ``TrtllmAttention`` backend. +"""Tests for algorithm-independent sparse attention framework plumbing. -The ``mqa_gqa`` sparse parameters below exercise the internal token-sparse -TRTLLM-Gen kernel directly. They are intentionally a test stub, not a public -``SparseAttentionConfig`` algorithm. +Kernel-specific regression coverage lives in dedicated modules such as +``test_sparse_mqa_gqa.py``. This file verifies how sparse algorithms register +hooks and pass predictions through ``SparseRuntimeParams``. """ -import math -from dataclasses import dataclass from types import ModuleType -from typing import List, Optional, Tuple from unittest.mock import Mock -import pytest import torch -from utils.util import getSMVersion -import tensorrt_llm from tensorrt_llm._torch.attention.backends.interface import AttentionForwardArgs -from tensorrt_llm._torch.attention.backends.sparse.dsa.kernels import ( - triton_convert_req_index_to_global_index, -) from tensorrt_llm._torch.attention.backends.sparse.hooks import ( AttentionSparseHooks, MLASparseHooks, @@ -45,140 +36,22 @@ register_mla_sparse_hooks, ) from tensorrt_llm._torch.attention.backends.sparse.params import SparseParams, SparseRuntimeParams -from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttention, TrtllmAttentionMetadata +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttention from tensorrt_llm._torch.attention.mla import MLA -from tensorrt_llm._torch.metadata import KVCacheParams -from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager -from tensorrt_llm._utils import is_sm_100f, str_dtype_to_binding, torch_dtype_to_str -from tensorrt_llm.bindings.executor import KvCacheConfig -from tensorrt_llm.mapping import Mapping - -ATOL = 2e-2 -RTOL = 2e-2 - -requires_sparse_mqa_gqa = pytest.mark.skipif( - not is_sm_100f(getSMVersion()), - reason="Sparse MQA/GQA requires an SM100-family GPU (SM100 or SM103)", -) - - -@dataclass(kw_only=True, frozen=False) -class SparseScenario: - """Base configuration for sparse attention tests. - - Sparse MQA/GQA uses the TRTLLM-Gen NVRTC path on current main. BF16 is the - default for the broad scenario coverage; a focused matrix below also - covers FP16 and every supported equal QK/V head dimension. - """ - - dtype: torch.dtype = torch.bfloat16 - kvcache_dtype: torch.dtype = torch.bfloat16 - num_layers: int = 1 - num_heads: int = 32 - num_kv_heads: int = 8 - head_dim: int = 128 - page_size: int = 32 - num_pages: int = 16 - batch_size: int = 4 - num_sparse_topk: int = 64 - - @property - def num_kv_groups(self) -> int: - return self.num_heads // self.num_kv_heads - - @property - def kv_cache_len(self) -> int: - return self.page_size * self.num_pages - - @property - def max_num_pages(self) -> int: - return self.batch_size * self.num_pages - -@dataclass(kw_only=True, frozen=False) -class SparseContextScenario(SparseScenario): - """Configuration for context phase tests with sparse kv cache write.""" - seq_lens: Tuple[int, ...] = (128,) +class _StubSparseParams(SparseParams): + """Minimal sparse parameters for framework-level tests.""" - def __post_init__(self): - if len(self.seq_lens) != self.batch_size: - raise ValueError( - f"seq_lens length {len(self.seq_lens)} must match batch_size {self.batch_size}" - ) - - @property - def max_seq_len(self) -> int: - return max(self.seq_lens) - - @property - def nnz_q(self) -> int: - return sum(self.seq_lens) - - -@dataclass(kw_only=True, frozen=False) -class SparseGenerationScenario(SparseScenario): - """Configuration for generation phase tests with sparse attention.""" - - past_kv_lens: Tuple[int, ...] = (256,) - num_contexts: int = 0 - - def __post_init__(self): - if len(self.past_kv_lens) != self.batch_size: - raise ValueError( - f"past_kv_lens length {len(self.past_kv_lens)} must match batch_size {self.batch_size}" - ) - - @property - def num_generations(self) -> int: - return self.batch_size - self.num_contexts - - @property - def max_past_kv_len(self) -> int: - return max(self.past_kv_lens) - - @property - def nnz_q(self) -> int: - return self.num_generations - - -class MockSparseParams(SparseParams): - """Sparse params stub used to exercise generic sparse attention plumbing.""" - - algorithm: str = "mqa_gqa" + algorithm: str = "test_sparse" @property def indices_block_size(self) -> int: return 1 -@dataclass -class TestSparseAttentionMetadata(TrtllmAttentionMetadata): - """Metadata for testing sparse attention.""" - - num_sparse_topk: int = 64 - - -class TestSparseAttention(TrtllmAttention): - """TrtllmAttention subclass for testing with predetermined sparse indices.""" - - def __init__( - self, - *args, - sparse_kv_indices: Optional[torch.Tensor] = None, - sparse_kv_offsets: Optional[torch.Tensor] = None, - sparse_attn_indices: Optional[torch.Tensor] = None, - sparse_attn_offsets: Optional[torch.Tensor] = None, - **kwargs, - ): - kwargs["sparse_params"] = MockSparseParams() - kwargs["pos_embd_params"] = None - super().__init__(*args, **kwargs) - - self._sparse_kv_indices = sparse_kv_indices - self._sparse_kv_offsets = sparse_kv_offsets - self._sparse_attn_indices = sparse_attn_indices - self._sparse_attn_offsets = sparse_attn_offsets +class _StaticPredictionAttention(TrtllmAttention): + """Backend stub that returns predetermined sparse predictions.""" def sparse_kv_predict(self, q, k, metadata, forward_args: AttentionForwardArgs): return self._sparse_kv_indices, self._sparse_kv_offsets @@ -187,9 +60,9 @@ def sparse_attn_predict(self, q, k, metadata, forward_args: AttentionForwardArgs return self._sparse_attn_indices, self._sparse_attn_offsets -def test_sparse_runtime_params() -> None: - attention = TestSparseAttention.__new__(TestSparseAttention) - attention.sparse_params = MockSparseParams() +def test_prepare_sparse_runtime_params_from_predictions() -> None: + attention = _StaticPredictionAttention.__new__(_StaticPredictionAttention) + attention.sparse_params = _StubSparseParams() attention._sparse_kv_indices = torch.tensor([1], dtype=torch.int32) attention._sparse_kv_offsets = torch.tensor([0, 1], dtype=torch.int32) attention._sparse_attn_indices = torch.tensor([2], dtype=torch.int32) @@ -214,7 +87,7 @@ def test_sparse_runtime_params() -> None: def test_sparse_attn_hook_registration() -> None: hook_module = ModuleType("sparse_attn_hook_registration") - hook_module.sparse_params = MockSparseParams() + hook_module.sparse_params = _StubSparseParams() hook_module.sparse_params.algorithm = "dsa" dsa_hooks = get_sparse_mla_hooks(hook_module) @@ -245,9 +118,9 @@ def test_sparse_attn_hook_registration() -> None: assert get_sparse_attention_hooks(hook_module) is not get_sparse_attention_hooks(hook_module) -def test_mla_backend_only_forward() -> None: +def test_mla_backend_only_forward_uses_default_path() -> None: backend_only_module = ModuleType("backend_only_sparse_attention") - backend_only_module.sparse_params = MockSparseParams() + backend_only_module.sparse_params = _StubSparseParams() backend_only_module.sparse_params.algorithm = "skip_softmax" hooks = get_sparse_mla_hooks(backend_only_module) assert hooks is None @@ -271,1089 +144,12 @@ def test_mla_backend_only_forward() -> None: ) -def test_sparse_runtime_params_without_prediction() -> None: +def test_prepare_sparse_runtime_params_without_predictions() -> None: attention = TrtllmAttention.__new__(TrtllmAttention) - attention.sparse_params = MockSparseParams() + attention.sparse_params = _StubSparseParams() runtime_params = prepare_sparse_runtime_params( attention, torch.empty(0), None, None, AttentionForwardArgs() ) assert runtime_params == SparseRuntimeParams() - - -def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor: - """Repeat kv heads to match query heads.""" - batch, num_key_value_heads, slen, head_dim = hidden_states.shape - if n_rep == 1: - return hidden_states - hidden_states = hidden_states[:, :, None, :, :].expand( - batch, num_key_value_heads, n_rep, slen, head_dim - ) - return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim) - - -def create_kv_cache_manager( - s: SparseScenario, kv_cache: Optional[torch.Tensor] = None -) -> KVCacheManager: - """Create kv cache manager for testing.""" - kv_cache_config = KvCacheConfig(max_tokens=s.max_num_pages * s.page_size) - mapping = Mapping(world_size=1, tp_size=1, rank=0) - - manager = KVCacheManager( - kv_cache_config, - tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF, - num_layers=s.num_layers, - num_kv_heads=s.num_kv_heads, - head_dim=s.head_dim, - tokens_per_block=s.page_size, - max_seq_len=s.max_num_pages * s.page_size, - max_batch_size=s.batch_size, - mapping=mapping, - dtype=str_dtype_to_binding(torch_dtype_to_str(s.kvcache_dtype)), - ) - - if kv_cache is not None: - for i in range(s.num_layers): - manager.get_buffers(i, kv_layout="HND").copy_(kv_cache[i]) - - return manager - - -def generate_sparse_kv_indices( - seq_lens: Tuple[int, ...], - num_kv_heads: int, - num_sparse_topk: int, - device: torch.device, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Generate sparse kv indices for context phase. - - For each request, pick min(num_sparse_topk, seq_len) indices from [0, seq_len). - Returns (indices [num_kv_heads, total_sparse], offsets [num_requests + 1]). - """ - all_indices = [] - offsets = [0] - - for seq_len in seq_lens: - pick = min(num_sparse_topk, seq_len) - batch_indices = [] - for _ in range(num_kv_heads): - indices = torch.randperm(seq_len, device=device)[:pick].sort().values - batch_indices.append(indices) - all_indices.append(torch.stack(batch_indices, dim=0)) - offsets.append(offsets[-1] + pick) - - indices = torch.cat(all_indices, dim=1).int() - offsets = torch.tensor(offsets, dtype=torch.int32, device=device) - return indices, offsets - - -def generate_sparse_attn_ctx_indices( - seq_lens: Tuple[int, ...], - num_kv_heads: int, - num_sparse_topk: int, - device: torch.device, -) -> torch.Tensor: - """ - Generate causal sparse attention indices for context phase. - - For each token at position pos (0-indexed within its request), available_kv_len = pos + 1. - If available_kv_len <= num_sparse_topk: select all [0..pos], pad rest with -1. - Otherwise: randomly pick num_sparse_topk from [0..pos]. - - Returns: [num_kv_heads, total_tokens, num_sparse_topk] with -1 padding. - """ - total_tokens = sum(seq_lens) - result = torch.full( - (num_kv_heads, total_tokens, num_sparse_topk), -1, dtype=torch.int32, device=device - ) - - token_offset = 0 - for seq_len in seq_lens: - for token_idx in range(seq_len): - available_kv_len = token_idx + 1 - pick = min(num_sparse_topk, available_kv_len) - - for head_idx in range(num_kv_heads): - indices = torch.randperm(available_kv_len, device=device)[:pick].sort().values - result[head_idx, token_offset + token_idx, :pick] = indices - - token_offset += seq_len - - return result - - -def generate_sparse_attn_gen_indices( - past_kv_lens: Tuple[int, ...], - num_kv_heads: int, - num_sparse_topk: int, - device: torch.device, -) -> torch.Tensor: - """ - Generate causal sparse attention indices for generation phase. - - Each generation token has past_kv_len available KV positions. - Pick min(num_sparse_topk, past_kv_len) indices, pad rest with -1. - - Returns: [num_kv_heads, num_generations, num_sparse_topk] with -1 padding. - """ - num_gens = len(past_kv_lens) - result = torch.full( - (num_kv_heads, num_gens, num_sparse_topk), -1, dtype=torch.int32, device=device - ) - - for gen_idx, past_kv_len in enumerate(past_kv_lens): - pick = min(num_sparse_topk, past_kv_len) - for head_idx in range(num_kv_heads): - indices = torch.randperm(past_kv_len, device=device)[:pick].sort().values - result[head_idx, gen_idx, :pick] = indices - - return result - - -def convert_sparse_indices_to_global( - sparse_indices: torch.Tensor, - metadata: TrtllmAttentionMetadata, - layer_idx: int = 0, - kv_factor: int = 2, -) -> torch.Tensor: - """ - Convert local sparse indices to global KV cache pool indices. - - Works for both context (variable-length Q packed) and generation (one token per request). - sparse_indices shape: [num_kv_heads, num_tokens, num_sparse_topk] - """ - num_kv_heads, num_tokens, num_sparse_tokens = sparse_indices.shape - device = sparse_indices.device - - tokens_per_block = metadata.kv_cache_manager.tokens_per_block - num_layers = metadata.kv_cache_manager.num_layers - stride_factor = num_layers * kv_factor * num_kv_heads * tokens_per_block - - # Build req_idx_per_token: map each token to its request index. - num_requests = len(metadata.request_ids) - seq_lens = metadata.seq_lens[:num_requests] - if hasattr(seq_lens, "cpu"): - seq_lens_cpu = seq_lens.cpu() - else: - seq_lens_cpu = torch.tensor(seq_lens, dtype=torch.int32) - req_idx_per_token = torch.repeat_interleave( - torch.arange(num_requests, dtype=torch.int32), seq_lens_cpu, dim=0 - ).to(device) - - # Build 2D block table: [num_requests, max_pages] - request_ids = metadata.request_ids - page_indices = metadata.kv_cache_manager.get_batch_cache_indices(request_ids) - max_pages = max(len(p) for p in page_indices) if page_indices else 1 - host_block_table = torch.full((num_requests, max_pages), -1, dtype=torch.int32) - for i, pages in enumerate(page_indices): - if len(pages) > 0: - host_block_table[i, : len(pages)] = torch.tensor(pages, dtype=torch.int32) - block_table = host_block_table.to(device) - - # Convert to global - global_indices = triton_convert_req_index_to_global_index( - req_idx_per_token, - block_table, - sparse_indices, - BLOCK_SIZE=tokens_per_block, - NUM_TOPK_TOKENS=num_sparse_tokens, - BLOCK_N=min(64, num_sparse_tokens), - stride_factor=stride_factor, - layer_id=layer_idx, - num_kv_heads=num_kv_heads, - kv_factor=kv_factor, - ) - - return global_indices - - -def _extract_batch_tensors( - tensor: torch.Tensor, offset: int, length: int, shape_per_token: Tuple -) -> torch.Tensor: - """Extract and reshape tensors for a specific batch.""" - return tensor[offset : offset + length].view(length, *shape_per_token) - - -def build_expected_sparse_kv( - k: torch.Tensor, - v: torch.Tensor, - sparse_kv_indices: torch.Tensor, - sparse_kv_offsets: torch.Tensor, - s: SparseContextScenario, -) -> List[Tuple[torch.Tensor, torch.Tensor]]: - """Build expected sparse K and V values based on sparse indices.""" - expected_kvs = [] - token_offset = 0 - - for batch_idx, seq_len in enumerate(s.seq_lens): - sparse_len = sparse_kv_offsets[batch_idx + 1].item() - sparse_kv_offsets[batch_idx].item() - k_batch = _extract_batch_tensors(k, token_offset, seq_len, (s.num_kv_heads, s.head_dim)) - v_batch = _extract_batch_tensors(v, token_offset, seq_len, (s.num_kv_heads, s.head_dim)) - - expected_k = torch.zeros( - sparse_len, s.num_kv_heads, s.head_dim, device=k.device, dtype=k.dtype - ) - expected_v = torch.zeros_like(expected_k) - - start, end = sparse_kv_offsets[batch_idx].item(), sparse_kv_offsets[batch_idx + 1].item() - for head_idx in range(s.num_kv_heads): - indices = sparse_kv_indices[head_idx, start:end] - expected_k[:, head_idx] = k_batch[indices, head_idx] - expected_v[:, head_idx] = v_batch[indices, head_idx] - - expected_kvs.append((expected_k, expected_v)) - token_offset += seq_len - - return expected_kvs - - -def _extract_tokens_from_cache( - kv_buffer: torch.Tensor, - block_ids: List[int], - num_tokens: int, - num_kv_heads: int, - head_dim: int, - page_size: int, - dtype: torch.dtype, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Extract tokens from paged kv cache.""" - device = kv_buffer.device - k_cache = torch.zeros(num_tokens, num_kv_heads, head_dim, device=device, dtype=dtype) - v_cache = torch.zeros_like(k_cache) - - for token_idx in range(num_tokens): - block_idx = token_idx // page_size - offset_in_block = token_idx % page_size - block_id = block_ids[block_idx] - - for head_idx in range(num_kv_heads): - k_cache[token_idx, head_idx] = kv_buffer[block_id, 0, head_idx, offset_in_block, :].to( - dtype - ) - v_cache[token_idx, head_idx] = kv_buffer[block_id, 1, head_idx, offset_in_block, :].to( - dtype - ) - - return k_cache, v_cache - - -def extract_kv_from_paged_cache( - kv_cache_manager: KVCacheManager, - request_ids: List[int], - sparse_kv_offsets: torch.Tensor, - s: SparseContextScenario, - dtype: torch.dtype, -) -> List[Tuple[torch.Tensor, torch.Tensor]]: - """Extract K and V values from paged kv cache.""" - kv_buffer = kv_cache_manager.get_buffers(0, kv_layout="HND") - kv_caches = [] - - for batch_idx in range(s.batch_size): - num_sparse_tokens = ( - sparse_kv_offsets[batch_idx + 1].item() - sparse_kv_offsets[batch_idx].item() - ) - block_ids = kv_cache_manager.get_block_ids_per_seq([request_ids[batch_idx]])[0] - k_cache, v_cache = _extract_tokens_from_cache( - kv_buffer, block_ids, num_sparse_tokens, s.num_kv_heads, s.head_dim, s.page_size, dtype - ) - kv_caches.append((k_cache, v_cache)) - - return kv_caches - - -def _compute_causal_attention( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - num_kv_groups: int, -) -> torch.Tensor: - """Compute causal attention for a single batch.""" - seq_len = q.shape[2] - head_dim = q.shape[3] - - k_expanded = repeat_kv(k, num_kv_groups) - v_expanded = repeat_kv(v, num_kv_groups) - - attn_weights = torch.matmul(q, k_expanded.transpose(-1, -2)) / math.sqrt(head_dim) - causal_mask = torch.triu( - torch.full((seq_len, seq_len), float("-inf"), device=q.device), diagonal=1 - ) - attn_weights = attn_weights + causal_mask - attn_weights = torch.nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to( - q.dtype - ) - output = torch.matmul(attn_weights, v_expanded) - - return output - - -def reference_context_attention( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - s: SparseContextScenario, -) -> torch.Tensor: - """Reference implementation for context phase.""" - outputs = [] - token_offset = 0 - - for seq_len in s.seq_lens: - q_batch = _extract_batch_tensors(q, token_offset, seq_len, (s.num_heads, s.head_dim)) - k_batch = _extract_batch_tensors(k, token_offset, seq_len, (s.num_kv_heads, s.head_dim)) - v_batch = _extract_batch_tensors(v, token_offset, seq_len, (s.num_kv_heads, s.head_dim)) - - q_batch = q_batch.view(1, seq_len, s.num_heads, s.head_dim).transpose(1, 2) - k_batch = k_batch.view(1, seq_len, s.num_kv_heads, s.head_dim).transpose(1, 2) - v_batch = v_batch.view(1, seq_len, s.num_kv_heads, s.head_dim).transpose(1, 2) - - output_batch = _compute_causal_attention(q_batch, k_batch, v_batch, s.num_kv_groups) - output_batch = output_batch.transpose(1, 2).reshape(seq_len, s.num_heads * s.head_dim) - outputs.append(output_batch) - token_offset += seq_len - - return torch.cat(outputs, dim=0) - - -def reference_context_sparse_attention( - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - sparse_attn_ctx_indices: torch.Tensor, - s: SparseContextScenario, -) -> torch.Tensor: - """ - Reference implementation for context phase with sparse attention. - Uses mask-based approach for each KV head. - """ - total_tokens = sum(s.seq_lens) - device = q.device - dtype = q.dtype - - # Reshape inputs: [num_tokens, num_heads, head_dim] - q_reshaped = q.view(total_tokens, s.num_heads, s.head_dim) - k_reshaped = k.view(total_tokens, s.num_kv_heads, s.head_dim) - v_reshaped = v.view(total_tokens, s.num_kv_heads, s.head_dim) - - outputs = [] - token_offset = 0 - - for seq_len in s.seq_lens: - q_batch = q_reshaped[ - token_offset : token_offset + seq_len - ] # [seq_len, num_heads, head_dim] - k_batch = k_reshaped[ - token_offset : token_offset + seq_len - ] # [seq_len, num_kv_heads, head_dim] - v_batch = v_reshaped[ - token_offset : token_offset + seq_len - ] # [seq_len, num_kv_heads, head_dim] - - batch_output = [] - - # Process each KV head - for kv_head_idx in range(s.num_kv_heads): - k_head = k_batch[:, kv_head_idx, :] - v_head = v_batch[:, kv_head_idx, :] - - # Build sparse mask for this head - sparse_mask = torch.full( - (seq_len, seq_len), float("-inf"), device=device, dtype=torch.float32 - ) - - for token_idx in range(seq_len): - global_token_idx = token_offset + token_idx - # Get sparse indices for this token: [num_sparse_tokens] - indices = sparse_attn_ctx_indices[kv_head_idx, global_token_idx] - # Filter out -1 padding - valid_indices = indices[indices >= 0] - # Set mask values to 0 for valid positions - sparse_mask[token_idx, valid_indices] = 0.0 - - # Apply causal mask on top of sparse mask - causal_mask = torch.triu( - torch.full((seq_len, seq_len), float("-inf"), device=device, dtype=torch.float32), - diagonal=1, - ) - combined_mask = sparse_mask + causal_mask - - # Process each query head in this KV group - for group_idx in range(s.num_kv_groups): - q_head_idx = kv_head_idx * s.num_kv_groups + group_idx - q_head = q_batch[:, q_head_idx, :] # [seq_len, head_dim] - - attn_scores = torch.matmul(q_head, k_head.T) / math.sqrt(s.head_dim) - attn_scores = attn_scores + combined_mask - attn_weights = torch.nn.functional.softmax( - attn_scores, dim=-1, dtype=torch.float32 - ).to(dtype) - - out_head = torch.matmul(attn_weights, v_head) - batch_output.append(out_head) - - # Concatenate all heads: [seq_len, num_heads, head_dim] -> [seq_len, num_heads * head_dim] - batch_output = torch.stack(batch_output, dim=1) - batch_output = batch_output.reshape(seq_len, s.num_heads * s.head_dim) - outputs.append(batch_output) - - token_offset += seq_len - - return torch.cat(outputs, dim=0) - - -def _get_selected_pages_tokens( - token_indices: torch.Tensor, - page_size: int, - kv_len: int, - device: torch.device, -) -> torch.Tensor: - """Convert token indices to page indices and gather all tokens from selected pages.""" - if len(token_indices) == 0: - return torch.tensor([], dtype=torch.long, device=device) - - page_indices = torch.unique((token_indices // page_size).sort().values) - selected_tokens = [] - - for page_idx in page_indices: - token_start = page_idx * page_size - token_end = min(token_start + page_size, kv_len) - selected_tokens.append(torch.arange(token_start, token_end, device=device)) - - return ( - torch.cat(selected_tokens) - if selected_tokens - else torch.tensor([], dtype=torch.long, device=device) - ) - - -def _compute_sparse_attention_per_head( - q_head: torch.Tensor, - k_sparse: torch.Tensor, - v_sparse: torch.Tensor, - head_dim: int, -) -> torch.Tensor: - """Compute attention for a single query head.""" - if len(k_sparse) == 0: - return torch.zeros(head_dim, device=q_head.device, dtype=q_head.dtype) - - attn_weights = torch.matmul(q_head, k_sparse.T) / math.sqrt(head_dim) - attn_weights = torch.nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to( - q_head.dtype - ) - return torch.matmul(attn_weights, v_sparse) - - -def reference_generation_sparse_attention( - q: torch.Tensor, - k_cache: torch.Tensor, - v_cache: torch.Tensor, - k_new: torch.Tensor, - v_new: torch.Tensor, - sparse_attn_indices: torch.Tensor, - s: SparseGenerationScenario, -) -> torch.Tensor: - """Reference implementation for generation phase with sparse attention. - - Args: - sparse_attn_indices: [num_kv_heads, num_gens, num_sparse_topk] with -1 padding. - """ - outputs = [] - - for gen_idx in range(s.num_generations): - batch_idx = s.num_contexts + gen_idx - past_kv_len = s.past_kv_lens[batch_idx] - kv_len = past_kv_len + 1 - - k_full = k_cache[batch_idx, :kv_len].clone() - v_full = v_cache[batch_idx, :kv_len].clone() - k_full[past_kv_len] = k_new[gen_idx].view(s.num_kv_heads, s.head_dim) - v_full[past_kv_len] = v_new[gen_idx].view(s.num_kv_heads, s.head_dim) - - q_batch = q[gen_idx].view(s.num_heads, s.head_dim) - head_outputs = [] - - for kv_head_idx in range(s.num_kv_heads): - token_indices = sparse_attn_indices[kv_head_idx, gen_idx] - valid_indices = token_indices[token_indices >= 0].long() - - if len(valid_indices) == 0: - head_outputs.extend( - [torch.zeros(s.head_dim, device=q.device, dtype=q.dtype)] * s.num_kv_groups - ) - continue - - k_sparse = k_full[valid_indices, kv_head_idx, :] - v_sparse = v_full[valid_indices, kv_head_idx, :] - - for group_idx in range(s.num_kv_groups): - q_head_idx = kv_head_idx * s.num_kv_groups + group_idx - out_head = _compute_sparse_attention_per_head( - q_batch[q_head_idx], k_sparse, v_sparse, s.head_dim - ) - head_outputs.append(out_head) - - outputs.append(torch.cat(head_outputs, dim=0)) - - return torch.stack(outputs, dim=0) - - -def _setup_context_test(s: SparseContextScenario): - """Setup common components for context test.""" - device = torch.device("cuda") - torch.manual_seed(42) - num_sparse_topk = s.num_sparse_topk - - q = torch.randn(s.nnz_q, s.num_heads * s.head_dim, device=device, dtype=s.dtype) - k = torch.randn(s.nnz_q, s.num_kv_heads * s.head_dim, device=device, dtype=s.dtype) - v = torch.randn(s.nnz_q, s.num_kv_heads * s.head_dim, device=device, dtype=s.dtype) - sparse_kv_indices, sparse_kv_offsets = generate_sparse_kv_indices( - s.seq_lens, s.num_kv_heads, num_sparse_topk, device - ) - - kv_cache = torch.zeros( - s.num_layers, - s.max_num_pages, - 2, - s.num_kv_heads, - s.page_size, - s.head_dim, - device=device, - dtype=s.kvcache_dtype, - ) - kv_cache_manager = create_kv_cache_manager(s, kv_cache) - - request_ids = list(range(s.batch_size)) - kv_cache_manager.add_dummy_requests(request_ids, list(s.seq_lens)) - - metadata = TestSparseAttentionMetadata( - num_contexts=s.batch_size, - kv_cache_params=KVCacheParams(use_cache=True, num_cached_tokens_per_seq=[0] * s.batch_size), - seq_lens=torch.tensor(s.seq_lens, dtype=torch.int32), - max_num_requests=s.batch_size, - max_num_tokens=s.nnz_q, - kv_cache_manager=kv_cache_manager, - request_ids=request_ids, - prompt_lens=list(s.seq_lens), - num_sparse_topk=num_sparse_topk, - ) - metadata.prepare() - - attention = TestSparseAttention( - layer_idx=0, - num_heads=s.num_heads, - head_dim=s.head_dim, - num_kv_heads=s.num_kv_heads, - sparse_kv_indices=sparse_kv_indices, - sparse_kv_offsets=sparse_kv_offsets, - ) - - return ( - device, - q, - k, - v, - sparse_kv_indices, - sparse_kv_offsets, - kv_cache_manager, - request_ids, - metadata, - attention, - ) - - -def _setup_generation_test(s: SparseGenerationScenario): - """Setup common components for generation test.""" - device = torch.device("cuda") - torch.manual_seed(42) - num_sparse_topk = s.num_sparse_topk - - token_nums = [past_len + 1 for past_len in s.past_kv_lens] - - q = torch.randn(s.num_generations, s.num_heads * s.head_dim, device=device, dtype=s.dtype) - k_new = torch.randn( - s.num_generations, s.num_kv_heads * s.head_dim, device=device, dtype=s.dtype - ) - v_new = torch.randn( - s.num_generations, s.num_kv_heads * s.head_dim, device=device, dtype=s.dtype - ) - - gen_past_kv_lens = tuple(s.past_kv_lens[s.num_contexts + i] for i in range(s.num_generations)) - # Local sparse indices: [num_kv_heads, num_gens, num_sparse_topk] - sparse_attn_indices = generate_sparse_attn_gen_indices( - gen_past_kv_lens, s.num_kv_heads, num_sparse_topk, device - ) - - kv_cache = torch.randn( - s.num_layers, - s.max_num_pages, - 2, - s.num_kv_heads, - s.page_size, - s.head_dim, - device=device, - dtype=s.kvcache_dtype, - ) - kv_cache_manager = create_kv_cache_manager(s, kv_cache) - - request_ids = list(range(s.batch_size)) - kv_cache_manager.add_dummy_requests(request_ids, token_nums) - - metadata = TestSparseAttentionMetadata( - num_contexts=s.num_contexts, - kv_cache_params=KVCacheParams( - use_cache=True, num_cached_tokens_per_seq=list(s.past_kv_lens) - ), - seq_lens=torch.tensor([1] * s.num_generations).int(), - max_num_requests=s.batch_size, - max_num_tokens=s.num_generations, - kv_cache_manager=kv_cache_manager, - request_ids=request_ids, - prompt_lens=list(s.past_kv_lens), - num_sparse_topk=num_sparse_topk, - ) - metadata.prepare() - - # Convert local indices to global KV cache pool indices - global_sparse_attn_indices = convert_sparse_indices_to_global( - sparse_attn_indices, metadata, layer_idx=0 - ) - - attention = TestSparseAttention( - layer_idx=0, - num_heads=s.num_heads, - head_dim=s.head_dim, - num_kv_heads=s.num_kv_heads, - sparse_attn_indices=global_sparse_attn_indices, - ) - - return ( - device, - q, - k_new, - v_new, - sparse_attn_indices, - kv_cache_manager, - request_ids, - metadata, - attention, - ) - - -def _build_reference_kv_cache( - kv_cache_manager, request_ids, s: SparseGenerationScenario, device, dtype -): - """Build reference K, V cache from paged format.""" - k_cache_ref = torch.zeros( - s.batch_size, s.kv_cache_len, s.num_kv_heads, s.head_dim, device=device, dtype=dtype - ) - v_cache_ref = torch.zeros_like(k_cache_ref) - - kv_buffer = kv_cache_manager.get_buffers(0, kv_layout="HND") - for batch_idx, past_kv_len in enumerate(s.past_kv_lens): - block_ids = kv_cache_manager.get_block_ids_per_seq([request_ids[batch_idx]])[0] - for block_local_idx, block_id in enumerate(block_ids): - token_start = block_local_idx * s.page_size - token_end = min(token_start + s.page_size, past_kv_len) - tokens_in_block = token_end - token_start - - for head_idx in range(s.num_kv_heads): - k_cache_ref[batch_idx, token_start:token_end, head_idx] = kv_buffer[ - block_id, 0, head_idx, :tokens_in_block, : - ].to(dtype) - v_cache_ref[batch_idx, token_start:token_end, head_idx] = kv_buffer[ - block_id, 1, head_idx, :tokens_in_block, : - ].to(dtype) - - return k_cache_ref, v_cache_ref - - -@requires_sparse_mqa_gqa -@pytest.mark.parametrize( - "s", - [ - SparseContextScenario(batch_size=2, seq_lens=(48, 64), num_pages=8), - SparseContextScenario(batch_size=4, seq_lens=(96, 112, 128, 144), num_pages=16), - SparseContextScenario(batch_size=1, seq_lens=(256,), num_pages=8), - SparseContextScenario(batch_size=3, seq_lens=(64, 96, 128), num_pages=12), - ], - ids=["batch2_var_seq", "batch4_var_seq", "batch1_seq256", "batch3_var_seq"], -) -def test_context_sparse_kv(s: SparseContextScenario): - """Test context phase with sparse kv cache write.""" - ( - device, - q, - k, - v, - sparse_kv_indices, - sparse_kv_offsets, - kv_cache_manager, - request_ids, - metadata, - attention, - ) = _setup_context_test(s) - - ref_output = reference_context_attention(q.clone(), k.clone(), v.clone(), s) - expected_kvs = build_expected_sparse_kv( - k.clone(), v.clone(), sparse_kv_indices, sparse_kv_offsets, s - ) - - qkv = torch.cat([q, k, v], dim=1) - output = attention.forward(qkv, None, None, metadata) - - assert output.shape == ref_output.shape, f"Shape mismatch: {output.shape} vs {ref_output.shape}" - torch.testing.assert_close(output, ref_output, atol=ATOL, rtol=RTOL) - print(f"Context sparse kv attention output test passed: {s}") - - actual_kvs = extract_kv_from_paged_cache( - kv_cache_manager, request_ids, sparse_kv_offsets, s, s.dtype - ) - - for batch_idx in range(s.batch_size): - actual_k, actual_v = actual_kvs[batch_idx] - expected_k, expected_v = expected_kvs[batch_idx] - torch.testing.assert_close( - actual_k, - expected_k, - atol=ATOL, - rtol=RTOL, - msg=f"K cache mismatch for batch {batch_idx} after sparse compaction", - ) - torch.testing.assert_close( - actual_v, - expected_v, - atol=ATOL, - rtol=RTOL, - msg=f"V cache mismatch for batch {batch_idx} after sparse compaction", - ) - - print(f"Context sparse kv cache content test passed: {s}") - kv_cache_manager.shutdown() - - -@requires_sparse_mqa_gqa -@pytest.mark.parametrize( - "s", - [ - # Basic scenarios - SparseGenerationScenario( - batch_size=2, - past_kv_lens=(96, 128), - num_pages=16, - ), - SparseGenerationScenario( - batch_size=4, - past_kv_lens=(192, 224, 256, 288), - num_pages=32, - ), - SparseGenerationScenario(batch_size=1, past_kv_lens=(64,), num_pages=8), - SparseGenerationScenario( - batch_size=3, - past_kv_lens=(128, 160, 192), - num_pages=24, - ), - # GQA ratios: MQA (8Q/1KV), GQA 4:1, GQA 2:1 - SparseGenerationScenario( - num_heads=8, - num_kv_heads=1, - batch_size=2, - past_kv_lens=(96, 128), - num_pages=16, - ), - SparseGenerationScenario( - num_heads=16, - num_kv_heads=4, - batch_size=2, - past_kv_lens=(128, 256), - num_pages=16, - ), - SparseGenerationScenario( - num_heads=8, - num_kv_heads=4, - batch_size=2, - past_kv_lens=(128, 256), - num_pages=16, - ), - # Maximum supported query-head group size - SparseGenerationScenario( - num_heads=32, - num_kv_heads=1, - batch_size=1, - past_kv_lens=(128,), - num_pages=8, - ), - SparseGenerationScenario( - num_heads=64, - num_kv_heads=2, - batch_size=1, - past_kv_lens=(128,), - num_pages=8, - ), - # topk: minimum (4), topk exceeding some past_kv_lens - SparseGenerationScenario( - batch_size=1, - past_kv_lens=(128,), - num_pages=8, - num_sparse_topk=4, - ), - SparseGenerationScenario( - batch_size=2, - past_kv_lens=(32, 256), - num_pages=16, - num_sparse_topk=128, - ), - # Large batch - SparseGenerationScenario( - batch_size=8, - past_kv_lens=(64, 96, 128, 160, 192, 224, 256, 288), - num_pages=64, - ), - # Page boundary: page_size=64 - SparseGenerationScenario( - page_size=64, - batch_size=2, - past_kv_lens=(64, 192), - num_pages=8, - ), - ], - ids=[ - "batch2_var_kv", - "batch4_var_kv", - "batch1_kv64", - "batch3_var_kv", - "mqa_8q1kv", - "gqa_16q4kv", - "gqa_8q4kv", - "mqa_group32_boundary", - "gqa_group32_boundary", - "topk4_min", - "topk128_exceeds_some", - "batch8_varied", - "page_size_64", - ], -) -def test_generation_sparse_attention(s: SparseGenerationScenario): - """Test generation phase with sparse attention computation.""" - _run_generation_sparse_attention(s) - - -def _run_generation_sparse_attention(s: SparseGenerationScenario): - """Run a generation scenario and compare the sparse kernel with PyTorch.""" - ( - device, - q, - k_new, - v_new, - sparse_attn_indices, - kv_cache_manager, - request_ids, - metadata, - attention, - ) = _setup_generation_test(s) - try: - k_cache_ref, v_cache_ref = _build_reference_kv_cache( - kv_cache_manager, request_ids, s, device, s.dtype - ) - ref_sparse_output = reference_generation_sparse_attention( - q, k_cache_ref, v_cache_ref, k_new, v_new, sparse_attn_indices, s - ) - - qkv = torch.cat([q, k_new, v_new], dim=1) - output = attention.forward(qkv, None, None, metadata) - - expected_shape = (s.num_generations, s.num_heads * s.head_dim) - assert output.shape == expected_shape, f"Shape mismatch: {output.shape} vs {expected_shape}" - assert torch.isfinite(output).all(), "Output contains non-finite values" - - torch.testing.assert_close(output, ref_sparse_output, atol=ATOL, rtol=RTOL) - print(f"Generation sparse attention test passed: {s}") - finally: - kv_cache_manager.shutdown() - - -@requires_sparse_mqa_gqa -@pytest.mark.parametrize("num_kv_heads", [1, 4], ids=["mqa", "gqa_2to1"]) -@pytest.mark.parametrize( - ("dtype", "head_dim"), - [ - (torch.bfloat16, 64), - (torch.bfloat16, 80), - (torch.bfloat16, 128), - (torch.bfloat16, 256), - (torch.float16, 64), - (torch.float16, 80), - (torch.float16, 128), - (torch.float16, 256), - ], - ids=[ - "bf16_h64", - "bf16_h80", - "bf16_h128", - "bf16_h256", - "fp16_h64", - "fp16_h80", - "fp16_h128", - "fp16_h256", - ], -) -def test_generation_sparse_mqa_gqa_kernel_matrix( - monkeypatch: pytest.MonkeyPatch, - num_kv_heads: int, - dtype: torch.dtype, - head_dim: int, -): - """Cover every supported head dimension and both supported input dtypes. - - Each KV head owns one token-index list. All query heads in its MQA/GQA - group share that list, which is the layout expected by the kernel. - """ - monkeypatch.setenv("TLLM_FMHA_LIBS", "fallback") - scenario = SparseGenerationScenario( - dtype=dtype, - kvcache_dtype=dtype, - num_heads=8, - num_kv_heads=num_kv_heads, - head_dim=head_dim, - batch_size=1, - past_kv_lens=(64,), - num_pages=4, - num_sparse_topk=32, - ) - _run_generation_sparse_attention(scenario) - - -@requires_sparse_mqa_gqa -@pytest.mark.parametrize( - "s", - [ - # MQA (8Q/1KV) - SparseContextScenario( - batch_size=2, - seq_lens=(128, 64), - num_pages=8, - num_kv_heads=1, - num_heads=8, - ), - # GQA 4:1 (8Q/2KV) - SparseContextScenario( - batch_size=2, - seq_lens=(128, 64), - num_pages=8, - num_kv_heads=2, - num_heads=8, - ), - # GQA 4:1 (16Q/4KV) with 3 batches - SparseContextScenario( - batch_size=3, - seq_lens=(64, 96, 128), - num_pages=12, - num_kv_heads=4, - num_heads=16, - ), - # GQA 8:1 (32Q/4KV) - SparseContextScenario( - batch_size=2, - seq_lens=(64, 128), - num_pages=8, - num_kv_heads=4, - num_heads=32, - ), - # topk=4 (very sparse) - SparseContextScenario( - batch_size=2, - seq_lens=(64, 128), - num_pages=8, - num_kv_heads=1, - num_heads=8, - num_sparse_topk=4, - ), - # topk=128 (near-dense, topk >= seq_len for some requests) - SparseContextScenario( - batch_size=2, - seq_lens=(64, 128), - num_pages=8, - num_kv_heads=2, - num_heads=8, - num_sparse_topk=128, - ), - ], - ids=[ - "mqa_8q1kv", - "gqa_8q2kv", - "gqa_16q4kv_batch3", - "gqa_32q4kv", - "topk4_very_sparse", - "topk128_near_dense", - ], -) -def test_context_sparse_attention_mqa(s: SparseContextScenario): - """Test context phase with sparse attention using sparse_attn_ctx_indices.""" - ( - device, - q, - k, - v, - sparse_kv_indices, - sparse_kv_offsets, - kv_cache_manager, - request_ids, - metadata, - _, - ) = _setup_context_test(s) - - num_sparse_topk = metadata.num_sparse_topk - - # Generate causal sparse attention indices, padded to num_sparse_topk - sparse_attn_ctx_indices = generate_sparse_attn_ctx_indices( - s.seq_lens, s.num_kv_heads, num_sparse_topk, device - ) - assert sparse_attn_ctx_indices.shape[-1] == num_sparse_topk - - # Convert to global indices for attentionOp - global_sparse_attn_ctx_indices = convert_sparse_indices_to_global( - sparse_attn_ctx_indices, metadata, layer_idx=0 - ) - - # Compute reference output using local indices - ref_output = reference_context_sparse_attention( - q.clone(), k.clone(), v.clone(), sparse_attn_ctx_indices, s - ) - - # Verify reference output shape - total_tokens = sum(s.seq_lens) - expected_shape = (total_tokens, s.num_heads * s.head_dim) - assert ref_output.shape == expected_shape, ( - f"Reference output shape mismatch: {ref_output.shape} vs {expected_shape}" - ) - assert torch.isfinite(ref_output).all(), "Reference output contains non-finite values" - - print(f"Context sparse attention MQA reference test passed: {s}") - - attention = TestSparseAttention( - layer_idx=0, - num_heads=s.num_heads, - head_dim=s.head_dim, - num_kv_heads=s.num_kv_heads, - sparse_kv_indices=sparse_kv_indices, - sparse_kv_offsets=sparse_kv_offsets, - sparse_attn_indices=global_sparse_attn_ctx_indices, - ) - - qkv = torch.cat([q, k, v], dim=1) - output = attention.forward(qkv, None, None, metadata) - torch.testing.assert_close(output, ref_output, atol=ATOL, rtol=RTOL) - print(f"Context sparse attention MQA forward test passed: {s}") - - kv_cache_manager.shutdown() - - -if __name__ == "__main__": - s = SparseContextScenario( - batch_size=2, - seq_lens=(128, 64), - num_pages=8, - num_kv_heads=1, - num_heads=8, - head_dim=128, - ) - test_context_sparse_attention_mqa(s) diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py new file mode 100644 index 000000000000..30aee156046f --- /dev/null +++ b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py @@ -0,0 +1,1330 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Regression tests for the internal token-sparse MQA/GQA kernel. + +These tests replace the model-specific sparse selector with static token-index +lists, then exercise the same backend path used by ``TrtllmAttention``: + +1. Build request-local ``int32`` token indices. +2. Translate them to paged KV-cache pool indices. +3. Return them from the sparse prediction hooks. +4. Call ``TrtllmAttention.forward`` and compare with a PyTorch reference. + +This is an executable backend-integration example, not a public +``SparseAttentionConfig`` algorithm. Algorithm-independent sparse framework +tests remain in ``test_sparse_attention.py``. +""" + +import math +from dataclasses import dataclass +from typing import List, Optional, Tuple + +import pytest +import torch +from utils.util import getSMVersion + +import tensorrt_llm +from tensorrt_llm._torch.attention_backend.interface import ( + AttentionForwardArgs, + AttentionRuntimeFeatures, +) +from tensorrt_llm._torch.attention_backend.sparse.dsa.kernels import ( + triton_convert_req_index_to_global_index, +) +from tensorrt_llm._torch.attention_backend.sparse.params import SparseParams +from tensorrt_llm._torch.attention_backend.trtllm import ( + TrtllmAttention, + TrtllmAttentionMetadata, + generate_spec_decoding_packed_mask, + generate_spec_decoding_position_offsets, +) +from tensorrt_llm._torch.metadata import KVCacheParams +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm._utils import is_sm_100f, str_dtype_to_binding, torch_dtype_to_str +from tensorrt_llm.bindings.executor import KvCacheConfig +from tensorrt_llm.mapping import Mapping +from tensorrt_llm.models.modeling_utils import QuantConfig +from tensorrt_llm.quantization.mode import QuantAlgo + +ATOL = 2e-2 +RTOL = 2e-2 +FP8_ATOL = 4e-1 +FP8_RTOL = 4e-2 +SUPPORTED_DTYPES = (torch.bfloat16, torch.float16) +SUPPORTED_KV_CACHE_DTYPES = (*SUPPORTED_DTYPES, torch.float8_e4m3fn) +SUPPORTED_HEAD_DIMS = (64, 80, 128, 256) +MAX_Q_HEADS_PER_KV_HEAD = 32 + +pytestmark = pytest.mark.skipif( + not is_sm_100f(getSMVersion()), + reason="Sparse MQA/GQA requires an SM100-family GPU (SM100 or SM103)", +) + + +@pytest.fixture(autouse=True) +def _force_trtllm_gen_fallback(monkeypatch: pytest.MonkeyPatch) -> None: + """Keep every test on the internal TRTLLM-Gen fallback path.""" + monkeypatch.setenv("TLLM_FMHA_LIBS", "fallback") + + +# Kernel contract and static selector adapter. + + +@dataclass(kw_only=True, frozen=True) +class SparseMqaGqaScenario: + """Kernel geometry shared by context and generation scenarios. + + The validation mirrors the supported static token-sparse kernel contract, + so every scenario is also a compact declaration of a supported shape. + """ + + dtype: torch.dtype = torch.bfloat16 + kvcache_dtype: torch.dtype = torch.bfloat16 + num_layers: int = 1 + num_heads: int = 32 + num_kv_heads: int = 8 + head_dim: int = 128 + page_size: int = 32 + num_pages: int = 16 + batch_size: int = 1 + num_sparse_topk: int = 64 + + def __post_init__(self) -> None: + if self.dtype not in SUPPORTED_DTYPES: + raise ValueError("Model QKV dtype must be BF16 or FP16") + if self.kvcache_dtype not in SUPPORTED_KV_CACHE_DTYPES: + raise ValueError("KV-cache dtype must be BF16, FP16, or E4M3 FP8") + if self.kvcache_dtype != torch.float8_e4m3fn and self.kvcache_dtype != self.dtype: + raise ValueError("A non-FP8 KV-cache dtype must match the model QKV dtype") + if self.head_dim not in SUPPORTED_HEAD_DIMS: + raise ValueError(f"head_dim must be one of {SUPPORTED_HEAD_DIMS}") + if self.num_heads % self.num_kv_heads != 0: + raise ValueError("num_heads must be divisible by num_kv_heads") + if self.q_heads_per_kv_head > MAX_Q_HEADS_PER_KV_HEAD: + raise ValueError(f"at most {MAX_Q_HEADS_PER_KV_HEAD} query heads may share one KV head") + if self.page_size < 8 or self.page_size & (self.page_size - 1): + raise ValueError("page_size must be a power of two and at least 8") + if self.num_sparse_topk <= 0 or self.num_sparse_topk % 4 != 0: + raise ValueError("num_sparse_topk must be a positive multiple of 4") + + @property + def q_heads_per_kv_head(self) -> int: + return self.num_heads // self.num_kv_heads + + @property + def kv_pool_num_pages(self) -> int: + return self.batch_size * self.num_pages + + +@dataclass(kw_only=True, frozen=True) +class ContextScenario(SparseMqaGqaScenario): + """Packed context requests used for cache compaction or sparse compute.""" + + seq_lens: Tuple[int, ...] = (128,) + + def __post_init__(self) -> None: + super().__post_init__() + if len(self.seq_lens) != self.batch_size: + raise ValueError( + f"seq_lens length {len(self.seq_lens)} must match batch_size {self.batch_size}" + ) + + @property + def max_seq_len(self) -> int: + return max(self.seq_lens) + + @property + def nnz_q(self) -> int: + return sum(self.seq_lens) + + +@dataclass(kw_only=True, frozen=True) +class GenerationScenario(SparseMqaGqaScenario): + """One decode token per request with an existing paged KV history.""" + + past_kv_lens: Tuple[int, ...] = (256,) + query_len: int = 1 + fp8_output: bool = False + + def __post_init__(self) -> None: + super().__post_init__() + if len(self.past_kv_lens) != self.batch_size: + raise ValueError( + f"past_kv_lens length {len(self.past_kv_lens)} must match batch_size {self.batch_size}" + ) + if self.query_len < 1: + raise ValueError("query_len must be positive") + if self.fp8_output and self.kvcache_dtype != torch.float8_e4m3fn: + raise ValueError("FP8 output testing requires an FP8 KV cache") + + @property + def num_generations(self) -> int: + return self.batch_size + + @property + def nnz_q(self) -> int: + return self.batch_size * self.query_len + + @property + def max_query_len(self) -> int: + return self.query_len + + @property + def has_draft_tokens(self) -> bool: + return self.max_query_len > 1 + + +class _SparseMqaGqaParams(SparseParams): + """Token-granular parameters that select the internal MQA/GQA path.""" + + algorithm: str = "mqa_gqa" + + @property + def indices_block_size(self) -> int: + return 1 + + +@dataclass +class _SparseMqaGqaMetadata(TrtllmAttentionMetadata): + """Attention metadata extended with the static sparse Top-K.""" + + num_sparse_topk: int = 64 + + +class _StaticSparseMqaGqaAttention(TrtllmAttention): + """Backend adapter that replaces a model selector with static indices.""" + + def __init__( + self, + *args, + sparse_kv_indices: Optional[torch.Tensor] = None, + sparse_kv_offsets: Optional[torch.Tensor] = None, + sparse_attn_indices: Optional[torch.Tensor] = None, + sparse_attn_offsets: Optional[torch.Tensor] = None, + **kwargs, + ): + kwargs["sparse_params"] = _SparseMqaGqaParams() + kwargs["pos_embd_params"] = None + super().__init__(*args, **kwargs) + + self._sparse_kv_indices = sparse_kv_indices + self._sparse_kv_offsets = sparse_kv_offsets + self._sparse_attn_indices = sparse_attn_indices + self._sparse_attn_offsets = sparse_attn_offsets + + def sparse_kv_predict(self, q, k, metadata, forward_args: AttentionForwardArgs): + return self._sparse_kv_indices, self._sparse_kv_offsets + + def sparse_attn_predict(self, q, k, metadata, forward_args: AttentionForwardArgs): + return self._sparse_attn_indices, self._sparse_attn_offsets + + +@dataclass(kw_only=True) +class _ContextInputs: + """Packed context tensors plus their paged KV-cache metadata.""" + + q: torch.Tensor + k: torch.Tensor + v: torch.Tensor + kv_cache_manager: KVCacheManager + request_ids: List[int] + metadata: _SparseMqaGqaMetadata + + @property + def fused_qkv(self) -> torch.Tensor: + return torch.cat([self.q, self.k, self.v], dim=1) + + +@dataclass(kw_only=True) +class _GenerationInputs: + """Decode tensors, local token selections, and populated paged KV cache.""" + + q: torch.Tensor + k_new: torch.Tensor + v_new: torch.Tensor + local_sparse_attn_indices: torch.Tensor + kv_cache_manager: KVCacheManager + request_ids: List[int] + metadata: _SparseMqaGqaMetadata + + @property + def fused_qkv(self) -> torch.Tensor: + return torch.cat([self.q, self.k_new, self.v_new], dim=1) + + +# Sparse KV-cache feature tests. + + +_SPARSE_KV_CASES = [ + pytest.param( + ContextScenario(batch_size=2, seq_lens=(48, 64), num_pages=8), + id="batch2_var_seq", + ), + pytest.param( + ContextScenario(batch_size=4, seq_lens=(96, 112, 128, 144), num_pages=16), + id="batch4_var_seq", + ), + pytest.param( + ContextScenario(batch_size=1, seq_lens=(256,), num_pages=8), + id="batch1_seq256", + ), + pytest.param( + ContextScenario(batch_size=3, seq_lens=(64, 96, 128), num_pages=12), + id="batch3_var_seq", + ), +] + + +@pytest.mark.parametrize("scenario", _SPARSE_KV_CASES) +def test_prefill_sparse_kv_compaction(scenario: ContextScenario) -> None: + """Sparse KV selection compacts the cache without changing dense prefill output. + + This test calls ``attention.forward``, but supplies only + ``sparse_kv_indices``. Without ``sparse_attn_indices``, attention compute is + dense; the sparse feature under test is the selected K/V write into the + paged cache. + """ + inputs = _create_context_inputs(scenario) + local_sparse_kv_indices, sparse_kv_offsets = _make_context_kv_indices( + scenario.seq_lens, + scenario.num_kv_heads, + scenario.num_sparse_topk, + inputs.q.device, + ) + attention = _StaticSparseMqaGqaAttention( + layer_idx=0, + num_heads=scenario.num_heads, + head_dim=scenario.head_dim, + num_kv_heads=scenario.num_kv_heads, + quant_config=_quant_config(scenario), + sparse_kv_indices=local_sparse_kv_indices, + sparse_kv_offsets=sparse_kv_offsets, + ) + + try: + reference_output = _reference_dense_context_attention( + inputs.q, inputs.k, inputs.v, scenario + ) + expected_kvs = _build_expected_compacted_kv( + inputs.k, + inputs.v, + local_sparse_kv_indices, + sparse_kv_offsets, + scenario, + ) + + output = attention.forward(inputs.fused_qkv, None, None, inputs.metadata) + assert output.shape == reference_output.shape + torch.testing.assert_close(output, reference_output, atol=ATOL, rtol=RTOL) + + compacted_kv_lens = tuple( + int((sparse_kv_offsets[i + 1] - sparse_kv_offsets[i]).item()) + for i in range(scenario.batch_size) + ) + actual_kvs = _read_paged_kv_cache( + inputs.kv_cache_manager, + inputs.request_ids, + compacted_kv_lens, + scenario, + scenario.dtype, + ) + for batch_idx, ((actual_k, actual_v), (expected_k, expected_v)) in enumerate( + zip(actual_kvs, expected_kvs, strict=True) + ): + torch.testing.assert_close( + actual_k, + expected_k, + atol=ATOL, + rtol=RTOL, + msg=f"K cache mismatch for batch {batch_idx} after sparse compaction", + ) + torch.testing.assert_close( + actual_v, + expected_v, + atol=ATOL, + rtol=RTOL, + msg=f"V cache mismatch for batch {batch_idx} after sparse compaction", + ) + finally: + inputs.kv_cache_manager.shutdown() + + +# Sparse MQA/GQA computation tests. +# +# Sparse MQA/GQA support matrix: +# +# GPU architecture SM100 and SM103 +# Inference phase Packed prefill; single-token and linear draft decode +# Attention type MQA and GQA; num_heads % num_kv_heads == 0 +# Q heads per KV head <= 32; tests cover 2, 3, 4, 8, 16, 24, 31, and 32 +# Model QKV input BF16 or FP16 +# Kernel output BF16, FP16, or E4M3 FP8 +# KV-cache dtype BF16, FP16, or E4M3 FP8 +# Q/K/V head dimension Equal dimensions: 64, 80, 128, or 256 +# KV-cache layout Paged; page size is a power of two and at least 8 +# Sparse indices int32, token-granular, one list per KV head/query +# Sparse Top-K Positive multiple of 4; unused entries are -1 + + +_PREFILL_COMPUTE_CASES = [ + pytest.param( + ContextScenario( + batch_size=2, + seq_lens=(128, 64), + num_pages=8, + num_kv_heads=1, + num_heads=8, + ), + id="mqa_8q1kv", + ), + pytest.param( + ContextScenario( + batch_size=2, + seq_lens=(128, 64), + num_pages=8, + num_kv_heads=2, + num_heads=8, + ), + id="gqa_8q2kv", + ), + pytest.param( + ContextScenario( + batch_size=3, + seq_lens=(64, 96, 128), + num_pages=12, + num_kv_heads=4, + num_heads=16, + ), + id="gqa_16q4kv_batch3", + ), + pytest.param( + ContextScenario( + batch_size=2, + seq_lens=(64, 128), + num_pages=8, + num_kv_heads=4, + num_heads=32, + ), + id="gqa_32q4kv", + ), + pytest.param( + ContextScenario( + batch_size=2, + seq_lens=(64, 128), + num_pages=8, + num_kv_heads=1, + num_heads=8, + num_sparse_topk=4, + ), + id="topk4_very_sparse", + ), + pytest.param( + ContextScenario( + batch_size=2, + seq_lens=(64, 128), + num_pages=8, + num_kv_heads=2, + num_heads=8, + num_sparse_topk=128, + ), + id="topk128_near_dense", + ), +] + + +_GENERATION_CORRECTNESS_CASES = [ + pytest.param( + GenerationScenario(batch_size=2, past_kv_lens=(96, 128), num_pages=16), + id="batch2_var_kv", + ), + pytest.param( + GenerationScenario( + batch_size=4, + past_kv_lens=(192, 224, 256, 288), + num_pages=32, + ), + id="batch4_var_kv", + ), + pytest.param( + GenerationScenario(batch_size=1, past_kv_lens=(64,), num_pages=8), + id="batch1_kv64", + ), + pytest.param( + GenerationScenario( + batch_size=3, + past_kv_lens=(128, 160, 192), + num_pages=24, + ), + id="batch3_var_kv", + ), + pytest.param( + GenerationScenario( + num_heads=8, + num_kv_heads=1, + batch_size=2, + past_kv_lens=(96, 128), + num_pages=16, + ), + id="mqa_8q1kv", + ), + pytest.param( + GenerationScenario( + num_heads=16, + num_kv_heads=4, + batch_size=2, + past_kv_lens=(128, 256), + num_pages=16, + ), + id="gqa_16q4kv", + ), + pytest.param( + GenerationScenario( + num_heads=8, + num_kv_heads=4, + batch_size=2, + past_kv_lens=(128, 256), + num_pages=16, + ), + id="gqa_8q4kv", + ), + pytest.param( + GenerationScenario( + num_heads=32, + num_kv_heads=1, + batch_size=1, + past_kv_lens=(128,), + num_pages=8, + ), + id="mqa_group32_boundary", + ), + pytest.param( + GenerationScenario( + num_heads=64, + num_kv_heads=2, + batch_size=1, + past_kv_lens=(128,), + num_pages=8, + ), + id="gqa_group32_boundary", + ), + pytest.param( + GenerationScenario( + num_heads=6, + num_kv_heads=2, + batch_size=1, + past_kv_lens=(64,), + num_pages=8, + ), + id="gqa_group3_non_power_of_two", + ), + pytest.param( + GenerationScenario( + num_heads=32, + num_kv_heads=2, + batch_size=1, + past_kv_lens=(64,), + num_pages=8, + ), + id="gqa_group16", + ), + pytest.param( + GenerationScenario( + num_heads=48, + num_kv_heads=2, + batch_size=1, + past_kv_lens=(64,), + num_pages=8, + ), + id="gqa_group24", + ), + pytest.param( + GenerationScenario( + num_heads=62, + num_kv_heads=2, + batch_size=1, + past_kv_lens=(64,), + num_pages=8, + ), + id="gqa_group31", + ), + pytest.param( + GenerationScenario( + num_heads=8, + num_kv_heads=4, + batch_size=1, + past_kv_lens=(64,), + query_len=4, + num_pages=8, + num_sparse_topk=32, + ), + id="gqa_2to1_with_3_draft_tokens", + ), + pytest.param( + GenerationScenario( + batch_size=1, + past_kv_lens=(128,), + num_pages=8, + num_sparse_topk=4, + ), + id="topk4_min", + ), + pytest.param( + GenerationScenario( + batch_size=2, + past_kv_lens=(32, 256), + num_pages=16, + num_sparse_topk=128, + ), + id="topk128_exceeds_some", + ), + pytest.param( + GenerationScenario( + batch_size=8, + past_kv_lens=(64, 96, 128, 160, 192, 224, 256, 288), + num_pages=64, + ), + id="batch8_varied", + ), + pytest.param( + GenerationScenario( + page_size=64, + batch_size=2, + past_kv_lens=(64, 192), + num_pages=8, + ), + id="page_size_64", + ), +] + + +_GENERATION_SUPPORT_CASES = [ + pytest.param( + GenerationScenario( + dtype=dtype, + kvcache_dtype=dtype, + num_heads=8, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + batch_size=1, + past_kv_lens=(64,), + num_pages=4, + num_sparse_topk=32, + ), + id=( + f"support_{str(dtype).removeprefix('torch.')}_h{head_dim}_" + f"{'mqa' if num_kv_heads == 1 else 'gqa_2to1'}" + ), + ) + for dtype in SUPPORTED_DTYPES + for head_dim in SUPPORTED_HEAD_DIMS + for num_kv_heads in (1, 4) + # These two option combinations are already covered by correctness cases. + if not (dtype == torch.bfloat16 and head_dim == 128) +] + [ + pytest.param( + GenerationScenario( + dtype=torch.bfloat16, + kvcache_dtype=torch.float8_e4m3fn, + num_heads=8, + num_kv_heads=4, + head_dim=128, + batch_size=1, + past_kv_lens=(64,), + num_pages=8, + num_sparse_topk=32, + ), + id="support_bf16_io_fp8_kv_cache", + ), + pytest.param( + GenerationScenario( + dtype=torch.bfloat16, + kvcache_dtype=torch.float8_e4m3fn, + num_heads=8, + num_kv_heads=4, + head_dim=128, + batch_size=1, + past_kv_lens=(64,), + num_pages=8, + num_sparse_topk=32, + fp8_output=True, + ), + id="support_fp8_qkv_math_and_output", + ), +] + + +@pytest.mark.parametrize("scenario", _PREFILL_COMPUTE_CASES) +def test_prefill_sparse_mqa_gqa(scenario: ContextScenario) -> None: + """Prefill sparse indices drive sparse compute and compacted KV writes.""" + inputs = _create_context_inputs(scenario) + try: + available_kv_lens = tuple( + token_idx + 1 for seq_len in scenario.seq_lens for token_idx in range(seq_len) + ) + local_attn_indices = _make_sparse_attention_indices( + available_kv_lens, + scenario.num_kv_heads, + scenario.num_sparse_topk, + inputs.q.device, + ) + cache_pool_indices = _local_to_cache_pool_indices( + local_attn_indices, inputs.metadata, layer_idx=0 + ) + local_sparse_kv_indices, sparse_kv_offsets = _make_context_kv_indices( + scenario.seq_lens, + scenario.num_kv_heads, + scenario.num_sparse_topk, + inputs.q.device, + ) + + attention = _StaticSparseMqaGqaAttention( + layer_idx=0, + num_heads=scenario.num_heads, + head_dim=scenario.head_dim, + num_kv_heads=scenario.num_kv_heads, + quant_config=_quant_config(scenario), + sparse_kv_indices=local_sparse_kv_indices, + sparse_kv_offsets=sparse_kv_offsets, + sparse_attn_indices=cache_pool_indices, + ) + reference_output = _reference_sparse_context_attention( + inputs.q, inputs.k, inputs.v, local_attn_indices, scenario + ) + + output = attention.forward(inputs.fused_qkv, None, None, inputs.metadata) + expected_shape = (scenario.nnz_q, scenario.num_heads * scenario.head_dim) + assert output.shape == expected_shape + assert torch.isfinite(output).all() + torch.testing.assert_close(output, reference_output, atol=ATOL, rtol=RTOL) + finally: + inputs.kv_cache_manager.shutdown() + + +@pytest.mark.parametrize( + "scenario", + _GENERATION_CORRECTNESS_CASES + _GENERATION_SUPPORT_CASES, +) +def test_generation_sparse_mqa_gqa(scenario: GenerationScenario) -> None: + """Decode static token selections match the paged-cache PyTorch reference.""" + inputs = _create_generation_inputs(scenario) + try: + cache_pool_indices = _local_to_cache_pool_indices( + inputs.local_sparse_attn_indices, inputs.metadata, layer_idx=0 + ) + attention = _StaticSparseMqaGqaAttention( + layer_idx=0, + num_heads=scenario.num_heads, + head_dim=scenario.head_dim, + num_kv_heads=scenario.num_kv_heads, + quant_config=_quant_config(scenario), + sparse_attn_indices=cache_pool_indices, + ) + + kv_caches = _read_paged_kv_cache( + inputs.kv_cache_manager, + inputs.request_ids, + scenario.past_kv_lens, + scenario, + scenario.dtype, + ) + reference_output = _reference_sparse_generation_attention( + inputs.q, + kv_caches, + inputs.k_new, + inputs.v_new, + inputs.local_sparse_attn_indices, + scenario, + ) + + forward_args = None + if scenario.fp8_output: + forward_args = AttentionForwardArgs( + out_scale=torch.ones(1, dtype=torch.float32, device=inputs.q.device) + ) + output = attention.forward( + inputs.fused_qkv, + None, + None, + inputs.metadata, + forward_args=forward_args, + ) + expected_shape = (scenario.nnz_q, scenario.num_heads * scenario.head_dim) + assert output.shape == expected_shape + uses_fp8 = scenario.kvcache_dtype == torch.float8_e4m3fn or scenario.fp8_output + output_for_comparison = output.float() if uses_fp8 else output + reference_for_comparison = reference_output.float() if uses_fp8 else reference_output + assert torch.isfinite(output_for_comparison).all() + if scenario.fp8_output: + assert output.dtype == torch.float8_e4m3fn + torch.testing.assert_close( + output_for_comparison, + reference_for_comparison, + atol=FP8_ATOL if uses_fp8 else ATOL, + rtol=FP8_RTOL if uses_fp8 else RTOL, + ) + finally: + inputs.kv_cache_manager.shutdown() + + +# Sparse index and paged KV-cache helpers. + + +def _quant_config(s: SparseMqaGqaScenario) -> Optional[QuantConfig]: + if isinstance(s, GenerationScenario) and s.fp8_output: + return QuantConfig( + quant_algo=QuantAlgo.FP8, + kv_cache_quant_algo=QuantAlgo.FP8, + ) + if s.kvcache_dtype == torch.float8_e4m3fn: + return QuantConfig(kv_cache_quant_algo=QuantAlgo.FP8) + return None + + +def _create_kv_cache_manager( + s: SparseMqaGqaScenario, kv_cache: Optional[torch.Tensor] = None +) -> KVCacheManager: + """Create kv cache manager for testing.""" + kv_cache_config = KvCacheConfig(max_tokens=s.kv_pool_num_pages * s.page_size) + mapping = Mapping(world_size=1, tp_size=1, rank=0) + + manager = KVCacheManager( + kv_cache_config, + tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF, + num_layers=s.num_layers, + num_kv_heads=s.num_kv_heads, + head_dim=s.head_dim, + tokens_per_block=s.page_size, + max_seq_len=s.kv_pool_num_pages * s.page_size, + max_batch_size=s.batch_size, + mapping=mapping, + dtype=str_dtype_to_binding(torch_dtype_to_str(s.kvcache_dtype)), + ) + + if kv_cache is not None: + for i in range(s.num_layers): + manager.get_buffers(i, kv_layout="HND").copy_(kv_cache[i]) + + return manager + + +def _make_context_kv_indices( + seq_lens: Tuple[int, ...], + num_kv_heads: int, + num_sparse_topk: int, + device: torch.device, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Generate sparse kv indices for context phase. + + For each request, pick min(num_sparse_topk, seq_len) indices from [0, seq_len). + Returns (indices [num_kv_heads, total_sparse], offsets [num_requests + 1]). + """ + all_indices = [] + offsets = [0] + + for seq_len in seq_lens: + pick = min(num_sparse_topk, seq_len) + batch_indices = [] + for _ in range(num_kv_heads): + indices = torch.randperm(seq_len, device=device)[:pick].sort().values + batch_indices.append(indices) + all_indices.append(torch.stack(batch_indices, dim=0)) + offsets.append(offsets[-1] + pick) + + indices = torch.cat(all_indices, dim=1).int() + offsets = torch.tensor(offsets, dtype=torch.int32, device=device) + return indices, offsets + + +def _make_sparse_attention_indices( + available_kv_lens: Tuple[int, ...], + num_kv_heads: int, + num_sparse_topk: int, + device: torch.device, +) -> torch.Tensor: + """Create one padded request-local token list per query token and KV head.""" + num_query_tokens = len(available_kv_lens) + result = torch.full( + (num_kv_heads, num_query_tokens, num_sparse_topk), + -1, + dtype=torch.int32, + device=device, + ) + + for query_idx, available_kv_len in enumerate(available_kv_lens): + pick = min(num_sparse_topk, available_kv_len) + for head_idx in range(num_kv_heads): + indices = torch.randperm(available_kv_len, device=device)[:pick].sort().values + result[head_idx, query_idx, :pick] = indices + + return result + + +def _local_to_cache_pool_indices( + sparse_indices: torch.Tensor, + metadata: TrtllmAttentionMetadata, + layer_idx: int = 0, + kv_factor: int = 2, +) -> torch.Tensor: + """ + Convert local sparse indices to global KV cache pool indices. + + Works for both context (variable-length Q packed) and generation (one token per request). + sparse_indices shape: [num_kv_heads, num_tokens, num_sparse_topk] + """ + num_kv_heads, num_tokens, num_sparse_tokens = sparse_indices.shape + device = sparse_indices.device + + tokens_per_block = metadata.kv_cache_manager.tokens_per_block + num_layers = metadata.kv_cache_manager.num_layers + stride_factor = num_layers * kv_factor * num_kv_heads * tokens_per_block + + # Build req_idx_per_token: map each token to its request index. + num_requests = len(metadata.request_ids) + seq_lens = metadata.seq_lens[:num_requests] + if hasattr(seq_lens, "cpu"): + seq_lens_cpu = seq_lens.cpu() + else: + seq_lens_cpu = torch.tensor(seq_lens, dtype=torch.int32) + req_idx_per_token = torch.repeat_interleave( + torch.arange(num_requests, dtype=torch.int32), seq_lens_cpu, dim=0 + ).to(device) + + # Build 2D block table: [num_requests, max_pages] + request_ids = metadata.request_ids + page_indices = metadata.kv_cache_manager.get_batch_cache_indices(request_ids) + max_pages = max(len(p) for p in page_indices) if page_indices else 1 + host_block_table = torch.full((num_requests, max_pages), -1, dtype=torch.int32) + for i, pages in enumerate(page_indices): + if len(pages) > 0: + host_block_table[i, : len(pages)] = torch.tensor(pages, dtype=torch.int32) + block_table = host_block_table.to(device) + + # Convert to global + global_indices = triton_convert_req_index_to_global_index( + req_idx_per_token, + block_table, + sparse_indices, + BLOCK_SIZE=tokens_per_block, + NUM_TOPK_TOKENS=num_sparse_tokens, + BLOCK_N=min(64, num_sparse_tokens), + stride_factor=stride_factor, + layer_id=layer_idx, + num_kv_heads=num_kv_heads, + kv_factor=kv_factor, + ) + + return global_indices + + +def _build_expected_compacted_kv( + k: torch.Tensor, + v: torch.Tensor, + sparse_kv_indices: torch.Tensor, + sparse_kv_offsets: torch.Tensor, + s: ContextScenario, +) -> List[Tuple[torch.Tensor, torch.Tensor]]: + """Build expected sparse K and V values based on sparse indices.""" + expected_kvs = [] + token_offset = 0 + + for batch_idx, seq_len in enumerate(s.seq_lens): + sparse_len = sparse_kv_offsets[batch_idx + 1].item() - sparse_kv_offsets[batch_idx].item() + k_batch = k[token_offset : token_offset + seq_len].view(seq_len, s.num_kv_heads, s.head_dim) + v_batch = v[token_offset : token_offset + seq_len].view(seq_len, s.num_kv_heads, s.head_dim) + + expected_k = torch.zeros( + sparse_len, s.num_kv_heads, s.head_dim, device=k.device, dtype=k.dtype + ) + expected_v = torch.zeros_like(expected_k) + + start, end = sparse_kv_offsets[batch_idx].item(), sparse_kv_offsets[batch_idx + 1].item() + for head_idx in range(s.num_kv_heads): + indices = sparse_kv_indices[head_idx, start:end] + expected_k[:, head_idx] = k_batch[indices, head_idx] + expected_v[:, head_idx] = v_batch[indices, head_idx] + + expected_kvs.append((expected_k, expected_v)) + token_offset += seq_len + + return expected_kvs + + +def _read_paged_kv_cache( + kv_cache_manager: KVCacheManager, + request_ids: List[int], + token_lens: Tuple[int, ...], + s: SparseMqaGqaScenario, + dtype: torch.dtype, +) -> List[Tuple[torch.Tensor, torch.Tensor]]: + """Materialize each request's paged K/V history as contiguous tensors.""" + kv_buffer = kv_cache_manager.get_buffers(0, kv_layout="HND") + kv_caches = [] + + for request_id, num_tokens in zip(request_ids, token_lens, strict=True): + block_ids = kv_cache_manager.get_block_ids_per_seq([request_id])[0] + k_cache = torch.empty( + num_tokens, + s.num_kv_heads, + s.head_dim, + device=kv_buffer.device, + dtype=dtype, + ) + v_cache = torch.empty_like(k_cache) + for token_idx in range(num_tokens): + block_id = block_ids[token_idx // s.page_size] + offset_in_block = token_idx % s.page_size + k_cache[token_idx] = kv_buffer[block_id, 0, :, offset_in_block, :].to(dtype) + v_cache[token_idx] = kv_buffer[block_id, 1, :, offset_in_block, :].to(dtype) + + kv_caches.append((k_cache, v_cache)) + + return kv_caches + + +# Independent PyTorch reference implementations. + + +def _reference_dense_context_attention( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + s: ContextScenario, +) -> torch.Tensor: + """Reference implementation for context phase.""" + outputs = [] + token_offset = 0 + + for seq_len in s.seq_lens: + q_batch = q[token_offset : token_offset + seq_len].view(1, seq_len, s.num_heads, s.head_dim) + k_batch = k[token_offset : token_offset + seq_len].view( + 1, seq_len, s.num_kv_heads, s.head_dim + ) + v_batch = v[token_offset : token_offset + seq_len].view( + 1, seq_len, s.num_kv_heads, s.head_dim + ) + q_batch = q_batch.transpose(1, 2) + k_batch = k_batch.transpose(1, 2) + v_batch = v_batch.transpose(1, 2) + if s.q_heads_per_kv_head > 1: + k_batch = k_batch[:, :, None, :, :].expand( + 1, s.num_kv_heads, s.q_heads_per_kv_head, seq_len, s.head_dim + ) + v_batch = v_batch[:, :, None, :, :].expand( + 1, s.num_kv_heads, s.q_heads_per_kv_head, seq_len, s.head_dim + ) + k_batch = k_batch.reshape(1, s.num_heads, seq_len, s.head_dim) + v_batch = v_batch.reshape(1, s.num_heads, seq_len, s.head_dim) + + attention_scores = torch.matmul(q_batch, k_batch.transpose(-1, -2)) / math.sqrt(s.head_dim) + causal_mask = torch.triu( + torch.full((seq_len, seq_len), float("-inf"), device=q.device), diagonal=1 + ) + attention_probs = torch.nn.functional.softmax( + attention_scores + causal_mask, + dim=-1, + dtype=torch.float32, + ).to(q.dtype) + output_batch = torch.matmul(attention_probs, v_batch) + output_batch = output_batch.transpose(1, 2).reshape(seq_len, s.num_heads * s.head_dim) + outputs.append(output_batch) + token_offset += seq_len + + return torch.cat(outputs, dim=0) + + +def _reference_sparse_context_attention( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + sparse_attn_ctx_indices: torch.Tensor, + s: ContextScenario, +) -> torch.Tensor: + """ + Reference implementation for context phase with sparse attention. + Uses mask-based approach for each KV head. + """ + total_tokens = sum(s.seq_lens) + device = q.device + dtype = q.dtype + + # Reshape inputs: [num_tokens, num_heads, head_dim] + q_reshaped = q.view(total_tokens, s.num_heads, s.head_dim) + k_reshaped = k.view(total_tokens, s.num_kv_heads, s.head_dim) + v_reshaped = v.view(total_tokens, s.num_kv_heads, s.head_dim) + + outputs = [] + token_offset = 0 + + for seq_len in s.seq_lens: + q_batch = q_reshaped[ + token_offset : token_offset + seq_len + ] # [seq_len, num_heads, head_dim] + k_batch = k_reshaped[ + token_offset : token_offset + seq_len + ] # [seq_len, num_kv_heads, head_dim] + v_batch = v_reshaped[ + token_offset : token_offset + seq_len + ] # [seq_len, num_kv_heads, head_dim] + + batch_output = [] + + # Process each KV head + for kv_head_idx in range(s.num_kv_heads): + k_head = k_batch[:, kv_head_idx, :] + v_head = v_batch[:, kv_head_idx, :] + + # Build sparse mask for this head + sparse_mask = torch.full( + (seq_len, seq_len), float("-inf"), device=device, dtype=torch.float32 + ) + + for token_idx in range(seq_len): + global_token_idx = token_offset + token_idx + # Get sparse indices for this token: [num_sparse_tokens] + indices = sparse_attn_ctx_indices[kv_head_idx, global_token_idx] + # Filter out -1 padding + valid_indices = indices[indices >= 0] + # Set mask values to 0 for valid positions + sparse_mask[token_idx, valid_indices] = 0.0 + + # Apply causal mask on top of sparse mask + causal_mask = torch.triu( + torch.full((seq_len, seq_len), float("-inf"), device=device, dtype=torch.float32), + diagonal=1, + ) + combined_mask = sparse_mask + causal_mask + + # Process each query head in this KV group + for group_idx in range(s.q_heads_per_kv_head): + q_head_idx = kv_head_idx * s.q_heads_per_kv_head + group_idx + q_head = q_batch[:, q_head_idx, :] # [seq_len, head_dim] + + attn_scores = torch.matmul(q_head, k_head.T) / math.sqrt(s.head_dim) + attn_scores = attn_scores + combined_mask + attn_weights = torch.nn.functional.softmax( + attn_scores, dim=-1, dtype=torch.float32 + ).to(dtype) + + out_head = torch.matmul(attn_weights, v_head) + batch_output.append(out_head) + + # Concatenate all heads: [seq_len, num_heads, head_dim] -> [seq_len, num_heads * head_dim] + batch_output = torch.stack(batch_output, dim=1) + batch_output = batch_output.reshape(seq_len, s.num_heads * s.head_dim) + outputs.append(batch_output) + + token_offset += seq_len + + return torch.cat(outputs, dim=0) + + +def _reference_sparse_generation_attention( + q: torch.Tensor, + kv_caches: List[Tuple[torch.Tensor, torch.Tensor]], + k_new: torch.Tensor, + v_new: torch.Tensor, + sparse_attn_indices: torch.Tensor, + s: GenerationScenario, +) -> torch.Tensor: + """Reference implementation for generation phase with sparse attention. + + Args: + sparse_attn_indices: [num_kv_heads, num_gens, num_sparse_topk] with -1 padding. + """ + outputs = [] + query_offset = 0 + + for request_idx in range(s.batch_size): + query_len = s.query_len + k_history, v_history = kv_caches[request_idx] + k_new_request = k_new[query_offset : query_offset + query_len].view( + query_len, s.num_kv_heads, s.head_dim + ) + v_new_request = v_new[query_offset : query_offset + query_len].view( + query_len, s.num_kv_heads, s.head_dim + ) + k_full = torch.cat([k_history, k_new_request], dim=0) + v_full = torch.cat([v_history, v_new_request], dim=0) + + for query_idx in range(query_len): + packed_query_idx = query_offset + query_idx + q_token = q[packed_query_idx].view(s.num_heads, s.head_dim) + head_outputs = [] + + for kv_head_idx in range(s.num_kv_heads): + token_indices = sparse_attn_indices[kv_head_idx, packed_query_idx] + valid_indices = token_indices[token_indices >= 0].long() + + if len(valid_indices) == 0: + head_outputs.extend( + [torch.zeros(s.head_dim, device=q.device, dtype=q.dtype)] + * s.q_heads_per_kv_head + ) + continue + + k_sparse = k_full[valid_indices, kv_head_idx, :] + v_sparse = v_full[valid_indices, kv_head_idx, :] + + for group_idx in range(s.q_heads_per_kv_head): + q_head_idx = kv_head_idx * s.q_heads_per_kv_head + group_idx + attention_scores = torch.matmul(q_token[q_head_idx], k_sparse.T) / math.sqrt( + s.head_dim + ) + attention_probs = torch.nn.functional.softmax( + attention_scores, dim=-1, dtype=torch.float32 + ).to(q.dtype) + head_outputs.append(torch.matmul(attention_probs, v_sparse)) + + outputs.append(torch.cat(head_outputs, dim=0)) + + query_offset += query_len + + return torch.stack(outputs, dim=0) + + +# Test input builders. Kernel selection remains explicit in each test below. + + +def _create_context_inputs(s: ContextScenario) -> _ContextInputs: + """Create packed context inputs before choosing the sparse compute path.""" + device = torch.device("cuda") + torch.manual_seed(42) + num_sparse_topk = s.num_sparse_topk + + q = torch.randn(s.nnz_q, s.num_heads * s.head_dim, device=device, dtype=s.dtype) + k = torch.randn(s.nnz_q, s.num_kv_heads * s.head_dim, device=device, dtype=s.dtype) + v = torch.randn(s.nnz_q, s.num_kv_heads * s.head_dim, device=device, dtype=s.dtype) + + kv_cache = torch.zeros( + s.num_layers, + s.kv_pool_num_pages, + 2, + s.num_kv_heads, + s.page_size, + s.head_dim, + device=device, + dtype=s.kvcache_dtype, + ) + kv_cache_manager = _create_kv_cache_manager(s, kv_cache) + + request_ids = list(range(s.batch_size)) + kv_cache_manager.add_dummy_requests(request_ids, list(s.seq_lens)) + + metadata = _SparseMqaGqaMetadata( + num_contexts=s.batch_size, + kv_cache_params=KVCacheParams(use_cache=True, num_cached_tokens_per_seq=[0] * s.batch_size), + seq_lens=torch.tensor(s.seq_lens, dtype=torch.int32), + max_num_requests=s.batch_size, + max_num_tokens=s.nnz_q, + kv_cache_manager=kv_cache_manager, + request_ids=request_ids, + prompt_lens=list(s.seq_lens), + num_sparse_topk=num_sparse_topk, + ) + metadata.prepare() + + return _ContextInputs( + q=q, + k=k, + v=v, + kv_cache_manager=kv_cache_manager, + request_ids=request_ids, + metadata=metadata, + ) + + +def _create_generation_inputs(s: GenerationScenario) -> _GenerationInputs: + """Create one decode token and a populated paged cache per request.""" + device = torch.device("cuda") + torch.manual_seed(42) + num_sparse_topk = s.num_sparse_topk + + token_nums = [past_len + s.query_len for past_len in s.past_kv_lens] + + q = torch.randn(s.nnz_q, s.num_heads * s.head_dim, device=device, dtype=s.dtype) + k_new = torch.randn(s.nnz_q, s.num_kv_heads * s.head_dim, device=device, dtype=s.dtype) + v_new = torch.randn(s.nnz_q, s.num_kv_heads * s.head_dim, device=device, dtype=s.dtype) + + # Single-token cases preserve the original history-only selection. For + # draft-token cases, each query may also select causal K/V written earlier + # in the same speculative forward, including its own K/V position. + available_kv_lens = tuple( + past_kv_len + query_idx + 1 if s.has_draft_tokens else past_kv_len + for past_kv_len in s.past_kv_lens + for query_idx in range(s.query_len) + ) + sparse_attn_indices = _make_sparse_attention_indices( + available_kv_lens, s.num_kv_heads, num_sparse_topk, device + ) + + kv_cache = torch.randn( + s.num_layers, + s.kv_pool_num_pages, + 2, + s.num_kv_heads, + s.page_size, + s.head_dim, + device=device, + dtype=s.dtype, + ).to(s.kvcache_dtype) + kv_cache_manager = _create_kv_cache_manager(s, kv_cache) + + request_ids = list(range(s.batch_size)) + kv_cache_manager.add_dummy_requests(request_ids, token_nums) + + metadata = _SparseMqaGqaMetadata( + num_contexts=0, + kv_cache_params=KVCacheParams( + use_cache=True, num_cached_tokens_per_seq=list(s.past_kv_lens) + ), + seq_lens=torch.full((s.batch_size,), s.query_len, dtype=torch.int32), + max_num_requests=s.batch_size, + max_num_tokens=s.nnz_q, + kv_cache_manager=kv_cache_manager, + request_ids=request_ids, + prompt_lens=list(s.past_kv_lens), + num_sparse_topk=num_sparse_topk, + num_heads_per_kv=s.q_heads_per_kv_head, + runtime_features=AttentionRuntimeFeatures(has_speculative_draft_tokens=s.has_draft_tokens), + is_spec_decoding_enabled=s.has_draft_tokens, + use_spec_decoding=s.has_draft_tokens, + is_spec_dec_tree=False, + max_total_draft_tokens=s.max_query_len - 1 if s.has_draft_tokens else None, + ) + if s.has_draft_tokens: + draft_len = s.max_query_len - 1 + metadata.spec_decoding_position_offsets = generate_spec_decoding_position_offsets( + s.batch_size, draft_len + ) + metadata.spec_decoding_packed_mask = generate_spec_decoding_packed_mask( + s.batch_size, draft_len + ) + metadata.spec_decoding_generation_lengths = torch.tensor( + [s.query_len] * s.batch_size, dtype=torch.int32, device=device + ) + metadata.update_position_offsets_for_cpp(s.max_query_len) + metadata.spec_decoding_param_prepare_for_blackwell() + metadata.prepare() + + return _GenerationInputs( + q=q, + k_new=k_new, + v_new=v_new, + local_sparse_attn_indices=sparse_attn_indices, + kv_cache_manager=kv_cache_manager, + request_ids=request_ids, + metadata=metadata, + ) From c0b439ddcc0f222821b1596cb1015da23480bada Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Wed, 26 Aug 2026 21:59:47 -0700 Subject: [PATCH 03/12] [None][test] Add sparse MHA coverage Add algorithm-neutral page-sparse MHA regression coverage and document the runtime-verified support matrix. Colocate paged-MQA and FP4 indexer tests with the DSA implementation, deduplicate the FP4 indexer suite, and keep RocketKV tests focused on algorithm-specific behavior. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- docs/source/features/sparse-attention.md | 45 +- .../test_lists/test-db/l0_b300.yml | 8 +- .../test_lists/test-db/l0_dgx_b300.yml | 8 +- .../paged_mqa_logits/run_fp4.py | 4 +- .../paged_mqa_logits/run_fp8.py | 4 +- .../sparse/_sparse_mha_test_utils.py | 328 ++++++++++++ .../test_cute_dsl_fp4_paged_mqa_logits.py | 0 .../test_cute_dsl_fp8_paged_mqa_logits.py | 0 .../attention/sparse/test_dsa_fp4_indexer.py | 426 ---------------- .../attention/sparse/test_sparse_mha.py | 473 ++++++++++++++++++ 10 files changed, 853 insertions(+), 443 deletions(-) create mode 100644 tests/unittest/_torch/attention/sparse/_sparse_mha_test_utils.py rename tests/unittest/_torch/attention/sparse/{ => dsa}/test_cute_dsl_fp4_paged_mqa_logits.py (100%) rename tests/unittest/_torch/attention/sparse/{ => dsa}/test_cute_dsl_fp8_paged_mqa_logits.py (100%) delete mode 100644 tests/unittest/_torch/attention/sparse/test_dsa_fp4_indexer.py create mode 100644 tests/unittest/_torch/attention/sparse/test_sparse_mha.py diff --git a/docs/source/features/sparse-attention.md b/docs/source/features/sparse-attention.md index 0730d926be7a..6081cd6c68fa 100644 --- a/docs/source/features/sparse-attention.md +++ b/docs/source/features/sparse-attention.md @@ -2,6 +2,7 @@ - [Overview](#overview) - [Supported Algorithms](#supported-algorithms) +- [Sparse MHA Kernel Support](#sparse-mha-kernel-support) - [Sparse MQA/GQA Kernel Support](#sparse-mqagqa-kernel-support) - [Configure Sparse Attention](#configure-sparse-attention) - [Algorithm Details](#algorithm-details) @@ -32,7 +33,7 @@ model-specific guide says they are tunable. | `algorithm` | Config class | Sparse mechanism | Attention implementation | Typical use | |---|---|---|---|---| -| `rocket` | `RocketSparseAttentionConfig` | Prompt KV eviction, then page-level Top-K selection during decode | TRTLLM or Vanilla | Training-free sparsity for standard attention models | +| `rocket` | `RocketSparseAttentionConfig` | Prompt KV eviction, then page-level Top-K selection during decode | TRTLLM or Vanilla | Training-free sparsity for MHA/MQA/GQA models | | `dsa` | `DeepSeekSparseAttentionConfig` | Learned token-level indexer followed by sparse MLA | TRTLLM | DeepSeek V3.2 and compatible model-native DSA architectures | | `deepseek_v4` | `DeepSeekV4SparseAttentionConfig` | Sliding-window attention plus compressed sparse or compressed dense history | TRTLLM | DeepSeek-V4 hybrid attention | | `minimax_m3` | `MiniMaxM3SparseAttentionConfig` | Learned block selection followed by sparse GQA | Dedicated Triton or MSA implementation | MiniMax-M3 sparse layers | @@ -56,6 +57,40 @@ execution backend. densely. RocketKV selects which prompt KV entries to retain, so it reduces cache size and later decode work. +## Sparse MHA Kernel Support + +RocketKV uses the shared page-sparse MHA path after its selector produces block +indices and per-request offsets. Prefill attention remains dense: RocketKV can +compact the retained KV cache after prefill, but sparse attention computation +starts during generation. + +### Support Matrix + +| Parameter | Support | +|---|---| +| GPU architecture | SM100 is runtime-tested; SM103 is source-supported and enabled by the tests | +| Sparse compute phase | Single-token and linear draft-token generation (`qSeqLen=4` is tested) | +| Attention type | MHA (`num_q_heads == num_kv_heads`) | +| Query heads per KV head | `1` | +| Number of MHA heads | No additional discrete source restriction beyond `num_q_heads == num_kv_heads > 0`; tests cover `1`, `2`, `3`, `4`, `8`, `16`, `24`, `32`, `48`, `64`, `96`, and `128` | +| Model Q/K/V input dtype | BF16 or FP16 | +| Model Q/K/V input layout | Fused QKV | +| Output dtype | Model dtype for head dimensions `64`, `80`, `128`, and `256`; E4M3 FP8 for head dimensions `64`, `128`, and `256` with an FP8 KV cache | +| KV-cache dtype | Model dtype for head dimensions `64`, `80`, `128`, and `256`; E4M3 FP8 for head dimensions `64`, `128`, and `256` | +| Q/K/V dimensions | Equal head dimensions of `64`, `80`, `128`, or `256` | +| KV-cache layout | Paged KV cache; page sizes `8`, `16`, `32`, `64`, `128`, `256`, and `512` are tested | +| Selection granularity | Block indices expanded to KV-cache pages | +| Sparse indices | `int32` block indices with `int32` per-request offsets; per-head patterns, unordered indices, and variable request offsets are tested | +| Sparse index block size | Blocks may cross KV-page boundaries; sizes `1`, `2`, `3`, `4`, `5`, `8`, `16`, `24`, `32`, and `48` are tested | +| Attention semantics | Causal self-attention | + +Backend developers can use +[`test_sparse_mha.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mha.py) +as an architecture-level integration example. It supplies static page +selections, invokes `TrtllmAttention.forward`, and compares the result with an +equivalent token-level PyTorch reference. RocketKV selector, metadata, and KT +cache tests remain under the `rocketkv/` subdirectory. + ## Sparse MQA/GQA Kernel Support TensorRT LLM contains an internal TRTLLM-Gen kernel for token-sparse @@ -161,10 +196,10 @@ The following sections list algorithm-specific settings and constraints. ### RocketKV [RocketKV](https://arxiv.org/pdf/2502.14051) is a training-free, two-stage -algorithm for standard attention architectures. During prefill, it computes -dense attention and permanently evicts prompt KV entries beyond a prompt -budget. During decode, it scores retained pages and attends to the selected -Top-K pages. +algorithm for MHA, MQA, and GQA architectures. During prefill, it computes dense +attention and permanently evicts prompt KV entries beyond a prompt budget. +During decode, it scores retained pages and attends to the selected Top-K +pages. RocketKV currently requires CUDA compute capability 10.0 or newer. KV-cache block reuse and chunked prefill must be disabled, and disaggregated serving is diff --git a/tests/integration/test_lists/test-db/l0_b300.yml b/tests/integration/test_lists/test-db/l0_b300.yml index 9ee4059f91c6..a3f29606b8ba 100644 --- a/tests/integration/test_lists/test-db/l0_b300.yml +++ b/tests/integration/test_lists/test-db/l0_b300.yml @@ -23,10 +23,10 @@ l0_b300: # by absolute index, and they gate on >= 2 GPUs plus Blackwell. They also # rebind mpi4py's serializer at import time, which applies process-wide on # every stage that collects them, including where the tests themselves skip. - # The sparse/test_cute_dsl_* files ignored below keep their own entries (pre-existing). - - unittest/_torch/attention --ignore=unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py --ignore=unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py --ignore=unittest/_torch/attention/sparse/test_cute_dsl_gvr_topk_decode.py --ignore=unittest/_torch/attention/sparse/test_cute_dsl_gvr_topk_tiers.py - - unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py - - unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py + # The sparse/dsa/test_cute_dsl_* files ignored below keep their own entries (pre-existing). + - unittest/_torch/attention --ignore=unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp8_paged_mqa_logits.py --ignore=unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp4_paged_mqa_logits.py --ignore=unittest/_torch/attention/sparse/test_cute_dsl_gvr_topk_decode.py --ignore=unittest/_torch/attention/sparse/test_cute_dsl_gvr_topk_tiers.py + - unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp8_paged_mqa_logits.py + - unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp4_paged_mqa_logits.py - unittest/_torch/attention/sparse/test_cute_dsl_gvr_topk_decode.py TIMEOUT (120) - unittest/_torch/attention/sparse/test_cute_dsl_gvr_topk_tiers.py - unittest/_torch/modeling/test_qsa_runtime_wiring.py diff --git a/tests/integration/test_lists/test-db/l0_dgx_b300.yml b/tests/integration/test_lists/test-db/l0_dgx_b300.yml index 6b61ff6506bc..7ed967b5398f 100644 --- a/tests/integration/test_lists/test-db/l0_dgx_b300.yml +++ b/tests/integration/test_lists/test-db/l0_dgx_b300.yml @@ -23,10 +23,10 @@ l0_dgx_b300: # by absolute index, and they gate on >= 2 GPUs plus Blackwell. They also # rebind mpi4py's serializer at import time, which applies process-wide on # every stage that collects them, including where the tests themselves skip. - # The sparse/test_cute_dsl_* files ignored below keep their own entries (pre-existing). - - unittest/_torch/attention --ignore=unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py --ignore=unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py --ignore=unittest/_torch/attention/sparse/test_cute_dsl_gvr_topk_decode.py TIMEOUT (120) - - unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py - - unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py + # The sparse/dsa/test_cute_dsl_* files ignored below keep their own entries (pre-existing). + - unittest/_torch/attention --ignore=unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp8_paged_mqa_logits.py --ignore=unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp4_paged_mqa_logits.py --ignore=unittest/_torch/attention/sparse/test_cute_dsl_gvr_topk_decode.py TIMEOUT (120) + - unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp8_paged_mqa_logits.py + - unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp4_paged_mqa_logits.py - unittest/_torch/attention/sparse/test_cute_dsl_gvr_topk_decode.py TIMEOUT (120) - unittest/_torch/executor - unittest/_torch/disaggregation diff --git a/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp4.py b/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp4.py index d9b781a57340..56da8ffa5a19 100644 --- a/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp4.py +++ b/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp4.py @@ -7,7 +7,7 @@ CI (matching the convention of other scripts under ``tests/scripts/cute_dsl_kernels/``). - Helpers (FP4 quant, KV cast, ref) are inlined from - tests/unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py. + tests/unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp4_paged_mqa_logits.py. - Schedule metadata is computed in pure Python (mirrors DeepGEMM's PagedMQALogitsScheduler), avoiding the deep_gemm C++ binding. - Compile + dispatch follows tensorrt_llm/_torch/custom_ops/cute_dsl_custom_ops.py @@ -393,7 +393,7 @@ def get_paged_mqa_logits_metadata_cute_dsl( # Element-wise tolerance keyed by (epi_dtype, output_dtype) — mirrors the # unit test's ELEM_TOL table at -# tests/unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py. +# tests/unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp4_paged_mqa_logits.py. _ELEM_TOL = { (torch.float32, torch.float32): (5e-5, 1e-5), (torch.bfloat16, torch.bfloat16): (1e-2, 1e-2), diff --git a/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp8.py b/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp8.py index abc9ffe436a7..eb93bf0eed2b 100644 --- a/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp8.py +++ b/tests/scripts/cute_dsl_kernels/paged_mqa_logits/run_fp8.py @@ -7,7 +7,7 @@ (matching the convention of other scripts under ``tests/scripts/cute_dsl_kernels/``). - Reference and data prep are inlined from - tests/unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py. + tests/unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp8_paged_mqa_logits.py. - Schedule metadata is computed in pure Python (mirrors DeepGEMM's PagedMQALogitsScheduler), avoiding the deep_gemm C++ binding. Same algorithm as run_fp4.py — both kernels use compute_block_kv=128 + NUM_MATH_WG=2 → @@ -51,7 +51,7 @@ } # Element-wise tolerance keyed by output_dtype — mirrors the unit test -# tests/unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py +# tests/unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp8_paged_mqa_logits.py # which sets atol/rtol purely from output_dtype (epi_dtype = acc_dtype = # output_dtype in that test). _ELEM_TOL = { diff --git a/tests/unittest/_torch/attention/sparse/_sparse_mha_test_utils.py b/tests/unittest/_torch/attention/sparse/_sparse_mha_test_utils.py new file mode 100644 index 000000000000..73b98ecad00b --- /dev/null +++ b/tests/unittest/_torch/attention/sparse/_sparse_mha_test_utils.py @@ -0,0 +1,328 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Input, cache, and reference helpers for page-sparse MHA tests.""" + +import math +from contextlib import ExitStack +from dataclasses import dataclass +from typing import Optional, Tuple + +import torch + +import tensorrt_llm +from tensorrt_llm._torch.attention_backend.interface import AttentionRuntimeFeatures +from tensorrt_llm._torch.attention_backend.trtllm import ( + TrtllmAttentionMetadata, + generate_spec_decoding_packed_mask, + generate_spec_decoding_position_offsets, +) +from tensorrt_llm._torch.metadata import KVCacheParams +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm._utils import str_dtype_to_binding, torch_dtype_to_str +from tensorrt_llm.bindings.executor import KvCacheConfig +from tensorrt_llm.mapping import Mapping +from tensorrt_llm.models.modeling_utils import QuantConfig +from tensorrt_llm.quantization.mode import QuantAlgo + +SUPPORTED_MODEL_DTYPES = (torch.bfloat16, torch.float16) +SUPPORTED_KV_CACHE_DTYPES = (*SUPPORTED_MODEL_DTYPES, torch.float8_e4m3fn) +SUPPORTED_MHA_HEAD_DIMS = (64, 80, 128, 256) + + +@dataclass(kw_only=True, frozen=True) +class MhaGenerationScenario: + """Generation inputs for page-sparse MHA computation.""" + + dtype: torch.dtype = torch.bfloat16 + kvcache_dtype: torch.dtype = torch.bfloat16 + num_layers: int = 1 + num_heads: int = 8 + num_kv_heads: int = 8 + head_dim: int = 128 + page_size: int = 32 + num_pages: int = 4 + batch_size: int = 1 + past_kv_lens: Tuple[int, ...] = (96,) + query_len: int = 1 + fp8_output: bool = False + + def __post_init__(self) -> None: + if self.dtype not in SUPPORTED_MODEL_DTYPES: + raise ValueError("Model QKV dtype must be BF16 or FP16") + if self.kvcache_dtype not in SUPPORTED_KV_CACHE_DTYPES: + raise ValueError("KV-cache dtype must be BF16, FP16, or E4M3 FP8") + if self.kvcache_dtype != torch.float8_e4m3fn and self.kvcache_dtype != self.dtype: + raise ValueError("A non-FP8 KV-cache dtype must match the model QKV dtype") + if self.num_heads <= 0 or self.num_heads != self.num_kv_heads: + raise ValueError("Page-sparse MHA requires equal positive Q and KV head counts") + if self.head_dim not in SUPPORTED_MHA_HEAD_DIMS: + raise ValueError(f"head_dim must be one of {SUPPORTED_MHA_HEAD_DIMS}") + if self.page_size < 8 or self.page_size & (self.page_size - 1): + raise ValueError("page_size must be a power of two and at least 8") + if len(self.past_kv_lens) != self.batch_size: + raise ValueError( + f"past_kv_lens length {len(self.past_kv_lens)} must match " + f"batch_size {self.batch_size}" + ) + if self.query_len < 1: + raise ValueError("query_len must be positive") + if self.fp8_output and self.kvcache_dtype != torch.float8_e4m3fn: + raise ValueError("FP8 output testing requires an FP8 KV cache") + for past_kv_len in self.past_kv_lens: + required_pages = math.ceil((past_kv_len + self.query_len) / self.page_size) + if required_pages > self.num_pages: + raise ValueError("num_pages does not cover the request KV length") + + @property + def nnz_q(self) -> int: + return self.batch_size * self.query_len + + @property + def max_query_len(self) -> int: + return self.query_len + + @property + def has_draft_tokens(self) -> bool: + return self.query_len > 1 + + @property + def kv_pool_num_pages(self) -> int: + return self.batch_size * self.num_pages + + +@dataclass(kw_only=True) +class MhaGenerationInputs: + """Generation tensors plus their populated paged KV cache.""" + + q: torch.Tensor + k_new: torch.Tensor + v_new: torch.Tensor + kv_cache_manager: KVCacheManager + request_ids: list[int] + metadata: TrtllmAttentionMetadata + + @property + def fused_qkv(self) -> torch.Tensor: + return torch.cat([self.q, self.k_new, self.v_new], dim=1) + + +def quant_config(scenario: MhaGenerationScenario) -> Optional[QuantConfig]: + """Build the quantization settings used by the attention backend.""" + if scenario.fp8_output: + return QuantConfig( + quant_algo=QuantAlgo.FP8, + kv_cache_quant_algo=QuantAlgo.FP8, + ) + if scenario.kvcache_dtype == torch.float8_e4m3fn: + return QuantConfig(kv_cache_quant_algo=QuantAlgo.FP8) + return None + + +def fp8_qdq(tensor: torch.Tensor) -> torch.Tensor: + """Apply unit-scale E4M3 quantize-dequantize to a reference tensor.""" + return tensor.to(torch.float8_e4m3fn).to(tensor.dtype) + + +def _create_kv_cache_manager( + scenario: MhaGenerationScenario, + kv_cache: torch.Tensor, +) -> KVCacheManager: + kv_cache_config = KvCacheConfig(max_tokens=scenario.kv_pool_num_pages * scenario.page_size) + manager = KVCacheManager( + kv_cache_config, + tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF, + num_layers=scenario.num_layers, + num_kv_heads=scenario.num_kv_heads, + head_dim=scenario.head_dim, + tokens_per_block=scenario.page_size, + max_seq_len=scenario.kv_pool_num_pages * scenario.page_size, + max_batch_size=scenario.batch_size, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=str_dtype_to_binding(torch_dtype_to_str(scenario.kvcache_dtype)), + ) + for layer_idx in range(scenario.num_layers): + manager.get_buffers(layer_idx, kv_layout="HND").copy_(kv_cache[layer_idx]) + return manager + + +def create_generation_inputs(scenario: MhaGenerationScenario) -> MhaGenerationInputs: + """Create generation inputs and release the cache if construction fails.""" + device = torch.device("cuda") + torch.manual_seed(42) + q = torch.randn( + scenario.nnz_q, + scenario.num_heads * scenario.head_dim, + device=device, + dtype=scenario.dtype, + ) + k_new = torch.randn( + scenario.nnz_q, + scenario.num_kv_heads * scenario.head_dim, + device=device, + dtype=scenario.dtype, + ) + v_new = torch.randn_like(k_new) + kv_cache = torch.randn( + scenario.num_layers, + scenario.kv_pool_num_pages, + 2, + scenario.num_kv_heads, + scenario.page_size, + scenario.head_dim, + device=device, + dtype=scenario.dtype, + ).to(scenario.kvcache_dtype) + + with ExitStack() as cleanup: + kv_cache_manager = _create_kv_cache_manager(scenario, kv_cache) + cleanup.callback(kv_cache_manager.shutdown) + request_ids = list(range(scenario.batch_size)) + token_nums = [past_kv_len + scenario.query_len for past_kv_len in scenario.past_kv_lens] + kv_cache_manager.add_dummy_requests(request_ids, token_nums) + metadata = TrtllmAttentionMetadata( + num_contexts=0, + kv_cache_params=KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=list(scenario.past_kv_lens), + ), + seq_lens=torch.full((scenario.batch_size,), scenario.query_len, dtype=torch.int32), + max_num_requests=scenario.batch_size, + max_num_tokens=scenario.nnz_q, + kv_cache_manager=kv_cache_manager, + request_ids=request_ids, + prompt_lens=list(scenario.past_kv_lens), + num_heads_per_kv=1, + runtime_features=AttentionRuntimeFeatures( + has_speculative_draft_tokens=scenario.has_draft_tokens + ), + is_spec_decoding_enabled=scenario.has_draft_tokens, + use_spec_decoding=scenario.has_draft_tokens, + is_spec_dec_tree=False, + max_total_draft_tokens=( + scenario.max_query_len - 1 if scenario.has_draft_tokens else None + ), + ) + if scenario.has_draft_tokens: + draft_len = scenario.max_query_len - 1 + metadata.spec_decoding_position_offsets = generate_spec_decoding_position_offsets( + scenario.batch_size, draft_len + ) + metadata.spec_decoding_packed_mask = generate_spec_decoding_packed_mask( + scenario.batch_size, draft_len + ) + metadata.spec_decoding_generation_lengths = torch.tensor( + [scenario.query_len] * scenario.batch_size, + dtype=torch.int32, + device=device, + ) + metadata.update_position_offsets_for_cpp(scenario.max_query_len) + metadata.spec_decoding_param_prepare_for_blackwell() + metadata.prepare() + cleanup.pop_all() + return MhaGenerationInputs( + q=q, + k_new=k_new, + v_new=v_new, + kv_cache_manager=kv_cache_manager, + request_ids=request_ids, + metadata=metadata, + ) + + +def read_paged_kv_cache( + inputs: MhaGenerationInputs, + scenario: MhaGenerationScenario, +) -> list[tuple[torch.Tensor, torch.Tensor]]: + """Materialize paged history using page-wise copies.""" + kv_buffer = inputs.kv_cache_manager.get_buffers(0, kv_layout="HND") + kv_caches = [] + for request_id, num_tokens in zip(inputs.request_ids, scenario.past_kv_lens, strict=True): + block_ids = inputs.kv_cache_manager.get_block_ids_per_seq([request_id])[0] + k_cache = torch.empty( + num_tokens, + scenario.num_kv_heads, + scenario.head_dim, + device=kv_buffer.device, + dtype=scenario.dtype, + ) + v_cache = torch.empty_like(k_cache) + for local_page_idx, block_id in enumerate(block_ids): + token_start = local_page_idx * scenario.page_size + token_end = min(token_start + scenario.page_size, num_tokens) + if token_start >= token_end: + break + num_page_tokens = token_end - token_start + k_cache[token_start:token_end] = ( + kv_buffer[block_id, 0, :, :num_page_tokens, :].transpose(0, 1).to(scenario.dtype) + ) + v_cache[token_start:token_end] = ( + kv_buffer[block_id, 1, :, :num_page_tokens, :].transpose(0, 1).to(scenario.dtype) + ) + kv_caches.append((k_cache, v_cache)) + return kv_caches + + +def reference_generation_attention( + q: torch.Tensor, + kv_caches: list[tuple[torch.Tensor, torch.Tensor]], + k_new: torch.Tensor, + v_new: torch.Tensor, + sparse_attn_indices: torch.Tensor, + scenario: MhaGenerationScenario, +) -> torch.Tensor: + """Compute page-sparse MHA from equivalent request-local token indices.""" + outputs = [] + query_offset = 0 + for request_idx in range(scenario.batch_size): + k_history, v_history = kv_caches[request_idx] + request_slice = slice(query_offset, query_offset + scenario.query_len) + k_full = torch.cat( + [ + k_history, + k_new[request_slice].view( + scenario.query_len, scenario.num_kv_heads, scenario.head_dim + ), + ], + dim=0, + ) + v_full = torch.cat( + [ + v_history, + v_new[request_slice].view( + scenario.query_len, scenario.num_kv_heads, scenario.head_dim + ), + ], + dim=0, + ) + for query_idx in range(scenario.query_len): + packed_query_idx = query_offset + query_idx + q_token = q[packed_query_idx].view(scenario.num_heads, scenario.head_dim) + head_outputs = [] + for head_idx in range(scenario.num_heads): + token_indices = sparse_attn_indices[head_idx, packed_query_idx] + valid_indices = token_indices[token_indices >= 0].long() + k_sparse = k_full[valid_indices, head_idx, :] + v_sparse = v_full[valid_indices, head_idx, :] + attention_scores = torch.matmul(q_token[head_idx], k_sparse.T) / math.sqrt( + scenario.head_dim + ) + attention_probs = torch.nn.functional.softmax( + attention_scores, dim=-1, dtype=torch.float32 + ).to(scenario.dtype) + head_outputs.append(torch.matmul(attention_probs, v_sparse)) + outputs.append(torch.cat(head_outputs, dim=0)) + query_offset += scenario.query_len + return torch.stack(outputs, dim=0) diff --git a/tests/unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py b/tests/unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp4_paged_mqa_logits.py similarity index 100% rename from tests/unittest/_torch/attention/sparse/test_cute_dsl_fp4_paged_mqa_logits.py rename to tests/unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp4_paged_mqa_logits.py diff --git a/tests/unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py b/tests/unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp8_paged_mqa_logits.py similarity index 100% rename from tests/unittest/_torch/attention/sparse/test_cute_dsl_fp8_paged_mqa_logits.py rename to tests/unittest/_torch/attention/sparse/dsa/test_cute_dsl_fp8_paged_mqa_logits.py diff --git a/tests/unittest/_torch/attention/sparse/test_dsa_fp4_indexer.py b/tests/unittest/_torch/attention/sparse/test_dsa_fp4_indexer.py deleted file mode 100644 index 09afc561f724..000000000000 --- a/tests/unittest/_torch/attention/sparse/test_dsa_fp4_indexer.py +++ /dev/null @@ -1,426 +0,0 @@ -# Copyright (c) 2022-2026, NVIDIA CORPORATION. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -"""Integration tests for the DSA FP4 indexer path (B200 / SM100 only). - -These tests drive the DeepGEMM FP4 MQA logits kernel through the TRT-LLM -Indexer's FP4 quantization op and the Indexer._call_mqa_logits dispatch. -Compared against the FP8 reference: -- Topk intersection rate between FP4 and FP8 should be >= 95% for the - same inputs, confirming the two indexer implementations pick - essentially the same candidate key tokens. -- The FP4 kernel must accept head_dim=128, the supported 64-head path, and - the packed int8/int32 layouts produced by torch.ops.trtllm.fused_cat_fp4. - -The DSA config validator rejects FP4 on SM<100 and on non-128 head_dim, -so skip when either precondition isn't met. -""" - -import pytest -import torch - -# Import tensorrt_llm to load C++ custom operators (registers trtllm::fused_cat_fp4). -import tensorrt_llm # noqa: F401 - -try: - from tensorrt_llm import deep_gemm -except ImportError: - # Only skip on actual module-missing failures — any other error should - # surface rather than silently turn into a test-wide skip. - HAS_DEEP_GEMM = False -else: - HAS_DEEP_GEMM = True - # If deep_gemm imports but fp8_fp4_mqa_logits is absent, the installed - # DeepGEMM version is wrong — fail loudly instead of silently skipping. - assert hasattr(deep_gemm, "fp8_fp4_mqa_logits"), ( - "deep_gemm imported but fp8_fp4_mqa_logits is missing; " - "check that the correct DeepGEMM version is installed" - ) - -from utils.util import skip_pre_blackwell - -from .dsa.test_dsa_indexer import _create_mock_metadata, create_dsa_cache_manager - -FP4_MQA_NUM_HEADS = [ - pytest.param( - 32, - marks=pytest.mark.skip( - reason="DeepGEMM fp8_fp4_mqa_logits JIT currently fails for 32 heads" - ), - ), - 64, -] - - -def _fp4_quantize_sf_transpose(x: torch.Tensor): - """Wrap trtllm::fused_cat_fp4 for tests that already hold a concatenated - tensor. The op takes (pe, nope); split at an arbitrary boundary since the - kernel reconstructs the concat internally. Returns shapes matching the - original helper: (*leading, head_dim//2) int8 packed and (*leading, 1) int32. - """ - head_dim = x.shape[-1] - assert head_dim == 128, f"expected head_dim=128, got {head_dim}" - pe, nope = x.split([head_dim // 2, head_dim // 2], dim=-1) - packed, scale = torch.ops.trtllm.fused_cat_fp4(pe, nope) - leading = x.shape[:-1] - return packed.view(*leading, head_dim // 2), scale.view(*leading, 1) - - -def _fp8_quantize_sf(x: torch.Tensor): - """Quantize along the sequence dim, mirroring test_dsa_indexer.""" - x_amax = x.abs().float().amax(dim=tuple(range(1, x.dim())), keepdim=True).clamp(1e-4) - sf = x_amax / 448.0 - x_scaled = (x * (1.0 / sf)).to(torch.float8_e4m3fn) - return x_scaled, sf.squeeze() - - -def _dense_context_bounds(seq_len: int, seq_len_kv: int, device): - """Causal attention window: token i attends to [0, seq_len_kv - seq_len + i).""" - cu_ks = torch.zeros(seq_len, dtype=torch.int32, device=device) - cu_ke = torch.arange(1, seq_len + 1, dtype=torch.int32, device=device) + (seq_len_kv - seq_len) - return cu_ks, cu_ke.to(torch.int32) - - -@pytest.mark.skipif(not HAS_DEEP_GEMM, reason="fp8_fp4_mqa_logits not available") -@skip_pre_blackwell -@pytest.mark.parametrize("num_heads", FP4_MQA_NUM_HEADS) -def test_fp4_mqa_logits_shape_and_topk_intersection(num_heads): - """FP4 MQA logits agree with FP8 on the top-k key selection.""" - torch.manual_seed(0) - head_dim = 128 - seq_len = 128 - seq_len_kv = 512 - - q = torch.randn(seq_len, num_heads, head_dim, device="cuda", dtype=torch.bfloat16) * 1.5 - k = torch.randn(seq_len_kv, head_dim, device="cuda", dtype=torch.bfloat16) - weights = torch.randn(seq_len, num_heads, device="cuda", dtype=torch.float32) - cu_ks, cu_ke = _dense_context_bounds(seq_len, seq_len_kv, q.device) - - # FP4 path: pack Q and K. _fp4_quantize_sf_transpose keeps a trailing - # num_blocks//4 dim to stay byte-identical with DeepGEMM's reference util, - # so squeeze it for the kernel (q_sf is 2D, kv_sf is 1D). - q_fp4, q_scale_full = _fp4_quantize_sf_transpose(q) - q_scale = q_scale_full.view(seq_len, num_heads) - k_fp4, k_scale_full = _fp4_quantize_sf_transpose(k) - k_scale_fp4 = k_scale_full.reshape(-1) - - # The FP4 kernel scales q internally; weights carry softmax_scale only. - softmax_scale = head_dim**-0.5 - n_heads_scale = num_heads**-0.5 - fp4_weights = weights * softmax_scale * n_heads_scale - fp4_logits = deep_gemm.fp8_fp4_mqa_logits( - (q_fp4, q_scale), - (k_fp4, k_scale_fp4), - fp4_weights, - cu_ks, - cu_ke, - False, # clean_logits - 0, # max_seqlen_k - torch.float32, # logits_dtype - ) - assert fp4_logits.shape == (seq_len, seq_len_kv) - assert fp4_logits.dtype == torch.float32 - - # FP8 reference: the legacy fp8_mqa_logits pre-scales weights with q_scale - # so the logits come out in the same numeric range. - q_fp8, q_scale_fp8 = _fp8_quantize_sf(q) - k_fp8, k_scale_fp8 = _fp8_quantize_sf(k) - fp8_weights = weights * q_scale_fp8.unsqueeze(-1) * softmax_scale * n_heads_scale - fp8_logits = deep_gemm.fp8_mqa_logits(q_fp8, (k_fp8, k_scale_fp8), fp8_weights, cu_ks, cu_ke) - - topk = 32 - fp4_valid = torch.where( - torch.arange(seq_len_kv, device="cuda").unsqueeze(0) < cu_ke.unsqueeze(1), - fp4_logits, - float("-inf"), - ) - fp8_valid = torch.where( - torch.arange(seq_len_kv, device="cuda").unsqueeze(0) < cu_ke.unsqueeze(1), - fp8_logits, - float("-inf"), - ) - fp4_top = fp4_valid.topk(topk, dim=-1).indices - fp8_top = fp8_valid.topk(topk, dim=-1).indices - - # Per-row intersection ratio between the two indexer variants. - intersections = [] - for i in range(seq_len): - a = set(fp4_top[i].tolist()) - b = set(fp8_top[i].tolist()) - if len(b) == 0: - continue - intersections.append(len(a & b) / len(b)) - mean_overlap = sum(intersections) / len(intersections) - # FP4 has 8 representable levels vs. FP8's ~240, so on synthetic random - # inputs the top-k lists diverge slightly even though the kernels are - # numerically consistent. The plan targets >= 95% intersection on real - # DSA traffic (where logit magnitudes are more polarized); for this - # shape-only sanity test a 0.80 floor catches gross regressions without - # flaking on random-seed noise. - assert mean_overlap >= 0.80, ( - f"FP4 vs FP8 topk overlap too low: {mean_overlap:.3f}. " - "Expect >= 0.80 mean overlap on synthetic inputs." - ) - - -@pytest.mark.skipif(not HAS_DEEP_GEMM, reason="fp8_fp4_mqa_logits not available") -@skip_pre_blackwell -def test_fp4_quantize_roundtrip_matches_bf16_kv(): - """Verify FP4 K quantize+dequantize preserves the dominant magnitudes. - - Sanity-checks the packing / scale recovery math outside the kernel so a - failure localizes between the Python quantizer and the DeepGEMM kernel. - """ - torch.manual_seed(7) - seq_len_kv = 128 - head_dim = 128 - k = torch.randn(seq_len_kv, head_dim, device="cuda", dtype=torch.bfloat16) * 2.0 - - k_fp4, scale = _fp4_quantize_sf_transpose(k) - - fp4_values = torch.tensor( - [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], - device="cuda", - dtype=torch.float32, - ) - packed_u8 = k_fp4.view(torch.uint8) - low = packed_u8 & 0x0F - high = (packed_u8 >> 4) & 0x0F - codes = torch.empty(seq_len_kv, head_dim, device="cuda", dtype=torch.uint8) - codes[:, 0::2] = low - codes[:, 1::2] = high - value_idx = (codes & 0x07).to(torch.int64) - sign = (codes & 0x08) != 0 - values = fp4_values[value_idx] - values = torch.where(sign & (value_idx != 0), -values, values) - scale_bytes = scale.view(torch.uint8).view(seq_len_kv, 4).to(torch.int32) - scale_fp32 = (scale_bytes << 23).view(torch.float32) - reconstructed = (values.view(seq_len_kv, 4, 32) * scale_fp32.unsqueeze(-1)).view( - seq_len_kv, head_dim - ) - - # MAE should be bounded by the FP4 step (~0.5 * max per block) — very loose, - # but clearly rules out catastrophic unpacking bugs. - mae = (reconstructed.float() - k.float()).abs().mean().item() - assert mae < 1.0, f"FP4 dequantize diverged from bf16 input: mae={mae:.3f}" - - -@pytest.mark.skipif(not HAS_DEEP_GEMM, reason="fp8_fp4_mqa_logits not available") -@skip_pre_blackwell -def test_fp4_indexer_k_cache_per_token_size_drops_to_68_bytes(): - """Evidence for the plan's primary goal: FP4 indexer K cache shrinks. - - The FP8 layout stores index_head_dim bytes of data + 4 bytes of float32 - scale per token (132 bytes at index_head_dim=128). The FP4 layout packs - two E2M1 codes per byte (index_head_dim // 2 = 64 bytes) and keeps the - same 4 scale bytes (UE8M0 x4 packed as one int32), for a total of 68 - bytes per token. - """ - # Simulate the pool allocation formula exactly as WindowBlockManager:: - # createIndexerKCachePools (kvCacheManager.cpp) and DSACacheManager:: - # get_indexer_k_cache_buffers (dsa.py) compute per-token size. - index_head_dim = 128 - quant_block_size = 128 - scale_bytes = index_head_dim // quant_block_size * 4 # 4 bytes either way - - fp8_data_bytes = index_head_dim - fp8_per_token = fp8_data_bytes + scale_bytes - - fp4_data_bytes = index_head_dim // 2 - fp4_per_token = fp4_data_bytes + scale_bytes - - assert fp8_per_token == 132, f"FP8 per-token size regressed from 132 to {fp8_per_token}" - assert fp4_per_token == 68, f"FP4 per-token size regressed from 68 to {fp4_per_token}" - assert fp4_per_token / fp8_per_token < 0.52, ( - f"FP4 pool did not shrink as expected: {fp4_per_token}/{fp8_per_token}" - ) - - -@skip_pre_blackwell -def test_indexer_k_dtype_survives_model_config_rebuild(): - """Regression guard: indexer_k_dtype must survive ModelConfig.from_pretrained. - - When loading DeepseekV32ForCausalLM / GlmMoeDsaForCausalLM, - ModelConfig.from_pretrained rebuilds the sparse_attention_config from the - user's fields plus pretrained-config defaults. A previous version of this - rebuild dropped indexer_k_dtype, silently forcing fp8 regardless of the - user's choice — the Pydantic validator and downstream DSACacheManager - then both saw "fp8" and the FP4 path was never taken. - - Exercise the rebuild with a stub pretrained_config instead of a real - checkpoint so the test is cheap (no weight load) and hermetic. - """ - from types import SimpleNamespace - from unittest.mock import patch - - from tensorrt_llm._torch.model_config import ModelConfig - from tensorrt_llm.llmapi.llm_args import DeepSeekSparseAttentionConfig - - stub_pretrained = SimpleNamespace( - architectures=["DeepseekV32ForCausalLM"], - index_n_heads=64, - index_head_dim=128, - index_topk=2048, - indexer_rope_interleave=False, - ) - user_config = DeepSeekSparseAttentionConfig( - index_head_dim=128, - indexer_k_dtype="fp4", - ) - - # Patch load_pretrained_config to return the stub, then exercise the - # DSV3.2 rebuild branch via the helper that actually rebuilds the - # sparse_attention_config. We don't call ModelConfig.from_pretrained - # end-to-end because it pulls in quantization/tokenizer machinery that - # needs a real on-disk checkpoint; instead we patch the one function - # whose return value the rebuild branch reads and invoke it directly. - rebuilt_kwargs: dict = {"sparse_attention_config": user_config} - with patch( - "tensorrt_llm._torch.model_config.load_pretrained_config", - return_value=stub_pretrained, - ): - # Inline the rebuild snippet from ModelConfig.from_pretrained so the - # test doesn't depend on checkpoint loaders. Keep in sync with - # ModelConfig.from_pretrained's DSV3.2 branch. - sparse_attn_config = rebuilt_kwargs["sparse_attention_config"] - rebuilt_kwargs["sparse_attention_config"] = DeepSeekSparseAttentionConfig( - index_n_heads=sparse_attn_config.index_n_heads or stub_pretrained.index_n_heads, - index_head_dim=sparse_attn_config.index_head_dim or stub_pretrained.index_head_dim, - index_topk=sparse_attn_config.index_topk or stub_pretrained.index_topk, - indexer_max_chunk_size=sparse_attn_config.indexer_max_chunk_size, - skip_indexer_for_short_seqs=sparse_attn_config.skip_indexer_for_short_seqs, - use_cute_dsl_topk=sparse_attn_config.use_cute_dsl_topk, - q_split_threshold=sparse_attn_config.q_split_threshold, - indexer_rope_interleave=stub_pretrained.indexer_rope_interleave, - enable_heuristic_topk=sparse_attn_config.enable_heuristic_topk, - indexer_k_dtype=sparse_attn_config.indexer_k_dtype, - ) - - rebuilt = rebuilt_kwargs["sparse_attention_config"] - assert rebuilt.indexer_k_dtype == "fp4", ( - f"indexer_k_dtype dropped during rebuild: got {rebuilt.indexer_k_dtype}" - ) - assert rebuilt.index_head_dim == 128 - # Static check: ensure the production rebuild branch actually forwards - # the field (regression guard for the pattern bug — not just the outcome - # of this test). If a future edit drops the keyword, this assert fails. - import inspect - - rebuild_src = inspect.getsource(ModelConfig.from_pretrained) - assert "indexer_k_dtype=indexer_k_dtype" in rebuild_src, ( - "ModelConfig.from_pretrained rebuild branch must forward " - "indexer_k_dtype to DeepSeekSparseAttentionConfig(...); otherwise " - "the user-visible FP4 knob will be silently dropped." - ) - - -@skip_pre_blackwell -def test_indexer_k_cache_scatter_custom_op_fp4(): - """FP4 variant: CUDA kernel vs Python reference for k_cache scatter. - - Under FP4 the data payload is head_dim//2 bytes (two packed E2M1 codes - per byte) and the scale is a single int32 per token. Verify the scatter - op handles the shorter per-token size correctly. - """ - torch.manual_seed(456) - - head_dim = 128 - fp4_data_dim = head_dim // 2 # 64 bytes packed - block_size = 64 - batch_size = 2 - num_tokens = 64 - max_seq_len = 512 - - layer_idx_cuda = 0 - layer_idx_python = 1 - - cache_manager, _ = create_dsa_cache_manager( - batch_size=batch_size, - head_dim=head_dim, - tokens_per_block=block_size, - max_seq_len=max_seq_len, - num_layers=3, - indexer_k_dtype="fp4", - ) - - request_ids = list(range(batch_size)) - tokens_per_req = [32, 32] - cache_manager.add_dummy_requests( - request_ids, tokens_per_req, is_gen=False, prepare_resource=True - ) - - metadata = _create_mock_metadata( - request_ids, - batch_size, - num_contexts=batch_size, - num_generations=0, - seq_lens=torch.tensor(tokens_per_req, dtype=torch.int32), - kv_lens=torch.tensor(tokens_per_req, dtype=torch.int32), - num_cached_tokens=[0] * batch_size, - cache_manager=cache_manager, - num_ctx_tokens=num_tokens, - num_tokens=num_tokens, - ) - - from tensorrt_llm._torch.attention.backends.sparse.dsa import Indexer - - Indexer.prepare(metadata) - - # FP4 packed data: [num_tokens, 64] int8; scale: [num_tokens, 1] int32 - k_fp4 = torch.randint(-128, 127, (num_tokens, fp4_data_dim), device="cuda", dtype=torch.int8) - k_scale = torch.randint(0, 2**31, (num_tokens, 1), device="cuda", dtype=torch.int32) - - scale_size = 4 # 1 int32 = 4 bytes - k_fp4_bytes = k_fp4.view(torch.uint8) - k_scale_bytes = k_scale.view(torch.uint8).view(num_tokens, scale_size) - - flat_indices_fp8 = metadata.slot_mapping_fp8[:num_tokens] - flat_indices_scale = metadata.slot_mapping_scale[:num_tokens] - - # CUDA path - k_cache_cuda = cache_manager.get_indexer_k_cache_buffers(layer_idx_cuda) - k_cache_cuda.zero_() - torch.ops.trtllm.indexer_k_cache_scatter_op( - k_fp4, - k_scale, - k_cache_cuda, - metadata.slot_mapping_fp8, - metadata.slot_mapping_scale, - num_tokens, - ) - torch.cuda.synchronize() - - # Python reference - k_cache_python = cache_manager.get_indexer_k_cache_buffers(layer_idx_python) - k_cache_python.zero_() - - def _unravel_indices(flat_indices, shape): - d3 = shape[3] - i3 = flat_indices % d3 - flat_indices = flat_indices // d3 - d2 = shape[2] - i2 = flat_indices % d2 - flat_indices = flat_indices // d2 - d1 = shape[1] - i1 = flat_indices % d1 - flat_indices = flat_indices // d1 - i0 = flat_indices - return i0, i1, i2, i3 - - byte_offsets = torch.arange(fp4_data_dim, device=k_cache_python.device).unsqueeze(0) - scatter_fp4 = flat_indices_fp8.unsqueeze(1) + byte_offsets - scatter_fp4 = _unravel_indices(scatter_fp4, k_cache_python.shape) - k_cache_python[scatter_fp4] = k_fp4_bytes - - byte_offsets = torch.arange(scale_size, device=k_cache_python.device).unsqueeze(0) - scatter_scale = flat_indices_scale.unsqueeze(1) + byte_offsets - scatter_scale = _unravel_indices(scatter_scale, k_cache_python.shape) - k_cache_python[scatter_scale] = k_scale_bytes - - assert torch.equal(k_cache_cuda, k_cache_python), ( - "FP4 scatter: CUDA kernel produced different results than Python reference" - ) diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_mha.py b/tests/unittest/_torch/attention/sparse/test_sparse_mha.py new file mode 100644 index 000000000000..58dd1b361ff5 --- /dev/null +++ b/tests/unittest/_torch/attention/sparse/test_sparse_mha.py @@ -0,0 +1,473 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Architecture-level regression tests for page-sparse MHA computation. + +The tests supply static block indices and per-request offsets, invoke +``TrtllmAttention.forward``, and compare the result with an equivalent +token-level PyTorch reference. + +Prefill attention is dense on this path; page-sparse MHA computation starts +during generation. +""" + +import math +from dataclasses import dataclass +from typing import Optional, Tuple + +import pytest +import torch +from utils.util import getSMVersion + +from tensorrt_llm._torch.attention.backends.interface import AttentionForwardArgs +from tensorrt_llm._torch.attention.backends.sparse.params import SparseParams +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttention + +from ._sparse_mha_test_utils import ( + SUPPORTED_MHA_HEAD_DIMS, + SUPPORTED_MODEL_DTYPES, + MhaGenerationScenario, + create_generation_inputs, + fp8_qdq, + quant_config, + read_paged_kv_cache, + reference_generation_attention, +) + +ATOL = 2e-2 +RTOL = 2e-2 +FP8_ATOL = 8e-2 +FP8_RTOL = 4e-2 +TESTED_MHA_HEAD_COUNTS = (1, 2, 3, 4, 8, 16, 24, 32, 48, 64, 96, 128) +TESTED_KV_PAGE_SIZES = (8, 16, 32, 64, 128, 256, 512) +SUPPORTED_SM_VERSIONS = (100, 103) + +pytestmark = pytest.mark.skipif( + getSMVersion() not in SUPPORTED_SM_VERSIONS, + reason="Page-sparse MHA requires SM100 or SM103", +) + + +@pytest.fixture(autouse=True) +def _force_trtllm_gen_fallback(monkeypatch: pytest.MonkeyPatch) -> None: + """Keep every test on the TRTLLM-Gen fallback path.""" + monkeypatch.setenv("TLLM_FMHA_LIBS", "fallback") + + +# Page-sparse MHA support matrix: +# +# GPU architecture SM100 and SM103 +# Sparse compute phase Single-token and linear draft-token generation +# Attention type MHA; num_heads == num_kv_heads +# Q heads per KV head 1 +# Number of MHA heads No discrete source restriction; tests cover +# 1, 2, 3, 4, 8, 16, 24, 32, 48, 64, 96, and 128 +# Model QKV input BF16 or FP16 +# Model QKV layout Fused QKV +# Kernel output Model dtype for H64/H80/H128/H256; E4M3 FP8 for +# H64/H128/H256 with an FP8 KV cache +# KV-cache dtype Model dtype for H64/H80/H128/H256; E4M3 FP8 for +# H64/H128/H256 +# Q/K/V head dimension Equal dimensions: 64, 80, 128, or 256 +# KV-cache layout Paged; page sizes 8, 16, 32, 64, 128, 256, and 512 +# Selection granularity Block indices expanded to KV-cache pages +# Sparse indices int32 block indices plus int32 request offsets +# Per-head patterns and variable request offsets +# Sparse index block Blocks may cross KV-page boundaries; sizes +# 1, 2, 3, 4, 5, 8, 16, 24, 32, and 48 are tested +# Attention semantics Causal self-attention + + +@dataclass(kw_only=True, frozen=True) +class PageSparseMhaScenario: + """Page-sparse MHA geometry layered on generation inputs.""" + + attention: MhaGenerationScenario + sparse_index_block_size: int = 4 + num_selected_sparse_blocks: int = 2 + + def __post_init__(self) -> None: + if self.attention.num_heads != self.attention.num_kv_heads: + raise ValueError("PageSparseMhaScenario requires MHA head geometry") + if self.sparse_index_block_size <= 0: + raise ValueError("sparse_index_block_size must be positive") + if self.num_selected_sparse_blocks <= 0: + raise ValueError("num_selected_sparse_blocks must be positive") + + +@dataclass(kw_only=True, frozen=True) +class _PageSparseMhaParams(SparseParams): + """Sparse parameters selecting block/page-granular attention.""" + + sparse_index_block_size: int + algorithm: str = "test_page_sparse_mha" + + @property + def indices_block_size(self) -> int: + return self.sparse_index_block_size + + +class _StaticPageSparseMhaAttention(TrtllmAttention): + """MHA backend adapter returning predetermined page selections.""" + + def __init__( + self, + *args, + sparse_index_block_size: int, + sparse_attn_indices: torch.Tensor, + sparse_attn_offsets: torch.Tensor, + **kwargs, + ) -> None: + kwargs["sparse_params"] = _PageSparseMhaParams( + sparse_index_block_size=sparse_index_block_size + ) + kwargs["pos_embd_params"] = None + super().__init__(*args, **kwargs) + self._sparse_attn_indices = sparse_attn_indices + self._sparse_attn_offsets = sparse_attn_offsets + + def sparse_kv_predict(self, q, k, metadata, forward_args: AttentionForwardArgs): + return None, None + + def sparse_attn_predict(self, q, k, metadata, forward_args: AttentionForwardArgs): + return self._sparse_attn_indices, self._sparse_attn_offsets + + +def _selected_sparse_blocks( + head_idx: int, + num_sparse_blocks: int, + num_selected_blocks: int, + page_size: int, + sparse_block_size: int, +) -> Tuple[int, ...]: + """Choose head-dependent blocks and retain the newest sparse block.""" + newest_block = num_sparse_blocks - 1 + if num_selected_blocks == 1: + return (newest_block,) + + older_blocks = list(range(newest_block)) + boundary_block = page_size // sparse_block_size + if head_idx % 2 and boundary_block in older_blocks: + older_blocks.remove(boundary_block) + older_blocks.insert(0, boundary_block) + selected = older_blocks[: num_selected_blocks - 1] + [newest_block] + if head_idx % 2: + selected.reverse() + return tuple(selected) + + +def _make_page_sparse_pattern( + scenario: PageSparseMhaScenario, + device: torch.device, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Build page indices plus equivalent request-local token indices for the reference.""" + attention = scenario.attention + sparse_indices_by_head = [[] for _ in range(attention.num_kv_heads)] + sparse_offsets = [0] + max_reference_tokens = max(attention.past_kv_lens) + attention.query_len + reference_indices = torch.full( + (attention.num_kv_heads, attention.nnz_q, max_reference_tokens), + -1, + dtype=torch.int32, + device=device, + ) + + query_offset = 0 + for past_kv_len in attention.past_kv_lens: + total_kv_len = past_kv_len + attention.query_len + num_sparse_blocks = math.ceil(total_kv_len / scenario.sparse_index_block_size) + num_selected_sparse_blocks = min(scenario.num_selected_sparse_blocks, num_sparse_blocks) + for head_idx in range(attention.num_kv_heads): + sparse_blocks = _selected_sparse_blocks( + head_idx, + num_sparse_blocks, + num_selected_sparse_blocks, + attention.page_size, + scenario.sparse_index_block_size, + ) + sparse_indices_by_head[head_idx].extend(sparse_blocks) + + for query_idx in range(attention.query_len): + available_kv_len = past_kv_len + query_idx + 1 + touched_pages = set() + for sparse_block_idx in sparse_blocks: + block_start = sparse_block_idx * scenario.sparse_index_block_size + block_end = min( + block_start + scenario.sparse_index_block_size, + available_kv_len, + ) + if block_start >= block_end: + continue + first_page = block_start // attention.page_size + last_page = (block_end - 1) // attention.page_size + touched_pages.update(range(first_page, last_page + 1)) + selected_tokens = [] + for page_idx in sorted(touched_pages): + page_start = page_idx * attention.page_size + page_end = min(page_start + attention.page_size, available_kv_len) + selected_tokens.extend(range(page_start, page_end)) + reference_indices[ + head_idx, + query_offset + query_idx, + : len(selected_tokens), + ] = torch.tensor(selected_tokens, dtype=torch.int32, device=device) + + sparse_offsets.append(sparse_offsets[-1] + num_selected_sparse_blocks) + query_offset += attention.query_len + + sparse_attn_indices = torch.tensor( + sparse_indices_by_head, + dtype=torch.int32, + device=device, + ) + sparse_attn_offsets = torch.tensor(sparse_offsets, dtype=torch.int32, device=device) + return sparse_attn_indices, sparse_attn_offsets, reference_indices + + +def _run_page_sparse_mha(scenario: PageSparseMhaScenario) -> None: + """Compare page-sparse generation with an equivalent PyTorch token reference.""" + attention_scenario = scenario.attention + inputs = create_generation_inputs(attention_scenario) + try: + sparse_attn_indices, sparse_attn_offsets, reference_indices = _make_page_sparse_pattern( + scenario, inputs.q.device + ) + attention = _StaticPageSparseMhaAttention( + layer_idx=0, + num_heads=attention_scenario.num_heads, + head_dim=attention_scenario.head_dim, + num_kv_heads=attention_scenario.num_kv_heads, + quant_config=quant_config(attention_scenario), + sparse_index_block_size=scenario.sparse_index_block_size, + sparse_attn_indices=sparse_attn_indices, + sparse_attn_offsets=sparse_attn_offsets, + ) + + kv_caches = read_paged_kv_cache(inputs, attention_scenario) + reference_q = inputs.q + reference_k_new = inputs.k_new + reference_v_new = inputs.v_new + if attention_scenario.kvcache_dtype == torch.float8_e4m3fn: + reference_q = fp8_qdq(reference_q) + reference_k_new = fp8_qdq(reference_k_new) + reference_v_new = fp8_qdq(reference_v_new) + reference_output = reference_generation_attention( + reference_q, + kv_caches, + reference_k_new, + reference_v_new, + reference_indices, + attention_scenario, + ) + + forward_args: Optional[AttentionForwardArgs] = None + if attention_scenario.fp8_output: + forward_args = AttentionForwardArgs( + out_scale=torch.ones(1, dtype=torch.float32, device=inputs.q.device) + ) + output = attention.forward( + inputs.fused_qkv, + None, + None, + inputs.metadata, + forward_args=forward_args, + ) + + expected_shape = ( + attention_scenario.nnz_q, + attention_scenario.num_heads * attention_scenario.head_dim, + ) + assert output.shape == expected_shape + expected_output_dtype = ( + torch.float8_e4m3fn if attention_scenario.fp8_output else attention_scenario.dtype + ) + assert output.dtype == expected_output_dtype + uses_fp8 = ( + attention_scenario.kvcache_dtype == torch.float8_e4m3fn or attention_scenario.fp8_output + ) + output_for_comparison = output.float() if uses_fp8 else output + if attention_scenario.fp8_output: + reference_output = reference_output.to(torch.float8_e4m3fn) + reference_for_comparison = reference_output.float() if uses_fp8 else reference_output + assert torch.isfinite(output_for_comparison).all() + torch.testing.assert_close( + output_for_comparison, + reference_for_comparison, + atol=FP8_ATOL if uses_fp8 else ATOL, + rtol=FP8_RTOL if uses_fp8 else RTOL, + ) + finally: + inputs.kv_cache_manager.shutdown() + + +_NUM_MHA_HEADS = 8 + +_PAGE_SPARSE_MHA_CASES = ( + [ + pytest.param( + PageSparseMhaScenario( + attention=MhaGenerationScenario( + dtype=dtype, + kvcache_dtype=dtype, + num_heads=_NUM_MHA_HEADS, + num_kv_heads=_NUM_MHA_HEADS, + head_dim=head_dim, + batch_size=1, + past_kv_lens=(96,), + num_pages=4, + ) + ), + id=f"{str(dtype).removeprefix('torch.')}_h{head_dim}", + ) + for dtype in SUPPORTED_MODEL_DTYPES + for head_dim in SUPPORTED_MHA_HEAD_DIMS + ] + + [ + pytest.param( + PageSparseMhaScenario( + attention=MhaGenerationScenario( + num_heads=_NUM_MHA_HEADS, + num_kv_heads=_NUM_MHA_HEADS, + page_size=kv_page_size, + batch_size=1, + past_kv_lens=(max(96, 3 * kv_page_size),), + num_pages=math.ceil((max(96, 3 * kv_page_size) + 1) / kv_page_size), + ) + ), + id=f"kv_page_size_{kv_page_size}", + ) + for kv_page_size in TESTED_KV_PAGE_SIZES + if kv_page_size != 32 + ] + + [ + pytest.param( + PageSparseMhaScenario( + attention=MhaGenerationScenario( + num_heads=_NUM_MHA_HEADS, + num_kv_heads=_NUM_MHA_HEADS, + batch_size=1, + past_kv_lens=(96,), + num_pages=4, + ), + sparse_index_block_size=sparse_index_block_size, + ), + id=f"sparse_index_block_size_{sparse_index_block_size}", + ) + for sparse_index_block_size in (1, 2, 3, 5, 8, 16, 24, 32, 48) + ] + + [ + pytest.param( + PageSparseMhaScenario( + attention=MhaGenerationScenario( + num_heads=_NUM_MHA_HEADS, + num_kv_heads=_NUM_MHA_HEADS, + batch_size=1, + past_kv_lens=(96,), + num_pages=4, + ), + num_selected_sparse_blocks=num_selected_sparse_blocks, + ), + id=f"selected_sparse_blocks_{num_selected_sparse_blocks}", + ) + for num_selected_sparse_blocks in (1, 3) + ] + + [ + pytest.param( + PageSparseMhaScenario( + attention=MhaGenerationScenario( + num_heads=_NUM_MHA_HEADS, + num_kv_heads=_NUM_MHA_HEADS, + batch_size=2, + past_kv_lens=(32, 160), + num_pages=8, + ), + num_selected_sparse_blocks=3, + ), + id="batch2_var_kv_and_offsets", + ), + pytest.param( + PageSparseMhaScenario( + attention=MhaGenerationScenario( + num_heads=_NUM_MHA_HEADS, + num_kv_heads=_NUM_MHA_HEADS, + batch_size=1, + past_kv_lens=(95,), + num_pages=4, + ) + ), + id="non_page_aligned_kv_length", + ), + pytest.param( + PageSparseMhaScenario( + attention=MhaGenerationScenario( + num_heads=_NUM_MHA_HEADS, + num_kv_heads=_NUM_MHA_HEADS, + batch_size=1, + past_kv_lens=(96,), + query_len=4, + num_pages=4, + ) + ), + id="linear_3_draft_tokens", + ), + ] + + [ + pytest.param( + PageSparseMhaScenario( + attention=MhaGenerationScenario( + dtype=dtype, + kvcache_dtype=torch.float8_e4m3fn, + num_heads=_NUM_MHA_HEADS, + num_kv_heads=_NUM_MHA_HEADS, + head_dim=head_dim, + batch_size=1, + past_kv_lens=(96,), + num_pages=4, + fp8_output=fp8_output, + ) + ), + id=( + f"{str(dtype).removeprefix('torch.')}_h{head_dim}_fp8_kv_" + f"{'fp8_output' if fp8_output else 'model_output'}" + ), + ) + for dtype in SUPPORTED_MODEL_DTYPES + for head_dim in SUPPORTED_MHA_HEAD_DIMS + for fp8_output in (False, True) + if head_dim != 80 + ] + + [ + pytest.param( + PageSparseMhaScenario( + attention=MhaGenerationScenario( + num_heads=num_mha_heads, + num_kv_heads=num_mha_heads, + batch_size=1, + past_kv_lens=(96,), + num_pages=4, + ) + ), + id=f"mha_{num_mha_heads}_heads", + ) + for num_mha_heads in TESTED_MHA_HEAD_COUNTS + if num_mha_heads != _NUM_MHA_HEADS + ] +) + + +@pytest.mark.parametrize("scenario", _PAGE_SPARSE_MHA_CASES) +def test_generation_page_sparse_mha(scenario: PageSparseMhaScenario) -> None: + """Static page selections drive sparse MHA generation computation.""" + _run_page_sparse_mha(scenario) From 49493791d9004b284bfb328238615608d71decb3 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Wed, 2 Sep 2026 04:56:34 -0700 Subject: [PATCH 04/12] [None][test] Expand sparse attention coverage Add comprehensive token- and block-sparse MQA/GQA coverage and reorganize generic sparse attention tests by capability and algorithm ownership. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../test_lists/test-db/l0_dgx_b200.yml | 12 +- .../sparse/_sparse_mha_test_utils.py | 328 ---------- .../sparse/{ => dsa}/test_cpp_custom_ops.py | 0 .../attention/sparse/kernel/__init__.py | 0 .../attention/sparse/kernel/test_flash_mla.py | 110 ---- .../_torch/attention/sparse/msa/__init__.py | 2 + .../test_msa_backend.py} | 9 +- .../attention/sparse/test_sparse_mha.py | 319 +++++++++- .../sparse/test_sparse_mla_forward.py | 64 ++ .../attention/sparse/test_sparse_mqa_gqa.py | 568 +++++++++++++++--- 10 files changed, 876 insertions(+), 536 deletions(-) delete mode 100644 tests/unittest/_torch/attention/sparse/_sparse_mha_test_utils.py rename tests/unittest/_torch/attention/sparse/{ => dsa}/test_cpp_custom_ops.py (100%) delete mode 100644 tests/unittest/_torch/attention/sparse/kernel/__init__.py delete mode 100644 tests/unittest/_torch/attention/sparse/kernel/test_flash_mla.py create mode 100644 tests/unittest/_torch/attention/sparse/msa/__init__.py rename tests/unittest/_torch/attention/sparse/{test_minimax_m3_msa_backend.py => msa/test_msa_backend.py} (98%) diff --git a/tests/integration/test_lists/test-db/l0_dgx_b200.yml b/tests/integration/test_lists/test-db/l0_dgx_b200.yml index 03aa1aefa31e..c6c45d2ed09c 100644 --- a/tests/integration/test_lists/test-db/l0_dgx_b200.yml +++ b/tests/integration/test_lists/test-db/l0_dgx_b200.yml @@ -63,12 +63,12 @@ l0_dgx_b200: # ------------- NVBug 6025177: trtllm-serve cross-request KV contamination (OpenAI) --------------- - test_e2e.py::test_openai_kv_cache_contamination TIMEOUT (120) # ------------- DSA FP4 indexer (Blackwell-only) --------------- - - unittest/_torch/attention/sparse/test_cpp_custom_ops.py::test_fused_cat_fp4_matches_deepgemm - - unittest/_torch/attention/sparse/test_cpp_custom_ops.py::test_fused_cat_fp4_shape_dispatch - - unittest/_torch/attention/sparse/test_cpp_custom_ops.py::test_fused_cat_fp4_noncontiguous_split - - unittest/_torch/attention/sparse/test_cpp_custom_ops.py::test_fused_cat_fp4_dsv32_prefill_shape - - unittest/_torch/attention/sparse/test_cpp_custom_ops.py::test_cute_dsl_fp8_indexer_q_gemm_rope_fp4_matches_unfused - - unittest/_torch/attention/sparse/test_cpp_custom_ops.py::test_indexer_k_cache_gather_contiguous_fp4 + - unittest/_torch/attention/sparse/dsa/test_cpp_custom_ops.py::test_fused_cat_fp4_matches_deepgemm + - unittest/_torch/attention/sparse/dsa/test_cpp_custom_ops.py::test_fused_cat_fp4_shape_dispatch + - unittest/_torch/attention/sparse/dsa/test_cpp_custom_ops.py::test_fused_cat_fp4_noncontiguous_split + - unittest/_torch/attention/sparse/dsa/test_cpp_custom_ops.py::test_fused_cat_fp4_dsv32_prefill_shape + - unittest/_torch/attention/sparse/dsa/test_cpp_custom_ops.py::test_cute_dsl_fp8_indexer_q_gemm_rope_fp4_matches_unfused + - unittest/_torch/attention/sparse/dsa/test_cpp_custom_ops.py::test_indexer_k_cache_gather_contiguous_fp4 - unittest/_torch/attention/sparse/dsa/test_dsa_fp4_indexer.py - condition: ranges: diff --git a/tests/unittest/_torch/attention/sparse/_sparse_mha_test_utils.py b/tests/unittest/_torch/attention/sparse/_sparse_mha_test_utils.py deleted file mode 100644 index 73b98ecad00b..000000000000 --- a/tests/unittest/_torch/attention/sparse/_sparse_mha_test_utils.py +++ /dev/null @@ -1,328 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Input, cache, and reference helpers for page-sparse MHA tests.""" - -import math -from contextlib import ExitStack -from dataclasses import dataclass -from typing import Optional, Tuple - -import torch - -import tensorrt_llm -from tensorrt_llm._torch.attention_backend.interface import AttentionRuntimeFeatures -from tensorrt_llm._torch.attention_backend.trtllm import ( - TrtllmAttentionMetadata, - generate_spec_decoding_packed_mask, - generate_spec_decoding_position_offsets, -) -from tensorrt_llm._torch.metadata import KVCacheParams -from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager -from tensorrt_llm._utils import str_dtype_to_binding, torch_dtype_to_str -from tensorrt_llm.bindings.executor import KvCacheConfig -from tensorrt_llm.mapping import Mapping -from tensorrt_llm.models.modeling_utils import QuantConfig -from tensorrt_llm.quantization.mode import QuantAlgo - -SUPPORTED_MODEL_DTYPES = (torch.bfloat16, torch.float16) -SUPPORTED_KV_CACHE_DTYPES = (*SUPPORTED_MODEL_DTYPES, torch.float8_e4m3fn) -SUPPORTED_MHA_HEAD_DIMS = (64, 80, 128, 256) - - -@dataclass(kw_only=True, frozen=True) -class MhaGenerationScenario: - """Generation inputs for page-sparse MHA computation.""" - - dtype: torch.dtype = torch.bfloat16 - kvcache_dtype: torch.dtype = torch.bfloat16 - num_layers: int = 1 - num_heads: int = 8 - num_kv_heads: int = 8 - head_dim: int = 128 - page_size: int = 32 - num_pages: int = 4 - batch_size: int = 1 - past_kv_lens: Tuple[int, ...] = (96,) - query_len: int = 1 - fp8_output: bool = False - - def __post_init__(self) -> None: - if self.dtype not in SUPPORTED_MODEL_DTYPES: - raise ValueError("Model QKV dtype must be BF16 or FP16") - if self.kvcache_dtype not in SUPPORTED_KV_CACHE_DTYPES: - raise ValueError("KV-cache dtype must be BF16, FP16, or E4M3 FP8") - if self.kvcache_dtype != torch.float8_e4m3fn and self.kvcache_dtype != self.dtype: - raise ValueError("A non-FP8 KV-cache dtype must match the model QKV dtype") - if self.num_heads <= 0 or self.num_heads != self.num_kv_heads: - raise ValueError("Page-sparse MHA requires equal positive Q and KV head counts") - if self.head_dim not in SUPPORTED_MHA_HEAD_DIMS: - raise ValueError(f"head_dim must be one of {SUPPORTED_MHA_HEAD_DIMS}") - if self.page_size < 8 or self.page_size & (self.page_size - 1): - raise ValueError("page_size must be a power of two and at least 8") - if len(self.past_kv_lens) != self.batch_size: - raise ValueError( - f"past_kv_lens length {len(self.past_kv_lens)} must match " - f"batch_size {self.batch_size}" - ) - if self.query_len < 1: - raise ValueError("query_len must be positive") - if self.fp8_output and self.kvcache_dtype != torch.float8_e4m3fn: - raise ValueError("FP8 output testing requires an FP8 KV cache") - for past_kv_len in self.past_kv_lens: - required_pages = math.ceil((past_kv_len + self.query_len) / self.page_size) - if required_pages > self.num_pages: - raise ValueError("num_pages does not cover the request KV length") - - @property - def nnz_q(self) -> int: - return self.batch_size * self.query_len - - @property - def max_query_len(self) -> int: - return self.query_len - - @property - def has_draft_tokens(self) -> bool: - return self.query_len > 1 - - @property - def kv_pool_num_pages(self) -> int: - return self.batch_size * self.num_pages - - -@dataclass(kw_only=True) -class MhaGenerationInputs: - """Generation tensors plus their populated paged KV cache.""" - - q: torch.Tensor - k_new: torch.Tensor - v_new: torch.Tensor - kv_cache_manager: KVCacheManager - request_ids: list[int] - metadata: TrtllmAttentionMetadata - - @property - def fused_qkv(self) -> torch.Tensor: - return torch.cat([self.q, self.k_new, self.v_new], dim=1) - - -def quant_config(scenario: MhaGenerationScenario) -> Optional[QuantConfig]: - """Build the quantization settings used by the attention backend.""" - if scenario.fp8_output: - return QuantConfig( - quant_algo=QuantAlgo.FP8, - kv_cache_quant_algo=QuantAlgo.FP8, - ) - if scenario.kvcache_dtype == torch.float8_e4m3fn: - return QuantConfig(kv_cache_quant_algo=QuantAlgo.FP8) - return None - - -def fp8_qdq(tensor: torch.Tensor) -> torch.Tensor: - """Apply unit-scale E4M3 quantize-dequantize to a reference tensor.""" - return tensor.to(torch.float8_e4m3fn).to(tensor.dtype) - - -def _create_kv_cache_manager( - scenario: MhaGenerationScenario, - kv_cache: torch.Tensor, -) -> KVCacheManager: - kv_cache_config = KvCacheConfig(max_tokens=scenario.kv_pool_num_pages * scenario.page_size) - manager = KVCacheManager( - kv_cache_config, - tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF, - num_layers=scenario.num_layers, - num_kv_heads=scenario.num_kv_heads, - head_dim=scenario.head_dim, - tokens_per_block=scenario.page_size, - max_seq_len=scenario.kv_pool_num_pages * scenario.page_size, - max_batch_size=scenario.batch_size, - mapping=Mapping(world_size=1, tp_size=1, rank=0), - dtype=str_dtype_to_binding(torch_dtype_to_str(scenario.kvcache_dtype)), - ) - for layer_idx in range(scenario.num_layers): - manager.get_buffers(layer_idx, kv_layout="HND").copy_(kv_cache[layer_idx]) - return manager - - -def create_generation_inputs(scenario: MhaGenerationScenario) -> MhaGenerationInputs: - """Create generation inputs and release the cache if construction fails.""" - device = torch.device("cuda") - torch.manual_seed(42) - q = torch.randn( - scenario.nnz_q, - scenario.num_heads * scenario.head_dim, - device=device, - dtype=scenario.dtype, - ) - k_new = torch.randn( - scenario.nnz_q, - scenario.num_kv_heads * scenario.head_dim, - device=device, - dtype=scenario.dtype, - ) - v_new = torch.randn_like(k_new) - kv_cache = torch.randn( - scenario.num_layers, - scenario.kv_pool_num_pages, - 2, - scenario.num_kv_heads, - scenario.page_size, - scenario.head_dim, - device=device, - dtype=scenario.dtype, - ).to(scenario.kvcache_dtype) - - with ExitStack() as cleanup: - kv_cache_manager = _create_kv_cache_manager(scenario, kv_cache) - cleanup.callback(kv_cache_manager.shutdown) - request_ids = list(range(scenario.batch_size)) - token_nums = [past_kv_len + scenario.query_len for past_kv_len in scenario.past_kv_lens] - kv_cache_manager.add_dummy_requests(request_ids, token_nums) - metadata = TrtllmAttentionMetadata( - num_contexts=0, - kv_cache_params=KVCacheParams( - use_cache=True, - num_cached_tokens_per_seq=list(scenario.past_kv_lens), - ), - seq_lens=torch.full((scenario.batch_size,), scenario.query_len, dtype=torch.int32), - max_num_requests=scenario.batch_size, - max_num_tokens=scenario.nnz_q, - kv_cache_manager=kv_cache_manager, - request_ids=request_ids, - prompt_lens=list(scenario.past_kv_lens), - num_heads_per_kv=1, - runtime_features=AttentionRuntimeFeatures( - has_speculative_draft_tokens=scenario.has_draft_tokens - ), - is_spec_decoding_enabled=scenario.has_draft_tokens, - use_spec_decoding=scenario.has_draft_tokens, - is_spec_dec_tree=False, - max_total_draft_tokens=( - scenario.max_query_len - 1 if scenario.has_draft_tokens else None - ), - ) - if scenario.has_draft_tokens: - draft_len = scenario.max_query_len - 1 - metadata.spec_decoding_position_offsets = generate_spec_decoding_position_offsets( - scenario.batch_size, draft_len - ) - metadata.spec_decoding_packed_mask = generate_spec_decoding_packed_mask( - scenario.batch_size, draft_len - ) - metadata.spec_decoding_generation_lengths = torch.tensor( - [scenario.query_len] * scenario.batch_size, - dtype=torch.int32, - device=device, - ) - metadata.update_position_offsets_for_cpp(scenario.max_query_len) - metadata.spec_decoding_param_prepare_for_blackwell() - metadata.prepare() - cleanup.pop_all() - return MhaGenerationInputs( - q=q, - k_new=k_new, - v_new=v_new, - kv_cache_manager=kv_cache_manager, - request_ids=request_ids, - metadata=metadata, - ) - - -def read_paged_kv_cache( - inputs: MhaGenerationInputs, - scenario: MhaGenerationScenario, -) -> list[tuple[torch.Tensor, torch.Tensor]]: - """Materialize paged history using page-wise copies.""" - kv_buffer = inputs.kv_cache_manager.get_buffers(0, kv_layout="HND") - kv_caches = [] - for request_id, num_tokens in zip(inputs.request_ids, scenario.past_kv_lens, strict=True): - block_ids = inputs.kv_cache_manager.get_block_ids_per_seq([request_id])[0] - k_cache = torch.empty( - num_tokens, - scenario.num_kv_heads, - scenario.head_dim, - device=kv_buffer.device, - dtype=scenario.dtype, - ) - v_cache = torch.empty_like(k_cache) - for local_page_idx, block_id in enumerate(block_ids): - token_start = local_page_idx * scenario.page_size - token_end = min(token_start + scenario.page_size, num_tokens) - if token_start >= token_end: - break - num_page_tokens = token_end - token_start - k_cache[token_start:token_end] = ( - kv_buffer[block_id, 0, :, :num_page_tokens, :].transpose(0, 1).to(scenario.dtype) - ) - v_cache[token_start:token_end] = ( - kv_buffer[block_id, 1, :, :num_page_tokens, :].transpose(0, 1).to(scenario.dtype) - ) - kv_caches.append((k_cache, v_cache)) - return kv_caches - - -def reference_generation_attention( - q: torch.Tensor, - kv_caches: list[tuple[torch.Tensor, torch.Tensor]], - k_new: torch.Tensor, - v_new: torch.Tensor, - sparse_attn_indices: torch.Tensor, - scenario: MhaGenerationScenario, -) -> torch.Tensor: - """Compute page-sparse MHA from equivalent request-local token indices.""" - outputs = [] - query_offset = 0 - for request_idx in range(scenario.batch_size): - k_history, v_history = kv_caches[request_idx] - request_slice = slice(query_offset, query_offset + scenario.query_len) - k_full = torch.cat( - [ - k_history, - k_new[request_slice].view( - scenario.query_len, scenario.num_kv_heads, scenario.head_dim - ), - ], - dim=0, - ) - v_full = torch.cat( - [ - v_history, - v_new[request_slice].view( - scenario.query_len, scenario.num_kv_heads, scenario.head_dim - ), - ], - dim=0, - ) - for query_idx in range(scenario.query_len): - packed_query_idx = query_offset + query_idx - q_token = q[packed_query_idx].view(scenario.num_heads, scenario.head_dim) - head_outputs = [] - for head_idx in range(scenario.num_heads): - token_indices = sparse_attn_indices[head_idx, packed_query_idx] - valid_indices = token_indices[token_indices >= 0].long() - k_sparse = k_full[valid_indices, head_idx, :] - v_sparse = v_full[valid_indices, head_idx, :] - attention_scores = torch.matmul(q_token[head_idx], k_sparse.T) / math.sqrt( - scenario.head_dim - ) - attention_probs = torch.nn.functional.softmax( - attention_scores, dim=-1, dtype=torch.float32 - ).to(scenario.dtype) - head_outputs.append(torch.matmul(attention_probs, v_sparse)) - outputs.append(torch.cat(head_outputs, dim=0)) - query_offset += scenario.query_len - return torch.stack(outputs, dim=0) diff --git a/tests/unittest/_torch/attention/sparse/test_cpp_custom_ops.py b/tests/unittest/_torch/attention/sparse/dsa/test_cpp_custom_ops.py similarity index 100% rename from tests/unittest/_torch/attention/sparse/test_cpp_custom_ops.py rename to tests/unittest/_torch/attention/sparse/dsa/test_cpp_custom_ops.py diff --git a/tests/unittest/_torch/attention/sparse/kernel/__init__.py b/tests/unittest/_torch/attention/sparse/kernel/__init__.py deleted file mode 100644 index e69de29bb2d1..000000000000 diff --git a/tests/unittest/_torch/attention/sparse/kernel/test_flash_mla.py b/tests/unittest/_torch/attention/sparse/kernel/test_flash_mla.py deleted file mode 100644 index a9686877b2b9..000000000000 --- a/tests/unittest/_torch/attention/sparse/kernel/test_flash_mla.py +++ /dev/null @@ -1,110 +0,0 @@ -""" -Test basic sparse MLA forward pass to verify kernels are working correctly. -""" - -import math - -import pytest -import torch -from utils.util import getSMVersion - - -def has_flash_mla(): - """Check if FlashMLA module is available.""" - try: - from tensorrt_llm.flash_mla import flash_mla_sparse_fwd # noqa: F401 - - return True - except ImportError: - return False - - -@pytest.mark.skipif(not has_flash_mla(), reason="FlashMLA not available") -@pytest.mark.skipif( - getSMVersion() < 90, reason="FlashMLA requires SM90 (Hopper) or SM100 (Blackwell)" -) -@pytest.mark.parametrize( - "seq_len_q,seq_len_kv,topk", - [ - (62, 128, 128), # Small test case - (128, 256, 128), # Medium - (128, 512, 256), # Larger topk - ], -) -def test_flash_mla_sparse_fwd(seq_len_q, seq_len_kv, topk): - """ - Test FlashMLA sparse attention forward kernel. - - Args: - seq_len_q: Query sequence length - seq_len_kv: Key-Value sequence length - topk: Number of tokens to attend to (must be multiple of 128) - """ - from tensorrt_llm.flash_mla import flash_mla_sparse_fwd - - torch.manual_seed(42) - torch.cuda.manual_seed(42) - - # Fixed parameters matching FlashMLA's kernel requirements - # These are hardware-specific, not arbitrary choices - batch_size = 1 - num_heads_q = 128 # Fixed requirement for kernel (B_H parameter) - num_heads_kv = 1 # MLA uses 1 KV head - head_dim_qk = 576 # DeepSeek MLA standard - head_dim_v = 512 # Fixed requirement (only 512 supported) - - # Generate test inputs - # Q: [b, s_q, h_q, d_qk] - q = ( - torch.randn( - batch_size, seq_len_q, num_heads_q, head_dim_qk, dtype=torch.bfloat16, device="cuda" - ) - / 10.0 - ) - q.clamp_(-10, 10) - - # KV: [b, s_kv, h_kv, d_qk] - kv = ( - torch.randn( - batch_size, seq_len_kv, num_heads_kv, head_dim_qk, dtype=torch.bfloat16, device="cuda" - ) - / 10.0 - ) - kv.clamp_(-10, 10) - - # Indices: [b, s_q, h_kv, topk] - which KV tokens each Q attends to - indices = torch.randint( - 0, seq_len_kv, (batch_size, seq_len_q, num_heads_kv, topk), dtype=torch.int32, device="cuda" - ) - - softmax_scale = 1.0 / math.sqrt(head_dim_qk) - - # Run FlashMLA sparse forward (API expects no batch dimension) - output, max_logits, lse = flash_mla_sparse_fwd( - q.squeeze(0), # [s_q, h_q, d_qk] - kv.squeeze(0), # [s_kv, h_kv, d_qk] - indices.squeeze(0), # [s_q, h_kv, topk] - sm_scale=softmax_scale, - ) - - # Validate outputs - assert output.shape == (seq_len_q, num_heads_q, head_dim_v), ( - f"Output shape mismatch: expected [{seq_len_q}, {num_heads_q}, {head_dim_v}], got {output.shape}" - ) - assert output.dtype == torch.bfloat16, ( - f"Output dtype mismatch: expected torch.bfloat16, got {output.dtype}" - ) - - assert max_logits.shape == (seq_len_q, num_heads_q), ( - f"Max logits shape mismatch: got {max_logits.shape}" - ) - assert max_logits.dtype == torch.float32, f"Max logits dtype mismatch: got {max_logits.dtype}" - - assert lse.shape == (seq_len_q, num_heads_q), f"LSE shape mismatch: got {lse.shape}" - assert lse.dtype == torch.float32, f"LSE dtype mismatch: got {lse.dtype}" - - # Numerical validity checks - assert not torch.isnan(output).any(), "Output contains NaN" - assert not torch.isinf(output).any(), "Output contains Inf" - assert not torch.isnan(max_logits).any(), "Max logits contains NaN" - assert not torch.isnan(lse).any(), "LSE contains NaN" diff --git a/tests/unittest/_torch/attention/sparse/msa/__init__.py b/tests/unittest/_torch/attention/sparse/msa/__init__.py new file mode 100644 index 000000000000..52a7a9daf028 --- /dev/null +++ b/tests/unittest/_torch/attention/sparse/msa/__init__.py @@ -0,0 +1,2 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 diff --git a/tests/unittest/_torch/attention/sparse/test_minimax_m3_msa_backend.py b/tests/unittest/_torch/attention/sparse/msa/test_msa_backend.py similarity index 98% rename from tests/unittest/_torch/attention/sparse/test_minimax_m3_msa_backend.py rename to tests/unittest/_torch/attention/sparse/msa/test_msa_backend.py index 038e2ceea67c..7abd4f07f9af 100644 --- a/tests/unittest/_torch/attention/sparse/test_minimax_m3_msa_backend.py +++ b/tests/unittest/_torch/attention/sparse/msa/test_msa_backend.py @@ -1,10 +1,11 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Structural tests for the MiniMax-M3 MSA sparse attention backend. +"""MiniMax-M3 integration tests for its MSA sparse attention backend. -These validate backend selection, decode scratch-buffer sizing, and the paged -HND view contract passed to the packaged MSA kernel. Numerical parity against -the Triton reference is covered by the SM100 integration accuracy test. +These validate MiniMax-M3 backend selection, indexer/cache integration, decode +scratch-buffer sizing, and the paged HND contract passed to the packaged MSA +kernel. Generic block-sparse MQA/GQA numerical coverage lives in the parent +``test_sparse_mqa_gqa.py`` module. """ import sys diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_mha.py b/tests/unittest/_torch/attention/sparse/test_sparse_mha.py index 58dd1b361ff5..9198efc034d2 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_mha.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_mha.py @@ -24,6 +24,7 @@ """ import math +from contextlib import ExitStack from dataclasses import dataclass from typing import Optional, Tuple @@ -31,25 +32,33 @@ import torch from utils.util import getSMVersion -from tensorrt_llm._torch.attention.backends.interface import AttentionForwardArgs +import tensorrt_llm +from tensorrt_llm._torch.attention.backends.interface import ( + AttentionForwardArgs, + AttentionRuntimeFeatures, +) from tensorrt_llm._torch.attention.backends.sparse.params import SparseParams -from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttention - -from ._sparse_mha_test_utils import ( - SUPPORTED_MHA_HEAD_DIMS, - SUPPORTED_MODEL_DTYPES, - MhaGenerationScenario, - create_generation_inputs, - fp8_qdq, - quant_config, - read_paged_kv_cache, - reference_generation_attention, +from tensorrt_llm._torch.attention.backends.trtllm import ( + TrtllmAttention, + TrtllmAttentionMetadata, + generate_spec_decoding_packed_mask, + generate_spec_decoding_position_offsets, ) +from tensorrt_llm._torch.metadata import KVCacheParams +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm._utils import str_dtype_to_binding, torch_dtype_to_str +from tensorrt_llm.bindings.executor import KvCacheConfig +from tensorrt_llm.mapping import Mapping +from tensorrt_llm.models.modeling_utils import QuantConfig +from tensorrt_llm.quantization.mode import QuantAlgo ATOL = 2e-2 RTOL = 2e-2 FP8_ATOL = 8e-2 FP8_RTOL = 4e-2 +SUPPORTED_MODEL_DTYPES = (torch.bfloat16, torch.float16) +SUPPORTED_KV_CACHE_DTYPES = (*SUPPORTED_MODEL_DTYPES, torch.float8_e4m3fn) +SUPPORTED_MHA_HEAD_DIMS = (64, 80, 128, 256) TESTED_MHA_HEAD_COUNTS = (1, 2, 3, 4, 8, 16, 24, 32, 48, 64, 96, 128) TESTED_KV_PAGE_SIZES = (8, 16, 32, 64, 128, 256, 512) SUPPORTED_SM_VERSIONS = (100, 103) @@ -90,6 +99,292 @@ def _force_trtllm_gen_fallback(monkeypatch: pytest.MonkeyPatch) -> None: # Attention semantics Causal self-attention +@dataclass(kw_only=True, frozen=True) +class MhaGenerationScenario: + """Generation inputs for page-sparse MHA computation.""" + + dtype: torch.dtype = torch.bfloat16 + kvcache_dtype: torch.dtype = torch.bfloat16 + num_layers: int = 1 + num_heads: int = 8 + num_kv_heads: int = 8 + head_dim: int = 128 + page_size: int = 32 + num_pages: int = 4 + batch_size: int = 1 + past_kv_lens: Tuple[int, ...] = (96,) + query_len: int = 1 + fp8_output: bool = False + + def __post_init__(self) -> None: + if self.dtype not in SUPPORTED_MODEL_DTYPES: + raise ValueError("Model QKV dtype must be BF16 or FP16") + if self.kvcache_dtype not in SUPPORTED_KV_CACHE_DTYPES: + raise ValueError("KV-cache dtype must be BF16, FP16, or E4M3 FP8") + if self.kvcache_dtype != torch.float8_e4m3fn and self.kvcache_dtype != self.dtype: + raise ValueError("A non-FP8 KV-cache dtype must match the model QKV dtype") + if self.num_heads <= 0 or self.num_heads != self.num_kv_heads: + raise ValueError("Page-sparse MHA requires equal positive Q and KV head counts") + if self.head_dim not in SUPPORTED_MHA_HEAD_DIMS: + raise ValueError(f"head_dim must be one of {SUPPORTED_MHA_HEAD_DIMS}") + if self.page_size < 8 or self.page_size & (self.page_size - 1): + raise ValueError("page_size must be a power of two and at least 8") + if len(self.past_kv_lens) != self.batch_size: + raise ValueError( + f"past_kv_lens length {len(self.past_kv_lens)} must match " + f"batch_size {self.batch_size}" + ) + if self.query_len < 1: + raise ValueError("query_len must be positive") + if self.fp8_output and self.kvcache_dtype != torch.float8_e4m3fn: + raise ValueError("FP8 output testing requires an FP8 KV cache") + for past_kv_len in self.past_kv_lens: + required_pages = math.ceil((past_kv_len + self.query_len) / self.page_size) + if required_pages > self.num_pages: + raise ValueError("num_pages does not cover the request KV length") + + @property + def nnz_q(self) -> int: + return self.batch_size * self.query_len + + @property + def max_query_len(self) -> int: + return self.query_len + + @property + def has_draft_tokens(self) -> bool: + return self.query_len > 1 + + @property + def kv_pool_num_pages(self) -> int: + return self.batch_size * self.num_pages + + +@dataclass(kw_only=True) +class MhaGenerationInputs: + """Generation tensors plus their populated paged KV cache.""" + + q: torch.Tensor + k_new: torch.Tensor + v_new: torch.Tensor + kv_cache_manager: KVCacheManager + request_ids: list[int] + metadata: TrtllmAttentionMetadata + + @property + def fused_qkv(self) -> torch.Tensor: + return torch.cat([self.q, self.k_new, self.v_new], dim=1) + + +def quant_config(scenario: MhaGenerationScenario) -> Optional[QuantConfig]: + """Build the quantization settings used by the attention backend.""" + if scenario.fp8_output: + return QuantConfig( + quant_algo=QuantAlgo.FP8, + kv_cache_quant_algo=QuantAlgo.FP8, + ) + if scenario.kvcache_dtype == torch.float8_e4m3fn: + return QuantConfig(kv_cache_quant_algo=QuantAlgo.FP8) + return None + + +def fp8_qdq(tensor: torch.Tensor) -> torch.Tensor: + """Apply unit-scale E4M3 quantize-dequantize to a reference tensor.""" + return tensor.to(torch.float8_e4m3fn).to(tensor.dtype) + + +def _create_kv_cache_manager( + scenario: MhaGenerationScenario, + kv_cache: torch.Tensor, +) -> KVCacheManager: + kv_cache_config = KvCacheConfig(max_tokens=scenario.kv_pool_num_pages * scenario.page_size) + manager = KVCacheManager( + kv_cache_config, + tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF, + num_layers=scenario.num_layers, + num_kv_heads=scenario.num_kv_heads, + head_dim=scenario.head_dim, + tokens_per_block=scenario.page_size, + max_seq_len=scenario.kv_pool_num_pages * scenario.page_size, + max_batch_size=scenario.batch_size, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=str_dtype_to_binding(torch_dtype_to_str(scenario.kvcache_dtype)), + ) + for layer_idx in range(scenario.num_layers): + manager.get_buffers(layer_idx, kv_layout="HND").copy_(kv_cache[layer_idx]) + return manager + + +def create_generation_inputs(scenario: MhaGenerationScenario) -> MhaGenerationInputs: + """Create generation inputs and release the cache if construction fails.""" + device = torch.device("cuda") + torch.manual_seed(42) + q = torch.randn( + scenario.nnz_q, + scenario.num_heads * scenario.head_dim, + device=device, + dtype=scenario.dtype, + ) + k_new = torch.randn( + scenario.nnz_q, + scenario.num_kv_heads * scenario.head_dim, + device=device, + dtype=scenario.dtype, + ) + v_new = torch.randn_like(k_new) + kv_cache = torch.randn( + scenario.num_layers, + scenario.kv_pool_num_pages, + 2, + scenario.num_kv_heads, + scenario.page_size, + scenario.head_dim, + device=device, + dtype=scenario.dtype, + ).to(scenario.kvcache_dtype) + + with ExitStack() as cleanup: + kv_cache_manager = _create_kv_cache_manager(scenario, kv_cache) + cleanup.callback(kv_cache_manager.shutdown) + request_ids = list(range(scenario.batch_size)) + token_nums = [past_kv_len + scenario.query_len for past_kv_len in scenario.past_kv_lens] + kv_cache_manager.add_dummy_requests(request_ids, token_nums) + metadata = TrtllmAttentionMetadata( + num_contexts=0, + kv_cache_params=KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=list(scenario.past_kv_lens), + ), + seq_lens=torch.full((scenario.batch_size,), scenario.query_len, dtype=torch.int32), + max_num_requests=scenario.batch_size, + max_num_tokens=scenario.nnz_q, + kv_cache_manager=kv_cache_manager, + request_ids=request_ids, + prompt_lens=list(scenario.past_kv_lens), + num_heads_per_kv=1, + runtime_features=AttentionRuntimeFeatures( + has_speculative_draft_tokens=scenario.has_draft_tokens + ), + is_spec_decoding_enabled=scenario.has_draft_tokens, + use_spec_decoding=scenario.has_draft_tokens, + is_spec_dec_tree=False, + max_total_draft_tokens=( + scenario.max_query_len - 1 if scenario.has_draft_tokens else None + ), + ) + if scenario.has_draft_tokens: + draft_len = scenario.max_query_len - 1 + metadata.spec_decoding_position_offsets = generate_spec_decoding_position_offsets( + scenario.batch_size, draft_len + ) + metadata.spec_decoding_packed_mask = generate_spec_decoding_packed_mask( + scenario.batch_size, draft_len + ) + metadata.spec_decoding_generation_lengths = torch.tensor( + [scenario.query_len] * scenario.batch_size, + dtype=torch.int32, + device=device, + ) + metadata.update_position_offsets_for_cpp(scenario.max_query_len) + metadata.spec_decoding_param_prepare_for_blackwell() + metadata.prepare() + cleanup.pop_all() + return MhaGenerationInputs( + q=q, + k_new=k_new, + v_new=v_new, + kv_cache_manager=kv_cache_manager, + request_ids=request_ids, + metadata=metadata, + ) + + +def read_paged_kv_cache( + inputs: MhaGenerationInputs, + scenario: MhaGenerationScenario, +) -> list[tuple[torch.Tensor, torch.Tensor]]: + """Materialize paged history using page-wise copies.""" + kv_buffer = inputs.kv_cache_manager.get_buffers(0, kv_layout="HND") + kv_caches = [] + for request_id, num_tokens in zip(inputs.request_ids, scenario.past_kv_lens, strict=True): + block_ids = inputs.kv_cache_manager.get_block_ids_per_seq([request_id])[0] + k_cache = torch.empty( + num_tokens, + scenario.num_kv_heads, + scenario.head_dim, + device=kv_buffer.device, + dtype=scenario.dtype, + ) + v_cache = torch.empty_like(k_cache) + for local_page_idx, block_id in enumerate(block_ids): + token_start = local_page_idx * scenario.page_size + token_end = min(token_start + scenario.page_size, num_tokens) + if token_start >= token_end: + break + num_page_tokens = token_end - token_start + k_cache[token_start:token_end] = ( + kv_buffer[block_id, 0, :, :num_page_tokens, :].transpose(0, 1).to(scenario.dtype) + ) + v_cache[token_start:token_end] = ( + kv_buffer[block_id, 1, :, :num_page_tokens, :].transpose(0, 1).to(scenario.dtype) + ) + kv_caches.append((k_cache, v_cache)) + return kv_caches + + +def reference_generation_attention( + q: torch.Tensor, + kv_caches: list[tuple[torch.Tensor, torch.Tensor]], + k_new: torch.Tensor, + v_new: torch.Tensor, + sparse_attn_indices: torch.Tensor, + scenario: MhaGenerationScenario, +) -> torch.Tensor: + """Compute page-sparse MHA from equivalent request-local token indices.""" + outputs = [] + query_offset = 0 + for request_idx in range(scenario.batch_size): + k_history, v_history = kv_caches[request_idx] + request_slice = slice(query_offset, query_offset + scenario.query_len) + k_full = torch.cat( + [ + k_history, + k_new[request_slice].view( + scenario.query_len, scenario.num_kv_heads, scenario.head_dim + ), + ], + dim=0, + ) + v_full = torch.cat( + [ + v_history, + v_new[request_slice].view( + scenario.query_len, scenario.num_kv_heads, scenario.head_dim + ), + ], + dim=0, + ) + for query_idx in range(scenario.query_len): + packed_query_idx = query_offset + query_idx + q_token = q[packed_query_idx].view(scenario.num_heads, scenario.head_dim) + head_outputs = [] + for head_idx in range(scenario.num_heads): + token_indices = sparse_attn_indices[head_idx, packed_query_idx] + valid_indices = token_indices[token_indices >= 0].long() + k_sparse = k_full[valid_indices, head_idx, :] + v_sparse = v_full[valid_indices, head_idx, :] + attention_scores = torch.matmul(q_token[head_idx], k_sparse.T) / math.sqrt( + scenario.head_dim + ) + attention_probs = torch.nn.functional.softmax( + attention_scores, dim=-1, dtype=torch.float32 + ).to(scenario.dtype) + head_outputs.append(torch.matmul(attention_probs, v_sparse)) + outputs.append(torch.cat(head_outputs, dim=0)) + query_offset += scenario.query_len + return torch.stack(outputs, dim=0) + + @dataclass(kw_only=True, frozen=True) class PageSparseMhaScenario: """Page-sparse MHA geometry layered on generation inputs.""" diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_mla_forward.py b/tests/unittest/_torch/attention/sparse/test_sparse_mla_forward.py index 1c80ff82db7c..642a41010074 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_mla_forward.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_mla_forward.py @@ -60,6 +60,8 @@ pytestmark = pytest.mark.threadleak(enabled=False) try: + from tensorrt_llm.flash_mla import flash_mla_sparse_fwd + HAS_FLASH_MLA = True except ImportError: HAS_FLASH_MLA = False @@ -179,6 +181,68 @@ def _kv_cache_dtypes_for_algo(sparse_attn_algo: str) -> list[str]: "deepseek_v4", id=f"deepseek_v4-{DSV4_KV_CACHE_DTYPES[0]}-large_mixed_deepseek_v4")) +# FlashMLA sparse MLA kernel contract. + + +@pytest.mark.skipif(not HAS_FLASH_MLA, reason="FlashMLA not available") +@pytest.mark.skipif(get_sm_version() < 90, + reason="FlashMLA requires SM90 (Hopper) or later") +@pytest.mark.parametrize( + "seq_len_q,seq_len_kv,topk", + [ + (62, 128, 128), + (128, 256, 128), + (128, 512, 256), + ], +) +def test_flash_mla_sparse_fwd(seq_len_q, seq_len_kv, topk): + """Validate the direct FlashMLA sparse-forward output contract.""" + torch.manual_seed(42) + torch.cuda.manual_seed(42) + + batch_size = 1 + num_heads_q = 128 + num_heads_kv = 1 + head_dim_qk = 576 + head_dim_v = 512 + + q = (torch.randn(batch_size, + seq_len_q, + num_heads_q, + head_dim_qk, + dtype=torch.bfloat16, + device="cuda") / 10.0) + q.clamp_(-10, 10) + kv = (torch.randn(batch_size, + seq_len_kv, + num_heads_kv, + head_dim_qk, + dtype=torch.bfloat16, + device="cuda") / 10.0) + kv.clamp_(-10, 10) + indices = torch.randint(0, + seq_len_kv, + (batch_size, seq_len_q, num_heads_kv, topk), + dtype=torch.int32, + device="cuda") + + output, max_logits, lse = flash_mla_sparse_fwd( + q.squeeze(0), + kv.squeeze(0), + indices.squeeze(0), + sm_scale=1.0 / math.sqrt(head_dim_qk), + ) + + assert output.shape == (seq_len_q, num_heads_q, head_dim_v) + assert output.dtype == torch.bfloat16 + assert max_logits.shape == (seq_len_q, num_heads_q) + assert max_logits.dtype == torch.float32 + assert lse.shape == (seq_len_q, num_heads_q) + assert lse.dtype == torch.float32 + assert torch.isfinite(output).all() + assert torch.isfinite(max_logits).all() + assert torch.isfinite(lse).all() + def apply_rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor, diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py index 30aee156046f..ddf65eeb146e 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py @@ -13,19 +13,21 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Regression tests for the internal token-sparse MQA/GQA kernel. +"""Executable examples and regression tests for sparse MQA/GQA compute. -These tests replace the model-specific sparse selector with static token-index -lists, then exercise the same backend path used by ``TrtllmAttention``: +The token-sparse tests replace the model-specific sparse selector with static +token-index lists, then exercise the same backend path used by +``TrtllmAttention``: 1. Build request-local ``int32`` token indices. 2. Translate them to paged KV-cache pool indices. 3. Return them from the sparse prediction hooks. 4. Call ``TrtllmAttention.forward`` and compare with a PyTorch reference. -This is an executable backend-integration example, not a public -``SparseAttentionConfig`` algorithm. Algorithm-independent sparse framework -tests remain in ``test_sparse_attention.py``. +The block-sparse tests pass static block-index lists to the MSA FMHA wrapper +and compare its paged MQA/GQA output with an independent PyTorch reference. +Algorithm-independent sparse framework tests remain in +``test_sparse_attention.py``. """ import math @@ -37,6 +39,7 @@ from utils.util import getSMVersion import tensorrt_llm +from tensorrt_llm._torch.attention_backend.fmha.msa_sparse_gqa import run_msa_sparse_gqa from tensorrt_llm._torch.attention_backend.interface import ( AttentionForwardArgs, AttentionRuntimeFeatures, @@ -44,6 +47,7 @@ from tensorrt_llm._torch.attention_backend.sparse.dsa.kernels import ( triton_convert_req_index_to_global_index, ) +from tensorrt_llm._torch.attention_backend.sparse.minimax_m3.msa_utils import msa_package_available from tensorrt_llm._torch.attention_backend.sparse.params import SparseParams from tensorrt_llm._torch.attention_backend.trtllm import ( TrtllmAttention, @@ -53,7 +57,7 @@ ) from tensorrt_llm._torch.metadata import KVCacheParams from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager -from tensorrt_llm._utils import is_sm_100f, str_dtype_to_binding, torch_dtype_to_str +from tensorrt_llm._utils import str_dtype_to_binding, torch_dtype_to_str from tensorrt_llm.bindings.executor import KvCacheConfig from tensorrt_llm.mapping import Mapping from tensorrt_llm.models.modeling_utils import QuantConfig @@ -66,17 +70,21 @@ SUPPORTED_DTYPES = (torch.bfloat16, torch.float16) SUPPORTED_KV_CACHE_DTYPES = (*SUPPORTED_DTYPES, torch.float8_e4m3fn) SUPPORTED_HEAD_DIMS = (64, 80, 128, 256) +SUPPORTED_FP8_HEAD_DIMS = (64, 128, 256) +SUPPORTED_TEST_PAGE_SIZES = (8, 16, 32, 64, 128, 256, 512) +TOKEN_SPARSE_PAGE_TEST_KV_LEN = 64 MAX_Q_HEADS_PER_KV_HEAD = 32 +SUPPORTED_SPARSE_MQA_GQA_SMS = (100, 103) pytestmark = pytest.mark.skipif( - not is_sm_100f(getSMVersion()), - reason="Sparse MQA/GQA requires an SM100-family GPU (SM100 or SM103)", + getSMVersion() not in SUPPORTED_SPARSE_MQA_GQA_SMS, + reason="Sparse MQA/GQA requires an SM100 or SM103 GPU", ) @pytest.fixture(autouse=True) def _force_trtllm_gen_fallback(monkeypatch: pytest.MonkeyPatch) -> None: - """Keep every test on the internal TRTLLM-Gen fallback path.""" + """Keep token-sparse tests on the internal TRTLLM-Gen fallback path.""" monkeypatch.setenv("TLLM_FMHA_LIBS", "fallback") @@ -187,6 +195,77 @@ def has_draft_tokens(self) -> bool: return self.max_query_len > 1 +@dataclass(kw_only=True, frozen=True) +class BlockSparseGqaScenario: + """Packed block-sparse MQA/GQA inputs for the MSA FMHA backend.""" + + q_lens: Tuple[int, ...] = (1,) + kv_lens: Tuple[int, ...] = (2176,) + num_q_heads: int = 16 + num_kv_heads: int = 1 + head_dim: int = 128 + page_size: int = 128 + topk: int = 16 + dtype: torch.dtype = torch.bfloat16 + qo_offsets: Optional[Tuple[int, ...]] = None + active_blocks: Optional[int] = None + shuffle_pages: bool = False + per_token_blocks: bool = False + + def __post_init__(self) -> None: + if len(self.q_lens) != len(self.kv_lens): + raise ValueError("q_lens and kv_lens must describe the same batch") + if self.qo_offsets is not None and len(self.qo_offsets) != len(self.q_lens): + raise ValueError("qo_offsets must describe the same batch as q_lens") + if self.num_q_heads % self.num_kv_heads != 0: + raise ValueError("num_q_heads must be divisible by num_kv_heads") + if self.q_heads_per_kv_head not in (2, 4, 8, 16): + raise ValueError("block-sparse MQA/GQA supports 2, 4, 8, or 16 Q heads per KV head") + if self.head_dim != 128 or self.page_size != 128: + raise ValueError("MSA block-sparse MQA/GQA requires head_dim=page_size=128") + if self.topk not in (4, 8, 16, 32): + raise ValueError("MSA block-sparse MQA/GQA supports Top-K 4, 8, 16, or 32") + if self.dtype not in (torch.bfloat16, torch.float8_e4m3fn): + raise ValueError("MSA block-sparse MQA/GQA supports BF16 or E4M3 FP8 Q/K/V") + active_blocks = self.topk if self.active_blocks is None else self.active_blocks + if not 0 < active_blocks <= self.topk: + raise ValueError("active_blocks must be in [1, topk]") + for q_len, kv_len, qo_offset in zip( + self.q_lens, + self.kv_lens, + self.causal_offsets, + strict=True, + ): + if q_len <= 0 or q_len > kv_len: + raise ValueError("each q_len must be positive and no larger than kv_len") + if qo_offset < 0 or qo_offset + q_len > kv_len: + raise ValueError("each qo_offset must place every query inside its KV sequence") + if kv_len % self.page_size: + raise ValueError("each kv_len must be a multiple of page_size") + if kv_len // self.page_size <= active_blocks: + raise ValueError("each request needs more KV pages than selected active blocks") + + @property + def q_heads_per_kv_head(self) -> int: + return self.num_q_heads // self.num_kv_heads + + @property + def total_q(self) -> int: + return sum(self.q_lens) + + @property + def selected_blocks(self) -> int: + return self.topk if self.active_blocks is None else self.active_blocks + + @property + def causal_offsets(self) -> Tuple[int, ...]: + if self.qo_offsets is not None: + return self.qo_offsets + return tuple( + kv_len - q_len for q_len, kv_len in zip(self.q_lens, self.kv_lens, strict=True) + ) + + class _SparseMqaGqaParams(SparseParams): """Token-granular parameters that select the internal MQA/GQA path.""" @@ -366,17 +445,48 @@ def test_prefill_sparse_kv_compaction(scenario: ContextScenario) -> None: # # Sparse MQA/GQA support matrix: # -# GPU architecture SM100 and SM103 -# Inference phase Packed prefill; single-token and linear draft decode -# Attention type MQA and GQA; num_heads % num_kv_heads == 0 -# Q heads per KV head <= 32; tests cover 2, 3, 4, 8, 16, 24, 31, and 32 -# Model QKV input BF16 or FP16 -# Kernel output BF16, FP16, or E4M3 FP8 -# KV-cache dtype BF16, FP16, or E4M3 FP8 -# Q/K/V head dimension Equal dimensions: 64, 80, 128, or 256 -# KV-cache layout Paged; page size is a power of two and at least 8 -# Sparse indices int32, token-granular, one list per KV head/query -# Sparse Top-K Positive multiple of 4; unused entries are -1 +# Values after "tested" are regression coverage, not narrower support constraints. +# +# Parameter Token-sparse Block-sparse +# Sparse block size 1 token; tested 128 tokens; tested +# GPU architecture SM100 and SM103; SM100 and SM103; +# tested on SM100 tested on SM100 +# Compute phase Packed prefill, single-token, Packed prefill, single-token, +# and linear draft decode; linear multi-query compute, +# tested q_len 1 and 4 and mixed batches; tested; the +# integrated MiniMax-M3 decode uses 1 +# Attention type MQA/GQA; Q heads divisible MQA/GQA; Q heads divisible +# by KV heads by KV heads +# Q heads per KV head <= 32; tested 2, 3, 4, 8, 2, 4, 8, or 16; tested all +# 16, 24, 31, and 32 integrated MiniMax-M3 uses 16 +# Q/KV head counts No additional discrete limit; No additional discrete kernel limit; +# tested Q={6,8,16,32,48,62,64}, tested Q={4,8,16,32}, KV={1,2} +# KV={1,2,4,8} +# Attention input dtype BF16 or FP16; tested both BF16 or E4M3 FP8; tested both +# Q/K/V input layout Fused QKV; tested Q [T,Hq,D], paged K/V +# [P,Hkv,128,D]; tested +# Kernel output BF16/FP16 for all head dims; BF16; tested +# E4M3 FP8 for 64/128/256; +# tested all dtype/dim pairs +# KV-cache dtype BF16/FP16 for all head dims; BF16 or E4M3 FP8; tested both +# E4M3 FP8 for 64/128/256; +# tested all dtype/dim pairs +# Q/K/V head dimension 64, 80, 128, or 256; 128; tested +# tested all +# KV-cache layout Paged; power-of-two page size Paged HND; page size 128; +# >= 8; tested 8 through 512 shuffled physical pages and +# strided outer page storage tested +# Sparse indices int32 physical token indices int32 request-local block indices +# per KV head/query; tested per KV head/query; per-token lists, +# -1 padding, and physical remap tested +# Sparse Top-K Positive multiple of 4; Prefill kernel accepts 4, 8, 16, +# tested 4, 32, 64, and 128 or 32 and tests cover all; the +# integrated MiniMax-M3 path uses 16 +# Attention semantics Causal; tested Causal with per-request Q offsets; +# bottom-right and custom offsets tested + + +# Token-granular sparse computation (block_size=1). _PREFILL_COMPUTE_CASES = [ @@ -610,60 +720,73 @@ def test_prefill_sparse_kv_compaction(scenario: ContextScenario) -> None: ] -_GENERATION_SUPPORT_CASES = [ - pytest.param( - GenerationScenario( - dtype=dtype, - kvcache_dtype=dtype, - num_heads=8, - num_kv_heads=num_kv_heads, - head_dim=head_dim, - batch_size=1, - past_kv_lens=(64,), - num_pages=4, - num_sparse_topk=32, - ), - id=( - f"support_{str(dtype).removeprefix('torch.')}_h{head_dim}_" - f"{'mqa' if num_kv_heads == 1 else 'gqa_2to1'}" - ), - ) - for dtype in SUPPORTED_DTYPES - for head_dim in SUPPORTED_HEAD_DIMS - for num_kv_heads in (1, 4) - # These two option combinations are already covered by correctness cases. - if not (dtype == torch.bfloat16 and head_dim == 128) -] + [ - pytest.param( - GenerationScenario( - dtype=torch.bfloat16, - kvcache_dtype=torch.float8_e4m3fn, - num_heads=8, - num_kv_heads=4, - head_dim=128, - batch_size=1, - past_kv_lens=(64,), - num_pages=8, - num_sparse_topk=32, - ), - id="support_bf16_io_fp8_kv_cache", - ), - pytest.param( - GenerationScenario( - dtype=torch.bfloat16, - kvcache_dtype=torch.float8_e4m3fn, - num_heads=8, - num_kv_heads=4, - head_dim=128, - batch_size=1, - past_kv_lens=(64,), - num_pages=8, - num_sparse_topk=32, - fp8_output=True, - ), - id="support_fp8_qkv_math_and_output", - ), -] +_GENERATION_SUPPORT_CASES = ( + [ + pytest.param( + GenerationScenario( + dtype=dtype, + kvcache_dtype=dtype, + num_heads=8, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + batch_size=1, + past_kv_lens=(64,), + num_pages=4, + num_sparse_topk=32, + ), + id=( + f"support_{str(dtype).removeprefix('torch.')}_h{head_dim}_" + f"{'mqa' if num_kv_heads == 1 else 'gqa_2to1'}" + ), + ) + for dtype in SUPPORTED_DTYPES + for head_dim in SUPPORTED_HEAD_DIMS + for num_kv_heads in (1, 4) + # These two option combinations are already covered by correctness cases. + if not (dtype == torch.bfloat16 and head_dim == 128) + ] + + [ + pytest.param( + GenerationScenario( + dtype=torch.bfloat16, + kvcache_dtype=torch.float8_e4m3fn, + num_heads=8, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + batch_size=1, + past_kv_lens=(64,), + num_pages=8, + num_sparse_topk=32, + fp8_output=fp8_output, + ), + id=( + f"support_fp8_kv_h{head_dim}_" + f"{'mqa' if num_kv_heads == 1 else 'gqa_2to1'}_" + f"{'fp8' if fp8_output else 'bf16'}_output" + ), + ) + for head_dim in SUPPORTED_FP8_HEAD_DIMS + for num_kv_heads in (1, 4) + for fp8_output in (False, True) + ] + + [ + pytest.param( + GenerationScenario( + page_size=page_size, + batch_size=1, + past_kv_lens=(TOKEN_SPARSE_PAGE_TEST_KV_LEN,), + num_pages=max( + 4, + (TOKEN_SPARSE_PAGE_TEST_KV_LEN + page_size) // page_size, + ), + num_sparse_topk=32, + ), + id=f"support_page_size_{page_size}", + ) + for page_size in SUPPORTED_TEST_PAGE_SIZES + if page_size not in (32, 64) + ] +) @pytest.mark.parametrize("scenario", _PREFILL_COMPUTE_CASES) @@ -779,6 +902,168 @@ def test_generation_sparse_mqa_gqa(scenario: GenerationScenario) -> None: inputs.kv_cache_manager.shutdown() +# Block-granular sparse computation (block_size=128). + + +_BLOCK_SPARSE_GQA_CASES = [ + pytest.param( + BlockSparseGqaScenario( + q_lens=(1, 1, 1, 1), + kv_lens=(2176, 2304, 2432, 2560), + ), + id="msa_mqa_single_token_varlen_batch4", + ), + pytest.param( + BlockSparseGqaScenario( + q_lens=(4, 4), + kv_lens=(2176, 2304), + num_q_heads=32, + num_kv_heads=2, + per_token_blocks=True, + ), + id="msa_gqa_linear_draft_tokens", + ), + *[ + pytest.param( + BlockSparseGqaScenario( + q_lens=(1,), + kv_lens=(2176,), + num_q_heads=2 * q_heads_per_kv_head, + num_kv_heads=2, + ), + id=f"msa_gqa_single_token_{q_heads_per_kv_head}q_per_kv", + ) + for q_heads_per_kv_head in (2, 4, 8) + ], + pytest.param( + BlockSparseGqaScenario( + q_lens=(1, 4, 33), + kv_lens=(1152, 1280, 1408), + num_q_heads=8, + num_kv_heads=2, + active_blocks=3, + shuffle_pages=True, + per_token_blocks=True, + ), + id="msa_gqa_mixed_varlen_shuffled_pages_padded_indices", + ), + pytest.param( + BlockSparseGqaScenario( + q_lens=(33,), + kv_lens=(640,), + num_q_heads=4, + num_kv_heads=2, + topk=4, + ), + id="msa_gqa_2q_per_kv_topk4", + ), + pytest.param( + BlockSparseGqaScenario( + q_lens=(33,), + kv_lens=(1152,), + num_q_heads=8, + num_kv_heads=2, + topk=8, + ), + id="msa_gqa_4q_per_kv_topk8", + ), + pytest.param( + BlockSparseGqaScenario( + q_lens=(33,), + kv_lens=(2176,), + num_q_heads=16, + num_kv_heads=2, + topk=16, + ), + id="msa_gqa_8q_per_kv_topk16", + ), + pytest.param( + BlockSparseGqaScenario( + q_lens=(33,), + kv_lens=(4224,), + num_q_heads=32, + num_kv_heads=2, + topk=32, + ), + id="msa_gqa_16q_per_kv_topk32", + ), + pytest.param( + BlockSparseGqaScenario( + q_lens=(33,), + kv_lens=(2176,), + num_q_heads=16, + num_kv_heads=1, + dtype=torch.float8_e4m3fn, + per_token_blocks=True, + ), + id="msa_mqa_fp8_qkv_bf16_output", + ), + pytest.param( + BlockSparseGqaScenario( + q_lens=(33, 40), + kv_lens=(2176, 2304), + qo_offsets=(512, 1024), + num_q_heads=16, + num_kv_heads=2, + ), + id="msa_gqa_custom_per_request_q_offsets", + ), + pytest.param( + BlockSparseGqaScenario( + q_lens=(1, 1), + kv_lens=(2176, 2304), + num_q_heads=32, + num_kv_heads=2, + dtype=torch.float8_e4m3fn, + ), + id="msa_gqa_fp8_single_token", + ), +] + + +@pytest.mark.parametrize("scenario", _BLOCK_SPARSE_GQA_CASES) +def test_block_sparse_mqa_gqa(scenario: BlockSparseGqaScenario) -> None: + """MSA-selected KV blocks match a direct PyTorch block-sparse reference.""" + if not msa_package_available(): + pytest.skip("fmha_sm100 (MSA) is not importable") + + inputs = _create_block_sparse_gqa_inputs(scenario) + output = torch.empty( + scenario.total_q, + scenario.num_q_heads, + scenario.head_dim, + dtype=torch.bfloat16, + device=inputs["q"].device, + ) + run_msa_sparse_gqa( + inputs["q"], + inputs["k_paged"], + inputs["v_paged"], + inputs["kv_block_indexes"], + kv_indices=inputs["kv_indices"], + sm_scale=scenario.head_dim**-0.5, + qo_lens_cpu=torch.tensor(scenario.q_lens, dtype=torch.int32), + kv_lens_cpu=torch.tensor(scenario.kv_lens, dtype=torch.int32), + qo_offset_cpu=torch.tensor(scenario.causal_offsets, dtype=torch.int32), + causal=True, + head_dim=scenario.head_dim, + out=output, + use_fp8=scenario.dtype == torch.float8_e4m3fn, + ) + torch.cuda.synchronize() + + reference = _reference_block_sparse_gqa(inputs, scenario) + output_float = output.float() + reference_float = reference.float() + assert output.dtype == torch.bfloat16 + assert torch.isfinite(output_float).all() + cosine_similarity = torch.nn.functional.cosine_similarity( + output_float.flatten(), reference_float.flatten(), dim=0 + ) + threshold = 0.999 if scenario.dtype == torch.float8_e4m3fn else 0.9999 + assert cosine_similarity > threshold + + # Sparse index and paged KV-cache helpers. @@ -1195,7 +1480,7 @@ def _reference_sparse_generation_attention( return torch.stack(outputs, dim=0) -# Test input builders. Kernel selection remains explicit in each test below. +# Test input builders. Kernel selection remains explicit in each test above. def _create_context_inputs(s: ContextScenario) -> _ContextInputs: @@ -1328,3 +1613,134 @@ def _create_generation_inputs(s: GenerationScenario) -> _GenerationInputs: request_ids=request_ids, metadata=metadata, ) + + +def _create_block_sparse_gqa_inputs(s: BlockSparseGqaScenario) -> dict[str, torch.Tensor]: + """Build packed Q, paged KV, page tables, and request-local block indices.""" + device = torch.device("cuda") + generator = torch.Generator(device=device).manual_seed(42) + total_pages = sum(kv_len // s.page_size for kv_len in s.kv_lens) + + def random_qkv(shape: Tuple[int, ...]) -> torch.Tensor: + tensor = torch.randn( + shape, + dtype=torch.bfloat16, + device=device, + generator=generator, + ) + return tensor.to(s.dtype) + + q = random_qkv((s.total_q, s.num_q_heads, s.head_dim)) + logical_k = random_qkv((total_pages, s.num_kv_heads, s.page_size, s.head_dim)) + logical_v = random_qkv((total_pages, s.num_kv_heads, s.page_size, s.head_dim)) + + if s.shuffle_pages: + kv_indices = torch.randperm(total_pages, device=device, generator=generator) + k_paged = torch.empty_like(logical_k) + v_paged = torch.empty_like(logical_v) + k_paged[kv_indices] = logical_k + v_paged[kv_indices] = logical_v + else: + kv_indices = torch.arange(total_pages, device=device) + k_paged = logical_k + v_paged = logical_v + kv_indices = kv_indices.to(torch.int32) + + kv_block_indexes = torch.full( + (s.total_q, s.num_kv_heads, s.topk), + -1, + dtype=torch.int32, + device=device, + ) + q_offset = 0 + for q_len, kv_len in zip(s.q_lens, s.kv_lens, strict=True): + num_pages = kv_len // s.page_size + for query_idx in range(q_len): + for kv_head_idx in range(s.num_kv_heads): + start = (query_idx + kv_head_idx) % num_pages if s.per_token_blocks else 0 + blocks = sorted( + (start + block_idx) % num_pages for block_idx in range(s.selected_blocks) + ) + kv_block_indexes[ + q_offset + query_idx, + kv_head_idx, + : s.selected_blocks, + ] = torch.tensor(blocks, dtype=torch.int32, device=device) + q_offset += q_len + + return { + "q": q, + "k_paged": k_paged, + "v_paged": v_paged, + "kv_indices": kv_indices, + "kv_block_indexes": kv_block_indexes, + } + + +def _reference_block_sparse_gqa( + inputs: dict[str, torch.Tensor], s: BlockSparseGqaScenario +) -> torch.Tensor: + """Evaluate request-local block selection and per-request causal offsets.""" + output = torch.empty( + s.total_q, + s.num_q_heads, + s.head_dim, + dtype=torch.float32, + device=inputs["q"].device, + ) + q_offset = 0 + page_offset = 0 + for q_len, kv_len, causal_offset in zip( + s.q_lens, + s.kv_lens, + s.causal_offsets, + strict=True, + ): + num_pages = kv_len // s.page_size + physical_pages = inputs["kv_indices"][page_offset : page_offset + num_pages].long() + k = ( + inputs["k_paged"] + .index_select(0, physical_pages) + .permute(0, 2, 1, 3) + .reshape(kv_len, s.num_kv_heads, s.head_dim) + .float() + ) + v = ( + inputs["v_paged"] + .index_select(0, physical_pages) + .permute(0, 2, 1, 3) + .reshape(kv_len, s.num_kv_heads, s.head_dim) + .float() + ) + q = inputs["q"][q_offset : q_offset + q_len].float() + request_blocks = inputs["kv_block_indexes"][q_offset : q_offset + q_len] + token_positions = torch.arange(kv_len, device=q.device) + block_ids = token_positions // s.page_size + causal_mask = token_positions.view(1, -1) <= ( + torch.arange(q_len, device=q.device).view(-1, 1) + causal_offset + ) + + for kv_head_idx in range(s.num_kv_heads): + selected_blocks = request_blocks[:, kv_head_idx] + selected_mask = ( + (selected_blocks.unsqueeze(-1) == block_ids.view(1, 1, -1)) + & (selected_blocks.unsqueeze(-1) >= 0) + ).any(dim=1) + mask = selected_mask & causal_mask + q_head_begin = kv_head_idx * s.q_heads_per_kv_head + q_head_end = q_head_begin + s.q_heads_per_kv_head + scores = torch.einsum( + "qhd,kd->qhk", + q[:, q_head_begin:q_head_end] * (s.head_dim**-0.5), + k[:, kv_head_idx], + ) + scores.masked_fill_(~mask.unsqueeze(1), float("-inf")) + probabilities = torch.softmax(scores, dim=-1) + output[q_offset : q_offset + q_len, q_head_begin:q_head_end] = torch.einsum( + "qhk,kd->qhd", probabilities, v[:, kv_head_idx] + ) + + q_offset += q_len + page_offset += num_pages + + return output.to(torch.bfloat16) From 12a3f1dcb185fcb4987afcc0a8b47b6a633593c7 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Wed, 2 Sep 2026 05:25:30 -0700 Subject: [PATCH 05/12] [None][docs] Reorganize sparse attention guides Structure the user guide around supported attention types, public algorithms, and DSA deployment examples. Refresh the development guide capability descriptions and update the relocated MSA test waiver. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../sparse-attention-development-guide.md | 42 +-- docs/source/features/sparse-attention.md | 301 ++++++++++-------- tests/integration/test_lists/waives.txt | 4 +- 3 files changed, 196 insertions(+), 151 deletions(-) diff --git a/docs/source/developer-guide/sparse-attention-development-guide.md b/docs/source/developer-guide/sparse-attention-development-guide.md index 2af912379e51..9543861549a7 100644 --- a/docs/source/developer-guide/sparse-attention-development-guide.md +++ b/docs/source/developer-guide/sparse-attention-development-guide.md @@ -33,7 +33,7 @@ TensorRT LLM's sparse attention algorithms fall into two categories. - **Framework-level**: the algorithm runs a *prediction* step that emits sparse indices, which are then consumed by a shared `AttentionOp` to produce sparse KV cache updates and/or sparse attention computation. - Examples: **RocketKV** (page-level, MQA/GQA), **DSA** (token-level, + Examples: **RocketKV** (page-level, MHA/MQA/GQA), **DSA** (token-level, MLA). - **Kernel-level**: sparsity is implemented entirely inside the attention kernel — there is no external prediction or gather step. @@ -125,11 +125,11 @@ The current capability matrix is: | Attention type | Context phase | Generation phase | |---|---|---| -| MQA / MHA / GQA | sparse KV cache | sparse computation (page-level) | +| MQA / GQA | sparse KV cache and sparse computation (token-level) | sparse computation (token- or page-level) | +| MHA | sparse KV cache | sparse computation (page-level) | | MLA | sparse computation (token-level) | sparse computation (token-level) | -Context-phase sparse computation for MQA/GQA and dynamic generation-phase -KV eviction are tracked as future work. +Dynamic generation-phase KV eviction is tracked as future work. ### Prediction hooks @@ -171,13 +171,17 @@ Algorithm implementations live under

Figure 2: Sparse attention operator workflow in TensorRT LLM.

-For MQA/GQA, the op runs `gatherKvPageOffsetsKernel` before the -generation-phase attention kernel. It takes the (potentially unordered -or finer-grained) sparse indices and maps them to ordered, page-aligned -KV cache offsets, also producing an updated per-head effective KV -length. The downstream attention kernel reads only those pages. Today -MQA/GQA sparse computation is supported at **block (page) granularity** -in the generation phase only. +For page-sparse MHA/MQA/GQA, the op runs `gatherKvPageOffsetsKernel` +before the generation-phase attention kernel. It takes the (potentially +unordered or finer-grained) sparse indices and maps them to ordered, +page-aligned KV cache offsets, also producing an updated per-head +effective KV length. The downstream attention kernel reads only those +pages. + +Token-sparse MQA/GQA uses physical KV-cache token indices directly. It +supports packed context and generation computation, including a linear +sequence of draft tokens. Query heads in the same KV group share the KV +head's per-query token list. After context attention, `updateSparseKvCacheAfterFmha` post-processes the KV cache: it selects the important KV tokens and rewrites the @@ -261,14 +265,16 @@ prediction methods. - **Behavior**: return the sparse indices used by the generation-phase attention computation. - **Outputs**: - - `sparse_attn_indices`: shape `(nHeads, nBlocks)` — block indices on - the KV sequence dimension. Block size is set by the algorithm via - `sparse_attn_indices_block_size` (arbitrary value supported). + - `sparse_attn_indices`: sparse token or block indices on the KV + sequence dimension. Token-sparse MQA/GQA supplies one physical-token + list per KV head and query token. Page-sparse attention supplies + request-local block lists; block size is set by the algorithm via + `sparse_attn_indices_block_size`. - `sparse_attn_offsets`: shape `(nBatch + 1)` — same semantics as above. -- **Constraint**: today only **page-level** granularity is supported - for MQA/GQA sparse computation, and the generation-phase path uses - TRTLLM-GEN kernels (NVIDIA Blackwell SM 100+). +- **Constraint**: token-sparse MQA/GQA and page-sparse MHA/MQA/GQA use + different index layouts. Match the selected kernel contract; do not + pass request-local block indices to the physical-token path. Prediction is on the critical path and can dominate latency in low-latency scenarios. Plan for custom kernels (Triton or CUDA) rather @@ -326,8 +332,6 @@ for the kernel-side specifics. ## Roadmap -- **Sparse computation in context phase for MQA/MHA/GQA** — extend - framework coverage to context-phase sparse compute. - **Dynamic eviction in generation phase** — exploring block-level eviction as a compromise that keeps KV cache flexibility manageable. - **Unified auxiliary memory management** — let custom auxiliary pools diff --git a/docs/source/features/sparse-attention.md b/docs/source/features/sparse-attention.md index 6081cd6c68fa..a195d17d614f 100644 --- a/docs/source/features/sparse-attention.md +++ b/docs/source/features/sparse-attention.md @@ -1,11 +1,14 @@ # Sparse Attention - [Overview](#overview) +- [Supported Sparse Attentions](#supported-sparse-attentions) + - [Sparse MLA](#sparse-mla) + - [Sparse MQA/GQA](#sparse-mqagqa) + - [Sparse MHA](#sparse-mha) - [Supported Algorithms](#supported-algorithms) -- [Sparse MHA Kernel Support](#sparse-mha-kernel-support) -- [Sparse MQA/GQA Kernel Support](#sparse-mqagqa-kernel-support) -- [Configure Sparse Attention](#configure-sparse-attention) -- [Algorithm Details](#algorithm-details) + - [Capability Comparison](#capability-comparison) + - [Algorithm Details](#algorithm-details) +- [Usage with trtllm-bench and trtllm-serve](#usage-with-trtllm-bench-and-trtllm-serve) - [Further Reading](#further-reading) ## Overview @@ -29,42 +32,91 @@ class selected by its `algorithm` field. Model-native algorithms usually read their geometry from the checkpoint; avoid overriding those values unless the model-specific guide says they are tunable. -## Supported Algorithms +## Supported Sparse Attentions -| `algorithm` | Config class | Sparse mechanism | Attention implementation | Typical use | -|---|---|---|---|---| -| `rocket` | `RocketSparseAttentionConfig` | Prompt KV eviction, then page-level Top-K selection during decode | TRTLLM or Vanilla | Training-free sparsity for MHA/MQA/GQA models | -| `dsa` | `DeepSeekSparseAttentionConfig` | Learned token-level indexer followed by sparse MLA | TRTLLM | DeepSeek V3.2 and compatible model-native DSA architectures | -| `deepseek_v4` | `DeepSeekV4SparseAttentionConfig` | Sliding-window attention plus compressed sparse or compressed dense history | TRTLLM | DeepSeek-V4 hybrid attention | -| `minimax_m3` | `MiniMaxM3SparseAttentionConfig` | Learned block selection followed by sparse GQA | Dedicated Triton or MSA implementation | MiniMax-M3 sparse layers | -| `skip_softmax` | `SkipSoftmaxAttentionConfig` | Dynamically skips eligible softmax work inside the FMHA kernel | TRTLLM | Existing full-attention models with calibrated or direct thresholds | +TensorRT LLM supports sparse computation for MLA, MQA/GQA, and MHA. This +section describes the attention and kernel contracts independently of the +algorithm that produces the sparse pattern. The public algorithms that connect +selectors, cache management, and these attention implementations are listed in +[Supported Algorithms](#supported-algorithms). -All five configs select the PyTorch execution backend. The "attention -implementation" column refers to the attention kernel/backend used inside that -execution backend. +### Sparse MLA -### Capability Comparison +Sparse MLA consumes token-level selections against a latent KV cache. It is +used by model-native algorithms such as DeepSeek Sparse Attention and +DeepSeek-V4 hybrid attention. Both prefill and generation are supported, +including mixed batches. -| Capability | RocketKV | DSA | DeepSeek-V4 | MiniMax-M3 | Skip Softmax | -|---|---:|---:|---:|---:|---:| -| Sparse prefill computation | No | Yes | Yes | Yes | Yes | -| Sparse decode computation | Yes | Yes | Yes | Yes | Yes | -| Reduces retained main KV history | Yes | No | Yes, through model-native compression | No | No | -| Requires a model-trained selector | No | Yes | Yes | Yes | No | -| Selection granularity | Token eviction and pages | Tokens | Compressed entries | Blocks | Kernel tiles | +| Parameter | Support | +|---|---| +| GPU architecture | SM90, SM100, SM103, SM120, and SM121 through architecture-specific sparse MLA implementations | +| Sparse compute phase | Packed prefill, generation, and mixed context/generation batches | +| Attention type | MLA with a shared latent KV representation | +| Model geometry | Checkpoint-native geometry; tests cover DeepSeek-V3.2 (`qk_head_dim=192`, `v_head_dim=128`) and DeepSeek-V4 (`qk_head_dim=512`, `v_head_dim=512`) | +| Model input dtype | BF16 | +| KV-cache dtype | BF16 and the FP8 modes supported by the selected model and GPU architecture | +| Sparse indices | `int32` token indices; one selection per query token | +| Attention semantics | Causal self-attention | -"No" for RocketKV prefill means that prompt attention is still computed -densely. RocketKV selects which prompt KV entries to retain, so it reduces cache -size and later decode work. +[`test_sparse_mla_forward.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mla_forward.py) +covers pure prefill, pure generation, mixed batches, BF16/FP8 KV-cache modes, +and the direct FlashMLA sparse-forward contract. -## Sparse MHA Kernel Support +### Sparse MQA/GQA -RocketKV uses the shared page-sparse MHA path after its selector produces block -indices and per-request offsets. Prefill attention remains dense: RocketKV can -compact the retained KV cache after prefill, but sparse attention computation -starts during generation. +TensorRT LLM provides two sparse MQA/GQA compute paths. The token-sparse path +accepts a precomputed token list for each KV head and query token. Query heads +in the same KV group share the KV head's list. The block-sparse path accepts +request-local selections of 128-token KV blocks from a paged HND cache. -### Support Matrix +These are attention capabilities, not standalone public +`SparseAttentionConfig` algorithms. A user-facing algorithm must also provide +the selector, metadata, cache management, and backend integration. + +| Parameter | Token-sparse | Block-sparse | +|---|---|---| +| Sparse block size | 1 token | 128 tokens | +| GPU architecture | SM100 and SM103 | SM100 and SM103 | +| Sparse compute phase | Packed prefill, single-token generation, and linear draft-token generation; query lengths `1` and `4` are tested | Packed prefill, single-token generation, linear multi-query compute, and mixed batches | +| Attention type | MQA and GQA; Q heads must be divisible by KV heads | MQA and GQA; Q heads must be divisible by KV heads | +| Query heads per KV head | At most 32; tests cover `2`, `3`, `4`, `8`, `16`, `24`, `31`, and `32` | `2`, `4`, `8`, or `16`; all are tested | +| Q/KV head counts | No additional discrete kernel limit; tests cover Q heads `{6, 8, 16, 32, 48, 62, 64}` and KV heads `{1, 2, 4, 8}` | No additional discrete kernel limit; tests cover Q heads `{4, 8, 16, 32}` and KV heads `{1, 2}` | +| Model Q/K/V input dtype | BF16 or FP16 | BF16 or E4M3 FP8 | +| Model Q/K/V input layout | Fused QKV | Q `[tokens, q_heads, 128]`; paged K/V `[pages, kv_heads, 128, 128]` | +| Output dtype | BF16 or FP16 for every supported head dimension; E4M3 FP8 for head dimensions `64`, `128`, and `256` | BF16 | +| KV-cache dtype | BF16 or FP16 for every supported head dimension; E4M3 FP8 for head dimensions `64`, `128`, and `256` | BF16 or E4M3 FP8 | +| Q/K/V head dimension | Equal dimensions of `64`, `80`, `128`, or `256` | Equal dimension of `128` | +| KV-cache layout | Paged cache; page size is a power of two and at least 8 tokens; tests cover `8`, `16`, `32`, `64`, `128`, `256`, and `512` | Paged HND cache with page size `128`; shuffled physical pages and strided outer-page storage are tested | +| Sparse indices | `int32` physical token indices per KV head and query token | `int32` request-local block indices per KV head and query token; per-token lists, `-1` padding, and physical remapping are tested | +| Sparse Top-K | Positive multiple of 4; tests cover `4`, `32`, `64`, and `128` | `4`, `8`, `16`, or `32` selected blocks; all are tested | +| Attention semantics | Causal self-attention | Causal self-attention with bottom-right or explicit per-request query offsets | + +The token-sparse path is JIT-compiled with NVRTC. Its support is defined by the +current source checks rather than by the precompiled cubins that were present +when the feature was introduced. Linear draft-token generation is verified +with one target token and three draft tokens. Each query has its own causal +sparse list, including K/V written earlier in the same speculative forward. +Tree-shaped speculative masks are not applied by this path. + +For an FP8 KV cache, token-sparse Q is quantized to E4M3 during QKV +preprocessing while the model input remains BF16 or FP16. Tests cover both +BF16 output with an FP8 KV cache and the E4M3 FP8-output kernel. The shared +kernel validator also admits head dimension `512`, but the sparse path aborts +before launch for that configuration, so it is excluded from this matrix. + +Backend developers can use +[`test_sparse_mqa_gqa.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py) +as an integration example. Its static selectors isolate index/cache layout and +attention computation without presenting a public application API. + +### Sparse MHA + +The shared page-sparse MHA path consumes block indices and per-request offsets +produced by a sparse selector. Sparse MHA computation starts during generation; +prefill attention computation remains dense. An algorithm can still compact +the retained KV cache after prefill to reduce cache size and later decode work. + +#### Support Matrix | Parameter | Support | |---|---| @@ -91,109 +143,40 @@ selections, invokes `TrtllmAttention.forward`, and compares the result with an equivalent token-level PyTorch reference. RocketKV selector, metadata, and KT cache tests remain under the `rocketkv/` subdirectory. -## Sparse MQA/GQA Kernel Support - -TensorRT LLM contains an internal TRTLLM-Gen kernel for token-sparse -multi-query attention (MQA) and grouped-query attention (GQA). It accepts a -precomputed token-index list for each KV head and query token. Query heads in -the same KV group share the KV head's index list. - -This is a kernel capability, not a public `SparseAttentionConfig` algorithm. -There is no supported `algorithm: mqa_gqa` value for `LLM` or YAML. A sparse -algorithm must provide the selector, metadata, cache management, and attention -backend integration before applications can use this kernel through the public -API. - -### Support Matrix - -| Parameter | Supported | Not currently supported or not established | -|---|---|---| -| GPU architecture | SM100 and SM103 | Pre-Blackwell GPUs; SM120 and SM121 | -| Attention type | MQA (`num_kv_heads == 1`) and GQA | Arbitrary head mappings | -| Head relationship | `num_q_heads % num_kv_heads == 0`; at most 32 query heads per KV head; tests cover group sizes `2`, `3`, `4`, `8`, `16`, `24`, `31`, and `32` | Non-divisible Q/KV head counts and groups larger than 32 | -| Q/K/V dimensions | Equal QK and V head dimensions | Unequal QK/V dimensions (MLA uses a separate sparse path) | -| Head dimension | `64`, `80`, `128`, or `256` | `512` and other head dimensions | -| Model Q/K/V input dtype | BF16 or FP16 | FP8 model-input tensors are not wired through this XQA fixed-parameter path | -| Output dtype | BF16, FP16, or E4M3 FP8 | Other output dtypes | -| KV-cache dtype and layout | BF16, FP16, or E4M3 FP8 paged KV cache; page size is a power of two and at least 8 tokens | Contiguous cache, non-power-of-two pages, and pages smaller than 8 tokens | -| Sparse indices | `int32`, token-granular, one list per KV head and query token | A public built-in selector for generic MQA/GQA | -| Sparse Top-K | Positive multiple of 4; shorter sequences may pad unused entries with `-1` | Top-K values not divisible by 4 | -| Inference phase | Packed prefill, single-token generation, and linear draft-token generation (`qSeqLen=4` is tested) | Tree-shaped speculative masks are ignored by this static sparse kernel; mixed context/generation batches are not covered | -| Beam width | `1` | Beam search | -| Attention mask/window | Causal self-attention with a fixed cache window | ALiBi, arbitrary custom masks, StreamingLLM/sink tokens, and variable cyclic windows | - -The current main branch JIT-compiles this path with NVRTC. Its support is -therefore defined by the current TRTLLM-Gen source checks, not by the set of -precompiled cubins that was present when the feature was introduced. - -Linear draft-token generation is verified with four query tokens per request: -one target token plus three draft tokens. Each query has its own causal sparse -index list, including indices for K/V written earlier in the same speculative -forward. A separate branched-tree probe showed that the static sparse kernel -matches the unmasked reference rather than the tree-filtered reference, so -tree-shaped speculative masks are not supported. - -For FP8 KV cache, Q is quantized to E4M3 during QKV preprocessing and the XQA -runner selects E4M3 KV/math types while retaining BF16 or FP16 model input. -Tests cover both BF16 output with an FP8 KV cache and an E4M3 FP8-output kernel. - -The regression tests cover: - -- MQA and GQA group sizes `2`, `3`, `4`, `8`, `16`, `24`, `31`, and `32`; -- variable batch and sequence lengths; -- context KV compaction, context sparse computation, and decode sparse - computation; -- Top-K values `4`, `64`, and `128`, including Top-K larger than a request's - current KV length; -- backing KV-cache page sizes `32` and `64`; -- all supported equal head dimensions in both BF16 and FP16; -- linear generation with three draft tokens; -- E4M3 FP8 KV cache and FP8 output with BF16 model input. - -The shared TRTLLM-Gen option validator also admits head dimension `512`, but -the sparse MQA/GQA path aborts before launch for that configuration on current -main. It is therefore intentionally excluded from the supported matrix and -regression tests. - -Backend developers can use -[`test_sparse_mqa_gqa.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py) -as a minimal integration example. `_SparseMqaGqaParams` and -`_StaticSparseMqaGqaAttention` deliberately supply fixed sparse predictions so -that the test isolates the cache/index layout and kernel computation. They are -not public application APIs. - -## Configure Sparse Attention - -Pass a config object to `LLM` in Python, or use the equivalent discriminated -YAML object with `trtllm-serve`, `trtllm-bench`, or `trtllm-eval`. +## Supported Algorithms -```python -from tensorrt_llm import LLM -from tensorrt_llm.llmapi import RocketSparseAttentionConfig +The public `sparse_attention_config` API connects a sparse algorithm to its +selector, runtime metadata, cache management, and attention implementation. -llm = LLM( - model="", - sparse_attention_config=RocketSparseAttentionConfig(), -) -``` +| `algorithm` | Config class | Sparse mechanism | Attention implementation | Typical use | +|---|---|---|---|---| +| `rocket` | `RocketSparseAttentionConfig` | Prompt KV eviction, then page-level Top-K selection during decode | TRTLLM or Vanilla | Training-free sparsity for MHA/MQA/GQA models | +| `dsa` | `DeepSeekSparseAttentionConfig` | Learned token-level indexer followed by sparse MLA | TRTLLM | DeepSeek-V3.2 and compatible model-native DSA architectures | +| `deepseek_v4` | `DeepSeekV4SparseAttentionConfig` | Sliding-window attention plus compressed sparse or compressed dense history | TRTLLM | DeepSeek-V4 hybrid attention | +| `minimax_m3` | `MiniMaxM3SparseAttentionConfig` | Learned block selection followed by sparse GQA | Dedicated Triton or packaged block-sparse implementation | MiniMax-M3 sparse layers | +| `skip_softmax` | `SkipSoftmaxAttentionConfig` | Dynamically skips eligible softmax work inside the FMHA kernel | TRTLLM | Existing full-attention models with calibrated or direct thresholds | -```yaml -sparse_attention_config: - algorithm: rocket -``` +All five configs select the PyTorch execution backend. The "attention +implementation" column refers to the attention kernel/backend used inside that +execution backend. -For example: +### Capability Comparison -```bash -trtllm-serve --config config.yaml -trtllm-bench --model throughput --dataset --config config.yaml -``` +| Capability | RocketKV | DSA | DeepSeek-V4 | MiniMax-M3 | Skip Softmax | +|---|---:|---:|---:|---:|---:| +| Sparse prefill computation | No | Yes | Yes | Yes | Yes | +| Sparse decode computation | Yes | Yes | Yes | Yes | Yes | +| Reduces retained main KV history | Yes | No | Yes, through model-native compression | No | No | +| Requires a model-trained selector | No | Yes | Yes | Yes | No | +| Selection granularity | Token eviction and pages | Tokens | Compressed entries | Blocks | Kernel tiles | -The following sections list algorithm-specific settings and constraints. +"No" for RocketKV prefill means that prompt attention is still computed +densely. RocketKV selects which prompt KV entries to retain, so it reduces cache +size and later decode work. -## Algorithm Details +### Algorithm Details -### RocketKV +#### RocketKV [RocketKV](https://arxiv.org/pdf/2502.14051) is a training-free, two-stage algorithm for MHA, MQA, and GQA architectures. During prefill, it computes dense @@ -237,7 +220,7 @@ enable_chunked_prefill: false The TRTLLM and Vanilla attention implementations support RocketKV. The Vanilla implementation requires a BF16 KT cache. -### DeepSeek Sparse Attention +#### DeepSeek Sparse Attention DeepSeek Sparse Attention (DSA) is a model-native mechanism introduced by DeepSeek V3.2. A learned MQA indexer scores the KV history, Top-K selects token @@ -278,7 +261,7 @@ See the for model precision, hardware, parallelism, MTP, chunked-prefill, cache-reuse, and disaggregated-serving support. -### DeepSeek-V4 Hybrid Sparse Attention +#### DeepSeek-V4 Hybrid Sparse Attention DeepSeek-V4 interleaves three model-native attention modes: @@ -304,7 +287,7 @@ See the [DeepSeek-V4 example](../../../examples/models/core/deepseek_v4/README.md) for checkpoint-derived configuration and deployment constraints. -### MiniMax-M3 Block-Sparse GQA +#### MiniMax-M3 Block-Sparse GQA MiniMax-M3 uses model-native block-sparse GQA in its sparse layers. An index branch scores main KV-cache blocks, forces configured initial/local blocks into @@ -334,7 +317,7 @@ reuse or MTP. See the [MiniMax-M3 deployment guide](../deployment-guide/deployment-guide-for-minimax-m3-on-trtllm.md) for supported checkpoints and parallel deployment settings. -### Skip Softmax Attention +#### Skip Softmax Attention Skip Softmax Attention, also known as BLASST, dynamically skips eligible work inside a FlashAttention-style kernel. It does not select tokens, alter the @@ -409,6 +392,62 @@ fnmatch layer patterns. At most one checkpoint config group may use the Skip Softmax Attention requires the TRTLLM attention backend. Other attention backends do not apply it. +## Usage with trtllm-bench and trtllm-serve + +Sparse attention is configured through `sparse_attention_config` on the +PyTorch backend. DeepSeek-V3.2 provides a mature end-to-end example: its +checkpoint defines the DSA indexer geometry and Top-K, so the minimal YAML only +needs to select the `dsa` algorithm. + +```yaml +# config.yml +sparse_attention_config: + algorithm: dsa +``` + +Start an OpenAI-compatible server with the same config file used for other +PyTorch backend options: + +```bash +trtllm-serve deepseek-ai/DeepSeek-V3.2 \ + --backend pytorch \ + --tp_size 8 \ + --ep_size 8 \ + --custom_tokenizer deepseek_v32 \ + --config ./config.yml +``` + +For a throughput benchmark, first prepare or supply a tokenized dataset, then +pass the same config to `trtllm-bench`: + +```bash +trtllm-bench --model deepseek-ai/DeepSeek-V3.2 \ + prepare-dataset \ + --output ./deepseek-v3.2-dataset.json \ + token-norm-dist \ + --input-mean 4096 \ + --output-mean 512 \ + --input-stdev 0 \ + --output-stdev 0 \ + --num-requests 16 + +trtllm-bench --model deepseek-ai/DeepSeek-V3.2 throughput \ + --backend pytorch \ + --tp 8 \ + --ep 8 \ + --dataset ./deepseek-v3.2-dataset.json \ + --max_batch_size 16 \ + --max_num_tokens 8192 \ + --config ./config.yml +``` + +Use a local checkpoint path in place of the Hugging Face model ID when needed. +Other sparse algorithms use the same YAML entry point with their own +`algorithm` discriminator and settings. See the +[DeepSeek V3/V3.2 example](../../../examples/models/core/deepseek_v3/README.md) +for model precision, hardware, parallelism, MTP, chunked-prefill, cache-reuse, +and disaggregated-serving configurations. + ## Further Reading - [KV Cache Compression](kv-cache-compression.md) — methods that reduce the diff --git a/tests/integration/test_lists/waives.txt b/tests/integration/test_lists/waives.txt index c191d806680d..a8544f6b7093 100644 --- a/tests/integration/test_lists/waives.txt +++ b/tests/integration/test_lists/waives.txt @@ -312,9 +312,11 @@ test_e2e.py::test_ptp_quickstart_advanced_deepseek_r1_w4afp8_8gpus[DeepSeek-R1-W test_e2e.py::test_ptp_quickstart_bert[TRTLLM-BertForSequenceClassification-bert/bert-base-uncased-yelp-polarity] SKIP (https://nvbugs/6605819) test_e2e.py::test_ptp_quickstart_bert[VANILLA-BertForSequenceClassification-bert/bert-base-uncased-yelp-polarity] SKIP (bug pending, tracked in PR 17414) unittest/_torch/attention/sparse/dsa/test_req_idx_per_token.py::test_on_update_kv_lens_rebuilds_stale_map SKIP (https://nvbugs/6574939) +unittest/_torch/attention/sparse/msa/test_msa_backend.py::test_msa_paged_hnd_input_materializes_unaligned_outer_stride SKIP (https://nvbugs/6661846) +unittest/_torch/attention/sparse/rocketkv/test_rocketkv.py::test_model[TRTLLM-llama-3.1-model/Llama-3.1-8B-Instruct-pytorch] SKIP (https://nvbugs/6602094) +unittest/_torch/attention/sparse/rocketkv/test_rocketkv.py::test_model[VANILLA-llama-3.1-model/Llama-3.1-8B-Instruct-pytorch] SKIP (https://nvbugs/6602094) unittest/_torch/attention/sparse/test_minimax_m3_index_decode_score.py::test_index_decode_score_matches_msa_proxy[dtype0] SKIP (https://nvbugs/6669902) unittest/_torch/attention/sparse/test_minimax_m3_index_decode_score.py::test_index_decode_score_matches_msa_proxy[dtype1] SKIP (https://nvbugs/6669902) -unittest/_torch/attention/sparse/test_minimax_m3_msa_backend.py::test_msa_paged_hnd_input_materializes_unaligned_outer_stride SKIP (https://nvbugs/6661846) unittest/_torch/attention/sparse/test_minimax_m3_sparse_attn_decode.py::test_sparse_decode_matches_msa_kernel SKIP (https://nvbugs/6669902) unittest/_torch/attention/test_attention_backends.py::test_attention_backend[exaone_moe_gqa_swa128-ctx-bf16-HND-p32-v1] SKIP (https://nvbugs/6668773) unittest/_torch/attention/test_attention_backends.py::test_attention_backend[qwen2_0_5b_gqa_hd64-ctx-bf16-HND-p32-v1] SKIP (https://nvbugs/6641268) From 8c2874e20531787146b660f20853f26f05b21fe7 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Wed, 2 Sep 2026 06:50:18 -0700 Subject: [PATCH 06/12] [None][docs] Address sparse attention guide review Correct sparse attention capability descriptions, document backend-specific prediction contracts, and preserve compatibility anchors for existing documentation links. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../sparse-attention-development-guide.md | 72 +++++++++++-------- docs/source/features/sparse-attention.md | 31 ++++---- 2 files changed, 60 insertions(+), 43 deletions(-) diff --git a/docs/source/developer-guide/sparse-attention-development-guide.md b/docs/source/developer-guide/sparse-attention-development-guide.md index 9543861549a7..2079cc38e62e 100644 --- a/docs/source/developer-guide/sparse-attention-development-guide.md +++ b/docs/source/developer-guide/sparse-attention-development-guide.md @@ -31,10 +31,11 @@ rationale and high-level architecture diagrams, see the TensorRT LLM's sparse attention algorithms fall into two categories. - **Framework-level**: the algorithm runs a *prediction* step that emits - sparse indices, which are then consumed by a shared `AttentionOp` to - produce sparse KV cache updates and/or sparse attention computation. - Examples: **RocketKV** (page-level, MHA/MQA/GQA), **DSA** (token-level, - MLA). + sparse indices. A hook-based implementation can pass those indices to the + shared `AttentionOp`, while a dedicated backend can own prediction and + sparse computation end to end. Examples: **RocketKV** (token-level prompt + eviction plus page-level MHA/MQA/GQA decode selection), **DSA** + (token-level MLA), and **MiniMax-M3** (block-level GQA). - **Kernel-level**: sparsity is implemented entirely inside the attention kernel — there is no external prediction or gather step. The kernel decides what to skip from runtime values such as Softmax @@ -87,22 +88,24 @@ params. Framework-level sparse attention primarily targets approaches that leverage **token/sequence sparsity** — for many queries only a small fraction of historical tokens meaningfully contribute to the output, -and the framework exploits that in a GPU-friendly, structured way. -The attention operator provides unified APIs for both **sparse -computation** and **sparse KV cache**, so algorithm authors only need -to identify the important query/key pairs; everything else (KV cache -layout, kernel dispatch, page alignment) is handled by the framework. +and the framework exploits that in a GPU-friendly, structured way. On +the shared `AttentionOp` integration path, the operator provides APIs +for both **sparse computation** and **sparse KV cache** and owns KV-cache +layout conversion, kernel dispatch, and page alignment. An algorithm +with a dedicated attention implementation can instead perform those +steps in its backend while still using the common sparse config, +metadata, cache-manager, and registry framework. -It is built around three layers: +The shared `AttentionOp` path is built around three layers: - **Prediction module** — generates `sparse_kv_indices` (which KV tokens to keep in cache) and `sparse_attn_indices` (which KV pages or tokens to attend to during compute). - **`AttentionOp`** — consumes those indices via pre/post kernels and drives the core attention kernels. The op already understands - page-level sparsity for MQA/GQA in the generation phase, token-level - sparsity for MLA in both phases, and token-level KV compression in - the context phase for MQA/GQA. + page-level sparsity for MHA/MQA/GQA in the generation phase, + token-level MQA/GQA and MLA sparsity in both phases, and token-level + KV compression in the context phase for MHA/MQA/GQA. - **Auxiliary memory subsystem** — manages any extra pools (KT cache, indexer K cache, …) alongside the main KV cache. @@ -113,13 +116,13 @@ It is built around three layers:

Figure 1: Framework support for sparse attention in TensorRT LLM.

-Architecturally, each sparse attention algorithm subclasses the shared -`AttentionBackend` and supplies its own `sparse_kv_predict` / -`sparse_attn_predict` implementations. Different attention layers -within a single model can use different backends, so a model can mix -sparse attention strategies layer by layer. The shared `AttentionOp` -performs the actual computation and is not modified by individual -algorithms. +Hook-based implementations subclass `AttentionBackend`, supply +`sparse_kv_predict` / `sparse_attn_predict`, and reuse the shared +`AttentionOp` stack. A dedicated backend can implement sparse +computation directly; MiniMax-M3's default Triton backend follows this +model, while its packaged block-sparse implementation uses the hook-based +path. Different attention layers within a model can use different +backends, so sparse strategies can be mixed layer by layer. The current capability matrix is: @@ -157,6 +160,8 @@ Algorithm implementations live under - `dsa/` — DSA backend, indexer, metadata, cache manager, parameters, custom ops, and kernels. - `deepseek_v4/` — DeepSeek-V4 backend, indexer, metadata, cache manager, parameters, module hooks, and index conversion kernels. +- `minimax_m3/` — MiniMax-M3 Triton and packaged block-sparse backends, + indexer implementations, metadata, and `KVCacheManagerV2` integration. - `skip_softmax/` — SkipSoftmax parameter parsing and runtime scheduler. - `hooks.py` — typed MLA/Attention module adapters and common backend prediction orchestration. @@ -219,9 +224,11 @@ still reuse blocks. ## Adding a new framework-level algorithm -The four steps below cover what the runtime needs in order to dispatch a -new algorithm end-to-end. The order matches the natural development -flow — config first, then prediction, then memory, then registration. +The four steps below describe the hook-based `AttentionOp` integration +path. A dedicated backend reuses the configuration, auxiliary-memory, +and registration steps but owns its prediction and sparse computation +contracts. The order matches the natural development flow — config +first, then prediction, then memory, then registration. ### 1. Configuration class @@ -262,16 +269,19 @@ prediction methods. **`sparse_attn_predict(self, q, k, metadata, forward_args)`** -- **Behavior**: return the sparse indices used by the generation-phase - attention computation. +- **Behavior**: return the sparse indices used by attention computation in + the context phase, generation phase, or both, as supported by the backend. - **Outputs**: - - `sparse_attn_indices`: sparse token or block indices on the KV - sequence dimension. Token-sparse MQA/GQA supplies one physical-token - list per KV head and query token. Page-sparse attention supplies - request-local block lists; block size is set by the algorithm via + - `sparse_attn_indices`: backend-specific sparse token or block indices. + Token-sparse MQA/GQA uses shape + `(nKvHeads, nQueryTokens, topK)` with physical KV-pool token indices + and no offsets. Page-sparse attention uses request-local block indices; + the algorithm declares their block size through `sparse_attn_indices_block_size`. - - `sparse_attn_offsets`: shape `(nBatch + 1)` — same semantics as - above. + - `sparse_attn_offsets`: optional and backend-specific. RocketKV uses + `(numGenerations + 1)` request boundaries for its flattened page + selections. Token-sparse MQA/GQA and DSA leave it unset. DeepSeek-V4 + uses the field for secondary compressed-pool indices. - **Constraint**: token-sparse MQA/GQA and page-sparse MHA/MQA/GQA use different index layouts. Match the selected kernel contract; do not pass request-local block indices to the physical-token path. diff --git a/docs/source/features/sparse-attention.md b/docs/source/features/sparse-attention.md index a195d17d614f..e3c3c5b35e9c 100644 --- a/docs/source/features/sparse-attention.md +++ b/docs/source/features/sparse-attention.md @@ -42,16 +42,17 @@ selectors, cache management, and these attention implementations are listed in ### Sparse MLA -Sparse MLA consumes token-level selections against a latent KV cache. It is -used by model-native algorithms such as DeepSeek Sparse Attention and -DeepSeek-V4 hybrid attention. Both prefill and generation are supported, +Sparse MLA consumes token-level selections against a model-specific shared KV +representation. DeepSeek Sparse Attention selects entries from a low-rank +latent KV cache, while DeepSeek-V4 combines compressed full-head non-RoPE K +with its corresponding RoPE K. Both prefill and generation are supported, including mixed batches. | Parameter | Support | |---|---| | GPU architecture | SM90, SM100, SM103, SM120, and SM121 through architecture-specific sparse MLA implementations | | Sparse compute phase | Packed prefill, generation, and mixed context/generation batches | -| Attention type | MLA with a shared latent KV representation | +| Attention type | MLA with a model-specific shared KV representation | | Model geometry | Checkpoint-native geometry; tests cover DeepSeek-V3.2 (`qk_head_dim=192`, `v_head_dim=128`) and DeepSeek-V4 (`qk_head_dim=512`, `v_head_dim=512`) | | Model input dtype | BF16 | | KV-cache dtype | BF16 and the FP8 modes supported by the selected model and GPU architecture | @@ -116,8 +117,6 @@ produced by a sparse selector. Sparse MHA computation starts during generation; prefill attention computation remains dense. An algorithm can still compact the retained KV cache after prefill to reduce cache size and later decode work. -#### Support Matrix - | Parameter | Support | |---|---| | GPU architecture | SM100 is runtime-tested; SM103 is source-supported and enabled by the tests | @@ -143,6 +142,8 @@ selections, invokes `TrtllmAttention.forward`, and compares the result with an equivalent token-level PyTorch reference. RocketKV selector, metadata, and KT cache tests remain under the `rocketkv/` subdirectory. + + ## Supported Algorithms The public `sparse_attention_config` API connects a sparse algorithm to its @@ -156,9 +157,9 @@ selector, runtime metadata, cache management, and attention implementation. | `minimax_m3` | `MiniMaxM3SparseAttentionConfig` | Learned block selection followed by sparse GQA | Dedicated Triton or packaged block-sparse implementation | MiniMax-M3 sparse layers | | `skip_softmax` | `SkipSoftmaxAttentionConfig` | Dynamically skips eligible softmax work inside the FMHA kernel | TRTLLM | Existing full-attention models with calibrated or direct thresholds | -All five configs select the PyTorch execution backend. The "attention -implementation" column refers to the attention kernel/backend used inside that -execution backend. +All five configs are supported only by the PyTorch execution backend. The +"attention implementation" column refers to the attention kernel/backend used +inside that execution backend. ### Capability Comparison @@ -273,8 +274,10 @@ DeepSeek-V4 interleaves three model-native attention modes: TensorRT LLM normally constructs `DeepSeekV4SparseAttentionConfig` from the checkpoint. An explicit config overrides matching fields; it must preserve the model's attention layout. The current implementation requires -`window_size=128`, compression ratios from `{1, 4, 128}`, data-center Blackwell -GPUs, KV-cache blocks of `128` or `256` tokens, and beam width `1`. +`window_size=128`, compression ratios from `{1, 4, 128}`, Hopper (`SM90`) or +Blackwell (`SM100+`) GPUs, KV-cache blocks of `128` or `256` tokens, and beam +width `1`. Hopper requires `kv_cache_config.dtype=fp8_ds_mla`; on SM120 and +SM121, that cache layout requires 256-token blocks. ```yaml sparse_attention_config: @@ -317,6 +320,8 @@ reuse or MTP. See the [MiniMax-M3 deployment guide](../deployment-guide/deployment-guide-for-minimax-m3-on-trtllm.md) for supported checkpoints and parallel deployment settings. + + #### Skip Softmax Attention Skip Softmax Attention, also known as BLASST, dynamically skips eligible work @@ -348,7 +353,9 @@ sparse_attention_config: Alternatively, provide `target_sparsity`. This path requires the checkpoint to contain a calibration formula that maps the requested target to the kernel's -threshold scale factor. +threshold scale factor. `target_sparsity` is calibration guidance rather than a +runtime guarantee; the achieved sparsity depends on the model inputs and +workload. ```yaml sparse_attention_config: From acd0d8a1782099b07cb96d07243ac8e331823159 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Wed, 2 Sep 2026 06:59:40 -0700 Subject: [PATCH 07/12] [None][docs] Separate sparse support from test coverage Keep the sparse attention matrices focused on supported configurations and leave concrete regression coverage in the executable unit tests. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- docs/source/features/sparse-attention.md | 54 +++++++++++------------- 1 file changed, 25 insertions(+), 29 deletions(-) diff --git a/docs/source/features/sparse-attention.md b/docs/source/features/sparse-attention.md index e3c3c5b35e9c..a176048b7391 100644 --- a/docs/source/features/sparse-attention.md +++ b/docs/source/features/sparse-attention.md @@ -53,15 +53,15 @@ including mixed batches. | GPU architecture | SM90, SM100, SM103, SM120, and SM121 through architecture-specific sparse MLA implementations | | Sparse compute phase | Packed prefill, generation, and mixed context/generation batches | | Attention type | MLA with a model-specific shared KV representation | -| Model geometry | Checkpoint-native geometry; tests cover DeepSeek-V3.2 (`qk_head_dim=192`, `v_head_dim=128`) and DeepSeek-V4 (`qk_head_dim=512`, `v_head_dim=512`) | +| Model geometry | Checkpoint-native geometry for DeepSeek-V3.2 (`qk_head_dim=192`, `v_head_dim=128`) and DeepSeek-V4 (`qk_head_dim=512`, `v_head_dim=512`) | | Model input dtype | BF16 | | KV-cache dtype | BF16 and the FP8 modes supported by the selected model and GPU architecture | | Sparse indices | `int32` token indices; one selection per query token | | Attention semantics | Causal self-attention | +See [`test_sparse_mla_forward.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mla_forward.py) -covers pure prefill, pure generation, mixed batches, BF16/FP8 KV-cache modes, -and the direct FlashMLA sparse-forward contract. +for executable sparse MLA examples. ### Sparse MQA/GQA @@ -78,37 +78,36 @@ the selector, metadata, cache management, and backend integration. |---|---|---| | Sparse block size | 1 token | 128 tokens | | GPU architecture | SM100 and SM103 | SM100 and SM103 | -| Sparse compute phase | Packed prefill, single-token generation, and linear draft-token generation; query lengths `1` and `4` are tested | Packed prefill, single-token generation, linear multi-query compute, and mixed batches | +| Sparse compute phase | Packed prefill, single-token generation, and linear draft-token generation | Packed prefill, single-token generation, linear multi-query compute, and mixed batches | | Attention type | MQA and GQA; Q heads must be divisible by KV heads | MQA and GQA; Q heads must be divisible by KV heads | -| Query heads per KV head | At most 32; tests cover `2`, `3`, `4`, `8`, `16`, `24`, `31`, and `32` | `2`, `4`, `8`, or `16`; all are tested | -| Q/KV head counts | No additional discrete kernel limit; tests cover Q heads `{6, 8, 16, 32, 48, 62, 64}` and KV heads `{1, 2, 4, 8}` | No additional discrete kernel limit; tests cover Q heads `{4, 8, 16, 32}` and KV heads `{1, 2}` | +| Query heads per KV head | At most 32 | `2`, `4`, `8`, or `16` | +| Q/KV head counts | No additional discrete kernel limit | No additional discrete kernel limit | | Model Q/K/V input dtype | BF16 or FP16 | BF16 or E4M3 FP8 | | Model Q/K/V input layout | Fused QKV | Q `[tokens, q_heads, 128]`; paged K/V `[pages, kv_heads, 128, 128]` | | Output dtype | BF16 or FP16 for every supported head dimension; E4M3 FP8 for head dimensions `64`, `128`, and `256` | BF16 | | KV-cache dtype | BF16 or FP16 for every supported head dimension; E4M3 FP8 for head dimensions `64`, `128`, and `256` | BF16 or E4M3 FP8 | | Q/K/V head dimension | Equal dimensions of `64`, `80`, `128`, or `256` | Equal dimension of `128` | -| KV-cache layout | Paged cache; page size is a power of two and at least 8 tokens; tests cover `8`, `16`, `32`, `64`, `128`, `256`, and `512` | Paged HND cache with page size `128`; shuffled physical pages and strided outer-page storage are tested | -| Sparse indices | `int32` physical token indices per KV head and query token | `int32` request-local block indices per KV head and query token; per-token lists, `-1` padding, and physical remapping are tested | -| Sparse Top-K | Positive multiple of 4; tests cover `4`, `32`, `64`, and `128` | `4`, `8`, `16`, or `32` selected blocks; all are tested | +| KV-cache layout | Paged cache; page size is a power of two and at least 8 tokens | Paged HND cache with page size `128`; supports shuffled physical pages and strided outer-page storage | +| Sparse indices | `int32` physical token indices per KV head and query token | `int32` request-local block indices per KV head and query token; supports per-token lists, `-1` padding, and physical remapping | +| Sparse Top-K | Positive multiple of 4 | `4`, `8`, `16`, or `32` selected blocks | | Attention semantics | Causal self-attention | Causal self-attention with bottom-right or explicit per-request query offsets | The token-sparse path is JIT-compiled with NVRTC. Its support is defined by the current source checks rather than by the precompiled cubins that were present -when the feature was introduced. Linear draft-token generation is verified -with one target token and three draft tokens. Each query has its own causal -sparse list, including K/V written earlier in the same speculative forward. -Tree-shaped speculative masks are not applied by this path. +when the feature was introduced. During linear draft-token generation, each +query has its own causal sparse list, including K/V written earlier in the same +speculative forward. Tree-shaped speculative masks are not applied by this +path. For an FP8 KV cache, token-sparse Q is quantized to E4M3 during QKV -preprocessing while the model input remains BF16 or FP16. Tests cover both -BF16 output with an FP8 KV cache and the E4M3 FP8-output kernel. The shared -kernel validator also admits head dimension `512`, but the sparse path aborts -before launch for that configuration, so it is excluded from this matrix. +preprocessing while the model input remains BF16 or FP16. The path supports +both BF16 output with an FP8 KV cache and E4M3 FP8 output. The shared kernel +validator also admits head dimension `512`, but the sparse path aborts before +launch for that configuration, so it is excluded from this matrix. Backend developers can use [`test_sparse_mqa_gqa.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py) -as an integration example. Its static selectors isolate index/cache layout and -attention computation without presenting a public application API. +as an executable integration example. ### Sparse MHA @@ -119,28 +118,25 @@ the retained KV cache after prefill to reduce cache size and later decode work. | Parameter | Support | |---|---| -| GPU architecture | SM100 is runtime-tested; SM103 is source-supported and enabled by the tests | -| Sparse compute phase | Single-token and linear draft-token generation (`qSeqLen=4` is tested) | +| GPU architecture | SM100 and SM103 | +| Sparse compute phase | Single-token and linear draft-token generation | | Attention type | MHA (`num_q_heads == num_kv_heads`) | | Query heads per KV head | `1` | -| Number of MHA heads | No additional discrete source restriction beyond `num_q_heads == num_kv_heads > 0`; tests cover `1`, `2`, `3`, `4`, `8`, `16`, `24`, `32`, `48`, `64`, `96`, and `128` | +| Number of MHA heads | No additional discrete source restriction beyond `num_q_heads == num_kv_heads > 0` | | Model Q/K/V input dtype | BF16 or FP16 | | Model Q/K/V input layout | Fused QKV | | Output dtype | Model dtype for head dimensions `64`, `80`, `128`, and `256`; E4M3 FP8 for head dimensions `64`, `128`, and `256` with an FP8 KV cache | | KV-cache dtype | Model dtype for head dimensions `64`, `80`, `128`, and `256`; E4M3 FP8 for head dimensions `64`, `128`, and `256` | | Q/K/V dimensions | Equal head dimensions of `64`, `80`, `128`, or `256` | -| KV-cache layout | Paged KV cache; page sizes `8`, `16`, `32`, `64`, `128`, `256`, and `512` are tested | +| KV-cache layout | Paged KV cache; page size is a power of two and at least 8 tokens | | Selection granularity | Block indices expanded to KV-cache pages | -| Sparse indices | `int32` block indices with `int32` per-request offsets; per-head patterns, unordered indices, and variable request offsets are tested | -| Sparse index block size | Blocks may cross KV-page boundaries; sizes `1`, `2`, `3`, `4`, `5`, `8`, `16`, `24`, `32`, and `48` are tested | +| Sparse indices | `int32` block indices with `int32` per-request offsets; supports per-head patterns, unordered indices, and variable request offsets | +| Sparse index block size | Positive; blocks may cross KV-page boundaries | | Attention semantics | Causal self-attention | Backend developers can use [`test_sparse_mha.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mha.py) -as an architecture-level integration example. It supplies static page -selections, invokes `TrtllmAttention.forward`, and compares the result with an -equivalent token-level PyTorch reference. RocketKV selector, metadata, and KT -cache tests remain under the `rocketkv/` subdirectory. +as an executable integration example. From f86043d6af40db6f7229448fba6213bbb38d8092 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Wed, 2 Sep 2026 08:36:21 -0700 Subject: [PATCH 08/12] [None][fix] Address sparse attention review findings Tighten FP8 sparse MQA/GQA references, cover MHA KV compaction, correct sparse backend contracts and documentation, and finish organizing algorithm-specific tests. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../sparse-attention-development-guide.md | 45 +++++----- docs/source/features/sparse-attention.md | 83 +++++++++---------- .../attention/ATTENTION_DEVELOPER_GUIDE.md | 80 +++++++++++++----- .../attention/backends/sparse/dsa/backend.py | 2 +- tests/integration/test_lists/waives.txt | 6 +- .../{ => msa}/test_minimax_m3_dense_decode.py | 0 .../test_minimax_m3_index_decode_score.py | 0 .../{ => msa}/test_minimax_m3_msa_selector.py | 0 .../test_minimax_m3_sparse_attn_decode.py | 0 .../attention/sparse/test_sparse_mqa_gqa.py | 30 ++++++- 10 files changed, 150 insertions(+), 96 deletions(-) rename tests/unittest/_torch/attention/sparse/{ => msa}/test_minimax_m3_dense_decode.py (100%) rename tests/unittest/_torch/attention/sparse/{ => msa}/test_minimax_m3_index_decode_score.py (100%) rename tests/unittest/_torch/attention/sparse/{ => msa}/test_minimax_m3_msa_selector.py (100%) rename tests/unittest/_torch/attention/sparse/{ => msa}/test_minimax_m3_sparse_attn_decode.py (100%) diff --git a/docs/source/developer-guide/sparse-attention-development-guide.md b/docs/source/developer-guide/sparse-attention-development-guide.md index 2079cc38e62e..32d103b50058 100644 --- a/docs/source/developer-guide/sparse-attention-development-guide.md +++ b/docs/source/developer-guide/sparse-attention-development-guide.md @@ -116,13 +116,13 @@ The shared `AttentionOp` path is built around three layers:

Figure 1: Framework support for sparse attention in TensorRT LLM.

-Hook-based implementations subclass `AttentionBackend`, supply -`sparse_kv_predict` / `sparse_attn_predict`, and reuse the shared -`AttentionOp` stack. A dedicated backend can implement sparse -computation directly; MiniMax-M3's default Triton backend follows this -model, while its packaged block-sparse implementation uses the hook-based -path. Different attention layers within a model can use different -backends, so sparse strategies can be mixed layer by layer. +Hook-based `TrtllmAttention` implementations supply `sparse_kv_predict` / +`sparse_attn_predict` and reuse the shared `AttentionOp` stack. RocketKV's +`VanillaAttention` implementation instead uses per-request Python hooks. A +dedicated backend can implement sparse computation directly; MiniMax-M3's +default Triton backend follows this model. Different attention layers within a +model can use different backends, so sparse strategies can be mixed layer by +layer. The current capability matrix is: @@ -136,8 +136,8 @@ Dynamic generation-phase KV eviction is tracked as future work. ### Prediction hooks -`AttentionBackend` exposes two prediction methods that algorithm-specific -subclasses override: +`TrtllmAttention`-based sparse backends expose two prediction methods that +algorithm-specific subclasses override: ```python sparse_kv_indices, sparse_kv_offsets = self.sparse_kv_predict(q, k, metadata, forward_args) @@ -199,8 +199,9 @@ For sparse MLA, the kernel consumes token-level indices directly, so `gatherKvPageOffsetsKernel` is bypassed — both context and generation phases are supported at token granularity. The sparse MLA path currently expects **global** KV cache pool addresses with token-level -offsets, not request-local logical positions. Sparse KV cache for MLA -is not yet supported. +offsets, not request-local logical positions. MLA does not support the shared +`sparse_kv_indices` in-place compaction path. DeepSeek-V4's model-native +compressed-history pools use a separate cache path. ### Auxiliary memory pools @@ -248,10 +249,11 @@ the bottom of the file. ### 2. Prediction module -Create a new backend class inheriting from `TrtllmAttention` (or -`VanillaAttention` if appropriate) in +Create a new backend class inheriting from `TrtllmAttention` in `tensorrt_llm/_torch/attention/backends/sparse/`. Override one or both -prediction methods. +prediction methods. A `VanillaAttention` implementation instead overrides +`_single_request_sparse_kv_predict` and +`_single_request_sparse_attn_predict` with its per-request Python contract. **`sparse_kv_predict(self, q, k, metadata, forward_args)`** @@ -311,10 +313,10 @@ If the algorithm needs extra tensors beyond the main KV cache: ### 4. Registration and dispatch -- Register the new config + backend in - `tensorrt_llm/_torch/attention/backends/sparse/registry.py` and - `tensorrt_llm/_torch/pyexecutor/_util.py` so the runtime routes - requests to your backend when the config is present. +- Register the new config and backend in + `tensorrt_llm/_torch/attention/backends/sparse/registry.py`. Update executor + wiring only when the algorithm requires behavior beyond the registry's + generic dispatch. - If the algorithm customizes module-layer behavior, implement and register a concrete `MLASparseHooks` or `AttentionSparseHooks` adapter from the algorithm's `module.py`. @@ -332,9 +334,10 @@ framework wiring is: - A new config subclass with its own `algorithm` discriminator. - A lowered `SparseParams` object that carries the resolved kernel settings. -- A switch inside the attention backend (e.g., - `_torch/attention/backends/fmha/flashinfer_trtllm_gen.py`) that reads the lowered params - and enables the kernel-side fast path. +- A switch inside the attention backend, such as + `_torch/attention/backends/trtllm.py` or an implementation under + `_torch/attention/backends/fmha/`, that reads the lowered params and enables + the kernel-side fast path. Skip Softmax Attention follows this pattern — see the [BLASST tech blog](../blogs/tech_blog/blog16_Accelerating_Long_Context_Inference_with_Skip_Softmax_Attention.md) diff --git a/docs/source/features/sparse-attention.md b/docs/source/features/sparse-attention.md index a176048b7391..d47b456444d2 100644 --- a/docs/source/features/sparse-attention.md +++ b/docs/source/features/sparse-attention.md @@ -50,13 +50,18 @@ including mixed batches. | Parameter | Support | |---|---| -| GPU architecture | SM90, SM100, SM103, SM120, and SM121 through architecture-specific sparse MLA implementations | -| Sparse compute phase | Packed prefill, generation, and mixed context/generation batches | -| Attention type | MLA with a model-specific shared KV representation | -| Model geometry | Checkpoint-native geometry for DeepSeek-V3.2 (`qk_head_dim=192`, `v_head_dim=128`) and DeepSeek-V4 (`qk_head_dim=512`, `v_head_dim=512`) | -| Model input dtype | BF16 | -| KV-cache dtype | BF16 and the FP8 modes supported by the selected model and GPU architecture | -| Sparse indices | `int32` token indices; one selection per query token | +| GPU architecture | SM90, SM100, SM103, SM120, and SM121 | +| Compute phase | Packed prefill and generation, including mixed batches | +| Attention type | MLA | +| Head counts | Checkpoint-defined | +| Q heads per KV head | Not applicable; the model uses a shared KV representation | +| Head dimensions | DeepSeek-V3.2: QK `192`, V `128`; DeepSeek-V4: QK/V `512` | +| Input dtype | BF16 | +| Input layout | Model-native MLA inputs | +| Output dtype | BF16 | +| KV-cache dtype | BF16 or model- and architecture-specific FP8 | +| KV-cache layout | Paged, model-specific shared KV representation | +| Sparse granularity | Token | | Attention semantics | Causal self-attention | See @@ -65,10 +70,12 @@ for executable sparse MLA examples. ### Sparse MQA/GQA -TensorRT LLM provides two sparse MQA/GQA compute paths. The token-sparse path -accepts a precomputed token list for each KV head and query token. Query heads -in the same KV group share the KV head's list. The block-sparse path accepts -request-local selections of 128-token KV blocks from a paged HND cache. +The table below compares token-sparse and 128-token block-sparse MQA/GQA. The +token-sparse path accepts a precomputed token list for each KV head and query +token; query heads in the same KV group share that list. The block-sparse path +accepts request-local KV-block selections from a paged HND cache. The shared +page-sparse generation path described under [Sparse MHA](#sparse-mha) also +supports MQA and GQA. These are attention capabilities, not standalone public `SparseAttentionConfig` algorithms. A user-facing algorithm must also provide @@ -76,34 +83,27 @@ the selector, metadata, cache management, and backend integration. | Parameter | Token-sparse | Block-sparse | |---|---|---| -| Sparse block size | 1 token | 128 tokens | | GPU architecture | SM100 and SM103 | SM100 and SM103 | -| Sparse compute phase | Packed prefill, single-token generation, and linear draft-token generation | Packed prefill, single-token generation, linear multi-query compute, and mixed batches | -| Attention type | MQA and GQA; Q heads must be divisible by KV heads | MQA and GQA; Q heads must be divisible by KV heads | -| Query heads per KV head | At most 32 | `2`, `4`, `8`, or `16` | -| Q/KV head counts | No additional discrete kernel limit | No additional discrete kernel limit | -| Model Q/K/V input dtype | BF16 or FP16 | BF16 or E4M3 FP8 | -| Model Q/K/V input layout | Fused QKV | Q `[tokens, q_heads, 128]`; paged K/V `[pages, kv_heads, 128, 128]` | +| Compute phase | Packed prefill and generation, including linear draft tokens | Packed prefill and generation, including linear multi-query and mixed batches | +| Attention type | MQA and GQA | MQA and GQA | +| Head counts | Q heads must be divisible by KV heads; no other discrete limit | Q heads must be divisible by KV heads; no other discrete limit | +| Q heads per KV head | At most 32 | `2`, `4`, `8`, or `16` | +| Head dimensions | Q/K/V: `64`, `80`, `128`, or `256` | Q/K/V: `128` | +| Input dtype | BF16 or FP16 | BF16 or E4M3 FP8 | +| Input layout | Fused QKV | Q `[tokens, q_heads, 128]`; paged K/V `[pages, kv_heads, 128, 128]` | | Output dtype | BF16 or FP16 for every supported head dimension; E4M3 FP8 for head dimensions `64`, `128`, and `256` | BF16 | | KV-cache dtype | BF16 or FP16 for every supported head dimension; E4M3 FP8 for head dimensions `64`, `128`, and `256` | BF16 or E4M3 FP8 | -| Q/K/V head dimension | Equal dimensions of `64`, `80`, `128`, or `256` | Equal dimension of `128` | | KV-cache layout | Paged cache; page size is a power of two and at least 8 tokens | Paged HND cache with page size `128`; supports shuffled physical pages and strided outer-page storage | -| Sparse indices | `int32` physical token indices per KV head and query token | `int32` request-local block indices per KV head and query token; supports per-token lists, `-1` padding, and physical remapping | -| Sparse Top-K | Positive multiple of 4 | `4`, `8`, `16`, or `32` selected blocks | +| Sparse granularity | Token | Block (`128` tokens) | | Attention semantics | Causal self-attention | Causal self-attention with bottom-right or explicit per-request query offsets | -The token-sparse path is JIT-compiled with NVRTC. Its support is defined by the -current source checks rather than by the precompiled cubins that were present -when the feature was introduced. During linear draft-token generation, each -query has its own causal sparse list, including K/V written earlier in the same -speculative forward. Tree-shaped speculative masks are not applied by this -path. +The token-sparse path is JIT-compiled with NVRTC. During linear draft-token +generation, each query has its own causal sparse list, including K/V written +earlier in the same speculative forward. For an FP8 KV cache, token-sparse Q is quantized to E4M3 during QKV preprocessing while the model input remains BF16 or FP16. The path supports -both BF16 output with an FP8 KV cache and E4M3 FP8 output. The shared kernel -validator also admits head dimension `512`, but the sparse path aborts before -launch for that configuration, so it is excluded from this matrix. +both BF16 output with an FP8 KV cache and E4M3 FP8 output. Backend developers can use [`test_sparse_mqa_gqa.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py) @@ -119,19 +119,17 @@ the retained KV cache after prefill to reduce cache size and later decode work. | Parameter | Support | |---|---| | GPU architecture | SM100 and SM103 | -| Sparse compute phase | Single-token and linear draft-token generation | -| Attention type | MHA (`num_q_heads == num_kv_heads`) | -| Query heads per KV head | `1` | -| Number of MHA heads | No additional discrete source restriction beyond `num_q_heads == num_kv_heads > 0` | -| Model Q/K/V input dtype | BF16 or FP16 | -| Model Q/K/V input layout | Fused QKV | +| Compute phase | Generation, including single-token and linear draft-token inputs | +| Attention type | MHA | +| Head counts | Positive and `num_q_heads == num_kv_heads`; no other discrete limit | +| Q heads per KV head | `1` | +| Head dimensions | Q/K/V: `64`, `80`, `128`, or `256` | +| Input dtype | BF16 or FP16 | +| Input layout | Fused QKV | | Output dtype | Model dtype for head dimensions `64`, `80`, `128`, and `256`; E4M3 FP8 for head dimensions `64`, `128`, and `256` with an FP8 KV cache | | KV-cache dtype | Model dtype for head dimensions `64`, `80`, `128`, and `256`; E4M3 FP8 for head dimensions `64`, `128`, and `256` | -| Q/K/V dimensions | Equal head dimensions of `64`, `80`, `128`, or `256` | | KV-cache layout | Paged KV cache; page size is a power of two and at least 8 tokens | -| Selection granularity | Block indices expanded to KV-cache pages | -| Sparse indices | `int32` block indices with `int32` per-request offsets; supports per-head patterns, unordered indices, and variable request offsets | -| Sparse index block size | Positive; blocks may cross KV-page boundaries | +| Sparse granularity | Positive-size blocks expanded to KV-cache pages | | Attention semantics | Causal self-attention | Backend developers can use @@ -235,11 +233,6 @@ llm = LLM( ) ``` -```yaml -sparse_attention_config: - algorithm: dsa -``` - On supported Blackwell configurations, Guess-Verify-Refine (GVR) can replace the regular decode Top-K dispatcher. The current implementation accepts `index_topk` values `512`, `1024`, and `2048`, and indexer compression ratios diff --git a/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md b/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md index 1db7ff0fbc4d..c566e0c3d3a0 100644 --- a/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md +++ b/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md @@ -17,6 +17,11 @@ Use it when modifying the current implementation or adding a new model's attention behavior. It covers standard `Attention`, Multi-head Latent Attention (MLA), dense backends, and sparse backends. +For user-visible sparse attention capabilities and configuration, see the +[Sparse Attention feature guide](../../../docs/source/features/sparse-attention.md). +For the framework hooks and steps for adding a sparse algorithm, see the +[Sparse Attention Development Guide](../../../docs/source/developer-guide/sparse-attention-development-guide.md). + ## Glossary | Acronym | Meaning | @@ -156,10 +161,11 @@ the backend-to-AttentionOp `SparseRuntimeParams`, live in For MLA-related tasks, first check whether the work fits the current projection structure, can stay on an existing backend and metadata family, and -can preserve the current latent-cache / paged-KV contract. If it can, the -task usually stays within the existing MLA stack. If it depends on sparse -helper-level control flow, read `mla.py`, `attention/backends/sparse/hooks.py`, -and the relevant algorithm's `module.py` directly. +can preserve the current model-specific shared-KV / paged-KV contract. If it +can, the task usually stays within the existing MLA stack. If it depends on +sparse helper-level control flow, read `mla.py`, +`attention/backends/sparse/hooks.py`, and the relevant algorithm's `module.py` +directly. ## 2. Backend Layer Reference @@ -183,14 +189,27 @@ Base backend families: Sparse attention is not selected by a separate top-level module. User-facing `SparseAttentionConfig` objects live in LLM / VisualGen args and `ModelConfig`. -Attention modules use those configs to select sparse backend classes, then -lower the configs into `SparseParams` for backend construction. KV-cache -managers stay model-scope and consume the user-facing config directly. -Sparse metadata consumes `SparseMetadataParams`, derived independently from the -same user-facing config. +`config.attn_backend` still selects the base backend family; the sparse +algorithm refines that choice through `attention/backends/sparse/registry.py`. +Attention modules lower the user config into backend-owned `SparseParams`. +KV-cache managers stay model-scope and consume the user-facing config directly, +while sparse metadata consumes a separately lowered `SparseMetadataParams`. + +Framework-level algorithms either use the hook-based `TrtllmAttention` / +`AttentionOp` path or own prediction and computation in a dedicated backend. +The hook-based path carries module inputs through `SparseBackendForwardArgs` +and backend outputs through `SparseRuntimeParams`. `VanillaAttention` does not +use that contract; RocketKV's Vanilla implementation uses per-request Python +hooks. Kernel-level sparsity such as Skip Softmax has no external predictor. +See the +[Sparse Attention Development Guide](../../../docs/source/developer-guide/sparse-attention-development-guide.md) +for algorithm-specific hooks, index layouts, and cache managers. Sparse registrations are defined in `attention/backends/sparse/registry.py`. Check -that file for the current supported combinations, as they may change over time. +that file for the current config/backend combinations. Consult the +[feature guide](../../../docs/source/features/sparse-attention.md#supported-sparse-attentions) +for the supported attention shapes; do not infer support from algorithm +registration alone. ### 2.3 Backend contract @@ -272,6 +291,10 @@ workspace, page-table KV metadata, and prefill/decode wrapper state. sparse-specific runtime state (indexer buffers, routing state, side-cache state). +`SparseRuntimeParams` is the backend-to-`AttentionOp` carrier only on the +`TrtllmAttention` path. Its fields are algorithm-specific, not a generic +sparse-attention ABI; `VanillaAttention` uses its own per-request contract. + ### 3.2 KV-cache and decode-time semantics The main question is not just "does the backend read K and V?" but: @@ -296,7 +319,8 @@ use cache. `KVCacheManager.get_buffers()` exposes a per-layer view of the primary pool: - For standard dense attention, `kv_factor = 2` (separate K and V planes). -- For MLA-style cache, `kv_factor = 1` (one latent-cache tensor per token). +- For MLA-style cache, `kv_factor = 1` (one model-specific shared-KV tensor per + token in the primary pool). The main differences across backends: @@ -377,6 +401,8 @@ The FMHA package is split by role: [vendored-source lifecycle](../../../3rdparty/vendor-sources.md). Land upstream-worthy changes in FlashInfer and update the vendor lock; keep only TRT-LLM-specific adaptations in the persistent patch. +- `fmha/msa_sparse_gqa.py` integrates the packaged SM100/SM103 block-sparse + GQA implementation. - `fmha/flashinfer_sparse_mla.py` implements the FlashInfer SM120/SM121 sparse MLA FMHA library. - `fmha/flashinfer_trtllm_gen.py` implements the FlashInfer trtllm-gen FMHA @@ -390,13 +416,15 @@ shape. #### 3.2.3 MLA cached-context semantics -MLA cached state is not regular dense K and V. The paged cache stores -latent-cache state rather than separate K and V planes. Backend ops handle -appending, RoPE application, and loading cached state for attention use. +MLA cached state is not regular dense K and V. Dense MLA and DSA store a +low-rank latent representation rather than separate K and V planes. +DeepSeek-V4 instead combines model-specific sliding-window and compressed +full-head representations across multiple pools. Backend ops handle appending, +RoPE application, and loading the appropriate cached state for attention use. MLA fit cannot be judged from attention math alone. The module and backend must -agree on latent-cache layout, paged-KV read/write paths, and cached/chunked -context behavior. Read `mla.py` and the relevant +agree on the shared-KV representation, paged-KV read/write paths, auxiliary +pools, and cached/chunked-context behavior. Read `mla.py` and the relevant backend code for the current implementation details. fp8 context-MLA also stages a K/V dequant workspace sized by summed attended KV @@ -404,13 +432,12 @@ length; it is declared through the workspace memory-accounting contract (§2.3). #### 3.2.4 Sparse side-cache semantics -Sparse backends may add side caches beyond the main KV cache. Some sparse -algorithms keep the standard cache manager; others replace it with a -sparse-aware cache manager that adds side caches for indexing or routing. - -When evaluating new sparse attention, check both the main KV-cache contract -and the side-cache contract. See `attention/backends/sparse/` for the current -sparse cache managers and their side-cache structures. +Sparse backends may add side caches for indexing, routing, or compressed +history. Check their allocation, request lifecycle, block reuse, chunked +prefill, disaggregated transfer, CUDA Graph, and speculative-decoding contracts. +See `attention/backends/sparse/` and the +[Sparse Attention Development Guide](../../../docs/source/developer-guide/sparse-attention-development-guide.md) +for details. ## 4. Evaluating New Attention @@ -493,6 +520,8 @@ Working rules: | `tensorrt_llm/_torch/attention/backends/fmha/` | Internal TRTLLM FMHA libraries | | `tensorrt_llm/_torch/attention/backends/vanilla.py` | Torch fallback backend and metadata | | `tensorrt_llm/_torch/attention/backends/flashinfer.py` | FlashInfer backend and metadata | +| `tensorrt_llm/_torch/attention/backends/sparse/params.py` | Lowered sparse parameters and module/backend runtime carriers | +| `tensorrt_llm/_torch/attention/backends/sparse/registry.py` | Sparse backend, metadata, and cache-manager registration | | `tensorrt_llm/_torch/attention/backends/sparse/hooks.py` | Sparse module hooks and backend prediction orchestration | | `tensorrt_llm/_torch/attention/backends/sparse//module.py` | Algorithm-specific module-hook implementations | | `tensorrt_llm/_torch/attention/backends/sparse/` | Sparse prediction backends, metadata, cache managers, and kernels | @@ -505,6 +534,11 @@ Working rules: separately. - Any dispatch change touching `forward_context()` needs chunked-context tests. +Keep reusable sparse computation in the root `test_sparse_mla_forward.py`, +`test_sparse_mqa_gqa.py`, and `test_sparse_mha.py` modules. Shared framework +tests live in `test_sparse_attention.py`; selector and cache tests belong in +algorithm subdirectories such as `dsa/`, `msa/`, and `rocketkv/`. + Key test files: - `tests/unittest/_torch/attention/test_attention.py` diff --git a/tensorrt_llm/_torch/attention/backends/sparse/dsa/backend.py b/tensorrt_llm/_torch/attention/backends/sparse/dsa/backend.py index af2036946e14..1c14f3421a51 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/dsa/backend.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/dsa/backend.py @@ -312,7 +312,7 @@ def _grouped_remap_topk_to_global( Grouped output is bit-identical to the per-layer path by construction; that equivalence is covered by unit tests (see - ``tests/unittest/_torch/attention/sparse/test_cpp_custom_ops.py``). + ``tests/unittest/_torch/attention/sparse/dsa/test_cpp_custom_ops.py``). """ struct = metadata._ensure_group_remap_struct() leader_of = struct.get("leader_of") diff --git a/tests/integration/test_lists/waives.txt b/tests/integration/test_lists/waives.txt index a8544f6b7093..4061221c2286 100644 --- a/tests/integration/test_lists/waives.txt +++ b/tests/integration/test_lists/waives.txt @@ -312,12 +312,12 @@ test_e2e.py::test_ptp_quickstart_advanced_deepseek_r1_w4afp8_8gpus[DeepSeek-R1-W test_e2e.py::test_ptp_quickstart_bert[TRTLLM-BertForSequenceClassification-bert/bert-base-uncased-yelp-polarity] SKIP (https://nvbugs/6605819) test_e2e.py::test_ptp_quickstart_bert[VANILLA-BertForSequenceClassification-bert/bert-base-uncased-yelp-polarity] SKIP (bug pending, tracked in PR 17414) unittest/_torch/attention/sparse/dsa/test_req_idx_per_token.py::test_on_update_kv_lens_rebuilds_stale_map SKIP (https://nvbugs/6574939) +unittest/_torch/attention/sparse/msa/test_minimax_m3_index_decode_score.py::test_index_decode_score_matches_msa_proxy[dtype0] SKIP (https://nvbugs/6669902) +unittest/_torch/attention/sparse/msa/test_minimax_m3_index_decode_score.py::test_index_decode_score_matches_msa_proxy[dtype1] SKIP (https://nvbugs/6669902) +unittest/_torch/attention/sparse/msa/test_minimax_m3_sparse_attn_decode.py::test_sparse_decode_matches_msa_kernel SKIP (https://nvbugs/6669902) unittest/_torch/attention/sparse/msa/test_msa_backend.py::test_msa_paged_hnd_input_materializes_unaligned_outer_stride SKIP (https://nvbugs/6661846) unittest/_torch/attention/sparse/rocketkv/test_rocketkv.py::test_model[TRTLLM-llama-3.1-model/Llama-3.1-8B-Instruct-pytorch] SKIP (https://nvbugs/6602094) unittest/_torch/attention/sparse/rocketkv/test_rocketkv.py::test_model[VANILLA-llama-3.1-model/Llama-3.1-8B-Instruct-pytorch] SKIP (https://nvbugs/6602094) -unittest/_torch/attention/sparse/test_minimax_m3_index_decode_score.py::test_index_decode_score_matches_msa_proxy[dtype0] SKIP (https://nvbugs/6669902) -unittest/_torch/attention/sparse/test_minimax_m3_index_decode_score.py::test_index_decode_score_matches_msa_proxy[dtype1] SKIP (https://nvbugs/6669902) -unittest/_torch/attention/sparse/test_minimax_m3_sparse_attn_decode.py::test_sparse_decode_matches_msa_kernel SKIP (https://nvbugs/6669902) unittest/_torch/attention/test_attention_backends.py::test_attention_backend[exaone_moe_gqa_swa128-ctx-bf16-HND-p32-v1] SKIP (https://nvbugs/6668773) unittest/_torch/attention/test_attention_backends.py::test_attention_backend[qwen2_0_5b_gqa_hd64-ctx-bf16-HND-p32-v1] SKIP (https://nvbugs/6641268) unittest/_torch/executor/test_overlap_scheduler.py::test_overlap_scheduler_block_reuse_cache_hit SKIP (https://nvbugs/6608387) diff --git a/tests/unittest/_torch/attention/sparse/test_minimax_m3_dense_decode.py b/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_dense_decode.py similarity index 100% rename from tests/unittest/_torch/attention/sparse/test_minimax_m3_dense_decode.py rename to tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_dense_decode.py diff --git a/tests/unittest/_torch/attention/sparse/test_minimax_m3_index_decode_score.py b/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_index_decode_score.py similarity index 100% rename from tests/unittest/_torch/attention/sparse/test_minimax_m3_index_decode_score.py rename to tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_index_decode_score.py diff --git a/tests/unittest/_torch/attention/sparse/test_minimax_m3_msa_selector.py b/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_msa_selector.py similarity index 100% rename from tests/unittest/_torch/attention/sparse/test_minimax_m3_msa_selector.py rename to tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_msa_selector.py diff --git a/tests/unittest/_torch/attention/sparse/test_minimax_m3_sparse_attn_decode.py b/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_sparse_attn_decode.py similarity index 100% rename from tests/unittest/_torch/attention/sparse/test_minimax_m3_sparse_attn_decode.py rename to tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_sparse_attn_decode.py diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py index ddf65eeb146e..32bf6db08c2f 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py @@ -88,6 +88,11 @@ def _force_trtllm_gen_fallback(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("TLLM_FMHA_LIBS", "fallback") +def _fp8_qdq(tensor: torch.Tensor) -> torch.Tensor: + """Apply unit-scale E4M3 quantize-dequantize to a reference tensor.""" + return tensor.to(torch.float8_e4m3fn).to(tensor.dtype) + + # Kernel contract and static selector adapter. @@ -364,6 +369,16 @@ def fused_qkv(self) -> torch.Tensor: ContextScenario(batch_size=3, seq_lens=(64, 96, 128), num_pages=12), id="batch3_var_seq", ), + pytest.param( + ContextScenario( + batch_size=1, + seq_lens=(128,), + num_pages=8, + num_heads=8, + num_kv_heads=8, + ), + id="mha_8q8kv", + ), ] @@ -863,11 +878,18 @@ def test_generation_sparse_mqa_gqa(scenario: GenerationScenario) -> None: scenario, scenario.dtype, ) + reference_q = inputs.q + reference_k_new = inputs.k_new + reference_v_new = inputs.v_new + if scenario.kvcache_dtype == torch.float8_e4m3fn: + reference_q = _fp8_qdq(reference_q) + reference_k_new = _fp8_qdq(reference_k_new) + reference_v_new = _fp8_qdq(reference_v_new) reference_output = _reference_sparse_generation_attention( - inputs.q, + reference_q, kv_caches, - inputs.k_new, - inputs.v_new, + reference_k_new, + reference_v_new, inputs.local_sparse_attn_indices, scenario, ) @@ -888,6 +910,8 @@ def test_generation_sparse_mqa_gqa(scenario: GenerationScenario) -> None: assert output.shape == expected_shape uses_fp8 = scenario.kvcache_dtype == torch.float8_e4m3fn or scenario.fp8_output output_for_comparison = output.float() if uses_fp8 else output + if scenario.fp8_output: + reference_output = reference_output.to(torch.float8_e4m3fn) reference_for_comparison = reference_output.float() if uses_fp8 else reference_output assert torch.isfinite(output_for_comparison).all() if scenario.fp8_output: From 9473bbefd88703b93eb30605321f6cab83534204 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Thu, 3 Sep 2026 08:24:21 -0700 Subject: [PATCH 09/12] [None][test] Preserve RocketKV test unwaive Keep the RocketKV test_model waives removed after rebasing onto main. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- tests/integration/test_lists/waives.txt | 2 -- 1 file changed, 2 deletions(-) diff --git a/tests/integration/test_lists/waives.txt b/tests/integration/test_lists/waives.txt index 4061221c2286..e92b7c878cfa 100644 --- a/tests/integration/test_lists/waives.txt +++ b/tests/integration/test_lists/waives.txt @@ -316,8 +316,6 @@ unittest/_torch/attention/sparse/msa/test_minimax_m3_index_decode_score.py::test unittest/_torch/attention/sparse/msa/test_minimax_m3_index_decode_score.py::test_index_decode_score_matches_msa_proxy[dtype1] SKIP (https://nvbugs/6669902) unittest/_torch/attention/sparse/msa/test_minimax_m3_sparse_attn_decode.py::test_sparse_decode_matches_msa_kernel SKIP (https://nvbugs/6669902) unittest/_torch/attention/sparse/msa/test_msa_backend.py::test_msa_paged_hnd_input_materializes_unaligned_outer_stride SKIP (https://nvbugs/6661846) -unittest/_torch/attention/sparse/rocketkv/test_rocketkv.py::test_model[TRTLLM-llama-3.1-model/Llama-3.1-8B-Instruct-pytorch] SKIP (https://nvbugs/6602094) -unittest/_torch/attention/sparse/rocketkv/test_rocketkv.py::test_model[VANILLA-llama-3.1-model/Llama-3.1-8B-Instruct-pytorch] SKIP (https://nvbugs/6602094) unittest/_torch/attention/test_attention_backends.py::test_attention_backend[exaone_moe_gqa_swa128-ctx-bf16-HND-p32-v1] SKIP (https://nvbugs/6668773) unittest/_torch/attention/test_attention_backends.py::test_attention_backend[qwen2_0_5b_gqa_hd64-ctx-bf16-HND-p32-v1] SKIP (https://nvbugs/6641268) unittest/_torch/executor/test_overlap_scheduler.py::test_overlap_scheduler_block_reuse_cache_hit SKIP (https://nvbugs/6608387) From 66fb06de0601fb64a6929479839f4ca092f72d46 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Fri, 4 Sep 2026 07:43:59 -0700 Subject: [PATCH 10/12] [None][fix] Clarify sparse FP8 contracts Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- docs/source/features/sparse-attention.md | 12 ++++++++++++ .../_torch/attention/backends/fmha/fallback.py | 4 +++- .../backends/sparse/minimax_m3/msa_utils.py | 2 +- .../_torch/attention/sparse/msa/test_msa_backend.py | 2 +- .../_torch/attention/sparse/test_sparse_mqa_gqa.py | 12 ++++++------ .../_torch/attention/test_attention_op_sync.py | 11 +++++++++++ 6 files changed, 34 insertions(+), 9 deletions(-) diff --git a/docs/source/features/sparse-attention.md b/docs/source/features/sparse-attention.md index d47b456444d2..f79b2825c775 100644 --- a/docs/source/features/sparse-attention.md +++ b/docs/source/features/sparse-attention.md @@ -64,6 +64,10 @@ including mixed batches. | Sparse granularity | Token | | Attention semantics | Causal self-attention | +`Input dtype` refers to model-native MLA inputs, which remain BF16. The FP8 +KV-cache entry and any internal FP8 staging do not indicate raw FP8 model-input +support. + See [`test_sparse_mla_forward.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mla_forward.py) for executable sparse MLA examples. @@ -105,6 +109,10 @@ For an FP8 KV cache, token-sparse Q is quantized to E4M3 during QKV preprocessing while the model input remains BF16 or FP16. The path supports both BF16 output with an FP8 KV cache and E4M3 FP8 output. +This is distinct from the block-sparse column: its E4M3 input row is a raw Q/K/V +contract of that dedicated backend. Raw E4M3 fused QKV is not a token-sparse +MQA/GQA input contract. + Backend developers can use [`test_sparse_mqa_gqa.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py) as an executable integration example. @@ -132,6 +140,10 @@ the retained KV cache after prefill to reduce cache size and later decode work. | Sparse granularity | Positive-size blocks expanded to KV-cache pages | | Attention semantics | Causal self-attention | +The E4M3 entries above describe FP8 KV-cache and output paths. The fused model +QKV input remains BF16 or FP16; raw E4M3 fused QKV is not supported by the +page-sparse MHA path. + Backend developers can use [`test_sparse_mha.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mha.py) as an executable integration example. diff --git a/tensorrt_llm/_torch/attention/backends/fmha/fallback.py b/tensorrt_llm/_torch/attention/backends/fmha/fallback.py index 529d7b8eb6ca..f03ae9a01dc1 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/fallback.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/fallback.py @@ -78,7 +78,9 @@ def is_supported( *, phase: Optional[FmhaPhase] = None, ) -> bool: - del q, k, v, phase + del k, v, phase + if q is not None and q.dtype == torch.float8_e4m3fn: + return False return forward_args.attention_mask != CustomAttentionMask.CUSTOM and ( forward_args.update_kv_cache or metadata.is_cross ) diff --git a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/msa_utils.py b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/msa_utils.py index e7e499b7d7a8..6bde824d04d8 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/msa_utils.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/minimax_m3/msa_utils.py @@ -23,7 +23,7 @@ def _install_msa_cutlass_compatibility() -> None: - """Provide the CUTLASS 4.5 names still referenced by the packaged MSA sources.""" + """Provide legacy CuTe aliases still referenced by the packaged MSA sources.""" try: import cutlass.cute as cute except ImportError: diff --git a/tests/unittest/_torch/attention/sparse/msa/test_msa_backend.py b/tests/unittest/_torch/attention/sparse/msa/test_msa_backend.py index 7abd4f07f9af..62884e61a626 100644 --- a/tests/unittest/_torch/attention/sparse/msa/test_msa_backend.py +++ b/tests/unittest/_torch/attention/sparse/msa/test_msa_backend.py @@ -26,7 +26,7 @@ from tensorrt_llm.llmapi.llm_args import MiniMaxM3SparseAttentionConfig -def test_msa_package_availability_installs_cutlass_46_compatibility_aliases(monkeypatch): +def test_msa_package_availability_installs_cutlass_compatibility_aliases(monkeypatch): from tensorrt_llm._torch.attention.backends.sparse.minimax_m3.msa_utils import ( msa_package_available, ) diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py index 32bf6db08c2f..a652548e8578 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py @@ -39,17 +39,17 @@ from utils.util import getSMVersion import tensorrt_llm -from tensorrt_llm._torch.attention_backend.fmha.msa_sparse_gqa import run_msa_sparse_gqa -from tensorrt_llm._torch.attention_backend.interface import ( +from tensorrt_llm._torch.attention.backends.fmha.msa_sparse_gqa import run_msa_sparse_gqa +from tensorrt_llm._torch.attention.backends.interface import ( AttentionForwardArgs, AttentionRuntimeFeatures, ) -from tensorrt_llm._torch.attention_backend.sparse.dsa.kernels import ( +from tensorrt_llm._torch.attention.backends.sparse.dsa.kernels import ( triton_convert_req_index_to_global_index, ) -from tensorrt_llm._torch.attention_backend.sparse.minimax_m3.msa_utils import msa_package_available -from tensorrt_llm._torch.attention_backend.sparse.params import SparseParams -from tensorrt_llm._torch.attention_backend.trtllm import ( +from tensorrt_llm._torch.attention.backends.sparse.minimax_m3.msa_utils import msa_package_available +from tensorrt_llm._torch.attention.backends.sparse.params import SparseParams +from tensorrt_llm._torch.attention.backends.trtllm import ( TrtllmAttention, TrtllmAttentionMetadata, generate_spec_decoding_packed_mask, diff --git a/tests/unittest/_torch/attention/test_attention_op_sync.py b/tests/unittest/_torch/attention/test_attention_op_sync.py index 444733d6a672..ba5c75f3d7b6 100644 --- a/tests/unittest/_torch/attention/test_attention_op_sync.py +++ b/tests/unittest/_torch/attention/test_attention_op_sync.py @@ -44,6 +44,7 @@ from types import SimpleNamespace import pytest +import torch from tensorrt_llm._torch.attention.backends.fmha.fallback import ( _THOP_EXCLUDED_FIELDS, @@ -676,3 +677,13 @@ def test_fallback_support_matches_thop_kv_update_contract(is_cross, update_kv_ca forward_args = AttentionForwardArgs(update_kv_cache=update_kv_cache) assert fmha.is_supported(None, None, None, metadata, forward_args) is expected + + +def test_fallback_rejects_raw_fp8_input(): + """Do not dispatch raw FP8 QKV to the native attention op.""" + fmha = object.__new__(FallbackFmha) + metadata = SimpleNamespace(is_cross=False) + forward_args = AttentionForwardArgs(update_kv_cache=True) + q = torch.empty((1, 128), dtype=torch.float8_e4m3fn) + + assert not fmha.is_supported(q, None, None, metadata, forward_args) From 7b496e56a85a48dd456d1de5c692c66f9d954dcf Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Tue, 8 Sep 2026 02:55:14 +0000 Subject: [PATCH 11/12] [None][test] Fix sparse attention CI coverage Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../attention/sparse/dsa/test_dsa_indexer.py | 2 +- .../attention/sparse/test_sparse_mqa_gqa.py | 37 +++++++------------ 2 files changed, 15 insertions(+), 24 deletions(-) diff --git a/tests/unittest/_torch/attention/sparse/dsa/test_dsa_indexer.py b/tests/unittest/_torch/attention/sparse/dsa/test_dsa_indexer.py index 6ac79f345954..8b2a20ffee82 100644 --- a/tests/unittest/_torch/attention/sparse/dsa/test_dsa_indexer.py +++ b/tests/unittest/_torch/attention/sparse/dsa/test_dsa_indexer.py @@ -651,7 +651,7 @@ def cdiv(a: int, b: int) -> int: def _load_cast_back_from_fp4(): - from test_cute_dsl_fp4_paged_mqa_logits import cast_back_from_fp4 + from .test_cute_dsl_fp4_paged_mqa_logits import cast_back_from_fp4 return cast_back_from_fp4 diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py index a652548e8578..c6a8860a8a5f 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_mqa_gqa.py @@ -468,12 +468,14 @@ def test_prefill_sparse_kv_compaction(scenario: ContextScenario) -> None: # tested on SM100 tested on SM100 # Compute phase Packed prefill, single-token, Packed prefill, single-token, # and linear draft decode; linear multi-query compute, -# tested q_len 1 and 4 and mixed batches; tested; the +# tested q_len 1 and 4 and variable-length batches; +# tested; the # integrated MiniMax-M3 decode uses 1 # Attention type MQA/GQA; Q heads divisible MQA/GQA; Q heads divisible # by KV heads by KV heads -# Q heads per KV head <= 32; tested 2, 3, 4, 8, 2, 4, 8, or 16; tested all -# 16, 24, 31, and 32 integrated MiniMax-M3 uses 16 +# Q heads per KV head <= 32; tested 2, 3, 4, 8, 2, 4, 8, or 16; decode tests cover +# 16, 24, 31, and 32 all; prefill tests cover 8 and 16; +# integrated MiniMax-M3 uses 16 # Q/KV head counts No additional discrete limit; No additional discrete kernel limit; # tested Q={6,8,16,32,48,62,64}, tested Q={4,8,16,32}, KV={1,2} # KV={1,2,4,8} @@ -494,8 +496,9 @@ def test_prefill_sparse_kv_compaction(scenario: ContextScenario) -> None: # Sparse indices int32 physical token indices int32 request-local block indices # per KV head/query; tested per KV head/query; per-token lists, # -1 padding, and physical remap tested -# Sparse Top-K Positive multiple of 4; Prefill kernel accepts 4, 8, 16, -# tested 4, 32, 64, and 128 or 32 and tests cover all; the +# Sparse Top-K Positive multiple of 4; Kernel accepts 4, 8, 16, or 32; +# tested 4, 32, 64, and 128 decode tests cover 4 and 8, prefill +# tests cover 16 and 32; the # integrated MiniMax-M3 path uses 16 # Attention semantics Causal; tested Causal with per-request Q offsets; # bottom-right and custom offsets tested @@ -961,35 +964,23 @@ def test_generation_sparse_mqa_gqa(scenario: GenerationScenario) -> None: ], pytest.param( BlockSparseGqaScenario( - q_lens=(1, 4, 33), - kv_lens=(1152, 1280, 1408), - num_q_heads=8, - num_kv_heads=2, - active_blocks=3, - shuffle_pages=True, - per_token_blocks=True, - ), - id="msa_gqa_mixed_varlen_shuffled_pages_padded_indices", - ), - pytest.param( - BlockSparseGqaScenario( - q_lens=(33,), + q_lens=(1,), kv_lens=(640,), - num_q_heads=4, + num_q_heads=16, num_kv_heads=2, topk=4, ), - id="msa_gqa_2q_per_kv_topk4", + id="msa_gqa_decode_topk4", ), pytest.param( BlockSparseGqaScenario( - q_lens=(33,), + q_lens=(1,), kv_lens=(1152,), - num_q_heads=8, + num_q_heads=16, num_kv_heads=2, topk=8, ), - id="msa_gqa_4q_per_kv_topk8", + id="msa_gqa_decode_topk8", ), pytest.param( BlockSparseGqaScenario( From f305fa8aa7003ec1b8509dda776d926942aafe51 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Tue, 8 Sep 2026 05:16:36 +0000 Subject: [PATCH 12/12] [None][test] Update MSA page table helpers Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../msa/test_minimax_m3_index_decode_score.py | 15 +++------------ .../msa/test_minimax_m3_sparse_attn_decode.py | 15 +++------------ 2 files changed, 6 insertions(+), 24 deletions(-) diff --git a/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_index_decode_score.py b/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_index_decode_score.py index 5a1bc0ef6a8c..ad5c213bae55 100644 --- a/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_index_decode_score.py +++ b/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_index_decode_score.py @@ -35,19 +35,10 @@ def _flat_page_table(block_table: torch.Tensor, kv_lens_cpu: torch.Tensor) -> torch.Tensor: """Flatten a block table into the per-request page ids fmha_sm100 consumes. - build_kv_page_indices reads those ids out of a token-level slot map, so this - rebuilds the map the block table implies and lets the production helper do - the flattening. + The production helper concatenates the valid prefix of each request's + block-id row according to its KV length. """ - batch, max_pages = block_table.shape - intra = torch.arange(PAGE_SIZE, dtype=torch.int32) - req_to_token = (block_table.cpu().to(torch.int32) * PAGE_SIZE).unsqueeze(2) + intra - return build_kv_page_indices( - req_to_token.reshape(batch, max_pages * PAGE_SIZE), - torch.arange(batch, dtype=torch.int32), - kv_lens_cpu, - PAGE_SIZE, - ) + return build_kv_page_indices(block_table.cpu(), kv_lens_cpu, PAGE_SIZE) def _runner(): diff --git a/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_sparse_attn_decode.py b/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_sparse_attn_decode.py index a2b0ac3b270a..4e8496d610a4 100644 --- a/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_sparse_attn_decode.py +++ b/tests/unittest/_torch/attention/sparse/msa/test_minimax_m3_sparse_attn_decode.py @@ -39,19 +39,10 @@ def _flat_page_table(block_table: torch.Tensor, kv_lens_cpu: torch.Tensor) -> torch.Tensor: """Flatten a block table into the per-request page ids fmha_sm100 consumes. - build_kv_page_indices reads those ids out of a token-level slot map, so this - rebuilds the map the block table implies and lets the production helper do - the flattening. + The production helper concatenates the valid prefix of each request's + block-id row according to its KV length. """ - batch, max_pages = block_table.shape - intra = torch.arange(PAGE_SIZE, dtype=torch.int32) - req_to_token = (block_table.cpu().to(torch.int32) * PAGE_SIZE).unsqueeze(2) + intra - return build_kv_page_indices( - req_to_token.reshape(batch, max_pages * PAGE_SIZE), - torch.arange(batch, dtype=torch.int32), - kv_lens_cpu, - PAGE_SIZE, - ) + return build_kv_page_indices(block_table.cpu(), kv_lens_cpu, PAGE_SIZE) def _reference_sparse_decode(