Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 11 additions & 11 deletions 3rdparty/vendor_patches/flashinfer-prims-ts.patch
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
diff --git a/block_sparse.py b/block_sparse.py
index 77a9fc4542ab5fb82aa0df6f7c73e56b7d5fbe89..af132c6a69ce0e4a35903b8e4428f481e949f6b8 100644
index c7e881ca73dfb95314da532912db39fe7b05eb94..e19dc7b1de7fca03c44b8e98217f8d10fb17e26b 100644
--- a/block_sparse.py
+++ b/block_sparse.py
@@ -1,3 +1,4 @@
Expand All @@ -12,15 +12,15 @@ index 77a9fc4542ab5fb82aa0df6f7c73e56b7d5fbe89..af132c6a69ce0e4a35903b8e4428f481

from flashinfer.api_logging import flashinfer_api
-from flashinfer.trace.templates.attention import (
- prims_ts_block_sparse_trace,
- prims_ts_block_sparse_trace_dispatch,
- prims_ts_block_sparse_wrapper_trace_dispatch,
- prims_ts_paged_block_sparse_trace_dispatch,
- prims_ts_paged_block_sparse_wrapper_trace_dispatch,
-)

from ._block_sparse.common import _validate_contiguous_route_mode
from ._block_sparse.config import _validate_block_sparse_static_profile
from ._block_sparse.inspection import (
@@ -220,7 +215,7 @@ class BlockSparseTSWrapper(_BlockSparseWrapperBase):
@@ -239,7 +234,7 @@ class BlockSparseTSWrapper(_BlockSparseWrapperBase):
# previously published revision intact and runnable.
self._plan_state = candidate

Expand All @@ -29,16 +29,16 @@ index 77a9fc4542ab5fb82aa0df6f7c73e56b7d5fbe89..af132c6a69ce0e4a35903b8e4428f481
def run(
self,
q: torch.Tensor,
@@ -302,7 +297,7 @@ class BlockSparseTSWrapper(_BlockSparseWrapperBase):
@@ -361,7 +356,7 @@ class BlockSparseTSWrapper(_BlockSparseWrapperBase):
return self._launch_validated_run(state, run_args, run_stream)


-@flashinfer_api(trace=prims_ts_block_sparse_trace)
-@flashinfer_api(trace=prims_ts_block_sparse_trace_dispatch)
+@flashinfer_api
def block_sparse_attention(
q: torch.Tensor,
k: torch.Tensor,
@@ -512,7 +507,7 @@ class BlockSparsePagedTSWrapper(_BlockSparseWrapperBase):
@@ -607,7 +602,7 @@ class BlockSparsePagedTSWrapper(_BlockSparseWrapperBase):
)
self._plan_state = candidate

Expand All @@ -47,7 +47,7 @@ index 77a9fc4542ab5fb82aa0df6f7c73e56b7d5fbe89..af132c6a69ce0e4a35903b8e4428f481
def run(
self,
q: torch.Tensor,
@@ -618,7 +613,7 @@ class BlockSparsePagedTSWrapper(_BlockSparseWrapperBase):
@@ -727,7 +722,7 @@ class BlockSparsePagedTSWrapper(_BlockSparseWrapperBase):
return self._launch_validated_run(state, run_args, run_stream)


Expand All @@ -57,10 +57,10 @@ index 77a9fc4542ab5fb82aa0df6f7c73e56b7d5fbe89..af132c6a69ce0e4a35903b8e4428f481
q: torch.Tensor,
paged_kv_cache: PagedKVCache,
diff --git a/context.py b/context.py
index 7245bee1a0f725086171c9c5002115757e425d84..47996c867a962a3684c78e07d3a14eebf34b8452 100644
index cea5a41438d9d72e152175999453881cc9c3e5e6..3a93019ff4a5995968c5efa9000e8ebd52c8b424 100644
--- a/context.py
+++ b/context.py
@@ -29,8 +29,7 @@ position is ``q + (S_kv - S_q)`` and ``window_left`` is measured from that
@@ -31,8 +31,7 @@ position is ``q + (S_kv - S_q)`` and ``window_left`` is measured from that
position.

PrimTS context entry points are intentionally excluded from ``fi_trace`` for
Expand All @@ -70,7 +70,7 @@ index 7245bee1a0f725086171c9c5002115757e425d84..47996c867a962a3684c78e07d3a14eeb
"""

from dataclasses import dataclass
@@ -424,7 +423,7 @@ def _validate_device(device: torch.device) -> int:
@@ -426,7 +425,7 @@ def _validate_device(device: torch.device) -> int:
# Rubin runs through the sm_100f family target; a CuTe DSL older than 4.8
# cannot emit for it unless CUTE_DSL_ARCH=sm_100f is set before import.
if capability == (10, 7):
Expand Down
8 changes: 4 additions & 4 deletions 3rdparty/vendor_sources.lock.yaml
Original file line number Diff line number Diff line change
@@ -1,13 +1,13 @@
schema_version: 1
vendors:
flashinfer-prims-ts:
url: https://github.com/yuxianq/flashinfer.git
url: https://github.com/heyuhhh/flashinfer.git
branch: trtllm-prims-ts-dev
commit: bad2bdb15aac3553934e7a7a164dcb2ca4fda7f1
commit: ce4657c2ddfa9c9ba22adbdd1bc5ef9860c96285
source: flashinfer/attention/prims_ts
destination: tensorrt_llm/_torch/attention/backends/prims_ts
include:
- '**/*.py'
patch: 3rdparty/vendor_patches/flashinfer-prims-ts.patch
patch_digest: sha256:0e2f58c6633f57fee03df42049bc78d4d038063b810ad3a2f0ba6d62f8183887
digest: sha256-tree-v1:5766418e4de80e50f6c64a9f3d37076e9b7e815118b5d8a7a22c27d321fd7576
patch_digest: sha256:b590a3e86c8a2a54a8da5a2268aa9e67f8425f401f6b98972841bfe5484b5732
digest: sha256-tree-v1:0946241eec648bba1e95d5dbbe06735e26c9e7e2608461a5b5ec4483cb0cf69c
Original file line number Diff line number Diff line change
Expand Up @@ -117,7 +117,10 @@ The shared `AttentionOp` path is built around three layers:
<p align="center"><sub><em>Figure 1: Framework support for sparse attention in TensorRT LLM.</em></sub></p>

Hook-based `TrtllmAttention` implementations supply `sparse_kv_predict` /
`sparse_attn_predict` and reuse the shared `AttentionOp` stack. RocketKV's
`sparse_attn_predict` and reuse the shared `AttentionOp` stack. Algorithms that
select whole KV blocks instead supply `block_sparse_attn_predict`; their routes
bypass `AttentionOp` and run on the generic block-sparse FMHA described in the
[feature guide](../features/sparse-attention.md#block-sparse-mha-mqa-gqa). 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
Expand All @@ -131,24 +134,55 @@ The current capability matrix is:
| 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) |
| Block-sparse MHA / MQA / GQA | sparse computation (block-level, contiguous Q/K/V without a KV cache) | sparse computation (block-level, paged) |

Dynamic generation-phase KV eviction is tracked as future work.

### Prediction hooks

`TrtllmAttention`-based sparse backends expose two prediction methods that
`TrtllmAttention`-based sparse backends expose three prediction methods that
algorithm-specific subclasses override:

```python
sparse_kv_indices, sparse_kv_offsets = self.sparse_kv_predict(q, k, metadata, forward_args)
sparse_attn_indices, sparse_attn_offsets = self.sparse_attn_predict(q, k, metadata, forward_args)
block_sparse_inputs = self.block_sparse_attn_predict(q, k, v, metadata, forward_args)
```

`hooks.py` writes these results to `SparseRuntimeParams`. SkipSoftmax writes
its thresholds to the same runtime interface consumed by `AttentionOp`.
`AttentionForwardArgs.sparse_backend_args` carries algorithm inputs from the
module to the backend, while `sparse_runtime_params` carries lowered inputs
from the backend to `AttentionOp`.
`prepare_sparse_runtime_params` in `sparse/hooks.py` runs all three hooks once
per call regardless of whether the backend carries `SparseParams`, applies the
SkipSoftmax threshold schedule when the backend carries `SkipSoftmaxParams`,
and returns a new per-call `SparseRuntimeParams` built from the caller's
`AttentionForwardArgs.sparse_runtime_params` plus the hook results. The core
forward assigns the returned carrier back to that field before FMHA dispatch.
Backends that need runtime state outside the three hooks (DSA's auxiliary pool
pointer, DeepSeek-V4's per-token KV lengths) write it into the caller's carrier
before or inside their hooks, and `prepare_sparse_runtime_params` carries those
fields over.
`AttentionForwardArgs.sparse_backend_args` carries
algorithm inputs from the module to the backend, while
`AttentionForwardArgs.sparse_runtime_params` carries the complete lowered state
from the backend through FMHA dispatch to `AttentionOp`.

`SparseRuntimeParams.block_sparse_inputs` is the nested carrier for optional,
algorithm-neutral `BlockSparseForwardInputs`. Only libraries that declare
`supports_block_sparse_inputs` are offered a request that carries routes, so a
dense kernel never silently ignores them; the selected block-sparse FMHA then
validates and consumes the field. `AttentionForwardArgs` defaults
the field to an empty `SparseRuntimeParams()`; the core forward always
overwrites it with the carrier prepared for the current call.

`block_sparse_attn_predict` runs even when the backend has no `SparseParams`.
Its default implementation hands through
`SparseBackendForwardArgs.block_sparse_inputs`, so an attention module that
predicts routes before the core forward only needs to place the complete
payload in `sparse_backend_args`. Algorithms that predict inside the backend
override the hook, read `metadata` for the batch layout and `forward_args` for
per-call state such as `timestep`, and return `None` for dense phases.

The core contract owns this runtime transport and general block-sparse FMHA
execution. Algorithm integrations own their prediction policy, effective Q/K/V
preparation, and any post-processing around the normal core forward.

Different KV heads are allowed to emit different sparse index sets; Q
heads that map to the same KV head share the KV head's sparse pattern.
Expand Down Expand Up @@ -250,8 +284,8 @@ the bottom of the file.
### 2. Prediction module

Create a new backend class inheriting from `TrtllmAttention` in
`tensorrt_llm/_torch/attention/backends/sparse/`. Override one or both
prediction methods. A `VanillaAttention` implementation instead overrides
`tensorrt_llm/_torch/attention/backends/sparse/`. Override one or more of the
three prediction methods. A `VanillaAttention` implementation instead overrides
`_single_request_sparse_kv_predict` and
`_single_request_sparse_attn_predict` with its per-request Python contract.

Expand Down Expand Up @@ -288,6 +322,19 @@ prediction methods. A `VanillaAttention` implementation instead overrides
different index layouts. Match the selected kernel contract; do not
pass request-local block indices to the physical-token path.

**`block_sparse_attn_predict(self, q, k, v, metadata, forward_args)`**

- **Behavior**: return the `BlockSparseForwardInputs` consumed by the
general block-sparse FMHA, or `None` for a dense call.
Comment thread
coderabbitai[bot] marked this conversation as resolved.
- **Outputs**: block geometry plus exactly one route representation
(BSR `block_indptr`/`block_indices` or a packed `exact_block_bits`
bitmask), optional K/V summaries for proxy routes, and optional
`kv_valid_bits` masking ragged KV tails.
- **Default**: hands through `SparseBackendForwardArgs.block_sparse_inputs`,
so modules that predict before the core forward do not override it.
Override it to predict inside the backend from the flattened Q/K/V,
the batch layout in `metadata`, and per-call state in `forward_args`.

Prediction is on the critical path and can dominate latency in
low-latency scenarios. Plan for custom kernels (Triton or CUDA) rather
than relying on generic PyTorch ops.
Expand Down
47 changes: 47 additions & 0 deletions docs/source/features/sparse-attention.md
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,53 @@ Backend developers can use
[`test_sparse_mha.py`](../../../tests/unittest/_torch/attention/sparse/test_sparse_mha.py)
as an executable integration example.

<a id="block-sparse-mha-mqa-gqa"></a>

### Block-sparse MHA/MQA/GQA

The generic block-sparse path executes attention over KV blocks that a sparse
algorithm selects for each KV head. The algorithm hands its routes to the core
forward through the `block_sparse_attn_predict` hook as
`BlockSparseForwardInputs`: canonical BSR (`block_indptr` plus `block_indices`)
or a packed block bitmask, optionally with K/V block summaries so unselected
blocks contribute a proxy instead of being dropped. Only the block-sparse FMHA
library declares `supports_block_sparse_inputs`, so a request that carries
routes is never served by a dense kernel. Both the contiguous prefill path,
used by diffusion models that keep no KV cache, and the paged generation path
are provided by the vendored PrimTS kernels.

This is an attention capability, not a standalone public
`SparseAttentionConfig` algorithm. A user-facing algorithm must also provide
the selector, metadata, and backend integration.

| Parameter | Contiguous prefill | Paged generation |
|---|---|---|
| GPU architecture | SM100 and SM103 | SM100 and SM103 |
| Compute phase | Prefill with separate Q/K/V and no KV cache | Generation with a fixed per-request query length |
| Attention type | MHA, MQA, and GQA | MHA, 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 | Any divisor of the Q head count | Any divisor of the Q head count |
| Head dimensions | Q/K/V: `128` | Q/K/V: `128` |
| Input dtype | BF16 or FP16 | BF16 or FP16 |
| Input layout | Separate Q `[tokens, q_heads, 128]` and K/V `[tokens, kv_heads, 128]` | Fused QKV |
| Output dtype | Model dtype | Model dtype |
| KV-cache dtype | No KV cache | Model dtype |
| KV-cache layout | No KV cache | Paged HND cache; page size `64` or `128` |
| Sparse granularity | Q blocks of `q_block_size` tokens by KV blocks of `8`, `16`, `32`, or a positive multiple of `64` tokens, selected per KV head | KV blocks of a positive multiple of `64` tokens, selected per KV head |
| Sparse routes | BSR or packed bitmask; optional K/V block summaries (proxy routes); optional packed `kv_valid_bits` for ragged KV tails | BSR; live per-request KV lengths and page tables |
| Attention semantics | Dense or causal self-attention; proxy routes require dense | Causal self-attention |

Routes are validated against the kernel's static profile on every call, and an
unsupported request raises instead of degrading to dense attention. A paged
request needs a page that holds at least one 64-token route fragment, which is
why page sizes below `64` are rejected.

Backend developers can use
[`test_prims_ts_block_sparse.py`](../../../tests/unittest/_torch/attention/sparse/test_prims_ts_block_sparse.py)
as an executable integration example; it covers MHA, GQA, and MQA head
topologies, both model dtypes, KV block sizes `64` and `128`, page sizes `64`
and `128`, proxy routes, token-validity masks, and CUDA Graph replay.

<a id="framework-level-sparse-attention"></a>

## Supported Algorithms
Expand Down
61 changes: 54 additions & 7 deletions tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md
Original file line number Diff line number Diff line change
Expand Up @@ -155,9 +155,12 @@ their module-to-backend inputs in a `SparseBackendForwardArgs` subclass and
pass it through the registered `AttentionForwardArgs.sparse_backend_args`
field. For example, DSA owns `DSABackendForwardArgs`, whose indexer
intermediates are consumed by `DSATrtllmAttention.sparse_attn_predict`.
Shared sparse carriers, including `SparseBackendForwardArgs.topk_indices` and
the backend-to-AttentionOp `SparseRuntimeParams`, live in
`attention/backends/sparse/params.py`.
Shared sparse carriers, including `SparseBackendForwardArgs.topk_indices`,
`SparseBackendForwardArgs.block_sparse_inputs`, and the
backend-to-FMHA/`AttentionOp` `SparseRuntimeParams`, live in
`attention/backends/sparse/params.py`. The latter is carried by
`AttentionForwardArgs.sparse_runtime_params` and nests optional general
block-sparse inputs in `SparseRuntimeParams.block_sparse_inputs`.

For MLA-related tasks, first check whether the work fits the current
projection structure, can stay on an existing backend and metadata family, and
Expand Down Expand Up @@ -211,6 +214,35 @@ that file for the current config/backend combinations. Consult the
for the supported attention shapes; do not infer support from algorithm
registration alone.

Block-sparse FMHA is a kernel-library contract rather than a sparse algorithm.
Algorithms lower their live routing state to an algorithm-neutral
`BlockSparseForwardInputs`, nested at
`SparseRuntimeParams.block_sparse_inputs`: block geometry plus either canonical
BSR routes or an exact packed bitmask. Optional K/V summaries enable proxy
routes, and optional token-validity bits mask ragged KV tails. Plans contain
only static format, proxy, geometry, and capacity choices; every run receives
the live routes, summaries, validity bits, page tables, and sequence lengths.

`PrimsTSBlockSparseFmha` owns its wrapper-plan cache by default. Integrations
whose attention layers execute serially may explicitly bind a model-scoped
cache to reuse graph-stable route workspaces across compatible layers. The
cache must not be shared by concurrent forwards; each independent model
component must own separate state.

`TrtllmAttention.block_sparse_attn_predict(q, k, v, metadata, forward_args)`
is the backend hook that produces this payload; `prepare_sparse_runtime_params`
calls it even when the backend has no `SparseParams`. The default hands through
`SparseBackendForwardArgs.block_sparse_inputs`, which lets an attention module
predict routes before the core forward and pass the complete payload in
`AttentionForwardArgs.sparse_backend_args`. Algorithms that predict inside the
backend override the hook and return `None` for dense phases.

The core library owns this general planning, validation, and execution
contract. Algorithm integrations own the surrounding lifecycle: prediction
policy, effective Q/K/V preparation before the core forward, plus any
algorithm-specific post-processing afterward. They route their payload through
these hooks instead of adding algorithm-specific FMHA libraries.

### 2.3 Backend contract

All backends implement the `AttentionBackend` interface.
Expand Down Expand Up @@ -343,13 +375,24 @@ starting with an empty selection cache. `TrtllmAttention` prepares the complete
per-forward state, passes itself to the manager for selection, and then executes
the selected library.

`TLLM_FMHA_LIBS` controls the ordered selection. PrimTS is opt-in because it may
add host overhead; use `TLLM_FMHA_LIBS=+prims_ts` to add it to the defaults or
`TLLM_FMHA_LIBS=fallback` to force the fallback path. Delta entries update the
`TLLM_FMHA_LIBS` controls the ordered selection. Dense PrimTS is opt-in because
it may add host overhead; use `TLLM_FMHA_LIBS=+prims_ts` to add it to the
defaults or `TLLM_FMHA_LIBS=fallback` to force the fallback path. Generic
block-sparse PrimTS remains enabled by default because a dense fallback cannot
preserve its routing semantics. Delta entries update the
default membership and follow canonical registry order, while an exact list
preserves the user-specified order. Each FMHA library exposes `is_available()`
for module/static environment checks and `is_supported()` for per-forward
request checks. For mixed non-MLA batches, the manager checks each active phase
request checks. `AttentionForwardArgs.sparse_runtime_params` is the sole
per-call lowered sparse runtime carrier and defaults to an empty
`SparseRuntimeParams()`. The core forward overwrites that field with the carrier
that `prepare_sparse_runtime_params` builds from the caller's carrier plus the
hook results. The carrier holds both flat `AttentionOp` parameters and optional
`BlockSparseForwardInputs` in its nested `block_sparse_inputs` field. A library
declares `supports_block_sparse_inputs` to receive requests with routes; the
manager skips every other library for such requests, so no dense kernel can
silently drop the routing semantics.
For mixed non-MLA batches, the manager checks each active phase
independently with `is_supported(..., phase=...)`; a phased library accepts only
phases backed by its corresponding `run_*()` entry point.

Expand All @@ -367,6 +410,10 @@ The FMHA package is split by role:
`TrtllmAttention` can pair it with a later causal-generation provider through
`CombinedFmha`.
- `fmha/cute_dsl_mla.py` implements the CuTe DSL MLA decode FMHA library.
- `fmha/prims_ts_block_sparse.py` adapts generic block-sparse requests to the
vendored PrimTS contiguous and paged wrappers. Paged generation passes a
live, zero-copy 2D K-page-table view with its TRT-LLM padded row stride; it
does not stage page tables through CSR metadata.
- `fmha/prims_ts.py` adapts TRT-LLM inputs and paged-cache metadata to the
vendored PrimTS kernels. Before changing the managed source under
`backends/prims_ts`, read the
Expand Down
Loading
Loading