Skip to content
Draft
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
10 changes: 5 additions & 5 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
branch: trtllm-prims-ts-dev
commit: e500966b575ab83db7c0e84e5a0f8fde6a4f3505
url: https://github.com/heyuhhh/flashinfer.git
branch: yuhangh/tmp-sol-attn-trtllm-dev
commit: 61454c5ce6e9f020a2158f24059d0b91322e3e77
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:e9af5482f6406af3128e711c3fb2d044359fb6e1d1bc9907743dc86d006d23ac
patch_digest: sha256:b590a3e86c8a2a54a8da5a2268aa9e67f8425f401f6b98972841bfe5484b5732
digest: sha256-tree-v1:e89b89471aac59e2689ed0fb8312cef77b08e1b46115e7668d751b9f402a29ba
Original file line number Diff line number Diff line change
Expand Up @@ -136,19 +136,48 @@ 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 never assigns that field; it dispatches with the returned carrier.
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`. The selected general
block-sparse FMHA validates and consumes that field; dense FMHA libraries reject
it instead of silently ignoring its routes. `AttentionForwardArgs` defaults
the field to an empty `SparseRuntimeParams()`; the core forward always
dispatches 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 @@ -288,6 +317,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.
- **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
15 changes: 9 additions & 6 deletions docs/source/models/visual-generation.md
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ Models are auto-detected from the checkpoint directory. Diffusers-format models

[^1]: FLUX models use embedded guidance and do not have a separate negative prompt path, so CFG parallelism is not applicable.

[^2]: `FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers` — VSA-fine-tuned checkpoint with learned sparse-attention gates. Requires `CUTEDSL` on Blackwell sm_100+ (falls back to dense SDPA on older hardware). Ring and Attention2D not supported (no LSE output); Ulysses supported.
[^2]: `FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers` — VSA-fine-tuned checkpoint with learned sparse-attention gates. Supports the `CUTEDSL` and `TRTLLM` attention backends; each uses its block-sparse fine stage when supported and otherwise uses its dense path with compact Q/K/V. Ring and Attention2D are not supported (no LSE output); Ulysses is supported.

[^3]: Wan 2.2 has two stage transformers; TeaCache requires explicit `teacache.coefficients` (high-noise) and `teacache.coefficients_2` (low-noise). There is no built-in coefficient table for Wan 2.2.

Expand Down Expand Up @@ -408,12 +408,15 @@ args = VisualGenArgs(

### Video Sparse Attention (VSA)

VSA reduces the compute cost of self-attention in video diffusion models by selectively attending to only the most relevant spatial-temporal blocks. It uses a two-branch design: a lightweight coarse mean-pool branch computes block-level attention scores to identify the top-K most relevant token blocks, then a fine branch runs a block-sparse CuTe kernel over only those blocks. The two outputs are blended with learned gates.
VSA reduces the compute cost of self-attention in video diffusion models by selectively attending to only the most relevant spatial-temporal blocks. It uses a two-branch design: a lightweight coarse mean-pool branch computes block-level attention scores to identify the top-K most relevant token blocks, then a fine branch runs the selected backend's block-sparse kernel over only those blocks. The two outputs are blended with learned gates.

VisualGen owns VSA route prediction and coarse/fine post-processing. With the `TRTLLM` backend, it nests the predicted routes in `SparseRuntimeParams.block_sparse_inputs` and passes those precomputed runtime parameters through the normal core attention forward. The core `PrimsTSBlockSparseFmha` owns the general block-sparse execution contract; it does not own VSA-specific prediction or blending.

**Requirements:**
- VSA-fine-tuned checkpoint: [`FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers`](https://huggingface.co/FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers). Standard Wan checkpoints do not have the learned VSA gates.
- Blackwell GPU (sm_100+) for the CuTe JIT kernel. Falls back to dense SDPA on older hardware with no accuracy loss.
- `CUTEDSL` attention backend.
- `CUTEDSL` or `TRTLLM` attention backend. `CUTEDSL` uses the CuTe DSL fine-stage kernel; `TRTLLM` lowers the selected blocks through the generic PrimTS block-sparse FMHA contract.
- A supported CUDA device and tensor shape for the selected block-sparse kernel. When that kernel is unavailable or the input is outside its supported envelope, the fine branch uses the selected backend's compact dense path (`SDPA` for `CUTEDSL`, TRTLLM attention for `TRTLLM`).
- VSA cannot be combined with `quant_attention_config`.
- Not compatible with Ring attention or Attention2D (VSA does not produce per-split LSE). Ulysses is supported.

**`vsa_sparsity`** controls the fraction of K/V blocks skipped in the fine branch (0.0 = dense, 0.9 = 90% blocks skipped). Higher sparsity gives more speedup at the cost of some quality.
Expand All @@ -427,7 +430,7 @@ from tensorrt_llm.visual_gen.args import AttentionConfig, VideoSparseAttentionCo
args = VisualGenArgs(
model="FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers",
attention_config=AttentionConfig(
backend="CUTEDSL",
backend="TRTLLM", # Use "CUTEDSL" for the CuTe DSL fine-stage kernel.
sparse_attention_config=VideoSparseAttentionConfig(vsa_sparsity=0.9),
),
)
Expand All @@ -437,7 +440,7 @@ YAML (for use with `--visual_gen_args` or `trtllm-serve`):

```yaml
attention_config:
backend: CUTEDSL
backend: TRTLLM # CUTEDSL is also supported.
sparse_attention_config:
algorithm: vsa
vsa_sparsity: 0.90
Expand Down
82 changes: 79 additions & 3 deletions docs/source/visual-gen/features/sparse-attention.md
Original file line number Diff line number Diff line change
Expand Up @@ -7,20 +7,22 @@ This page is an unindexed draft until the VisualGen documentation hub is introdu
- [Overview](#overview)
- [Algorithms](#algorithms)
- [Skip Softmax Attention](#skip-softmax-attention)
- [SOL Attention](#sol-attention)
- [Video Sparse Attention (VSA)](#video-sparse-attention-vsa)

## Overview

Visual generation models naturally operate on long image or video token sequences. Each denoising step is closer to a full-context prefill pass than to autoregressive decoding, and attention can dominate runtime for high-resolution image generation or long video generation.

Sparse attention in VisualGen is configured through `VisualGenArgs.attention_config.sparse_attention_config`. The user-facing config stays in VisualGen args or model config. Checkpoint calibration metadata remains internal and is lowered into per-attention-backend `SparseParams` when each attention module is constructed.
Sparse attention in VisualGen is configured through `VisualGenArgs.attention_config.sparse_attention_config`. The user-facing config stays in VisualGen args or model config, while `attention_config.backend` selects the kernel family. Algorithms produce their block-sparse routes through the core `block_sparse_attn_predict` hook: a backend either predicts inside that hook from the flattened Q/K/V, or predicts before the core forward and hands the complete `BlockSparseForwardInputs` through `AttentionForwardArgs.sparse_backend_args`, which the default hook passes through. `SparseRuntimeParams` is the single lowered runtime carrier passed as `AttentionForwardArgs.sparse_runtime_params`; its optional `block_sparse_inputs` field nests the algorithm-neutral routes for the general block-sparse FMHA. `None` means prediction has not run, while an empty `SparseRuntimeParams()` records that prediction ran without a sparse payload.

### Algorithms

| `algorithm` | Config class | Status |
|---|---|---|
| `skip_softmax` | `SkipSoftmaxAttentionConfig` | Supported |
| VSA | TBD | TODO |
| `vsa` | `VideoSparseAttentionConfig` | Supported (`CUTEDSL`, `TRTLLM`) |
| `sol_attn` | `SolAttentionConfig` | Experimental (`TRTLLM`) |

## Skip Softmax Attention

Expand Down Expand Up @@ -214,6 +216,80 @@ attention_config:

Graphs are captured lazily. The first denoising step seen for a given tensor shape and sparse-attention phase captures a graph; later steps with the same shape and phase replay that graph. When denoising crosses the cutoff, the phase key changes, so VisualGen captures a second graph for the enabled phase instead of replaying the graph from the disabled phase.

## SOL Attention

SOL is a two-stage self-attention algorithm. A TRT-LLM-owned predictor first
produces an exact block bitmask and K/V proxy summaries from Q/K/V. The shared
`PrimsTSBlockSparseFmha` library then executes that route from the shared
`SparseRuntimeParams`. `SOLTrtllmAttention` is only the VisualGen bridge;
SOL does not add an algorithm-specific core attention backend or FMHA library.

Configure SOL with `SolAttentionConfig` and the `TRTLLM` backend:

```python
from tensorrt_llm.visual_gen import AttentionConfig, SolAttentionConfig

attention_config = AttentionConfig(
backend="TRTLLM",
sparse_attention_config=SolAttentionConfig(
tau=1.0,
disabled_until_timestep=0.6,
dense_layers="0,2-4",
),
)
```

The equivalent YAML is:

```yaml
attention_config:
backend: TRTLLM
sparse_attention_config:
algorithm: sol_attn
tau: 1.0
disabled_until_timestep: 0.6
dense_layers: "0,2-4"
```

- `tau` is the finite float32 routing threshold consumed by the predictor.
- `disabled_until_timestep` is an optional normalized cutoff in `(0, 1]`.
Attention stays dense while the current timestep is greater than or equal to
the cutoff and switches to SOL below it.
- `dense_layers` is an optional comma-separated list of zero-based layer indices
and inclusive ranges that always use dense attention.

The initial SOL envelope is full-mask BF16 self-attention on SM100 or SM103,
with 4-D BSHD Q/K/V tensors, equal Q/K/V shapes, and head dimension 128.
`SOLTrtllmAttention` overrides the core `block_sparse_attn_predict` hook, so
prediction runs inside the core forward from the flattened Q/K/V, the batch
layout in the attention metadata, and the `timestep` in the forward arguments;
dense layers and dense timestep phases return no routes. The VisualGen wrapper
compacts fused projection split views once and shares those tensors between
prediction and the block-sparse FMHA. Cross-attention, context parallelism,
attention quantization, and unsupported tensor envelopes raise an error instead
of silently falling back to dense attention.
SOL uses a host-side graph break to prepare and own predictor plans, so
`torch_compile_config.enable_fullgraph=True` is not supported; keep the default
`False` setting.

When a cutoff is configured, VisualGen includes the dense-or-sparse phase in
the CUDA Graph key. Each SOL backend prepares that phase during graph warmup
and reuses it during capture, while predictor route buffers remain stable for
replay. A dense capture therefore cannot be reused for the sparse phase.

## Video Sparse Attention (VSA)

TODO
VSA combines a coarse mean-pooled branch with a top-K block-sparse fine branch. Select either `CUTEDSL` for the CuTe DSL kernel or `TRTLLM` for PrimTS block-sparse attention. If the selected sparse kernel is unavailable or the known VSA tensor envelope is not met, the fine branch uses the compact Q/K/V tensors with that backend's dense path. VSA cannot be combined with `quant_attention_config`.

VSA retains shape-dependent metadata and route tensors so CUDA Graph replay can reuse stable addresses. A pipeline instance accepts up to 16 distinct VSA shape profiles; reuse configured resolution/frame profiles or restart the pipeline before serving additional shapes.

Both VSA backends share one VisualGen-owned predictor and identical
post-processing. The `TRTLLM` path runs the coarse stage before the core
forward, hands the predicted `BlockSparseForwardInputs` (including the
tile-padding validity bits only the VSA predictor knows) through
`sparse_backend_args`, lets the default core prediction hook pass them to the
general block-sparse FMHA, and then blends the fine and coarse outputs. Its
compact dense fallback passes no sparse inputs, so the core runs dense attention
and VSA post-processing still runs. `CUTEDSL` retains only its backend-specific
fine-attention execution. The core FMHA registry owns the reusable block-sparse
implementation rather than a VSA-specific lifecycle.
Loading
Loading