diff --git a/3rdparty/vendor_patches/flashinfer-prims-ts.patch b/3rdparty/vendor_patches/flashinfer-prims-ts.patch index 44f3e7a4581d..b5d4bc75155c 100644 --- a/3rdparty/vendor_patches/flashinfer-prims-ts.patch +++ b/3rdparty/vendor_patches/flashinfer-prims-ts.patch @@ -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 @@ @@ -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 @@ -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 @@ -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) @@ -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 @@ -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): diff --git a/3rdparty/vendor_sources.lock.yaml b/3rdparty/vendor_sources.lock.yaml index fbfbdf4b2e49..8edf64bc8852 100644 --- a/3rdparty/vendor_sources.lock.yaml +++ b/3rdparty/vendor_sources.lock.yaml @@ -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 diff --git a/docs/source/developer-guide/sparse-attention-development-guide.md b/docs/source/developer-guide/sparse-attention-development-guide.md index 32d103b50058..cfcbae3d41de 100644 --- a/docs/source/developer-guide/sparse-attention-development-guide.md +++ b/docs/source/developer-guide/sparse-attention-development-guide.md @@ -117,7 +117,10 @@ The shared `AttentionOp` path is built around three layers:

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

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 @@ -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. @@ -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. @@ -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. +- **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. diff --git a/docs/source/features/sparse-attention.md b/docs/source/features/sparse-attention.md index f79b2825c775..28302a107c19 100644 --- a/docs/source/features/sparse-attention.md +++ b/docs/source/features/sparse-attention.md @@ -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. + + +### 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. + ## Supported Algorithms diff --git a/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md b/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md index e56bbb5fa24a..a4513b172642 100644 --- a/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md +++ b/tensorrt_llm/_torch/attention/ATTENTION_DEVELOPER_GUIDE.md @@ -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 @@ -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. @@ -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. @@ -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 diff --git a/tensorrt_llm/_torch/attention/backends/fmha/fallback.py b/tensorrt_llm/_torch/attention/backends/fmha/fallback.py index f03ae9a01dc1..2b30e94b9f9b 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/fallback.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/fallback.py @@ -40,10 +40,11 @@ _THOP_EXCLUDED_FIELDS: frozenset = frozenset( { "sparse_backend_args", # consumed by sparse prediction before the attention op + "block_sparse_inputs", # consumed by the selected block-sparse FMHA "attention_mask_data", # custom-mask code path "out_scale_sf", # promoted into ``out_scale`` in ``TrtllmAttention.forward`` for NVFP4 path "skip_mla_rope_generation", # handled in ``TrtllmAttention.forward`` for the test-only MLA path - "timestep", # used to populate skip-softmax params in ``TrtllmAttention.forward`` + "timestep", # consumed by sparse prediction before FMHA dispatch } ) @@ -81,9 +82,11 @@ def is_supported( 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 - ) + if forward_args.attention_mask == CustomAttentionMask.CUSTOM: + return False + if not forward_args.update_kv_cache and not metadata.is_cross: + return False + return True def forward( self, @@ -218,7 +221,7 @@ def forward( forward_args.sparse_runtime_params.sparse_attn_indices_block_size ), sparse_attn_kv_lens=forward_args.sparse_runtime_params.sparse_attn_kv_lens, - aux_kv_cache_pool_ptr=(forward_args.sparse_runtime_params.aux_kv_cache_pool_ptr), + aux_kv_cache_pool_ptr=forward_args.sparse_runtime_params.aux_kv_cache_pool_ptr, skip_softmax_threshold_scale_factor_prefill=( forward_args.sparse_runtime_params.threshold_scale_factor_prefill ), diff --git a/tensorrt_llm/_torch/attention/backends/fmha/interface.py b/tensorrt_llm/_torch/attention/backends/fmha/interface.py index 1e5bce7161aa..c576a43c8c4b 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/interface.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/interface.py @@ -40,6 +40,7 @@ class Fmha(ABC): """Common runtime contract for TRT-LLM attention FMHA libraries.""" supports_skip_correction = False + supports_block_sparse_inputs = False def __init__(self, attn: "TrtllmAttention"): self._attn_ref: weakref.ReferenceType["TrtllmAttention"] = weakref.ref(attn) diff --git a/tensorrt_llm/_torch/attention/backends/fmha/manager.py b/tensorrt_llm/_torch/attention/backends/fmha/manager.py index 07837c39eadd..7488c97c59f4 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/manager.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/manager.py @@ -134,6 +134,7 @@ class _FmhaCacheKey(NamedTuple): generation_seq_len_q: int attention_mask_type: AttentionMaskType use_spec_decoding: bool + has_block_sparse_inputs: bool # LoRA can change the effective output from packed NVFP4 to unpacked BF16 # without changing the request shape. Keep those selection regimes apart. output_dtype: torch.dtype | None @@ -366,12 +367,14 @@ def _make_cache_key( generation_seq_len_q, _FMHA_CACHE_SEQ_LEN_Q_GRID ) + block_sparse_inputs = forward_args.sparse_runtime_params.block_sparse_inputs return _FmhaCacheKey( context_batch_size=context_batch_size, generation_batch_size=generation_batch_size, generation_seq_len_q=generation_seq_len_q, attention_mask_type=attention_mask_type, use_spec_decoding=metadata.use_spec_decoding, + has_block_sparse_inputs=block_sparse_inputs is not None, output_dtype=output_dtype, output_sf_dtype=output_sf_dtype, ) @@ -452,9 +455,12 @@ def _select_non_mla( if not has_context and not has_generation: return None + has_block_sparse_inputs = forward_args.sparse_runtime_params.block_sparse_inputs is not None context_fmha = None generation_fmha = None for fmha in self.fmha_libs: + if has_block_sparse_inputs and not fmha.supports_block_sparse_inputs: + continue if fmha.is_supported(q, k, v, metadata, forward_args): return fmha diff --git a/tensorrt_llm/_torch/attention/backends/fmha/prims_ts.py b/tensorrt_llm/_torch/attention/backends/fmha/prims_ts.py index afad6479067d..3c176085f30b 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/prims_ts.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/prims_ts.py @@ -20,7 +20,7 @@ import math from importlib import import_module from importlib.metadata import PackageNotFoundError, version -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING, Any, Optional import torch from packaging.version import InvalidVersion, Version @@ -40,7 +40,7 @@ from .interface import FmhaPhase from .phased import FmhaParams, PhasedFmha -from .utils import get_kv_page_offset +from .utils import get_kv_page_offset, get_multi_processor_count_for_device if TYPE_CHECKING: from tensorrt_llm._torch.attention.backends.prims_ts.context import BatchPrefillPagedTSWrapper @@ -59,6 +59,96 @@ _WORKSPACE_ALIGNMENT = 32 +def get_paged_kv_storage_unsupported_reason( + attn: "TrtllmAttention", + metadata: "TrtllmAttentionMetadata", +) -> Optional[str]: + """Return why the TRT-LLM paged KV storage cannot feed a fixed page-table kernel.""" + if metadata.kv_cache_manager is None: + return "a KV cache manager is required." + if metadata.kv_cache_block_offsets is None: + return "paged KV-cache block offsets are required." + if metadata.host_kv_cache_pool_pointers is None: + return "KV-cache pool pointers are required." + pool_mapping = metadata.host_kv_cache_pool_mapping + if pool_mapping is None: + return "KV-cache pool mapping is required." + if metadata.kv_layout != "HND": + return "only HND KV-cache layout is supported." + manager = metadata.kv_cache_manager + if isinstance(manager, KVCacheManagerV2): + if manager.enable_swa_scratch_reuse: + return "KVCacheManagerV2 SWA scratch reuse is not supported." + elif isinstance(manager, KVCacheManager): + if manager.num_pools != 1: + return "KVCacheManagerV1 with multiple memory pools is not supported." + local_layer_idx = attn.local_layer_idx + if ( + pool_mapping.ndim != 2 + or pool_mapping.shape[1] < 2 + or local_layer_idx is None + or not 0 <= local_layer_idx < pool_mapping.shape[0] + ): + return "KVCacheManagerV1 has an invalid layer-to-pool mapping." + pool_index = int(pool_mapping[local_layer_idx, 0]) + layer_idx_in_pool = int(pool_mapping[local_layer_idx, 1]) + if pool_index != 0 or not 0 <= layer_idx_in_pool < manager.num_local_layers: + return "KVCacheManagerV1 has an invalid layer-to-pool mapping." + else: + return f"unsupported KV cache manager {type(manager).__name__}." + return None + + +def get_paged_kv_policy_unsupported_reason( + attn: "TrtllmAttention", + metadata: "TrtllmAttentionMetadata", +) -> Optional[str]: + """Return why the request's decoding policy is outside the fixed page-table envelope.""" + if metadata.beam_width != 1: + return "beam search is not supported." + if ( + metadata.is_spec_decoding_enabled + or metadata.use_spec_decoding + or metadata.is_spec_dec_tree + or metadata.is_spec_dec_dynamic_tree + ): + return "speculative decoding is not supported by the initial adapter." + position_embedding_type = int(attn.position_embedding_type) + if position_embedding_type in (4, 5, 6, 7, 10): + return f"position embedding type {position_embedding_type} is not supported." + try: + quant_mode = QuantMode(attn.quant_mode) + except (TypeError, ValueError): + return "invalid KV-cache quantization mode." + if quant_mode.has_kv_cache_quant(): + return "quantized KV cache is not supported by the initial adapter." + return None + + +def get_attention_feature_unsupported_reason( + metadata: "TrtllmAttentionMetadata", + forward_args: "AttentionForwardArgs", +) -> Optional[str]: + """Return which optional attention feature the fused-kernel adapters do not implement.""" + if metadata.helix_position_offsets is not None: + return "Helix parallelism is not supported." + if forward_args.relative_attention_bias is not None: + return "relative attention bias is not supported." + if forward_args.attention_sinks is not None: + return "attention sinks are not supported." + if forward_args.attention_mask_data is not None: + return "custom attention masks are not supported." + if forward_args.enable_dsv4_epilogue_fusion: + return "DSv4 epilogue fusion is not supported." + if ( + forward_args.sage_attn_num_elts_per_blk_q > 0 + or forward_args.sage_attn_num_elts_per_blk_k > 0 + or forward_args.sage_attn_num_elts_per_blk_v > 0 + ): + return "SageAttention is not supported." + return None + + class PrimsTSFmha(PhasedFmha): """Blackwell task-scheduled paged context and decode FMHA library.""" @@ -180,6 +270,7 @@ def _is_supported_with_reason( phase: Optional[FmhaPhase] = None, ) -> tuple[bool, str]: """Return a conservative, side-effect-free whole-request support decision.""" + sparse_runtime_params = fwd.sparse_runtime_params # PrimTS prepares workspace for every active request phase before # dispatch. Accept the phased dispatcher keyword, but do not narrow # support until that preparation is phase-aware too. @@ -194,40 +285,9 @@ def _is_supported_with_reason( return False, "only fused QKV input is supported." if meta.is_cross: return False, "cross attention is not supported." - if meta.kv_cache_manager is None: - return False, "a KV cache manager is required." - if meta.kv_cache_block_offsets is None: - return False, "paged KV-cache block offsets are required." - if meta.host_kv_cache_pool_pointers is None: - return False, "KV-cache pool pointers are required." - if meta.host_kv_cache_pool_mapping is None: - return False, "KV-cache pool mapping is required." - if meta.kv_layout != "HND": - return False, "only HND KV-cache layout is supported." - kv_cache_manager = meta.kv_cache_manager - if isinstance(kv_cache_manager, KVCacheManagerV2): - if kv_cache_manager.enable_swa_scratch_reuse: - return False, "KVCacheManagerV2 SWA scratch reuse is not supported." - elif isinstance(kv_cache_manager, KVCacheManager): - if kv_cache_manager.num_pools != 1: - return False, "KVCacheManagerV1 with multiple memory pools is not supported." - pool_mapping = meta.host_kv_cache_pool_mapping - local_layer_idx = attn.local_layer_idx - num_local_layers = kv_cache_manager.num_local_layers - if ( - pool_mapping.ndim != 2 - or pool_mapping.shape[1] < 2 - or local_layer_idx is None - or local_layer_idx < 0 - or local_layer_idx >= pool_mapping.shape[0] - ): - return False, "KVCacheManagerV1 has an invalid layer-to-pool mapping." - pool_index = int(pool_mapping[local_layer_idx, 0]) - layer_idx_in_pool = int(pool_mapping[local_layer_idx, 1]) - if pool_index != 0 or not 0 <= layer_idx_in_pool < num_local_layers: - return False, "KVCacheManagerV1 has an invalid layer-to-pool mapping." - else: - return False, f"unsupported KV cache manager {type(kv_cache_manager).__name__}." + storage_reason = get_paged_kv_storage_unsupported_reason(attn, meta) + if storage_reason is not None: + return False, storage_reason output = fwd.output if output is None: @@ -240,38 +300,15 @@ def _is_supported_with_reason( if attn.sparse_params is not None: return False, "sparse attention is not supported." if ( - fwd.sparse_runtime_params.sparse_kv_indices is not None - or fwd.sparse_runtime_params.sparse_attn_indices is not None + sparse_runtime_params.sparse_kv_indices is not None + or sparse_runtime_params.sparse_attn_indices is not None ): return False, "sparse attention metadata is not supported." if meta.num_sparse_topk > 0: return False, "sparse attention metadata is not supported." - if meta.helix_position_offsets is not None: - return False, "Helix parallelism is not supported." - if fwd.relative_attention_bias is not None: - return False, "relative attention bias is not supported." - if fwd.attention_sinks is not None: - return False, "attention sinks are not supported." - if fwd.attention_mask_data is not None: - return False, "custom attention masks are not supported." - if fwd.enable_dsv4_epilogue_fusion: - return False, "DSv4 epilogue fusion is not supported." - if ( - fwd.sage_attn_num_elts_per_blk_q > 0 - or fwd.sage_attn_num_elts_per_blk_k > 0 - or fwd.sage_attn_num_elts_per_blk_v > 0 - ): - return False, "SageAttention is not supported." - - if meta.beam_width != 1: - return False, "beam search is not supported." - if ( - meta.is_spec_decoding_enabled - or meta.use_spec_decoding - or meta.is_spec_dec_tree - or meta.is_spec_dec_dynamic_tree - ): - return False, "speculative decoding is not supported by the initial adapter." + feature_reason = get_attention_feature_unsupported_reason(meta, fwd) + if feature_reason is not None: + return False, feature_reason try: mask_type = AttentionMaskType(fwd.mask_type) @@ -279,17 +316,9 @@ def _is_supported_with_reason( return False, "the attention mask is not causal or dense." if mask_type not in (AttentionMaskType.causal, AttentionMaskType.padding): return False, f"attention mask type {mask_type} is not supported." - - position_embedding_type = int(attn.position_embedding_type) - if position_embedding_type in (4, 5, 6, 7, 10): - return False, f"position embedding type {position_embedding_type} is not supported." - - try: - quant_mode = QuantMode(attn.quant_mode) - except (TypeError, ValueError): - return False, "invalid KV-cache quantization mode." - if quant_mode.has_kv_cache_quant(): - return False, "quantized KV cache is not supported by the initial adapter." + policy_reason = get_paged_kv_policy_unsupported_reason(attn, meta) + if policy_reason is not None: + return False, policy_reason input_type = fwd.attention_input_type if input_type not in ( @@ -432,6 +461,27 @@ def _get_fixed_block_tables( ) return block_tables[:batch_size, 0, :] + def _get_generation_workspace_layout( + self, + dtype: torch.dtype, + num_requests: int, + num_tokens: int, + ) -> dict[str, int]: + """Return the shared TRT-LLM generation preprocessing layout.""" + + return thop.get_trtllm_gen_generation_workspace_layout( + dtype, + num_requests, + num_tokens, + self.attn.num_heads, + self.attn.head_dim, + self.attn.rope_dim, + self.attn.num_kv_heads, + 0, + False, + skip_fmha_workspace=True, + ) + @staticmethod def _get_sequence_lengths( sequence_lengths: torch.Tensor, @@ -626,9 +676,7 @@ def prepare_workspace( metadata.num_generations > 0 and input_type != AttentionInputType.context_only ) if self._multi_processor_count is None: - self._multi_processor_count = torch.cuda.get_device_properties( - q.device - ).multi_processor_count + self._multi_processor_count = get_multi_processor_count_for_device(q.device.index) required_preprocess_bytes = 0 if has_context and not self.attn.is_mla_enable: @@ -652,17 +700,10 @@ def prepare_workspace( if input_type == AttentionInputType.generation_only else q.shape[0] - int(metadata.num_ctx_tokens) ) - generation_layout = thop.get_trtllm_gen_generation_workspace_layout( + generation_layout = self._get_generation_workspace_layout( q.dtype, int(metadata.num_generations), num_gen_tokens_for_layout, - self.attn.num_heads, - self.attn.head_dim, - self.attn.rope_dim, - self.attn.num_kv_heads, - 0, - False, - skip_fmha_workspace=True, ) required_preprocess_bytes = max( required_preprocess_bytes, int(generation_layout["total_size"]) @@ -956,34 +997,17 @@ def run_context(self, params: FmhaParams) -> None: skip_fmha_workspace=True, ) - def run_generation(self, params: FmhaParams) -> None: - if params.qkv_input is None or params.context_buf is None: - raise RuntimeError("PrimTS decode requires QKV input and an output buffer.") - if params.sequence_lengths is None: - raise RuntimeError("PrimTS decode requires sequence lengths.") - if self._multi_processor_count is None: - raise RuntimeError("PrimTS decode workspace was not prepared.") + def _run_generation_preprocess(self, params: FmhaParams) -> tuple[Any, ...]: + """Run the shared TRT-LLM generation QKV and cache preprocessing.""" + if self._multi_processor_count is None: + raise RuntimeError("PrimTS generation workspace was not prepared.") attn = params.attn meta = params.meta fwd = params.fwd rope_params = attn.rope_params - batch_size = params.batch_size attention_chunk_size = attn.attention_chunk_size or 0 - ( - q_processed, - kv_pool, - block_tables, - _kv_scale_pool, - _bmm1_scale, - _bmm2_scale, - fmha_workspace, - _cu_seqlens, - _max_q_len, - _max_kv_len, - window_left, - is_multi_token_gen, - ) = thop.trtllm_gen_generation_preprocess( + return thop.trtllm_gen_generation_preprocess( params.qkv_input, params.workspace, params.sequence_lengths, @@ -1008,7 +1032,7 @@ def run_generation(self, params: FmhaParams) -> None: params.max_attention_window_size, params.cyclic_attention_window_size, params.num_tokens, - batch_size, + params.batch_size, params.input_seq_length, params.max_past_kv_length, rope_params.dim, @@ -1029,6 +1053,33 @@ def run_generation(self, params: FmhaParams) -> None: False, skip_fmha_workspace=True, ) + + def run_generation(self, params: FmhaParams) -> None: + if params.qkv_input is None or params.context_buf is None: + raise RuntimeError("PrimTS decode requires QKV input and an output buffer.") + if params.sequence_lengths is None: + raise RuntimeError("PrimTS decode requires sequence lengths.") + if self._multi_processor_count is None: + raise RuntimeError("PrimTS decode workspace was not prepared.") + + attn = params.attn + meta = params.meta + fwd = params.fwd + batch_size = params.batch_size + ( + q_processed, + kv_pool, + block_tables, + _kv_scale_pool, + _bmm1_scale, + _bmm2_scale, + fmha_workspace, + _cu_seqlens, + _max_q_len, + _max_kv_len, + window_left, + is_multi_token_gen, + ) = self._run_generation_preprocess(params) if fmha_workspace.numel() != 0: raise RuntimeError("PrimTS generation preprocessing returned an FMHA workspace.") if is_multi_token_gen: diff --git a/tensorrt_llm/_torch/attention/backends/fmha/prims_ts_block_sparse.py b/tensorrt_llm/_torch/attention/backends/fmha/prims_ts_block_sparse.py new file mode 100644 index 000000000000..87ecb7989e11 --- /dev/null +++ b/tensorrt_llm/_torch/attention/backends/fmha/prims_ts_block_sparse.py @@ -0,0 +1,586 @@ +# 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. + +"""TRT-LLM FMHA adapter for the vendored PrimTS block-sparse kernels.""" + +import math +from dataclasses import dataclass, fields +from typing import TYPE_CHECKING, Literal, cast + +import torch + +from tensorrt_llm._torch.attention.backends.interface import ( + AttentionForwardArgs, + AttentionInputType, + PredefinedAttentionMask, +) +from tensorrt_llm._torch.attention.backends.prims_ts._block_sparse.config import ( + _validate_block_sparse_static_profile, +) +from tensorrt_llm._torch.attention.backends.sparse.params import BlockSparseForwardInputs +from tensorrt_llm.logger import logger + +from .interface import FmhaPhase +from .phased import FmhaParams +from .prims_ts import ( + PrimsTSFmha, + get_attention_feature_unsupported_reason, + get_paged_kv_policy_unsupported_reason, + get_paged_kv_storage_unsupported_reason, +) +from .utils import get_kv_page_offset, get_multi_processor_count_for_device + +if TYPE_CHECKING: + from tensorrt_llm._torch.attention.backends.prims_ts import ( + BlockSparsePagedTSWrapper, + BlockSparseTSWrapper, + ) + from tensorrt_llm._torch.attention.backends.trtllm import ( + TrtllmAttention, + TrtllmAttentionMetadata, + ) + +from tensorrt_llm._torch.attention.backends.prims_ts import ( + BlockSparsePagedTSWrapper as _BlockSparsePagedTSWrapper, +) +from tensorrt_llm._torch.attention.backends.prims_ts import ( + BlockSparseTSWrapper as _BlockSparseTSWrapper, +) + + +@dataclass(frozen=True, slots=True) +class _BlockSparsePlanKey: + """Static wrapper profile shared by compatible attention layers. + + The key is the single description of a plan: support checks validate it + against the kernel library and the wrapper cache plans from it. Every field + is an argument of ``plan()``, so two requests share a planned wrapper + exactly when the kernel could serve both with one plan. Per-layer constants + such as the head topology and dtype stay in the key because a bound cache + may be shared by layers with different geometries. + """ + + device: torch.device + batch_size: int + seq_len_q: int + kv_capacity: int + num_heads: int + num_kv_heads: int + head_dim: int + page_size: int | None + q_block_size: int + kv_block_size: int + max_blocks_per_row: int + mask_type: Literal["dense", "causal"] + dtype: torch.dtype + use_kv_valid_bits: bool + sparse_format: Literal["bsr", "bitmask"] + use_proxy_routes: bool + + def unsupported_reason(self) -> str | None: + try: + _validate_block_sparse_static_profile( + batch_size=self.batch_size, + seq_len_q=self.seq_len_q, + seq_len_kv=self.kv_capacity, + num_qo_heads=self.num_heads, + num_kv_heads=self.num_kv_heads, + head_dim=self.head_dim, + q_block_size=self.q_block_size, + kv_block_size=self.kv_block_size, + use_kv_valid_bits=self.use_kv_valid_bits, + mask_type=self.mask_type, + q_dtype=self.dtype, + kv_dtype=self.dtype, + output_dtype=self.dtype, + max_blocks_per_row=self.max_blocks_per_row, + page_size=self.page_size, + ) + except (ValueError, NotImplementedError, OverflowError) as error: + return str(error) + return None + + def plan(self) -> "BlockSparseTSWrapper | BlockSparsePagedTSWrapper": + paged = self.page_size is not None + wrapper_type = _BlockSparsePagedTSWrapper if paged else _BlockSparseTSWrapper + assert wrapper_type is not None + wrapper = wrapper_type() + plan_args = ( + self.batch_size, + self.seq_len_q, + self.kv_capacity, + self.num_heads, + self.num_kv_heads, + self.head_dim, + self.q_block_size, + self.kv_block_size, + ) + plan_kwargs = { + "device": self.device, + "max_blocks_per_row": self.max_blocks_per_row, + "use_kv_valid_bits": self.use_kv_valid_bits, + "mask_type": self.mask_type, + "q_data_type": self.dtype, + "kv_data_type": self.dtype, + "o_data_type": self.dtype, + } + if paged: + plan_args += (self.page_size,) + else: + plan_kwargs.update( + sparse_format=self.sparse_format, + use_proxy_routes=self.use_proxy_routes, + ) + wrapper.plan(*plan_args, **plan_kwargs) + return wrapper + + +def _get_block_sparse_inputs( + forward_args: AttentionForwardArgs, +) -> BlockSparseForwardInputs | None: + return forward_args.sparse_runtime_params.block_sparse_inputs + + +def _has_other_sparse_runtime(forward_args: AttentionForwardArgs) -> bool: + """Whether the runtime carrier holds any sparse state besides block-sparse routes.""" + params = forward_args.sparse_runtime_params + for field in fields(params): + if field.name == "block_sparse_inputs": + continue + value = getattr(params, field.name) + if isinstance(value, torch.Tensor) or (value is not None and value != 0): + return True + return False + + +def _route_batch_size(inputs: BlockSparseForwardInputs) -> int: + routes = inputs.block_indptr if inputs.sparse_format == "bsr" else inputs.exact_block_bits + return int(routes.shape[0]) + + +def _uniform_seq_len_q( + q: torch.Tensor, + metadata: "TrtllmAttentionMetadata", + batch_size: int, +) -> int | None: + """Return the fixed per-request query length, or ``None`` if the batch is ragged.""" + seq_lens = metadata.seq_lens + if batch_size <= 0 or q.shape[0] % batch_size: + return None + if seq_lens is None or seq_lens.numel() < batch_size: + return None + seq_len_q = int(q.shape[0]) // batch_size + if not bool(seq_lens[:batch_size].eq(seq_len_q).all()): + return None + return seq_len_q + + +class PrimsTSBlockSparseFmha(PrimsTSFmha): + """Contiguous context and fixed-Q paged generation block-sparse FMHA.""" + + supports_block_sparse_inputs = True + + def __init__(self, attn: "TrtllmAttention") -> None: + super().__init__(attn) + self._contiguous_wrappers: dict[_BlockSparsePlanKey, "BlockSparseTSWrapper"] = {} + self._paged_wrappers: dict[_BlockSparsePlanKey, "BlockSparsePagedTSWrapper"] = {} + + def bind_plan_cache(self, cache_state: dict[str, object]) -> None: + """Share planned wrappers with every adapter bound to ``cache_state``. + + Attention layers that execute serially, such as the blocks of one + diffusion transformer, see identical static profiles. Binding them to + one model-scoped container plans each profile once and allocates its + route workspace once. Call before the first forward. + """ + + self._contiguous_wrappers = cast( + dict[_BlockSparsePlanKey, "BlockSparseTSWrapper"], + cache_state.setdefault("contiguous_wrappers", {}), + ) + self._paged_wrappers = cast( + dict[_BlockSparsePlanKey, "BlockSparsePagedTSWrapper"], + cache_state.setdefault("paged_wrappers", {}), + ) + + @classmethod + def is_available(cls, attn: "TrtllmAttention") -> bool: + return super().is_available(attn) + + def is_supported( + self, + q: torch.Tensor, + k: torch.Tensor | None, + v: torch.Tensor | None, + metadata: "TrtllmAttentionMetadata", + forward_args: AttentionForwardArgs, + *, + phase: FmhaPhase | None = None, + ) -> bool: + supported, reason = self._is_supported_with_reason( + q, + k, + v, + metadata, + forward_args, + phase=phase, + ) + if not supported: + logger.debug(f"PrimTS block-sparse FMHA does not support request: {reason}") + return supported + + def _is_supported_with_reason( + self, + q: torch.Tensor, + k: torch.Tensor | None, + v: torch.Tensor | None, + metadata: "TrtllmAttentionMetadata", + forward_args: AttentionForwardArgs, + *, + phase: FmhaPhase | None = None, + ) -> tuple[bool, str]: + reason = self._common_unsupported_reason(metadata, forward_args) + if reason is None: + paged = metadata.kv_cache_manager is not None + expected_phase = FmhaPhase.GENERATION if paged else FmhaPhase.CONTEXT + if phase not in (None, expected_phase): + storage = "paged" if paged else "contiguous" + reason = ( + f"{storage} block-sparse attention only supports the " + f"{expected_phase.name.lower()} phase" + ) + elif paged: + reason = self._paged_unsupported_reason(q, metadata, forward_args) + else: + reason = self._contiguous_unsupported_reason(q, k, v, metadata, forward_args) + return reason is None, reason or "" + + def _common_unsupported_reason( + self, + metadata: "TrtllmAttentionMetadata", + forward_args: AttentionForwardArgs, + ) -> str | None: + """Gates shared by the contiguous and paged block-sparse paths.""" + if _get_block_sparse_inputs(forward_args) is None: + return "block-sparse forward inputs are required" + if metadata.is_cross: + return "cross attention is not supported" + if self.attn.is_mla_enable: + return "MLA is not supported" + if metadata.num_sparse_topk > 0 or _has_other_sparse_runtime(forward_args): + return "legacy sparse attention cannot be combined with block-sparse inputs" + feature_reason = get_attention_feature_unsupported_reason(metadata, forward_args) + if feature_reason is not None: + return feature_reason + if forward_args.softmax_stats_tensor is not None: + return "softmax statistics output is not supported" + if ( + forward_args.output_sf is not None + or forward_args.out_scale is not None + or forward_args.out_scale_sf is not None + ): + return "quantized output is not supported" + if forward_args.attention_mask not in ( + PredefinedAttentionMask.FULL, + PredefinedAttentionMask.CAUSAL, + ): + return "only full and causal masks are supported" + return None + + def _make_plan_key( + self, + q: torch.Tensor, + inputs: BlockSparseForwardInputs, + *, + batch_size: int, + seq_len_q: int, + kv_capacity: int, + page_size: int | None, + mask_type: Literal["dense", "causal"], + ) -> _BlockSparsePlanKey: + max_blocks_per_row = inputs.max_blocks_per_row + if max_blocks_per_row is None: + max_blocks_per_row = math.ceil(kv_capacity / inputs.kv_block_size) + return _BlockSparsePlanKey( + device=q.device, + batch_size=batch_size, + seq_len_q=seq_len_q, + kv_capacity=kv_capacity, + num_heads=self.attn.num_heads, + num_kv_heads=self.attn.num_kv_heads, + head_dim=self.attn.head_dim, + page_size=page_size, + q_block_size=inputs.q_block_size, + kv_block_size=inputs.kv_block_size, + max_blocks_per_row=max_blocks_per_row, + mask_type=mask_type, + dtype=q.dtype, + use_kv_valid_bits=inputs.kv_valid_bits is not None, + sparse_format=inputs.sparse_format, + use_proxy_routes=inputs.use_proxy_routes, + ) + + def _get_or_plan_wrapper( + self, + key: _BlockSparsePlanKey, + ) -> "BlockSparseTSWrapper | BlockSparsePagedTSWrapper": + cache = self._paged_wrappers if key.page_size is not None else self._contiguous_wrappers + wrapper = cache.get(key) + if wrapper is None: + wrapper = key.plan() + cache[key] = wrapper + return wrapper + + def _contiguous_unsupported_reason( + self, + q: torch.Tensor, + k: torch.Tensor | None, + v: torch.Tensor | None, + metadata: "TrtllmAttentionMetadata", + forward_args: AttentionForwardArgs, + ) -> str | None: + if forward_args.is_fused_qkv or k is None or v is None: + return "contiguous block-sparse attention requires separate Q, K, and V" + if self.attn.position_embedding_type != 0 or forward_args.mrope_position_deltas is not None: + return "contiguous Q/K/V must have position embedding applied before attention" + if forward_args.cu_q_seqlens is not None or forward_args.cu_kv_seqlens is not None: + return "packed variable-length Q/KV inputs are not supported" + inputs = _get_block_sparse_inputs(forward_args) + assert inputs is not None + mask_type = self._get_prims_mask_type(forward_args) + if inputs.use_proxy_routes and mask_type != "dense": + return "block-sparse proxy routes require mask_type='dense'" + batch_size = _route_batch_size(inputs) + seq_len_q = _uniform_seq_len_q(q, metadata, batch_size) + if seq_len_q is None or k.shape[0] % batch_size: + return "query and KV token counts must be batch-uniform over the route batch size" + key = self._make_plan_key( + q, + inputs, + batch_size=batch_size, + seq_len_q=seq_len_q, + kv_capacity=int(k.shape[0]) // batch_size, + page_size=None, + mask_type=mask_type, + ) + return key.unsupported_reason() + + def _paged_unsupported_reason( + self, + q: torch.Tensor, + metadata: "TrtllmAttentionMetadata", + forward_args: AttentionForwardArgs, + ) -> str | None: + inputs = _get_block_sparse_inputs(forward_args) + assert inputs is not None + if inputs.sparse_format != "bsr" or inputs.use_proxy_routes: + return "paged block-sparse attention only supports BSR exact routes" + if not forward_args.is_fused_qkv: + return "paged block-sparse attention requires fused QKV input" + if ( + forward_args.attention_input_type == AttentionInputType.context_only + or metadata.num_contexts != 0 + ): + return "paged block-sparse attention requires a generation-only batch" + reason = get_paged_kv_storage_unsupported_reason( + self.attn, metadata + ) or get_paged_kv_policy_unsupported_reason(self.attn, metadata) + if reason is not None: + return reason + if metadata.tokens_per_block not in self.SUPPORTED_PAGE_SIZES: + return f"page size {metadata.tokens_per_block} is unsupported" + if self.attn.attention_chunk_size: + return "chunked attention is not supported" + if get_kv_page_offset(self.attn, metadata, 0, cache=self._kv_page_offset_cache) is None: + return "the K-to-V page displacement could not be resolved" + + batch_size = int(metadata.num_generations) + seq_len_q = _uniform_seq_len_q(q, metadata, batch_size) + if seq_len_q is None: + return "query lengths must be batch-uniform and match the fixed query shape" + block_tables = metadata.kv_cache_block_offsets + if block_tables.shape[1] < batch_size: + return "paged KV-cache block offsets must cover the generation batch" + page_size = int(metadata.tokens_per_block) + kv_capacity = int(block_tables.shape[-1]) * page_size + logical_max_seq_len = int(metadata.max_seq_len) + if logical_max_seq_len > kv_capacity: + return "logical maximum sequence length must fit the page-table capacity" + attention_window_size = forward_args.attention_window_size + if ( + attention_window_size is None + or attention_window_size < logical_max_seq_len + or attention_window_size > kv_capacity + ): + return "attention window must fit the non-cyclic page-table capacity" + host_seq_lens = metadata.kv_lens_runtime[:batch_size] + min_seq_len_kv = int(host_seq_lens.min()) + if min_seq_len_kv <= 0: + return "every active request must contain at least one KV token" + mask_type = self._get_prims_mask_type(forward_args) + if mask_type == "causal" and min_seq_len_kv < seq_len_q: + return "causal KV lengths must be at least the fixed query length" + if int(host_seq_lens.max()) > logical_max_seq_len: + return "an active KV length exceeds the logical maximum sequence length" + key = self._make_plan_key( + q, + inputs, + batch_size=batch_size, + seq_len_q=seq_len_q, + kv_capacity=kv_capacity, + page_size=page_size, + mask_type=mask_type, + ) + return key.unsupported_reason() + + def prepare_workspace( + self, + q: torch.Tensor, + k: torch.Tensor | None, + v: torch.Tensor | None, + metadata: "TrtllmAttentionMetadata", + forward_args: AttentionForwardArgs, + workspace: torch.Tensor, + ) -> None: + del k, v, forward_args + with torch.cuda.device(q.device): + # Contiguous requests run without a KV cache and never touch the + # generation preprocessing workspace. + if metadata.kv_cache_manager is not None: + layout = self._get_generation_workspace_layout( + q.dtype, + int(metadata.num_generations), + int(q.shape[0]), + ) + required_bytes = int(layout["total_size"]) + if workspace.numel() * workspace.element_size() < required_bytes: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "TRT-LLM QKV preprocessing workspace must be sized before " + "CUDA Graph capture" + ) + workspace.resize_((math.ceil(required_bytes / workspace.element_size()),)) + if self._multi_processor_count is None: + self._multi_processor_count = get_multi_processor_count_for_device(q.device.index) + + def run_generation(self, params: FmhaParams) -> None: + q = params.qkv_input + output_buffer = params.context_buf + sequence_lengths = params.sequence_lengths + assert q is not None and output_buffer is not None and sequence_lengths is not None + metadata = params.meta + forward_args = params.fwd + inputs = _get_block_sparse_inputs(forward_args) + assert inputs is not None + batch_size = params.num_requests + seq_len_q = params.input_seq_length + page_size = params.tokens_per_block + block_offsets = metadata.kv_cache_block_offsets + assert block_offsets is not None + preprocess = self._run_generation_preprocess(params) + q_processed, kv_pool, block_tables = preprocess[:3] + fmha_workspace = preprocess[6] + if fmha_workspace.numel() != 0: + raise RuntimeError("PrimTS block-sparse preprocessing returned an FMHA workspace.") + if q_processed is None or kv_pool is None or block_tables is None: + raise RuntimeError("TRT-LLM preprocessing did not return paged PrimTS metadata") + kv_page_offset = get_kv_page_offset( + params.attn, + metadata, + params.seq_offset, + cache=self._kv_page_offset_cache, + ) + if kv_page_offset is None: + raise RuntimeError("PrimTS could not resolve the K-to-V page displacement") + k_cache, v_cache = self._standard_kv_views(kv_pool, kv_page_offset) + query = q_processed.view( + batch_size, + seq_len_q, + self.attn.num_heads, + self.attn.head_dim, + ) + key = self._make_plan_key( + query, + inputs, + batch_size=batch_size, + seq_len_q=seq_len_q, + kv_capacity=int(block_offsets.shape[-1]) * page_size, + page_size=page_size, + mask_type=self._get_prims_mask_type(forward_args), + ) + wrapper = cast("BlockSparsePagedTSWrapper", self._get_or_plan_wrapper(key)) + wrapper.run( + query, + (k_cache, v_cache), + block_tables=self._get_fixed_block_tables(block_tables, batch_size), + seq_lens_kv=self._get_sequence_lengths(sequence_lengths, batch_size), + block_indptr=inputs.block_indptr, + block_indices=inputs.block_indices, + kv_valid_bits=inputs.kv_valid_bits, + sm_scale=self._get_bmm1_scale(self.attn), + out=output_buffer.view_as(query), + ) + + def _forward_contiguous( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + forward_args: AttentionForwardArgs, + ) -> None: + inputs = _get_block_sparse_inputs(forward_args) + assert inputs is not None + assert forward_args.output is not None + batch_size = _route_batch_size(inputs) + query = q.view(batch_size, -1, self.attn.num_heads, self.attn.head_dim) + key_states = k.view(batch_size, -1, self.attn.num_kv_heads, self.attn.head_dim) + value_states = v.view_as(key_states) + key = self._make_plan_key( + query, + inputs, + batch_size=batch_size, + seq_len_q=int(query.shape[1]), + kv_capacity=int(key_states.shape[1]), + page_size=None, + mask_type=self._get_prims_mask_type(forward_args), + ) + wrapper = cast("BlockSparseTSWrapper", self._get_or_plan_wrapper(key)) + wrapper.run( + query, + key_states, + value_states, + block_indptr=inputs.block_indptr, + block_indices=inputs.block_indices, + exact_block_bits=inputs.exact_block_bits, + k_summary=inputs.k_summary, + v_summary=inputs.v_summary, + kv_valid_bits=inputs.kv_valid_bits, + sm_scale=self._get_bmm1_scale(self.attn), + out=forward_args.output.view_as(query), + ) + + def forward( + self, + q: torch.Tensor, + k: torch.Tensor | None, + v: torch.Tensor | None, + metadata: "TrtllmAttentionMetadata", + forward_args: AttentionForwardArgs, + ) -> None: + if metadata.kv_cache_manager is None: + assert k is not None and v is not None + self._forward_contiguous(q, k, v, forward_args) + return + super().forward(q, k, v, metadata, forward_args) diff --git a/tensorrt_llm/_torch/attention/backends/fmha/registry.py b/tensorrt_llm/_torch/attention/backends/fmha/registry.py index 4e545924e63e..240e5242feb0 100644 --- a/tensorrt_llm/_torch/attention/backends/fmha/registry.py +++ b/tensorrt_llm/_torch/attention/backends/fmha/registry.py @@ -35,6 +35,7 @@ def init_fmha_libs() -> dict[str, "FmhaCls"]: """ from .flashinfer_sparse_mla import FlashInferSparseMlaFmha from .msa_sparse_gqa import MsaSparseGqaFmha + from .prims_ts_block_sparse import PrimsTSBlockSparseFmha return { "triton_custom_mask": TritonCustomMaskFmha, @@ -42,6 +43,7 @@ def init_fmha_libs() -> dict[str, "FmhaCls"]: "msa_sparse_gqa": MsaSparseGqaFmha, "flashinfer_sparse_mla": FlashInferSparseMlaFmha, "prims_ts": PrimsTSFmha, + "prims_ts_block_sparse": PrimsTSBlockSparseFmha, "flashinfer_trtllm_gen": FlashInferTrtllmGenFmha, "fallback": FallbackFmha, } diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/common.py b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/common.py index a82deafc7a3f..13cf6038efd9 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/common.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/common.py @@ -22,6 +22,20 @@ _BLOCK_SPARSE_MAX_HEADS_Q_PER_KV = 32 +def _validate_contiguous_route_mode( + sparse_format: object, + use_proxy_routes: object, +) -> None: + """Validate the two public continuous-route axes before device work.""" + + if not isinstance(sparse_format, str): + raise TypeError("sparse_format must be 'bsr' or 'bitmask'") + if sparse_format not in ("bsr", "bitmask"): + raise ValueError("sparse_format must be 'bsr' or 'bitmask'") + if type(use_proxy_routes) is not bool: + raise TypeError("use_proxy_routes must be a bool") + + def _validate_sparse_q_block_size(value: object) -> int: """Return a positive semantic Q block size representable by the ABI.""" @@ -110,6 +124,17 @@ def _block_sparse_kv_atom_size(kv_block_size: int) -> int: ) +def _block_sparse_proxy_summary_geometry( + seq_len_kv: int, + kv_block_size: int, +) -> tuple[int, int]: + """Return the summary count and final summary's represented-token mass.""" + + num_summaries = (seq_len_kv + kv_block_size - 1) // kv_block_size + tail_mass = seq_len_kv - (num_summaries - 1) * kv_block_size + return num_summaries, tail_mass + + def _prepared_kv_routes_are_block_aligned( kv_block_size: int, kv_route_size: int, @@ -117,3 +142,24 @@ def _prepared_kv_routes_are_block_aligned( """Return whether each prepared route stays within one semantic BSR block.""" return _validate_sparse_kv_block_size(kv_block_size) % kv_route_size == 0 + + +def _block_sparse_contiguous_kv_copy_geometry( + *, + kv_block_size: int, + kv_route_size: int, +) -> tuple[int, int, bool]: + """Return source-independent primary/atom TensorMap geometry. + + Exact and proxy routes address different logical matrices, but a route's + physical copy shape depends only on its semantic block and physical route + sizes. Coarse routes prefer KV128 copies and keep a KV64 descriptor only + when KV256 staging or runtime adjacency requires it. + """ + + atom_size = _block_sparse_kv_atom_size(kv_block_size) + primary_box_size = 2 * atom_size if atom_size == 64 else atom_size + needs_aux_atom = atom_size == 64 and ( + kv_route_size == 256 or kv_block_size % kv_route_size != 0 + ) + return primary_box_size, atom_size, needs_aux_atom diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/compiler.py b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/compiler.py index a73f26060753..fcaa0735436f 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/compiler.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/compiler.py @@ -36,31 +36,46 @@ def _compile_block_sparse(key: _BlockSparseCompileKey) -> Callable[..., object]: from ..kernels.fmha_decode.fmha_decode_config import FmhaDecodeConfig from ..kernels.fmha_decode.block_sparse_prepare import ( - _PrepareBlockSparseRoutes, + _PrepareBitmaskRoutes, + _PrepareBsrRoutes, ) from ..kernels.fmha_decode.fmha_decode_kernel import ( fmha_block_sparse_launch, ) config = _make_block_sparse_config(key) - prepare_routes = _PrepareBlockSparseRoutes( - batch_size=key.batch_size, - num_kv_heads=key.num_kv_heads, - seq_len_q=key.seq_len_q, - seq_len_kv=key.seq_len_kv, - q_block_size=key.q_block_size, - kv_block_size=key.kv_block_size, - kv_route_size=key.kv_route_size, - has_token_bits=key.use_kv_valid_bits, - page_size=key.page_size, - mask_type=key.mask_type, - ) + prepare_kwargs = { + "batch_size": key.batch_size, + "num_kv_heads": key.num_kv_heads, + "seq_len_q": key.seq_len_q, + "seq_len_kv": key.seq_len_kv, + "q_block_size": key.q_block_size, + "kv_block_size": key.kv_block_size, + "kv_route_size": key.kv_route_size, + "use_proxy_routes": key.use_proxy_routes, + "use_causal_mask": key.mask_type == "causal", + "apply_token_mask": key.use_kv_valid_bits, + "store_score_words": config.uses_prepared_score_keep_words, + } + if key.page_size is not None: + if key.sparse_format != "bsr" or key.use_proxy_routes: + raise AssertionError("paged block-sparse supports exact BSR routes only") + prepare_kwargs["page_size"] = key.page_size + + if key.sparse_format == "bsr": + prepare_routes = _PrepareBsrRoutes(**prepare_kwargs) + elif key.sparse_format == "bitmask": + prepare_routes = _PrepareBitmaskRoutes(**prepare_kwargs) + else: + raise AssertionError("sparse_format must be 'bsr' or 'bitmask'") + + route_metadata_base = prepare_routes.route_metadata_base_word_offset Int32 = cutlass.Int32 Int64 = cutlass.Int64 Float32 = cutlass.Float32 @cute.jit - def contiguous_tensor_adapter( + def exact_bsr_adapter( q: cute.Tensor, k: cute.Tensor, v: cute.Tensor, @@ -86,7 +101,7 @@ def contiguous_tensor_adapter( kv_valid_bits, None, None, - None, + Int64(0), Int64(0), row_route_offsets, route_workspace, @@ -95,9 +110,7 @@ def contiguous_tensor_adapter( ) # Live per-row route counts occupy the first words of run scratch. row_route_counts = route_workspace.iterator - route_metadata = route_workspace.iterator + Int32( - prepare_routes.route_metadata_base_word_offset - ) + route_metadata = route_workspace.iterator + Int32(route_metadata_base) fmha_block_sparse_launch( ( Int32(static_batch_size), @@ -109,6 +122,8 @@ def contiguous_tensor_adapter( q.iterator, k.iterator, v.iterator, + k.iterator, + v.iterator, out.iterator, row_route_offsets.iterator, row_route_counts, @@ -119,6 +134,169 @@ def contiguous_tensor_adapter( static_seq_len_kv, ) + @cute.jit + def exact_bitmask_adapter( + q: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + out: cute.Tensor, + exact_block_bits: cute.Tensor, + kv_valid_bits: cute.Tensor, + row_route_offsets: cute.Tensor, + route_workspace: cute.Tensor, + max_blocks_per_row: cutlass.Int32, + sm_scale: cutlass.Float32, + stream: cuda_drv.CUstream, + static_config: cutlass.Constexpr[FmhaDecodeConfig], + static_batch_size: cutlass.Constexpr[int], + static_seq_len_kv: cutlass.Constexpr[int], + static_num_qo_heads: cutlass.Constexpr[int], + static_num_kv_heads: cutlass.Constexpr[int], + static_head_dim: cutlass.Constexpr[int], + ) -> None: + prepare_routes( + exact_block_bits, + kv_valid_bits, + row_route_offsets, + route_workspace, + max_blocks_per_row, + stream, + ) + fmha_block_sparse_launch( + ( + Int32(static_batch_size), + Int32(static_num_qo_heads), + Int32(static_num_kv_heads), + Int32(static_seq_len_kv), + Int32(static_head_dim), + ), + q.iterator, + k.iterator, + v.iterator, + k.iterator, + v.iterator, + out.iterator, + row_route_offsets.iterator, + route_workspace.iterator, + route_workspace.iterator + Int32(route_metadata_base), + sm_scale, + stream, + static_config, + static_seq_len_kv, + ) + + @cute.jit + def proxy_bsr_adapter( + q: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + k_summary: cute.Tensor, + v_summary: cute.Tensor, + out: cute.Tensor, + block_indptr: cute.Tensor, + block_indices: cute.Tensor, + kv_valid_bits: cute.Tensor, + row_route_offsets: cute.Tensor, + route_workspace: cute.Tensor, + max_blocks_per_row: cutlass.Int32, + sm_scale: cutlass.Float32, + stream: cuda_drv.CUstream, + static_config: cutlass.Constexpr[FmhaDecodeConfig], + static_batch_size: cutlass.Constexpr[int], + static_seq_len_kv: cutlass.Constexpr[int], + static_num_qo_heads: cutlass.Constexpr[int], + static_num_kv_heads: cutlass.Constexpr[int], + static_head_dim: cutlass.Constexpr[int], + ) -> None: + prepare_routes( + block_indptr, + block_indices, + kv_valid_bits, + None, + None, + Int64(0), + Int64(0), + row_route_offsets, + route_workspace, + max_blocks_per_row, + stream, + ) + fmha_block_sparse_launch( + ( + Int32(static_batch_size), + Int32(static_num_qo_heads), + Int32(static_num_kv_heads), + Int32(static_seq_len_kv), + Int32(static_head_dim), + ), + q.iterator, + k.iterator, + v.iterator, + k_summary.iterator, + v_summary.iterator, + out.iterator, + row_route_offsets.iterator, + route_workspace.iterator, + route_workspace.iterator + Int32(route_metadata_base), + sm_scale, + stream, + static_config, + static_seq_len_kv, + ) + + @cute.jit + def proxy_bitmask_adapter( + q: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + k_summary: cute.Tensor, + v_summary: cute.Tensor, + out: cute.Tensor, + exact_block_bits: cute.Tensor, + kv_valid_bits: cute.Tensor, + row_route_offsets: cute.Tensor, + route_workspace: cute.Tensor, + max_blocks_per_row: cutlass.Int32, + sm_scale: cutlass.Float32, + stream: cuda_drv.CUstream, + static_config: cutlass.Constexpr[FmhaDecodeConfig], + static_batch_size: cutlass.Constexpr[int], + static_seq_len_kv: cutlass.Constexpr[int], + static_num_qo_heads: cutlass.Constexpr[int], + static_num_kv_heads: cutlass.Constexpr[int], + static_head_dim: cutlass.Constexpr[int], + ) -> None: + prepare_routes( + exact_block_bits, + kv_valid_bits, + row_route_offsets, + route_workspace, + max_blocks_per_row, + stream, + ) + fmha_block_sparse_launch( + ( + Int32(static_batch_size), + Int32(static_num_qo_heads), + Int32(static_num_kv_heads), + Int32(static_seq_len_kv), + Int32(static_head_dim), + ), + q.iterator, + k.iterator, + v.iterator, + k_summary.iterator, + v_summary.iterator, + out.iterator, + row_route_offsets.iterator, + route_workspace.iterator, + route_workspace.iterator + Int32(route_metadata_base), + sm_scale, + stream, + static_config, + static_seq_len_kv, + ) + @cute.jit def paged_tensor_adapter( q: cute.Tensor, @@ -128,13 +306,13 @@ def paged_tensor_adapter( block_indptr: cute.Tensor, block_indices: cute.Tensor, kv_valid_bits: cute.Tensor, - paged_kv_indptr: cute.Tensor, - paged_kv_indices: cute.Tensor, + block_tables: cute.Tensor, seq_lens_kv: cute.Tensor, row_route_offsets: cute.Tensor, route_workspace: cute.Tensor, max_blocks_per_row: cutlass.Int32, num_physical_kv_pages: cutlass.Int64, + block_table_row_stride: cutlass.Int64, k_page_stride: cutlass.Int64, v_page_stride: cutlass.Int64, sm_scale: cutlass.Float32, @@ -151,18 +329,16 @@ def paged_tensor_adapter( block_indices, kv_valid_bits, seq_lens_kv, - paged_kv_indptr, - paged_kv_indices, + block_tables, num_physical_kv_pages, + block_table_row_stride, row_route_offsets, route_workspace, max_blocks_per_row, stream, ) row_route_counts = route_workspace.iterator - route_metadata = route_workspace.iterator + Int32( - prepare_routes.route_metadata_base_word_offset - ) + route_metadata = route_workspace.iterator + Int32(route_metadata_base) fmha_block_sparse_launch( ( Int32(static_batch_size), @@ -174,6 +350,8 @@ def paged_tensor_adapter( q.iterator, k_cache.iterator, v_cache.iterator, + k_cache.iterator, + v_cache.iterator, out.iterator, row_route_offsets.iterator, row_route_counts, @@ -203,6 +381,7 @@ def fake_compact( logical_workspace_words = cute.sym_int() q_shape = (key.batch_size, key.seq_len_q, key.num_qo_heads, key.head_dim) num_q_blocks = ceil_div(key.seq_len_q, key.q_block_size) + num_kv_blocks = ceil_div(key.seq_len_kv, key.kv_block_size) indptr_fake = fake_compact( Int32, (key.batch_size, key.num_kv_heads, num_q_blocks + 1), @@ -238,24 +417,96 @@ def fake_compact( ) k_fake = fake_compact(config.kv_dtype, kv_shape, 16) v_fake = fake_compact(config.kv_dtype, kv_shape, 16) - tensor_adapter = contiguous_tensor_adapter - dynamic_args = ( - q_fake, - k_fake, - v_fake, - out_fake, - indptr_fake, - indices_fake, + exact_bits_fake = fake_compact( + cutlass.Uint32, + ( + key.batch_size, + key.num_kv_heads, + num_q_blocks, + ceil_div(num_kv_blocks, 32), + ), + 4, + ) + common_tail = ( valid_bits_fake, row_route_offsets_fake, route_workspace_fake, Int32(0), Float32(1.0), ) + if key.sparse_format == "bsr" and not key.use_proxy_routes: + tensor_adapter = exact_bsr_adapter + dynamic_args = ( + q_fake, + k_fake, + v_fake, + out_fake, + indptr_fake, + indices_fake, + *common_tail, + ) + elif key.sparse_format == "bitmask" and not key.use_proxy_routes: + tensor_adapter = exact_bitmask_adapter + dynamic_args = ( + q_fake, + k_fake, + v_fake, + out_fake, + exact_bits_fake, + *common_tail, + ) + elif key.sparse_format == "bsr" and key.use_proxy_routes: + summary_shape = ( + key.batch_size, + num_kv_blocks, + key.num_kv_heads, + key.head_dim, + ) + k_summary_fake = fake_compact(config.kv_dtype, summary_shape, 16) + v_summary_fake = fake_compact(config.kv_dtype, summary_shape, 16) + proxy_prefix = ( + q_fake, + k_fake, + v_fake, + k_summary_fake, + v_summary_fake, + out_fake, + ) + tensor_adapter = proxy_bsr_adapter + dynamic_args = ( + *proxy_prefix, + indptr_fake, + indices_fake, + *common_tail, + ) + elif key.sparse_format == "bitmask" and key.use_proxy_routes: + summary_shape = ( + key.batch_size, + num_kv_blocks, + key.num_kv_heads, + key.head_dim, + ) + k_summary_fake = fake_compact(config.kv_dtype, summary_shape, 16) + v_summary_fake = fake_compact(config.kv_dtype, summary_shape, 16) + tensor_adapter = proxy_bitmask_adapter + dynamic_args = ( + q_fake, + k_fake, + v_fake, + k_summary_fake, + v_summary_fake, + out_fake, + exact_bits_fake, + *common_tail, + ) + else: + raise AssertionError("continuous sparse_format must be 'bsr' or 'bitmask'") else: page_size = key.page_size + assert page_size is not None physical_pages = cute.sym_int() - logical_pages = cute.sym_int() + runtime_page_columns = cute.sym_int() + runtime_page_row_stride = cute.sym_int64(divisibility=1) k_outer_stride = cute.sym_int64(divisibility=1) v_outer_stride = cute.sym_int64(divisibility=1) kv_shape = ( @@ -286,12 +537,12 @@ def fake_compact( ), assumed_align=16, ) - paged_kv_indptr_fake = fake_compact( + block_tables_fake = cute.runtime.make_fake_tensor( Int32, - (key.batch_size + 1,), - 4, + (key.batch_size, runtime_page_columns), + stride=(runtime_page_row_stride, 1), + assumed_align=4, ) - paged_kv_indices_fake = fake_compact(Int32, (logical_pages,), 4) seq_lens_kv_fake = fake_compact(Int32, (key.batch_size,), 4) tensor_adapter = paged_tensor_adapter dynamic_args = ( @@ -302,8 +553,7 @@ def fake_compact( indptr_fake, indices_fake, valid_bits_fake, - paged_kv_indptr_fake, - paged_kv_indices_fake, + block_tables_fake, seq_lens_kv_fake, row_route_offsets_fake, route_workspace_fake, @@ -311,6 +561,7 @@ def fake_compact( Int64(1), Int64(1), Int64(1), + Int64(1), Float32(1.0), ) diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/config.py b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/config.py index 47d7c60340b5..cc4458a4298a 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/config.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/config.py @@ -64,7 +64,7 @@ @dataclass(frozen=True) class _BlockSparseCompileKey: - """Named, hashable inputs that determine one compiled adapter.""" + """Named, hashable inputs that determine one compiled sparse adapter.""" device_index: int batch_size: int @@ -81,6 +81,8 @@ class _BlockSparseCompileKey: use_kv_valid_bits: bool use_persistent_scheduler: bool use_parallel_sparse_kv_loads: bool + sparse_format: Literal["bsr", "bitmask"] = "bsr" + use_proxy_routes: bool = False page_size: int | None = None @@ -90,6 +92,9 @@ class _BlockSparseLaunchSpec: policy: tuple[tuple[str, object], ...] compile_key: _BlockSparseCompileKey + # Whether prepared routes carry K32 score-validity words; the decode + # config owns this rule and the plan sizes its route storage from it. + prepares_score_words: bool _CAPACITY_UNSET = object() @@ -194,12 +199,12 @@ def _select_block_sparse_scheduler( use_kv_valid_bits: bool, max_row_route_capacity: int, ) -> tuple[int, bool]: - """Select the Q tile and scheduler without depending on KV storage.""" + """Select the Q tile and scheduler without depending on KV storage. - from ..kernels.fmha_decode.fmha_decode_config import ( - _select_auto_launch_mode, - make_q_tile_geometry, - ) + Proxy routes add one summary route per row on top of the exact routes and + see the same per-tile fixed cost the persistent scheduler amortizes, so + the selection does not depend on the route kind. + """ heads_q_per_kv = num_qo_heads // num_kv_heads q_tile_size = _select_block_sparse_q_tile_size( @@ -216,6 +221,11 @@ def _select_block_sparse_scheduler( ): return q_tile_size, False + from ..kernels.fmha_decode.fmha_decode_config import ( + _select_auto_launch_mode, + make_q_tile_geometry, + ) + q_geometry = make_q_tile_geometry( rows_per_cta=q_tile_size, heads_q_per_kv=heads_q_per_kv, @@ -352,11 +362,13 @@ def _validate_block_sparse_static_profile( kv_block_size=kv_block_size, ) if page_size is not None: + # Validate the paged route geometry with a capacity-free layout; the + # score-word slots do not take part in the page/atom checks. _BlockSparseRouteLayout.create( kv_route_size=kv_route_size, kv_block_size=kv_block_size, page_size=page_size, - has_token_bits=use_kv_valid_bits, + has_token_bits=False, route_metadata_capacity=0, num_rows=1, ) @@ -413,6 +425,8 @@ def _make_block_sparse_config(key: _BlockSparseCompileKey) -> "FmhaDecodeConfig" } if key.use_persistent_scheduler: config_args["use_persistent_scheduler"] = True + if key.use_proxy_routes: + config_args["use_block_sparse_proxy_routes"] = True layout_args: dict[str, object] if key.page_size is None: layout_args = {"qkv_layout": "contiguousKv"} @@ -455,14 +469,17 @@ def _resolve_block_sparse_launch_spec( mask_type: Literal["dense", "causal"], use_kv_valid_bits: bool, max_row_route_capacity: int, + sparse_format: Literal["bsr", "bitmask"] = "bsr", + use_proxy_routes: bool = False, page_size: int | None = None, ) -> _BlockSparseLaunchSpec: """Resolve and cache one validated static or CLC launch. ``max_row_route_capacity`` is a conservative prepared-route bound. Live index values and physical-tail morphology never specialize this cache - entry. If the selected persistent profile is unsupported, retain the valid - static profile instead. + entry. Proxy and exact routes share one scheduler selection. An + unsupported persistent profile falls back to its valid static + counterpart. """ q_tile_size, use_persistent_scheduler = _select_block_sparse_scheduler( @@ -499,10 +516,12 @@ def _resolve_block_sparse_launch_spec( max_row_route_capacity=max_row_route_capacity, use_persistent_scheduler=use_persistent_scheduler, ), + sparse_format=sparse_format, + use_proxy_routes=use_proxy_routes, page_size=page_size, ) try: - _make_block_sparse_config(compile_key) + config = _make_block_sparse_config(compile_key) except ValueError: if not compile_key.use_persistent_scheduler: raise @@ -516,11 +535,15 @@ def _resolve_block_sparse_launch_spec( use_persistent_scheduler=False, ), ) - _make_block_sparse_config(compile_key) + config = _make_block_sparse_config(compile_key) policy_entries: list[tuple[str, object]] = [ ("tile_size_q", q_tile_size), ("tile_size_kv", kv_route_size), + ( + "scheduler", + "persistent" if compile_key.use_persistent_scheduler else "static", + ), ] if page_size is not None: policy_entries.append(("page_size", page_size)) @@ -538,6 +561,7 @@ def _resolve_block_sparse_launch_spec( return _BlockSparseLaunchSpec( policy=tuple(policy_entries), compile_key=compile_key, + prepares_score_words=config.uses_prepared_score_keep_words, ) diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/inspection.py b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/inspection.py index 39c0c22d6f64..7abf22e52792 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/inspection.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/inspection.py @@ -64,12 +64,7 @@ def _raise_for_invalid_paged_metadata( return reason = { 4: (f"seq_lens_kv values must lie in [{minimum_seq_len_kv}, {max_seq_len_kv}]"), - 5: ( - "paged_kv_indptr must start at zero and each row must be " - "bounded and monotone" - ), - 6: "paged_kv_indptr rows must contain enough pages for seq_lens_kv", - 7: "paged_kv_indices must contain an in-range physical page ID", + 5: "block_tables must contain an in-range physical page ID for every live page", }.get(error_code) if reason is None: reason = ( @@ -152,8 +147,7 @@ def launch(summary: torch.Tensor, device_index: int) -> None: def _inspect_paged_block_sparse_metadata( block_indptr: torch.Tensor, block_indices: torch.Tensor, - paged_kv_indptr: torch.Tensor, - paged_kv_indices: torch.Tensor, + block_tables: torch.Tensor, seq_lens_kv: torch.Tensor, *, static: _BlockSparseStaticProfile, @@ -186,8 +180,8 @@ def launch(summary: torch.Tensor, device_index: int) -> None: inspect_metadata( block_indptr, block_indices, - paged_kv_indptr, - paged_kv_indices, + block_tables, + block_tables.stride(0), seq_lens_kv, num_physical_kv_pages, summary, diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/plan.py b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/plan.py index e70ed415fc97..453e8e35134a 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/plan.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/plan.py @@ -18,13 +18,13 @@ from dataclasses import dataclass import functools import _thread -from typing import Concatenate, ParamSpec, Protocol, TypeVar, cast +from typing import Concatenate, Literal, ParamSpec, Protocol, TypeVar, cast import torch from flashinfer.utils import ceil_div -from .common import _SIGNED_INT32_MAX +from .common import _SIGNED_INT32_MAX, _block_sparse_proxy_summary_geometry from .compiler import _get_compiled_block_sparse from .config import ( _BlockSparseStaticProfile, @@ -85,6 +85,8 @@ class _BlockSparsePlanState: cannot prevent in-place modification or replace graph ownership. """ + sparse_format: Literal["bsr", "bitmask"] + use_proxy_routes: bool device: torch.device batch_size: int seq_len_q: int @@ -93,6 +95,7 @@ class _BlockSparsePlanState: num_kv_heads: int head_dim: int q_block_size: int + kv_block_size: int q_dtype: torch.dtype kv_dtype: torch.dtype output_dtype: torch.dtype @@ -179,27 +182,31 @@ def _build_block_sparse_plan_state( device: torch.device, device_index: int, plan_stream: torch.cuda.Stream, + sparse_format: Literal["bsr", "bitmask"] = "bsr", + use_proxy_routes: bool = False, ) -> _BlockSparsePlanState: - """Build and close one complete state after storage validation.""" + """Build one format- and route-specialized plan atomically.""" assert static.max_blocks_per_row is not None - max_row_route_capacity = ceil_div( - static.max_blocks_per_row * static.kv_block_size, - static.kv_route_size, - ) + if static.page_size is not None: + assert sparse_format == "bsr" and not use_proxy_routes num_rows = ( static.batch_size * static.num_kv_heads * ceil_div(static.seq_len_q, static.q_block_size) ) - route_layout = _BlockSparseRouteLayout.create( - kv_route_size=static.kv_route_size, - kv_block_size=static.kv_block_size, - page_size=static.page_size, - has_token_bits=static.use_kv_valid_bits, - route_metadata_capacity=num_rows * max_row_route_capacity, - num_rows=num_rows, + if num_rows > _SIGNED_INT32_MAX: + raise OverflowError("row_count must fit in signed int32") + max_row_route_capacity = ceil_div( + static.max_blocks_per_row * static.kv_block_size, + static.kv_route_size, ) + if use_proxy_routes: + num_summaries, _ = _block_sparse_proxy_summary_geometry( + static.seq_len_kv, + static.kv_block_size, + ) + max_row_route_capacity += ceil_div(num_summaries, static.kv_route_size) with torch.cuda.device(device_index), torch.cuda.stream(plan_stream): spec = _resolve_block_sparse_launch_spec( device_index=device_index, @@ -217,11 +224,21 @@ def _build_block_sparse_plan_state( mask_type=static.mask_type, use_kv_valid_bits=static.use_kv_valid_bits, max_row_route_capacity=max_row_route_capacity, + sparse_format=sparse_format, + use_proxy_routes=use_proxy_routes, ) policy = ( *spec.policy, ("max_blocks_per_row", static.max_blocks_per_row), ) + route_layout = _BlockSparseRouteLayout.create( + kv_route_size=static.kv_route_size, + kv_block_size=static.kv_block_size, + page_size=static.page_size, + has_token_bits=spec.prepares_score_words, + route_metadata_capacity=num_rows * max_row_route_capacity, + num_rows=num_rows, + ) compiled = _get_compiled_block_sparse(spec.compile_key) dummy_kv_valid_bits = ( None @@ -240,6 +257,8 @@ def _build_block_sparse_plan_state( ready_event = _record_block_sparse_plan_ready_event(plan_stream) return _BlockSparsePlanState( + sparse_format=sparse_format, + use_proxy_routes=use_proxy_routes, device=device, batch_size=static.batch_size, seq_len_q=static.seq_len_q, @@ -248,6 +267,7 @@ def _build_block_sparse_plan_state( num_kv_heads=static.num_kv_heads, head_dim=static.head_dim, q_block_size=static.q_block_size, + kv_block_size=static.kv_block_size, q_dtype=static.q_dtype, kv_dtype=static.kv_dtype, output_dtype=static.output_dtype, diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/prepared.py b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/prepared.py index 091becbd33b8..66815f890a55 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/prepared.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/prepared.py @@ -22,6 +22,7 @@ _SECTION_ALIGNMENT_WORDS = 4 _PREPARED_ROUTE_IS_FULL_FLAG = 1 << 0 +_PREPARED_ROUTE_IS_PROXY_FLAG = 1 << 1 _SUPPORTED_KV_ROUTE_SIZES = (128, 256) _SUPPORTED_PAGED_KV_PAGE_SIZES = (16, 32, 64, 128) @@ -91,13 +92,16 @@ class _BlockSparseRouteLayout: Each route's metadata stores logical KV-token atom origins, optional physical page IDs, one atom-valid-mask word, one route-flags word, and - optional token-valid words. ``page_size is None`` selects the contiguous + optional token-mask words. ``page_size is None`` selects the contiguous record; otherwise the paged record adds one page-ID word per logical origin. Logical origins remain independent of the K/V storage locator used - by the attention load path. An invalid logical origin is encoded as - ``-1``. Bit ``i`` of the atom-valid mask corresponds to logical origin - ``i``. ``_PREPARED_ROUTE_IS_FULL_FLAG`` (bit 0) states that the route is - both structurally full and, when token bits are present, token-full. + by the attention load path. Exact routes address raw-token origins, while + proxy routes address summary-token origins and set + ``_PREPARED_ROUTE_IS_PROXY_FLAG`` (bit 1). An invalid logical origin is + encoded as ``-1``. Bit ``i`` of the atom-valid mask corresponds to logical + origin ``i``. ``_PREPARED_ROUTE_IS_FULL_FLAG`` (bit 0) states that the + route is both structurally full and, when token-mask bits are present, + mask-full. """ # Store semantic inputs plus the three validated allocation values. All @@ -238,6 +242,17 @@ def token_words_word_offset(self) -> int | None: return self.route_flags_word_offset + 1 if self.has_token_bits else None + @property + def uses_one_warp_transport(self) -> bool: + """Whether this layout uses the continuous one-warp transport.""" + + token_words_word_offset = self.token_words_word_offset + return ( + not self.is_paged + and token_words_word_offset is not None + and token_words_word_offset + self.token_words_per_route <= 32 + ) + @property def route_metadata_capacity(self) -> int: """Number of routes whose metadata fits in the mutable workspace.""" @@ -249,5 +264,6 @@ def route_metadata_capacity(self) -> int: __all__ = [ "_PREPARED_ROUTE_IS_FULL_FLAG", + "_PREPARED_ROUTE_IS_PROXY_FLAG", "_BlockSparseRouteLayout", ] diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/runtime.py b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/runtime.py index 43e8ff439d31..4f24524e3d4a 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/runtime.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/_block_sparse/runtime.py @@ -16,7 +16,7 @@ from dataclasses import dataclass import math -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Literal import torch @@ -25,6 +25,7 @@ PagedKVCache, _normalize_paged_kv_cache, _validate_16byte_alignment, + _validate_block_table_metadata, _validate_exact_compact_strides, _validate_scale, ) @@ -47,8 +48,7 @@ class _PagedKVStorage: """Paged K/V storage and request metadata consumed by one live run.""" paged_kv_cache: PagedKVCache - paged_kv_indptr: torch.Tensor - paged_kv_indices: torch.Tensor + block_tables: torch.Tensor seq_lens_kv: torch.Tensor @@ -56,8 +56,8 @@ class _PagedKVStorage: class _PagedKVLaunchPayload: """Launch-only live paged metadata derived during shared validation.""" - paged_kv_indptr: torch.Tensor - paged_kv_indices: torch.Tensor + block_tables: torch.Tensor + block_table_row_stride: int seq_lens_kv: torch.Tensor num_physical_kv_pages: int k_page_stride: int @@ -72,12 +72,15 @@ class _BlockSparseRunArgs: k: torch.Tensor v: torch.Tensor out: torch.Tensor - block_indptr: torch.Tensor - block_indices: torch.Tensor + block_indptr: torch.Tensor | None + block_indices: torch.Tensor | None kv_valid_bits: torch.Tensor kv_valid_bits_is_live: bool sm_scale: float paged_kv: _PagedKVLaunchPayload | None + exact_block_bits: torch.Tensor | None = None + k_summary: torch.Tensor | None = None + v_summary: torch.Tensor | None = None def _validate_metadata_tensor( @@ -113,36 +116,71 @@ def _validate_metadata_tensor( def validate_block_sparse_metadata( - block_indptr: torch.Tensor, - block_indices: torch.Tensor, - kv_valid_bits: torch.Tensor | None, *, + sparse_format: Literal["bsr", "bitmask"], + block_indptr: torch.Tensor | None, + block_indices: torch.Tensor | None, + exact_block_bits: torch.Tensor | None, + kv_valid_bits: torch.Tensor | None, device: torch.device, batch_size: int, seq_len_q: int, seq_len_kv: int, num_kv_heads: int, q_block_size: int, + kv_block_size: int, use_kv_valid_bits: bool, ) -> None: - """Validate raw runtime routing without reading device-side values.""" + """Validate the planned route frontend without reading tensor values.""" num_q_blocks = (seq_len_q + q_block_size - 1) // q_block_size - _validate_metadata_tensor( - block_indptr, - "block_indptr", - ndim=3, - dtype=torch.int32, - expected_device=device, - expected_shape=(batch_size, num_kv_heads, num_q_blocks + 1), - ) - _validate_metadata_tensor( - block_indices, - "block_indices", - ndim=1, - dtype=torch.int32, - expected_device=device, - ) + if sparse_format == "bsr": + if block_indptr is None or block_indices is None: + raise ValueError( + "block_indptr and block_indices are required by a BSR plan" + ) + if exact_block_bits is not None: + raise ValueError("exact_block_bits is valid only for a bitmask plan") + _validate_metadata_tensor( + block_indptr, + "block_indptr", + ndim=3, + dtype=torch.int32, + expected_device=device, + expected_shape=(batch_size, num_kv_heads, num_q_blocks + 1), + ) + _validate_metadata_tensor( + block_indices, + "block_indices", + ndim=1, + dtype=torch.int32, + expected_device=device, + ) + elif sparse_format == "bitmask": + if ( + exact_block_bits is None + or block_indptr is not None + or block_indices is not None + ): + raise ValueError( + "runtime route inputs must match planned sparse_format='bitmask'" + ) + num_kv_blocks = (seq_len_kv + kv_block_size - 1) // kv_block_size + _validate_metadata_tensor( + exact_block_bits, + "exact_block_bits", + ndim=4, + dtype=torch.uint32, + expected_device=device, + expected_shape=( + batch_size, + num_kv_heads, + num_q_blocks, + (num_kv_blocks + 31) // 32, + ), + ) + else: + raise AssertionError(f"unsupported sparse format {sparse_format!r}") if use_kv_valid_bits: if kv_valid_bits is None: @@ -159,31 +197,6 @@ def validate_block_sparse_metadata( raise ValueError("kv_valid_bits must be None when use_kv_valid_bits=False") -def validate_paged_kv_metadata( - paged_kv_indptr: torch.Tensor, - paged_kv_indices: torch.Tensor, - seq_lens_kv: torch.Tensor, - *, - device: torch.device, - batch_size: int, -) -> None: - """Validate the shared structural ABI for live paged request metadata.""" - - for tensor, name, shape in ( - (paged_kv_indptr, "paged_kv_indptr", (batch_size + 1,)), - (paged_kv_indices, "paged_kv_indices", None), - (seq_lens_kv, "seq_lens_kv", (batch_size,)), - ): - _validate_metadata_tensor( - tensor, - name, - ndim=1, - dtype=torch.int32, - expected_device=device, - expected_shape=shape, - ) - - def _validate_bshd_tensor( tensor: torch.Tensor, name: str, @@ -214,8 +227,11 @@ def validate_block_sparse_run( kv_storage: _ContiguousKVStorage | _PagedKVStorage, *, state: "_BlockSparsePlanState", - block_indptr: torch.Tensor, - block_indices: torch.Tensor, + block_indptr: torch.Tensor | None, + block_indices: torch.Tensor | None, + exact_block_bits: torch.Tensor | None = None, + k_summary: torch.Tensor | None = None, + v_summary: torch.Tensor | None = None, kv_valid_bits: torch.Tensor | None, sm_scale: float | None, out: torch.Tensor | None, @@ -229,18 +245,44 @@ def validate_block_sparse_run( input. ``sm_scale=None`` is materialized as ``1 / sqrt(D)``. """ + use_proxy_routes = state.use_proxy_routes + num_kv_blocks = (state.seq_len_kv + state.kv_block_size - 1) // state.kv_block_size validate_block_sparse_metadata( - block_indptr, - block_indices, - kv_valid_bits, + sparse_format=state.sparse_format, + block_indptr=block_indptr, + block_indices=block_indices, + exact_block_bits=exact_block_bits, + kv_valid_bits=kv_valid_bits, device=state.device, batch_size=state.batch_size, seq_len_q=state.seq_len_q, seq_len_kv=state.seq_len_kv, num_kv_heads=state.num_kv_heads, q_block_size=state.q_block_size, + kv_block_size=state.kv_block_size, use_kv_valid_bits=state.use_kv_valid_bits, ) + + if use_proxy_routes: + if k_summary is None or v_summary is None: + raise ValueError("K/V summaries are required when proxy routes are enabled") + summary_shape = ( + state.batch_size, + num_kv_blocks, + state.num_kv_heads, + state.head_dim, + ) + for tensor, name in ((k_summary, "k_summary"), (v_summary, "v_summary")): + _validate_bshd_tensor( + tensor, + name, + expected_shape=summary_shape, + expected_dtype=state.kv_dtype, + expected_device=state.device, + ) + elif k_summary is not None or v_summary is not None: + raise ValueError("summaries are valid only when proxy routes are enabled") + if state.use_kv_valid_bits: assert kv_valid_bits is not None effective_kv_valid_bits = kv_valid_bits @@ -258,7 +300,7 @@ def validate_block_sparse_run( expected_device=state.device, ) paged_kv: _PagedKVLaunchPayload | None = None - overlap_inputs: tuple[tuple[str, torch.Tensor], ...] + overlap_inputs: list[tuple[str, torch.Tensor]] if isinstance(kv_storage, _ContiguousKVStorage): if state.page_size is not None: raise TypeError("contiguous K/V storage requires a contiguous plan state") @@ -278,16 +320,11 @@ def validate_block_sparse_run( ) k = kv_storage.k v = kv_storage.v - overlap_inputs = ( + overlap_inputs = [ ("q", q), ("k", k), ("v", v), - ("block_indptr", block_indptr), - ("block_indices", block_indices), - ("kv_valid_bits", effective_kv_valid_bits), - ("row_route_offsets", state.row_route_offsets), - ("route_workspace", state.route_workspace), - ) + ] elif isinstance(kv_storage, _PagedKVStorage): page_size = state.page_size if page_size is None: @@ -324,36 +361,60 @@ def validate_block_sparse_run( raise ValueError( f"K/V dtype must match the plan ({state.kv_dtype}), got {k.dtype}" ) - validate_paged_kv_metadata( - kv_storage.paged_kv_indptr, - kv_storage.paged_kv_indices, - kv_storage.seq_lens_kv, - device=state.device, - batch_size=state.batch_size, + metadata_device, metadata_batch_size, table_capacity = ( + _validate_block_table_metadata( + kv_storage.block_tables, + kv_storage.seq_lens_kv, + ) ) + if metadata_device != state.device: + raise ValueError( + f"per-run metadata must be on {state.device}, got {metadata_device}" + ) + if metadata_batch_size != state.batch_size: + raise ValueError( + "per-run metadata batch size must match the plan " + f"({state.batch_size}), got {metadata_batch_size}" + ) + if table_capacity * page_size < state.seq_len_kv: + raise ValueError( + "block_tables must cover the planned K/V capacity: expected at " + f"least {(state.seq_len_kv + page_size - 1) // page_size} columns, " + f"got {table_capacity}" + ) paged_kv = _PagedKVLaunchPayload( - paged_kv_indptr=kv_storage.paged_kv_indptr, - paged_kv_indices=kv_storage.paged_kv_indices, + block_tables=kv_storage.block_tables, + block_table_row_stride=kv_storage.block_tables.stride(0), seq_lens_kv=kv_storage.seq_lens_kv, num_physical_kv_pages=num_physical_kv_pages, k_page_stride=k_page_stride, v_page_stride=v_page_stride, ) - overlap_inputs = ( + overlap_inputs = [ ("q", q), ("k_cache", k), ("v_cache", v), - ("block_indptr", block_indptr), - ("block_indices", block_indices), - ("kv_valid_bits", effective_kv_valid_bits), - ("paged_kv_indptr", kv_storage.paged_kv_indptr), - ("paged_kv_indices", kv_storage.paged_kv_indices), + ("block_tables", kv_storage.block_tables), ("seq_lens_kv", kv_storage.seq_lens_kv), + ] + else: + raise TypeError("kv_storage must be _ContiguousKVStorage or _PagedKVStorage") + + if block_indptr is not None and block_indices is not None: + overlap_inputs.extend( + (("block_indptr", block_indptr), ("block_indices", block_indices)) + ) + if exact_block_bits is not None: + overlap_inputs.append(("exact_block_bits", exact_block_bits)) + if k_summary is not None and v_summary is not None: + overlap_inputs.extend((("k_summary", k_summary), ("v_summary", v_summary))) + overlap_inputs.extend( + ( + ("kv_valid_bits", effective_kv_valid_bits), ("row_route_offsets", state.row_route_offsets), ("route_workspace", state.route_workspace), ) - else: - raise TypeError("kv_storage must be _ContiguousKVStorage or _PagedKVStorage") + ) effective_scale = _validate_scale( 1.0 / math.sqrt(state.head_dim) if sm_scale is None else sm_scale, @@ -377,6 +438,9 @@ def validate_block_sparse_run( out=out, block_indptr=block_indptr, block_indices=block_indices, + exact_block_bits=exact_block_bits, + k_summary=k_summary, + v_summary=v_summary, kv_valid_bits=effective_kv_valid_bits, kv_valid_bits_is_live=state.use_kv_valid_bits, sm_scale=effective_scale, @@ -384,23 +448,100 @@ def validate_block_sparse_run( ) +def prepare_block_sparse_run_unchecked( + q: torch.Tensor, + kv_storage: _ContiguousKVStorage | _PagedKVStorage, + *, + state: "_BlockSparsePlanState", + block_indptr: torch.Tensor | None, + block_indices: torch.Tensor | None, + exact_block_bits: torch.Tensor | None = None, + k_summary: torch.Tensor | None = None, + v_summary: torch.Tensor | None = None, + kv_valid_bits: torch.Tensor | None, + sm_scale: float | None, + out: torch.Tensor | None, +) -> _BlockSparseRunArgs: + """Canonicalize one trusted run without invoking explicit validators. + + Only the work every launch needs happens here: K/V view selection, the + plan-owned dummy token mask when the plan disabled token bits, the default + softmax scale, and allocation of an omitted output tensor. + """ + + paged_kv: _PagedKVLaunchPayload | None = None + if isinstance(kv_storage, _ContiguousKVStorage): + k = kv_storage.k + v = kv_storage.v + else: + paged_kv_cache = kv_storage.paged_kv_cache + if isinstance(paged_kv_cache, torch.Tensor): + k = paged_kv_cache[:, 0] + v = paged_kv_cache[:, 1] + else: + k, v = paged_kv_cache + paged_kv = _PagedKVLaunchPayload( + block_tables=kv_storage.block_tables, + block_table_row_stride=kv_storage.block_tables.stride(0), + seq_lens_kv=kv_storage.seq_lens_kv, + num_physical_kv_pages=int(k.shape[0]), + k_page_stride=int(k.stride(0)), + v_page_stride=int(v.stride(0)), + ) + if state.use_kv_valid_bits: + effective_kv_valid_bits = kv_valid_bits + else: + effective_kv_valid_bits = state.dummy_kv_valid_bits + assert effective_kv_valid_bits is not None + if out is None: + out = torch.empty( + (state.batch_size, state.seq_len_q, state.num_qo_heads, state.head_dim), + device=state.device, + dtype=state.output_dtype, + ) + return _BlockSparseRunArgs( + q=q, + k=k, + v=v, + out=out, + block_indptr=block_indptr, + block_indices=block_indices, + exact_block_bits=exact_block_bits, + k_summary=k_summary, + v_summary=v_summary, + kv_valid_bits=effective_kv_valid_bits, + kv_valid_bits_is_live=state.use_kv_valid_bits, + sm_scale=1.0 / math.sqrt(state.head_dim) + if sm_scale is None + else float(sm_scale), + paged_kv=paged_kv, + ) + + def record_block_sparse_run_args( run_args: _BlockSparseRunArgs, stream: torch.cuda.Stream, ) -> None: """Extend tensor lifetimes for the asynchronous launch currently in flight.""" - run_args.q.record_stream(stream) - run_args.k.record_stream(stream) - run_args.v.record_stream(stream) + for tensor in (run_args.q, run_args.k, run_args.v): + tensor.record_stream(stream) + if run_args.k_summary is not None: + run_args.k_summary.record_stream(stream) + assert run_args.v_summary is not None + run_args.v_summary.record_stream(stream) run_args.out.record_stream(stream) - run_args.block_indptr.record_stream(stream) - run_args.block_indices.record_stream(stream) + if run_args.block_indptr is not None: + run_args.block_indptr.record_stream(stream) + assert run_args.block_indices is not None + run_args.block_indices.record_stream(stream) + else: + assert run_args.exact_block_bits is not None + run_args.exact_block_bits.record_stream(stream) if run_args.kv_valid_bits_is_live: run_args.kv_valid_bits.record_stream(stream) if run_args.paged_kv is not None: - run_args.paged_kv.paged_kv_indptr.record_stream(stream) - run_args.paged_kv.paged_kv_indices.record_stream(stream) + run_args.paged_kv.block_tables.record_stream(stream) run_args.paged_kv.seq_lens_kv.record_stream(stream) @@ -409,9 +550,13 @@ def launch_block_sparse( *, state: "_BlockSparsePlanState", ) -> torch.Tensor: - """Invoke the exact contiguous or paged ABI chosen by validated payload.""" + """Invoke the layout- and route-specific ABI chosen by the frozen plan.""" - if run_args.paged_kv is None: + sparse_format = state.sparse_format + use_proxy_routes = state.use_proxy_routes + if run_args.paged_kv is not None: + assert run_args.block_indptr is not None + assert run_args.block_indices is not None state.compiled( run_args.q, run_args.k, @@ -420,12 +565,20 @@ def launch_block_sparse( run_args.block_indptr, run_args.block_indices, run_args.kv_valid_bits, + run_args.paged_kv.block_tables, + run_args.paged_kv.seq_lens_kv, state.row_route_offsets, state.route_workspace, state.max_blocks_per_row, + run_args.paged_kv.num_physical_kv_pages, + run_args.paged_kv.block_table_row_stride, + run_args.paged_kv.k_page_stride, + run_args.paged_kv.v_page_stride, run_args.sm_scale, ) - else: + elif sparse_format == "bsr" and not use_proxy_routes: + assert run_args.block_indptr is not None + assert run_args.block_indices is not None state.compiled( run_args.q, run_args.k, @@ -434,17 +587,65 @@ def launch_block_sparse( run_args.block_indptr, run_args.block_indices, run_args.kv_valid_bits, - run_args.paged_kv.paged_kv_indptr, - run_args.paged_kv.paged_kv_indices, - run_args.paged_kv.seq_lens_kv, state.row_route_offsets, state.route_workspace, state.max_blocks_per_row, - run_args.paged_kv.num_physical_kv_pages, - run_args.paged_kv.k_page_stride, - run_args.paged_kv.v_page_stride, run_args.sm_scale, ) + elif sparse_format == "bitmask" and not use_proxy_routes: + assert run_args.exact_block_bits is not None + state.compiled( + run_args.q, + run_args.k, + run_args.v, + run_args.out, + run_args.exact_block_bits, + run_args.kv_valid_bits, + state.row_route_offsets, + state.route_workspace, + state.max_blocks_per_row, + run_args.sm_scale, + ) + elif sparse_format == "bsr" and use_proxy_routes: + assert run_args.block_indptr is not None + assert run_args.block_indices is not None + assert run_args.k_summary is not None + assert run_args.v_summary is not None + state.compiled( + run_args.q, + run_args.k, + run_args.v, + run_args.k_summary, + run_args.v_summary, + run_args.out, + run_args.block_indptr, + run_args.block_indices, + run_args.kv_valid_bits, + state.row_route_offsets, + state.route_workspace, + state.max_blocks_per_row, + run_args.sm_scale, + ) + elif sparse_format == "bitmask" and use_proxy_routes: + assert run_args.exact_block_bits is not None + assert run_args.k_summary is not None + assert run_args.v_summary is not None + state.compiled( + run_args.q, + run_args.k, + run_args.v, + run_args.k_summary, + run_args.v_summary, + run_args.out, + run_args.exact_block_bits, + run_args.kv_valid_bits, + state.row_route_offsets, + state.route_workspace, + state.max_blocks_per_row, + run_args.sm_scale, + ) + else: + raise AssertionError("frozen block-sparse plan has an unsupported route mode") return run_args.out @@ -454,8 +655,8 @@ def launch_block_sparse( "_PagedKVLaunchPayload", "_PagedKVStorage", "launch_block_sparse", + "prepare_block_sparse_run_unchecked", "record_block_sparse_run_args", "validate_block_sparse_metadata", "validate_block_sparse_run", - "validate_paged_kv_metadata", ] diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/block_sparse.py b/tensorrt_llm/_torch/attention/backends/prims_ts/block_sparse.py index af132c6a69ce..e19dc7b1de7f 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/block_sparse.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/block_sparse.py @@ -37,6 +37,7 @@ from flashinfer.api_logging import flashinfer_api +from ._block_sparse.common import _validate_contiguous_route_mode from ._block_sparse.config import _validate_block_sparse_static_profile from ._block_sparse.inspection import ( _inspect_block_sparse_bsr, @@ -53,15 +54,16 @@ _ContiguousKVStorage, _PagedKVStorage, launch_block_sparse as _launch_block_sparse, + prepare_block_sparse_run_unchecked as _prepare_block_sparse_run_unchecked, record_block_sparse_run_args as _record_block_sparse_run_args, validate_block_sparse_metadata as _validate_block_sparse_metadata, validate_block_sparse_run as _validate_block_sparse_run, - validate_paged_kv_metadata as _validate_paged_kv_metadata, ) from .decode import ( PagedKVCache, _normalize_paged_kv_cache, _resolve_cuda_device, + _validate_block_table_metadata, ) @@ -111,10 +113,12 @@ class BlockSparseTSWrapper(_BlockSparseWrapperBase): Q is ``[B, Sq, Hq, D]`` and K/V are ``[B, Skv, Hkv, D]``. Sparse rows are owned per batch, KV head, and query block, so every Q head in one grouped KV head consumes the same sparse row. A plan fixes geometry and a per-row - capacity; every run supplies its own BSR and optional token mask. - Callers must keep those tensors alive and immutable until the queued run or - captured graph finishes using them. CUDA Graph capture pins plan-owned - state only, so captured routing storage remains the caller's responsibility. + capacity; every run supplies either BSR or a packed exact-block bitmask. + Proxy-enabled plans additionally consume caller-owned K/V summaries, while + an optional token mask applies only to exact routes. Callers must keep those + tensors alive and immutable until the queued run or captured graph finishes + using them. CUDA Graph capture pins plan-owned state only, so captured + routing storage remains the caller's responsibility. One plan revision owns one mutable route workspace. Its runs must be ordered on one stream or externally synchronized; unordered concurrent runs require @@ -136,6 +140,8 @@ def plan( device: torch.device | str | int, max_blocks_per_row: int, use_kv_valid_bits: bool, + sparse_format: Literal["bsr", "bitmask"] = "bsr", + use_proxy_routes: bool = False, mask_type: Literal["dense", "causal"] = "dense", q_data_type: torch.dtype = torch.float16, kv_data_type: torch.dtype | None = None, @@ -144,11 +150,16 @@ def plan( """Choose a legal profile and allocate reusable routing capacity. The plan owns immutable geometry and a uniform route workspace, not a - sparse pattern. ``max_blocks_per_row`` bounds each runtime BSR row in - semantic ``kv_block_size`` blocks. ``use_kv_valid_bits`` selects whether - every :meth:`run` must supply the shared batch token mask. Callers may - pass different routing tensor identities and index extents to each run - as long as they fit this declared capacity. + sparse pattern. ``max_blocks_per_row`` bounds each runtime sparse row in + semantic ``kv_block_size`` blocks. ``sparse_format="bsr"`` consumes + canonical CSR-style rows, while ``"bitmask"`` consumes packed exact- + block bits. Enabling proxy routes represents unselected blocks through + caller-provided K/V summaries and currently requires + ``mask_type="dense"``. Exact-only plans continue to support causal + masking. ``use_kv_valid_bits`` selects whether every :meth:`run` must + supply the shared batch token mask. Callers may pass different routing + tensor identities and index extents to each run as long as they fit + this declared capacity. MHA, GQA, and MQA are supported with ``Hq / Hkv`` a power of two no greater than 32 and ``D=128``. Q, K, V, and O use one matching @@ -162,10 +173,13 @@ def plan( respectively. ``kv_block_size`` may be 8, 16, 32, or a positive multiple of 64. The Q tile groups complete Q-head groups and as many Q tokens as fit without crossing a semantic Q-block row, up to Q128; - fine KV blocks cap this at a SWAPAB Q32 tile. Every run prepares - per-KV-head canonical BSR into compact, profile-selected fixed-width - route metadata, and the attention core consumes only that metadata. - This remains true when every KV block is selected; + fine KV blocks cap this at a SWAPAB Q32 tile. Proxy routes reuse the + same Q-tile, KV-route, and MMA geometry as exact routes, but currently + use the direct scheduler because reusable planning cannot observe live + exact-route work. Every run prepares its selected BSR or bitmask into + compact, profile-selected fixed-width route metadata, and the attention + core consumes only that metadata. This remains true when every KV block + is selected; callers that know a pattern is dense should choose the dense FMHA API explicitly. @@ -182,6 +196,7 @@ def plan( runs require distinct wrappers. """ + _validate_contiguous_route_mode(sparse_format, use_proxy_routes) static = _validate_block_sparse_static_profile( batch_size=batch_size, seq_len_q=seq_len_q, @@ -198,6 +213,8 @@ def plan( output_dtype=o_data_type, max_blocks_per_row=max_blocks_per_row, ) + if use_proxy_routes and static.mask_type != "dense": + raise ValueError("block-sparse proxy routes require mask_type='dense'") device, device_index = _resolve_cuda_device(device) plan_stream = torch.cuda.current_stream(device) with torch.cuda.device(device_index), torch.cuda.stream(plan_stream): @@ -210,6 +227,8 @@ def plan( device=device, device_index=device_index, plan_stream=plan_stream, + sparse_format=sparse_format, + use_proxy_routes=use_proxy_routes, ) # This is the only wrapper mutation. Every failure above leaves the # previously published revision intact and runnable. @@ -221,12 +240,16 @@ def run( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, - block_indptr: torch.Tensor, - block_indices: torch.Tensor, + block_indptr: torch.Tensor | None = None, + block_indices: torch.Tensor | None = None, *, + exact_block_bits: torch.Tensor | None = None, + k_summary: torch.Tensor | None = None, + v_summary: torch.Tensor | None = None, kv_valid_bits: torch.Tensor | None = None, sm_scale: float | None = None, out: torch.Tensor | None = None, + validate: bool = True, ) -> torch.Tensor: """Launch the current plan on the caller's current CUDA stream. @@ -238,17 +261,34 @@ def run( Only O is returned; this PrimTS API does not return LSE. The launch is enqueued asynchronously on the caller's current CUDA stream. - ``block_indptr`` is compact Int32 - ``[B, Hkv, ceil(Sq / q_block_size) + 1]`` and indexes compact - ``block_indices``. Every row must fit the planned semantic-block - capacity; referenced block IDs must be strictly increasing, unique, - and in range. Reusable runs trust these device-side values. CuTe DSL - assertions can diagnose violations when enabled before compilation; - otherwise invalid values have undefined behavior and may access out of - bounds. A masked plan requires + ``validate=True`` performs structural, plan-geometry, and alias + validation without reading tensor values; it is the safe public + default. ``validate=False`` treats every run argument as a trusted + binding and performs no explicit wrapper validation. K/V view + selection, scale forwarding, and optional output allocation are + unavoidable in both modes. + + A BSR plan consumes compact Int32 ``block_indptr`` with shape + ``[B, Hkv, ceil(Sq / q_block_size) + 1]`` and compact Int32 + ``block_indices``. A bitmask plan instead requires both BSR arguments + to be ``None`` and consumes packed UInt32 ``exact_block_bits`` with + shape ``[B, Hkv, ceil(Sq / q_block_size), ceil(num_kv_blocks / 32)]``. + Bit ``r`` of word ``w`` selects block ``32 * w + r``; final-word + padding bits are ignored. A proxy plan additionally consumes compact + ``k_summary`` and ``v_summary`` with shape + ``[B, num_kv_blocks, Hkv, D]``. K summaries are block means and V + summaries are block sums; the final partial block covers only its + structural tokens. + + Every row must fit the planned semantic-block capacity. Reusable runs + trust routing values. CuTe DSL assertions can diagnose violations when + enabled before compilation; otherwise invalid values have undefined + behavior and may access out of bounds. A masked plan requires ``kv_valid_bits`` with shape ``[B, ceil(Skv / 32)]`` and dtype UInt32; an unmasked plan requires - ``None``. Routing tensors may have different identities on every run. + ``None``. The mask applies only to raw exact routes; proxy summaries and + their represented-token mass remain caller-defined. Routing tensors may + have different identities on every run. Keep this wrapper alive until every captured CUDA Graph is destroyed. @@ -261,12 +301,18 @@ def run( v : torch.Tensor Compact value tensor with the same shape, dtype, and strides as ``k``. - block_indptr : torch.Tensor + block_indptr : torch.Tensor, optional Contiguous Int32 BSR row offsets with shape - ``[B, Hkv, ceil(Sq / q_block_size) + 1]``. - block_indices : torch.Tensor + ``[B, Hkv, ceil(Sq / q_block_size) + 1]``. Required by BSR plans. + block_indices : torch.Tensor, optional Contiguous Int32 semantic KV-block IDs referenced by - ``block_indptr``. + ``block_indptr``. Required by BSR plans. + exact_block_bits : torch.Tensor, optional + Compact packed UInt32 exact-block bitmap required by bitmask plans. + k_summary : torch.Tensor, optional + Per-block mean K tensor required by proxy plans. + v_summary : torch.Tensor, optional + Per-block summed V tensor required by proxy plans. kv_valid_bits : torch.Tensor, optional Contiguous UInt32 token-validity bitmap ``[B, ceil(Skv / 32)]``. Supply it exactly when the plan enabled token validity bits. @@ -275,6 +321,9 @@ def run( out : torch.Tensor, optional Caller-owned compact output buffer ``[B, Sq, Hq, D]`` with the planned output dtype. + validate : bool + Whether to validate tensor structure, plan geometry, and aliasing + before launching. Defaults to ``True``. Returns ------- @@ -283,17 +332,27 @@ def run( """ state = self._require_run_state() - run_stream = torch.cuda.current_stream(state.device) - run_args = _validate_block_sparse_run( + if not isinstance(validate, bool): + raise TypeError("validate must be a bool") + prepare_run = ( + _validate_block_sparse_run + if validate + else _prepare_block_sparse_run_unchecked + ) + run_args = prepare_run( q, _ContiguousKVStorage(k=k, v=v), state=state, block_indptr=block_indptr, block_indices=block_indices, + exact_block_bits=exact_block_bits, + k_summary=k_summary, + v_summary=v_summary, kv_valid_bits=kv_valid_bits, sm_scale=sm_scale, out=out, ) + run_stream = torch.cuda.current_stream(state.device) return self._launch_validated_run(state, run_args, run_stream) @@ -302,23 +361,29 @@ def block_sparse_attention( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, - block_indptr: torch.Tensor, - block_indices: torch.Tensor, + block_indptr: torch.Tensor | None, + block_indices: torch.Tensor | None, q_block_size: int, kv_block_size: int, *, + exact_block_bits: torch.Tensor | None = None, + k_summary: torch.Tensor | None = None, + v_summary: torch.Tensor | None = None, kv_valid_bits: torch.Tensor | None = None, + sparse_format: Literal["bsr", "bitmask"] = "bsr", + use_proxy_routes: bool = False, mask_type: Literal["dense", "causal"] = "dense", sm_scale: float | None = None, out: torch.Tensor | None = None, ) -> torch.Tensor: """Plan and run one compact-BSHD block-sparse attention launch. - This one-shot form synchronously inspects canonical BSR, derives its largest - semantic row, creates a capacity-only plan, and passes the original routing - tensors to :meth:`BlockSparseTSWrapper.run`. It therefore cannot be invoked - inside CUDA Graph capture; plan a wrapper outside capture and capture only - ``run()`` instead. + This one-shot form creates a capacity-only plan and passes the original + routing tensors to :meth:`BlockSparseTSWrapper.run`. BSR inputs are + synchronously inspected to validate canonical rows and derive their maximum + width. Bitmask inputs use the structural KV-block count as a conservative + capacity bound. It therefore cannot be invoked inside CUDA Graph capture; + plan a wrapper outside capture and capture only ``run()`` instead. Parameters ---------- @@ -328,11 +393,12 @@ def block_sparse_attention( Compact key tensor ``[B, Skv, Hkv, D]``. v : torch.Tensor Compact value tensor with the same shape, dtype, and strides as ``k``. - block_indptr : torch.Tensor + block_indptr : torch.Tensor, optional Contiguous Int32 BSR row offsets with shape - ``[B, Hkv, ceil(Sq / q_block_size) + 1]``. - block_indices : torch.Tensor + ``[B, Hkv, ceil(Sq / q_block_size) + 1]``. Required in BSR mode. + block_indices : torch.Tensor, optional Contiguous Int32 semantic KV-block IDs referenced by ``block_indptr``. + Required in BSR mode. q_block_size : int Positive number of logical query tokens represented by one BSR row. The product with ``Hq / Hkv`` must be divisible by 8 so a physical Q @@ -340,8 +406,19 @@ def block_sparse_attention( kv_block_size : int Number of logical KV tokens represented by one BSR block ID; it must be 8, 16, 32, or a positive multiple of 64. + exact_block_bits : torch.Tensor, optional + Compact UInt32 exact-block bitmap required in bitmask mode. + k_summary : torch.Tensor, optional + Per-block mean K tensor required when proxy routes are enabled. + v_summary : torch.Tensor, optional + Per-block summed V tensor required when proxy routes are enabled. kv_valid_bits : torch.Tensor, optional Contiguous UInt32 token-validity bitmap ``[B, ceil(Skv / 32)]``. + sparse_format : {"bsr", "bitmask"}, optional + Runtime sparse representation. Defaults to ``"bsr"``. + use_proxy_routes : bool, optional + Represent unselected blocks through K/V summaries. Proxy routes + currently require dense masking. mask_type : {"dense", "causal"}, optional Attention mask applied inside each selected sparse block. sm_scale : float, optional @@ -370,6 +447,7 @@ def block_sparse_attention( raise ValueError("K and V must have identical shapes") use_kv_valid_bits = kv_valid_bits is not None + _validate_contiguous_route_mode(sparse_format, use_proxy_routes) static = _validate_block_sparse_static_profile( batch_size=batch_size, seq_len_q=seq_len_q, @@ -385,25 +463,37 @@ def block_sparse_attention( kv_dtype=k.dtype, output_dtype=q.dtype if out is None else out.dtype, ) + if use_proxy_routes and static.mask_type != "dense": + raise ValueError("block-sparse proxy routes require mask_type='dense'") device, _ = _resolve_cuda_device(q.device) _validate_block_sparse_metadata( - block_indptr, - block_indices, - kv_valid_bits, + sparse_format=sparse_format, + block_indptr=block_indptr, + block_indices=block_indices, + exact_block_bits=exact_block_bits, + kv_valid_bits=kv_valid_bits, device=device, batch_size=static.batch_size, seq_len_q=static.seq_len_q, seq_len_kv=static.seq_len_kv, num_kv_heads=static.num_kv_heads, q_block_size=static.q_block_size, + kv_block_size=static.kv_block_size, use_kv_valid_bits=static.use_kv_valid_bits, ) - max_blocks_per_row = _inspect_block_sparse_bsr( - block_indptr, - block_indices, - static=static, - stream=torch.cuda.current_stream(device), - ) + if sparse_format == "bsr": + max_blocks_per_row = _inspect_block_sparse_bsr( + block_indptr, + block_indices, + static=static, + stream=torch.cuda.current_stream(device), + ) + elif sparse_format == "bitmask": + max_blocks_per_row = ( + static.seq_len_kv + static.kv_block_size - 1 + ) // static.kv_block_size + else: + raise AssertionError(f"unsupported sparse format {sparse_format!r}") wrapper = BlockSparseTSWrapper() wrapper.plan( @@ -418,6 +508,8 @@ def block_sparse_attention( device=device, max_blocks_per_row=max_blocks_per_row, use_kv_valid_bits=static.use_kv_valid_bits, + sparse_format=sparse_format, + use_proxy_routes=use_proxy_routes, mask_type=static.mask_type, q_data_type=static.q_dtype, kv_data_type=static.kv_dtype, @@ -429,6 +521,9 @@ def block_sparse_attention( v, block_indptr, block_indices, + exact_block_bits=exact_block_bits, + k_summary=k_summary, + v_summary=v_summary, kv_valid_bits=kv_valid_bits, sm_scale=sm_scale, out=out, @@ -512,8 +607,7 @@ def run( self, q: torch.Tensor, paged_kv_cache: PagedKVCache, - paged_kv_indptr: torch.Tensor, - paged_kv_indices: torch.Tensor, + block_tables: torch.Tensor, seq_lens_kv: torch.Tensor, block_indptr: torch.Tensor, block_indices: torch.Tensor, @@ -521,6 +615,7 @@ def run( kv_valid_bits: torch.Tensor | None = None, sm_scale: float | None = None, out: torch.Tensor | None = None, + validate: bool = True, ) -> torch.Tensor: """Launch with live lengths, page tables, and sparse routes. @@ -529,15 +624,20 @@ def run( tuple whose members are ``[P, Hkv, page, D]`` with compact inner HND strides and arbitrary non-overlapping outer page strides. - ``paged_kv_indptr`` is compact Int32 ``[B + 1]``; - ``paged_kv_indices`` is compact Int32 with capacity at least its live - final offset; and ``seq_lens_kv`` is compact Int32 ``[B]``. All values - are read on device. The caller must keep every dense length in - ``[1, max_seq_len_kv]`` and every causal length in - ``[Sq, max_seq_len_kv]``. ``paged_kv_indptr`` must start at zero and - contain bounded, monotone rows with at least - ``ceil(seq_lens_kv[b] / page_size)`` entries. Every page ID in its live - prefix must lie in ``[0, P)``. Each BSR row must contain strictly + ``validate=True`` performs structural, plan-geometry, and alias + validation without reading tensor values; it is the safe public + default. ``validate=False`` treats every run argument as a trusted + binding and performs no explicit wrapper validation. K/V view + selection, scale forwarding, and optional output allocation are + unavoidable in both modes. + + ``block_tables`` is Int32 ``[B, C]``, contiguous within each row but + permitted to use a padded outer row stride; ``seq_lens_kv`` is compact + Int32 ``[B]``. All values are read on device. The caller must keep + every dense length in ``[1, max_seq_len_kv]`` and every causal length + in ``[Sq, max_seq_len_kv]``. Every page-table row must contain at least + ``ceil(seq_lens_kv[b] / page_size)`` live entries. Every page ID in its + live prefix must lie in ``[0, P)``. Each BSR row must contain strictly increasing, unique block IDs whose final block starts before that request's live K/V length, and its width must not exceed the planned ``max_blocks_per_row``. Reusable runs trust all of these device-side @@ -564,10 +664,10 @@ def run( Either a combined cache ``[P, 2, Hkv, page_size, D]`` or a ``(K, V)`` tuple whose tensors are ``[P, Hkv, page_size, D]``. - paged_kv_indptr : torch.Tensor - Contiguous Int32 live request offsets with shape ``[B + 1]``. - paged_kv_indices : torch.Tensor - Contiguous Int32 physical-page ID capacity. + block_tables : torch.Tensor + Live Int32 physical page IDs with shape ``[B, C]``. Entries are + contiguous within each row; padded, non-overlapping row strides are + supported and inactive tail entries are ignored. seq_lens_kv : torch.Tensor Contiguous Int32 live logical K/V lengths with shape ``[B]``. Values must satisfy the dense or causal bounds above. @@ -586,6 +686,9 @@ def run( out : torch.Tensor, optional Caller-owned compact output buffer ``[B, Sq, Hq, D]`` with the planned output dtype. + validate : bool + Whether to validate tensor structure, plan geometry, and aliasing + before launching. Defaults to ``True``. Returns ------- @@ -594,13 +697,19 @@ def run( """ state = self._require_run_state() + if not isinstance(validate, bool): + raise TypeError("validate must be a bool") run_stream = torch.cuda.current_stream(state.device) - run_args = _validate_block_sparse_run( + prepare_run = ( + _validate_block_sparse_run + if validate + else _prepare_block_sparse_run_unchecked + ) + run_args = prepare_run( q, _PagedKVStorage( paged_kv_cache=paged_kv_cache, - paged_kv_indptr=paged_kv_indptr, - paged_kv_indices=paged_kv_indices, + block_tables=block_tables, seq_lens_kv=seq_lens_kv, ), state=state, @@ -617,15 +726,14 @@ def run( def block_sparse_attention_with_paged_kv_cache( q: torch.Tensor, paged_kv_cache: PagedKVCache, - paged_kv_indptr: torch.Tensor, - paged_kv_indices: torch.Tensor, + block_tables: torch.Tensor, + seq_lens_kv: torch.Tensor, block_indptr: torch.Tensor, block_indices: torch.Tensor, q_block_size: int, kv_block_size: int, *, max_seq_len_kv: int, - seq_lens_kv: torch.Tensor, kv_valid_bits: torch.Tensor | None = None, mask_type: Literal["dense", "causal"] = "dense", sm_scale: float | None = None, @@ -633,12 +741,11 @@ def block_sparse_attention_with_paged_kv_cache( ) -> torch.Tensor: """Plan and run one fixed-Q paged block-sparse attention launch. - This convenience entry point synchronously validates live page and sparse - metadata, including the complete live physical-page-ID prefix, creates a - capacity-only temporary plan, then forwards the inspected tensors through - the trusted live run API. It cannot run during CUDA Graph capture; plan a - wrapper outside capture and capture only - :meth:`BlockSparsePagedTSWrapper.run` instead. + This convenience entry point synchronously validates the live page tables, + K/V lengths, and sparse metadata, creates a capacity-only temporary plan, + then forwards the inspected tensors through the trusted live run API. It + cannot run during CUDA Graph capture; plan a wrapper outside capture and + capture only :meth:`BlockSparsePagedTSWrapper.run` instead. Parameters ---------- @@ -647,11 +754,13 @@ def block_sparse_attention_with_paged_kv_cache( paged_kv_cache : PagedKVCache Either a combined cache ``[P, 2, Hkv, page_size, D]`` or a ``(K, V)`` tuple whose tensors are ``[P, Hkv, page_size, D]``. - paged_kv_indptr : torch.Tensor - Contiguous Int32 request offsets into ``paged_kv_indices``, with shape - ``[B + 1]``. - paged_kv_indices : torch.Tensor - Contiguous Int32 physical page IDs referenced by ``paged_kv_indptr``. + block_tables : torch.Tensor + Int32 physical page IDs ``[B, C]``, contiguous within each row and free + to use a padded outer row stride. ``C * page_size`` must cover + ``max_seq_len_kv``; only the first ``ceil(seq_lens_kv[b] / page_size)`` + entries of each row are read. + seq_lens_kv : torch.Tensor + Contiguous Int32 per-request logical KV lengths with shape ``[B]``. block_indptr : torch.Tensor Contiguous Int32 BSR row offsets with shape ``[B, Hkv, ceil(Sq / q_block_size) + 1]``. @@ -666,8 +775,6 @@ def block_sparse_attention_with_paged_kv_cache( be 8, 16, 32, or a positive multiple of 64. max_seq_len_kv : int Static maximum logical K/V length used for planning. - seq_lens_kv : torch.Tensor - Contiguous Int32 per-request logical KV lengths with shape ``[B]``. kv_valid_bits : torch.Tensor, optional Contiguous UInt32 logical-token validity bitmap ``[B, ceil(max_seq_len_kv / 32)]``. @@ -693,13 +800,16 @@ def block_sparse_attention_with_paged_kv_cache( batch_size, seq_len_q, num_qo_heads, head_dim = map(int, q.shape) metadata_device, _ = _resolve_cuda_device(q.device) - _validate_paged_kv_metadata( - paged_kv_indptr, - paged_kv_indices, - seq_lens_kv, - device=metadata_device, - batch_size=batch_size, + table_device, table_batch_size, table_capacity = _validate_block_table_metadata( + block_tables, seq_lens_kv ) + if table_device != q.device: + raise ValueError(f"paged-KV metadata must be on {q.device}, got {table_device}") + if table_batch_size != batch_size: + raise ValueError( + "seq_lens_kv must have one entry per request: " + f"expected {batch_size}, got {table_batch_size}" + ) ( k_cache, @@ -734,23 +844,31 @@ def block_sparse_attention_with_paged_kv_cache( kv_dtype=k_cache.dtype, output_dtype=q.dtype if out is None else out.dtype, ) + if table_capacity * page_size < static.seq_len_kv: + raise ValueError( + "block_tables must cover the planned K/V capacity: expected at " + f"least {(static.seq_len_kv + page_size - 1) // page_size} columns, " + f"got {table_capacity}" + ) _validate_block_sparse_metadata( - block_indptr, - block_indices, - kv_valid_bits, + sparse_format="bsr", + block_indptr=block_indptr, + block_indices=block_indices, + exact_block_bits=None, + kv_valid_bits=kv_valid_bits, device=metadata_device, batch_size=static.batch_size, seq_len_q=static.seq_len_q, seq_len_kv=static.seq_len_kv, num_kv_heads=static.num_kv_heads, q_block_size=static.q_block_size, + kv_block_size=static.kv_block_size, use_kv_valid_bits=static.use_kv_valid_bits, ) max_blocks_per_row = _inspect_paged_block_sparse_metadata( block_indptr, block_indices, - paged_kv_indptr, - paged_kv_indices, + block_tables, seq_lens_kv, static=static, num_physical_kv_pages=num_physical_kv_pages, @@ -780,8 +898,7 @@ def block_sparse_attention_with_paged_kv_cache( return wrapper.run( q, paged_kv_cache, - paged_kv_indptr, - paged_kv_indices, + block_tables, seq_lens_kv, block_indptr, block_indices, diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/block_sparse_inspect.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/block_sparse_inspect.py index de53ccb9c7de..5728a7c6d373 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/block_sparse_inspect.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/block_sparse_inspect.py @@ -20,10 +20,10 @@ their caller. Token-mask contents belong to the run-time prepare kernel and are not read here. -Paged inspection first validates live sequence lengths and page rows, then -four warps validate four BSR Q-block rows per CTA. Both publish one validation -status plus the maximum row width in one Int64 summary; no route payload is -constructed. +Paged inspection first validates live sequence lengths and the live prefix of +every page-table row, then four warps validate four BSR Q-block rows per CTA. +Both publish one validation status plus the maximum row width in one Int64 +summary; no route payload is constructed. """ import functools @@ -51,9 +51,7 @@ _BSR_ERROR_INDEX_OUT_OF_RANGE = 2 _BSR_ERROR_INVALID_INDPTR = 3 _ERROR_INVALID_SEQ_LEN = 4 -_ERROR_INVALID_PAGE_INDPTR = 5 -_ERROR_INSUFFICIENT_PAGE_CAPACITY = 6 -_ERROR_INVALID_PHYSICAL_PAGE_ID = 7 +_ERROR_INVALID_PHYSICAL_PAGE_ID = 5 @cute.jit @@ -296,7 +294,7 @@ def __call__( class _InspectPagedKvMetadata: - """Validate live lengths and page rows with one warp per request.""" + """Validate live lengths and page-table rows with one warp per request.""" def __init__( self, @@ -314,16 +312,16 @@ def __init__( @cute.jit def __call__( self, - paged_kv_indptr: cute.Tensor, - paged_kv_indices: cute.Tensor, + block_tables: cute.Tensor, + block_table_row_stride: cutlass.Int64, seq_lens_kv: cute.Tensor, num_physical_kv_pages: cutlass.Int64, summary: cute.Tensor, stream: cuda_drv.CUstream, ) -> None: self.kernel( - paged_kv_indptr, - paged_kv_indices, + block_tables, + block_table_row_stride, seq_lens_kv, num_physical_kv_pages, summary, @@ -340,8 +338,8 @@ def __call__( @cute.kernel def kernel( self, - paged_kv_indptr: cute.Tensor, - paged_kv_indices: cute.Tensor, + block_tables: cute.Tensor, + block_table_row_stride: cutlass.Int64, seq_lens_kv: cute.Tensor, num_physical_kv_pages: cutlass.Int64, summary: cute.Tensor, @@ -353,9 +351,7 @@ def kernel( batch_idx = block_idx * _WARPS_PER_CTA + warp_idx request_is_valid = batch_idx < self.batch_size - request_begin = cutlass.Int32(0) - request_end = cutlass.Int32(0) - request_range_is_valid = cutlass.Int32(0) + live_pages = cutlass.Int32(0) error_code = cutlass.Int32(_BSR_ERROR_NONE) if lane_idx == 0 and request_is_valid: seq_len_kv = cutlass.Int32(seq_lens_kv[batch_idx]) @@ -363,38 +359,20 @@ def kernel( seq_len_kv >= cutlass.Int32(self.minimum_seq_len_kv) and seq_len_kv <= cutlass.Int32(self.max_seq_len_kv) ) - if not seq_len_is_valid: - error_code = cutlass.Int32(_ERROR_INVALID_SEQ_LEN) - - request_begin = cutlass.Int32(paged_kv_indptr[batch_idx]) - request_end = cutlass.Int32(paged_kv_indptr[batch_idx + 1]) - num_page_indices = cutlass.Int32(cute.size(paged_kv_indices)) - request_range_is_valid = cutlass.Int32( - paged_kv_indptr[cutlass.Int32(0)] == cutlass.Int32(0) - and request_begin >= cutlass.Int32(0) - and request_begin <= request_end - and request_end <= num_page_indices - ) - if request_range_is_valid == cutlass.Int32(0): - error_code = cutlass.Int32(_ERROR_INVALID_PAGE_INDPTR) - elif seq_len_is_valid: - required_pages = (seq_len_kv - cutlass.Int32(1)) // cutlass.Int32( + if seq_len_is_valid: + live_pages = (seq_len_kv - cutlass.Int32(1)) // cutlass.Int32( self.page_size ) + cutlass.Int32(1) - if request_end - request_begin < required_pages: - error_code = cutlass.Int32(_ERROR_INSUFFICIENT_PAGE_CAPACITY) + else: + error_code = cutlass.Int32(_ERROR_INVALID_SEQ_LEN) - request_begin = _warp_broadcast_i32(request_begin, 0) - request_end = _warp_broadcast_i32(request_end, 0) - request_range_is_valid = _warp_broadcast_i32(request_range_is_valid, 0) - if request_is_valid and request_range_is_valid != cutlass.Int32(0): + live_pages = _warp_broadcast_i32(live_pages, 0) + if request_is_valid: + row_begin = cutlass.Int64(batch_idx) * block_table_row_stride page_offset = cutlass.Int64(lane_idx) - request_page_count = cutlass.Int64(request_end) - cutlass.Int64( - request_begin - ) - while page_offset < request_page_count: - page_position = cutlass.Int64(request_begin) + page_offset - physical_page_id = cutlass.Int32(paged_kv_indices[page_position]) + while page_offset < cutlass.Int64(live_pages): + page_position = cutlass.Int64(row_begin + page_offset) + physical_page_id = cutlass.Int32(block_tables.iterator[page_position]) if ( physical_page_id < cutlass.Int32(0) or cutlass.Int64(physical_page_id) >= num_physical_kv_pages @@ -433,16 +411,16 @@ def __call__( self, block_indptr: cute.Tensor, block_indices: cute.Tensor, - paged_kv_indptr: cute.Tensor, - paged_kv_indices: cute.Tensor, + block_tables: cute.Tensor, + block_table_row_stride: cutlass.Int64, seq_lens_kv: cute.Tensor, num_physical_kv_pages: cutlass.Int64, summary: cute.Tensor, stream: cuda_drv.CUstream, ) -> None: self.inspect_requests( - paged_kv_indptr, - paged_kv_indices, + block_tables, + block_table_row_stride, seq_lens_kv, num_physical_kv_pages, summary, @@ -532,8 +510,9 @@ def compile_paged_block_sparse_metadata_inspection( """Compile one paged metadata entry that launches request then live-BSR.""" num_q_block_rows = (seq_len_q + q_block_size - 1) // q_block_size - logical_page_capacity = cute.sym_int() logical_nnz = cute.sym_int() + runtime_page_columns = cute.sym_int() + runtime_page_row_stride = cute.sym_int64(divisibility=1) stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True) inspect_requests = _InspectPagedKvMetadata( @@ -564,8 +543,13 @@ def compile_paged_block_sparse_metadata_inspection( alignment=4, ), _fake_compact(cutlass.Int32, (logical_nnz,), alignment=4), - _fake_compact(cutlass.Int32, (batch_size + 1,), alignment=4), - _fake_compact(cutlass.Int32, (logical_page_capacity,), alignment=4), + cute.runtime.make_fake_tensor( + cutlass.Int32, + (batch_size, runtime_page_columns), + stride=(runtime_page_row_stride, 1), + assumed_align=4, + ), + cutlass.Int64(1), _fake_compact(cutlass.Int32, (batch_size,), alignment=4), cutlass.Int64(1), _fake_compact(cutlass.Int64, (_SUMMARY_FIELDS,), alignment=8), diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/block_sparse_prepare.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/block_sparse_prepare.py index 6a61d0be86cd..08f2872809cb 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/block_sparse_prepare.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/block_sparse_prepare.py @@ -12,13 +12,15 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Prepare live canonical BSR rows for the PrimTS FMHA route consumer. +"""Prepare exact-first sparse routes for the PrimTS FMHA route consumer. -The kernel converts caller-owned semantic KV blocks into fixed-stride route -metadata on every run. Route origins are logical KV-token coordinates; -the paged specialization also resolves each origin to a physical page ID for -the attention load path. One warp handles one BSR row and iterates only that -row's live routes, while four warps share a CTA. +The BSR frontend is shared by continuous exact/proxy and paged exact storage. +It validates each canonical row once, emits the same logical exact records, +then either resolves paged locators or appends continuous proxy records. The +bitmask frontend shares the record geometry and emitters but remains limited +to continuous storage. Proxy suffixes contain one stable record per summary +group; a fully exact group remains present with zero score words. One warp owns +one sparse row and four warps share a CTA. ``row_route_offsets`` is a separate plan-owned immutable Int32 tensor. ``route_workspace`` contains only mutable row counts and route metadata @@ -27,21 +29,20 @@ semantic BSR-block limit, which remains distinct from packed-route capacity. """ -import math from dataclasses import dataclass import cutlass import cutlass.cute as cute from cuda.bindings import driver as cuda_drv from cutlass.cute.testing import assert_ as runtime_assert -from cutlass.experimental import primitives as prims from ..._block_sparse.prepared import ( _PREPARED_ROUTE_IS_FULL_FLAG, + _PREPARED_ROUTE_IS_PROXY_FLAG, _BlockSparseRouteLayout, ) -from .block_sparse_inspect import _validate_bsr_row_lane from .fmha_decode_resources.helpers_common import _warp_broadcast_i32 +from .block_sparse_inspect import _validate_bsr_row_lane _WARPS_PER_CTA = 4 @@ -50,22 +51,27 @@ @dataclass(frozen=True) -class _PreparedRouteConfig: - """Compile-time geometry shared by contiguous and paged route packing.""" +class _RouteConfig: + """Compile-time route geometry shared across sparse input/storage modes.""" num_kv_heads: int - num_q_block_rows: int + num_q_blocks: int num_kv_blocks: int + num_exact_words: int + num_proxy_groups: int num_rows: int seq_len_kv: int kv_block_size: int atom_size: int + atoms_per_block: int logical_origins_per_route: int token_words_per_route: int atom_valid_mask_word_offset: int route_flags_word_offset: int token_words_word_offset: int - has_token_bits: bool + stores_score_words: bool + apply_token_mask: bool + use_proxy_routes: bool route_metadata_stride_words: int route_metadata_base_word_offset: int @@ -78,18 +84,30 @@ def create( seq_len_kv: int, q_block_size: int, kv_block_size: int, - ) -> "_PreparedRouteConfig": - """Build shared prepare geometry without adding a storage-mode flag.""" + apply_token_mask: bool, + use_proxy_routes: bool, + ) -> "_RouteConfig": + """Build storage-independent route geometry and policy flags.""" - num_q_block_rows = (seq_len_q + q_block_size - 1) // q_block_size - return _PreparedRouteConfig( + stores_score_words = layout.token_words_word_offset is not None + if apply_token_mask and not stores_score_words: + raise ValueError("token masking requires prepared score words") + if use_proxy_routes and not stores_score_words: + raise ValueError("proxy routes require prepared score words") + num_q_blocks = (seq_len_q + q_block_size - 1) // q_block_size + num_kv_blocks = (seq_len_kv + kv_block_size - 1) // kv_block_size + return _RouteConfig( num_kv_heads=num_kv_heads, - num_q_block_rows=num_q_block_rows, - num_kv_blocks=(seq_len_kv + kv_block_size - 1) // kv_block_size, + num_q_blocks=num_q_blocks, + num_kv_blocks=num_kv_blocks, + num_exact_words=(num_kv_blocks + _WARP_SIZE - 1) // _WARP_SIZE, + num_proxy_groups=(num_kv_blocks + layout.kv_route_size - 1) + // layout.kv_route_size, num_rows=layout.num_rows, seq_len_kv=seq_len_kv, kv_block_size=kv_block_size, atom_size=layout.atom_size, + atoms_per_block=kv_block_size // layout.atom_size, logical_origins_per_route=layout.logical_origins_per_route, token_words_per_route=layout.token_words_per_route, atom_valid_mask_word_offset=layout.atom_valid_mask_word_offset, @@ -99,7 +117,9 @@ def create( if layout.token_words_word_offset is not None else 0 ), - has_token_bits=layout.has_token_bits, + stores_score_words=stores_score_words, + apply_token_mask=apply_token_mask, + use_proxy_routes=use_proxy_routes, route_metadata_stride_words=layout.route_metadata_stride_words, route_metadata_base_word_offset=layout.route_metadata_base_word_offset, ) @@ -115,35 +135,44 @@ def _positive_i32_ceil_div( @cute.jit -def _retained_atom_count( - block_indices: cute.Tensor, - row_begin: cutlass.Int32, - row_end: cutlass.Int32, - kv_block_size: cutlass.Constexpr[int], - atom_size: cutlass.Constexpr[int], - seq_len_kv: cutlass.Int32, +def _prepared_route_counts( + selected_block_count: cutlass.Int32, + cfg: cutlass.Constexpr[_RouteConfig], +) -> tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32]: + """Return exact atoms, exact routes, and total prepared routes for one row.""" + + exact_atom_count = selected_block_count * cutlass.Int32(cfg.atoms_per_block) + exact_route_count = ( + exact_atom_count + cutlass.Int32(cfg.logical_origins_per_route - 1) + ) // cutlass.Int32(cfg.logical_origins_per_route) + total_route_count = exact_route_count + if cutlass.const_expr(cfg.use_proxy_routes): + total_route_count += cutlass.Int32(cfg.num_proxy_groups) + return exact_atom_count, exact_route_count, total_route_count + + +@cute.jit +def _prepared_row_route_begin( + row_route_offsets: cute.Tensor, + linear_row_idx: cutlass.Int32, + lane_idx: cutlass.Int32, + row_is_valid: cutlass.Boolean, + total_route_count: cutlass.Int32, ) -> cutlass.Int32: - """Count selected atoms whose logical origin precedes ``seq_len_kv``.""" - - row_nnz = row_end - row_begin - retained_atoms = cutlass.Int32(0) - if row_nnz > cutlass.Int32(0): - atoms_per_block = kv_block_size // atom_size - retained_atoms = (row_nnz - cutlass.Int32(1)) * cutlass.Int32(atoms_per_block) - last_block_idx = cutlass.Int32(block_indices[row_end - cutlass.Int32(1)]) - last_block_origin = last_block_idx * cutlass.Int32(kv_block_size) - remaining_tokens = cutlass.Int32(seq_len_kv) - last_block_origin + """Load and validate one row's plan-owned prepared-route span.""" + + row_route_begin = cutlass.Int32(0) + if lane_idx == cutlass.Int32(0) and row_is_valid: + row_route_begin = cutlass.Int32(row_route_offsets[linear_row_idx]) + row_route_end = cutlass.Int32(row_route_offsets[linear_row_idx + 1]) + row_capacity = row_route_end - row_route_begin runtime_assert( - remaining_tokens > cutlass.Int32(0), - "block_indices row exceeds the live KV block range", + row_route_begin >= cutlass.Int32(0) + and row_capacity >= cutlass.Int32(0) + and total_route_count <= row_capacity, + "prepared routes exceed planned row capacity", ) - retained_last_atoms = (remaining_tokens - cutlass.Int32(1)) // cutlass.Int32( - atom_size - ) + cutlass.Int32(1) - if retained_last_atoms > cutlass.Int32(atoms_per_block): - retained_last_atoms = cutlass.Int32(atoms_per_block) - retained_atoms = retained_atoms + retained_last_atoms - return retained_atoms + return _warp_broadcast_i32(row_route_begin, 0) @cute.jit @@ -177,71 +206,72 @@ def _resolve_route_logical_atom_origin( @cute.jit -def _load_coarse_token_word( - block_indices: cute.Tensor, - kv_valid_bits: cute.Tensor, - row_begin: cutlass.Int32, - row_end: cutlass.Int32, - route_idx: cutlass.Int32, - logical_word_idx: cutlass.Int32, - batch_idx: cutlass.Int32, - kv_block_size: cutlass.Constexpr[int], - atom_size: cutlass.Constexpr[int], - logical_origins_per_route: cutlass.Constexpr[int], - seq_len_kv: cutlass.Int32, -) -> cutlass.Uint32: - """Load one logical K32 word from a coarse atom larger than K32.""" - - logical_word = cutlass.Uint32(0) - words_per_atom = atom_size // 32 - atom_in_route = logical_word_idx // cutlass.Int32(words_per_atom) - word_in_atom = logical_word_idx % cutlass.Int32(words_per_atom) - logical_origin, valid = _resolve_route_logical_atom_origin( - block_indices, - row_begin, - row_end, - route_idx, - atom_in_route, - kv_block_size, - atom_size, - logical_origins_per_route, - seq_len_kv, - ) - logical_word_origin = logical_origin + word_in_atom * cutlass.Int32(32) - if valid and logical_word_origin < cutlass.Int32(seq_len_kv): - valid_bits_word_idx = logical_word_origin >> cutlass.Int32(5) - logical_word = cutlass.Uint32(kv_valid_bits[batch_idx, valid_bits_word_idx]) - remaining_tokens = cutlass.Int32(seq_len_kv) - logical_word_origin - if remaining_tokens < cutlass.Int32(32): - logical_word = logical_word & ( - (cutlass.Uint32(1) << remaining_tokens) - cutlass.Uint32(1) - ) - return logical_word +def _low_bits_mask(valid_bits: cutlass.Int32) -> cutlass.Uint32: + """Return a Uint32 mask with its lowest clamped bit count set.""" + + mask = cutlass.Uint32(0) + if valid_bits >= cutlass.Int32(_WARP_SIZE): + mask = cutlass.Uint32(0xFFFFFFFF) + elif valid_bits > cutlass.Int32(0): + mask = (cutlass.Uint32(1) << valid_bits) - cutlass.Uint32(1) + return mask @cute.jit -def _load_atom_token_chunk( +def _load_exact_score_word( + route_workspace: cute.Tensor, kv_valid_bits: cute.Tensor, + route_metadata_word_index: cutlass.Int32, + logical_word_idx: cutlass.Int32, batch_idx: cutlass.Int32, - logical_origin: cutlass.Int32, - origin_is_valid: cutlass.Boolean, - atom_size: cutlass.Constexpr[int], seq_len_kv: cutlass.Int32, + cfg: cutlass.Constexpr[_RouteConfig], ) -> cutlass.Uint32: - """Load the <=K32 mask chunk owned by one resolved-origin lane.""" - - token_chunk = cutlass.Uint32(0) - if origin_is_valid: - valid_bits_word_idx = logical_origin >> cutlass.Int32(5) - source_word = cutlass.Uint32(kv_valid_bits[batch_idx, valid_bits_word_idx]) - token_chunk = source_word >> (logical_origin & cutlass.Int32(31)) - token_chunk = token_chunk & cutlass.Uint32((1 << atom_size) - 1) - remaining_tokens = cutlass.Int32(seq_len_kv) - logical_origin - if remaining_tokens < cutlass.Int32(atom_size): - token_chunk = token_chunk & ( - (cutlass.Uint32(1) << remaining_tokens) - cutlass.Uint32(1) - ) - return token_chunk + """Build one exact score word with optional caller-token masking.""" + + token_word = cutlass.Uint32(0) + if cutlass.const_expr(cfg.atom_size <= _WARP_SIZE): + atoms_per_word = _WARP_SIZE // cfg.atom_size + first_atom_idx = logical_word_idx * cutlass.Int32(atoms_per_word) + for atom_in_word in cutlass.range_constexpr(atoms_per_word): + atom_idx = first_atom_idx + cutlass.Int32(atom_in_word) + if atom_idx < cutlass.Int32(cfg.logical_origins_per_route): + origin = cutlass.Int32( + route_workspace[route_metadata_word_index + atom_idx] + ) + atom_word = cutlass.Uint32(0) + if origin >= cutlass.Int32(0): + if cutlass.const_expr(cfg.apply_token_mask): + source_word_idx = origin >> cutlass.Int32(5) + atom_word = cutlass.Uint32( + kv_valid_bits[batch_idx, source_word_idx] + ) + atom_word = atom_word >> (origin & cutlass.Int32(31)) + atom_word = atom_word & cutlass.Uint32((1 << cfg.atom_size) - 1) + atom_word = atom_word & _low_bits_mask(seq_len_kv - origin) + else: + atom_word = _low_bits_mask( + seq_len_kv - origin + ) & cutlass.Uint32((1 << cfg.atom_size) - 1) + token_word = token_word | ( + atom_word << cutlass.Int32(atom_in_word * cfg.atom_size) + ) + else: + words_per_atom = cfg.atom_size // _WARP_SIZE + atom_idx = logical_word_idx // cutlass.Int32(words_per_atom) + word_in_atom = logical_word_idx % cutlass.Int32(words_per_atom) + origin = cutlass.Int32(route_workspace[route_metadata_word_index + atom_idx]) + word_origin = origin + word_in_atom * cutlass.Int32(_WARP_SIZE) + if origin >= cutlass.Int32(0): + if cutlass.const_expr(cfg.apply_token_mask): + if word_origin < seq_len_kv: + source_word_idx = word_origin >> cutlass.Int32(5) + token_word = cutlass.Uint32( + kv_valid_bits[batch_idx, source_word_idx] + ) & _low_bits_mask(seq_len_kv - word_origin) + else: + token_word = _low_bits_mask(seq_len_kv - word_origin) + return token_word @cute.jit @@ -251,7 +281,7 @@ def _resolve_prepared_bsr_row( linear_row_idx: cutlass.Int32, lane_idx: cutlass.Int32, row_is_valid: cutlass.Boolean, - cfg: cutlass.Constexpr[_PreparedRouteConfig], + cfg: cutlass.Constexpr[_RouteConfig], ) -> tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32]: """Resolve one trusted canonical runtime BSR row.""" @@ -259,8 +289,8 @@ def _resolve_prepared_bsr_row( row_end = cutlass.Int32(0) batch_idx = cutlass.Int32(0) if lane_idx == cutlass.Int32(0) and row_is_valid: - q_block_row_idx = linear_row_idx % cfg.num_q_block_rows - linear_batch_head_idx = linear_row_idx // cfg.num_q_block_rows + q_block_row_idx = linear_row_idx % cfg.num_q_blocks + linear_batch_head_idx = linear_row_idx // cfg.num_q_blocks kv_head_idx = linear_batch_head_idx % cfg.num_kv_heads batch_idx = linear_batch_head_idx // cfg.num_kv_heads row_begin = cutlass.Int32(block_indptr[batch_idx, kv_head_idx, q_block_row_idx]) @@ -297,209 +327,256 @@ def _resolve_prepared_bsr_row( @cute.jit -def _publish_prepared_route_count( - block_indices: cute.Tensor, - row_route_offsets: cute.Tensor, +def _finalize_exact_route( route_workspace: cute.Tensor, - row_begin: cutlass.Int32, - row_end: cutlass.Int32, - linear_row_idx: cutlass.Int32, - lane_idx: cutlass.Int32, - row_is_valid: cutlass.Boolean, - max_blocks_per_row: cutlass.Int32, - seq_len_kv: cutlass.Int32, - cfg: cutlass.Constexpr[_PreparedRouteConfig], -) -> tuple[cutlass.Int32, cutlass.Int32]: - """Assert semantic capacity, publish the header, and return its live span.""" - - row_route_begin = cutlass.Int32(0) - required_route_count = cutlass.Int32(0) - if lane_idx == cutlass.Int32(0) and row_is_valid: - row_route_begin = cutlass.Int32(row_route_offsets[linear_row_idx]) - selected_block_count = row_end - row_begin - runtime_assert( - selected_block_count <= max_blocks_per_row, - "selected BSR blocks exceed planned semantic capacity", - ) - retained_atom_count = _retained_atom_count( - block_indices, - row_begin, - row_end, - cfg.kv_block_size, - cfg.atom_size, - seq_len_kv, - ) - required_route_count = ( - retained_atom_count + cutlass.Int32(cfg.logical_origins_per_route - 1) - ) // cutlass.Int32(cfg.logical_origins_per_route) - route_workspace[linear_row_idx] = required_route_count - row_route_begin = _warp_broadcast_i32(row_route_begin, 0) - required_route_count = _warp_broadcast_i32(required_route_count, 0) - return required_route_count, row_route_begin - - -@cute.jit -def _store_prepared_route_validity( - block_indices: cute.Tensor, kv_valid_bits: cute.Tensor, - route_workspace: cute.Tensor, - row_begin: cutlass.Int32, - row_end: cutlass.Int32, - route_idx: cutlass.Int32, + route_metadata_word_index: cutlass.Int32, batch_idx: cutlass.Int32, lane_idx: cutlass.Int32, - logical_origin: cutlass.Int32, - logical_origin_is_valid: cutlass.Boolean, - stored_atom_is_full: cutlass.Boolean, - route_metadata_word_index: cutlass.Int32, + atom_is_valid: cutlass.Boolean, seq_len_kv: cutlass.Int32, - cfg: cutlass.Constexpr[_PreparedRouteConfig], + cfg: cutlass.Constexpr[_RouteConfig], ) -> None: - """Store storage-independent atom, token, and route validity metadata.""" + """Finalize an exact record after its logical origins are stored.""" - stored_atom_valid_mask = cutlass.Int32( - cute.arch.vote_ballot_sync(logical_origin_is_valid) - ) - structural_route_is_full = cute.arch.vote_all_sync( - lane_idx >= cutlass.Int32(cfg.logical_origins_per_route) or stored_atom_is_full + atom_is_full = cutlass.Boolean(False) + if lane_idx < cutlass.Int32(cfg.logical_origins_per_route): + origin = cutlass.Int32(route_workspace[route_metadata_word_index + lane_idx]) + atom_is_full = cutlass.Boolean( + atom_is_valid and origin <= seq_len_kv - cutlass.Int32(cfg.atom_size) + ) + atom_valid_mask = cutlass.Int32(cute.arch.vote_ballot_sync(atom_is_valid)) + structural_full = cute.arch.vote_all_sync( + lane_idx >= cutlass.Int32(cfg.logical_origins_per_route) or atom_is_full ) - route_is_full = structural_route_is_full - if cutlass.const_expr(cfg.has_token_bits): - token_word = cutlass.Uint32(0) - if cutlass.const_expr(cfg.atom_size <= 32): - token_chunk = _load_atom_token_chunk( + + score_words_are_full = cutlass.Boolean(True) + if cutlass.const_expr(cfg.stores_score_words): + score_word = cutlass.Uint32(0) + if lane_idx < cutlass.Int32(cfg.token_words_per_route): + score_word = _load_exact_score_word( + route_workspace, kv_valid_bits, + route_metadata_word_index, + lane_idx, batch_idx, - logical_origin, - logical_origin_is_valid, - cfg.atom_size, seq_len_kv, + cfg, ) - atoms_per_word = 32 // cfg.atom_size - if lane_idx < cutlass.Int32(cfg.logical_origins_per_route): - atom_in_word = lane_idx % cutlass.Int32(atoms_per_word) - token_word = token_chunk << ( - atom_in_word * cutlass.Int32(cfg.atom_size) - ) - active_origin_lanes = (1 << cfg.logical_origins_per_route) - 1 - for shuffle_step in cutlass.range_constexpr( - int(math.log2(atoms_per_word)) - ): - peer_word = cutlass.Uint32( - prims.shfl_sync( - thread_mask=active_origin_lanes, - val=token_word, - offset=1 << shuffle_step, - mask_and_clamp=0x1F, - kind=prims.Shfl.BFLY, - ) - ) - token_word = token_word | peer_word - if atom_in_word == cutlass.Int32(0): - logical_word_idx = lane_idx // cutlass.Int32(atoms_per_word) - route_workspace[ - route_metadata_word_index - + cutlass.Int32(cfg.token_words_word_offset) - + logical_word_idx - ] = cutlass.Int32(token_word) - full_atom_mask = cutlass.Uint32((1 << cfg.atom_size) - 1) - token_route_is_full = cute.arch.vote_all_sync( - lane_idx >= cutlass.Int32(cfg.logical_origins_per_route) - or token_chunk == full_atom_mask - ) - else: - if lane_idx < cutlass.Int32(cfg.token_words_per_route): - token_word = _load_coarse_token_word( - block_indices, - kv_valid_bits, - row_begin, - row_end, - route_idx, - lane_idx, - batch_idx, - cfg.kv_block_size, - cfg.atom_size, - cfg.logical_origins_per_route, - seq_len_kv, - ) - route_workspace[ - route_metadata_word_index - + cutlass.Int32(cfg.token_words_word_offset) - + lane_idx - ] = cutlass.Int32(token_word) - token_route_is_full = cute.arch.vote_all_sync( - lane_idx >= cutlass.Int32(cfg.token_words_per_route) - or token_word == cutlass.Uint32(0xFFFFFFFF) - ) - route_is_full = cutlass.Boolean( - structural_route_is_full and token_route_is_full + route_workspace[ + route_metadata_word_index + + cutlass.Int32(cfg.token_words_word_offset) + + lane_idx + ] = cutlass.Int32(score_word) + score_words_are_full = cute.arch.vote_all_sync( + lane_idx >= cutlass.Int32(cfg.token_words_per_route) + or score_word == cutlass.Uint32(0xFFFFFFFF) ) - if lane_idx == cutlass.Int32(0): route_workspace[ route_metadata_word_index + cutlass.Int32(cfg.atom_valid_mask_word_offset) - ] = stored_atom_valid_mask + ] = atom_valid_mask route_workspace[ route_metadata_word_index + cutlass.Int32(cfg.route_flags_word_offset) ] = ( cutlass.Int32(_PREPARED_ROUTE_IS_FULL_FLAG) - if route_is_full + if structural_full and score_words_are_full else cutlass.Int32(0) ) -@cute.jit -def _paged_request_page_range_is_valid( - request_begin: cutlass.Int32, - request_end: cutlass.Int32, - num_indices: cutlass.Int32, - required_pages: cutlass.Int32, -) -> cutlass.Boolean: - """Validate one request's page-table range before any index load.""" - - return cutlass.Boolean( - request_begin >= cutlass.Int32(0) - and request_begin <= request_end - and request_end <= num_indices - and request_end - request_begin >= required_pages - ) - - @cute.jit def _resolve_paged_route_atom_page_id( - paged_kv_indices: cute.Tensor, - request_begin: cutlass.Int32, + block_tables: cute.Tensor, + batch_idx: cutlass.Int32, + block_table_row_stride: cutlass.Int64, logical_origin: cutlass.Int32, logical_origin_is_valid: cutlass.Boolean, lane_idx: cutlass.Int32, page_size: cutlass.Constexpr[int], num_physical_kv_pages: cutlass.Int64, ) -> cutlass.Int32: - """Resolve one trusted selected logical atom to its raw physical page ID.""" + """Resolve one trusted selected logical atom to its physical page ID.""" physical_page_id = cutlass.Int32(-1) page_id_is_valid = cutlass.Boolean(True) if logical_origin_is_valid: logical_page_idx = logical_origin // cutlass.Int32(page_size) - candidate_page_id = cutlass.Int32( - paged_kv_indices[request_begin + logical_page_idx] + physical_page_id = cutlass.Int32( + block_tables.iterator[ + cutlass.Int64(batch_idx) * block_table_row_stride + + cutlass.Int64(logical_page_idx) + ] ) - physical_page_id = candidate_page_id page_id_is_valid = cutlass.Boolean( - candidate_page_id >= cutlass.Int32(0) - and cutlass.Int64(candidate_page_id) < num_physical_kv_pages + physical_page_id >= cutlass.Int32(0) + and cutlass.Int64(physical_page_id) < num_physical_kv_pages ) page_ids_are_valid = cute.arch.vote_all_sync(page_id_is_valid) if lane_idx == cutlass.Int32(0): runtime_assert( page_ids_are_valid, - "paged_kv_indices contains an out-of-range physical page ID", + "block_tables contains an out-of-range physical page ID", ) return physical_page_id -class _PrepareBlockSparseRoutes: - """Prepare contiguous or paged sparse routes for one static geometry.""" +@cute.jit +def _exact_lane_rank( + exact_ballot: cutlass.Uint32, + lane_idx: cutlass.Int32, + exact_prefix: cutlass.Int32, +) -> cutlass.Int32: + """Return one exact lane's global semantic-block rank.""" + + lower_lane_mask = (cutlass.Uint32(1) << lane_idx) - cutlass.Uint32(1) + return exact_prefix + cutlass.Int32(cute.arch.popc(exact_ballot & lower_lane_mask)) + + +@cute.jit +def _emit_exact_block_atoms( + route_workspace: cute.Tensor, + row_route_begin: cutlass.Int32, + semantic_block_idx: cutlass.Int32, + exact_block_rank: cutlass.Int32, + cfg: cutlass.Constexpr[_RouteConfig], +) -> None: + """Expand one bitmask-selected block into fixed row-global atom slots.""" + + first_atom_rank = exact_block_rank * cutlass.Int32(cfg.atoms_per_block) + atom_in_block = cutlass.Int32(0) + while atom_in_block < cutlass.Int32(cfg.atoms_per_block): + atom_rank = first_atom_rank + atom_in_block + route_idx = atom_rank // cutlass.Int32(cfg.logical_origins_per_route) + atom_in_route = atom_rank % cutlass.Int32(cfg.logical_origins_per_route) + route_word_index = cutlass.Int32(cfg.route_metadata_base_word_offset) + ( + (row_route_begin + route_idx) + * cutlass.Int32(cfg.route_metadata_stride_words) + ) + logical_origin = semantic_block_idx * cutlass.Int32( + cfg.kv_block_size + ) + atom_in_block * cutlass.Int32(cfg.atom_size) + stored_origin = cutlass.Int32(-1) + if logical_origin < cutlass.Int32(cfg.seq_len_kv): + stored_origin = logical_origin + route_workspace[route_word_index + atom_in_route] = stored_origin + atom_in_block += cutlass.Int32(1) + + +@cute.jit +def _load_bitmask_word( + exact_block_bits: cute.Tensor, + batch_idx: cutlass.Int32, + kv_head_idx: cutlass.Int32, + q_block_idx: cutlass.Int32, + logical_word_idx: cutlass.Int32, + cfg: cutlass.Constexpr[_RouteConfig], + for_proxy: cutlass.Constexpr[bool], +) -> cutlass.Uint32: + """Load one in-range exact or proxy semantic-block word.""" + + valid_word = _low_bits_mask( + cutlass.Int32(cfg.num_kv_blocks) - logical_word_idx * cutlass.Int32(_WARP_SIZE) + ) + selected_word = cutlass.Uint32( + exact_block_bits[batch_idx, kv_head_idx, q_block_idx, logical_word_idx] + ) + if cutlass.const_expr(for_proxy): + selected_word = ~selected_word + return valid_word & selected_word + + +@cute.jit +def _load_bsr_proxy_word( + block_indices: cute.Tensor, + row_begin: cutlass.Int32, + row_end: cutlass.Int32, + logical_word_idx: cutlass.Int32, + cfg: cutlass.Constexpr[_RouteConfig], +) -> cutlass.Uint32: + """Build one proxy word from a canonical sorted-BSR interval.""" + + word_begin = logical_word_idx * cutlass.Int32(_WARP_SIZE) + valid_word = _low_bits_mask(cutlass.Int32(cfg.num_kv_blocks) - word_begin) + selected_word = cutlass.Uint32(0) + lower = row_begin + upper = row_end + while lower < upper: + middle = lower + (upper - lower) // cutlass.Int32(2) + if cutlass.Int32(block_indices[middle]) < word_begin: + lower = middle + cutlass.Int32(1) + else: + upper = middle + cursor = lower + word_end = word_begin + cutlass.Int32(_WARP_SIZE) + scanning = cutlass.Boolean(True) + while cursor < row_end and scanning: + block_idx = cutlass.Int32(block_indices[cursor]) + if block_idx < word_end: + selected_word = selected_word | ( + cutlass.Uint32(1) << (block_idx - word_begin) + ) + cursor += cutlass.Int32(1) + else: + scanning = cutlass.Boolean(False) + return valid_word & ~selected_word + + +@cute.jit +def _emit_proxy_route( + route_workspace: cute.Tensor, + row_route_begin: cutlass.Int32, + exact_route_count: cutlass.Int32, + group_idx: cutlass.Int32, + proxy_word: cutlass.Uint32, + lane_idx: cutlass.Int32, + cfg: cutlass.Constexpr[_RouteConfig], +) -> None: + """Emit one fixed summary-group proxy record, including an empty mask.""" + + route_metadata_word_index = cutlass.Int32(cfg.route_metadata_base_word_offset) + ( + row_route_begin + exact_route_count + group_idx + ) * cutlass.Int32(cfg.route_metadata_stride_words) + group_start = group_idx * cutlass.Int32(cfg.token_words_per_route * _WARP_SIZE) + group_size = cutlass.Int32(cfg.num_kv_blocks) - group_start + if group_size > cutlass.Int32(cfg.token_words_per_route * _WARP_SIZE): + group_size = cutlass.Int32(cfg.token_words_per_route * _WARP_SIZE) + origin_is_valid = cutlass.Boolean(False) + if lane_idx < cutlass.Int32(cfg.logical_origins_per_route): + summary_origin = group_start + lane_idx * cutlass.Int32(cfg.atom_size) + origin_is_valid = cutlass.Boolean(summary_origin < cfg.num_kv_blocks) + stored_origin = cutlass.Int32(-1) + if origin_is_valid: + stored_origin = summary_origin + route_workspace[route_metadata_word_index + lane_idx] = stored_origin + atom_valid_mask = cutlass.Int32(cute.arch.vote_ballot_sync(origin_is_valid)) + if lane_idx < cutlass.Int32(cfg.token_words_per_route): + route_workspace[ + route_metadata_word_index + + cutlass.Int32(cfg.token_words_word_offset) + + lane_idx + ] = cutlass.Int32(proxy_word) + score_full = cute.arch.vote_all_sync( + lane_idx >= cutlass.Int32(cfg.token_words_per_route) + or proxy_word == cutlass.Uint32(0xFFFFFFFF) + ) + if lane_idx == cutlass.Int32(0): + route_workspace[ + route_metadata_word_index + cutlass.Int32(cfg.atom_valid_mask_word_offset) + ] = atom_valid_mask + proxy_is_full = cutlass.Boolean( + group_size == cutlass.Int32(cfg.token_words_per_route * _WARP_SIZE) + and score_full + ) + route_workspace[ + route_metadata_word_index + cutlass.Int32(cfg.route_flags_word_offset) + ] = cutlass.Int32(_PREPARED_ROUTE_IS_PROXY_FLAG) | ( + cutlass.Int32(proxy_is_full) * cutlass.Int32(_PREPARED_ROUTE_IS_FULL_FLAG) + ) + + +class _PrepareRoutesBase: + """Own shared route geometry and compile-time storage/policy flags.""" def __init__( self, @@ -511,38 +588,58 @@ def __init__( q_block_size: int, kv_block_size: int, kv_route_size: int, - has_token_bits: bool, + use_proxy_routes: bool, + use_causal_mask: bool = False, + apply_token_mask: bool = False, + store_score_words: bool = False, page_size: int | None = None, - mask_type: str, ) -> None: - if mask_type not in ("dense", "causal"): - raise ValueError(f"unsupported mask_type: {mask_type}") - num_q_block_rows = (seq_len_q + q_block_size - 1) // q_block_size - num_rows = batch_size * num_kv_heads * num_q_block_rows + if not isinstance(use_proxy_routes, bool): + raise TypeError("use_proxy_routes must be a bool") + if not isinstance(apply_token_mask, bool): + raise TypeError("apply_token_mask must be a bool") + if not isinstance(store_score_words, bool): + raise TypeError("store_score_words must be a bool") + if not isinstance(use_causal_mask, bool): + raise TypeError("use_causal_mask must be a bool") + if use_proxy_routes and page_size is not None: + raise ValueError("paged KV does not support proxy routes") + + num_q_blocks = (seq_len_q + q_block_size - 1) // q_block_size + num_rows = batch_size * num_kv_heads * num_q_blocks + # Structural score words (sequence tail, invalid atoms) can be stored + # without a caller token mask; proxy routes and token masks require them. + stores_score_words = use_proxy_routes or apply_token_mask or store_score_words layout = _BlockSparseRouteLayout.create( kv_route_size=kv_route_size, kv_block_size=kv_block_size, page_size=page_size, - has_token_bits=has_token_bits, + has_token_bits=stores_score_words, route_metadata_capacity=0, num_rows=num_rows, ) - self.cfg = _PreparedRouteConfig.create( + self.route_layout = layout + self.cfg = _RouteConfig.create( layout=layout, num_kv_heads=num_kv_heads, seq_len_q=seq_len_q, seq_len_kv=seq_len_kv, q_block_size=q_block_size, kv_block_size=kv_block_size, + apply_token_mask=apply_token_mask, + use_proxy_routes=use_proxy_routes, ) - self.route_layout = layout self.page_size = page_size if page_size is not None else 1 - self.minimum_seq_len_kv = seq_len_q if mask_type == "causal" else 1 + self.minimum_seq_len_kv = seq_len_q if use_causal_mask else 1 self.physical_page_ids_word_offset = ( layout.physical_page_ids_word_offset if layout.is_paged else 0 ) self.route_metadata_base_word_offset = layout.route_metadata_base_word_offset + +class _PrepareBsrRoutes(_PrepareRoutesBase): + """Prepare continuous exact/proxy or paged exact routes from one BSR flow.""" + @cute.jit def __call__( self, @@ -550,24 +647,24 @@ def __call__( block_indices: cute.Tensor, kv_valid_bits: cute.Tensor, seq_lens_kv: cute.Tensor | None, - paged_kv_indptr: cute.Tensor | None, - paged_kv_indices: cute.Tensor | None, + block_tables: cute.Tensor | None, num_physical_kv_pages: cutlass.Int64, + block_table_row_stride: cutlass.Int64, row_route_offsets: cute.Tensor, route_workspace: cute.Tensor, max_blocks_per_row: cutlass.Int32, stream: cuda_drv.CUstream, ) -> None: - """Launch four independent row preparers per CTA.""" + """Launch four independent BSR row preparers per CTA.""" self.kernel( block_indptr, block_indices, kv_valid_bits, seq_lens_kv, - paged_kv_indptr, - paged_kv_indices, + block_tables, num_physical_kv_pages, + block_table_row_stride, row_route_offsets, route_workspace, max_blocks_per_row, @@ -588,14 +685,14 @@ def kernel( block_indices: cute.Tensor, kv_valid_bits: cute.Tensor, seq_lens_kv: cute.Tensor | None, - paged_kv_indptr: cute.Tensor | None, - paged_kv_indices: cute.Tensor | None, + block_tables: cute.Tensor | None, num_physical_kv_pages: cutlass.Int64, + block_table_row_stride: cutlass.Int64, row_route_offsets: cute.Tensor, route_workspace: cute.Tensor, max_blocks_per_row: cutlass.Int32, ) -> None: - """Pack logical routes and, when paged, translate physical locators.""" + """Assert trusted inputs, emit routes, resolve storage, then publish.""" thread_idx, _, _ = cute.arch.thread_idx() block_idx, _, _ = cute.arch.block_idx() @@ -613,8 +710,8 @@ def kernel( self.cfg, ) - request_begin = cutlass.Int32(0) live_seq_len_kv = cutlass.Int32(self.cfg.seq_len_kv) + selected_block_count = row_end - row_begin if cutlass.const_expr(self.route_layout.is_paged): raw_seq_len_kv = cutlass.Int32(self.cfg.seq_len_kv) if lane_idx == cutlass.Int32(0) and row_is_valid: @@ -624,110 +721,366 @@ def kernel( and raw_seq_len_kv <= cutlass.Int32(self.cfg.seq_len_kv), "seq_lens_kv is outside the planned live-length range", ) - raw_seq_len_kv = _warp_broadcast_i32(raw_seq_len_kv, 0) - live_seq_len_kv = raw_seq_len_kv + live_seq_len_kv = _warp_broadcast_i32(raw_seq_len_kv, 0) + + if ( + lane_idx == cutlass.Int32(0) + and row_is_valid + and selected_block_count > cutlass.Int32(0) + ): + last_block_idx = cutlass.Int32( + block_indices[row_end - cutlass.Int32(1)] + ) + runtime_assert( + last_block_idx * cutlass.Int32(self.cfg.kv_block_size) + < live_seq_len_kv, + "block_indices row exceeds the live KV block range", + ) if lane_idx == cutlass.Int32(0) and row_is_valid: required_pages = _positive_i32_ceil_div( live_seq_len_kv, self.page_size, ) - request_begin = cutlass.Int32(paged_kv_indptr[batch_idx]) - request_end = cutlass.Int32( - paged_kv_indptr[batch_idx + cutlass.Int32(1)] - ) - metadata_starts_at_zero = cutlass.Boolean( - paged_kv_indptr[cutlass.Int32(0)] == cutlass.Int32(0) - ) runtime_assert( - metadata_starts_at_zero - and _paged_request_page_range_is_valid( - request_begin, - request_end, - cutlass.Int32(cute.size(paged_kv_indices)), - required_pages, - ), - "paged_kv_indptr row lacks the required live page capacity", + required_pages <= cutlass.Int32(block_tables.shape[1]), + "block_tables row lacks the required live page capacity", ) - request_begin = _warp_broadcast_i32(request_begin, 0) - route_count, row_route_begin = _publish_prepared_route_count( - block_indices, + _, exact_route_count, total_route_count = _prepared_route_counts( + selected_block_count, + self.cfg, + ) + if lane_idx == cutlass.Int32(0) and row_is_valid: + runtime_assert( + selected_block_count <= max_blocks_per_row, + "selected BSR blocks exceed planned semantic capacity", + ) + row_route_begin = _prepared_row_route_begin( row_route_offsets, - route_workspace, - row_begin, - row_end, linear_row_idx, lane_idx, row_is_valid, - max_blocks_per_row, - live_seq_len_kv, - self.cfg, + total_route_count, ) - route_idx = cutlass.Int32(0) - while route_idx < route_count: - route_ordinal = row_route_begin + route_idx - route_metadata_word_index = cutlass.Int32( - self.cfg.route_metadata_base_word_offset - ) + route_ordinal * cutlass.Int32(self.cfg.route_metadata_stride_words) - logical_origin = cutlass.Int32(-1) - logical_origin_is_valid = cutlass.Boolean(False) - physical_page_id = cutlass.Int32(-1) - atom_is_full = cutlass.Boolean(False) - if lane_idx < cutlass.Int32(self.cfg.logical_origins_per_route): - ( - logical_origin, - logical_origin_is_valid, - ) = _resolve_route_logical_atom_origin( - block_indices, - row_begin, - row_end, - route_idx, + if row_is_valid: + route_idx = cutlass.Int32(0) + while route_idx < exact_route_count: + route_word_index = cutlass.Int32( + self.cfg.route_metadata_base_word_offset + ) + (row_route_begin + route_idx) * cutlass.Int32( + self.cfg.route_metadata_stride_words + ) + logical_origin = cutlass.Int32(-1) + logical_origin_is_valid = cutlass.Boolean(False) + physical_page_id = cutlass.Int32(-1) + if lane_idx < cutlass.Int32(self.cfg.logical_origins_per_route): + ( + logical_origin, + logical_origin_is_valid, + ) = _resolve_route_logical_atom_origin( + block_indices, + row_begin, + row_end, + route_idx, + lane_idx, + self.cfg.kv_block_size, + self.cfg.atom_size, + self.cfg.logical_origins_per_route, + live_seq_len_kv, + ) + if cutlass.const_expr(self.route_layout.is_paged): + physical_page_id = _resolve_paged_route_atom_page_id( + block_tables, + batch_idx, + block_table_row_stride, + logical_origin, + logical_origin_is_valid, + lane_idx, + self.page_size, + num_physical_kv_pages, + ) + if lane_idx < cutlass.Int32(self.cfg.logical_origins_per_route): + route_workspace[route_word_index + lane_idx] = logical_origin + if cutlass.const_expr(self.route_layout.is_paged): + route_workspace[ + route_word_index + + cutlass.Int32(self.physical_page_ids_word_offset) + + lane_idx + ] = physical_page_id + cute.arch.sync_warp() + + _finalize_exact_route( + route_workspace, + kv_valid_bits, + route_word_index, + batch_idx, lane_idx, - self.cfg.kv_block_size, - self.cfg.atom_size, - self.cfg.logical_origins_per_route, + logical_origin_is_valid, live_seq_len_kv, + self.cfg, ) - if cutlass.const_expr(self.route_layout.is_paged): - physical_page_id = _resolve_paged_route_atom_page_id( - paged_kv_indices, - request_begin, - logical_origin, - logical_origin_is_valid, - lane_idx, - self.page_size, - num_physical_kv_pages, + route_idx += cutlass.Int32(1) + + if cutlass.const_expr(self.cfg.use_proxy_routes): + group_idx = cutlass.Int32(0) + while group_idx < cutlass.Int32(self.cfg.num_proxy_groups): + proxy_word = cutlass.Uint32(0) + logical_word_idx = ( + group_idx * cutlass.Int32(self.cfg.token_words_per_route) + + lane_idx + ) + if lane_idx < cutlass.Int32(self.cfg.token_words_per_route): + if logical_word_idx < cutlass.Int32(self.cfg.num_exact_words): + proxy_word = _load_bsr_proxy_word( + block_indices, + row_begin, + row_end, + logical_word_idx, + self.cfg, + ) + _emit_proxy_route( + route_workspace, + row_route_begin, + exact_route_count, + group_idx, + proxy_word, + lane_idx, + self.cfg, + ) + group_idx += cutlass.Int32(1) + + if lane_idx == cutlass.Int32(0) and row_is_valid: + route_workspace[linear_row_idx] = total_route_count + + +class _PrepareBitmaskRoutes(_PrepareRoutesBase): + """Lower packed exact-block bits to continuous exact-first routes.""" + + def __init__( + self, + *, + batch_size: int, + num_kv_heads: int, + seq_len_q: int, + seq_len_kv: int, + q_block_size: int, + kv_block_size: int, + kv_route_size: int, + use_proxy_routes: bool, + use_causal_mask: bool = False, + apply_token_mask: bool = False, + store_score_words: bool = False, + ) -> None: + super().__init__( + batch_size=batch_size, + num_kv_heads=num_kv_heads, + seq_len_q=seq_len_q, + seq_len_kv=seq_len_kv, + q_block_size=q_block_size, + kv_block_size=kv_block_size, + kv_route_size=kv_route_size, + use_proxy_routes=use_proxy_routes, + use_causal_mask=use_causal_mask, + apply_token_mask=apply_token_mask, + store_score_words=store_score_words, + ) + + @cute.jit + def __call__( + self, + exact_block_bits: cute.Tensor, + kv_valid_bits: cute.Tensor, + row_route_offsets: cute.Tensor, + route_workspace: cute.Tensor, + max_blocks_per_row: cutlass.Int32, + stream: cuda_drv.CUstream, + ) -> None: + self.kernel( + exact_block_bits, + kv_valid_bits, + row_route_offsets, + route_workspace, + max_blocks_per_row, + ).launch( + grid=[ + (self.cfg.num_rows + _WARPS_PER_CTA - 1) // _WARPS_PER_CTA, + 1, + 1, + ], + block=[_THREADS_PER_CTA, 1, 1], + stream=stream, + ) + + @cute.kernel + def kernel( + self, + exact_block_bits: cute.Tensor, + kv_valid_bits: cute.Tensor, + row_route_offsets: cute.Tensor, + route_workspace: cute.Tensor, + max_blocks_per_row: cutlass.Int32, + ) -> None: + """Pack one bitmask row after proving its complete payload fits.""" + + thread_idx, _, _ = cute.arch.thread_idx() + block_idx, _, _ = cute.arch.block_idx() + warp_idx = thread_idx // _WARP_SIZE + lane_idx = thread_idx % _WARP_SIZE + linear_row_idx = block_idx * _WARPS_PER_CTA + warp_idx + row_is_valid = linear_row_idx < self.cfg.num_rows + q_block_idx = linear_row_idx % self.cfg.num_q_blocks + linear_batch_head_idx = linear_row_idx // self.cfg.num_q_blocks + kv_head_idx = linear_batch_head_idx % self.cfg.num_kv_heads + batch_idx = linear_batch_head_idx // self.cfg.num_kv_heads + + lane_exact_count = cutlass.Int32(0) + word_idx = lane_idx + while word_idx < cutlass.Int32(self.cfg.num_exact_words): + if row_is_valid: + exact_word = _load_bitmask_word( + exact_block_bits, + batch_idx, + kv_head_idx, + q_block_idx, + word_idx, + self.cfg, + for_proxy=False, + ) + lane_exact_count += cutlass.Int32(cute.arch.popc(exact_word)) + word_idx += cutlass.Int32(_WARP_SIZE) + exact_block_count = cutlass.Int32( + cute.arch.warp_redux_sync(lane_exact_count, "add") + ) + + exact_atom_count, exact_route_count, total_route_count = _prepared_route_counts( + exact_block_count, + self.cfg, + ) + if lane_idx == cutlass.Int32(0) and row_is_valid: + runtime_assert( + exact_block_count <= max_blocks_per_row, + "selected bitmask blocks exceed planned semantic capacity", + ) + row_route_begin = _prepared_row_route_begin( + row_route_offsets, + linear_row_idx, + lane_idx, + row_is_valid, + total_route_count, + ) + + if row_is_valid: + exact_prefix = cutlass.Int32(0) + word_idx = cutlass.Int32(0) + while word_idx < cutlass.Int32(self.cfg.num_exact_words): + exact_word_i32 = cutlass.Int32(0) + if lane_idx == cutlass.Int32(0): + exact_word_i32 = _load_bitmask_word( + exact_block_bits, + batch_idx, + kv_head_idx, + q_block_idx, + word_idx, + self.cfg, + for_proxy=False, + ).bitcast(cutlass.Int32) + exact_word = _warp_broadcast_i32(exact_word_i32, 0).bitcast( + cutlass.Uint32 ) - if logical_origin_is_valid: - atom_is_full = cutlass.Boolean( - logical_origin - <= live_seq_len_kv - cutlass.Int32(self.cfg.atom_size) + is_exact = cutlass.Boolean( + (exact_word & (cutlass.Uint32(1) << lane_idx)) != cutlass.Uint32(0) ) - if lane_idx < cutlass.Int32(self.cfg.logical_origins_per_route): - route_workspace[route_metadata_word_index + lane_idx] = logical_origin - if cutlass.const_expr(self.route_layout.is_paged): - route_workspace[ - route_metadata_word_index - + cutlass.Int32(self.physical_page_ids_word_offset) - + lane_idx - ] = physical_page_id + exact_ballot = cute.arch.vote_ballot_sync(is_exact).bitcast( + cutlass.Uint32 + ) + exact_rank = _exact_lane_rank(exact_ballot, lane_idx, exact_prefix) + if is_exact: + _emit_exact_block_atoms( + route_workspace, + row_route_begin, + word_idx * cutlass.Int32(_WARP_SIZE) + lane_idx, + exact_rank, + self.cfg, + ) + exact_prefix += cutlass.Int32(cute.arch.popc(exact_ballot)) + word_idx += cutlass.Int32(1) - _store_prepared_route_validity( - block_indices, - kv_valid_bits, - route_workspace, - row_begin, - row_end, - route_idx, - batch_idx, - lane_idx, - logical_origin, - logical_origin_is_valid, - atom_is_full, - route_metadata_word_index, - live_seq_len_kv, - self.cfg, + final_route_atom_count = exact_atom_count % cutlass.Int32( + self.cfg.logical_origins_per_route ) - route_idx = route_idx + cutlass.Int32(1) + if final_route_atom_count != cutlass.Int32(0): + if lane_idx >= final_route_atom_count and lane_idx < cutlass.Int32( + self.cfg.logical_origins_per_route + ): + final_route_word_index = cutlass.Int32( + self.cfg.route_metadata_base_word_offset + ) + (row_route_begin + exact_route_count - cutlass.Int32(1)) * ( + cutlass.Int32(self.cfg.route_metadata_stride_words) + ) + route_workspace[final_route_word_index + lane_idx] = cutlass.Int32( + -1 + ) + cute.arch.sync_warp() + + route_idx = cutlass.Int32(0) + while route_idx < exact_route_count: + route_word_index = cutlass.Int32( + self.cfg.route_metadata_base_word_offset + ) + (row_route_begin + route_idx) * cutlass.Int32( + self.cfg.route_metadata_stride_words + ) + atom_is_valid = cutlass.Boolean(False) + if lane_idx < cutlass.Int32(self.cfg.logical_origins_per_route): + atom_is_valid = cutlass.Boolean( + cutlass.Int32(route_workspace[route_word_index + lane_idx]) + >= cutlass.Int32(0) + ) + _finalize_exact_route( + route_workspace, + kv_valid_bits, + route_word_index, + batch_idx, + lane_idx, + atom_is_valid, + cutlass.Int32(self.cfg.seq_len_kv), + self.cfg, + ) + route_idx += cutlass.Int32(1) + + if cutlass.const_expr(self.cfg.use_proxy_routes): + group_idx = cutlass.Int32(0) + while group_idx < cutlass.Int32(self.cfg.num_proxy_groups): + proxy_word = cutlass.Uint32(0) + logical_word_idx = ( + group_idx * cutlass.Int32(self.cfg.token_words_per_route) + + lane_idx + ) + if lane_idx < cutlass.Int32(self.cfg.token_words_per_route): + if logical_word_idx < cutlass.Int32(self.cfg.num_exact_words): + proxy_word = _load_bitmask_word( + exact_block_bits, + batch_idx, + kv_head_idx, + q_block_idx, + logical_word_idx, + self.cfg, + for_proxy=True, + ) + _emit_proxy_route( + route_workspace, + row_route_begin, + exact_route_count, + group_idx, + proxy_word, + lane_idx, + self.cfg, + ) + group_idx += cutlass.Int32(1) + + if lane_idx == cutlass.Int32(0) and row_is_valid: + route_workspace[linear_row_idx] = total_route_count + + +__all__ = ["_PrepareBitmaskRoutes", "_PrepareBsrRoutes"] diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_config.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_config.py index c67e4c83bf10..0626410d1f10 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_config.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_config.py @@ -102,6 +102,16 @@ (Float16, Float16, Float16, 256, 128, 1, 1), } +# Per-thread register budgets for the Q64/KV256 warp groups once the launch +# bound enables ``setmaxnreg``. The MMA, load, and scheduler warps keep 56, so +# the two softmax groups and the correction group share the remainder: +# 8 * softmax + 4 * correction = 65536 / 32 - 4 * 56. An even split measured +# fastest on B200 for both the static grid and the persistent scheduler; the +# correction group needs the extra room for its persistent bookkeeping and +# the KV256 tail merge, while the rolled softmax fragment loop needs less. +KV_TILE_256_SOFTMAX_TASK_REGISTERS = 152 +KV_TILE_256_CORRECTION_TASK_REGISTERS = 152 + _KV_TILE_256_PHYSICAL_DEFAULTS: Mapping[str, ConfigValue] = { "tmem_s_cols": 128, "tmem_stats_cols": 32, @@ -505,6 +515,10 @@ class FmhaDecodeConfig: # restricted to 8/16/32 or positive multiples of 64 and are assembled into # a profile-selected fixed KV128 or KV256 route. use_block_sparse: bool = False + # Interpret prepared records as a typed proxy/exact stream. Proxy records + # source semantic-block summaries while exact records retain the K/V + # atom path. Source selection is orthogonal to the physical Q/KV profile. + use_block_sparse_proxy_routes: bool = False q_block_size: int = 0 kv_block_size: int = 0 # Optional batch-wide physical-token validity metadata shared by every head @@ -644,13 +658,13 @@ def uses_task_register_reallocation(self) -> bool: def softmax_task_num_registers(self) -> int | None: if not self.uses_task_register_reallocation: return None - return 176 if self.tile_size_kv == 256 else 184 + return KV_TILE_256_SOFTMAX_TASK_REGISTERS if self.tile_size_kv == 256 else 184 @property def correction_task_num_registers(self) -> int | None: if not self.uses_task_register_reallocation: return None - return 104 if self.tile_size_kv == 256 else 88 + return KV_TILE_256_CORRECTION_TASK_REGISTERS if self.tile_size_kv == 256 else 88 @property def mma_load_task_num_registers(self) -> int | None: @@ -1066,12 +1080,13 @@ def num_s_regs_per_thread(self) -> int: def softmax_score_fragment_regs(self) -> int: """Return the maximum score fragment kept live in registers. - KV256 owns 128 score values per lane but streams them as four native - 32-register LDTM atoms. Other profiles retain their complete score - fragment, so this property is intentionally distinct from - ``num_s_regs_per_thread`` (the total logical ownership). + Streamed profiles own 128 score values per lane but process them as + four native 32-register LDTM atoms. Other profiles retain their + complete score fragment, so this property is intentionally distinct + from ``num_s_regs_per_thread`` (the total logical ownership). The + selection of streamed profiles lives in ``streams_tmem_p_fragments``. """ - if self.tile_size_kv == 256: + if self.streams_tmem_p_fragments: return 32 return self.num_s_regs_per_thread @@ -1080,6 +1095,30 @@ def num_softmax_score_fragments(self) -> int: """Return score fragments used to cover one logical KV tile.""" return self.num_s_regs_per_thread // self.softmax_score_fragment_regs + @property + def block_sparse_kv_atom_size(self) -> int: + """Return the K token span of one block-sparse route origin.""" + assert self.use_block_sparse + return _block_sparse_kv_atom_size(self.kv_block_size) + + @property + def softmax_fragments_per_route_atom(self) -> int: + """Return the streamed score fragments that share one route origin. + + Route origins are staged per K64 atom, so a 128-token KV block spans + two origins; the fragment-to-origin mapping follows the atom. + """ + return self.block_sparse_kv_atom_size // self.softmax_score_fragment_regs + + @property + def uses_ws_2x2_datapath(self) -> bool: + """Whether QK and PV issue the WS 2x2 instruction over two lane halves. + + KV256 exposes two spatial KV128 partials per logical Q row; every other + Keeps profile issues the plain CTA-local instruction. + """ + return self.tile_size_kv == 256 + @property def num_packed_p_regs(self) -> int: """Return packed P registers stored by each softmax producer lane.""" @@ -1246,6 +1285,10 @@ def validate_paged_kv_staging_config(self) -> None: def validate_block_sparse_profile(self, *, heads_q_per_kv: int) -> None: """Validate the qualified host profile for block-sparse.""" + if self.use_block_sparse_proxy_routes and not self.use_block_sparse: + raise ValueError("proxy routes require block-sparse attention") + if self.use_block_sparse_proxy_routes and self.mask_type != DENSE: + raise ValueError("block-sparse proxy routes require mask_type='dense'") if not self.use_block_sparse: if self.use_parallel_sparse_kv_loads: raise ValueError( @@ -1283,6 +1326,13 @@ def validate_block_sparse_profile(self, *, heads_q_per_kv: int) -> None: "block-sparse tile_size_kv=256 requires the Q64 16-bit Keeps " "profile with coarse KV blocks and one load task" ) + if self.use_keeps_mma_ab and not self.streams_tmem_p_fragments: + # The block-sparse Keeps softmax and P passes exist only in their + # streamed K32-fragment form. + raise ValueError( + "block-sparse KeepsMmaAb requires a streamed TMEM-P profile " + "(Q64/KV256 or 16-bit Q128/KV128)" + ) if self.tile_size_q != selected_q_tile: raise ValueError( "block-sparse tile_size_q must match its grouped-Q geometry" @@ -1332,6 +1382,42 @@ def compile_signature(self) -> tuple[tuple[str, object], ...]: for config_field in fields(self) ) + @property + def uses_prepared_score_keep_words(self) -> bool: + """Whether prepared routes carry BMM1 score-column validity words. + + Dense block-sparse Keeps plans prepare them even without a caller + token mask: the streamed max pass trusts the words directly, which is + cheaper than deriving each fragment's visible range in the softmax + warps. The plan sizes its route storage and the prepare kernels store + the words from this same property, via the resolved launch spec. + """ + + return ( + self.use_kv_valid_bits + or self.use_block_sparse_proxy_routes + or ( + self.use_block_sparse + and self.use_keeps_mma_ab + and self.mask_type == DENSE + ) + ) + + @property + def trusts_prepared_score_words(self) -> bool: + """Whether prepared words fully describe dense score-column validity. + + Dense prepared routes have already combined structural tail validity + with any caller-provided exact-token bits. Their K32 words therefore + apply to exact and proxy sources alike. + """ + + return ( + self.use_block_sparse + and self.uses_prepared_score_keep_words + and self.mask_type == DENSE + ) + @property def uses_q_desc_ref(self) -> bool: """Whether QK derives Q's descriptor from shared resource state.""" @@ -1435,10 +1521,10 @@ def has_static_dense_full_kv_tiles(self) -> bool: @property def uses_ordered_softmax_barrier(self) -> bool: """Whether this profile selects the ordered P0/P1 softmax barrier.""" - if self.tile_size_kv == 256: - # KV256 uses independent four-stage P-fragment pipelines. Ordering - # the two softmax groups would serialize fragment production and - # defeat the intended P/PV overlap. + if self.streams_tmem_p_fragments: + # Streamed profiles use independent per-fragment P pipelines. + # Ordering the two softmax groups would serialize fragment + # production and defeat the intended P/PV overlap. return False if self.ordered_softmax_barrier_mode == 2: return True @@ -1554,10 +1640,11 @@ def uses_staged_one_inst_tmem_p(self) -> bool: def uses_two_inst_tmem_p(self) -> bool: """Whether a two-instance Keeps profile uses the TMEM-P overlay. - Q128/KV128 and sparse Q64/KV128 publish a complete packed-P row per - pipeline token. Q64/KV256 uses the same S-to-P aliasing contract but - streams four independently ready K32 fragments. Dense Q64/KV128 keeps - the base kernel's faster SMEM-P cadence. + FP8 Q128/KV128 and dense 16-bit Q128/KV128 publish a complete + packed-P row per pipeline token. Q64/KV256 and block-sparse 16-bit + Q128/KV128 use the same S-to-P aliasing contract but stream four + independently ready K32 fragments (see ``streams_tmem_p_fragments``). + Q64/KV128 keeps the base kernel's faster SMEM-P cadence. """ # Two-instance Keeps keeps stats outside S, so both static and persistent # work tiles can overlay P on the consumed S instance. The split K/V @@ -1568,11 +1655,6 @@ def uses_two_inst_tmem_p(self) -> bool: and ( (self.tile_size_q == 128 and self.tile_size_kv == 128) or (self.tile_size_q == 64 and self.tile_size_kv == 256) - or ( - self.use_block_sparse - and self.tile_size_q == 64 - and self.tile_size_kv == 128 - ) ) and self.head_dim_per_stage_kv == 0 and self.num_insts_kv == 2 @@ -1582,42 +1664,51 @@ def uses_two_inst_tmem_p(self) -> bool: @property def streams_tmem_p_fragments(self) -> bool: - """Whether P is published as independently ready TMEM fragments.""" - return self.uses_two_inst_tmem_p and self.num_softmax_score_fragments > 1 - - @property - def matches_kv256_task_topology(self) -> bool: - """Whether task roles match KV256's validated 16-warp layout.""" - return all( - getattr(self, field) == expected - for field, expected in _KV_TILE_256_TASK_TOPOLOGY_DEFAULTS.items() + """Whether P is published as independently ready TMEM fragments. + + Streamed profiles produce their K32 fragments from one rolled runtime + loop: the max pass writes masked scores back to TMEM, so the P pass + reloads each fragment without mask logic and the exponentiation body + exists once in the instruction stream. Each published fragment lets + the MMA warp start its PV k-slice before the row is complete, at the + cost of one barrier round per fragment. + + Streaming is limited to the 16-bit two-instance profiles whose route + loop waits on the K/V loads, where the earlier PV start hides load + latency: Q64/KV256 and block-sparse Q128/KV128. Dense Q128/KV128 + keeps the complete row because its route loop is not load-bound, so + the per-fragment barriers are not compensated. FP8 Q128 keeps the + complete row because its P publication packs four values per column + into one store. + """ + return ( + self.uses_two_inst_tmem_p + and not self.use_fp8_qkv + and (self.tile_size_kv == 256 or self.use_block_sparse) ) @property - def uses_rotating_kv256_exchange(self) -> bool: - """Whether this profile selects KV-ring scratch for correction. + def defers_softmax_anchor_updates(self) -> bool: + """Whether small row-max increases keep the previous exponent anchor. - Persistent direct output can overlap the next work tile's first two - K loads with correction by placing its exchange in the third, drained - KV stage. Split-KV and attention sinks retain the fixed exchange because - their tail storage and lifetime differ from direct output. + Keeps correction skips the in-place O rescale whenever the anchor is + unchanged, so keeping the prior anchor within + ``SOFTMAX_RESCALE_THRESHOLD_LOG2`` trades a bounded 16-bit P range + (2**8) for fewer TMEM rescales. The profiles listed here are the ones + where that trade was measured to pay: KV256 tiles and block-sparse + routes, whose row maximum moves often but rarely by much. """ - selects_persistent_kv256 = ( - self.streams_tmem_p_fragments - and self.tile_size_q == 64 - and self.tile_size_kv == 256 - and self.use_persistent_scheduler + return self.use_keeps_mma_ab and ( + self.tile_size_kv == 256 or self.use_block_sparse ) - if not selects_persistent_kv256: - return False - has_rotating_kv_ring = ( - self.num_head_dim_stages_kv == 1 - and self.kv_stages == KV_TILE_256_SHARED_FIFO_STAGES - and self.load_num_warps == 1 + @property + def matches_kv256_task_topology(self) -> bool: + """Whether task roles match KV256's validated 16-warp layout.""" + return all( + getattr(self, field) == expected + for field, expected in _KV_TILE_256_TASK_TOPOLOGY_DEFAULTS.items() ) - has_direct_output_lifetime = not (self.use_split_kv or self.use_attention_sinks) - return has_rotating_kv_ring and has_direct_output_lifetime @property def keeps_separates_tmem_s_and_stats(self) -> bool: @@ -2196,17 +2287,6 @@ def _require_python_int(field_name: str) -> int: f"{pipeline_smem_bytes} bytes, limit is " f"{pipeline_smem_budget_bytes} bytes" ) - if ( - cfg.use_persistent_scheduler - and not cfg.use_split_kv - and not cfg.use_attention_sinks - and cfg.kv_stages != KV_TILE_256_SHARED_FIFO_STAGES - ): - raise ValueError( - "persistent KV256 requires kv_stages=" - f"{KV_TILE_256_SHARED_FIFO_STAGES} for the rotating shared-KV " - f"exchange, got {cfg.kv_stages}" - ) if not cfg.supports_grouped_keeps: raise ValueError( "KV256 currently supports only the qualified Q64 FP16/BF16/D128 " diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_constants.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_constants.py index 68686febfbbb..8bf94760cf31 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_constants.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_constants.py @@ -45,8 +45,12 @@ # Keep the old maximum as the exponent reference while a new maximum is at # most eight log2 units larger. This avoids an output-correction round without # letting an intermediate probability exceed 2**8; the softmax identity is -# unchanged apart from normal finite-precision rounding. -KV_TILE_256_RESCALE_THRESHOLD_LOG2 = 8.0 +# unchanged apart from normal finite-precision rounding. As in the +# FlashInfer/TRT-LLM policy, this assumes normal model logits rather than +# adversarial values outside the qualified probability bound. Streamed KV256 +# and block-sparse Keeps profiles apply it; see +# ``FmhaDecodeConfig.defers_softmax_anchor_updates``. +SOFTMAX_RESCALE_THRESHOLD_LOG2 = 8.0 # A launch bound makes ptxas honor warpgroup ``setmaxnreg`` allocations, but # the resulting register hand-off has a fixed cost. Paired B200 measurements diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_kernel.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_kernel.py index 929696c30d3e..d2540ec61920 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_kernel.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_kernel.py @@ -56,8 +56,8 @@ from cutlass.experimental.task_scheduling.task_manager import TaskManager from ..._block_sparse.common import ( - _block_sparse_kv_atom_size, - _prepared_kv_routes_are_block_aligned, + _block_sparse_contiguous_kv_copy_geometry, + _block_sparse_proxy_summary_geometry, ) from ..._block_sparse.prepared import _BlockSparseRouteLayout from ..tensor_map import ( @@ -95,7 +95,7 @@ from .fmha_decode_tasks import ( PackedDecodeWorkQueue, ScheduleTokenThrottleResource, - SmemKvReuseCreditResource, + _prefetch_prepared_sparse_row, create_block_sparse_load_tasks_per_inst, create_correction_task, create_correction_task_one_inst_qkv, @@ -113,7 +113,6 @@ create_softmax0_task, create_softmax1_task, ) - from .reduction import ( # noqa: F401 decode_gen_separate_reduction_kernel, fmha_decode_separate_reduction_launch, @@ -322,6 +321,10 @@ def _build_decode_gen_schedule( tma_desc_v: cutlass.Pointer | None = None, tma_desc_k_atom: cutlass.Pointer | None = None, tma_desc_v_atom: cutlass.Pointer | None = None, + tma_desc_k_summary: cutlass.Pointer | None = None, + tma_desc_v_summary: cutlass.Pointer | None = None, + tma_desc_k_summary_atom: cutlass.Pointer | None = None, + tma_desc_v_summary_atom: cutlass.Pointer | None = None, page_idx_kv: cute.Pointer | None = None, h_k_idx: Int32 | None = None, b_idx: Int32 | None = None, @@ -341,6 +344,8 @@ def _build_decode_gen_schedule( sparse_row_route_offsets: cute.Pointer | None = None, sparse_row_route_counts: cute.Pointer | None = None, sparse_route_metadata: cute.Pointer | None = None, + sparse_row_route_begin: Int32 | None = None, + sparse_route_count: Int32 | None = None, ) -> tuple[ list[Task], dict[MemoryResource, list[MemoryResource]], @@ -408,6 +413,15 @@ def _build_decode_gen_schedule( "tma_desc_k_atom": tma_desc_k_atom, "tma_desc_v_atom": tma_desc_v_atom, } + if cfg.use_block_sparse_proxy_routes: + segment_tensormaps.update( + { + "tma_desc_k_summary": tma_desc_k_summary, + "tma_desc_v_summary": tma_desc_v_summary, + "tma_desc_k_summary_atom": tma_desc_k_summary_atom, + "tma_desc_v_summary_atom": tma_desc_v_summary_atom, + } + ) for name, descriptor in segment_tensormaps.items(): if descriptor is None: raise ValueError( @@ -741,7 +755,6 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: # ------------------------------------------------------------------ work_queue = None schedule_token_throttle = None - smem_kv_reuse_credit = None # CLC remains the single persistent policy for every supported topology. # The stock static WorkQueue advances and decodes coordinates separately # in every task, which regresses multi-wave decode workloads. CLC computes @@ -799,17 +812,6 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: ), name="schedule_token_throttle", ) - if cfg.uses_rotating_kv256_exchange: - smem_kv_reuse_credit = SmemKvReuseCreditResource( - cfg=cfg, - pipeline_config=PipelineConfig.create_async_async_pipeline_cfg( - num_stages=1, - producer_group=load_grp, - consumer_group=correction_grp, - cta_layout_vmnk=cta_layout, - ), - name="smem_kv_reuse_credit", - ) smem_q = SmemQResource( pipeline_config=smem_q_cfg, cfg=cfg, @@ -873,10 +875,12 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: sparse_softmax_metadata0 = None sparse_softmax_metadata1 = None if cfg.use_block_sparse: + # This selects the prepared-record storage ABI. Causal consumers still + # intersect these column-validity words with each Q row's causal mask. prepared_route_layout = _BlockSparseRouteLayout.create( kv_route_size=cfg.tile_size_kv, kv_block_size=cfg.kv_block_size, - has_token_bits=cfg.use_kv_valid_bits, + has_token_bits=cfg.uses_prepared_score_keep_words, route_metadata_capacity=0, num_rows=1, page_size=cfg.num_tokens_per_page if cfg.use_paged_kv else None, @@ -928,6 +932,10 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: tma_desc_v=tma_desc_v, tma_desc_k_atom=tma_desc_k_atom, tma_desc_v_atom=tma_desc_v_atom, + tma_desc_k_summary=tma_desc_k_summary, + tma_desc_v_summary=tma_desc_v_summary, + tma_desc_k_summary_atom=tma_desc_k_summary_atom, + tma_desc_v_summary_atom=tma_desc_v_summary_atom, sparse_kv_metadata=sparse_kv_metadata0, page_offsets_kv=smem_page_offsets, seqlens_kv=kv_seqlens, @@ -947,6 +955,10 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: tma_desc_v=tma_desc_v, tma_desc_k_atom=tma_desc_k_atom, tma_desc_v_atom=tma_desc_v_atom, + tma_desc_k_summary=tma_desc_k_summary, + tma_desc_v_summary=tma_desc_v_summary, + tma_desc_k_summary_atom=tma_desc_k_summary_atom, + tma_desc_v_summary_atom=tma_desc_v_summary_atom, sparse_kv_metadata=sparse_kv_metadata1, page_offsets_kv=smem_page_offsets, seqlens_kv=kv_seqlens, @@ -966,6 +978,10 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: tma_desc_v=tma_desc_v, tma_desc_k_atom=tma_desc_k_atom, tma_desc_v_atom=tma_desc_v_atom, + tma_desc_k_summary=tma_desc_k_summary, + tma_desc_v_summary=tma_desc_v_summary, + tma_desc_k_summary_atom=tma_desc_k_summary_atom, + tma_desc_v_summary_atom=tma_desc_v_summary_atom, sparse_kv_metadata=sparse_kv_metadata0, page_offsets_kv=smem_page_offsets_v or smem_page_offsets, seqlens_kv=kv_seqlens, @@ -985,6 +1001,10 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: tma_desc_v=tma_desc_v, tma_desc_k_atom=tma_desc_k_atom, tma_desc_v_atom=tma_desc_v_atom, + tma_desc_k_summary=tma_desc_k_summary, + tma_desc_v_summary=tma_desc_v_summary, + tma_desc_k_summary_atom=tma_desc_k_summary_atom, + tma_desc_v_summary_atom=tma_desc_v_summary_atom, sparse_kv_metadata=sparse_kv_metadata1, page_offsets_kv=smem_page_offsets_v or smem_page_offsets, seqlens_kv=kv_seqlens, @@ -1005,6 +1025,10 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: tma_desc_v=tma_desc_v, tma_desc_k_atom=tma_desc_k_atom, tma_desc_v_atom=tma_desc_v_atom, + tma_desc_k_summary=tma_desc_k_summary, + tma_desc_v_summary=tma_desc_v_summary, + tma_desc_k_summary_atom=tma_desc_k_summary_atom, + tma_desc_v_summary_atom=tma_desc_v_summary_atom, sparse_kv_metadata0=sparse_kv_metadata0, sparse_kv_metadata1=sparse_kv_metadata1, page_offsets_kv=smem_page_offsets, @@ -1213,6 +1237,8 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: "seq_len_q": seq_len_q, "sparse_row_route_offsets": sparse_row_route_offsets, "sparse_row_route_counts": sparse_row_route_counts, + "sparse_row_route_begin": sparse_row_route_begin, + "sparse_route_count": sparse_route_count, "num_heads_kv": num_heads_kv, } if use_one_inst_qkv: @@ -1283,7 +1309,6 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: smem_kv, work_queue, schedule_token_throttle, - smem_kv_reuse_credit, cfg, domain=load_domain, domain_bias=0, @@ -1431,7 +1456,6 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: tmem_corr0, tmem_corr1, work_queue, - smem_kv_reuse_credit, cfg, domain=corr_domain, tmem_stats_done0=tmem_stats_done0, @@ -1607,10 +1631,6 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: ) if schedule_token_throttle is not None: resource_dependency_graph[schedule_token_throttle] = [work_queue] - if smem_kv_reuse_credit is not None: - # A self-edge models the one-slot ownership token: Load produces it - # for the current tile and Correction consumes it before the next Load. - resource_dependency_graph[smem_kv_reuse_credit] = [smem_kv_reuse_credit] dma_consumer_release_labels: dict[ tuple[MemoryResource, MemoryResource], set[str] ] = {} @@ -1662,8 +1682,6 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: smem_allocator.add_resource(work_queue) if schedule_token_throttle is not None: smem_allocator.add_resource(schedule_token_throttle) - if smem_kv_reuse_credit is not None: - smem_allocator.add_resource(smem_kv_reuse_credit) smem_allocator.add_resource(smem_q) if smem_page_offsets is not None: smem_allocator.add_resource(smem_page_offsets) @@ -1701,21 +1719,10 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: smem_allocator.add_resource(tmem_corr0) if not use_one_inst_qkv: smem_allocator.add_resource(tmem_corr1) - if cfg.tile_size_kv == 256: - # KV256 direct-output correction rotates one compact 35,840-byte - # payload through the shared 192-KiB K/V ring. Split-KV retains the - # fixed full exchange. Neither path increases the CTA SMEM footprint. - smem_allocator.add_alias_group( - [ - smem_kv.get_smem_requirements(), - tmem_corr1.get_smem_requirements(), - ] - ) smem_allocator.add_tmem_ptr( SmemAllocation("fmha_tmem_ptr_i32", dtype=cutlass.Int32, alignment=4) ) smem_allocator.compute_layout() - tmem_allocator = TmemAllocator() if cfg.use_keeps_mma_ab: if use_one_inst_qkv: @@ -1774,13 +1781,15 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: p0_alloc = smem_p0.get_tmem_requirements()[0] p1_alloc = smem_p1.get_tmem_requirements()[0] o_alloc = tmem_o.get_tmem_requirements()[0] - if cfg.tile_size_kv == 256: - # KV256 keeps O in the low 256 columns and overlays packed P on - # each S region from its first column. Softmax streams K32 - # fragments in order, so every 16-column P store only overwrites - # scores that have already been consumed. Starting P after the - # nominal stats columns would instead clobber the next unread S - # fragment; KV256 keeps its softmax stats in SMEM. + if cfg.streams_tmem_p_fragments: + # Streamed profiles keep O in the low 256 columns and overlay + # packed P on each S region from its first column. Softmax streams + # K32 fragments in order, so every 16-column P store only + # overwrites scores that have already been consumed. Starting P + # after the nominal stats columns would instead clobber the next + # unread S fragment; streamed profiles keep their softmax stats in + # SMEM. + assert cfg.keeps_stats_via_smem o_alloc.offset = 0 s0_alloc.offset = 2 * cfg.tmem_o_stage_cols s1_alloc.offset = s0_alloc.offset + cfg.tmem_s_cols @@ -1835,12 +1844,8 @@ def _make_page_offsets_cfg(num_stages: int | None = None) -> PipelineConfig: eager_init_resources = ( [tmem_corr0] if use_one_inst_qkv else [tmem_corr0, tmem_corr1] ) - if smem_kv_reuse_credit is not None: - # Initialize the persistent ring cursor under the same CTA-wide fence - # and barrier used by other manually managed SMEM control state. - eager_init_resources.append(smem_kv_reuse_credit) - if cfg.tile_size_kv == 256: - # KV256's TMEM P operands use one-way per-fragment ready barriers. + if cfg.streams_tmem_p_fragments: + # Streamed TMEM P operands use one-way per-fragment ready barriers. # Initialize them beside correction's manually managed SMEM state. eager_init_resources.extend([smem_p0, smem_p1]) @@ -1863,10 +1868,10 @@ def _has_unmodeled_tmem_p_alias_protocol(cfg: FmhaDecodeConfig) -> bool: """Whether exhaustive TS checking would report a known false P/S race. The staged D256 path selects one of two physical P/S stages at runtime. - Static KV256 instead orders streamed P fragments with private mbarriers and + Static streamed profiles instead order P fragments with private mbarriers and reuses the matching TmemO-full barrier as the next-QK overwrite credit. Those intra-work protocols are below TaskManager's resource transitions, - so its allocation-level checker cannot prove them. Persistent KV256 has + so its allocation-level checker cannot prove them. Persistent streaming has enough task-level ordering for the checker and remains covered. """ return cfg.uses_staged_one_inst_tmem_p or ( @@ -1982,6 +1987,10 @@ def _run_decode_gen_active( g_sparse_row_route_offsets: cute.Pointer | None = None, g_sparse_row_route_counts: cute.Pointer | None = None, g_sparse_route_metadata: cute.Pointer | None = None, + tma_desc_k_summary: cutlass.GridConstant[cuda.TensorMap] | None = None, + tma_desc_v_summary: cutlass.GridConstant[cuda.TensorMap] | None = None, + tma_desc_k_summary_atom: cutlass.GridConstant[cuda.TensorMap] | None = None, + tma_desc_v_summary_atom: cutlass.GridConstant[cuda.TensorMap] | None = None, ) -> None: """Run the complete decode body for one runtime-valid Q tile. @@ -2021,28 +2030,43 @@ def _run_decode_gen_active( else Int32(cfg.static_seq_len_kv) ) use_clc_dynamic_scheduler = cfg.use_persistent_scheduler + tma_desc_k_summary_ptr = None + tma_desc_v_summary_ptr = None + tma_desc_k_summary_atom_ptr = None + tma_desc_v_summary_atom_ptr = None + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + assert tma_desc_k_summary is not None + assert tma_desc_v_summary is not None + assert tma_desc_k_summary_atom is not None + assert tma_desc_v_summary_atom is not None + tma_desc_k_summary_ptr = tma_desc_k_summary.get_ptr() + tma_desc_v_summary_ptr = tma_desc_v_summary.get_ptr() + tma_desc_k_summary_atom_ptr = tma_desc_k_summary_atom.get_ptr() + tma_desc_v_summary_atom_ptr = tma_desc_v_summary_atom.get_ptr() # Prefetch TMA + uses_atom_desc = False + if cutlass.const_expr(cfg.use_block_sparse): + _, _, uses_atom_desc = _block_sparse_contiguous_kv_copy_geometry( + kv_block_size=cfg.kv_block_size, + kv_route_size=cfg.tile_size_kv, + ) init_warp = 1 if warp_idx == init_warp: prims.prefetch_tensormap(tma_desc_q.get_ptr()) prims.prefetch_tensormap(tma_desc_k.get_ptr()) prims.prefetch_tensormap(tma_desc_v.get_ptr()) - if cutlass.const_expr( - cfg.use_block_sparse - and _block_sparse_kv_atom_size(cfg.kv_block_size) == 64 - and ( - cfg.tile_size_kv == 256 - or not _prepared_kv_routes_are_block_aligned( - cfg.kv_block_size, - cfg.tile_size_kv, - ) - ) - ): - # KV256 always issues semantic KV64 atoms. KV128 needs this second - # descriptor only for non-aligned coarse routes. + if cutlass.const_expr(cfg.use_block_sparse and uses_atom_desc): + # KV256 and non-aligned coarse KV128 may select the exact atom maps. prims.prefetch_tensormap(tma_desc_k_atom.get_ptr()) prims.prefetch_tensormap(tma_desc_v_atom.get_ptr()) + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + if cutlass.const_expr(cfg.tile_size_kv != 256): + prims.prefetch_tensormap(tma_desc_k_summary_ptr) + prims.prefetch_tensormap(tma_desc_v_summary_ptr) + if cutlass.const_expr(uses_atom_desc): + prims.prefetch_tensormap(tma_desc_k_summary_atom_ptr) + prims.prefetch_tensormap(tma_desc_v_summary_atom_ptr) init_warp += 1 clc_response_ptr = None @@ -2055,6 +2079,27 @@ def _run_decode_gen_active( if cutlass.const_expr(cfg.max_seq_len_q > 1): q_output_rows = g_h_r * Int32(cfg.max_seq_len_q) + # Static block-sparse tiles read their prepared row header here, before + # TMEM allocation and barrier setup, so that global-memory round trip is + # hidden instead of stalling every task at its first schedule step. + sparse_row_route_begin = None + sparse_route_count = None + if cutlass.const_expr( + cfg.use_block_sparse + and not cfg.use_persistent_scheduler + and g_sparse_row_route_offsets is not None + and g_sparse_row_route_counts is not None + ): + sparse_row_route_begin, sparse_route_count = _prefetch_prepared_sparse_row( + cfg, + g_sparse_row_route_offsets, + g_sparse_row_route_counts, + q_group_idx, + h_k_idx, + b_idx, + g_h_k, + ) + ( task_list, dep_graph, @@ -2083,6 +2128,10 @@ def _run_decode_gen_active( tma_desc_v=tma_desc_v.get_ptr(), tma_desc_k_atom=tma_desc_k_atom.get_ptr(), tma_desc_v_atom=tma_desc_v_atom.get_ptr(), + tma_desc_k_summary=tma_desc_k_summary_ptr, + tma_desc_v_summary=tma_desc_v_summary_ptr, + tma_desc_k_summary_atom=tma_desc_k_summary_atom_ptr, + tma_desc_v_summary_atom=tma_desc_v_summary_atom_ptr, page_idx_kv=g_page_idx_kv, h_k_idx=h_k_idx, b_idx=b_idx, @@ -2102,6 +2151,8 @@ def _run_decode_gen_active( sparse_row_route_offsets=g_sparse_row_route_offsets, sparse_row_route_counts=g_sparse_row_route_counts, sparse_route_metadata=g_sparse_route_metadata, + sparse_row_route_begin=sparse_row_route_begin, + sparse_route_count=sparse_route_count, ) smem_allocator.allocate() @@ -2264,6 +2315,10 @@ def _run_decode_gen_runtime_prefix( g_sparse_row_route_offsets: cute.Pointer | None, g_sparse_row_route_counts: cute.Pointer | None, g_sparse_route_metadata: cute.Pointer | None, + tma_desc_k_summary: cutlass.GridConstant[cuda.TensorMap] | None = None, + tma_desc_v_summary: cutlass.GridConstant[cuda.TensorMap] | None = None, + tma_desc_k_summary_atom: cutlass.GridConstant[cuda.TensorMap] | None = None, + tma_desc_v_summary_atom: cutlass.GridConstant[cuda.TensorMap] | None = None, ) -> None: """Run the general runtime split-prefix producer or retire its suffix.""" @@ -2328,6 +2383,10 @@ def _run_decode_gen_runtime_prefix( g_sparse_row_route_offsets, g_sparse_row_route_counts, g_sparse_route_metadata, + tma_desc_k_summary=tma_desc_k_summary, + tma_desc_v_summary=tma_desc_v_summary, + tma_desc_k_summary_atom=tma_desc_k_summary_atom, + tma_desc_v_summary_atom=tma_desc_v_summary_atom, ) else: _run_decode_gen_inactive_cluster_rank() @@ -2371,6 +2430,10 @@ def _run_decode_gen_runtime_prefix( g_sparse_row_route_offsets, g_sparse_row_route_counts, g_sparse_route_metadata, + tma_desc_k_summary=tma_desc_k_summary, + tma_desc_v_summary=tma_desc_v_summary, + tma_desc_k_summary_atom=tma_desc_k_summary_atom, + tma_desc_v_summary_atom=tma_desc_v_summary_atom, ) else: _signal_padded_pdl_producer(cfg) @@ -2409,6 +2472,10 @@ def decode_gen_kernel( g_sparse_row_route_counts: cute.Pointer | None = None, g_sparse_route_metadata: cute.Pointer | None = None, static_full_split_prefix: cutlass.Constexpr[bool] = False, + tma_desc_k_summary: cutlass.GridConstant[cuda.TensorMap] | None = None, + tma_desc_v_summary: cutlass.GridConstant[cuda.TensorMap] | None = None, + tma_desc_k_summary_atom: cutlass.GridConstant[cuda.TensorMap] | None = None, + tma_desc_v_summary_atom: cutlass.GridConstant[cuda.TensorMap] | None = None, ) -> None: """Dispatch one static Q/split tile and drain padded launch slots safely.""" q_group_cta_idx, h_k_idx, b_idx = cute.arch.block_idx() @@ -2470,6 +2537,10 @@ def decode_gen_kernel( g_sparse_row_route_offsets, g_sparse_row_route_counts, g_sparse_route_metadata, + tma_desc_k_summary=tma_desc_k_summary, + tma_desc_v_summary=tma_desc_v_summary, + tma_desc_k_summary_atom=tma_desc_k_summary_atom, + tma_desc_v_summary_atom=tma_desc_v_summary_atom, ) else: _run_decode_gen_runtime_prefix( @@ -2510,6 +2581,10 @@ def decode_gen_kernel( g_sparse_row_route_offsets, g_sparse_row_route_counts, g_sparse_route_metadata, + tma_desc_k_summary=tma_desc_k_summary, + tma_desc_v_summary=tma_desc_v_summary, + tma_desc_k_summary_atom=tma_desc_k_summary_atom, + tma_desc_v_summary_atom=tma_desc_v_summary_atom, ) else: # Packed-Q grids use a batch-wide maximum envelope. These Q CTAs own no @@ -2754,7 +2829,6 @@ def fmha_decode_launch( tma_desc_q, tma_desc_k, tma_desc_v, - # Dense/paged profiles never inspect the 64-token descriptor slots. tma_desc_k, tma_desc_v, o_iter, @@ -2799,6 +2873,8 @@ def fmha_block_sparse_launch( q_iter: cute.Pointer, k_iter: cute.Pointer, v_iter: cute.Pointer, + k_summary_iter: cute.Pointer, + v_summary_iter: cute.Pointer, o_iter: cute.Pointer, row_route_offsets_iter: cute.Pointer, row_route_counts_iter: cute.Pointer, @@ -2813,14 +2889,18 @@ def fmha_block_sparse_launch( k_page_stride: Int64 = 0, v_page_stride: Int64 = 0, ) -> None: - """Launch attention over contiguous or paged prepared KV routes. + """Launch attention over exact and typed exact/proxy prepared KV routes. A preceding prepare kernel has already resolved each BSR row into compact logical atom origins, storage locators, validity flags, and optional token - words. Both layouts execute the same ``decode_gen_kernel`` schedule. + words. Exact routes address K/V; proxy routes address summary K/V. Both + layouts execute the same ``decode_gen_kernel`` schedule and + physical copy policy. Exact builds constexpr-elide summary TensorMaps. """ if cutlass.const_expr(not cfg.use_block_sparse): - raise ValueError("fmha_block_sparse_launch requires cfg.use_block_sparse=True") + raise ValueError("fmha_block_sparse_launch requires block-sparse config") + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes and cfg.use_paged_kv): + raise ValueError("block-sparse proxy routes require contiguous K/V") log2_e = math.log2(math.e) b, h_q, h_k, s_k, d = problem_shape @@ -2855,7 +2935,18 @@ def fmha_block_sparse_launch( swizzle=tma_swizzle, ) - kv_atom_size = _block_sparse_kv_atom_size(cfg.kv_block_size) + ( + primary_kv_box_size, + kv_atom_size, + uses_atom_desc, + ) = _block_sparse_contiguous_kv_copy_geometry( + kv_block_size=cfg.kv_block_size, + kv_route_size=cfg.tile_size_kv, + ) + k_desc_summary_primary = None + v_desc_summary_primary = None + k_desc_summary_atom = None + v_desc_summary_atom = None if cutlass.const_expr(cfg.use_paged_kv): # Paged HND storage is addressed as (D, token-in-page, Hkv, page). # Prepared routes already contain each atom's physical page ID, so no @@ -2899,9 +2990,10 @@ def fmha_block_sparse_launch( k_desc_primary = k_desc_atom v_desc_primary = v_desc_atom else: - # Contiguous sparse coordinates retain the logical (D, S, H, B) - # order and the established primary/atom descriptor split. - primary_kv_box_size = 2 * kv_atom_size if kv_atom_size == 64 else kv_atom_size + # Exact and summary tensors form one logical segmented KV coordinate + # space. Each physical source owns the same primary/atom descriptor + # pair; the prepared route kind selects the pair, while the loader + # retains the existing KV128/fine/KV256 copy policy. kv_dims = (d, s_k, h_k, b) k_desc_primary = create_tensor_map_tiled( global_address=k_iter.toint(), @@ -2921,16 +3013,7 @@ def fmha_block_sparse_launch( ) k_desc_atom = k_desc_primary v_desc_atom = v_desc_primary - if cutlass.const_expr( - kv_atom_size == 64 - and ( - cfg.tile_size_kv == 256 - or not _prepared_kv_routes_are_block_aligned( - cfg.kv_block_size, - cfg.tile_size_kv, - ) - ) - ): + if cutlass.const_expr(uses_atom_desc): # KV256 always stages four semantic KV64 atoms. KV128 needs this # map only when a route may join unrelated BSR entries. k_desc_atom = create_tensor_map_tiled( @@ -2950,6 +3033,55 @@ def fmha_block_sparse_launch( swizzle=tma_swizzle, ) + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + num_kv_blocks, _ = _block_sparse_proxy_summary_geometry( + seq_len_kv, + cfg.kv_block_size, + ) + _, summary_kv_strides = _block_sparse_bshd_tma_strides( + q_seq=q_seq, + h_q=h_q, + h_k=h_k, + s_k=num_kv_blocks, + d=d, + ) + summary_dims = (d, num_kv_blocks, h_k, b) + k_desc_summary_primary = create_tensor_map_tiled( + global_address=k_summary_iter.toint(), + dtype=cfg.kv_dtype, + global_dims=summary_dims, + global_strides=summary_kv_strides, + box_dims=(tma_box0, primary_kv_box_size, 1, 1), + swizzle=tma_swizzle, + ) + v_desc_summary_primary = create_tensor_map_tiled( + global_address=v_summary_iter.toint(), + dtype=cfg.kv_dtype, + global_dims=summary_dims, + global_strides=summary_kv_strides, + box_dims=(tma_box0, primary_kv_box_size, 1, 1), + swizzle=tma_swizzle, + ) + k_desc_summary_atom = k_desc_summary_primary + v_desc_summary_atom = v_desc_summary_primary + if cutlass.const_expr(uses_atom_desc): + k_desc_summary_atom = create_tensor_map_tiled( + global_address=k_summary_iter.toint(), + dtype=cfg.kv_dtype, + global_dims=summary_dims, + global_strides=summary_kv_strides, + box_dims=(tma_box0, kv_atom_size, 1, 1), + swizzle=tma_swizzle, + ) + v_desc_summary_atom = create_tensor_map_tiled( + global_address=v_summary_iter.toint(), + dtype=cfg.kv_dtype, + global_dims=summary_dims, + global_strides=summary_kv_strides, + box_dims=(tma_box0, kv_atom_size, 1, 1), + swizzle=tma_swizzle, + ) + q_groups = Int32( (cfg.max_seq_len_q + cfg.q_tokens_per_cta - 1) // cfg.q_tokens_per_cta ) @@ -3008,6 +3140,10 @@ def fmha_block_sparse_launch( row_route_counts_iter, route_metadata_iter, False, # static_full_split_prefix + tma_desc_k_summary=k_desc_summary_primary, + tma_desc_v_summary=v_desc_summary_primary, + tma_desc_k_summary_atom=k_desc_summary_atom, + tma_desc_v_summary_atom=v_desc_summary_atom, ).launch( grid=grid, block=[cfg.threads_per_cta, 1, 1], diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/helpers_common.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/helpers_common.py index e3190f09097c..77412cddca7b 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/helpers_common.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/helpers_common.py @@ -25,6 +25,8 @@ import cutlass import cutlass.cute as cute from cutlass import BFloat16, Float16, Float32, Int32, Int64, Uint32 +from cutlass._mlir.dialects import llvm +from cutlass.cutlass_dsl import dsl_user_op from cutlass.experimental import primitives as prims from cutlass.experimental.task_scheduling.resources import ( @@ -59,6 +61,22 @@ ) ResourceVars = dict[str, ResourceVarValue] + +@dsl_user_op +def _assume_nonnegative_i32(value: Int32, *, loc=None, ip=None) -> Int32: + """Express a caller-guaranteed nonnegative Int32 contract to codegen.""" + + condition = cutlass.Boolean(value >= Int32(0)) + llvm.intr_assume( + condition.ir_value(loc=loc, ip=ip), + [], + [], + loc=loc, + ip=ip, + ) + return value + + # Offsets into DecodeGenTask.make_task_cache(). Keeping these symbolic makes # resource code explicit about which task-local lane or address value it needs. _TASK_CACHE_TMEM_BASE_OFFSET = 0 @@ -109,6 +127,33 @@ def _warp_broadcast_i32(value: Int32, source_lane: Constexpr[int]) -> Int32: ) +@cute.jit +def _swaps_routed_coordinate( + cfg: Constexpr[FmhaDecodeConfig], + lane_k_offset: Int32, + origin0: Int32, + origin1: Int32, + origin2: Int32, + origin3: Int32, + *, + token_group_idx: Constexpr[int], +) -> tuple[Int32, Int32]: + """Map one SWAP register group to its staged atom and logical coordinate.""" + + atom_size = min(cfg.kv_block_size, 32) + groups_per_atom = atom_size // 8 + origin_idx = token_group_idx // groups_per_atom + atom_origin = origin0 + if cutlass.const_expr(origin_idx == 1): + atom_origin = origin1 + elif cutlass.const_expr(origin_idx == 2): + atom_origin = origin2 + elif cutlass.const_expr(origin_idx == 3): + atom_origin = origin3 + token_offset = (token_group_idx % groups_per_atom) * 8 + return atom_origin, atom_origin + Int32(token_offset) + lane_k_offset + + def _mma_kind_for_qkv(cfg: FmhaDecodeConfig) -> prims.Tcgen05MMAKind: """Select the tcgen05 MMA opcode family used for Q/K/V operands.""" return prims.Tcgen05MMAKind.F8F6F4 if cfg.use_fp8_qkv else prims.Tcgen05MMAKind.F16 diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/helpers_softmax.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/helpers_softmax.py index 4585213cd80b..2dbe3c7e7ff8 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/helpers_softmax.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/helpers_softmax.py @@ -71,6 +71,70 @@ def _pack_float4_to_fp8_e4m3_inline( ) +@cute.jit +def _combine_int_frac_ex2( + x_rounded: Float32, frac_ex2: Float32, *, loc=None, ip=None +) -> Float32: + """Scale ``frac_ex2`` by ``2**floor(x)`` through the FP32 exponent bits. + + ``x_rounded`` still carries the magic rounding constant, so its low + mantissa bits hold ``floor(x)``; shifting them into the exponent field and + adding the bits of the polynomial result multiplies by the integer power. + """ + return cute.arch.inline_ptx( + "{\n" + " .reg .b32 xi;\n" + " .reg .b32 fi;\n" + " .reg .b32 xe;\n" + " .reg .b32 oi;\n" + " mov.b32 xi, {$r0};\n" + " mov.b32 fi, {$r1};\n" + " shl.b32 xe, xi, 23;\n" + " add.s32 oi, xe, fi;\n" + " mov.b32 {$w0}, oi;\n" + "}", + write_only_types=[Float32], + read_only_args=[x_rounded, frac_ex2], + loc=loc, + ip=ip, + ) + + +@cute.jit +def _ex2_emulation_packed_f32x2(x: Float32, y: Float32) -> tuple[Float32, Float32]: + """Evaluate ``2**x`` and ``2**y`` on the FMA pipe instead of MUFU. + + Inputs are non-positive scaled scores minus the row maximum. The integer + part is split off with a magic-constant rounding add, the fraction in + [0, 1) goes through a degree-3 minimax polynomial, and the two parts are + recombined through the exponent bits. The relative error stays below the + BF16 rounding of the P operand, matching the emulation used by the dense + Blackwell FMHA kernels. + """ + fp32_round_int = float(2**23 + 2**22) + xy_clamped = (cute.arch.fmax(x, -127.0), cute.arch.fmax(y, -127.0)) + xy_rounded = cute.arch.add_packed_f32x2( + xy_clamped, (fp32_round_int, fp32_round_int), rnd="rm" + ) + xy_rounded_back = cute.arch.sub_packed_f32x2( + xy_rounded, (fp32_round_int, fp32_round_int) + ) + xy_frac = cute.arch.sub_packed_f32x2(xy_clamped, xy_rounded_back) + coeff = ( + 1.0, + 0.695146143436431884765625, + 0.227564394474029541015625, + 0.077119089663028717041015625, + ) + out = (coeff[3], coeff[3]) + for degree in cutlass.range_constexpr(2, -1, -1): + out = cute.arch.fma_packed_f32x2(out, xy_frac, (coeff[degree], coeff[degree])) + return ( + _combine_int_frac_ex2(xy_rounded[0], out[0]), + _combine_int_frac_ex2(xy_rounded[1], out[1]), + ) + + @cute.jit def _compute_fp8_p_regs_and_local_sums( scale_softmax_log2: Float32, diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_block_sparse_metadata.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_block_sparse_metadata.py index 3254de0e79e7..f79f3a1e0566 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_block_sparse_metadata.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_block_sparse_metadata.py @@ -43,6 +43,7 @@ from ...._block_sparse.prepared import ( _PREPARED_ROUTE_IS_FULL_FLAG, + _PREPARED_ROUTE_IS_PROXY_FLAG, _BlockSparseRouteLayout, ) from ...placeholder_helpers import _placeholder_smem_array @@ -56,7 +57,6 @@ DecodeGenResourceBase, ResourceVars, _decode_gen_task_cache, - _keeps_col_base, _sparse_task_cache_route_begin, _sparse_task_cache_route_count, _warp_broadcast_i32, @@ -65,13 +65,20 @@ # Keeps staging uses the low four bits for structural KV64 validity. Bit 4 # carries the conservative prepared summary that token masking can be skipped; -# structural, tail, and causal masking remain independent. +# structural, tail, and causal masking remain independent. The streamed Keeps +# max pass derives its keep words from the token words directly, so the bit is +# currently staged for the consumer but not read. _SOFTMAX_TOKEN_MASK_IS_FULL_FLAG = 1 << 4 - -# B8 SWAP origins are eight-token aligned, so bit 0 is free while the route is -# in Softmax's private staging payload. Reusing it avoids adding a word to every -# pipeline stage merely to forward prepare's route-full summary. -_SWAPS_PACKED_ROUTE_FULL_CLEAR_MASK = ~_PREPARED_ROUTE_IS_FULL_FLAG +# Keeps reserves bit 5 for the prepared route kind. The low four structural +# validity bits and bit 4 keep their existing meaning. +_SOFTMAX_ROUTE_IS_PROXY_FLAG = 1 << 5 + +# SWAP origins are at least eight-token aligned, so their low two bits are free +# while the route is in Softmax's private staging payload. Reusing them avoids +# adding a word to every pipeline stage for prepared FULL/PROXY route flags. +_SWAPS_PACKED_ROUTE_FLAGS_CLEAR_MASK = ~( + _PREPARED_ROUTE_IS_FULL_FLAG | _PREPARED_ROUTE_IS_PROXY_FLAG +) @cute.jit @@ -108,7 +115,7 @@ def _swaps_forwards_packed_route_full(cfg: FmhaDecodeConfig) -> bool: return ( cfg.tile_size_q == 8 and cfg.kv_block_size == 8 - and not cfg.use_kv_valid_bits + and not cfg.uses_prepared_score_keep_words and not cfg.uses_uniform_causal_mask and not cfg.uses_per_row_causal_mask ) @@ -116,12 +123,16 @@ def _swaps_forwards_packed_route_full(cfg: FmhaDecodeConfig) -> bool: def _kv_retained_route_words( route_layout: _BlockSparseRouteLayout, + *, + retain_proxy_kind: bool = False, ) -> int: """Return the aligned SMEM words retained from K issue through V. - Contiguous routes retain their existing load-origin payload. Paged routes - retain parallel logical-origin and physical-page-ID arrays so every atom - has an independent storage locator; invalid entries use ``(-1, -1)``. + Contiguous routes retain their load-origin payload. Paged routes retain + parallel logical-origin and physical-page-ID arrays so every atom has an + independent storage locator. A two-origin contiguous route additionally + keeps its atom-valid mask. Proxy-capable exact/proxy routes reserve the + final aligned word for an explicit source kind. """ payload_words = route_layout.logical_origins_per_route @@ -129,6 +140,8 @@ def _kv_retained_route_words( payload_words *= 2 elif route_layout.logical_origins_per_route == 2: payload_words += 1 + if retain_proxy_kind: + payload_words += 1 return ((payload_words + 3) // 4) * 4 @@ -139,9 +152,9 @@ class _BlockSparseSoftmaxStagingLayout: Keeps retains all route origins, a flags word, alignment padding, and the optional K32 token words. KV256 consumers then select the four words owned by their spatial half. SWAP stores execution-ordered origins followed by - optional logical K32 token words, one for each consumer warp. Its - noncausal Q8/B8 profile without token bits packs route-full into the - otherwise-zero low bit of each warp's first aligned origin. + optional logical K32 token words, one for each consumer warp. Selected + prepared route flags travel in the otherwise-zero low bits of each warp's + first aligned origin. """ # Logical-origin scalars staged for one complete KV route. @@ -206,14 +219,23 @@ class SmemBlockSparseKvMetadataResource(DecodeGenResourceBase): """Pipeline-free route metadata retained from one K issue through V. ``route_metadata`` points at the first prepared GMEM record. Resolution - returns logical origins to masking consumers. The private SMEM copy keeps - contiguous load origins or paged ``(logical origin, physical page ID)`` - pairs through the matching V issue. Invalid atoms are retained as safe - storage-specific OOB coordinates. + returns logical/source-domain origins to masking consumers. The private + SMEM copy keeps contiguous load origins or paged ``(logical origin, + physical page ID)`` pairs through the matching V issue. Invalid atoms are + retained as safe storage-specific OOB coordinates. The proxy-capable + contiguous specialization interprets origins in the selected source domain + (summary tokens for proxy routes, K/V tokens for exact routes) and + retains the prepared route kind in a separate aligned word. Exact-only + specializations keep their original allocation. """ _task_local_specs: ClassVar[tuple[tuple, ...]] = ( - ("resolved_origin0_slot", Int32, Int32(0), "First logical origin."), + ( + "resolved_record_word_slot", + Int32, + Int32(0), + "Lane-owned record word; locator lanes carry route origins.", + ), ("resolved_origin1_slot", Int32, Int32(0), "Second logical origin."), ( "resolved_atom_validity_slot", @@ -227,6 +249,18 @@ class SmemBlockSparseKvMetadataResource(DecodeGenResourceBase): Int32(-1), "Metadata-relative record offset, or -1 for a dummy route.", ), + ( + "prefetched_record_word_slot", + Int32, + Int32(0), + "Lane-owned record word loaded one resolution ahead.", + ), + ( + "prefetched_record_offset_slot", + Int32, + Int32(-1), + "Record offset of the prefetched route, or -1 for a dummy route.", + ), ) cfg: Constexpr[FmhaDecodeConfig] = None inst_id: Constexpr[int] = 0 @@ -236,7 +270,7 @@ class SmemBlockSparseKvMetadataResource(DecodeGenResourceBase): _retained_route_words: Constexpr[int] = 0 _alloc: Constexpr[SmemAllocation | None] = None _smem_words: cutlass.Array = None - resolved_origin0_slot: Constexpr[TaskLocalVariable] = ( + resolved_record_word_slot: Constexpr[TaskLocalVariable] = ( TaskLocalVariable.uninitialized() ) resolved_origin1_slot: Constexpr[TaskLocalVariable] = ( @@ -248,13 +282,25 @@ class SmemBlockSparseKvMetadataResource(DecodeGenResourceBase): route_record_word_offset_slot: Constexpr[TaskLocalVariable] = ( TaskLocalVariable.uninitialized() ) + prefetched_record_word_slot: Constexpr[TaskLocalVariable] = ( + TaskLocalVariable.uninitialized() + ) + prefetched_record_offset_slot: Constexpr[TaskLocalVariable] = ( + TaskLocalVariable.uninitialized() + ) def __post_init__(self) -> None: """Derive the retained K/V payload from the prepared route layout.""" assert self.route_layout is not None assert self.route_layout.is_paged == self.cfg.use_paged_kv - self._retained_route_words = _kv_retained_route_words(self.route_layout) + if self.cfg.use_block_sparse_proxy_routes: + assert not self.route_layout.is_paged + assert self.route_layout.uses_one_warp_transport + self._retained_route_words = _kv_retained_route_words( + self.route_layout, + retain_proxy_kind=self.cfg.use_block_sparse_proxy_routes, + ) super().__post_init__() def _init_placeholder_state(self) -> None: @@ -332,9 +378,81 @@ def _prepared_route_physical_page_id_if_valid( ) return physical_page_id + @cute.jit + def _route_record_word_offset( + self, stage_info: StageInfo, route_idx: Int32 + ) -> Int32: + """Return the record offset of one route index, or -1 past the row.""" + + task_cache = _decode_gen_task_cache(stage_info) + row_route_begin = _sparse_task_cache_route_begin(task_cache) + route_count = _sparse_task_cache_route_count(task_cache) + route_record_word_offset = Int32(-1) + if route_idx < route_count: + route_record_word_offset = (row_route_begin + route_idx) * Int32( + self.route_layout.route_metadata_stride_words + ) + return cute.arch.make_warp_uniform(route_record_word_offset) + + @consumer_work( + returns=( + prefetched_record_word_slot, + prefetched_record_offset_slot, + ) + ) + @cute.jit + def prefetch_route( + self, stage_info: StageInfo, *, target: Constexpr[str] + ) -> tuple[Int32, Int32]: + """Issue the record load for a route that ``resolve_route`` uses later. + + ``target`` selects the route relative to the calling section: + ``"head"`` is this instance's HEAD route, ``"first_loop"`` the route of + LOOP iteration 0 (called from HEAD), ``"current_loop"`` the route of + the calling LOOP iteration (no pipelining), and ``"next_loop"`` the + route of the following LOOP iteration. Only the lane-distributed load is issued + here; the warp broadcasts happen in ``resolve_route`` so the global + memory latency overlaps the TMA issue of the current route instead of + stalling the load warp. Layouts without one-warp transport keep their + loads in ``resolve_route`` and get placeholder values here. + """ + + assert self.route_metadata is not None + num_insts = Int32(self.cfg.num_insts_kv) + if cutlass.const_expr(target == "head"): + route_idx = Int32(self.inst_id) + elif cutlass.const_expr(target == "first_loop"): + route_idx = num_insts + Int32(self.inst_id) + elif cutlass.const_expr(target == "current_loop"): + route_idx = (stage_info.loop_offset + Int32(1)) * num_insts + Int32( + self.inst_id + ) + else: + route_idx = (stage_info.loop_offset + Int32(2)) * num_insts + Int32( + self.inst_id + ) + route_record_word_offset = self._route_record_word_offset(stage_info, route_idx) + record_word = Int32(0) + if cutlass.const_expr(self.route_layout.uses_one_warp_transport): + assert self.route_layout.token_words_word_offset is not None + meaningful_words = ( + self.route_layout.token_words_word_offset + + self.route_layout.token_words_per_route + ) + lane_idx = cute.arch.thread_idx()[0] & Int32(0x1F) + if lane_idx < Int32(self.route_layout.logical_origins_per_route): + record_word = Int32(-1) + if route_record_word_offset >= Int32(0) and lane_idx < Int32( + meaningful_words + ): + record_word = Int32( + self.route_metadata[route_record_word_offset + lane_idx] + ) + return record_word, route_record_word_offset + @consumer_work( returns=( - resolved_origin0_slot, + resolved_record_word_slot, resolved_origin1_slot, resolved_atom_validity_slot, route_record_word_offset_slot, @@ -342,14 +460,23 @@ def _prepared_route_physical_page_id_if_valid( ) @cute.jit def resolve_route( - self, stage_info: StageInfo, *, section: Constexpr[FmhaStage] + self, + stage_info: StageInfo, + *, + section: Constexpr[FmhaStage], + prefetched_record_word_slot: Int32, + prefetched_record_offset_slot: Int32, ) -> tuple[Int32, Int32, Int32, Int32]: - """Load this resource instance's real or dummy prepared KV route.""" + """Resolve this instance's real or dummy prepared KV route. + + One-warp-transport layouts consume the words that ``prefetch_route`` + loaded earlier; other layouts load their record here. The routed + inputs carry the task-local slot names so that every ``prefetch_route`` + call, including the one at the end of the previous LOOP iteration, + updates the value read here. + """ assert self.route_metadata is not None - task_cache = _decode_gen_task_cache(stage_info) - row_route_begin = _sparse_task_cache_route_begin(task_cache) - route_count = _sparse_task_cache_route_count(task_cache) # HEAD publishes one route per instruction. LOOP starts after those # two publications, hence the one-based loop offset below. Keeping the # constexpr branch local lets the task scheduler specialize each work @@ -361,12 +488,12 @@ def resolve_route( self.cfg.num_insts_kv ) + Int32(self.inst_id) lane_idx = cute.arch.thread_idx()[0] & Int32(0x1F) - route_record_word_offset = Int32(-1) - if route_idx < route_count: - route_record_word_offset = (row_route_begin + route_idx) * Int32( - self.route_layout.route_metadata_stride_words + if cutlass.const_expr(self.route_layout.uses_one_warp_transport): + route_record_word_offset = prefetched_record_offset_slot + else: + route_record_word_offset = self._route_record_word_offset( + stage_info, route_idx ) - route_record_word_offset = cute.arch.make_warp_uniform(route_record_word_offset) num_logical_origins = self.route_layout.logical_origins_per_route uses_two_fragment_route = num_logical_origins == 2 @@ -378,13 +505,37 @@ def resolve_route( atom_valid_mask = Int32(0) route_record_is_valid = route_record_word_offset >= Int32(0) + if cutlass.const_expr(self.route_layout.uses_one_warp_transport): + resolved_record_word = prefetched_record_word_slot + atom_valid_mask = _warp_broadcast_i32( + resolved_record_word, + self.route_layout.atom_valid_mask_word_offset, + ) + if cutlass.const_expr(uses_two_fragment_route): + return ( + resolved_record_word, + _warp_broadcast_i32(resolved_record_word, 1), + atom_valid_mask, + route_record_word_offset, + ) + atom_is_valid = cutlass.Boolean( + lane_idx < Int32(num_logical_origins) + and (atom_valid_mask & (Int32(1) << lane_idx)) != Int32(0) + ) + return ( + resolved_record_word, + Int32(0), + Int32(atom_is_valid), + route_record_word_offset, + ) + if route_record_is_valid: if lane_idx < Int32(num_logical_origins): logical_origin = self._prepared_route_logical_origin( route_record_word_offset, lane_idx, ) - if cutlass.const_expr(self.cfg.use_kv_valid_bits): + if cutlass.const_expr(self.cfg.uses_prepared_score_keep_words): if lane_idx == Int32(valid_mask_lane): atom_valid_mask = Int32( self.route_metadata[ @@ -396,7 +547,7 @@ def resolve_route( if cutlass.const_expr(uses_two_fragment_route): origin0 = _warp_broadcast_i32(logical_origin, 0) origin1 = _warp_broadcast_i32(logical_origin, 1) - if cutlass.const_expr(self.cfg.use_kv_valid_bits): + if cutlass.const_expr(self.cfg.uses_prepared_score_keep_words): # The validity word shares the prepared record's cache line # with fields consumed shortly afterward by Softmax. atom_valid_mask = _warp_broadcast_i32(atom_valid_mask, valid_mask_lane) @@ -412,7 +563,7 @@ def resolve_route( # Wider routes stay lane-distributed: each active lane carries only # its origin and validity through the existing three-scalar K/V ABI. valid = cutlass.Boolean(False) - if cutlass.const_expr(self.cfg.use_kv_valid_bits): + if cutlass.const_expr(self.cfg.uses_prepared_score_keep_words): atom_valid_mask = _warp_broadcast_i32(atom_valid_mask, valid_mask_lane) if lane_idx < Int32(num_logical_origins): valid = (atom_valid_mask & (Int32(1) << lane_idx)) != Int32(0) @@ -427,7 +578,7 @@ def store_route( self, stage_info: StageInfo, *, - resolved_origin0: Int32, + resolved_record_word: Int32, resolved_origin1: Int32, resolved_atom_validity: Int32, route_record_word_offset: Int32, @@ -439,11 +590,11 @@ def store_route( num_origins = self.route_layout.logical_origins_per_route if cutlass.const_expr(self.route_layout.is_paged): if lane_idx < Int32(num_origins): - logical_origin = Int32(resolved_origin0) + logical_origin = Int32(resolved_record_word) atom_is_valid = resolved_atom_validity != Int32(0) if cutlass.const_expr(num_origins == 2): if lane_idx == Int32(0): - logical_origin = resolved_origin0 + logical_origin = resolved_record_word else: logical_origin = resolved_origin1 atom_is_valid = ( @@ -463,14 +614,14 @@ def store_route( ] = physical_page_id elif cutlass.const_expr(num_origins == 2): if lane_idx == Int32(0): - self._smem_words[Int32(0)] = resolved_origin0 + self._smem_words[Int32(0)] = resolved_record_word self._smem_words[Int32(1)] = resolved_origin1 self._smem_words[ Int32(self.route_layout.atom_valid_mask_word_offset) ] = resolved_atom_validity else: - if lane_idx < Int32(self.route_layout.logical_origins_per_route): - load_origin = Int32(resolved_origin0) + if lane_idx < Int32(num_origins): + load_origin = Int32(resolved_record_word) if resolved_atom_validity == Int32(0): # Fine-route K and V both consume this retained value. # Materialize their TensorMap OOB coordinate once here @@ -478,6 +629,12 @@ def store_route( # every atom copy in both producer passes. load_origin = Int32(self.tma_oob_origin) self._smem_words[lane_idx] = load_origin + if cutlass.const_expr(self.cfg.use_block_sparse_proxy_routes): + if lane_idx == Int32(self.route_layout.route_flags_word_offset): + prepared_route_flags = Int32(resolved_record_word) + self._smem_words[Int32(self._retained_route_words - 1)] = Int32( + prepared_route_flags + ) & Int32(_PREPARED_ROUTE_IS_PROXY_FLAG) # K consumes this slot immediately, while V consumes it at the start # of the next cadence. Both execute in this warp, so a warp fence is # sufficient; no cross-warp mbarrier belongs here. @@ -523,6 +680,20 @@ def route_atom_valid_mask(self) -> Int32: self._smem_words[Int32(self.route_layout.atom_valid_mask_word_offset)] ) + @cute.jit + def route_is_proxy(self) -> cutlass.Boolean: + """Return the retained prepared route kind for the current K/V pair.""" + + if cutlass.const_expr(not self.cfg.use_block_sparse_proxy_routes): + return cutlass.Boolean(False) + return cutlass.Boolean( + ( + Int32(self._smem_words[Int32(self._retained_route_words - 1)]) + & Int32(_PREPARED_ROUTE_IS_PROXY_FLAG) + ) + != Int32(0) + ) + @dataclass(kw_only=True) class SmemBlockSparseSoftmaxMetadataResource(DecodeGenResourceBase): @@ -533,8 +704,8 @@ class SmemBlockSparseSoftmaxMetadataResource(DecodeGenResourceBase): producer passes the resolved payload explicitly instead of recomputing it. For Keeps, every route token word moves through SMEM without a data-dependent branch; each consumer receives at most four words through - the stable task-local ABI. A runtime route-full bit can skip per-score token - predicates while leaving structural masking independent. + the stable task-local ABI. Runtime route flags carry the conservative FULL + summary and, for proxy-capable builds, the route source kind. """ _task_local_specs: ClassVar[tuple[tuple, ...]] = ( @@ -562,7 +733,7 @@ class SmemBlockSparseSoftmaxMetadataResource(DecodeGenResourceBase): "softmax_token_word2_slot", Uint32, Uint32(0xFFFFFFFF), - "Loaded third Keeps token word or SWAP's B8 route-full summary.", + "Loaded third Keeps token word or SWAP's packed route flags.", ), ( "softmax_token_word3_slot", @@ -605,6 +776,8 @@ def __post_init__(self) -> None: assert self.route_layout is not None assert self.route_layout.is_paged == self.cfg.use_paged_kv + if self.cfg.use_block_sparse_proxy_routes: + assert self.route_layout.uses_one_warp_transport self.staging_layout = _BlockSparseSoftmaxStagingLayout.create( use_keeps_mma_ab=self.cfg.use_keeps_mma_ab, route_layout=self.route_layout, @@ -680,57 +853,82 @@ def _consumer_stage_base(self) -> Int32: def _store_route_swaps( self, stage_info: StageInfo, - resolved_origin0: Int32, + resolved_record_word: Int32, resolved_origin1: Int32, resolved_atom_validity: Int32, route_record_word_offset: Int32, ) -> None: """Stage SWAP origins and optional logical-K32 token metadata. - The noncausal Q8/B8 profile without token bits also packs prepare's - route-full summary into bit 0 of each warp's first aligned origin. + Selected prepared route flags use the free low bits of each warp's + first aligned origin. """ lane_idx = cute.arch.thread_idx()[0] & Int32(0x1F) stage_base = self._producer_stage_base(stage_info) task_cache = _decode_gen_task_cache(stage_info) seq_len_kv = Int32(task_cache[_TASK_CACHE_SEQ_LEN_KV]) - route_record_is_valid = route_record_word_offset >= Int32(0) - - packed_route_full = Int32(0) - if cutlass.const_expr(_swaps_forwards_packed_route_full(self.cfg)): - if lane_idx == Int32(0) and route_record_is_valid: - packed_route_full = Int32( - self.route_metadata[ - route_record_word_offset - + Int32(self.route_layout.route_flags_word_offset) - ] - ) & Int32(_PREPARED_ROUTE_IS_FULL_FLAG) - packed_route_full = _warp_broadcast_i32(packed_route_full, 0) + uses_one_warp_transport = self.route_layout.uses_one_warp_transport + if cutlass.const_expr( + not uses_one_warp_transport + and ( + _swaps_forwards_packed_route_full(self.cfg) + or self.cfg.uses_prepared_score_keep_words + ) + ): + route_record_is_valid = route_record_word_offset >= Int32(0) + + packed_route_flags = Int32(0) + if cutlass.const_expr( + _swaps_forwards_packed_route_full(self.cfg) + or self.cfg.use_block_sparse_proxy_routes + ): + if cutlass.const_expr(uses_one_warp_transport): + packed_route_flags = _warp_broadcast_i32( + resolved_record_word, + self.route_layout.route_flags_word_offset, + ) + else: + if lane_idx == Int32(0) and route_record_is_valid: + packed_route_flags = Int32( + self.route_metadata[ + route_record_word_offset + + Int32(self.route_layout.route_flags_word_offset) + ] + ) & Int32(_PREPARED_ROUTE_IS_FULL_FLAG) + packed_route_flags = _warp_broadcast_i32(packed_route_flags, 0) softmax_origin = Int32(-1) if cutlass.const_expr(self.cfg.kv_block_size < 64): if lane_idx < Int32(self.staging_layout.num_origin_words): - softmax_origin = Int32(resolved_origin0) + softmax_origin = Int32(resolved_record_word) if resolved_atom_validity == Int32(0): softmax_origin = Int32(-1) - if cutlass.const_expr(_swaps_forwards_packed_route_full(self.cfg)): - # Replicate route-full in each K32 slice's first origin; - # B8 alignment leaves bit 0 free for the summary. + if cutlass.const_expr( + _swaps_forwards_packed_route_full(self.cfg) + or self.cfg.use_block_sparse_proxy_routes + ): + # Every SWAP atom is at least B8 aligned. Replicate the + # route flags in each K32 slice's first origin so the + # established seven-slot Softmax ABI also carries source + # kind without growing the staged payload. if lane_idx % Int32(self.staging_layout.origins_per_warp) == Int32( 0 ): softmax_origin = ( - softmax_origin & Int32(_SWAPS_PACKED_ROUTE_FULL_CLEAR_MASK) - ) | packed_route_full + softmax_origin & Int32(_SWAPS_PACKED_ROUTE_FLAGS_CLEAR_MASK) + ) | packed_route_flags self._smem_words[stage_base + lane_idx] = softmax_origin else: # SWAP with a coarse KV atom expands the two resolved KV64 # fragments into the four logical K32 origins consumed by its # four softmax warps. + coarse_origin0 = Int32(resolved_record_word) + if cutlass.const_expr(uses_one_warp_transport): + coarse_origin0 = _warp_broadcast_i32(resolved_record_word, 0) if lane_idx < Int32(4): fragment_idx = lane_idx >> Int32(1) - softmax_origin = Int32(resolved_origin0) + softmax_origin = coarse_origin0 valid = (resolved_atom_validity & Int32(1)) != Int32(0) if fragment_idx == Int32(1): softmax_origin = Int32(resolved_origin1) @@ -738,13 +936,30 @@ def _store_route_swaps( softmax_origin = softmax_origin + (lane_idx & Int32(1)) * Int32(32) if not valid or softmax_origin >= seq_len_kv: softmax_origin = Int32(-1) + if cutlass.const_expr(self.cfg.use_block_sparse_proxy_routes): + # Coarse SWAP expands KV64 atoms to K32-aligned origins; + # their low bits carry the same typed-route flags as the + # fine-route representation above. + softmax_origin = ( + softmax_origin & Int32(_SWAPS_PACKED_ROUTE_FLAGS_CLEAR_MASK) + ) | packed_route_flags self._smem_words[stage_base + lane_idx] = softmax_origin - if cutlass.const_expr(self.cfg.use_kv_valid_bits): + if cutlass.const_expr(self.cfg.uses_prepared_score_keep_words): assert self.route_metadata is not None assert self.route_layout.token_words_word_offset is not None assert self.staging_layout.token_words_word_offset is not None - if lane_idx < Int32(self.route_layout.token_words_per_route): + if cutlass.const_expr(uses_one_warp_transport): + token_begin = Int32(self.route_layout.token_words_word_offset) + token_end = token_begin + Int32(self.route_layout.token_words_per_route) + if lane_idx >= token_begin and lane_idx < token_end: + self._smem_words[ + stage_base + + Int32(self.staging_layout.token_words_word_offset) + + lane_idx + - token_begin + ] = Int32(resolved_record_word) + elif lane_idx < Int32(self.route_layout.token_words_per_route): logical_word = Uint32(0) if route_record_is_valid: logical_word = Uint32( @@ -765,7 +980,7 @@ def _store_route_swaps( def _store_route_keeps( self, stage_info: StageInfo, - resolved_origin0: Int32, + resolved_record_word: Int32, resolved_origin1: Int32, resolved_atom_validity: Int32, route_record_word_offset: Int32, @@ -779,8 +994,12 @@ def _store_route_keeps( assert self.staging_layout.route_flags_word_offset is not None lane_idx = cute.arch.thread_idx()[0] & Int32(0x1F) - route_record_is_valid = route_record_word_offset >= Int32(0) num_origins = self.route_layout.logical_origins_per_route + uses_one_warp_transport = self.route_layout.uses_one_warp_transport + if cutlass.const_expr( + not uses_one_warp_transport and self.cfg.uses_prepared_score_keep_words + ): + route_record_is_valid = route_record_word_offset >= Int32(0) route_flags = Int32(resolved_atom_validity) if cutlass.const_expr(num_origins > 2): route_flags = Int32( @@ -788,47 +1007,73 @@ def _store_route_keeps( lane_idx < Int32(num_origins) and resolved_atom_validity != Int32(0) ) ) + token_word = Uint32(0) route_token_mask_is_full = cutlass.Boolean(False) - if cutlass.const_expr(self.cfg.use_kv_valid_bits): - assert self.route_metadata is not None + if cutlass.const_expr(self.cfg.uses_prepared_score_keep_words): assert self.route_layout.token_words_word_offset is not None - gmem_route_flags = Int32(0) - if lane_idx == Int32(0) and route_record_is_valid: - gmem_route_flags = Int32( - self.route_metadata[ - route_record_word_offset - + Int32(self.route_layout.route_flags_word_offset) - ] - ) - gmem_route_flags = _warp_broadcast_i32(gmem_route_flags, 0) - # Prepared bit 0 summarizes the whole route. Staged low bits are - # already fragment validity, so remap the summary above them. - route_token_mask_is_full = cutlass.Boolean( - (gmem_route_flags & Int32(_PREPARED_ROUTE_IS_FULL_FLAG)) != Int32(0) - ) - if ( - lane_idx < Int32(self.route_layout.token_words_per_route) - and route_record_is_valid - ): - token_word = Uint32( - self.route_metadata[ - route_record_word_offset - + Int32(self.route_layout.token_words_word_offset) - + lane_idx - ] + assert self.route_metadata is not None + if cutlass.const_expr(not uses_one_warp_transport): + gmem_route_flags = Int32(0) + if lane_idx == Int32(0) and route_record_is_valid: + gmem_route_flags = Int32( + self.route_metadata[ + route_record_word_offset + + Int32(self.route_layout.route_flags_word_offset) + ] + ) + gmem_route_flags = _warp_broadcast_i32(gmem_route_flags, 0) + # Prepared bit 0 summarizes the whole route. Staged low bits + # already hold fragment validity, so remap it above them. + route_token_mask_is_full = cutlass.Boolean( + (gmem_route_flags & Int32(_PREPARED_ROUTE_IS_FULL_FLAG)) != Int32(0) ) + if ( + lane_idx < Int32(self.route_layout.token_words_per_route) + and route_record_is_valid + ): + token_word = Uint32( + self.route_metadata[ + route_record_word_offset + + Int32(self.route_layout.token_words_word_offset) + + lane_idx + ] + ) stage_base = self._producer_stage_base(stage_info) if cutlass.const_expr(num_origins == 2): if lane_idx == Int32(0): - self._smem_words[stage_base] = Int32(resolved_origin0) + self._smem_words[stage_base] = Int32(resolved_record_word) self._smem_words[stage_base + Int32(1)] = Int32(resolved_origin1) - else: - if lane_idx < Int32(num_origins): - self._smem_words[stage_base + lane_idx] = Int32(resolved_origin0) - if lane_idx == Int32(0): - if cutlass.const_expr(self.cfg.use_kv_valid_bits): + elif lane_idx < Int32(num_origins): + self._smem_words[stage_base + lane_idx] = Int32(resolved_record_word) + + if cutlass.const_expr(uses_one_warp_transport): + if lane_idx == Int32(self.route_layout.route_flags_word_offset): + prepared_route_flags = Int32(resolved_record_word) + route_flags = route_flags | ( + Int32( + (prepared_route_flags & Int32(_PREPARED_ROUTE_IS_FULL_FLAG)) + != Int32(0) + ) + * Int32(_SOFTMAX_TOKEN_MASK_IS_FULL_FLAG) + ) + if cutlass.const_expr(self.cfg.use_block_sparse_proxy_routes): + route_flags = route_flags | ( + Int32( + ( + prepared_route_flags + & Int32(_PREPARED_ROUTE_IS_PROXY_FLAG) + ) + != Int32(0) + ) + * Int32(_SOFTMAX_ROUTE_IS_PROXY_FLAG) + ) + self._smem_words[ + stage_base + Int32(self.staging_layout.route_flags_word_offset) + ] = route_flags + elif lane_idx == Int32(0): + if cutlass.const_expr(self.cfg.uses_prepared_score_keep_words): route_flags = route_flags | ( Int32(route_token_mask_is_full) * Int32(_SOFTMAX_TOKEN_MASK_IS_FULL_FLAG) @@ -836,9 +1081,20 @@ def _store_route_keeps( self._smem_words[ stage_base + Int32(self.staging_layout.route_flags_word_offset) ] = route_flags - if cutlass.const_expr(self.cfg.use_kv_valid_bits): + + if cutlass.const_expr(self.cfg.uses_prepared_score_keep_words): assert self.staging_layout.token_words_word_offset is not None - if lane_idx < Int32(self.route_layout.token_words_per_route): + if cutlass.const_expr(uses_one_warp_transport): + token_begin = Int32(self.route_layout.token_words_word_offset) + token_end = token_begin + Int32(self.route_layout.token_words_per_route) + if lane_idx >= token_begin and lane_idx < token_end: + self._smem_words[ + stage_base + + Int32(self.staging_layout.token_words_word_offset) + + lane_idx + - token_begin + ] = Int32(resolved_record_word) + elif lane_idx < Int32(self.route_layout.token_words_per_route): self._smem_words[ stage_base + Int32(self.staging_layout.token_words_word_offset) @@ -852,7 +1108,7 @@ def store_route( self, stage_info: StageInfo, *, - resolved_origin0: Int32, + resolved_record_word: Int32, resolved_origin1: Int32, resolved_atom_validity: Int32, route_record_word_offset: Int32, @@ -862,7 +1118,7 @@ def store_route( if cutlass.const_expr(self.cfg.use_keeps_mma_ab): self._store_route_keeps( stage_info, - resolved_origin0, + resolved_record_word, resolved_origin1, resolved_atom_validity, route_record_word_offset, @@ -870,7 +1126,7 @@ def store_route( else: self._store_route_swaps( stage_info, - resolved_origin0, + resolved_record_word, resolved_origin1, resolved_atom_validity, route_record_word_offset, @@ -886,8 +1142,8 @@ def _load_route_swaps_values( this Softmax warp's logical K32 slice; unused or invalid origins are negative. To preserve the shared seven-slot task ABI, origin 2/3 subsequently travel through the shared route-flags/token-word-0 slots. - Token-word 1 carries the logical K32 mask, token-word 2 carries - route-full, and token-word 3 is unused. + Token-word 1 carries the logical K32 mask, token-word 2 carries packed + route flags, and token-word 3 is unused. """ stage_base = self._consumer_stage_base() @@ -907,7 +1163,7 @@ def _load_route_swaps_values( origin3 = Int32(self._smem_words[warp_origin_base + Int32(3)]) token_word = Uint32(0xFFFFFFFF) - if cutlass.const_expr(self.cfg.use_kv_valid_bits): + if cutlass.const_expr(self.cfg.uses_prepared_score_keep_words): assert self.staging_layout.token_words_word_offset is not None token_word = Uint32( self._smem_words[ @@ -917,9 +1173,25 @@ def _load_route_swaps_values( ] ) route_flags = Uint32(0) - if cutlass.const_expr(_swaps_forwards_packed_route_full(self.cfg)): - route_flags = Uint32(origin0 & Int32(1)) - origin0 = origin0 & Int32(_SWAPS_PACKED_ROUTE_FULL_CLEAR_MASK) + if cutlass.const_expr( + _swaps_forwards_packed_route_full(self.cfg) + or self.cfg.use_block_sparse_proxy_routes + ): + packed_route_flags = origin0 & Int32( + _PREPARED_ROUTE_IS_FULL_FLAG | _PREPARED_ROUTE_IS_PROXY_FLAG + ) + route_flags = Uint32( + packed_route_flags & Int32(_PREPARED_ROUTE_IS_FULL_FLAG) + ) + if cutlass.const_expr(self.cfg.use_block_sparse_proxy_routes): + route_flags = route_flags | Uint32( + Int32( + (packed_route_flags & Int32(_PREPARED_ROUTE_IS_PROXY_FLAG)) + != Int32(0) + ) + * Int32(_SOFTMAX_ROUTE_IS_PROXY_FLAG) + ) + origin0 = origin0 & Int32(_SWAPS_PACKED_ROUTE_FLAGS_CLEAR_MASK) return ( origin0, origin1, @@ -986,17 +1258,21 @@ def load_route( valid0 = (stored_route_flags >> origin0_idx) & Int32(1) valid1 = (stored_route_flags >> origin1_idx) & Int32(1) route_flags = valid0 | (valid1 << Int32(1)) - if cutlass.const_expr(self.cfg.use_kv_valid_bits): + if cutlass.const_expr(self.cfg.uses_prepared_score_keep_words): route_flags = route_flags | ( stored_route_flags & Int32(_SOFTMAX_TOKEN_MASK_IS_FULL_FLAG) ) + if cutlass.const_expr(self.cfg.use_block_sparse_proxy_routes): + route_flags = route_flags | ( + stored_route_flags & Int32(_SOFTMAX_ROUTE_IS_PROXY_FLAG) + ) origin0 = Int32(self._smem_words[stage_base + origin0_idx]) origin1 = Int32(self._smem_words[stage_base + origin1_idx]) token_word0 = Uint32(0xFFFFFFFF) token_word1 = Uint32(0xFFFFFFFF) token_word2 = Uint32(0xFFFFFFFF) token_word3 = Uint32(0xFFFFFFFF) - if cutlass.const_expr(self.cfg.use_kv_valid_bits): + if cutlass.const_expr(self.cfg.uses_prepared_score_keep_words): assert self.staging_layout.token_words_word_offset is not None if cutlass.const_expr(self.route_layout.kv_route_size == 256): token_base = Int32(self.staging_layout.token_words_word_offset) @@ -1014,28 +1290,6 @@ def load_route( token_word3 = Uint32( self._smem_words[stage_base + token_base + word1_idx + Int32(1)] ) - elif cutlass.const_expr(self.cfg.tile_size_q == 64): - lane_idx = cute.arch.thread_idx()[0] & Int32(0x1F) - local_word_base = _keeps_col_base( - self.cfg, - lane_idx, - self.cfg.num_s_regs_per_thread, - ) >> Int32(5) - token_word0 = Uint32( - self._smem_words[ - stage_base - + Int32(self.staging_layout.token_words_word_offset) - + local_word_base - ] - ) - token_word1 = Uint32( - self._smem_words[ - stage_base - + Int32(self.staging_layout.token_words_word_offset) - + local_word_base - + Int32(1) - ] - ) else: token_word0 = Uint32( self._smem_words[ diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_p.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_p.py index bb27c3c553f0..330f15b05691 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_p.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_p.py @@ -25,7 +25,7 @@ import cutlass import cutlass.cute as cute -from cutlass import Float32, Int32, Int64 +from cutlass import Float32, Int32, Int64, Uint32 from cutlass.experimental import primitives as prims from cutlass.experimental.task_scheduling.memory import ( @@ -42,6 +42,7 @@ ) from ..fmha_decode_config import FmhaDecodeConfig +from ...._block_sparse.common import _block_sparse_proxy_summary_geometry from ...placeholder_helpers import _placeholder_smem_array from .helpers_common import ( Constexpr, @@ -58,11 +59,13 @@ _is_last_loop_iteration, _keeps_col_base, _keeps_row_idx, + _keeps_tcgen05_ld, _keeps_tcgen05_st, _named_barrier_arrive, _neg_max_f32, _pack_float2_to_bf16, _pack_float2_to_fp16, + _swaps_routed_coordinate, _wait_for_mbarrier_phase, ) from .helpers_output import ( @@ -76,11 +79,27 @@ _compute_fp8_p_regs_and_local_sums, _compute_fp8_p_regs_and_local_sums_dense, _compute_p_values_and_local_sums_dense, + _ex2_emulation_packed_f32x2, _pack_float4_to_fp8_e4m3, _pack_float4_to_fp8_e4m3_inline, ) +from .smem_block_sparse_metadata import _SOFTMAX_ROUTE_IS_PROXY_FLAG from .tmem_s import TmemSResource +# Tunable: number of score pairs per streamed fragment whose exponentials run +# as FMA polynomials instead of MUFU. The MUFU issue rate bounds the fragment +# otherwise, while the FMA pipe is nearly idle in the softmax warps. Larger +# shares grow the fragment body and the softmax warps become instruction-fetch +# bound again, so one quarter of the 16 pairs is the measured optimum. +KV_TILE_256_EX2_EMULATED_PAIRS = 4 + + +def _pair_uses_ex2_emulation(pair_idx: int, pairs_per_fragment: int) -> bool: + """Spread the emulated pairs evenly across a fragment's score pairs.""" + count = KV_TILE_256_EX2_EMULATED_PAIRS + pairs = pairs_per_fragment + return ((pair_idx + 1) * count) // pairs != (pair_idx * count) // pairs + @dataclass(kw_only=True) class SmemPResource(DecodeGenResourceBase): @@ -88,7 +107,7 @@ class SmemPResource(DecodeGenResourceBase): Softmax producers convert S to P, store it in the profile's TMEM or SMEM layout, and publish local sums back to TmemS. Most profiles use the generic - full/empty P pipeline. KV256 instead publishes four independently ready + full/empty P pipeline. Streamed profiles instead publish four independently ready K32 TMEM fragments; BMM2 consumes those fragments in order, while the matching TmemO full barrier prevents the next QK from overwriting aliased P. """ @@ -152,7 +171,7 @@ def _init_placeholder_state(self) -> None: ) def get_smem_requirements(self) -> list[SmemAllocation]: - """Allocate P storage or the KV256 fragment-ready barriers.""" + """Allocate P storage or the streamed fragment-ready barriers.""" if self.cfg.streams_tmem_p_fragments: if self._fragment_ready_alloc is None: self._fragment_ready_alloc = SmemAllocation( @@ -173,7 +192,7 @@ def get_smem_requirements(self) -> list[SmemAllocation]: @cute.jit def _bind_fragment_ready(self, context: ResourceContext | None = None) -> None: - """Bind the one-way KV256 P-ready barriers from the SMEM context.""" + """Bind the one-way streamed P-ready barriers from the SMEM context.""" if cutlass.const_expr( self.cfg.streams_tmem_p_fragments and context is not None @@ -191,7 +210,7 @@ def _bind_fragment_ready(self, context: ResourceContext | None = None) -> None: def create_function_variables( self, context: ResourceContext | None = None ) -> ResourceVars: - """Bind and initialize KV256's per-fragment ready barriers.""" + """Bind and initialize the streamed per-fragment ready barriers.""" self._bind_fragment_ready(context) if cutlass.const_expr(self.cfg.streams_tmem_p_fragments): tidx, _, _ = cute.arch.thread_idx() @@ -299,43 +318,203 @@ def init_descriptor_state(self, stage_info: StageInfo) -> None: # work can publish a valid descriptor or TMEM address for this tile. self._create_initial_task_locals(stage_info.context) + @cute.jit + def _apply_proxy_route_denominator_mass( + self, + local_sum: Float32, + tail_p: Float32, + route_is_proxy: Int32, + ) -> Float32: + """Weight only a proxy route's softmax denominator by block mass.""" + + if cutlass.const_expr(not self.cfg.use_block_sparse_proxy_routes): + return local_sum + if route_is_proxy != Int32(0): + _, tail_len = _block_sparse_proxy_summary_geometry( + self.cfg.static_seq_len_kv, + self.cfg.kv_block_size, + ) + local_sum *= Float32(self.cfg.kv_block_size) + tail_delta = tail_len - self.cfg.kv_block_size + if cutlass.const_expr(tail_delta != 0): + local_sum += Float32(tail_delta) * tail_p + return local_sum + @producer_work @cute.jit - def compute_p_fragment( + def compute_p_fragments( self, stage_info: StageInfo, *, - fragment_idx: Constexpr[int], new_max_arr: cutlass.Array, - s_arr: cutlass.Array, ) -> None: - """Convert one KV256 K32 score fragment and publish its TMEM P slice.""" + """Stream every ordinary K32 fragment from one rolled loop.""" + self._compute_p_fragments_impl( + stage_info, + new_max_arr=new_max_arr, + route_is_proxy=Int32(0), + route_origin0=Int32(0), + route_origin1=Int32(0), + ) + + @producer_work + @cute.jit + def compute_proxy_route_p_fragments( + self, + stage_info: StageInfo, + *, + new_max_arr: cutlass.Array, + route_flags: Int32, + route_origin0: Int32, + route_origin1: Int32, + ) -> None: + """Stream every proxy-capable KV256 K32 fragment from one rolled loop.""" + assert self.cfg.use_block_sparse_proxy_routes + route_is_proxy = Int32( + (route_flags & Int32(_SOFTMAX_ROUTE_IS_PROXY_FLAG)) != Int32(0) + ) + self._compute_p_fragments_impl( + stage_info, + new_max_arr=new_max_arr, + route_is_proxy=route_is_proxy, + route_origin0=route_origin0, + route_origin1=route_origin1, + ) + + @cute.jit + def _compute_p_fragments_impl( + self, + stage_info: StageInfo, + *, + new_max_arr: cutlass.Array, + route_is_proxy: Int32, + route_origin0: Int32, + route_origin1: Int32, + ) -> None: + """Reload, exponentiate, and publish all K32 fragments in a rolled loop. + + The fragment index is a runtime loop variable, so the exponentiation + body exists once in the instruction stream and only the TMEM column + offset, the fragment barrier, and the proxy tail bookkeeping depend on + it. Unrolling the fragments would replicate that body for every + fragment and both softmax instances and leave the softmax warps + instruction-fetch bound. The max pass has already written masked + scores back to TMEM, so the reload needs no mask logic of its own. + """ + _ = stage_info cfg = self.cfg assert cfg.streams_tmem_p_fragments - assert not cfg.use_fp8_qkv and cfg.uses_two_inst_tmem_p - assert cfg.softmax_score_fragment_regs == 32 + assert self._tmem_alloc.offset == self.tmem_s_ref._alloc.offset + # One FP32 score per column, two packed 16-bit probabilities per column. + fragment_regs = cfg.softmax_score_fragment_regs + fragment_cols = fragment_regs // 2 new_max = new_max_arr[0] safe_new_max = new_max if safe_new_max == _neg_max_f32(): safe_new_max = Float32(0.0) minus_max_scale = Float32(-self.scale_softmax_log2 * safe_new_max) + tmem_base = self._tmem_base_addr + Int32(self._tmem_alloc.offset) + tidx, _, _ = cute.arch.thread_idx() + publishes_fragment = (tidx & Int32(31)) == Int32(0) + + total_sum = Float32(0.0) + for fragment_idx in cutlass.range(cfg.num_softmax_score_fragments, unroll=1): + fragment = Int32(fragment_idx) + loaded = _keeps_tcgen05_ld( + cfg, + prims.make_tmem_ptr( + tmem_base + fragment * Int32(fragment_regs), Float32 + ), + num=fragment_regs, + offset=cfg.tile_size_kv // 2, + ) + prims.tcgen05_wait(kind=prims.Tcgen05Wait.LOAD) + s_arr = cutlass.Array( + Float32, fragment_regs, space=cutlass.AddressSpace.rmem + ) + for score_idx in cutlass.range_constexpr(fragment_regs): + s_arr[score_idx] = loaded[score_idx] + + local_sum = self._exponentiate_fragment_pairs(s_arr, minus_max_scale) + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + if route_is_proxy != Int32(0): + local_sum = self._proxy_fragment_sum( + local_sum, + s_arr, + fragment_origin=self._runtime_fragment_origin( + fragment, route_origin0, route_origin1 + ), + ) + + packed_p = ( + s_arr.data_ptr() + .load(count=fragment_regs, alignment=4) + .to(cfg.q_dtype) + .bitcast(Int32) + ) + _keeps_tcgen05_st( + cfg, + prims.make_tmem_ptr(tmem_base + fragment * Int32(fragment_cols), Int32), + packed_p, + offset=cfg.tmem_p_cols_per_inst, + ) + cute.arch.fence_view_async_tmem_store() + prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) + if publishes_fragment: + prims.mbarrier_arrive(self._fragment_ready.data_ptr() + fragment) + total_sum += local_sum + self.tmem_s_ref.store_p_local_sum(0, total_sum) + + @cute.jit + def _runtime_fragment_origin( + self, fragment: Int32, route_origin0: Int32, route_origin1: Int32 + ) -> Int32: + """Return the token origin of a fragment selected at runtime. + + Each lane's fragments cover two K64 route atoms in order: the first + atom's fragments start at ``route_origin0``, the second atom's at + ``route_origin1``, and consecutive fragments within an atom advance by + one fragment width. + """ + cfg = self.cfg + fragment_regs = cfg.softmax_score_fragment_regs + fragments_per_origin = cfg.softmax_fragments_per_route_atom + fragment_origin = Int32(route_origin0) + if fragment >= Int32(fragments_per_origin): + fragment_origin = Int32(route_origin1) + return fragment_origin + (fragment % Int32(fragments_per_origin)) * Int32( + fragment_regs + ) + + @cute.jit + def _exponentiate_fragment_pairs( + self, s_arr: cutlass.Array, minus_max_scale: Float32 + ) -> Float32: + """Turn one fragment of scaled scores into probabilities in place. - # Eight independent chains keep the denominator update off one long - # dependency chain. Reuse s_arr for probabilities so only one K32 score - # fragment remains live while P is packed. + Returns the fragment's probability sum. Eight independent chains keep + the denominator update off one long dependency chain, and a configurable + subset of pairs runs its exponentials on the FMA pipe. + """ + pairs_per_fragment = self.cfg.softmax_score_fragment_regs // 2 sum_chains = cutlass.Array(Float32, 8, space=cutlass.AddressSpace.rmem) for chain_idx in cutlass.range_constexpr(8): sum_chains[chain_idx] = Float32(0.0) - for pair_idx in cutlass.range_constexpr(16): + for pair_idx in cutlass.range_constexpr(pairs_per_fragment): value_idx = pair_idx * 2 p0, p1 = cute.arch.fma_packed_f32x2( (Float32(s_arr[value_idx]), Float32(s_arr[value_idx + 1])), (self.scale_softmax_log2, self.scale_softmax_log2), (minus_max_scale, minus_max_scale), ) - p0 = Float32(cute.math.exp2(p0, fastmath=True)) - p1 = Float32(cute.math.exp2(p1, fastmath=True)) + if cutlass.const_expr( + _pair_uses_ex2_emulation(pair_idx, pairs_per_fragment) + ): + p0, p1 = _ex2_emulation_packed_f32x2(p0, p1) + else: + p0 = Float32(cute.math.exp2(p0, fastmath=True)) + p1 = Float32(cute.math.exp2(p1, fastmath=True)) s_arr[value_idx] = p0 s_arr[value_idx + 1] = p1 chain_idx = (pair_idx & 3) * 2 @@ -345,10 +524,6 @@ def compute_p_fragment( (p0, p1), ) ) - - # Collapse the eight reduction chains before packing P and publishing - # its barrier. This keeps only one sum scalar live across STTM instead - # of overlapping the full reduction state with packed P and addresses. sum01 = cute.arch.add_packed_f32x2( (sum_chains[0], sum_chains[1]), (sum_chains[2], sum_chains[3]), @@ -358,40 +533,43 @@ def compute_p_fragment( (sum_chains[6], sum_chains[7]), ) total_pair = cute.arch.add_packed_f32x2(sum01, sum23) - local_sum = Float32(total_pair[0] + total_pair[1]) + return Float32(total_pair[0] + total_pair[1]) - packed_p = ( - s_arr.data_ptr().load(count=32, alignment=4).to(cfg.q_dtype).bitcast(Int32) - ) + @cute.jit + def _proxy_fragment_sum( + self, + local_sum: Float32, + s_arr: cutlass.Array, + *, + fragment_origin: Int32, + ) -> Float32: + """Weight a proxy fragment's sum by the token mass each summary stands for. - fragment_cols = cfg.softmax_score_fragment_regs // 2 - p_tmem_addr = ( - self._tmem_base_addr - + Int32(self._tmem_alloc.offset) - + Int32(fragment_idx * fragment_cols) - ) - _keeps_tcgen05_st( - cfg, - prims.make_tmem_ptr(p_tmem_addr, Int32), - packed_p, - offset=cfg.tmem_p_cols_per_inst, + KC stores one mean K vector per semantic KV block while VC stores its V + sum. P itself stays unweighted for PV; only the denominator accounts + for the represented token count, with the final summary covering the + shorter tail block. + """ + cfg = self.cfg + fragment_regs = cfg.softmax_score_fragment_regs + num_summaries, tail_len = _block_sparse_proxy_summary_geometry( + cfg.static_seq_len_kv, + cfg.kv_block_size, ) - # This lowers to the warp-collective tcgen05.wait::st. The explicit - # proxy fence then makes every lane's completed STTM visible through - # the lane-0 mbarrier publication consumed by the MMA warp. - cute.arch.fence_view_async_tmem_store() - prims.tcgen05_fence(prims.Tcgen05Fence.BEFORE_THREAD_SYNC) - - # KV256 aliases P with the score tile that produced it. Each softmax - # warp publishes its own rows after the TMEM store drains; BMM2 waits - # for all producer warps before consuming the fragment. - tidx, _, _ = cute.arch.thread_idx() - if (tidx & Int32(31)) == Int32(0): - prims.mbarrier_arrive(self._fragment_ready.data_ptr() + Int32(fragment_idx)) - - if cutlass.const_expr(fragment_idx != 0): - local_sum += self.tmem_s_ref.load_p_local_sum(0) - self.tmem_s_ref.store_p_local_sum(0, local_sum) + local_sum *= Float32(cfg.kv_block_size) + tail_delta = tail_len - cfg.kv_block_size + if cutlass.const_expr(tail_delta != 0): + final_summary_idx = num_summaries - 1 + final_summary_offset = Int32(final_summary_idx) - fragment_origin + if final_summary_offset >= Int32(0) and final_summary_offset < Int32( + fragment_regs + ): + # Proxy fragment origins are fragment-aligned in summary + # coordinates, so the tail's in-fragment lane is a compile-time + # constant even though route ownership is decided at runtime. + tail_lane = final_summary_idx % fragment_regs + local_sum += Float32(tail_delta) * Float32(s_arr[tail_lane]) + return local_sum @cute.jit def _compute_keeps_p( @@ -401,16 +579,18 @@ def _compute_keeps_p( new_max_arr: cutlass.Array, s_arr: cutlass.Array, ) -> None: - """Materialize one non-KV256 row-major Keeps probability tile. + """Materialize one complete-row Keeps probability tile. TQ128 gives each warp-group thread a complete 128-column row. TQ64 gives paired lanes the low/high 64-column halves of one row. Each lane writes disjoint packed blocks into the TMEM or SMEM layout consumed by - BMM2. + BMM2. Streamed profiles, including every block-sparse Keeps profile, + produce P through the rolled fragment loop instead. """ cfg = self.cfg - # KV256 uses compute_p_fragment so only one K32 score fragment is live. - assert not cfg.streams_tmem_p_fragments + # Every block-sparse Keeps profile streams P; only dense complete rows + # reach this path. + assert not cfg.streams_tmem_p_fragments and not cfg.use_block_sparse task_cache = _decode_gen_task_cache(stage_info) warp_grp_thread_idx = task_cache[_TASK_CACHE_WARP_GRP_THREAD_IDX] lane_idx = task_cache[_TASK_CACHE_LANE_IDX] @@ -439,7 +619,6 @@ def _compute_keeps_p( # without keeping a second 16-value P array live beside the S row. local_sum_pair_01 = (Float32(0.0), Float32(0.0)) local_sum_pair_23 = (Float32(0.0), Float32(0.0)) - # Each vector block is exactly 16 bytes after conversion. Compute and # pack adjacent pairs directly into their final register payload. packed_p_regs = cfg.num_packed_p_regs if cfg.uses_two_inst_tmem_p else 4 @@ -538,10 +717,10 @@ def _compute_keeps_p( packed_p.data_ptr().load(count=4, alignment=4), alignment=16 ) if cutlass.const_expr(cfg.uses_two_inst_tmem_p): - # FP8 publishes a complete row with one x16/x32 STTM. FP16/BF16 - # uses x16 slices to limit Softmax register pressure. This is the - # complete-row Q128/KV128 path; KV256 publishes K32 fragments. - assert cfg.num_packed_p_regs in (16, 32, 64) + # FP8 publishes the complete row with one x32 STTM. Dense 16-bit + # Q128/KV128 uses x16 slices to limit Softmax register pressure. + # Block-sparse two-instance profiles stream K32 fragments instead. + assert cfg.num_packed_p_regs in (32, 64) regs_per_store = cfg.num_packed_p_regs if cfg.use_fp8_qkv else 16 assert cfg.num_packed_p_regs % regs_per_store == 0 for store_idx in cutlass.range_constexpr( @@ -591,14 +770,18 @@ def _compute_keeps_p( # point and no extra named barrier is needed here. cute.arch.fence_view_async_shared() - @producer_work @cute.jit - def compute_p( + def _compute_p_impl( self, stage_info: StageInfo, *, new_max_arr: cutlass.Array, s_arr: cutlass.Array, + route_is_proxy: Int32, + route_origin0: Int32, + route_origin1: Int32, + route_origin2: Int32, + route_origin3: Int32, ) -> None: """Compute P from S, stage its BMM2 operand, and publish local sums.""" cfg = self.cfg @@ -616,6 +799,7 @@ def compute_p( # warp/lane ownership for SMEM offsets and STSM swizzles. task_cache = _decode_gen_task_cache(stage_info) warp_grp_thread_idx = task_cache[_TASK_CACHE_WARP_GRP_THREAD_IDX] + lane_idx = Int32(task_cache[_TASK_CACHE_LANE_IDX]) if cutlass.const_expr(cfg.tile_size_q == 32 and cfg.use_fp8_qkv): # Tile-Q=32 FP8 fast path: compute E4M3 P registers in the # same order consumed by the STSM helper, while also capturing @@ -796,10 +980,14 @@ def compute_p( local_sums = cutlass.Array( Float32, num_scale_groups, space=cutlass.AddressSpace.rmem ) + proxy_tail_p = cutlass.Array( + Float32, num_scale_groups, space=cutlass.AddressSpace.rmem + ) for idx in cutlass.range_constexpr(num_s_regs): p_vals[idx] = Float32(0.0) for idx in cutlass.range_constexpr(num_scale_groups): local_sums[idx] = Float32(0.0) + proxy_tail_p[idx] = Float32(0.0) for scale_idx in cutlass.range_constexpr(num_scale_groups): # Convert each softmax scale group from S to P. Masked rows have @@ -828,9 +1016,32 @@ def compute_p( ) p_vals[s_idx] = p_val local_sums[scale_idx] += p_val + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + num_summaries, _ = _block_sparse_proxy_summary_geometry( + cfg.static_seq_len_kv, + cfg.kv_block_size, + ) + atom_origin, logical_summary = _swaps_routed_coordinate( + cfg, + lane_idx >> Int32(2), + route_origin0, + route_origin1, + route_origin2, + route_origin3, + token_group_idx=k_pair_idx, + ) + if atom_origin >= Int32(0) and logical_summary == Int32( + num_summaries - 1 + ): + proxy_tail_p[scale_idx] = p_val # Hand off denominator contributions through TmemS. P remains a pure # MMA operand in SMEM; sums are not reloaded from the P tile. for scale_idx in cutlass.range_constexpr(num_scale_groups): + local_sums[scale_idx] = self._apply_proxy_route_denominator_mass( + local_sums[scale_idx], + proxy_tail_p[scale_idx], + route_is_proxy, + ) self.tmem_s_ref.store_p_local_sum(scale_idx, local_sums[scale_idx]) if cutlass.const_expr(cfg.use_fp8_qkv): @@ -1057,6 +1268,33 @@ def compute_p( p_vals[p_base + 4] = p_pair[1] local_sum[scale_idx] += p_pair[0] local_sum[scale_idx] += p_pair[1] + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + num_summaries, _ = _block_sparse_proxy_summary_geometry( + cfg.static_seq_len_kv, + cfg.kv_block_size, + ) + for scale_idx in cutlass.range_constexpr(cfg.num_softmax_scale_groups): + proxy_tail_p = Float32(0.0) + for token_group_idx in cutlass.range_constexpr(4): + atom_origin, logical_summary = _swaps_routed_coordinate( + cfg, + lane_idx >> Int32(2), + route_origin0, + route_origin1, + route_origin2, + route_origin3, + token_group_idx=token_group_idx, + ) + if atom_origin >= Int32(0) and logical_summary == Int32( + num_summaries - 1 + ): + p_idx = scale_idx + token_group_idx * 2 + proxy_tail_p = Float32(p_vals[p_idx]) + local_sum[scale_idx] = self._apply_proxy_route_denominator_mass( + local_sum[scale_idx], + proxy_tail_p, + route_is_proxy, + ) # Pack the P scalars to match the dtype consumed by BMM2. regs_p = cutlass.Array( Int32, cfg.num_packed_p_regs, space=cutlass.AddressSpace.rmem @@ -1107,6 +1345,77 @@ def compute_p( # BMM2 cannot observe a partially written P tile. prims.barrier_cta_sync(4 + self.inst_id, thread_count=128) + @producer_work + @cute.jit + def compute_p( + self, + stage_info: StageInfo, + *, + new_max_arr: cutlass.Array, + s_arr: cutlass.Array, + ) -> None: + """Compute an exact/dense P tile without typed-route metadata.""" + + self._compute_p_impl( + stage_info, + new_max_arr=new_max_arr, + s_arr=s_arr, + route_is_proxy=Int32(0), + route_origin0=Int32(0), + route_origin1=Int32(0), + route_origin2=Int32(0), + route_origin3=Int32(0), + ) + + @producer_work + @cute.jit + def compute_proxy_route_p( + self, + stage_info: StageInfo, + *, + new_max_arr: cutlass.Array, + s_arr: cutlass.Array, + route_origin0: Int32, + route_origin1: Int32, + keeps_route_flags_or_swaps_origin2: Int32, + swaps_route_origin3_bits: Uint32, + swaps_route_flags: Uint32, + ) -> None: + """Normalize the active Keeps/SWAP metadata view and compute P. + + The shared Int32 input is Keeps route flags or SWAP origin2. SWAP's + origin3 and flags stay bit-preserving Uint32 values until this work + boundary because schedule-level dataflow tokens cannot be cast. + """ + + assert self.cfg.use_block_sparse_proxy_routes + route_origin2 = Int32(0) + route_origin3 = Int32(0) + if cutlass.const_expr(self.cfg.use_keeps_mma_ab): + route_is_proxy = Int32( + ( + keeps_route_flags_or_swaps_origin2 + & Int32(_SOFTMAX_ROUTE_IS_PROXY_FLAG) + ) + != Int32(0) + ) + else: + route_is_proxy = Int32( + (swaps_route_flags & Uint32(_SOFTMAX_ROUTE_IS_PROXY_FLAG)) != Uint32(0) + ) + route_origin2 = keeps_route_flags_or_swaps_origin2 + route_origin3 = swaps_route_origin3_bits.bitcast(Int32) + self._compute_p_impl( + stage_info, + new_max_arr=new_max_arr, + s_arr=s_arr, + route_is_proxy=route_is_proxy, + route_origin0=route_origin0, + route_origin1=route_origin1, + route_origin2=route_origin2, + route_origin3=route_origin3, + ) + @consumer_work( returns=( p_desc_0_slot, @@ -1136,7 +1445,7 @@ def p_operands( # stats-free columns of the corresponding S stage. p_stage_cols = cfg.tmem_s_cols if cutlass.const_expr(cfg.streams_tmem_p_fragments): - # KV256's four pipeline stages are K32 fragments of one P + # A streamed profile's four pipeline stages are K32 fragments of one P # operand, not four independent full S/P stages. p_stage_cols = cfg.softmax_score_fragment_regs // 2 p_tmem_addr = self._tmem_base_addr + Int32( @@ -1170,7 +1479,7 @@ def wait_p_fragment( *, fragment_idx: Constexpr[int], ) -> Int32: - """Wait for and return the next KV256 P-fragment TMEM address.""" + """Wait for and return the next streamed P-fragment TMEM address.""" cfg = self.cfg _ = stage_info assert cfg.streams_tmem_p_fragments @@ -1185,25 +1494,3 @@ def wait_p_fragment( self._tmem_alloc.offset + fragment_idx * fragment_cols ) return p_tmem_addr - - @consumer_work(work_attrs=WorkAttr.AUXILIARY) - @cute.jit - def wait_until_reusable_before_qk(self, stage_info: StageInfo) -> None: - """Wait until the previous same-instance PV has stopped reading P. - - KV256 aliases each streamed P instance with its next S accumulator. - The existing two-stage O pipeline commits stage ``inst_id`` only when - the matching PV completes, so its full barrier is also the P-reuse - credit. The S producer phase supplies the generation: the first QK - waits on the initially complete opposite parity, and every later QK - waits for the preceding PV without another commit or barrier. - """ - _ = stage_info - cfg = self.cfg - assert cfg.streams_tmem_p_fragments - assert cfg.o_stages == cfg.num_insts_kv == 2 - barrier = self.tmem_o_ref.pipeline.sync_object_full.get_barrier( - Int32(self.inst_id) - ) - _wait_for_mbarrier_phase(barrier, self.tmem_s_ref.producer_state.phase) - prims.tcgen05_fence(prims.Tcgen05Fence.AFTER_THREAD_SYNC) diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_resources.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_resources.py index 20bb78fb810c..173a13cb5447 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_resources.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/smem_resources.py @@ -497,6 +497,10 @@ class SmemKvTileResource(DecodeGenResourceBase): tma_desc_v: cutlass.Pointer | None = None tma_desc_k_atom: cutlass.Pointer | None = None tma_desc_v_atom: cutlass.Pointer | None = None + tma_desc_k_summary: cutlass.Pointer | None = None + tma_desc_v_summary: cutlass.Pointer | None = None + tma_desc_k_summary_atom: cutlass.Pointer | None = None + tma_desc_v_summary_atom: cutlass.Pointer | None = None sparse_kv_metadata: "SmemBlockSparseKvMetadataResource | None" = None page_offsets_kv: "SmemPageOffsetsKvResource | None" = None seqlens_kv: cute.Pointer | None = None @@ -703,15 +707,24 @@ def _producer_load( assert self.sparse_kv_metadata is not None assert self.tma_desc_k_atom is not None assert self.tma_desc_v_atom is not None - # The positional TensorMaps keep the decode ABI stable. The - # primary K/V descriptors are KV128 for coarse routes and one atom - # for fine routes. The auxiliary slots always expose the atom - # descriptor and alias the primary descriptor for fine routes. + # K/V and summary sources expose the same primary/atom descriptor + # pair. Route kind selects the source; the geometry below alone + # selects the physical copy policy. tma_desc_atom = ( self.tma_desc_v_atom if cutlass.const_expr(self.kv_kind == KV_KIND_V) else self.tma_desc_k_atom ) + tma_desc_summary = ( + self.tma_desc_v_summary + if cutlass.const_expr(self.kv_kind == KV_KIND_V) + else self.tma_desc_k_summary + ) + tma_desc_summary_atom = ( + self.tma_desc_v_summary_atom + if cutlass.const_expr(self.kv_kind == KV_KIND_V) + else self.tma_desc_k_summary_atom + ) kv_atom_size = _block_sparse_kv_atom_size(cfg.kv_block_size) head_dim_stage = cfg.head_dim_kv_stage head_dim_stage_offset = head_dim_stage_idx * head_dim_stage @@ -775,6 +788,12 @@ def _producer_load( # join unrelated entries and must prove physical adjacency. fragment_chunk_elems = chunk_hd * 64 if prims.elect_sync(): + route_tma_desc = tma_desc + route_tma_desc_atom = tma_desc_atom + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + if self.sparse_kv_metadata.route_is_proxy(): + route_tma_desc = tma_desc_summary + route_tma_desc_atom = tma_desc_summary_atom origin0, _ = self.sparse_kv_metadata.route_tma_coordinate( Int32(0), logical_b_idx, @@ -800,7 +819,7 @@ def _producer_load( local_tile_offset = chunk_idx * tile_chunk_elems prims.cp_async_bulk_tensor_shared_cta_global( stage_base.subview(local_tile_offset), - tma_desc, + route_tma_desc, ( Int32(global_head_dim_offset), origin0, @@ -832,7 +851,7 @@ def _producer_load( if adjacent: prims.cp_async_bulk_tensor_shared_cta_global( stage_base.subview(local_tile_offset), - tma_desc, + route_tma_desc, ( Int32(global_head_dim_offset), origin0, @@ -844,7 +863,7 @@ def _producer_load( else: prims.cp_async_bulk_tensor_shared_cta_global( stage_base.subview(local_tile_offset), - tma_desc_atom, + route_tma_desc_atom, ( Int32(global_head_dim_offset), origin0, @@ -857,7 +876,7 @@ def _producer_load( stage_base.subview( local_tile_offset + fragment_chunk_elems ), - tma_desc_atom, + route_tma_desc_atom, ( Int32(global_head_dim_offset), origin1, @@ -875,6 +894,10 @@ def _producer_load( atom_chunk_elems = chunk_hd * kv_atom_size atoms_per_route = cfg.tile_size_kv // kv_atom_size if prims.elect_sync(): + route_tma_desc_atom = tma_desc_atom + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + if self.sparse_kv_metadata.route_is_proxy(): + route_tma_desc_atom = tma_desc_summary_atom stage_base = self._stage_base(stage_info) # Reuse each retained origin across all head-dimension # chunks. The copies still target disjoint SMEM regions @@ -913,7 +936,7 @@ def _producer_load( stage_base.subview( local_tile_offset + atom_idx * atom_chunk_elems ), - tma_desc_atom, + route_tma_desc_atom, ( Int32(global_head_dim_offset), origin, @@ -1487,6 +1510,10 @@ class SmemKvResource(DecodeGenResourceBase): tma_desc_v: cutlass.Pointer | None = None tma_desc_k_atom: cutlass.Pointer | None = None tma_desc_v_atom: cutlass.Pointer | None = None + tma_desc_k_summary: cutlass.Pointer | None = None + tma_desc_v_summary: cutlass.Pointer | None = None + tma_desc_k_summary_atom: cutlass.Pointer | None = None + tma_desc_v_summary_atom: cutlass.Pointer | None = None sparse_kv_metadata0: "SmemBlockSparseKvMetadataResource | None" = None sparse_kv_metadata1: "SmemBlockSparseKvMetadataResource | None" = None page_offsets_kv: SmemPageOffsetsKvResource | None = None @@ -1922,6 +1949,29 @@ def _producer_load_kv_tile_256( ): assert self.page_offsets_kv is not None dense_page_ids = self.page_offsets_kv.page_ids(grouped_tile_idx) + # Select the logical source before the constexpr 4 x 2 loop so K/V + # and summary routes retain one physical KV256 staging body. + route_tma_desc = tma_desc + if cutlass.const_expr(cfg.use_block_sparse): + sparse_tma_desc = ( + self.tma_desc_v_atom + if cutlass.const_expr(kv_kind == KV_KIND_V) + else self.tma_desc_k_atom + ) + assert sparse_tma_desc is not None + route_tma_desc = sparse_tma_desc + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + assert sparse_kv_metadata is not None + summary_tma_desc = ( + self.tma_desc_v_summary_atom + if cutlass.const_expr(kv_kind == KV_KIND_V) + else self.tma_desc_k_summary_atom + ) + assert summary_tma_desc is not None + route_is_proxy = sparse_kv_metadata.route_is_proxy() + route_tma_desc = ( + summary_tma_desc if route_is_proxy else sparse_tma_desc + ) for semantic_block in cutlass.range_constexpr(4): token_coord = Int32(0) storage_coord = logical_b_idx @@ -1950,15 +2000,9 @@ def _producer_load_kv_tile_256( ) if cutlass.const_expr(cfg.use_block_sparse): - sparse_tma_desc = ( - self.tma_desc_v_atom - if cutlass.const_expr(kv_kind == KV_KIND_V) - else self.tma_desc_k_atom - ) - assert sparse_tma_desc is not None prims.cp_async_bulk_tensor_shared_cta_global( stage_base.subview(block_base), - sparse_tma_desc, + route_tma_desc, ( Int32(dim_half * 64), token_coord, diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_corr.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_corr.py index c59712aa8f3c..11b384a103bd 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_corr.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_corr.py @@ -91,7 +91,9 @@ _KV_TILE_256_CORRECTION_THREADS = 128 _KV_TILE_256_LOGICAL_OUTPUT_ROWS = 64 -_KV_TILE_256_EXCHANGE_ROW_STRIDE = 132 +# One D32 fragment per logical output row, padded by four floats so adjacent +# rows fall on different bank groups. +_KV_TILE_256_EXCHANGE_FRAGMENT_STRIDE = 36 _KV_TILE_256_STATS_PER_THREAD = 4 @@ -157,10 +159,10 @@ def get_o_stage_dtype_bytes(self) -> int: ) def _kv_tile_256_exchange_entries(self) -> int: - """Return 128 lane-local stats plus 64 logical output rows.""" + """Return 128 lane-local stats plus one D32 fragment per output row.""" return ( _KV_TILE_256_CORRECTION_THREADS * _KV_TILE_256_STATS_PER_THREAD - + _KV_TILE_256_LOGICAL_OUTPUT_ROWS * _KV_TILE_256_EXCHANGE_ROW_STRIDE + + _KV_TILE_256_LOGICAL_OUTPUT_ROWS * _KV_TILE_256_EXCHANGE_FRAGMENT_STRIDE ) def _init_placeholder_state(self) -> None: @@ -306,22 +308,14 @@ def get_smem_requirements(self) -> list[SmemAllocation]: ) if self.cfg.tile_size_kv == 256 and self._kv_tile_256_exchange_alloc is None: # Tail correction exchanges all lane-local stats, then pipelines - # D32 fragments through 64 logical output rows. Upper lanes publish - # one spatial half while lower lanes retain the matching fragment - # in registers. The dependency graph places this scratch after the - # shared KV ring so it can reuse the dead storage. - payload_bytes = self._kv_tile_256_exchange_entries() * 4 - exchange_bytes = payload_bytes - if self.cfg.uses_rotating_kv256_exchange: - assert payload_bytes <= self.cfg.smem_kv_tile_bytes - # Runtime selects one compact payload inside this explicit - # full-ring alias envelope. The envelope keeps every dynamic - # pointer within a declared allocation while the actual live - # exchange remains only 35,840 B in one 64-KiB stage. - exchange_bytes = self.cfg.smem_kv_tile_bytes * self.cfg.kv_stages + # D32 fragments through 64 logical output rows one fragment at a + # time. Upper lanes publish one spatial half while lower lanes + # retain the matching fragment in registers. The buffer is + # dedicated, so the shared KV ring keeps streaming the next tile's + # routes while the tail runs. self._kv_tile_256_exchange_alloc = SmemAllocation( name=f"{self.name}_kvTile256Exchange", - size_bytes=exchange_bytes, + size_bytes=self._kv_tile_256_exchange_entries() * 4, alignment=16, ) allocs = [] @@ -874,60 +868,145 @@ def _fold_split_o_vec8( return output_vals, sum_val, new_max, new_max @cute.jit - def _store_final_o_vec8( + def _merge_kv_tile_256_peer_fragment( + self, + own_vals: cutlass.Array, + peer_vals, + first_col: Constexpr[int], + count: Constexpr[int], + ) -> cutlass.Array: + """Add the peer spatial half to ``count`` own columns from ``first_col``.""" + merged_vals = cutlass.Array( + Float32, + count, + space=cutlass.AddressSpace.rmem, + ) + for elem in cutlass.range_constexpr(0, count, 2): + value_idx = first_col + elem + merged = fadd2( + (own_vals[value_idx], own_vals[value_idx + 1]), + ( + Float32(peer_vals[value_idx]), + Float32(peer_vals[value_idx + 1]), + ), + ) + merged_vals[elem] = merged[0] + merged_vals[elem + 1] = merged[1] + return merged_vals + + @cute.jit + def _store_final_o_columns( self, final_o_dst, output_vals: cutlass.Array, norm_scale: Float32, + *, + count: Constexpr[int], + sector_aligned: cutlass.Boolean, ) -> None: - """Pack one contiguous 8-element output fragment to the final O dtype.""" + """Scale, pack, and store ``count`` contiguous final output columns. + + FP8 output packs four values per register and writes eight columns per + 8-byte store. 16-bit output packs pairs; sixteen columns fill one + 32-byte sector and go out as a single 256-bit store when the + destination is sector aligned, otherwise every eight columns use one + 16-byte store. Callers choose ``count`` per path: the KV256 tail owns + whole rows per lane and pays for half-written sectors, while the + split-KV reducers write eight-column fragments. + """ cfg = self.cfg + assert count % 8 == 0 if cutlass.const_expr(cfg.use_fp8_output): - final_pairs = cutlass.Array(Float32, 8, space=cutlass.AddressSpace.rmem) - for pair_idx in cutlass.range_constexpr(4): - val_base = pair_idx * 2 - pair = fmul2( - (norm_scale, norm_scale), - (output_vals[val_base], output_vals[val_base + 1]), + for chunk_idx in cutlass.range_constexpr(count // 8): + fp8_regs = self._pack_fp8_output_quads( + output_vals, norm_scale, chunk_idx * 8 ) - final_pairs[val_base] = pair[0] - final_pairs[val_base + 1] = pair[1] - final_fp8_regs = cutlass.Array(Int32, 2, space=cutlass.AddressSpace.rmem) - final_fp8_regs[0] = _pack_float4_to_fp8_e4m3( - final_pairs[0], - final_pairs[1], - final_pairs[2], - final_pairs[3], - ) - final_fp8_regs[1] = _pack_float4_to_fp8_e4m3( - final_pairs[4], - final_pairs[5], - final_pairs[6], - final_pairs[7], - ) - final_o_dst.store( - final_fp8_regs.data_ptr().load(count=2, alignment=4), - alignment=8, - ) - else: - final_regs = cutlass.Array(Int32, 4, space=cutlass.AddressSpace.rmem) - for reg_idx in cutlass.range_constexpr(4): - pair = fmul2( - (norm_scale, norm_scale), - ( - output_vals[reg_idx * 2], - output_vals[reg_idx * 2 + 1], - ), + (final_o_dst + Int32(chunk_idx * 2)).store( + fp8_regs.data_ptr().load(count=2, alignment=4), + alignment=8, ) - if cutlass.const_expr(cfg.use_bf16_output): - final_regs[reg_idx] = _pack_float2_to_bf16(pair[0], pair[1]) + else: + final_regs = self._pack_final_o_regs(output_vals, norm_scale, count) + if cutlass.const_expr(count == 16): + if sector_aligned: + final_o_dst.store( + final_regs.data_ptr().load(count=8, alignment=4), + alignment=32, + ) else: - final_regs[reg_idx] = _pack_float2_to_fp16(pair[0], pair[1]) - final_o_dst.store( - final_regs.data_ptr().load(count=4, alignment=4), + self._store_16bit_output_chunks(final_o_dst, final_regs, count) + else: + self._store_16bit_output_chunks(final_o_dst, final_regs, count) + + @cute.jit + def _store_16bit_output_chunks( + self, + final_o_dst, + final_regs: cutlass.Array, + count: Constexpr[int], + ) -> None: + """Store packed 16-bit output columns as 16-byte chunks of eight columns.""" + for chunk_idx in cutlass.range_constexpr(count // 8): + (final_o_dst + Int32(chunk_idx * 4)).store( + (final_regs.data_ptr() + Int32(chunk_idx * 4)).load( + count=4, alignment=4 + ), alignment=16, ) + @cute.jit + def _pack_fp8_output_quads( + self, + output_vals: cutlass.Array, + norm_scale: Float32, + first_col: Constexpr[int], + ) -> cutlass.Array: + """Scale eight output columns from ``first_col`` into two FP8 registers.""" + final_pairs = cutlass.Array(Float32, 8, space=cutlass.AddressSpace.rmem) + for pair_idx in cutlass.range_constexpr(4): + val_base = pair_idx * 2 + pair = fmul2( + (norm_scale, norm_scale), + ( + output_vals[first_col + val_base], + output_vals[first_col + val_base + 1], + ), + ) + final_pairs[val_base] = pair[0] + final_pairs[val_base + 1] = pair[1] + fp8_regs = cutlass.Array(Int32, 2, space=cutlass.AddressSpace.rmem) + fp8_regs[0] = _pack_float4_to_fp8_e4m3( + final_pairs[0], final_pairs[1], final_pairs[2], final_pairs[3] + ) + fp8_regs[1] = _pack_float4_to_fp8_e4m3( + final_pairs[4], final_pairs[5], final_pairs[6], final_pairs[7] + ) + return fp8_regs + + @cute.jit + def _pack_final_o_regs( + self, + output_vals: cutlass.Array, + norm_scale: Float32, + count: Constexpr[int], + ) -> cutlass.Array: + """Scale ``count`` output columns and pack them as 16-bit pairs.""" + cfg = self.cfg + final_regs = cutlass.Array(Int32, count // 2, space=cutlass.AddressSpace.rmem) + for reg_idx in cutlass.range_constexpr(count // 2): + pair = fmul2( + (norm_scale, norm_scale), + ( + output_vals[reg_idx * 2], + output_vals[reg_idx * 2 + 1], + ), + ) + if cutlass.const_expr(cfg.use_bf16_output): + final_regs[reg_idx] = _pack_float2_to_bf16(pair[0], pair[1]) + else: + final_regs[reg_idx] = _pack_float2_to_fp16(pair[0], pair[1]) + return final_regs + @cute.jit def _store_softmax_normalized_o_vec8( self, @@ -953,7 +1032,13 @@ def _store_softmax_normalized_o_vec8( mem_space=1, dtype=Int32, ) - self._store_final_o_vec8(final_o_dst, output_vals, norm_scale) + self._store_final_o_columns( + final_o_dst, + output_vals, + norm_scale, + count=8, + sector_aligned=cutlass.Boolean(False), + ) @cute.jit def _softmax_output_row_state( @@ -2274,24 +2359,6 @@ def _publish_and_reduce_cluster_swaps_partials( full_prefix=False, ) - @cute.jit - def _kv_tile_256_exchange_for_stage( - self, - stage_info: StageInfo, - scratch_stage: Int32 | None, - ) -> cutlass.Array: - """Return the fixed exchange or its dynamically selected KV stage.""" - if cutlass.const_expr(scratch_stage is None): - return self._kv_tile_256_exchange - return cutlass.Array( - stage_info.context.smem_base.data_ptr() - + self._kv_tile_256_exchange_alloc.offset - + scratch_stage * Int32(self.cfg.smem_kv_tile_bytes), - dtype=Float32, - shape=(self._kv_tile_256_exchange_entries(),), - addrspace=3, - ) - @cute.jit def _kv_tile_256_temporal_fragment( self, @@ -2329,6 +2396,99 @@ def _kv_tile_256_temporal_fragment( ) return cutlass.Vector.from_elements(combined, Float32) + @cute.jit + def _store_kv_tile_256_direct_fragment( + self, + own_vals: cutlass.Array, + peer_vals, + *, + fragment_col: Constexpr[int], + dst_row_base: Int64, + norm_scale: Float32, + valid_output_row: cutlass.Boolean, + o_is_32b_aligned: cutlass.Boolean, + ) -> None: + """Merge one D32 fragment with its peer half and write the final output. + + Sixteen columns go out per store so each lane writes one full 32-byte + sector. Adjacent lanes own adjacent rows, so 16-byte stores would leave + every sector half-written twice. + """ + cfg = self.cfg + for vector_pair in cutlass.range_constexpr(2): + pair_col = vector_pair * 16 + merged_vals = self._merge_kv_tile_256_peer_fragment( + own_vals, + peer_vals, + pair_col, + 16, + ) + if valid_output_row: + output_col = fragment_col + pair_col + dst_offset = dst_row_base + Int32(output_col * cfg.o_dtype_bytes) + final_o_dst = cutlass.inttoptr( + self.o_ptr.toint() + cutlass.Int64(dst_offset), + mem_space=1, + dtype=Int32, + ) + self._store_final_o_columns( + final_o_dst, + merged_vals, + norm_scale, + count=16, + sector_aligned=o_is_32b_aligned, + ) + + @cute.jit + def _store_kv_tile_256_partial_fragment( + self, + own_vals: cutlass.Array, + peer_vals, + *, + fragment_col: Constexpr[int], + partial_row_base: Int64, + partial_scale: Float32, + valid_output_row: cutlass.Boolean, + ) -> None: + """Merge one D32 fragment with its peer half and write the split-KV partial.""" + cfg = self.cfg + partial_o_uses_bf16 = ( + cfg.use_bf16_separate_partial_o + if cfg.use_separate_reduction_kernel + else cfg.use_bf16_output + ) + for vector_idx in cutlass.range_constexpr(4): + vector_col = vector_idx * 8 + output_vals = self._merge_kv_tile_256_peer_fragment( + own_vals, + peer_vals, + vector_col, + 8, + ) + if valid_output_row: + output_col = fragment_col + vector_col + scaled_values: tuple = () + for elem in cutlass.range_constexpr(0, 8, 2): + scaled_values += fmul2( + (partial_scale, partial_scale), + (output_vals[elem], output_vals[elem + 1]), + ) + scaled_vector = cutlass.Vector.from_elements(scaled_values, Float32) + if cutlass.const_expr(partial_o_uses_bf16): + packed = scaled_vector.to(cutlass.BFloat16).bitcast(Int32) + else: + packed = scaled_vector.to(cutlass.Float16).bitcast(Int32) + # Split-KV partials are 16-bit, so the column offset follows + # the partial element width + partial_o_dst = cutlass.inttoptr( + self.partial_o_ptr.toint() + + partial_row_base + + Int64(output_col * 2), + mem_space=1, + dtype=Int32, + ) + partial_o_dst.store(packed, alignment=16) + @cute.jit def _kv_tile_256_merge_spatial_output( self, @@ -2357,18 +2517,13 @@ def _kv_tile_256_merge_spatial_output( serializing the complete D128 upper and lower halves. """ cfg = self.cfg - partial_o_uses_bf16 = ( - cfg.use_bf16_separate_partial_o - if cfg.use_separate_reduction_kernel - else cfg.use_bf16_output - ) output_exchange_base = Int32( _KV_TILE_256_CORRECTION_THREADS * _KV_TILE_256_STATS_PER_THREAD ) output_lane = exchange_idx < Int32(_KV_TILE_256_LOGICAL_OUTPUT_ROWS) exchange_row_idx = exchange_idx & Int32(_KV_TILE_256_LOGICAL_OUTPUT_ROWS - 1) output_exchange_row_base = output_exchange_base + exchange_row_idx * Int32( - _KV_TILE_256_EXCHANGE_ROW_STRIDE + _KV_TILE_256_EXCHANGE_FRAGMENT_STRIDE ) logical_output_row_idx = q_row_offset + exchange_row_idx valid_output_row = cutlass.Boolean(False) @@ -2394,6 +2549,9 @@ def _kv_tile_256_merge_spatial_output( else: dst_row_base = Int64(0) norm_scale = Float32(1.0) + # Row strides are multiples of 32 bytes for D128, so sector-wide + # stores are legal exactly when the output base pointer is. + o_is_32b_aligned = (self.o_ptr.toint() & Int64(31)) == Int64(0) if output_lane: valid_output_row = _q_row_is_valid_for_seq( cfg, @@ -2419,10 +2577,17 @@ def _kv_tile_256_merge_spatial_output( weight00=weight00, weight10=weight10, ) + if cutlass.const_expr(fragment != 0): + # The single fragment buffer is reused: lower lanes must have + # consumed the previous peer fragment before it is overwritten. + prims.barrier_cta_sync( + self.store_barrier_id, + thread_count=cfg.correction_barrier_threads, + ) if exchange_idx >= Int32(_KV_TILE_256_LOGICAL_OUTPUT_ROWS): - ( - exchange.data_ptr() + output_exchange_row_base + Int32(fragment_col) - ).store(own_vals, alignment=16) + (exchange.data_ptr() + output_exchange_row_base).store( + own_vals, alignment=16 + ) # Lower lanes keep ``own_vals`` live across the barrier. Once every # lane arrives, upper lanes may prepare the next fragment while @@ -2433,69 +2598,28 @@ def _kv_tile_256_merge_spatial_output( ) if output_lane: - peer_vals = ( - exchange.data_ptr() + output_exchange_row_base + Int32(fragment_col) - ).load(count=32, alignment=16) - for vector_idx in cutlass.range_constexpr(4): - vector_col = vector_idx * 8 - output_vals = cutlass.Array( - Float32, - 8, - space=cutlass.AddressSpace.rmem, + peer_vals = (exchange.data_ptr() + output_exchange_row_base).load( + count=32, alignment=16 + ) + if cutlass.const_expr(not cfg.use_split_kv): + self._store_kv_tile_256_direct_fragment( + own_vals, + peer_vals, + fragment_col=fragment_col, + dst_row_base=dst_row_base, + norm_scale=norm_scale, + valid_output_row=valid_output_row, + o_is_32b_aligned=o_is_32b_aligned, + ) + else: + self._store_kv_tile_256_partial_fragment( + own_vals, + peer_vals, + fragment_col=fragment_col, + partial_row_base=partial_row_base, + partial_scale=partial_scale, + valid_output_row=valid_output_row, ) - for elem in cutlass.range_constexpr(0, 8, 2): - value_idx = vector_col + elem - merged = fadd2( - (own_vals[value_idx], own_vals[value_idx + 1]), - ( - Float32(peer_vals[value_idx]), - Float32(peer_vals[value_idx + 1]), - ), - ) - output_vals[elem] = merged[0] - output_vals[elem + 1] = merged[1] - if valid_output_row: - output_col = fragment_col + vector_col - if cutlass.const_expr(cfg.use_split_kv): - scaled_values: tuple = () - for elem in cutlass.range_constexpr(0, 8, 2): - scaled_values += fmul2( - (partial_scale, partial_scale), - (output_vals[elem], output_vals[elem + 1]), - ) - scaled_vector = cutlass.Vector.from_elements( - scaled_values, Float32 - ) - if cutlass.const_expr(partial_o_uses_bf16): - packed = scaled_vector.to(cutlass.BFloat16).bitcast( - Int32 - ) - else: - packed = scaled_vector.to(cutlass.Float16).bitcast( - Int32 - ) - partial_o_dst = cutlass.inttoptr( - self.partial_o_ptr.toint() - + partial_row_base - + Int64(output_col * cfg.o_dtype_bytes), - mem_space=1, - dtype=Int32, - ) - partial_o_dst.store(packed, alignment=16) - else: - dst_offset = dst_row_base + Int32( - output_col * cfg.o_dtype_bytes - ) - final_o_dst = cutlass.inttoptr( - self.o_ptr.toint() + cutlass.Int64(dst_offset), - mem_space=1, - dtype=Int32, - ) - self._store_final_o_vec8( - final_o_dst, - output_vals, - norm_scale, - ) if cutlass.const_expr(cfg.use_split_kv): if valid_output_row: @@ -2538,7 +2662,6 @@ def _kv_tile_256_tail_epilogue( self, stage_info: StageInfo, *, - scratch_stage: Int32 | None, tail_o_stage_idx_0: Int32, tail_o_stage_idx_1: Int32, inst0_new_max_arr: cutlass.Array, @@ -2553,15 +2676,13 @@ def _kv_tile_256_tail_epilogue( The standard decode schedule still owns the two temporal instances. KV256 adds one physical spatial split per instance. Correction exchanges - their stats, stages one spatial half in SMEM after the shared KV ring is - dead, then publishes the ordinary logical Q64xD128 output. + their stats, stages one spatial half through its dedicated SMEM + exchange one D32 fragment at a time, then publishes the ordinary + logical Q64xD128 output. """ cfg = self.cfg assert cfg.headdim == 128 - exchange = self._kv_tile_256_exchange_for_stage( - stage_info, - scratch_stage, - ) + exchange = self._kv_tile_256_exchange exchange_idx = warp_grp_thread_idx peer_idx = exchange_idx ^ Int32(_KV_TILE_256_LOGICAL_OUTPUT_ROWS) @@ -4371,7 +4492,6 @@ def _correction_tail_epilogue_impl( self, stage_info: StageInfo, *, - scratch_stage: Int32 | None, o_stage_idx: Int32, tail_o_stage_idx_0: Int32, tail_o_stage_idx_1: Int32, @@ -4417,7 +4537,6 @@ def _correction_tail_epilogue_impl( if cutlass.const_expr(cfg.tile_size_kv == 256): self._kv_tile_256_tail_epilogue( stage_info, - scratch_stage=scratch_stage, tail_o_stage_idx_0=tail_o_stage_idx_0, tail_o_stage_idx_1=tail_o_stage_idx_1, inst0_new_max_arr=inst0_new_max_arr, @@ -4493,9 +4612,6 @@ def _correction_tail_epilogue_impl( ) return - # Task Scheduling routes every non-constexpr work argument as a required - # data-flow token. Keep separate fixed/rotating entry points so only the - # latter consumes ``scratch_stage``; both still share the implementation. @producer_work @cute.jit def correction_tail_epilogue( @@ -4512,42 +4628,9 @@ def correction_tail_epilogue( inst1_new_max_arr: cutlass.Array, inst1_sum_arr: cutlass.Array, ) -> None: - """Run the ordinary fixed-exchange tail epilogue.""" - self._correction_tail_epilogue_impl( - stage_info, - scratch_stage=None, - o_stage_idx=o_stage_idx, - tail_o_stage_idx_0=tail_o_stage_idx_0, - tail_o_stage_idx_1=tail_o_stage_idx_1, - old_max_arr=old_max_arr, - new_max_arr=new_max_arr, - inst0_new_max_arr=inst0_new_max_arr, - inst0_sum_arr=inst0_sum_arr, - inst1_new_max_arr=inst1_new_max_arr, - inst1_sum_arr=inst1_sum_arr, - ) - - @producer_work - @cute.jit - def correction_tail_epilogue_rotating_exchange( - self, - stage_info: StageInfo, - *, - scratch_stage: Int32, - o_stage_idx: Int32, - tail_o_stage_idx_0: Int32, - tail_o_stage_idx_1: Int32, - old_max_arr: cutlass.Array, - new_max_arr: cutlass.Array, - inst0_new_max_arr: cutlass.Array, - inst0_sum_arr: cutlass.Array, - inst1_new_max_arr: cutlass.Array, - inst1_sum_arr: cutlass.Array, - ) -> None: - """Run persistent direct output in the stage named by its credit.""" + """Normalize the final O stages and publish the output tile.""" self._correction_tail_epilogue_impl( stage_info, - scratch_stage=scratch_stage, o_stage_idx=o_stage_idx, tail_o_stage_idx_0=tail_o_stage_idx_0, tail_o_stage_idx_1=tail_o_stage_idx_1, diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_o.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_o.py index fcf2cbe71730..fa8928695a98 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_o.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_o.py @@ -59,7 +59,7 @@ def _pv_mma_operand_contract_for_config( cfg.headdim if cfg.head_dim_per_stage_kv == 0 else cfg.head_dim_kv_stage ) if cfg.use_keeps_mma_ab: - if cfg.tile_size_kv == 256: + if cfg.uses_ws_2x2_datapath: # The WS 2x2 PV instruction exposes two spatial D128 partials as # one physical KV256 operation. Correction merges those spatial # halves after the two temporal decode streams are complete. @@ -196,7 +196,7 @@ def vp_mma_loop_fragment( p_tmem_addr: Int32, fragment_idx: Constexpr[int], ) -> None: - """Issue one K32 fragment of a KV256 loop PV tile.""" + """Issue one K32 fragment of a streamed loop PV tile.""" self._vp_mma_fragment( stage_info, v_desc=v_desc, @@ -215,7 +215,7 @@ def vp_mma_tail_fragment( p_tmem_addr: Int32, fragment_idx: Constexpr[int], ) -> None: - """Issue one K32 fragment of the final KV256 PV tile.""" + """Issue one K32 fragment of the final streamed PV tile.""" self._vp_mma_fragment( stage_info, v_desc=v_desc, @@ -234,14 +234,16 @@ def _vp_mma_fragment( fragment_idx: Constexpr[int], initial_scale_d, ) -> None: - """Issue the two WS MMA steps covered by one KV256 P fragment. + """Issue the two MMA K-steps covered by one streamed P fragment. ``p_tmem_addr`` is already the base of the fragment selected by ``wait_p_fragment``. Only the two local K-step offsets are added here; - ``fragment_idx`` must not be applied to the TMEM address again. + ``fragment_idx`` must not be applied to the TMEM address again. KV256 + issues the WS 2x2 instruction over its two spatial halves; KV128 issues + the plain M=128 instruction and advances V by one K16 slice per step. """ cfg = self.cfg - assert cfg.tile_size_kv == 256 and cfg.uses_two_inst_tmem_p + assert cfg.streams_tmem_p_fragments v_desc = _freeze_smem_descriptor(v_desc) task_cache = _decode_gen_task_cache(stage_info) @@ -268,17 +270,32 @@ def _vp_mma_fragment( p_operand = prims.make_tmem_ptr( p_tmem_addr + Int32(local_k_step * 8), Int32 ) - iter_v_desc = v_desc + Int32( - (k_step // 4) * cfg.headdim * 16 + (k_step % 4) * 128 - ) - tcgen05_mma_ws( - _mma_kind_for_qkv(cfg), - tmem_col, - p_operand, - iter_v_desc, - idesc, - initial_scale_d or fragment_idx != 0 or local_k_step != 0, - ) + scale_d = initial_scale_d or fragment_idx != 0 or local_k_step != 0 + if cutlass.const_expr(cfg.uses_ws_2x2_datapath): + # V holds four K64 atoms; jump between atoms every four + # K16 steps. + iter_v_desc = v_desc + Int32( + (k_step // 4) * cfg.headdim * 16 + (k_step % 4) * 128 + ) + tcgen05_mma_ws( + _mma_kind_for_qkv(cfg), + tmem_col, + p_operand, + iter_v_desc, + idesc, + scale_d, + ) + else: + iter_v_desc = v_desc + Int32(k_step * 128) + prims.tcgen05_mma( + _mma_kind_for_qkv(cfg), + prims.CTAGroup.CTA_1, + tmem_col, + p_operand, + iter_v_desc, + idesc, + scale_d, + ) @cute.jit def _vp_mma( diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_s.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_s.py index 703134a5b61d..b010dfdb9d21 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_s.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_resources/tmem_s.py @@ -41,10 +41,9 @@ producer_work, ) -from ...._block_sparse.common import _MAX_KV_ATOM_SIZE from ...._block_sparse.prepared import _PREPARED_ROUTE_IS_FULL_FLAG from ..fmha_decode_config import CAUSAL, FmhaDecodeConfig -from ..fmha_decode_constants import KV_TILE_256_RESCALE_THRESHOLD_LOG2 +from ..fmha_decode_constants import SOFTMAX_RESCALE_THRESHOLD_LOG2 from ...tcgen05_compat import tcgen05_mma_ws from ...placeholder_helpers import ( _placeholder_local_array, @@ -76,13 +75,14 @@ _mma_kind_for_qkv, _neg_max_f32, _softmax_scale_pair_width, + _swaps_routed_coordinate, _q_row_is_valid_for_seq, _q_row_token_and_local_head, _q_group_token_base, _softmax_tile_idx, ) from .smem_block_sparse_metadata import ( - _SOFTMAX_TOKEN_MASK_IS_FULL_FLAG, + _SOFTMAX_ROUTE_IS_PROXY_FLAG, _swaps_forwards_packed_route_full, ) from .helpers_kv_tile_idx import ( @@ -103,20 +103,13 @@ _wspro_reduce_max4, ) -# A block-sparse route often changes the exact row maximum without changing it -# enough to justify rescaling the live O tile. Keeping the prior anchor within -# this bound makes the correction scale exactly one and bounds FP16/BF16 P by -# 2**8. As in the FlashInfer/TRT-LLM policy, this assumes normal model logits -# rather than adversarial values outside the qualified probability bound. -_BLOCK_SPARSE_RESCALE_THRESHOLD_LOG2 = 8.0 - def _swaps_uses_origin0_k32_full_guard(cfg: FmhaDecodeConfig) -> bool: """Whether one staged origin can prove this warp's K32 slice valid.""" return ( cfg.kv_block_size >= 32 - and not cfg.use_kv_valid_bits + and not cfg.uses_prepared_score_keep_words and not cfg.uses_uniform_causal_mask and not cfg.uses_per_row_causal_mask ) @@ -126,7 +119,7 @@ def _swaps_token_word_covers_kv_tail(cfg: FmhaDecodeConfig) -> bool: """Whether SWAP's prepared token word covers the logical KV tail.""" return ( - cfg.use_kv_valid_bits + cfg.uses_prepared_score_keep_words and not cfg.uses_uniform_causal_mask and not cfg.uses_per_row_causal_mask ) @@ -145,42 +138,34 @@ def _swaps_uses_token_only_score_validity(cfg: FmhaDecodeConfig) -> bool: @cute.jit -def _can_skip_sparse_keeps_structural_mask( - q_row_is_valid: Boolean, - origin0: Int32, - origin1: Int32, - valid0: Int32, - valid1: Int32, - seq_len_kv: Int32, - causal_end: Int32, +def _dense_fragment_keep_word( + rows_are_active: Boolean, + visible_start: Int32, + visible_end: Int32, *, - apply_causal_mask: cutlass.Constexpr[bool], -) -> Boolean: - """Return whether one Keeps row needs no Q/tail/causal predicate. + fragment_regs: cutlass.Constexpr[int], +) -> Uint32: + """Return the keep word of one dense K32 fragment. - Token-bit masking is independent. Comparing against the last complete - KV64 origin avoids overflowing an origin near the Int32 upper bound. + ``visible_start`` and ``visible_end`` are the visible token range relative + to the fragment's first column. Columns outside ``[start, end)`` are + masked; an inactive tile or Q row masks the whole fragment. """ - - fragment_size = Int32(_MAX_KV_ATOM_SIZE) - last_complete_origin = seq_len_kv - fragment_size - can_skip = Boolean( - q_row_is_valid - and valid0 != Int32(0) - and valid1 != Int32(0) - and origin0 <= last_complete_origin - and origin1 <= last_complete_origin - ) - if cutlass.const_expr(apply_causal_mask): - last_causal_origin = causal_end - fragment_size - can_skip = Boolean( - can_skip and origin0 <= last_causal_origin and origin1 <= last_causal_origin - ) - return can_skip + keep_word = Uint32(0) + if rows_are_active: + first_kept = cute.math.max(visible_start, Int32(0)) + end_kept = cute.math.min(visible_end, Int32(fragment_regs)) + if first_kept < end_kept: + # Both shift amounts stay strictly below the register width: + # 1 <= end_kept <= fragment_regs and 0 <= first_kept < end_kept. + keep_word = (Uint32(0xFFFFFFFF) >> (Int32(fragment_regs) - end_kept)) & ( + Uint32(0xFFFFFFFF) << first_kept + ) + return keep_word @cute.jit -def _sparse_k32_effective_keep_word( +def _sparse_effective_keep_word( q_row_is_valid: Boolean, fragment_origin: Int32, fragment_valid: Int32, @@ -667,7 +652,7 @@ def _qk_mma( a_desc, b_desc = q_desc, k_desc else: a_desc, b_desc = k_desc, q_desc - if cutlass.const_expr(cfg.tile_size_kv == 256): + if cutlass.const_expr(cfg.uses_ws_2x2_datapath): tcgen05_mma_ws( _mma_kind_for_qkv(cfg), tmem_col, @@ -842,6 +827,7 @@ def _resolve_keeps_tile_context(self, stage_info: StageInfo): is_valid_effective_tile, is_masked_final_wave, tile_is_unmasked, + tile_has_valid_scores, ) @cute.jit @@ -896,49 +882,32 @@ def _publish_keeps_softmax_state( ) -> None: """Publish a masked Keeps row and its updated softmax anchor.""" - new_anchor = cute.math.max(old_max, tile_max, ftz=True) - if cutlass.const_expr( - self.cfg.use_block_sparse and _BLOCK_SPARSE_RESCALE_THRESHOLD_LOG2 > 0.0 - ): - # Online softmax only requires a common finite reference for P, - # sum, and O; it does not require the exact row maximum. Defer a - # small anchor increase so correction can skip a TMEM O rescale. - rescale_log2 = (old_max - new_anchor) * self.scale_softmax_log2 - if (old_max != _neg_max_f32()) and ( - rescale_log2 >= Float32(-_BLOCK_SPARSE_RESCALE_THRESHOLD_LOG2) - ): - new_anchor = old_max old_max_arr[0] = old_max sum_arr[0] = running_sum - new_max_arr[0] = new_anchor + new_max_arr[0] = self._softmax_anchor(old_max, tile_max) for reg_idx in cutlass.range_constexpr(self.cfg.num_s_regs_per_thread): s_arr[reg_idx] = s_vals[reg_idx] @cute.jit - def _mask_and_store_sparse_keeps_atom( - self, - s_vals: cutlass.Array, - loaded: cutlass.Vector, - token_word: Uint32, - *, - atom_col: Constexpr[int], - token_mask_is_required: cutlass.Boolean, - ) -> None: - """Store one 32-score atom, applying its token word when required.""" - - if token_mask_is_required: - for atom_reg_idx in cutlass.range_constexpr(32): - score_idx = atom_col + atom_reg_idx - s_vals[score_idx] = loaded[atom_reg_idx] - token_bit_is_valid = ( - (token_word >> Int32(atom_reg_idx)) & Uint32(1) - ) != Uint32(0) - if not token_bit_is_valid: - s_vals[score_idx] = _neg_max_f32() - else: - for atom_reg_idx in cutlass.range_constexpr(32): - score_idx = atom_col + atom_reg_idx - s_vals[score_idx] = loaded[atom_reg_idx] + def _softmax_anchor(self, old_max: Float32, tile_max: Float32) -> Float32: + """Return the exponent reference max for the tile's P pass. + + Online softmax only requires a common finite reference for P, the + running sum, and O; it does not require the exact row maximum. + Profiles that defer anchor updates keep the previous reference while + the tile raises it by less than ``SOFTMAX_RESCALE_THRESHOLD_LOG2`` + log2 units, so correction can skip the in-place TMEM O rescale. The + 16-bit P path represents the bounded values above one, and the + numerator and denominator stay in the same scale frame. Larger jumps + still rebase to keep P comfortably in range. + """ + new_max = cute.math.max(old_max, tile_max, ftz=True) + if cutlass.const_expr(self.cfg.defers_softmax_anchor_updates): + if old_max != _neg_max_f32(): + max_delta_log2 = self.scale_softmax_log2 * (old_max - new_max) + if max_delta_log2 >= Float32(-SOFTMAX_RESCALE_THRESHOLD_LOG2): + new_max = old_max + return new_max @cute.jit def _load_keeps_fragment_impl( @@ -954,9 +923,8 @@ def _load_keeps_fragment_impl( is_masked_final_wave: cutlass.Boolean, *, apply_boundary_mask: Constexpr[bool], - fragment_idx: Constexpr[int] = 0, ) -> None: - """Load one Keeps score fragment with a compile-time mask policy. + """Load one complete-row Keeps score tile with a compile-time mask policy. The caller chooses the masked/unmasked path before TMEM load. Keeping the score fragment out of the branch condition avoids carrying 64/128 live @@ -966,8 +934,7 @@ def _load_keeps_fragment_impl( """ cfg = self.cfg task_cache = _decode_gen_task_cache(stage_info) - num_s_regs = cfg.softmax_score_fragment_regs - fragment_reg_base = fragment_idx * num_s_regs + num_s_regs = cfg.num_s_regs_per_thread base_addr = ( task_cache[_TASK_CACHE_TMEM_BASE_OFFSET] + Int32(self._alloc.offset) @@ -977,9 +944,7 @@ def _load_keeps_fragment_impl( atom_col = load_atom_idx * 32 loaded = _keeps_tcgen05_ld( cfg, - prims.make_tmem_ptr( - base_addr + Int32(fragment_reg_base + atom_col), Float32 - ), + prims.make_tmem_ptr(base_addr + Int32(atom_col), Float32), num=32, offset=cfg.tile_size_kv // 2, ) @@ -1023,7 +988,7 @@ def _load_keeps_fragment_impl( token_idx = tile_offset_k + _keeps_score_col( cfg, warp_grp_thread_idx, - fragment_reg_base + reg_idx, + reg_idx, col_base, ) if token_idx >= element_mask_end_idx: @@ -1051,7 +1016,7 @@ def _load_keeps_fragment_impl( score_col = _keeps_score_col( cfg, warp_grp_thread_idx, - fragment_reg_base + reg_idx, + reg_idx, col_base, ) if cutlass.const_expr(cfg.use_sliding_window_causal): @@ -1084,8 +1049,6 @@ def _load_keeps_fragment( is_valid_effective_tile: cutlass.Boolean, is_masked_final_wave: cutlass.Boolean, tile_is_unmasked: cutlass.Boolean, - *, - fragment_idx: Constexpr[int] = 0, ) -> None: """Select the masked or unmasked fragment loader before LDTM. @@ -1106,7 +1069,6 @@ def _load_keeps_fragment( is_valid_effective_tile, is_masked_final_wave, apply_boundary_mask=False, - fragment_idx=fragment_idx, ) else: self._load_keeps_fragment_impl( @@ -1120,275 +1082,8 @@ def _load_keeps_fragment( is_valid_effective_tile, is_masked_final_wave, apply_boundary_mask=True, - fragment_idx=fragment_idx, ) - @cute.jit - def _reduce_keeps_fragment_max(self, s_vals: cutlass.Array) -> Float32: - """Reduce the row maximum of a previously loaded Keeps fragment.""" - cfg = self.cfg - num_s_regs = cfg.softmax_score_fragment_regs - - max_chains = cutlass.Array(Float32, 4, space=cutlass.AddressSpace.rmem) - for chain_idx in cutlass.range_constexpr(4): - max_chains[chain_idx] = _neg_max_f32() - for reg_base in cutlass.range_constexpr(0, num_s_regs, 4): - for chain_idx in cutlass.range_constexpr(4): - max_chains[chain_idx] = cute.math.max( - max_chains[chain_idx], - s_vals[reg_base + chain_idx], - ftz=True, - ) - tile_max = cute.math.max( - cute.math.max(max_chains[0], max_chains[1], ftz=True), - cute.math.max(max_chains[2], max_chains[3], ftz=True), - ftz=True, - ) - if cutlass.const_expr(cfg.tile_size_q == 64 and cfg.tile_size_kv != 256): - return cute.math.max( - tile_max, - Float32( - prims.shfl_sync( - thread_mask=0xFFFFFFFF, - val=tile_max, - offset=16, - mask_and_clamp=0x1F, - kind=prims.Shfl.BFLY, - ) - ), - ftz=True, - ) - return tile_max - - @cute.jit - def _decode_sparse_mask_metadata( - self, - routed_origin0: Int32, - routed_origin1: Int32, - routed_route_flags: Int32, - routed_token_word0: Uint32, - routed_token_word1: Uint32, - routed_token_word2: Uint32, - routed_token_word3: Uint32, - ) -> tuple[Int32, Int32, Int32, Int32, cutlass.Array, cutlass.Boolean]: - """Decode one prepared, register-routed mask payload.""" - - origin0 = Int32(routed_origin0) - origin1 = Int32(routed_origin1) - route_flags = Int32(routed_route_flags) - valid0 = route_flags & Int32(1) - valid1 = (route_flags >> Int32(1)) & Int32(1) - route_token_mask_is_full = cutlass.Boolean(False) - if cutlass.const_expr(self.cfg.use_kv_valid_bits): - route_token_mask_is_full = cutlass.Boolean( - (route_flags & Int32(_SOFTMAX_TOKEN_MASK_IS_FULL_FLAG)) != Int32(0) - ) - - num_local_words = 4 if self.cfg.tile_size_q == 128 else 2 - local_token_words = cutlass.Array( - Uint32, - num_local_words, - space=cutlass.AddressSpace.rmem, - ) - for word_idx in cutlass.range_constexpr(num_local_words): - local_token_words[word_idx] = Uint32(0xFFFFFFFF) - if cutlass.const_expr(self.cfg.use_kv_valid_bits): - if not route_token_mask_is_full: - if cutlass.const_expr(self.cfg.tile_size_q == 128): - local_token_words[0] = Uint32(routed_token_word0) - local_token_words[1] = Uint32(routed_token_word1) - local_token_words[2] = Uint32(routed_token_word2) - local_token_words[3] = Uint32(routed_token_word3) - else: - local_word0 = Uint32(routed_token_word0) - local_word1 = Uint32(routed_token_word1) - local_token_words[0] = local_word0 - local_token_words[1] = local_word1 - return ( - origin0, - origin1, - valid0, - valid1, - local_token_words, - route_token_mask_is_full, - ) - - @cute.jit - def _compute_softmax_loop_sparse_keeps( - self, - stage_info: StageInfo, - *, - old_max_arr: cutlass.Array, - sum_arr: cutlass.Array, - new_max_arr: cutlass.Array, - s_arr: cutlass.Array, - routed_origin0: Int32, - routed_origin1: Int32, - routed_route_flags: Int32, - routed_token_word0: Uint32, - routed_token_word1: Uint32, - routed_token_word2: Uint32, - routed_token_word3: Uint32, - ) -> tuple[object, object, object, object]: - """Load Keeps scores and mask them in logical KV coordinates.""" - cfg = self.cfg - num_s_regs = cfg.num_s_regs_per_thread - old_max = new_max_arr[0] - running_sum = sum_arr[0] - s_vals = cutlass.Array(Float32, num_s_regs, space=cutlass.AddressSpace.rmem) - task_cache = _decode_gen_task_cache(stage_info) - seq_len_kv = _load_runtime_seq_len_kv( - self.seqlens_kv, - self.max_seq_len_kv, - stage_info, - Int32(0), - Int32(0), - ) - warp_grp_thread_idx = Int32(task_cache[_TASK_CACHE_WARP_GRP_THREAD_IDX]) - lane_idx = Int32(task_cache[_TASK_CACHE_LANE_IDX]) - tile_row_idx = _keeps_row_idx(cfg, warp_grp_thread_idx) - col_base = _keeps_col_base(cfg, lane_idx, num_s_regs) - ( - origin0, - origin1, - valid0, - valid1, - local_token_words, - route_token_mask_is_full, - ) = self._decode_sparse_mask_metadata( - routed_origin0=routed_origin0, - routed_origin1=routed_origin1, - routed_route_flags=routed_route_flags, - routed_token_word0=routed_token_word0, - routed_token_word1=routed_token_word1, - routed_token_word2=routed_token_word2, - routed_token_word3=routed_token_word3, - ) - - base_addr = ( - task_cache[_TASK_CACHE_TMEM_BASE_OFFSET] - + Int32(self._alloc.offset) - + self._softmax_loop_stage_slot_offset(stage_info) - ) - num_load_atoms = num_s_regs // 32 - if cutlass.const_expr(cfg.tile_size_q == 64 and cfg.use_kv_valid_bits): - token_mask_is_required = not route_token_mask_is_full - - # Keep each Q64 atom's load, wait, and mask together. A/B testing - # showed that hoisting both loads extends live fragment ranges and - # regresses the Q64 code generated by ptxas. - for load_atom_idx in cutlass.range_constexpr(2): - atom_col = load_atom_idx * 32 - loaded = _keeps_tcgen05_ld( - cfg, - prims.make_tmem_ptr(base_addr + Int32(atom_col), Float32), - num=32, - offset=cfg.tile_size_kv // 2, - ) - prims.tcgen05_wait(kind=prims.Tcgen05Wait.LOAD) - self._mask_and_store_sparse_keeps_atom( - s_vals, - loaded, - local_token_words[load_atom_idx], - atom_col=atom_col, - token_mask_is_required=token_mask_is_required, - ) - else: - for load_atom_idx in cutlass.range_constexpr(num_load_atoms): - atom_col = load_atom_idx * 32 - loaded = _keeps_tcgen05_ld( - cfg, - prims.make_tmem_ptr(base_addr + Int32(atom_col), Float32), - num=32, - offset=cfg.tile_size_kv // 2, - ) - prims.tcgen05_wait(kind=prims.Tcgen05Wait.LOAD) - for atom_reg_idx in cutlass.range_constexpr(32): - score_idx = atom_col + atom_reg_idx - s_vals[score_idx] = loaded[atom_reg_idx] - - logical_q_group_idx = _logical_q_group_idx(cfg, stage_info, self.q_group_idx) - q_token_idx, _ = _q_row_token_and_local_head( - cfg, - self.h_r, - logical_q_group_idx, - tile_row_idx, - ) - q_row_is_valid = _q_row_is_valid_for_seq( - cfg, - self.h_r, - logical_q_group_idx, - tile_row_idx, - self.seq_len_q, - ) - causal_end = seq_len_kv - self.seq_len_q + q_token_idx + Int32(1) - can_skip_structural_mask = _can_skip_sparse_keeps_structural_mask( - q_row_is_valid, - origin0, - origin1, - valid0, - valid1, - seq_len_kv, - causal_end, - apply_causal_mask=cfg.mask_type == CAUSAL, - ) - # This guard covers only route/Q/tail/causal structure. Q64 token - # holes were applied while materializing its two LDTM atoms; Q128 - # applies them in the post-pass below. - if not can_skip_structural_mask: - for reg_idx in cutlass.range_constexpr(num_s_regs): - fragment_offset = Int32(reg_idx) - logical_k = origin0 + fragment_offset - fragment_valid = valid0 - if cutlass.const_expr(cfg.tile_size_q == 128 and reg_idx >= 64): - fragment_offset = Int32(reg_idx - 64) - logical_k = origin1 + fragment_offset - fragment_valid = valid1 - elif cutlass.const_expr(cfg.tile_size_q == 64): - if col_base >= Int32(64): - logical_k = origin1 + fragment_offset - fragment_valid = valid1 - - score_is_valid = ( - q_row_is_valid - and fragment_valid != Int32(0) - and logical_k < seq_len_kv - ) - if cutlass.const_expr(cfg.mask_type == CAUSAL): - score_is_valid = score_is_valid and logical_k < causal_end - if not score_is_valid: - s_vals[reg_idx] = _neg_max_f32() - - # Q128 deliberately keeps all four LDTM atoms adjacent: unlike Q64, - # interleaving each load with mask control flow regresses its codegen. - # The post-pass follows structural masking; the producer's runtime - # route flag skips it only when all four current token words are full. - if cutlass.const_expr(cfg.tile_size_q == 128 and cfg.use_kv_valid_bits): - token_mask_is_required = not route_token_mask_is_full - if token_mask_is_required: - for word_idx in cutlass.range_constexpr(4): - token_word = local_token_words[word_idx] - for bit_idx in cutlass.range_constexpr(32): - reg_idx = word_idx * 32 + bit_idx - token_bit_is_valid = ( - (token_word >> Int32(bit_idx)) & Uint32(1) - ) != Uint32(0) - if not token_bit_is_valid: - s_vals[reg_idx] = _neg_max_f32() - - tile_max = self._reduce_keeps_row_max(s_vals) - self._publish_keeps_softmax_state( - s_vals, - tile_max, - old_max, - running_sum, - old_max_arr, - sum_arr, - new_max_arr, - s_arr, - ) - return old_max_arr, sum_arr, new_max_arr, s_arr - @cute.jit def _compute_softmax_loop_keeps( self, @@ -1407,8 +1102,26 @@ def _compute_softmax_loop_keeps( reduction, whose 16x256b register mapping is unrelated to Keeps. """ cfg = self.cfg + if cutlass.const_expr(cfg.streams_tmem_p_fragments): + # Streamed profiles share the fragment max pass with block-sparse + # routes; dense tiles describe their visible range as keep words. + return self._compute_softmax_loop_keeps_fragments( + stage_info, + old_max_arr=old_max_arr, + sum_arr=sum_arr, + new_max_arr=new_max_arr, + s_arr=s_arr, + use_sparse=False, + sparse_origin0=Int32(0), + sparse_origin1=Int32(0), + sparse_route_flags=Int32(0), + sparse_token_word0=Uint32(0xFFFFFFFF), + sparse_token_word1=Uint32(0xFFFFFFFF), + sparse_token_word2=Uint32(0xFFFFFFFF), + sparse_token_word3=Uint32(0xFFFFFFFF), + ) task_cache = _decode_gen_task_cache(stage_info) - num_s_regs = cfg.softmax_score_fragment_regs + num_s_regs = cfg.num_s_regs_per_thread old_max = new_max_arr[0] running_sum = sum_arr[0] s_vals = cutlass.Array(Float32, num_s_regs, space=cutlass.AddressSpace.rmem) @@ -1434,50 +1147,9 @@ def _compute_softmax_loop_keeps( is_valid_effective_tile, is_masked_final_wave, tile_is_unmasked, + tile_has_valid_scores, ) = self._resolve_keeps_tile_context(stage_info) - if cutlass.const_expr(cfg.tile_size_kv == 256): - # KV256 owns four physical K32 fragments per lane. Reduce the max - # one fragment at a time so only one native LDTM atom is live; the - # P pass reloads the same fragments after the reference max is - # known. - tile_max = _neg_max_f32() - for fragment_idx in cutlass.range_constexpr( - cfg.num_softmax_score_fragments - ): - self._load_keeps_fragment( - stage_info, - s_vals, - tile_offset_k, - element_mask_end_idx, - window_start_idx, - seq_len_kv, - logical_q_group_idx, - is_valid_effective_tile, - is_masked_final_wave, - tile_is_unmasked, - fragment_idx=fragment_idx, - ) - fragment_max = self._reduce_keeps_fragment_max(s_vals) - tile_max = cute.math.max(tile_max, fragment_max, ftz=True) - - new_max = cute.math.max(old_max, tile_max, ftz=True) - if old_max != _neg_max_f32(): - # Keeping the previous reference max avoids an in-place O - # rescale when the new tile raises it only modestly. The - # 16-bit P path can represent the bounded values above one; the - # numerator and denominator remain in the same scale frame. - # Large jumps still rebase to keep P comfortably in range. - max_delta_log2 = self.scale_softmax_log2 * (old_max - new_max) - if max_delta_log2 >= Float32(-KV_TILE_256_RESCALE_THRESHOLD_LOG2): - new_max = old_max - old_max_arr[0] = old_max - sum_arr[0] = running_sum - new_max_arr[0] = new_max - for reg_idx in cutlass.range_constexpr(num_s_regs): - s_arr[reg_idx] = s_vals[reg_idx] - return old_max_arr, sum_arr, new_max_arr, s_arr - if cutlass.const_expr(use_preload_mask_split): # Select the complete unmasked/masked TMEM load+max path before any S # registers are materialized. The shared predicate covers the @@ -1495,7 +1167,7 @@ def _compute_softmax_loop_keeps( is_masked_final_wave, tile_is_unmasked, ) - tile_max = self._reduce_keeps_fragment_max(s_vals) + tile_max = self._reduce_keeps_row_max(s_vals) self._publish_keeps_softmax_state( s_vals, @@ -1509,12 +1181,7 @@ def _compute_softmax_loop_keeps( ) return old_max_arr, sum_arr, new_max_arr, s_arr - should_load_s = ( - is_valid_effective_tile - and (tile_offset_k < seq_len_kv) - and not is_masked_final_wave - ) - if should_load_s: + if tile_has_valid_scores: base_addr = ( task_cache[_TASK_CACHE_TMEM_BASE_OFFSET] + Int32(self._alloc.offset) @@ -1600,73 +1267,6 @@ def _compute_softmax_loop_keeps( ) return old_max_arr, sum_arr, new_max_arr, s_arr - @consumer_work(returns=s_arr, work_attrs=WorkAttr.AUXILIARY) - @cute.jit - def load_softmax_p_fragment( - self, - stage_info: StageInfo, - *, - fragment_idx: Constexpr[int], - s_arr: cutlass.Array, - ) -> cutlass.Array: - """Reload and mask one KV256 K32 fragment for P materialization.""" - if cutlass.const_expr(self.cfg.use_block_sparse): - return self._load_block_sparse_softmax_p_fragment( - stage_info, - fragment_idx=fragment_idx, - s_arr=s_arr, - ) - ( - seq_len_kv, - logical_q_group_idx, - element_mask_end_idx, - tile_offset_k, - window_start_idx, - is_valid_effective_tile, - is_masked_final_wave, - tile_is_unmasked, - ) = self._resolve_keeps_tile_context(stage_info) - self._load_keeps_fragment( - stage_info, - s_arr, - tile_offset_k, - element_mask_end_idx, - window_start_idx, - seq_len_kv, - logical_q_group_idx, - is_valid_effective_tile, - is_masked_final_wave, - tile_is_unmasked, - fragment_idx=fragment_idx, - ) - return s_arr - - @cute.jit - def _sparse_swaps_logical_k( - self, - lane_k_offset: Int32, - sparse_origin0: Int32, - sparse_origin1: Int32, - sparse_origin2: Int32, - sparse_origin3: Int32, - *, - token_group_idx: Constexpr[int], - ) -> tuple[Int32, Int32]: - """Map one SWAP register group to its routed logical K position.""" - - atom_size = min(self.cfg.kv_block_size, 32) - groups_per_atom = atom_size // 8 - origin_idx = token_group_idx // groups_per_atom - atom_origin = sparse_origin0 - if cutlass.const_expr(origin_idx == 1): - atom_origin = sparse_origin1 - elif cutlass.const_expr(origin_idx == 2): - atom_origin = sparse_origin2 - elif cutlass.const_expr(origin_idx == 3): - atom_origin = sparse_origin3 - token_offset = (token_group_idx % groups_per_atom) * 8 - return atom_origin, atom_origin + Int32(token_offset) + lane_k_offset - @cute.jit def _compute_softmax_loop_swaps( self, @@ -1860,6 +1460,11 @@ def _compute_softmax_loop_swaps( s_vals[q_repeats * 4 + ld_base + 2] = loaded1[ld_base + 2] s_vals[q_repeats * 4 + ld_base + 3] = loaded1[ld_base + 3] + route_is_proxy = cutlass.Boolean(False) + if cutlass.const_expr(use_sparse and cfg.use_block_sparse_proxy_routes): + route_is_proxy = cutlass.Boolean( + (sparse_route_flags & Uint32(_SOFTMAX_ROUTE_IS_PROXY_FLAG)) != Uint32(0) + ) if cutlass.const_expr(use_sparse): # Route, KV-tail, uniform-causal, and token validity depend only on # K, so one predicate masks the adjacent pair of Q-row registers. @@ -1894,7 +1499,8 @@ def _compute_softmax_loop_swaps( lane_k_offset = Int32(task_cache[_TASK_CACHE_LANE_IDX]) >> Int32(2) token_word_covers_kv_tail = _swaps_token_word_covers_kv_tail(cfg) for token_group_idx in cutlass.range_constexpr(4): - atom_origin, logical_k = self._sparse_swaps_logical_k( + atom_origin, logical_k = _swaps_routed_coordinate( + cfg, lane_k_offset, sparse_origin0, sparse_origin1, @@ -1906,19 +1512,20 @@ def _compute_softmax_loop_swaps( # tail. Qualified profiles can therefore omit the local # atom-origin guard, independently of the K/V issuer warp. score_is_valid = cutlass.Boolean(True) - if cutlass.const_expr( - not _swaps_uses_token_only_score_validity(cfg) - ): - score_is_valid = cutlass.Boolean(atom_origin >= Int32(0)) - if cutlass.const_expr(not token_word_covers_kv_tail): - score_is_valid = cutlass.Boolean( - score_is_valid and logical_k < seq_len_kv - ) - if cutlass.const_expr(cfg.uses_uniform_causal_mask): - score_is_valid = cutlass.Boolean( - score_is_valid and logical_k < element_mask_end_idx - ) - if cutlass.const_expr(cfg.use_kv_valid_bits): + if not route_is_proxy: + if cutlass.const_expr( + not _swaps_uses_token_only_score_validity(cfg) + ): + score_is_valid = cutlass.Boolean(atom_origin >= Int32(0)) + if cutlass.const_expr(not token_word_covers_kv_tail): + score_is_valid = cutlass.Boolean( + score_is_valid and logical_k < seq_len_kv + ) + if cutlass.const_expr(cfg.uses_uniform_causal_mask): + score_is_valid = cutlass.Boolean( + score_is_valid and logical_k < element_mask_end_idx + ) + if cutlass.const_expr(cfg.uses_prepared_score_keep_words): token_bit_idx = Int32(token_group_idx * 8) + lane_k_offset token_is_valid = ( (sparse_token_word >> token_bit_idx) & Uint32(1) @@ -2061,7 +1668,8 @@ def _compute_softmax_loop_swaps( tile_offset_k + local_idx_k0 + Int32(token_group_idx * 8) ) if cutlass.const_expr(use_sparse): - _, token_idx = self._sparse_swaps_logical_k( + _, token_idx = _swaps_routed_coordinate( + cfg, lane_idx >> Int32(2), sparse_origin0, sparse_origin1, @@ -2351,7 +1959,7 @@ def reduce_sums( return sum_arr @cute.jit - def _compute_softmax_loop_sparse_keeps_kv256( + def _compute_softmax_loop_keeps_fragments( self, stage_info: StageInfo, *, @@ -2359,6 +1967,7 @@ def _compute_softmax_loop_sparse_keeps_kv256( sum_arr: cutlass.Array, new_max_arr: cutlass.Array, s_arr: cutlass.Array, + use_sparse: Constexpr[bool], sparse_origin0: Int32, sparse_origin1: Int32, sparse_route_flags: Int32, @@ -2367,77 +1976,141 @@ def _compute_softmax_loop_sparse_keeps_kv256( sparse_token_word2: Uint32, sparse_token_word3: Uint32, ) -> tuple[object, object, object, object]: - """Reduce one sparse KV256 route as four bounded K32 fragments. - - The full route path only loads and reduces scores. A partial route - predicates one native 32-score fragment at a time and writes it back - to TMEM, so the later P pass can replay masked scores without keeping - the logical 128-score tile live in registers. + """Mask streamed K32 score fragments in place and reduce their max. + + Every fragment gets one keep word. Block-sparse routes derive it from + their two K64 atom origins, validity flags and prepared token words; + dense tiles derive it from the tile's visible token range (sequence + end, uniform or per-row causal end, sliding-window start) and the Q + row's validity. Masked fragments are written back to TMEM so the P + pass can reload them without any mask logic. """ - cfg = self.cfg - assert cfg.tile_size_kv == 256 + assert cfg.streams_tmem_p_fragments + num_fragments = cfg.num_softmax_score_fragments + fragment_regs = cfg.softmax_score_fragment_regs + # The seven-slot softmax metadata ABI carries exactly four token words. + assert num_fragments == 4 and fragment_regs == 32 task_cache = _decode_gen_task_cache(stage_info) - token_words = ( - sparse_token_word0, - sparse_token_word1, - sparse_token_word2, - sparse_token_word3, + keep_words = cutlass.Array( + Uint32, num_fragments, space=cutlass.AddressSpace.rmem ) - keep_words = cutlass.Array(Uint32, 4, space=cutlass.AddressSpace.rmem) warp_group_thread_idx = Int32(task_cache[_TASK_CACHE_WARP_GRP_THREAD_IDX]) tile_row_idx = _keeps_row_idx(cfg, warp_group_thread_idx) - logical_q_group_idx = _logical_q_group_idx(cfg, stage_info, self.q_group_idx) - q_token_idx, _ = _q_row_token_and_local_head( - cfg, - self.h_r, - logical_q_group_idx, - tile_row_idx, - ) - q_row_is_valid = _q_row_is_valid_for_seq( - cfg, - self.h_r, - logical_q_group_idx, - tile_row_idx, - self.seq_len_q, - ) - seq_len_kv = _load_runtime_seq_len_kv( - self.seqlens_kv, - self.max_seq_len_kv, - stage_info, - Int32(0), - Int32(0), - ) - causal_end = seq_len_kv - self.seq_len_q + q_token_idx + Int32(1) - origin0 = Int32(sparse_origin0) - origin1 = Int32(sparse_origin1) - valid0 = sparse_route_flags & Int32(1) - valid1 = (sparse_route_flags >> Int32(1)) & Int32(1) - for fragment_idx in cutlass.range_constexpr(4): - fragment_origin = origin0 + Int32((fragment_idx % 2) * 32) - fragment_valid = valid0 - if cutlass.const_expr(fragment_idx >= 2): - fragment_origin = origin1 + Int32((fragment_idx % 2) * 32) - fragment_valid = valid1 - keep_words[fragment_idx] = _sparse_k32_effective_keep_word( - q_row_is_valid, - fragment_origin, - fragment_valid, - Uint32(token_words[fragment_idx]), - seq_len_kv, - causal_end, - apply_causal_mask=cfg.mask_type == CAUSAL, - apply_token_mask=cfg.use_kv_valid_bits, + if cutlass.const_expr(use_sparse): + token_words = ( + sparse_token_word0, + sparse_token_word1, + sparse_token_word2, + sparse_token_word3, + ) + logical_q_group_idx = _logical_q_group_idx( + cfg, stage_info, self.q_group_idx + ) + q_token_idx, _ = _q_row_token_and_local_head( + cfg, + self.h_r, + logical_q_group_idx, + tile_row_idx, + ) + q_row_is_valid = _q_row_is_valid_for_seq( + cfg, + self.h_r, + logical_q_group_idx, + tile_row_idx, + self.seq_len_q, + ) + seq_len_kv = _load_runtime_seq_len_kv( + self.seqlens_kv, + self.max_seq_len_kv, + stage_info, + Int32(0), + Int32(0), ) + causal_end = seq_len_kv - self.seq_len_q + q_token_idx + Int32(1) + origin0 = Int32(sparse_origin0) + origin1 = Int32(sparse_origin1) + valid0 = sparse_route_flags & Int32(1) + valid1 = (sparse_route_flags >> Int32(1)) & Int32(1) + fragments_per_origin = cfg.softmax_fragments_per_route_atom + for fragment_idx in cutlass.range_constexpr(num_fragments): + atom_offset = Int32( + (fragment_idx % fragments_per_origin) * fragment_regs + ) + fragment_origin = origin0 + atom_offset + fragment_valid = valid0 + if cutlass.const_expr(fragment_idx >= fragments_per_origin): + fragment_origin = origin1 + atom_offset + fragment_valid = valid1 + if cutlass.const_expr(cfg.trusts_prepared_score_words): + prepared_keep_word = Uint32(0) + if q_row_is_valid: + prepared_keep_word = Uint32(token_words[fragment_idx]) + keep_words[fragment_idx] = prepared_keep_word + else: + keep_words[fragment_idx] = _sparse_effective_keep_word( + q_row_is_valid, + fragment_origin, + fragment_valid, + Uint32(token_words[fragment_idx]), + seq_len_kv, + causal_end, + apply_causal_mask=cfg.mask_type == CAUSAL, + apply_token_mask=cfg.uses_prepared_score_keep_words, + ) - warp_scores_are_unmasked = cutlass.Boolean(True) - for fragment_idx in cutlass.range_constexpr(4): - warp_scores_are_unmasked = cutlass.Boolean( - warp_scores_are_unmasked - and keep_words[fragment_idx] == Uint32(0xFFFFFFFF) + else: + ( + seq_len_kv, + logical_q_group_idx, + element_mask_end_idx, + tile_offset_k, + window_start_idx, + _is_valid_effective_tile, + _is_masked_final_wave, + tile_is_unmasked, + rows_are_active, + ) = self._resolve_keeps_tile_context(stage_info) + if cutlass.const_expr(cfg.q_score_rows_need_mask): + rows_are_active = cutlass.Boolean( + rows_are_active + and _q_row_is_valid_for_seq( + cfg, + self.h_r, + logical_q_group_idx, + tile_row_idx, + self.seq_len_q, + ) + ) + visible_start = Int32(0) + visible_end = element_mask_end_idx + if cutlass.const_expr(cfg.uses_per_row_causal_mask): + q_token_idx, _ = _q_row_token_and_local_head( + cfg, + self.h_r, + logical_q_group_idx, + tile_row_idx, + ) + visible_end = seq_len_kv - self.seq_len_q + q_token_idx + Int32(1) + visible_start = _sliding_window_start_idx( + cfg, seq_len_kv, self.seq_len_q, q_token_idx + ) + elif cutlass.const_expr(cfg.use_sliding_window_causal): + visible_start = window_start_idx + # A tile that is unmasked for the whole Q group has all-ones keep + # words on every active row, so only masked tiles build them. + warp_scores_are_unmasked = cute.arch.vote_all_sync( + cutlass.Boolean(tile_is_unmasked and rows_are_active) ) - # The load/store branch must be uniform for each participating warp. - warp_scores_are_unmasked = cute.arch.vote_all_sync(warp_scores_are_unmasked) + if cutlass.const_expr(use_sparse): + warp_scores_are_unmasked = cutlass.Boolean(True) + for fragment_idx in cutlass.range_constexpr(num_fragments): + warp_scores_are_unmasked = cutlass.Boolean( + warp_scores_are_unmasked + and keep_words[fragment_idx] == Uint32(0xFFFFFFFF) + ) + # The load/store branch must be uniform for each participating warp. + warp_scores_are_unmasked = cute.arch.vote_all_sync(warp_scores_are_unmasked) score_tmem_addr = ( task_cache[_TASK_CACHE_TMEM_BASE_OFFSET] @@ -2449,17 +2122,17 @@ def _compute_softmax_loop_sparse_keeps_kv256( max_chains[chain_idx] = _neg_max_f32() if warp_scores_are_unmasked: - for fragment_idx in cutlass.range_constexpr(4): + for fragment_idx in cutlass.range_constexpr(num_fragments): loaded = _keeps_tcgen05_ld( cfg, prims.make_tmem_ptr( - score_tmem_addr + Int32(fragment_idx * 32), Float32 + score_tmem_addr + Int32(fragment_idx * fragment_regs), Float32 ), - num=32, + num=fragment_regs, offset=cfg.tile_size_kv // 2, ) prims.tcgen05_wait(kind=prims.Tcgen05Wait.LOAD) - for score_idx in cutlass.range_constexpr(32): + for score_idx in cutlass.range_constexpr(fragment_regs): chain_idx: Constexpr[int] = score_idx % 4 max_chains[chain_idx] = cute.math.max( max_chains[chain_idx], @@ -2467,19 +2140,35 @@ def _compute_softmax_loop_sparse_keeps_kv256( ftz=True, ) else: - for fragment_idx in cutlass.range_constexpr(4): - fragment_addr = score_tmem_addr + Int32(fragment_idx * 32) + if cutlass.const_expr(not use_sparse): + lane_idx = Int32(task_cache[_TASK_CACHE_LANE_IDX]) + col_base = _keeps_col_base(cfg, lane_idx, num_fragments * fragment_regs) + for fragment_idx in cutlass.range_constexpr(num_fragments): + fragment_token_base = tile_offset_k + _keeps_score_col( + cfg, + warp_group_thread_idx, + fragment_idx * fragment_regs, + col_base, + ) + keep_words[fragment_idx] = _dense_fragment_keep_word( + rows_are_active, + visible_start - fragment_token_base, + visible_end - fragment_token_base, + fragment_regs=fragment_regs, + ) + for fragment_idx in cutlass.range_constexpr(num_fragments): + fragment_addr = score_tmem_addr + Int32(fragment_idx * fragment_regs) loaded = _keeps_tcgen05_ld( cfg, prims.make_tmem_ptr(fragment_addr, Float32), - num=32, + num=fragment_regs, offset=cfg.tile_size_kv // 2, ) prims.tcgen05_wait(kind=prims.Tcgen05Wait.LOAD) masked_scores = cutlass.Array( - Float32, 32, space=cutlass.AddressSpace.rmem + Float32, fragment_regs, space=cutlass.AddressSpace.rmem ) - for score_idx in cutlass.range_constexpr(32): + for score_idx in cutlass.range_constexpr(fragment_regs): score = Float32(loaded[score_idx]) score_is_kept = ( (keep_words[fragment_idx] >> Int32(score_idx)) & Uint32(1) @@ -2494,7 +2183,7 @@ def _compute_softmax_loop_sparse_keeps_kv256( _keeps_tcgen05_st( cfg, prims.make_tmem_ptr(fragment_addr, Float32), - masked_scores.data_ptr().load(count=32, alignment=4), + masked_scores.data_ptr().load(count=fragment_regs, alignment=4), offset=cfg.tile_size_kv // 2, ) prims.tcgen05_wait(kind=prims.Tcgen05Wait.STORE) @@ -2506,44 +2195,11 @@ def _compute_softmax_loop_sparse_keeps_kv256( ftz=True, ) old_max = new_max_arr[0] - new_max = cute.math.max(old_max, tile_max, ftz=True) - if old_max != _neg_max_f32(): - max_delta_log2 = self.scale_softmax_log2 * (old_max - new_max) - if max_delta_log2 >= Float32(-KV_TILE_256_RESCALE_THRESHOLD_LOG2): - new_max = old_max + new_max = self._softmax_anchor(old_max, tile_max) old_max_arr[0] = old_max new_max_arr[0] = new_max return old_max_arr, sum_arr, new_max_arr, s_arr - @cute.jit - def _load_block_sparse_softmax_p_fragment( - self, - stage_info: StageInfo, - *, - fragment_idx: Constexpr[int], - s_arr: cutlass.Array, - ) -> cutlass.Array: - """Reload one full or already-predicated sparse KV256 fragment for P.""" - - assert self.cfg.tile_size_kv == 256 - task_cache = _decode_gen_task_cache(stage_info) - score_tmem_addr = ( - task_cache[_TASK_CACHE_TMEM_BASE_OFFSET] - + Int32(self._alloc.offset) - + self._softmax_loop_stage_slot_offset(stage_info) - + Int32(fragment_idx * 32) - ) - loaded = _keeps_tcgen05_ld( - self.cfg, - prims.make_tmem_ptr(score_tmem_addr, Float32), - num=32, - offset=self.cfg.tile_size_kv // 2, - ) - prims.tcgen05_wait(kind=prims.Tcgen05Wait.LOAD) - for score_idx in cutlass.range_constexpr(32): - s_arr[score_idx] = loaded[score_idx] - return s_arr - @consumer_work(returns=("old_max_arr", "sum_arr", "new_max_arr", "s_arr")) @cute.jit def compute_block_sparse_softmax_loop( @@ -2566,34 +2222,21 @@ def compute_block_sparse_softmax_loop( assert self.cfg.use_block_sparse if cutlass.const_expr(self.cfg.use_keeps_mma_ab): - if cutlass.const_expr(self.cfg.tile_size_kv == 256): - return self._compute_softmax_loop_sparse_keeps_kv256( - stage_info, - old_max_arr=old_max_arr, - sum_arr=sum_arr, - new_max_arr=new_max_arr, - s_arr=s_arr, - sparse_origin0=sparse_origin0, - sparse_origin1=sparse_origin1, - sparse_route_flags=sparse_route_flags, - sparse_token_word0=sparse_token_word0, - sparse_token_word1=sparse_token_word1, - sparse_token_word2=sparse_token_word2, - sparse_token_word3=sparse_token_word3, - ) - return self._compute_softmax_loop_sparse_keeps( + # Every block-sparse Keeps profile streams K32 fragments. + return self._compute_softmax_loop_keeps_fragments( stage_info, old_max_arr=old_max_arr, sum_arr=sum_arr, new_max_arr=new_max_arr, s_arr=s_arr, - routed_origin0=sparse_origin0, - routed_origin1=sparse_origin1, - routed_route_flags=sparse_route_flags, - routed_token_word0=sparse_token_word0, - routed_token_word1=sparse_token_word1, - routed_token_word2=sparse_token_word2, - routed_token_word3=sparse_token_word3, + use_sparse=True, + sparse_origin0=sparse_origin0, + sparse_origin1=sparse_origin1, + sparse_route_flags=sparse_route_flags, + sparse_token_word0=sparse_token_word0, + sparse_token_word1=sparse_token_word1, + sparse_token_word2=sparse_token_word2, + sparse_token_word3=sparse_token_word3, ) # SWAP reuses the Keeps seven-slot task ABI: all four origins remain # logical KV atom bases, but origin2 occupies the flags slot and diff --git a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_tasks.py b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_tasks.py index a84b92985204..0483ca9b3519 100644 --- a/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_tasks.py +++ b/tensorrt_llm/_torch/attention/backends/prims_ts/kernels/fmha_decode/fmha_decode_tasks.py @@ -32,16 +32,13 @@ import cutlass import cutlass.cute as cute -from cutlass import Int32 from cutlass.experimental import primitives as prims from cutlass.experimental.task_scheduling.memory import ( ResourceContext, - SmemAllocation, ) from cutlass.experimental.task_scheduling.resources import ( MemoryResource, StageInfo, - TaskLocalVariable, WorkQueue, consumer_work, producer_work, @@ -61,10 +58,9 @@ KV_INST1, KV_KIND_K, KV_KIND_V, - KV_TILE_256_SHARED_FIFO_STAGES, ) from .fmha_decode_resources.helpers_common import ( - ResourceVars, + _assume_nonnegative_i32, _q_group_token_base, _q_seq_bounds, _warp_broadcast_i32, @@ -122,14 +118,18 @@ def restore_slots(*resource_proxies: object) -> None: return traced -def _block_sparse_route_loop_domain( - route_count: cutlass.Int32, +def _loop_domain_after_head( + total_kv_tiles: cutlass.Int32, *, num_insts_kv: int, ) -> cutlass.Int32: - """Return LOOP iterations after HEAD reserves one candidate per instance.""" + """Return LOOP iterations after HEAD reserves one KV tile per instance. - remaining = route_count - cutlass.Int32(num_insts_kv) + Dense tiles and block-sparse routes share this recurrence; only the tile + count's source differs. + """ + + remaining = total_kv_tiles - cutlass.Int32(num_insts_kv) remaining = cute.math.max(remaining, cutlass.Int32(0)) insts = cutlass.Int32(num_insts_kv) return (remaining + insts - cutlass.Int32(1)) // insts @@ -152,102 +152,6 @@ def consume_schedule_token(self, stage_info: StageInfo) -> None: del stage_info -@dataclass(kw_only=True) -class SmemKvReuseCreditResource(MemoryResource): - """One-slot credit carrying the rotating KV256 exchange-stage index. - - Load publishes which drained 64-KiB physical K/V stage Correction may use - as tail scratch. The one-stage pipeline couples that payload to the same - ownership epoch: the following Load may use the other two physical stages, - but cannot publish a new alias until Correction releases this credit after - all output work completes. - """ - - cfg: cutlass.Constexpr[FmhaDecodeConfig] = None - _alloc: cutlass.Constexpr[SmemAllocation | None] = None - scratch_stage_slot: cutlass.Constexpr[TaskLocalVariable] = ( - TaskLocalVariable.uninitialized() - ) - - def __post_init__(self) -> None: - """Create the routed consumer slot for one physical K/V stage.""" - assert KV_TILE_256_SHARED_FIFO_STAGES == 3, ( - "KV256 reuse-credit rotation requires exactly three shared FIFO stages" - ) - if not self.cfg.uses_rotating_kv256_exchange: - raise ValueError( - "rotating KV scratch requires persistent direct Q64/KV256 " - "with two KV instructions, one head-dimension stage, and " - "one load warp" - ) - self.scratch_stage_slot = TaskLocalVariable( - dtype=Int32, - default=Int32(0), - docs="Physical shared-K/V stage reserved for KV256 tail exchange.", - ) - - def get_smem_requirements(self) -> list[SmemAllocation]: - """Allocate the one-word stage payload guarded by this pipeline.""" - if self._alloc is None: - self._alloc = SmemAllocation( - name=f"{self.name}_scratchStage", - size_bytes=4, - alignment=4, - ) - return [self._alloc] - - @cute.jit - def _payload(self, stage_info: StageInfo) -> cutlass.Array: - """Return the natural next-stage cursor owned by this credit.""" - return cutlass.Array( - stage_info.context.smem_base.data_ptr() + self._alloc.offset, - dtype=Int32, - shape=(1,), - addrspace=3, - ) - - @cute.jit - def create_function_variables( - self, - context: ResourceContext | None = None, - ) -> ResourceVars: - """Initialize the persistent ring cursor before TS tasks start.""" - if cutlass.const_expr(context is not None and context.smem_base is not None): - payload = cutlass.Array( - context.smem_base.data_ptr() + self._alloc.offset, - dtype=Int32, - shape=(1,), - addrspace=3, - ) - thread_idx, _, _ = cute.arch.thread_idx() - if thread_idx == Int32(0): - payload[0] = Int32(0) - return {} - - @producer_work - @cute.jit - def publish_scratch_stage(self, stage_info: StageInfo) -> None: - """Advance the persistent ring cursor and publish the drained stage.""" - num_stages = Int32(KV_TILE_256_SHARED_FIFO_STAGES) - if prims.elect_sync(): - payload = self._payload(stage_info) - # Each work commits T = 4 * (loop_end + 1) K/V transactions. - # Since 4 == 1 (mod 3), the cursor advances by loop_end + 1. - # loop_end is the resolved per-work domain, so heterogeneous - # runtime sequence lengths do not inherit a captured host bound. - payload[0] = ( - Int32(payload[0]) + stage_info.loop_end + Int32(1) - ) % num_stages - - @consumer_work(returns=scratch_stage_slot) - @cute.jit - def read_scratch_stage(self, stage_info: StageInfo) -> Int32: - """Read the alias only after the matching credit wait completes.""" - num_stages = Int32(KV_TILE_256_SHARED_FIFO_STAGES) - next_stage = Int32(self._payload(stage_info)[0]) - return (next_stage + num_stages - Int32(1)) % num_stages - - @dataclass(kw_only=True) class PackedDecodeWorkQueue(WorkQueue): """CLC work queue that drops packed-Q tiles beyond a batch's Q length.""" @@ -414,28 +318,23 @@ def _produce_staged_page_offsets( def _consume_staged_qk_mma( smem_kv: MemoryResource, tmem_s: MemoryResource, - aliased_p: MemoryResource, q_desc: Any, k_desc_label: str, qk_mma_label: str, section: FmhaStage, cfg: FmhaDecodeConfig, ) -> None: - """Consume all K head-dim stages for one QK MMA wave.""" + """Consume all K head-dim stages for one QK MMA wave. + + Streamed KV256 aliases P with the S columns this QK overwrites. The + preceding same-instance PV reads P as its TMEM A operand from the same + issuing thread, and the tensor core interlocks that read against a later + MMA's accumulator write, so no completion wait is needed before QK. + """ tmem_s.acquire() for head_dim_stage_idx in range(cfg.num_head_dim_stages_kv): smem_kv.wait() kv_desc = getattr(smem_kv, k_desc_label)() - if cutlass.const_expr( - cfg.streams_tmem_p_fragments - and head_dim_stage_idx == 0 - and (section == FmhaStage.Loop or cfg.use_persistent_scheduler) - ): - # Wait as late as possible: K staging overlaps the previous PV, - # but QK cannot overwrite the matching S/P alias until PV is done. - # Static HEAD has no previous tile; persistent HEAD may follow the - # same CTA's tail from another logical work tile and must wait. - aliased_p.wait_until_reusable_before_qk() if cfg.uses_q_desc_ref: getattr(tmem_s, f"{qk_mma_label}_from_q_ref")( kv_desc=kv_desc, @@ -451,6 +350,44 @@ def _consume_staged_qk_mma( tmem_s.commit() +def _consume_streamed_pv_fragments( + smem_kv: MemoryResource, + tmem_p: MemoryResource, + tmem_o: MemoryResource, + v_desc_label: str, + vp_mma_label: str, + cfg: FmhaDecodeConfig, +) -> None: + """Issue one PV wave as its K32 P fragments become ready. + + P fragment 0 is the earliest dependency: wait for it and for the + correction credit before holding the V stage. Later fragments may become + ready while the previous PV fragment is already executing; every slot + stays live through the complete async UMMA wave so the producer cannot + overwrite an operand prematurely. + """ + assert cfg.num_head_dim_stages_kv == 1 + fragment_label = f"{vp_mma_label}_fragment" + p_tmem_addr = tmem_p.wait_p_fragment(fragment_idx=0) + tmem_o.acquire() + smem_kv.wait() + v_desc = getattr(smem_kv, v_desc_label)() + getattr(tmem_o, fragment_label)( + v_desc=v_desc, + p_tmem_addr=p_tmem_addr, + fragment_idx=0, + ) + for fragment_idx in range(1, cfg.num_softmax_score_fragments): + p_tmem_addr = tmem_p.wait_p_fragment(fragment_idx=fragment_idx) + getattr(tmem_o, fragment_label)( + v_desc=v_desc, + p_tmem_addr=p_tmem_addr, + fragment_idx=fragment_idx, + ) + smem_kv.release() + tmem_o.commit() + + def _consume_staged_pv_mma( smem_kv: MemoryResource, tmem_p: MemoryResource, @@ -464,33 +401,9 @@ def _consume_staged_pv_mma( """Consume all V head-dim stages for one PV MMA wave.""" _ = section if cutlass.const_expr(cfg.streams_tmem_p_fragments): - assert cfg.num_head_dim_stages_kv == 1 - fragment_label = f"{vp_mma_label}_fragment" - - # P fragment 0 is the earliest dependency. Wait for it and for the - # correction credit before holding the shared V FIFO stage. - p_tmem_addr = tmem_p.wait_p_fragment(fragment_idx=0) - tmem_o.acquire() - smem_kv.wait() - v_desc = getattr(smem_kv, v_desc_label)() - getattr(tmem_o, fragment_label)( - v_desc=v_desc, - p_tmem_addr=p_tmem_addr, - fragment_idx=0, - ) - - # Later P fragments may become ready while the previous PV fragment is - # already executing. Keep every slot live through the complete async - # UMMA wave so the producer cannot overwrite an operand prematurely. - for fragment_idx in range(1, cfg.num_softmax_score_fragments): - p_tmem_addr = tmem_p.wait_p_fragment(fragment_idx=fragment_idx) - getattr(tmem_o, fragment_label)( - v_desc=v_desc, - p_tmem_addr=p_tmem_addr, - fragment_idx=fragment_idx, - ) - smem_kv.release() - tmem_o.commit() + _consume_streamed_pv_fragments( + smem_kv, tmem_p, tmem_o, v_desc_label, vp_mma_label, cfg + ) return tmem_p.wait() @@ -634,6 +547,47 @@ def _decode_work_tile_schedule_with_invariant_bridge( _work_queue_tail(work_queue) +@cute.jit +def _prepared_sparse_row_address( + cfg: cutlass.Constexpr[FmhaDecodeConfig], + q_group_idx: cutlass.Int32, + h_idx: cutlass.Int32, + b_idx: cutlass.Int32, + num_heads_kv: cutlass.Int32, +) -> cutlass.Int32: + """Map a (q_group, head, batch) tile to its prepared row header index.""" + + q_token_base = _q_group_token_base(cfg, q_group_idx) + q_block = q_token_base // cutlass.Int32(cfg.q_block_size) + num_q_blocks = (cfg.max_seq_len_q + cfg.q_block_size - 1) // cfg.q_block_size + return (b_idx * num_heads_kv + h_idx) * cutlass.Int32(num_q_blocks) + q_block + + +@cute.jit +def _prefetch_prepared_sparse_row( + cfg: cutlass.Constexpr[FmhaDecodeConfig], + row_route_offsets: cute.Pointer, + row_route_counts: cute.Pointer, + q_group_idx: cutlass.Int32, + h_idx: cutlass.Int32, + b_idx: cutlass.Int32, + num_heads_kv: cutlass.Int32, +) -> tuple[cutlass.Int32, cutlass.Int32]: + """Load one static tile's prepared row header before its tasks start. + + Every thread loads the same two words, so each warp issues one request + and the global-memory latency overlaps the CTA prologue (TMEM allocation + and barrier setup) instead of stalling every task at its first step. + """ + + row_address = _prepared_sparse_row_address( + cfg, q_group_idx, h_idx, b_idx, num_heads_kv + ) + row_route_begin = cutlass.Int32(row_route_offsets[row_address]) + route_count = _assume_nonnegative_i32(cutlass.Int32(row_route_counts[row_address])) + return row_route_begin, route_count + + @cute.jit def _load_prepared_sparse_row_warp( row_route_offsets: cute.Pointer, @@ -653,7 +607,7 @@ def _load_prepared_sparse_row_warp( loaded_row_route_begin = cutlass.Int32(row_route_offsets[row_address]) loaded_route_count = cutlass.Int32(row_route_counts[row_address]) row_route_begin = _warp_broadcast_i32(loaded_row_route_begin, 0) - route_count = _warp_broadcast_i32(loaded_route_count, 0) + route_count = _assume_nonnegative_i32(_warp_broadcast_i32(loaded_route_count, 0)) return row_route_begin, route_count @@ -667,6 +621,9 @@ def __init__(self, **kwargs: TaskKwarg) -> None: self.block_table_capacity = kwargs.pop("block_table_capacity", None) self.sparse_row_route_offsets = kwargs.pop("sparse_row_route_offsets", None) self.sparse_row_route_counts = kwargs.pop("sparse_row_route_counts", None) + # Static tiles may pass the already loaded row header instead. + self.sparse_row_route_begin = kwargs.pop("sparse_row_route_begin", None) + self.sparse_route_count = kwargs.pop("sparse_route_count", None) self.num_heads_kv = kwargs.pop("num_heads_kv", None) self.max_seq_len_kv = kwargs.pop("max_seq_len_kv", cutlass.Int32(0)) self.seq_len_q = kwargs.pop("seq_len_q", None) @@ -915,20 +872,20 @@ def get_domain(self, tile_coord: cute.Coord) -> cutlass.Int32 | int: q_group_idx = cutlass.Int32(tile_coord[0]) h_idx = cutlass.Int32(tile_coord[1]) b_idx = cutlass.Int32(tile_coord[2]) - q_token_base = _q_group_token_base(self.cfg, q_group_idx) - - q_block = q_token_base // self.cfg.q_block_size - num_q_blocks = ( - self.cfg.max_seq_len_q + self.cfg.q_block_size - 1 - ) // self.cfg.q_block_size - row_address = (b_idx * self.num_heads_kv + h_idx) * num_q_blocks + q_block - - row_route_begin, route_count = _load_prepared_sparse_row_warp( - row_route_offsets, - row_route_counts, - cutlass.Int32(row_address), - self._lane_idx, - ) + if self.sparse_row_route_begin is not None: + # The static kernel prologue already loaded this tile's header. + row_route_begin = self.sparse_row_route_begin + route_count = self.sparse_route_count + else: + row_address = _prepared_sparse_row_address( + self.cfg, q_group_idx, h_idx, b_idx, self.num_heads_kv + ) + row_route_begin, route_count = _load_prepared_sparse_row_warp( + row_route_offsets, + row_route_counts, + row_address, + self._lane_idx, + ) # Sparse route-span accessors share two underlying cache words # with paged KV. Clear dense/paged-only coordinates on every @@ -943,7 +900,7 @@ def get_domain(self, tile_coord: cute.Coord) -> cutlass.Int32 | int: self._kv_valid_tile_end = route_count self._kv_window_start = cutlass.Int32(0) - loop_domain = _block_sparse_route_loop_domain( + loop_domain = _loop_domain_after_head( route_count, num_insts_kv=self.cfg.num_insts_kv, ) @@ -989,14 +946,10 @@ def get_domain(self, tile_coord: cute.Coord) -> cutlass.Int32 | int: self._kv_window_start = cutlass.Int32(0) self._kv_valid_tile_end = total_kv_tiles self._kv_raw_tile_base = cutlass.Int32(0) - remaining_kv_tiles = cute.math.max( - total_kv_tiles - cutlass.Int32(self.cfg.num_insts_kv), - cutlass.Int32(0), + loop_domain = _loop_domain_after_head( + total_kv_tiles, + num_insts_kv=self.cfg.num_insts_kv, ) - num_insts_kv = cutlass.Int32(self.cfg.num_insts_kv) - loop_domain = ( - remaining_kv_tiles + num_insts_kv - cutlass.Int32(1) - ) // num_insts_kv return loop_domain + cutlass.Int32(self.domain_bias) # Decode the logical Q tile with the configured physical split fanout, @@ -1050,13 +1003,10 @@ def get_domain(self, tile_coord: cute.Coord) -> cutlass.Int32 | int: self._kv_raw_tile_base = skipped_tiles + split_idx * total_kv_tiles else: self._kv_raw_tile_base = skipped_tiles - remaining_kv_tiles = cute.math.max( - total_kv_tiles - cutlass.Int32(self.cfg.num_insts_kv), cutlass.Int32(0) + loop_domain = _loop_domain_after_head( + total_kv_tiles, + num_insts_kv=self.cfg.num_insts_kv, ) - num_insts_kv = cutlass.Int32(self.cfg.num_insts_kv) - loop_domain = ( - remaining_kv_tiles + num_insts_kv - cutlass.Int32(1) - ) // num_insts_kv # All tasks share the MMA-loop domain; tail-only tasks add a bias. return loop_domain + cutlass.Int32(self.domain_bias) @@ -1073,29 +1023,63 @@ def get_domain(self, tile_coord: cute.Coord) -> cutlass.Int32 | int: def _resolve_and_store_sparse_route( sparse_kv_metadata: MemoryResource | None, section: FmhaStage, -) -> tuple[Any, Any, Any, Any] | None: - """Resolve one prepared route and retain it for the matching K/V pair.""" + prefetch: tuple[Any, Any] | None = None, + *, + pipeline: bool = True, +) -> tuple[tuple[Any, Any, Any, Any] | None, tuple[Any, Any] | None]: + """Resolve one prepared route and retain it for the matching K/V pair. + + Returns ``(route, prefetch)``. With ``pipeline`` the record load is issued + one resolution ahead: HEAD loads its own record immediately, and every + resolution issues the load for the next one (LOOP iteration 0 from HEAD, + iteration i + 1 from iteration i) before the caller's K TMA burst, so the + global-memory latency overlaps that issue instead of stalling the load + warp. Callers pass the returned ``prefetch`` back into the next resolution + of the same instance, the way ``_staged_kv_load`` threads its cached page + IDs. Without ``pipeline`` the record is loaded where it is resolved and no + state is returned; the split-ring load variants use this because the + pipelined form measured slower for them. Dense profiles pass ``None`` and + get ``(None, None)``. + """ if sparse_kv_metadata is None: - return None + return None, None + if not pipeline: + prefetch = sparse_kv_metadata.prefetch_route( + target="head" if section == FmhaStage.Head else "current_loop" + ) + elif section == FmhaStage.Head: + prefetch = sparse_kv_metadata.prefetch_route(target="head") + assert prefetch is not None + prefetched_record_word, prefetched_record_offset = prefetch ( - resolved_origin0, + resolved_record_word, resolved_origin1, resolved_atom_validity, route_record_word_offset, - ) = sparse_kv_metadata.resolve_route(section=section) + ) = sparse_kv_metadata.resolve_route( + section=section, + prefetched_record_word_slot=prefetched_record_word, + prefetched_record_offset_slot=prefetched_record_offset, + ) sparse_kv_metadata.store_route( - resolved_origin0=resolved_origin0, + resolved_record_word=resolved_record_word, resolved_origin1=resolved_origin1, resolved_atom_validity=resolved_atom_validity, route_record_word_offset=route_record_word_offset, ) - return ( - resolved_origin0, + next_prefetch = None + if pipeline: + next_prefetch = sparse_kv_metadata.prefetch_route( + target="first_loop" if section == FmhaStage.Head else "next_loop" + ) + route = ( + resolved_record_word, resolved_origin1, resolved_atom_validity, route_record_word_offset, ) + return route, next_prefetch def _publish_sparse_softmax_route( @@ -1108,14 +1092,14 @@ def _publish_sparse_softmax_route( return assert route is not None ( - resolved_origin0, + resolved_record_word, resolved_origin1, resolved_atom_validity, route_record_word_offset, ) = route sparse_softmax_metadata.acquire() sparse_softmax_metadata.store_route( - resolved_origin0=resolved_origin0, + resolved_record_word=resolved_record_word, resolved_origin1=resolved_origin1, resolved_atom_validity=resolved_atom_validity, route_record_word_offset=route_record_word_offset, @@ -1128,7 +1112,6 @@ def create_load_task( smem_kv: MemoryResource, work_queue: WorkQueue | None, schedule_token_throttle: MemoryResource | None, - smem_kv_reuse_credit: MemoryResource | None, cfg: FmhaDecodeConfig, *, domain: int | cutlass.Int32, @@ -1150,7 +1133,6 @@ def load_schedule_body( smem_kv: MemoryResource, smem_page_offsets: MemoryResource | None, schedule_token_throttle: MemoryResource | None, - smem_kv_reuse_credit: MemoryResource | None, sparse_kv_metadata0: MemoryResource | None = None, sparse_kv_metadata1: MemoryResource | None = None, sparse_softmax_metadata0: MemoryResource | None = None, @@ -1200,26 +1182,22 @@ def _kv_load(label: str, section: FmhaStage) -> None: smem_page_offsets.wait() else: _page_offsets_consume(smem_page_offsets) - if sparse_kv_metadata0 is None: - for label in ("load_k0", "load_k1"): - _kv_load(label, FmhaStage.Head) - else: - route0 = _resolve_and_store_sparse_route( - sparse_kv_metadata0, FmhaStage.Head - ) - _kv_load("load_k0", FmhaStage.Head) - route1 = _resolve_and_store_sparse_route( - sparse_kv_metadata1, FmhaStage.Head - ) - _kv_load("load_k1", FmhaStage.Head) - # Issue both K tiles before either metadata FIFO can backpressure - # the load warp, matching the split-resource sparse cadence. - _publish_sparse_softmax_route(sparse_softmax_metadata0, route0) - _publish_sparse_softmax_route(sparse_softmax_metadata1, route1) - if smem_kv_reuse_credit is not None: - # K0/K1 occupy the two stages disjoint from the previous work's - # scratch. Acquire only before issuing the third K/V transaction. - smem_kv_reuse_credit.acquire() + # Dense profiles have no route metadata: the resolve and publish + # helpers are no-ops for ``None`` resources, so one cadence serves + # both dense and block-sparse loads. + route0, prefetch0 = _resolve_and_store_sparse_route( + sparse_kv_metadata0, FmhaStage.Head + ) + _kv_load("load_k0", FmhaStage.Head) + route1, prefetch1 = _resolve_and_store_sparse_route( + sparse_kv_metadata1, FmhaStage.Head + ) + _kv_load("load_k1", FmhaStage.Head) + # Issue both K tiles before either metadata FIFO can backpressure + # the load warp, matching the split-resource sparse cadence. + _publish_sparse_softmax_route(sparse_softmax_metadata0, route0) + _publish_sparse_softmax_route(sparse_softmax_metadata1, route1) + prefetch_by_label = {"load_k0": prefetch0, "load_k1": prefetch1} # LOOP: each iter prefetches the full ``num_insts_kv`` K/V pair set. # When P aliases the consumed S columns, MMA must consume each V/P pair @@ -1231,41 +1209,30 @@ def _kv_load(label: str, section: FmhaStage) -> None: else ("load_k0", "load_v0", "load_k1", "load_v1") ) with domain_loop(0, domain, 1, unroll=1): - if sparse_kv_metadata0 is None: - for label in loop_labels: - _kv_load(label, FmhaStage.Loop) - else: - # Follow the dense KV256 stage order exactly. Each V consumes - # its retained route before the matching K label replaces it. - loop_routes = [] - for label in loop_labels: - route = None - sparse_softmax_metadata = None - if label == "load_k0": - route = _resolve_and_store_sparse_route( - sparse_kv_metadata0, FmhaStage.Loop - ) - sparse_softmax_metadata = sparse_softmax_metadata0 - elif label == "load_k1": - route = _resolve_and_store_sparse_route( - sparse_kv_metadata1, FmhaStage.Loop - ) - sparse_softmax_metadata = sparse_softmax_metadata1 - _kv_load(label, FmhaStage.Loop) - if route is not None: - loop_routes.append((sparse_softmax_metadata, route)) - for sparse_softmax_metadata, route in loop_routes: - _publish_sparse_softmax_route(sparse_softmax_metadata, route) + # Generic V-first profiles consume their retained route before the + # matching K label replaces it. + route_metadata_by_label = { + "load_k0": (sparse_kv_metadata0, sparse_softmax_metadata0), + "load_k1": (sparse_kv_metadata1, sparse_softmax_metadata1), + } + loop_routes = [] + for label in loop_labels: + kv_metadata, softmax_metadata = route_metadata_by_label.get( + label, (None, None) + ) + route, prefetch_by_label[label] = _resolve_and_store_sparse_route( + kv_metadata, FmhaStage.Loop, prefetch_by_label.get(label) + ) + _kv_load(label, FmhaStage.Loop) + if route is not None: + loop_routes.append((softmax_metadata, route)) + for sparse_softmax_metadata, route in loop_routes: + _publish_sparse_softmax_route(sparse_softmax_metadata, route) # TAIL: after no more future K tiles are needed, load the final two V # tiles consumed by the final BMM2 calls. for label in ("load_v0", "load_v1"): _kv_load(label, FmhaStage.Tail) - if smem_kv_reuse_credit is not None: - # Publish the physical stage drained by this work together with - # the ownership token consumed by the correction tail. - smem_kv_reuse_credit.publish_scratch_stage() - smem_kv_reuse_credit.commit() if hold_page_window: _page_offsets_release(smem_page_offsets) @@ -1280,7 +1247,6 @@ def load_schedule( sparse_softmax_metadata1: MemoryResource | None, work_queue: WorkQueue | None, schedule_token_throttle: MemoryResource | None, - smem_kv_reuse_credit: MemoryResource | None, ) -> None: """Schedule shared-KV loads with only the resources in this profile.""" @@ -1292,7 +1258,6 @@ def load_schedule( smem_kv, smem_page_offsets, schedule_token_throttle, - smem_kv_reuse_credit, sparse_kv_metadata0, sparse_kv_metadata1, sparse_softmax_metadata0, @@ -1326,7 +1291,6 @@ def load_schedule( sparse_softmax_metadata1, work_queue, schedule_token_throttle, - smem_kv_reuse_credit, ) src = [] for sparse_kv_metadata in (sparse_kv_metadata0, sparse_kv_metadata1): @@ -1342,8 +1306,6 @@ def load_schedule( dst.append(sparse_resource) if schedule_token_throttle is not None: dst.append(schedule_token_throttle) - if smem_kv_reuse_credit is not None: - dst.append(smem_kv_reuse_credit) return task_class( src_resources=src, dst_resources=dst, @@ -1712,7 +1674,9 @@ def load_tile( ) in active_instances: if smem_k is None: continue - route = _resolve_and_store_sparse_route(sparse_kv_metadata, FmhaStage.Head) + route, _ = _resolve_and_store_sparse_route( + sparse_kv_metadata, FmhaStage.Head, pipeline=False + ) load_tile(smem_k, load_k, smem_page_offsets_k, FmhaStage.Head) head_routes.append((sparse_softmax_metadata, route)) # In the combined task, preserve both K issues ahead of Softmax @@ -1740,8 +1704,8 @@ def load_tile( smem_page_offsets_v_local, FmhaStage.Loop, ) - route = _resolve_and_store_sparse_route( - sparse_kv_metadata, FmhaStage.Loop + route, _ = _resolve_and_store_sparse_route( + sparse_kv_metadata, FmhaStage.Loop, pipeline=False ) load_tile(smem_k, load_k, smem_page_offsets_k, FmhaStage.Loop) loop_routes.append((sparse_softmax_metadata, route)) @@ -2167,26 +2131,9 @@ def pv_mma( section: FmhaStage, ) -> None: """Issue one scheduled PV wave using the selected phase work.""" - _ = section - tmem_p.wait() - p_desc_0, p_desc_1, p_tmem_addr_0, p_tmem_addr_1 = tmem_p.p_operands() - tmem_o.acquire() - for head_dim_stage_idx in range(cfg.num_head_dim_stages_kv): - smem_kv.wait() - v_desc = smem_kv.v_desc() - getattr(tmem_o, vp_mma_label)( - v_desc_0=v_desc, - v_desc_1=v_desc, - p_desc_0=p_desc_0, - p_desc_1=p_desc_1, - p_tmem_addr_0=p_tmem_addr_0, - p_tmem_addr_1=p_tmem_addr_1, - inst_idx=inst_idx, - head_dim_stage_idx=head_dim_stage_idx, - ) - smem_kv.release() - tmem_o.commit() - tmem_p.release() + _consume_staged_pv_mma( + smem_kv, tmem_p, tmem_o, "v_desc", vp_mma_label, inst_idx, section, cfg + ) qk_mma( smem_k0, @@ -2644,7 +2591,6 @@ def mma_schedule_body( _consume_staged_qk_mma( smem_kv, tmem_s0, - smem_p0, q_desc, "k_desc_0", "qk_mma_head", @@ -2654,7 +2600,6 @@ def mma_schedule_body( _consume_staged_qk_mma( smem_kv, tmem_s1, - smem_p1, q_desc, "k_desc_1", "qk_mma_head", @@ -2663,8 +2608,9 @@ def mma_schedule_body( ) # LOOP: consume aliased TMEM P before the next same-instance QK - # overwrites its S columns. SMEM-P profiles retain their established - # K-before-V cadence because P no longer depends on S lifetime. + # overwrites its S columns. Full-SMEM P remains score-dependent during + # Softmax replay, but owns independent storage once the replay commits + # and releases S; that completed handoff enables the QK-before-PV cadence. with domain_loop(0, domain, 1, unroll=1): if cfg.uses_two_inst_tmem_p: _consume_staged_pv_mma( @@ -2680,7 +2626,6 @@ def mma_schedule_body( _consume_staged_qk_mma( smem_kv, tmem_s0, - smem_p0, q_desc, "k_desc_0", "qk_mma_loop", @@ -2712,7 +2657,6 @@ def mma_schedule_body( _consume_staged_qk_mma( smem_kv, tmem_s1, - smem_p1, q_desc, "k_desc_1", "qk_mma_loop", @@ -2731,6 +2675,12 @@ def mma_schedule_body( cfg, ) + # Q is live for every BMM1 call, and the last BMM1 has been issued once + # the loop ends. Releasing here commits after those MMAs complete, so + # the next tile's Q load overlaps the final softmax and BMM2 waves + # instead of waiting for them. + smem_q.release() + # TAIL: no future K tiles remain, so only the final two BMM2 waves run. _consume_staged_pv_mma( smem_kv, @@ -2752,8 +2702,6 @@ def mma_schedule_body( FmhaStage.Tail, cfg, ) - # Q is live for every BMM1 call and can be released only after the loop. - smem_q.release() def mma_schedule_prelude( smem_q: MemoryResource, @@ -2876,9 +2824,9 @@ def softmax0_schedule_body( sparse_softmax_metadata.init_read_state() with domain_loop(0, domain, 1, unroll=1) as d: - # ConsWait/ConsWork: load S from TMEM and compute the tile max. - tmem_s0.wait() if sparse_softmax_metadata is not None: + # Consume the independent metadata stream first so its SMEM + # loads and release can overlap the subsequent score wait. sparse_softmax_metadata.wait() # Copy the complete payload to registers before release, so # masking cannot race the producer's next SMEM-stage reuse. @@ -2892,6 +2840,9 @@ def softmax0_schedule_body( sparse_token_word3, ) = sparse_softmax_metadata.load_route() sparse_softmax_metadata.release() + # ConsWait/ConsWork: load S from TMEM and compute the tile max. + tmem_s0.wait() + if sparse_softmax_metadata is not None: old_max_arr, sum_arr, new_max_arr, s_arr = ( tmem_s0.compute_block_sparse_softmax_loop( old_max_arr=old_max_arr, @@ -2929,18 +2880,17 @@ def softmax0_schedule_body( ) tmem_softmax_local0.commit() if cutlass.const_expr(cfg.streams_tmem_p_fragments): - # Publish one K32 probability fragment at a time so PV can - # consume early fragments while later scores are processed. - for fragment_idx in range(cfg.num_softmax_score_fragments): - s_arr = tmem_s0.load_softmax_p_fragment( - fragment_idx=fragment_idx, - s_arr=s_arr, - ) - smem_p0.compute_p_fragment( - fragment_idx=fragment_idx, + # One rolled loop streams every K32 probability fragment; the + # fragment body exists once in the instruction stream. + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + smem_p0.compute_proxy_route_p_fragments( new_max_arr=new_max_arr, - s_arr=s_arr, + route_flags=sparse_route_flags, + route_origin0=sparse_origin0, + route_origin1=sparse_origin1, ) + else: + smem_p0.compute_p_fragments(new_max_arr=new_max_arr) else: # Wait for a free P stage before entering the ordered window so # BMM2 backpressure on this group's P pipeline cannot extend the @@ -2951,16 +2901,27 @@ def softmax0_schedule_body( # ProdWork: compute P=exp(S-new_max), store it in the profile's # SMEM or staged-TMEM operand, and record local sums for the # running softmax sum update. - smem_p0.compute_p( - new_max_arr=new_max_arr, - s_arr=s_arr, - ) # publishes the local denominator through tmem_s0 + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + smem_p0.compute_proxy_route_p( + new_max_arr=new_max_arr, + s_arr=s_arr, + route_origin0=sparse_origin0, + route_origin1=sparse_origin1, + keeps_route_flags_or_swaps_origin2=sparse_route_flags, + swaps_route_origin3_bits=sparse_token_word0, + swaps_route_flags=sparse_token_word2, + ) + else: + smem_p0.compute_p( + new_max_arr=new_max_arr, + s_arr=s_arr, + ) # publishes the local denominator through tmem_s0 smem_p0.commit() if tmem_softmax_order is not None: tmem_softmax_order.release_softmax1() if cutlass.const_expr(cfg.use_keeps_mma_ab and cfg.uses_tmem_p): - # The TMEM-P store has consumed the aliased S columns, so the - # next QK wave can now overwrite them. + # TMEM-P has consumed the aliased S columns, so the next QK + # wave can now overwrite S. tmem_s0.release() # ProdWork: FP8 path applies the cross-resource sum correction # before TmemS.reduce_sums publishes the new running sums. @@ -3098,9 +3059,9 @@ def softmax1_schedule_body( sparse_softmax_metadata.init_read_state() with domain_loop(0, domain, 1, unroll=1) as d: - # ConsWait/ConsWork: load the second S instance and compute max. - tmem_s1.wait() if sparse_softmax_metadata is not None: + # Consume the independent metadata stream first so its SMEM + # loads and release can overlap the subsequent score wait. sparse_softmax_metadata.wait() # Copy to registers before release so the producer can reuse # the SMEM stage while this warp group applies the masks. @@ -3114,6 +3075,9 @@ def softmax1_schedule_body( sparse_token_word3, ) = sparse_softmax_metadata.load_route() sparse_softmax_metadata.release() + # ConsWait/ConsWork: load the second S instance and compute max. + tmem_s1.wait() + if sparse_softmax_metadata is not None: old_max_arr, sum_arr, new_max_arr, s_arr = ( tmem_s1.compute_block_sparse_softmax_loop( old_max_arr=old_max_arr, @@ -3149,16 +3113,15 @@ def softmax1_schedule_body( ) tmem_softmax_local1.commit() if cutlass.const_expr(cfg.streams_tmem_p_fragments): - for fragment_idx in range(cfg.num_softmax_score_fragments): - s_arr = tmem_s1.load_softmax_p_fragment( - fragment_idx=fragment_idx, - s_arr=s_arr, - ) - smem_p1.compute_p_fragment( - fragment_idx=fragment_idx, + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + smem_p1.compute_proxy_route_p_fragments( new_max_arr=new_max_arr, - s_arr=s_arr, + route_flags=sparse_route_flags, + route_origin0=sparse_origin0, + route_origin1=sparse_origin1, ) + else: + smem_p1.compute_p_fragments(new_max_arr=new_max_arr) else: # Wait for a free P stage before entering the ordered window so # BMM2 backpressure on this group's P pipeline cannot extend the @@ -3167,7 +3130,18 @@ def softmax1_schedule_body( if tmem_softmax_order is not None: tmem_softmax_order.wait_softmax1() # ProdWork: compute and publish P1 for BMM2. - smem_p1.compute_p(new_max_arr=new_max_arr, s_arr=s_arr) + if cutlass.const_expr(cfg.use_block_sparse_proxy_routes): + smem_p1.compute_proxy_route_p( + new_max_arr=new_max_arr, + s_arr=s_arr, + route_origin0=sparse_origin0, + route_origin1=sparse_origin1, + keeps_route_flags_or_swaps_origin2=sparse_route_flags, + swaps_route_origin3_bits=sparse_token_word0, + swaps_route_flags=sparse_token_word2, + ) + else: + smem_p1.compute_p(new_max_arr=new_max_arr, s_arr=s_arr) smem_p1.commit() if tmem_softmax_order is not None: tmem_softmax_order.release_softmax0() @@ -3282,7 +3256,6 @@ def create_correction_task( tmem_corr0: MemoryResource, tmem_corr1: MemoryResource, work_queue: WorkQueue | None, - smem_kv_reuse_credit: MemoryResource | None, cfg: FmhaDecodeConfig, *, domain: int | cutlass.Int32, @@ -3293,9 +3266,6 @@ def create_correction_task( ) -> Task: """Create the two-instance correction and output task.""" - if smem_kv_reuse_credit is not None and work_queue is None: - raise ValueError("KV reuse credit requires a work queue") - def correction_schedule_body( tmem_softmax_local0: MemoryResource, tmem_softmax_local1: MemoryResource, @@ -3304,7 +3274,6 @@ def correction_schedule_body( tmem_corr1: MemoryResource, tmem_stats_done0: MemoryResource | None, tmem_stats_done1: MemoryResource | None, - smem_kv_reuse_credit: MemoryResource | None, ) -> None: """Schedule two-instance O correction and final output normalization.""" @@ -3498,37 +3467,17 @@ def correct_o( tail_o_stage_idx_1=tail_1, inst_idx=KV_INST1, ) - if smem_kv_reuse_credit is None: - tmem_corr1.correction_tail_epilogue( - o_stage_idx=o_stage_idx, - tail_o_stage_idx_0=tail_0, - tail_o_stage_idx_1=tail_1, - old_max_arr=old_max_arr, - new_max_arr=new_max_arr, - inst0_new_max_arr=inst0_new_max_arr, - inst0_sum_arr=inst0_sum_arr, - inst1_new_max_arr=inst1_new_max_arr, - inst1_sum_arr=inst1_sum_arr, - ) - else: - # The stage selector and ownership token share one pipeline epoch. - # Wait before the first aliased access and release immediately - # after correction stops touching the selected KV-ring stage. - smem_kv_reuse_credit.wait() - scratch_stage = smem_kv_reuse_credit.read_scratch_stage() - tmem_corr1.correction_tail_epilogue_rotating_exchange( - scratch_stage=scratch_stage, - o_stage_idx=o_stage_idx, - tail_o_stage_idx_0=tail_0, - tail_o_stage_idx_1=tail_1, - old_max_arr=old_max_arr, - new_max_arr=new_max_arr, - inst0_new_max_arr=inst0_new_max_arr, - inst0_sum_arr=inst0_sum_arr, - inst1_new_max_arr=inst1_new_max_arr, - inst1_sum_arr=inst1_sum_arr, - ) - smem_kv_reuse_credit.release() + tmem_corr1.correction_tail_epilogue( + o_stage_idx=o_stage_idx, + tail_o_stage_idx_0=tail_0, + tail_o_stage_idx_1=tail_1, + old_max_arr=old_max_arr, + new_max_arr=new_max_arr, + inst0_new_max_arr=inst0_new_max_arr, + inst0_sum_arr=inst0_sum_arr, + inst1_new_max_arr=inst1_new_max_arr, + inst1_sum_arr=inst1_sum_arr, + ) # Inst1 final reduction consumes both O0 and O1, so defer O0 release # until after inst1 has finished reading it. tmem_o.release() @@ -3542,7 +3491,6 @@ def run_correction_schedule( tmem_corr1: MemoryResource, tmem_stats_done0: MemoryResource | None, tmem_stats_done1: MemoryResource | None, - smem_kv_reuse_credit: MemoryResource | None, work_queue: WorkQueue | None, ) -> None: """Wrap correction with optional stats lifetime gates.""" @@ -3557,7 +3505,6 @@ def run_correction_schedule( tmem_corr1, tmem_stats_done0, tmem_stats_done1, - smem_kv_reuse_credit, ), ) @@ -3569,7 +3516,6 @@ def correction_schedule( tmem_corr0: MemoryResource, tmem_corr1: MemoryResource, work_queue: WorkQueue | None = None, - smem_kv_reuse_credit: MemoryResource | None = None, ) -> None: """Capture the Swaps correction schedule.""" run_correction_schedule( @@ -3580,7 +3526,6 @@ def correction_schedule( tmem_corr1, None, None, - smem_kv_reuse_credit, work_queue, ) @@ -3594,7 +3539,6 @@ def correction_keeps_schedule( tmem_stats_done0: MemoryResource, tmem_stats_done1: MemoryResource, work_queue: WorkQueue | None = None, - smem_kv_reuse_credit: MemoryResource | None = None, ) -> None: """Capture Keeps correction with explicit stats lifetime gates.""" run_correction_schedule( @@ -3605,7 +3549,6 @@ def correction_keeps_schedule( tmem_corr1, tmem_stats_done0, tmem_stats_done1, - smem_kv_reuse_credit, work_queue, ) @@ -3618,15 +3561,6 @@ def correction_keeps_schedule( tmem_corr0, tmem_corr1, ) - elif smem_kv_reuse_credit is None: - captured_schedule = correction_schedule( - tmem_softmax_local0, - tmem_softmax_local1, - tmem_o, - tmem_corr0, - tmem_corr1, - work_queue, - ) else: captured_schedule = correction_schedule( tmem_softmax_local0, @@ -3635,7 +3569,6 @@ def correction_keeps_schedule( tmem_corr0, tmem_corr1, work_queue, - smem_kv_reuse_credit, ) src = [tmem_softmax_local0, tmem_softmax_local1, tmem_o] else: @@ -3649,17 +3582,6 @@ def correction_keeps_schedule( tmem_stats_done0, tmem_stats_done1, ) - elif smem_kv_reuse_credit is None: - captured_schedule = correction_keeps_schedule( - tmem_softmax_local0, - tmem_softmax_local1, - tmem_o, - tmem_corr0, - tmem_corr1, - tmem_stats_done0, - tmem_stats_done1, - work_queue, - ) else: captured_schedule = correction_keeps_schedule( tmem_softmax_local0, @@ -3670,7 +3592,6 @@ def correction_keeps_schedule( tmem_stats_done0, tmem_stats_done1, work_queue, - smem_kv_reuse_credit, ) src = [ tmem_softmax_local0, @@ -3681,8 +3602,6 @@ def correction_keeps_schedule( ] if work_queue is not None: src.append(work_queue) - if smem_kv_reuse_credit is not None: - src.append(smem_kv_reuse_credit) return task_class( src_resources=src, dst_resources=[tmem_corr0, tmem_corr1], diff --git a/tensorrt_llm/_torch/attention/backends/sparse/hooks.py b/tensorrt_llm/_torch/attention/backends/sparse/hooks.py index 47b390098d05..abddc8ad0b2d 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/hooks.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/hooks.py @@ -12,6 +12,8 @@ from importlib import import_module from typing import TYPE_CHECKING, Optional +from .skip_softmax import SkipSoftmaxParams + if TYPE_CHECKING: import torch @@ -224,26 +226,36 @@ def prepare_sparse_runtime_params( backend: "TrtllmAttention", q: "torch.Tensor", k: Optional["torch.Tensor"], + v: Optional["torch.Tensor"], metadata: "AttentionMetadata", forward_args: "AttentionForwardArgs", ) -> "SparseRuntimeParams": - """Run backend prediction hooks and update attention-op parameters.""" - runtime_params = forward_args.sparse_runtime_params - if backend.sparse_params is None: - return runtime_params - + """Predict all sparse inputs for one attention call. + + Runs the ``sparse_kv_predict``, ``sparse_attn_predict`` and + ``block_sparse_attn_predict`` hooks once each and returns a new + ``SparseRuntimeParams`` built from ``forward_args.sparse_runtime_params`` + plus the hook results. Fields a backend writes into that carrier outside + the hooks, such as an auxiliary pool pointer, are carried over. SkipSoftmax + backends receive their threshold schedule last. + """ kv_indices, kv_offsets = backend.sparse_kv_predict(q, k, metadata, forward_args) attn_indices, attn_offsets = backend.sparse_attn_predict(q, k, metadata, forward_args) - block_size = ( - backend.sparse_params.indices_block_size - if attn_indices is not None or attn_offsets is not None - else runtime_params.sparse_attn_indices_block_size - ) - return replace( - runtime_params, + block_sparse_inputs = backend.block_sparse_attn_predict(q, k, v, metadata, forward_args) + has_attn_indices = attn_indices is not None or attn_offsets is not None + sparse_params = backend.sparse_params + runtime_params = replace( + forward_args.sparse_runtime_params, sparse_kv_indices=kv_indices, sparse_kv_offsets=kv_offsets, sparse_attn_indices=attn_indices, sparse_attn_offsets=attn_offsets, - sparse_attn_indices_block_size=block_size, + sparse_attn_indices_block_size=sparse_params.indices_block_size if has_attn_indices else 0, + block_sparse_inputs=block_sparse_inputs, ) + if isinstance(sparse_params, SkipSoftmaxParams): + runtime_params = sparse_params.scheduler.get_runtime_params( + runtime_params=runtime_params, + timestep=forward_args.timestep, + ) + return runtime_params diff --git a/tensorrt_llm/_torch/attention/backends/sparse/params.py b/tensorrt_llm/_torch/attention/backends/sparse/params.py index 956fc8c6fdbd..8e6829f8c879 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/params.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/params.py @@ -15,7 +15,7 @@ """Shared sparse attention parameter types.""" from dataclasses import dataclass -from typing import Optional +from typing import Literal, Optional import torch @@ -36,11 +36,56 @@ class SparseBackendForwardArgs: # Shared by algorithms that accept precomputed top-k indices. topk_indices: Optional[torch.Tensor] = None + # Complete block-sparse routing payload predicted by the module before the + # core forward; the default backend hook hands it through unchanged. + block_sparse_inputs: Optional["BlockSparseForwardInputs"] = None + + +@dataclass(frozen=True, slots=True) +class BlockSparseForwardInputs: + """Block geometry and live routing payload for one attention call. + + Exactly one routing representation is present. Canonical BSR uses + ``block_indptr`` and ``block_indices``; packed bitmask routing uses + ``exact_block_bits``. Paired K/V summaries enable proxy routes without + encoding an algorithm name in this shared carrier. + """ + + q_block_size: int + kv_block_size: int + max_blocks_per_row: Optional[int] = None + block_indptr: Optional[torch.Tensor] = None + block_indices: Optional[torch.Tensor] = None + exact_block_bits: Optional[torch.Tensor] = None + k_summary: Optional[torch.Tensor] = None + v_summary: Optional[torch.Tensor] = None + kv_valid_bits: Optional[torch.Tensor] = None + + def __post_init__(self) -> None: + has_bsr = self.block_indptr is not None + if has_bsr != (self.block_indices is not None): + raise ValueError("block_indptr and block_indices must be provided together") + if has_bsr == (self.exact_block_bits is not None): + raise ValueError("exactly one route representation must be provided") + if has_bsr and self.max_blocks_per_row is None: + raise ValueError("BSR routes require max_blocks_per_row") + if (self.k_summary is None) != (self.v_summary is None): + raise ValueError("k_summary and v_summary must be provided together") + + @property + def sparse_format(self) -> Literal["bsr", "bitmask"]: + """Routing representation selected by the live payload.""" + return "bitmask" if self.exact_block_bits is not None else "bsr" + + @property + def use_proxy_routes(self) -> bool: + """Whether unselected blocks are represented by K/V summaries.""" + return self.k_summary is not None @dataclass(kw_only=True, slots=True) class SparseRuntimeParams: - """Flat optional sparse inputs passed from a backend to ``AttentionOp``.""" + """Complete per-attention sparse runtime state consumed by FMHA/``AttentionOp``.""" # Sparse index inputs shared by multiple algorithms. sparse_kv_indices: Optional[torch.Tensor] = None @@ -57,3 +102,13 @@ class SparseRuntimeParams: threshold_scale_factor_prefill: float = 0.0 # SkipSoftmax decode threshold; diffusion models leave it at zero. threshold_scale_factor_decode: float = 0.0 + block_sparse_inputs: Optional[BlockSparseForwardInputs] = None + + +__all__ = [ + "BlockSparseForwardInputs", + "SparseBackendForwardArgs", + "SparseMetadataParams", + "SparseParams", + "SparseRuntimeParams", +] diff --git a/tensorrt_llm/_torch/attention/backends/trtllm.py b/tensorrt_llm/_torch/attention/backends/trtllm.py index 914f6f3a654e..32f5b4d0a3cc 100644 --- a/tensorrt_llm/_torch/attention/backends/trtllm.py +++ b/tensorrt_llm/_torch/attention/backends/trtllm.py @@ -46,8 +46,7 @@ PredefinedAttentionMask, RopeParams, merge_attention_forward_args) from .sparse.hooks import prepare_sparse_runtime_params -from .sparse.params import SparseParams -from .sparse.skip_softmax import SkipSoftmaxParams +from .sparse.params import BlockSparseForwardInputs, SparseParams _SKIP_CORRECTION_SUPPORTED_SMS = frozenset((100, 103)) @@ -1929,7 +1928,7 @@ def forward( ) forward_args.sparse_runtime_params = prepare_sparse_runtime_params( - self, q, k, metadata, forward_args) + self, q, k, v, metadata, forward_args) # Compute FlashMLA tile-scheduler metadata once per forward pass. # The flag is invalidated whenever FlashMLA inputs change. The metadata @@ -2059,14 +2058,6 @@ def forward( if forward_args.kv_scale_quant_orig is None: forward_args.kv_scale_quant_orig = self.kv_scale_quant_orig - sparse_params = self.sparse_params - if isinstance(sparse_params, SkipSoftmaxParams): - forward_args.sparse_runtime_params = ( - sparse_params.scheduler.get_runtime_params( - runtime_params=forward_args.sparse_runtime_params, - timestep=forward_args.timestep, - )) - # max_context_q_len_override is only set when encoder CUDA graphs are enabled. if metadata.max_context_q_len_override is not None: assert metadata.is_cuda_graph @@ -2286,6 +2277,26 @@ def sparse_kv_predict( """Predict sparse KV indices when required by an algorithm.""" return None, None + def block_sparse_attn_predict( + self, + q: torch.Tensor, + k: Optional[torch.Tensor], + v: Optional[torch.Tensor], + metadata: TrtllmAttentionMetadata, + forward_args: AttentionForwardArgs, + ) -> Optional[BlockSparseForwardInputs]: + """Predict the block-sparse routing payload for one attention call. + + The default hands through routes that the attention module predicted + before the core forward via ``sparse_backend_args``. Algorithms that + predict inside the backend override this method and return ``None`` + for dense phases. + """ + backend_args = forward_args.sparse_backend_args + if backend_args is None: + return None + return backend_args.block_sparse_inputs + def sparse_attn_predict( self, q: torch.Tensor, diff --git a/tensorrt_llm/_torch/visual_gen/attention_backend/flash_attn4.py b/tensorrt_llm/_torch/visual_gen/attention_backend/flash_attn4.py index 5024a43b192a..a441ce7edac7 100644 --- a/tensorrt_llm/_torch/visual_gen/attention_backend/flash_attn4.py +++ b/tensorrt_llm/_torch/visual_gen/attention_backend/flash_attn4.py @@ -39,10 +39,64 @@ def _install_cutlass_dsl_compatibility() -> None: cute.make_fragment = cute.make_rmem_tensor +def _install_flash_attn_tile_scheduler_compatibility() -> None: + """Keep FA4's four-axis ``WorkTileInfo`` independent of CUTLASS task scheduling. + + Importing ``cutlass.experimental.task_scheduling`` rewrites the shared + ``cutlass.utils.WorkTileInfo`` class in place: its constructor unpacks + ``tile_idx`` into exactly three scalars and ``tile_idx`` / ``is_valid_tile`` + become properties over those scalars. The vendored PrimTS kernels import + that package, so once any PrimTS FMHA has been probed or planned in a + process, every later FA4 kernel trace fails with ``ValueError: too many + values to unpack (expected 3)``: FA4 subclasses the same CUTLASS class with + a (block, head, batch, split) coordinate but does not define its own + constructor. Installing the upstream tuple semantics directly on the FA4 + subclass makes it immune to the parent rewrite regardless of import order. + Remove once CUTLASS stops patching the shared class or FA4 owns these + members itself. + """ + try: + from flash_attn.cute import tile_scheduler + except (ImportError, OSError): + return + import cutlass.cute as cute + from cutlass.cutlass_dsl import Boolean, extract_mlir_values + + work_tile_info = tile_scheduler.WorkTileInfo + if "__init__" in vars(work_tile_info): + return + + def __init__(self, tile_idx: cute.Coord, is_valid_tile: Boolean) -> None: + self._tile_idx = tile_idx + self._is_valid_tile = Boolean(is_valid_tile) + self._tile_idx_num_values = None + + def __extract_mlir_values__(self) -> list: + tile_idx_values = extract_mlir_values(self._tile_idx) + valid_values = extract_mlir_values(self._is_valid_tile) + self._tile_idx_num_values = len(tile_idx_values) + return tile_idx_values + valid_values + + @cute.jit + def tile_idx(self) -> cute.Coord: + return self._tile_idx + + @cute.jit + def is_valid_tile(self) -> Boolean: + return self._is_valid_tile + + work_tile_info.__init__ = __init__ + work_tile_info.__extract_mlir_values__ = __extract_mlir_values__ + work_tile_info.tile_idx = property(tile_idx) + work_tile_info.is_valid_tile = property(is_valid_tile) + + _flash_attn_fwd_import_error = None try: _install_cutlass_dsl_compatibility() from flash_attn.cute.interface import _flash_attn_fwd + + _install_flash_attn_tile_scheduler_compatibility() except (ImportError, OSError) as e: _flash_attn_fwd = None _flash_attn_fwd_import_error = e diff --git a/tests/unittest/_torch/attention/sparse/test_prims_ts_block_sparse.py b/tests/unittest/_torch/attention/sparse/test_prims_ts_block_sparse.py new file mode 100644 index 000000000000..bec985ef374b --- /dev/null +++ b/tests/unittest/_torch/attention/sparse/test_prims_ts_block_sparse.py @@ -0,0 +1,905 @@ +# 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. + +import math +from contextlib import nullcontext +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import torch +from utils.util import isSM100Family + +from tensorrt_llm._torch.attention.backends import prims_ts +from tensorrt_llm._torch.attention.backends.fmha import prims_ts_block_sparse as block_sparse_fmha +from tensorrt_llm._torch.attention.backends.fmha.interface import FmhaPhase +from tensorrt_llm._torch.attention.backends.fmha.phased import FmhaParams +from tensorrt_llm._torch.attention.backends.interface import ( + AttentionForwardArgs, + AttentionInputType, + PredefinedAttentionMask, +) +from tensorrt_llm._torch.attention.backends.sparse.params import ( + BlockSparseForwardInputs, + SparseRuntimeParams, +) +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm.functional import PositionEmbeddingType + +pytestmark = pytest.mark.cpu_only + +_REQUIRES_PRIMTS_GPU = pytest.mark.skipif( + not isSM100Family(), + reason="PrimTS block-sparse attention requires SM100 or SM103", +) + + +def _bsr_inputs(*, kv_valid_bits: torch.Tensor | None = None): + return BlockSparseForwardInputs( + q_block_size=64, + kv_block_size=64, + max_blocks_per_row=2, + block_indptr=torch.tensor([[[0, 2]], [[2, 4]]], dtype=torch.int32), + block_indices=torch.tensor([0, 1, 2, 3], dtype=torch.int32), + kv_valid_bits=kv_valid_bits, + ) + + +def _bitmask_inputs(*, proxy: bool): + summaries = { + "k_summary": torch.zeros((2, 4, 1, 128), dtype=torch.bfloat16), + "v_summary": torch.zeros((2, 4, 1, 128), dtype=torch.bfloat16), + } + return BlockSparseForwardInputs( + q_block_size=64, + kv_block_size=64, + exact_block_bits=torch.ones((2, 1, 1, 1), dtype=torch.uint32), + **(summaries if proxy else {}), + ) + + +def _set_block_sparse_inputs( + forward_args: AttentionForwardArgs, + block_sparse_inputs, +) -> None: + forward_args.sparse_runtime_params = SparseRuntimeParams( + block_sparse_inputs=block_sparse_inputs + ) + + +def _get_block_sparse_inputs(forward_args: AttentionForwardArgs): + block_sparse_inputs = forward_args.sparse_runtime_params.block_sparse_inputs + assert block_sparse_inputs is not None + return block_sparse_inputs + + +def _pack_token_mask(mask: torch.Tensor) -> torch.Tensor: + shifts = torch.arange(32, dtype=torch.int64, device=mask.device) + weights = torch.ones_like(shifts).bitwise_left_shift_(shifts) + return (mask.view(1, -1, 32).to(torch.int64) * weights).sum(dim=-1).to(torch.uint32) + + +def _proxy_reference( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + k_summary: torch.Tensor, + v_summary: torch.Tensor, + exact_block: int, +) -> torch.Tensor: + block_size = 64 + exact_tokens = torch.arange( + exact_block * block_size, + (exact_block + 1) * block_size, + device=q.device, + ) + proxy_blocks = [block for block in range(k_summary.shape[1]) if block != exact_block] + q_rows = q[0, :, 0].float() + exact_logits = q_rows @ k[0, exact_tokens, 0].float().T + proxy_logits = q_rows @ k_summary[0, proxy_blocks, 0].float().T + logits = torch.cat((exact_logits, proxy_logits), dim=1) / math.sqrt(q.shape[-1]) + weights = torch.exp(logits - logits.amax(dim=1, keepdim=True)) + exact_weights, proxy_weights = weights.split((block_size, len(proxy_blocks)), dim=1) + numerator = exact_weights @ v[0, exact_tokens, 0].float() + numerator += proxy_weights @ v_summary[0, proxy_blocks, 0].float() + denominator = exact_weights.sum(dim=1, keepdim=True) + denominator += proxy_weights.sum(dim=1, keepdim=True) * block_size + return (numerator / denominator).to(q.dtype)[None, :, None] + + +class _Attention: + def __init__(self) -> None: + self.sparse_params = None + self.num_heads = 2 + self.num_kv_heads = 1 + self.head_dim = 128 + self.is_mla_enable = False + self.kv_lora_rank = None + self.qk_rope_head_dim = None + self.qk_nope_head_dim = None + self.v_head_dim = None + self.q_scaling = 1.0 + self.quant_mode = 0 + self.local_layer_idx = 0 + self.position_embedding_type = PositionEmbeddingType.learned_absolute + self.attention_chunk_size = 0 + + +def _contiguous_case(): + attention = _Attention() + fmha = block_sparse_fmha.PrimsTSBlockSparseFmha(attention) + q = torch.zeros((128, 256), dtype=torch.bfloat16) + k = torch.zeros((512, 128), dtype=torch.bfloat16) + v = torch.zeros_like(k) + metadata = SimpleNamespace( + is_cross=False, + num_sparse_topk=0, + helix_position_offsets=None, + kv_cache_manager=None, + seq_lens=torch.tensor([64, 64], dtype=torch.int32), + ) + args = AttentionForwardArgs( + output=torch.empty_like(q), + attention_input_type=AttentionInputType.context_only, + attention_mask=PredefinedAttentionMask.FULL, + sparse_runtime_params=SparseRuntimeParams(block_sparse_inputs=_bsr_inputs()), + ) + return attention, fmha, q, k, v, metadata, args + + +def _paged_metadata(): + batch_size, max_pages, page_size = 2, 4, 64 + key_pages = torch.arange(batch_size * max_pages, dtype=torch.int32).view(batch_size, max_pages) + block_offsets = torch.stack((key_pages, key_pages + 8), dim=1).unsqueeze(0) + manager = Mock(spec=KVCacheManager) + manager.dtype = torch.bfloat16 + manager.num_pools = manager.num_local_layers = 1 + manager.host_kv_cache_block_offsets = block_offsets + return SimpleNamespace( + is_cross=False, + num_sparse_topk=0, + helix_position_offsets=None, + num_contexts=0, + num_generations=batch_size, + seq_lens=torch.ones(batch_size, dtype=torch.int32), + beam_width=1, + is_spec_decoding_enabled=False, + use_spec_decoding=False, + is_spec_dec_tree=False, + is_spec_dec_dynamic_tree=False, + tokens_per_block=page_size, + max_seq_len=max_pages * page_size, + kv_layout="HND", + kv_lens_runtime=torch.tensor([129, 193], dtype=torch.int32), + kv_cache_block_offsets=block_offsets, + host_kv_cache_pool_pointers=torch.tensor([[1234, 5678]], dtype=torch.int64), + host_kv_cache_pool_mapping=torch.tensor([[0, 0]], dtype=torch.int32), + kv_cache_manager=manager, + ) + + +def _paged_case(): + attention = _Attention() + fmha = block_sparse_fmha.PrimsTSBlockSparseFmha(attention) + fmha._multi_processor_count = 1 + metadata = _paged_metadata() + q = torch.zeros((2, 512), dtype=torch.bfloat16) + args = AttentionForwardArgs( + output=torch.empty((2, 256), dtype=q.dtype), + attention_input_type=AttentionInputType.generation_only, + attention_mask=PredefinedAttentionMask.CAUSAL, + attention_window_size=metadata.max_seq_len, + is_fused_qkv=True, + sparse_runtime_params=SparseRuntimeParams(block_sparse_inputs=_bsr_inputs()), + ) + return attention, fmha, q, metadata, args + + +def test_block_sparse_route_mode_is_derived_from_payload() -> None: + bsr = _bsr_inputs() + exact = _bitmask_inputs(proxy=False) + proxy = _bitmask_inputs(proxy=True) + + assert (bsr.sparse_format, bsr.use_proxy_routes) == ("bsr", False) + assert (exact.sparse_format, exact.use_proxy_routes) == ("bitmask", False) + assert (proxy.sparse_format, proxy.use_proxy_routes) == ("bitmask", True) + + +@pytest.mark.parametrize( + ("overrides", "message"), + [ + ({"block_indices": None}, "block_indptr and block_indices"), + ({"max_blocks_per_row": None}, "max_blocks_per_row"), + ( + {"exact_block_bits": torch.ones((1, 1, 1, 1), dtype=torch.uint32)}, + "exactly one route representation", + ), + ({"k_summary": torch.empty(0)}, "k_summary and v_summary"), + ], +) +def test_block_sparse_payload_rejects_ambiguous_combinations(overrides, message) -> None: + kwargs = { + "q_block_size": 64, + "kv_block_size": 64, + "max_blocks_per_row": 1, + "block_indptr": torch.tensor([[[0, 1]]], dtype=torch.int32), + "block_indices": torch.tensor([0], dtype=torch.int32), + } + kwargs.update(overrides) + + with pytest.raises((TypeError, ValueError), match=message): + BlockSparseForwardInputs(**kwargs) + + +def test_block_sparse_support_is_phase_specific_and_paged_proxy_is_rejected( + monkeypatch, +) -> None: + _attention, contiguous, q, k, v, metadata, args = _contiguous_case() + monkeypatch.setattr(contiguous, "_common_unsupported_reason", Mock(return_value=None)) + assert contiguous.is_supported(q, k, v, metadata, args, phase=FmhaPhase.CONTEXT) + assert not contiguous.is_supported(q, k, v, metadata, args, phase=FmhaPhase.GENERATION) + + _attention, paged, q, metadata, args = _paged_case() + _set_block_sparse_inputs(args, _bitmask_inputs(proxy=True)) + _supported, reason = paged._is_supported_with_reason( + q, None, None, metadata, args, phase=FmhaPhase.GENERATION + ) + assert not _supported + assert reason == "paged block-sparse attention only supports BSR exact routes" + + +def test_contiguous_proxy_routes_reject_causal_mask_before_planning(monkeypatch) -> None: + _attention, fmha, q, k, v, metadata, args = _contiguous_case() + _set_block_sparse_inputs(args, _bitmask_inputs(proxy=True)) + args.attention_mask = PredefinedAttentionMask.CAUSAL + monkeypatch.setattr(fmha, "_common_unsupported_reason", Mock(return_value=None)) + + supported, reason = fmha._is_supported_with_reason( + q, + k, + v, + metadata, + args, + phase=FmhaPhase.CONTEXT, + ) + + assert not supported + assert reason == "block-sparse proxy routes require mask_type='dense'" + + +@pytest.mark.parametrize("paged", [False, True]) +def test_block_sparse_support_rejects_invalid_static_kernel_profile( + monkeypatch, + paged, +) -> None: + if paged: + attention, fmha, q, metadata, args = _paged_case() + attention.head_dim = 64 + q = torch.zeros((2, 256), dtype=torch.bfloat16) + args.output = torch.empty((2, 128), dtype=q.dtype) + monkeypatch.setattr(fmha, "_common_unsupported_reason", Mock(return_value=None)) + supported, reason = fmha._is_supported_with_reason( + q, + None, + None, + metadata, + args, + phase=FmhaPhase.GENERATION, + ) + else: + attention, fmha, _q, _k, _v, metadata, args = _contiguous_case() + attention.head_dim = 64 + q = torch.zeros((128, 128), dtype=torch.bfloat16) + k = torch.zeros((512, 64), dtype=torch.bfloat16) + v = torch.zeros_like(k) + args.output = torch.empty_like(q) + monkeypatch.setattr(fmha, "_common_unsupported_reason", Mock(return_value=None)) + supported, reason = fmha._is_supported_with_reason( + q, + k, + v, + metadata, + args, + phase=FmhaPhase.CONTEXT, + ) + + assert not supported + assert reason == "block-sparse requires head_dim=128" + + +def test_contiguous_wrappers_cache_static_profile_and_keep_routes_live(monkeypatch) -> None: + _attention, fmha, q, k, v, _metadata, args = _contiguous_case() + wrapper = Mock() + factory = Mock(return_value=wrapper) + monkeypatch.setattr(block_sparse_fmha, "_BlockSparseTSWrapper", factory) + + bsr_inputs = [ + _bsr_inputs(), + BlockSparseForwardInputs( + q_block_size=64, + kv_block_size=64, + max_blocks_per_row=2, + block_indptr=torch.tensor([[[0, 1]], [[1, 4]]], dtype=torch.int32), + block_indices=torch.tensor([3, 1, 0, 2], dtype=torch.int32), + ), + ] + for inputs in bsr_inputs: + _set_block_sparse_inputs(args, inputs) + fmha._forward_contiguous(q, k, v, args) + + proxy_inputs = [_bitmask_inputs(proxy=True), _bitmask_inputs(proxy=True)] + for inputs in proxy_inputs: + _set_block_sparse_inputs(args, inputs) + fmha._forward_contiguous(q, k, v, args) + + assert factory.call_count == 2 + assert wrapper.plan.call_count == 2 + assert wrapper.plan.call_args_list[0].kwargs["sparse_format"] == "bsr" + assert wrapper.plan.call_args_list[0].kwargs["use_proxy_routes"] is False + assert wrapper.plan.call_args_list[1].kwargs["sparse_format"] == "bitmask" + assert wrapper.plan.call_args_list[1].kwargs["use_proxy_routes"] is True + assert wrapper.plan.call_args_list[1].kwargs["max_blocks_per_row"] == 4 + assert wrapper.run.call_count == 4 + + for call, inputs in zip(wrapper.run.call_args_list[:2], bsr_inputs): + assert call.kwargs["block_indptr"] is inputs.block_indptr + assert call.kwargs["block_indices"] is inputs.block_indices + for call, inputs in zip(wrapper.run.call_args_list[2:], proxy_inputs): + assert call.kwargs["exact_block_bits"] is inputs.exact_block_bits + assert call.kwargs["k_summary"] is inputs.k_summary + assert call.kwargs["v_summary"] is inputs.v_summary + + +def test_block_sparse_plan_key_includes_attention_head_topology() -> None: + inputs = _bitmask_inputs(proxy=True) + q = torch.empty((128, 256), dtype=torch.bfloat16) + first_attention = _Attention() + second_attention = _Attention() + second_attention.num_heads = 4 + first = block_sparse_fmha.PrimsTSBlockSparseFmha(first_attention) + second = block_sparse_fmha.PrimsTSBlockSparseFmha(second_attention) + + def _key(fmha): + return fmha._make_plan_key( + q, + inputs, + batch_size=1, + seq_len_q=128, + kv_capacity=256, + page_size=None, + mask_type="dense", + ) + + assert _key(first) != _key(second) + + +def test_block_sparse_plan_cache_is_shared_only_when_explicitly_bound() -> None: + first = block_sparse_fmha.PrimsTSBlockSparseFmha(_Attention()) + second = block_sparse_fmha.PrimsTSBlockSparseFmha(_Attention()) + + assert first._contiguous_wrappers is not second._contiguous_wrappers + assert first._paged_wrappers is not second._paged_wrappers + + cache_state = {} + first.bind_plan_cache(cache_state) + second.bind_plan_cache(cache_state) + + assert first._contiguous_wrappers is second._contiguous_wrappers + assert first._paged_wrappers is second._paged_wrappers + assert cache_state == { + "contiguous_wrappers": {}, + "paged_wrappers": {}, + } + + +def test_paged_wrapper_uses_zero_copy_padded_row_stride_block_tables(monkeypatch) -> None: + attention, fmha, q, metadata, args = _paged_case() + wrapper = Mock() + monkeypatch.setattr(block_sparse_fmha, "_BlockSparsePagedTSWrapper", Mock(return_value=wrapper)) + monkeypatch.setattr(block_sparse_fmha, "get_kv_page_offset", Mock(return_value=8)) + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", Mock(return_value=False)) + q_processed = torch.zeros((2, 2, 128), dtype=torch.bfloat16) + kv_pool = torch.empty((16, 1, 64, 128), dtype=torch.bfloat16) + block_tables = metadata.kv_cache_block_offsets[0] + empty = torch.empty(0, dtype=torch.uint8) + preprocessed = (q_processed, kv_pool, block_tables, None, 1.0, 1.0) + ( + empty, + None, + 1, + 256, + -1, + False, + ) + monkeypatch.setattr(fmha, "_run_generation_preprocess", Mock(return_value=preprocessed)) + params = FmhaParams( + attn=attention, + meta=metadata, + fwd=args, + workspace=torch.empty(0, dtype=torch.uint8), + qkv_input=q, + context_buf=args.output, + sequence_lengths=torch.tensor([129, 193], dtype=torch.int32), + input_seq_length=1, + tokens_per_block=64, + num_requests=2, + ) + expected_block_tables = block_tables[:2, 0, :] + snapshots = [] + + def snapshot(*_args, **kwargs): + snapshots.append( + ( + kwargs["seq_lens_kv"].clone(), + kwargs["block_tables"], + kwargs["block_tables"].clone(), + kwargs["block_indptr"], + kwargs["block_indices"], + ) + ) + + wrapper.run.side_effect = snapshot + first_inputs = _get_block_sparse_inputs(args) + fmha.run_generation(params) + block_tables[:, 0].add_(10) + params.sequence_lengths = torch.tensor([130, 194], dtype=torch.int32) + _set_block_sparse_inputs( + args, + BlockSparseForwardInputs( + q_block_size=64, + kv_block_size=64, + max_blocks_per_row=2, + block_indptr=torch.tensor([[[0, 1]], [[1, 4]]], dtype=torch.int32), + block_indices=torch.tensor([3, 2, 1, 0], dtype=torch.int32), + ), + ) + fmha.run_generation(params) + + wrapper.plan.assert_called_once() + assert wrapper.run.call_count == 2 + torch.testing.assert_close(snapshots[0][0], torch.tensor([129, 193], dtype=torch.int32)) + torch.testing.assert_close(snapshots[1][0], torch.tensor([130, 194], dtype=torch.int32)) + assert snapshots[0][1].data_ptr() == expected_block_tables.data_ptr() + assert snapshots[1][1].data_ptr() == expected_block_tables.data_ptr() + assert snapshots[0][1].shape == (2, 4) + assert snapshots[0][1].stride() == (8, 1) + torch.testing.assert_close(snapshots[0][2], torch.arange(8, dtype=torch.int32).view(2, 4)) + torch.testing.assert_close(snapshots[1][2], torch.arange(8, dtype=torch.int32).view(2, 4) + 10) + assert snapshots[0][3] is first_inputs.block_indptr + assert snapshots[1][3] is _get_block_sparse_inputs(args).block_indptr + + +def test_paged_block_tables_remain_live_across_graph_replay(monkeypatch) -> None: + attention, fmha, q, metadata, args = _paged_case() + wrapper = Mock() + monkeypatch.setattr(block_sparse_fmha, "_BlockSparsePagedTSWrapper", Mock(return_value=wrapper)) + monkeypatch.setattr(block_sparse_fmha, "get_kv_page_offset", Mock(return_value=8)) + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", Mock(return_value=True)) + q_processed = torch.zeros((2, 2, 128), dtype=torch.bfloat16) + kv_pool = torch.empty((16, 1, 64, 128), dtype=torch.bfloat16) + block_tables = metadata.kv_cache_block_offsets[0] + empty = torch.empty(0, dtype=torch.uint8) + preprocessed = (q_processed, kv_pool, block_tables, None, 1.0, 1.0) + ( + empty, + None, + 1, + 256, + -1, + False, + ) + monkeypatch.setattr(fmha, "_run_generation_preprocess", Mock(return_value=preprocessed)) + params = FmhaParams( + attn=attention, + meta=metadata, + fwd=args, + workspace=torch.empty(0, dtype=torch.uint8), + qkv_input=q, + context_buf=args.output, + sequence_lengths=torch.tensor([129, 193], dtype=torch.int32), + input_seq_length=1, + tokens_per_block=64, + num_requests=2, + ) + seen = [] + + def snapshot(*_args, **kwargs): + seen.append((kwargs["block_tables"].data_ptr(), kwargs["block_tables"].clone())) + + wrapper.run.side_effect = snapshot + fmha.run_generation(params) + block_tables[:, 0, :].add_(10) + block_tables[:, 1, :].fill_(-1) + fmha.run_generation(params) + + assert seen[0][0] == seen[1][0] == block_tables.data_ptr() + torch.testing.assert_close(seen[0][1], torch.arange(8, dtype=torch.int32).view(2, 4)) + torch.testing.assert_close(seen[1][1], torch.arange(8, dtype=torch.int32).view(2, 4) + 10) + + +def test_prepare_workspace_checks_capture_before_resize(monkeypatch) -> None: + _attention, fmha, _q, _metadata, _args = _paged_case() + query_device = torch.device("cuda:1") + q = SimpleNamespace( + device=query_device, + dtype=torch.bfloat16, + shape=(2, 512), + ) + metadata = SimpleNamespace( + kv_cache_manager=object(), + kv_cache_block_offsets=SimpleNamespace(device=query_device, shape=(1, 2, 4)), + max_num_requests=2, + tokens_per_block=64, + num_generations=2, + ) + workspace = torch.empty(0, dtype=torch.uint8) + monkeypatch.setattr( + fmha, + "_get_generation_workspace_layout", + Mock(return_value={"total_size": 16}), + ) + fmha._multi_processor_count = 1 + device_scope = Mock(return_value=nullcontext()) + monkeypatch.setattr(torch.cuda, "device", device_scope) + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", Mock(return_value=True)) + + with pytest.raises(RuntimeError, match="workspace must be sized"): + fmha.prepare_workspace(q, None, None, metadata, _args, workspace) + + device_scope.assert_called_once_with(query_device) + assert workspace.numel() == 0 + + +def test_prepare_workspace_skips_generation_layout_for_contiguous_requests(monkeypatch) -> None: + _attention, fmha, q, _k, _v, metadata, args = _contiguous_case() + layout = Mock() + monkeypatch.setattr(fmha, "_get_generation_workspace_layout", layout) + monkeypatch.setattr(torch.cuda, "device", Mock(return_value=nullcontext())) + fmha._multi_processor_count = 1 + workspace = torch.empty(0, dtype=torch.uint8) + + fmha.prepare_workspace(q, None, None, metadata, args, workspace) + + layout.assert_not_called() + assert workspace.numel() == 0 + + +@_REQUIRES_PRIMTS_GPU +@torch.no_grad() +def test_real_gpu_raw_routes_and_token_mask_match_reference() -> None: + torch.manual_seed(1234) + q = torch.randn((1, 128, 1, 128), device="cuda", dtype=torch.float16) + k = torch.randn((1, 256, 1, 128), device="cuda", dtype=torch.float16) + v = torch.randn_like(k) + token_mask = torch.ones(256, device="cuda", dtype=torch.bool) + token_mask[[1, 63, 64, 95, 129, 190, 255]] = False + inputs = BlockSparseForwardInputs( + q_block_size=64, + kv_block_size=64, + max_blocks_per_row=3, + block_indptr=torch.tensor([[[0, 2, 5]]], device="cuda", dtype=torch.int32), + block_indices=torch.tensor([0, 2, 0, 1, 3], device="cuda", dtype=torch.int32), + kv_valid_bits=_pack_token_mask(token_mask), + ) + sm_scale = 128**-0.5 + + key_blocks = torch.arange(256, device="cuda") // 64 + allowed = torch.zeros((128, 256), device="cuda", dtype=torch.bool) + for row, selected_blocks in enumerate(((0, 2), (0, 1, 3))): + selected = torch.tensor(selected_blocks, device="cuda") + allowed[row * 64 : (row + 1) * 64] = torch.isin(key_blocks, selected) & token_mask + scores = (q[0, :, 0].float() @ k[0, :, 0].float().T) * sm_scale + expected = ( + torch.softmax(scores.masked_fill(~allowed, float("-inf")), dim=-1) @ v[0, :, 0].float() + ).to(q.dtype)[None, :, None, :] + + actual = prims_ts.block_sparse_attention( + q, + k, + v, + block_indptr=inputs.block_indptr, + block_indices=inputs.block_indices, + q_block_size=inputs.q_block_size, + kv_block_size=inputs.kv_block_size, + kv_valid_bits=inputs.kv_valid_bits, + sm_scale=sm_scale, + ) + torch.testing.assert_close(actual, expected, rtol=1e-2, atol=1e-2) + + +@_REQUIRES_PRIMTS_GPU +@torch.no_grad() +def test_real_gpu_proxy_adapter_replays_live_routes_and_summaries() -> None: + torch.manual_seed(20260901) + q = torch.randn((1, 64, 1, 128), device="cuda", dtype=torch.bfloat16) + k = torch.randn((1, 192, 1, 128), device="cuda", dtype=torch.bfloat16) + v = torch.randn_like(k) + k_blocks = k.float().view(1, 3, 64, 1, 128) + v_blocks = v.float().view(1, 3, 64, 1, 128) + initial_k_summary = k_blocks.mean(dim=2).to(k.dtype) + initial_v_summary = v_blocks.sum(dim=2).to(v.dtype) + live_k_summary = initial_k_summary.clone() + live_v_summary = initial_v_summary.clone() + live_exact_bits = torch.tensor([[[[1]]]], device="cuda", dtype=torch.uint32) + + attention = _Attention() + attention.num_heads = attention.num_kv_heads = 1 + fmha = block_sparse_fmha.PrimsTSBlockSparseFmha(attention) + output = torch.empty_like(q).view(64, 128) + args = AttentionForwardArgs( + output=output, + attention_input_type=AttentionInputType.context_only, + attention_mask=PredefinedAttentionMask.FULL, + sparse_runtime_params=SparseRuntimeParams( + block_sparse_inputs=BlockSparseForwardInputs( + q_block_size=64, + kv_block_size=64, + exact_block_bits=live_exact_bits, + k_summary=live_k_summary, + v_summary=live_v_summary, + ), + ), + ) + metadata = SimpleNamespace( + is_cross=False, + kv_cache_manager=None, + seq_lens=torch.tensor([64], dtype=torch.int32), + ) + flat_q, flat_k, flat_v = (tensor.flatten(0, 2) for tensor in (q, k, v)) + + fmha.forward(flat_q, flat_k, flat_v, metadata, args) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + fmha.forward(flat_q, flat_k, flat_v, metadata, args) + + graph.replay() + torch.cuda.synchronize() + expected = _proxy_reference(q, k, v, live_k_summary, live_v_summary, exact_block=0) + torch.testing.assert_close(output.view_as(q), expected, rtol=2e-2, atol=2e-2) + + live_exact_bits.fill_(1 << 2) + live_k_summary.copy_((initial_k_summary.float() * 0.5 + 0.125).to(k.dtype)) + live_v_summary.copy_((initial_v_summary.float() * -0.25).to(v.dtype)) + graph.replay() + torch.cuda.synchronize() + expected = _proxy_reference(q, k, v, live_k_summary, live_v_summary, exact_block=2) + torch.testing.assert_close(output.view_as(q), expected, rtol=2e-2, atol=2e-2) + + +@_REQUIRES_PRIMTS_GPU +@torch.no_grad() +def test_real_gpu_paged_routes_use_live_length_below_capacity() -> None: + torch.manual_seed(7) + q = torch.randn((1, 64, 1, 128), device="cuda", dtype=torch.float16) + k_cache = torch.randn((4, 1, 64, 128), device="cuda", dtype=torch.float16) + v_cache = torch.randn_like(k_cache) + page_indices = torch.tensor([2, 0, 3, 1], device="cuda", dtype=torch.int32) + seq_lens_kv = torch.tensor([160], device="cuda", dtype=torch.int32) + inputs = BlockSparseForwardInputs( + q_block_size=64, + kv_block_size=64, + max_blocks_per_row=2, + block_indptr=torch.tensor([[[0, 2]]], device="cuda", dtype=torch.int32), + block_indices=torch.tensor([0, 2], device="cuda", dtype=torch.int32), + ) + sm_scale = 128**-0.5 + + actual = prims_ts.block_sparse_attention_with_paged_kv_cache( + q, + (k_cache, v_cache), + block_tables=page_indices.view(1, 4), + seq_lens_kv=seq_lens_kv, + block_indptr=inputs.block_indptr, + block_indices=inputs.block_indices, + max_seq_len_kv=256, + q_block_size=inputs.q_block_size, + kv_block_size=inputs.kv_block_size, + sm_scale=sm_scale, + ) + + logical_k = k_cache.index_select(0, page_indices.long()).reshape(256, 1, 128) + logical_v = v_cache.index_select(0, page_indices.long()).reshape(256, 1, 128) + allowed = torch.zeros(256, device="cuda", dtype=torch.bool) + allowed[:64] = True + allowed[128:160] = True + scores = (q[0, :, 0].float() @ logical_k[:, 0].float().T) * sm_scale + expected = ( + torch.softmax(scores.masked_fill(~allowed, float("-inf")), dim=-1) @ logical_v[:, 0].float() + ).to(q.dtype)[None, :, None, :] + + torch.testing.assert_close(actual, expected, rtol=1e-2, atol=1e-2) + + +# Block-sparse FMHA support matrix (generic PrimTS block-sparse kernels): +# +# GPU architecture SM100 and SM103 +# Compute phase Contiguous prefill with separate Q/K/V and no KV cache; +# fixed-query generation over a paged HND KV cache +# Attention type MHA, MQA, and GQA; num_heads % num_kv_heads == 0 +# Q/K/V head dimension 128 +# Model dtype BF16 or FP16; Q, K/V, and output share one dtype +# Route format BSR (block_indptr [B, Hkv, num_q_blocks + 1] plus flat +# block IDs) or a packed block bitmask, optionally with +# K/V block summaries for proxy routes (contiguous only) +# KV block size 8, 16, 32, or a positive multiple of 64 (contiguous); +# a positive multiple of 64 (paged) +# KV-cache layout Paged HND; page size 64 or 128 (a page holds at least +# one 64-token route fragment) +# Attention semantics Dense or causal; proxy routes require dense + +_REAL_GPU_DTYPES = (torch.float16, torch.bfloat16) +# (num_heads, num_kv_heads): MHA, GQA, and MQA head topologies. +_REAL_GPU_HEAD_TOPOLOGIES = ((1, 1), (8, 8), (8, 2), (8, 1)) +_REAL_GPU_KV_BLOCK_SIZES = (64, 128) +_REAL_GPU_PAGE_SIZES = (64, 128) + + +def _bsr_routes_per_kv_head( + num_kv_heads: int, + num_q_blocks: int, + num_kv_blocks: int, + num_selected: int, +) -> tuple[torch.Tensor, torch.Tensor, list[list[list[int]]]]: + """Build head-dependent, increasing BSR routes and the selection per (head, q block).""" + selections: list[list[list[int]]] = [] + indptr = torch.zeros((1, num_kv_heads, num_q_blocks + 1), dtype=torch.int32) + indices: list[int] = [] + for head_idx in range(num_kv_heads): + head_selection = [] + indptr[0, head_idx, 0] = len(indices) + for q_block in range(num_q_blocks): + start = (head_idx + q_block) % num_kv_blocks + selected = sorted({(start + step) % num_kv_blocks for step in range(num_selected)}) + indices.extend(selected) + indptr[0, head_idx, q_block + 1] = len(indices) + head_selection.append(selected) + selections.append(head_selection) + return indptr.cuda(), torch.tensor(indices, dtype=torch.int32, device="cuda"), selections + + +def _reference_block_sparse( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + selections: list[list[list[int]]], + *, + q_block_size: int, + kv_block_size: int, + seq_len_kv: int, +) -> torch.Tensor: + """Dense reference with a per-(kv head, q block) block mask; Q/K/V are [1, S, H, D].""" + num_heads, num_kv_heads = q.shape[2], k.shape[2] + heads_per_kv = num_heads // num_kv_heads + key_blocks = torch.arange(seq_len_kv, device=q.device) // kv_block_size + outputs = [] + for head_idx in range(num_heads): + kv_head = head_idx // heads_per_kv + allowed = torch.zeros((q.shape[1], seq_len_kv), device=q.device, dtype=torch.bool) + for q_block, selected in enumerate(selections[kv_head]): + rows = slice(q_block * q_block_size, (q_block + 1) * q_block_size) + allowed[rows] = torch.isin(key_blocks, torch.tensor(selected, device=q.device)) + scores = (q[0, :, head_idx].float() @ k[0, :seq_len_kv, kv_head].float().T) * ( + q.shape[-1] ** -0.5 + ) + probs = torch.softmax(scores.masked_fill(~allowed, float("-inf")), dim=-1) + outputs.append(probs @ v[0, :seq_len_kv, kv_head].float()) + return torch.stack(outputs, dim=1).to(q.dtype)[None] + + +@_REQUIRES_PRIMTS_GPU +@torch.no_grad() +@pytest.mark.parametrize("dtype", _REAL_GPU_DTYPES, ids=lambda d: str(d).removeprefix("torch.")) +@pytest.mark.parametrize("num_heads,num_kv_heads", _REAL_GPU_HEAD_TOPOLOGIES) +@pytest.mark.parametrize("kv_block_size", _REAL_GPU_KV_BLOCK_SIZES) +def test_real_gpu_contiguous_routes_match_reference_across_topologies( + dtype: torch.dtype, num_heads: int, num_kv_heads: int, kv_block_size: int +) -> None: + torch.manual_seed(7) + seq_len_q, seq_len_kv, q_block_size = 128, 512, 64 + q = torch.randn((1, seq_len_q, num_heads, 128), device="cuda", dtype=dtype) + k = torch.randn((1, seq_len_kv, num_kv_heads, 128), device="cuda", dtype=dtype) + v = torch.randn_like(k) + num_kv_blocks = seq_len_kv // kv_block_size + block_indptr, block_indices, selections = _bsr_routes_per_kv_head( + num_kv_heads, seq_len_q // q_block_size, num_kv_blocks, num_selected=2 + ) + + actual = prims_ts.block_sparse_attention( + q, + k, + v, + block_indptr=block_indptr, + block_indices=block_indices, + q_block_size=q_block_size, + kv_block_size=kv_block_size, + ) + + expected = _reference_block_sparse( + q, + k, + v, + selections, + q_block_size=q_block_size, + kv_block_size=kv_block_size, + seq_len_kv=seq_len_kv, + ) + torch.testing.assert_close(actual, expected, rtol=2e-2, atol=2e-2) + + +@_REQUIRES_PRIMTS_GPU +@torch.no_grad() +@pytest.mark.parametrize("dtype", _REAL_GPU_DTYPES, ids=lambda d: str(d).removeprefix("torch.")) +@pytest.mark.parametrize("num_heads,num_kv_heads", _REAL_GPU_HEAD_TOPOLOGIES) +@pytest.mark.parametrize("kv_block_size", _REAL_GPU_KV_BLOCK_SIZES) +@pytest.mark.parametrize("page_size", _REAL_GPU_PAGE_SIZES) +def test_real_gpu_paged_routes_match_reference_across_topologies( + dtype: torch.dtype, num_heads: int, num_kv_heads: int, kv_block_size: int, page_size: int +) -> None: + torch.manual_seed(11) + max_seq_len_kv, seq_len_kv = 512, 400 + num_pages = max_seq_len_kv // page_size + q = torch.randn((1, 1, num_heads, 128), device="cuda", dtype=dtype) + k_cache = torch.randn((num_pages, num_kv_heads, page_size, 128), device="cuda", dtype=dtype) + v_cache = torch.randn_like(k_cache) + page_indices = torch.randperm(num_pages, device="cuda").to(torch.int32) + num_kv_blocks = -(-seq_len_kv // kv_block_size) + block_indptr, block_indices, selections = _bsr_routes_per_kv_head( + num_kv_heads, 1, num_kv_blocks, num_selected=3 + ) + + actual = prims_ts.block_sparse_attention_with_paged_kv_cache( + q, + (k_cache, v_cache), + block_tables=page_indices.view(1, num_pages), + seq_lens_kv=torch.tensor([seq_len_kv], device="cuda", dtype=torch.int32), + block_indptr=block_indptr, + block_indices=block_indices, + max_seq_len_kv=max_seq_len_kv, + q_block_size=64, + kv_block_size=kv_block_size, + ) + + logical_k = k_cache.index_select(0, page_indices.long()).permute(0, 2, 1, 3) + logical_k = logical_k.reshape(1, max_seq_len_kv, num_kv_heads, 128) + logical_v = v_cache.index_select(0, page_indices.long()).permute(0, 2, 1, 3) + logical_v = logical_v.reshape(1, max_seq_len_kv, num_kv_heads, 128) + expected = _reference_block_sparse( + q, + logical_k, + logical_v, + selections, + q_block_size=64, + kv_block_size=kv_block_size, + seq_len_kv=seq_len_kv, + ) + torch.testing.assert_close(actual, expected, rtol=2e-2, atol=2e-2) + + +def test_paged_support_rejects_pages_smaller_than_a_route_fragment() -> None: + page_size = 32 + _attention, fmha, q, metadata, args = _paged_case() + metadata.tokens_per_block = page_size + metadata.max_seq_len = 4 * page_size + metadata.kv_lens_runtime = torch.tensor([page_size + 1, 3 * page_size], dtype=torch.int32) + args = AttentionForwardArgs( + output=args.output, + attention_input_type=args.attention_input_type, + attention_mask=args.attention_mask, + attention_window_size=metadata.max_seq_len, + is_fused_qkv=True, + sparse_runtime_params=args.sparse_runtime_params, + ) + + reason = fmha._paged_unsupported_reason(q, metadata, args) + + assert reason == "atom_size must not exceed page_size" diff --git a/tests/unittest/_torch/attention/sparse/test_sparse_attention.py b/tests/unittest/_torch/attention/sparse/test_sparse_attention.py index 40a151daa252..6cf7b15fcc38 100644 --- a/tests/unittest/_torch/attention/sparse/test_sparse_attention.py +++ b/tests/unittest/_torch/attention/sparse/test_sparse_attention.py @@ -21,11 +21,16 @@ """ from types import ModuleType -from unittest.mock import Mock +from unittest.mock import Mock, patch +import pytest import torch -from tensorrt_llm._torch.attention.backends.interface import AttentionForwardArgs +from tensorrt_llm._torch.attention.backends import trtllm as trtllm_backend +from tensorrt_llm._torch.attention.backends.interface import ( + AttentionForwardArgs, + AttentionInputType, +) from tensorrt_llm._torch.attention.backends.sparse.hooks import ( AttentionSparseHooks, MLASparseHooks, @@ -35,8 +40,13 @@ register_attention_sparse_hooks, register_mla_sparse_hooks, ) -from tensorrt_llm._torch.attention.backends.sparse.params import SparseParams, SparseRuntimeParams -from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttention +from tensorrt_llm._torch.attention.backends.sparse.params import ( + BlockSparseForwardInputs, + SparseBackendForwardArgs, + SparseParams, + SparseRuntimeParams, +) +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttention, TrtllmAttentionMetadata from tensorrt_llm._torch.attention.mla import MLA @@ -72,7 +82,7 @@ def test_prepare_sparse_runtime_params_from_predictions() -> None: ) runtime_params = prepare_sparse_runtime_params( - attention, torch.empty(0), None, None, forward_args + attention, torch.empty(0), None, None, None, forward_args ) assert runtime_params.sparse_kv_indices is attention._sparse_kv_indices @@ -144,12 +154,249 @@ def test_mla_backend_only_forward_uses_default_path() -> None: ) -def test_prepare_sparse_runtime_params_without_predictions() -> None: +@pytest.mark.parametrize( + "sparse_params", [None, _StubSparseParams()], ids=["dense_backend", "sparse_backend"] +) +def test_prepare_sparse_runtime_params_without_predictions(sparse_params) -> None: attention = TrtllmAttention.__new__(TrtllmAttention) - attention.sparse_params = _StubSparseParams() + attention.sparse_params = sparse_params runtime_params = prepare_sparse_runtime_params( - attention, torch.empty(0), None, None, AttentionForwardArgs() + attention, torch.empty(0), None, None, None, AttentionForwardArgs() ) assert runtime_params == SparseRuntimeParams() + + +def test_prepare_sparse_runtime_params_runs_index_hooks_once() -> None: + attention = _StaticPredictionAttention.__new__(_StaticPredictionAttention) + attention.sparse_params = _StubSparseParams() + sparse_kv_indices = torch.tensor([1], dtype=torch.int32) + sparse_kv_offsets = torch.tensor([0, 1], dtype=torch.int32) + sparse_attn_indices = torch.tensor([2], dtype=torch.int32) + sparse_attn_offsets = torch.tensor([0, 1], dtype=torch.int32) + attention.sparse_kv_predict = Mock(return_value=(sparse_kv_indices, sparse_kv_offsets)) + attention.sparse_attn_predict = Mock(return_value=(sparse_attn_indices, sparse_attn_offsets)) + q = torch.empty((1, 4)) + k = torch.empty((1, 4)) + v = torch.empty((1, 4)) + metadata = Mock() + caller_kv_lens = torch.tensor([3]) + forward_args = AttentionForwardArgs( + sparse_runtime_params=SparseRuntimeParams(sparse_attn_kv_lens=caller_kv_lens) + ) + + runtime_params = prepare_sparse_runtime_params(attention, q, k, v, metadata, forward_args) + + assert isinstance(runtime_params, SparseRuntimeParams) + assert runtime_params.block_sparse_inputs is None + assert runtime_params.sparse_kv_indices is sparse_kv_indices + assert runtime_params.sparse_kv_offsets is sparse_kv_offsets + assert runtime_params.sparse_attn_indices is sparse_attn_indices + assert runtime_params.sparse_attn_offsets is sparse_attn_offsets + assert runtime_params.sparse_attn_indices_block_size == 1 + assert runtime_params.sparse_attn_kv_lens is caller_kv_lens + attention.sparse_kv_predict.assert_called_once_with(q, k, metadata, forward_args) + attention.sparse_attn_predict.assert_called_once_with(q, k, metadata, forward_args) + + +def test_prepare_sparse_runtime_params_schedules_skip_softmax_thresholds() -> None: + from tensorrt_llm._torch.attention.backends.sparse.skip_softmax import SkipSoftmaxParams + + attention = TrtllmAttention.__new__(TrtllmAttention) + attention.sparse_params = SkipSoftmaxParams() + scheduler = attention.sparse_params.scheduler + timestep = torch.tensor(0.5) + forward_args = AttentionForwardArgs(timestep=timestep) + + with patch.object( + scheduler, "get_runtime_params", wraps=scheduler.get_runtime_params + ) as schedule: + runtime_params = prepare_sparse_runtime_params( + attention, torch.empty(0), None, None, None, forward_args + ) + + schedule.assert_called_once_with(runtime_params=SparseRuntimeParams(), timestep=timestep) + assert runtime_params == scheduler.get_runtime_params(timestep=timestep) + + +def _make_block_sparse_inputs() -> BlockSparseForwardInputs: + return BlockSparseForwardInputs( + q_block_size=64, + kv_block_size=64, + exact_block_bits=torch.zeros((1, 1), dtype=torch.int32), + ) + + +def test_block_sparse_attn_predict_hands_through_backend_args() -> None: + attention = TrtllmAttention.__new__(TrtllmAttention) + attention.sparse_params = None + block_sparse_inputs = _make_block_sparse_inputs() + forward_args = AttentionForwardArgs( + sparse_backend_args=SparseBackendForwardArgs(block_sparse_inputs=block_sparse_inputs) + ) + + runtime_params = prepare_sparse_runtime_params( + attention, torch.empty(0), None, None, None, forward_args + ) + + assert runtime_params.block_sparse_inputs is block_sparse_inputs + assert runtime_params == SparseRuntimeParams(block_sparse_inputs=block_sparse_inputs) + + +def test_block_sparse_attn_predict_override_composes_with_index_predictors() -> None: + attention = _StaticPredictionAttention.__new__(_StaticPredictionAttention) + attention.sparse_params = _StubSparseParams() + sparse_attn_indices = torch.tensor([2], dtype=torch.int32) + sparse_attn_offsets = torch.tensor([0, 1], dtype=torch.int32) + attention.sparse_kv_predict = Mock(return_value=(None, None)) + attention.sparse_attn_predict = Mock(return_value=(sparse_attn_indices, sparse_attn_offsets)) + block_sparse_inputs = _make_block_sparse_inputs() + attention.block_sparse_attn_predict = Mock(return_value=block_sparse_inputs) + q = torch.empty((1, 4)) + k = torch.empty((1, 4)) + v = torch.empty((1, 4)) + metadata = Mock() + forward_args = AttentionForwardArgs() + + runtime_params = prepare_sparse_runtime_params(attention, q, k, v, metadata, forward_args) + + assert runtime_params.block_sparse_inputs is block_sparse_inputs + assert runtime_params.sparse_attn_indices is sparse_attn_indices + assert runtime_params.sparse_attn_indices_block_size == 1 + attention.block_sparse_attn_predict.assert_called_once_with(q, k, v, metadata, forward_args) + + +def test_attention_forward_args_default_to_empty_sparse_runtime_params() -> None: + assert AttentionForwardArgs().sparse_runtime_params == SparseRuntimeParams() + + +class _StopAfterShapeValidation(Exception): + pass + + +def _make_sparse_prediction_forward_backend() -> TrtllmAttention: + attention = _StaticPredictionAttention.__new__(_StaticPredictionAttention) + attention.sparse_params = None + attention.is_mla_enable = False + attention.num_heads = 1 + attention.num_kv_heads = 1 + attention.head_dim = 4 + attention.get_local_layer_idx = Mock(return_value=1) + attention._ensure_rope_table_size = Mock(side_effect=_StopAfterShapeValidation) + return attention + + +def _make_sparse_prediction_forward_metadata() -> TrtllmAttentionMetadata: + metadata = object.__new__(TrtllmAttentionMetadata) + seq_lens = torch.tensor([2], dtype=torch.int32) + metadata._seq_lens = seq_lens + metadata._seq_lens_kv = seq_lens + metadata._seq_lens_cuda = None + metadata.kv_cache_manager = None + metadata._max_seq_len_storage = 4 + metadata.use_paged_context_fmha = False + metadata.cu_q_seqlens = None + metadata.cu_kv_seqlens = None + metadata.enable_flash_mla = False + metadata.spec_bl_tree_first_sparse_mask_offset_kv = None + metadata.kv_lens_cuda_runtime = torch.tensor([2], dtype=torch.int32) + metadata.kv_lens_runtime = torch.tensor([2], dtype=torch.int32) + metadata.prompt_lens_cuda_runtime = torch.tensor([2], dtype=torch.int32) + metadata.prompt_lens_cpu_runtime = torch.tensor([2], dtype=torch.int32) + metadata.host_request_types_runtime = torch.tensor([0], dtype=torch.int32) + return metadata + + +def test_forward_materializes_dynamic_block_sparse_prediction_before_shape_validation() -> None: + attention = _make_sparse_prediction_forward_backend() + block_sparse_inputs = _make_block_sparse_inputs() + prediction = SparseRuntimeParams( + sparse_attn_kv_lens=torch.tensor([4]), + block_sparse_inputs=block_sparse_inputs, + ) + q = torch.empty((2, 4)) + k = torch.empty((2, 4)) + v = torch.empty((2, 4)) + metadata = _make_sparse_prediction_forward_metadata() + forward_args = AttentionForwardArgs( + output=torch.empty_like(q), + attention_input_type=AttentionInputType.context_only, + ) + + with patch.object( + trtllm_backend, "prepare_sparse_runtime_params", return_value=prediction + ) as prepare: + for _ in range(2): + with pytest.raises(_StopAfterShapeValidation): + attention.forward(q, k, v, metadata, forward_args) + + assert prepare.call_count == 2 + prepare.assert_called_with(attention, q, k, v, metadata, forward_args) + assert attention._ensure_rope_table_size.call_count == 2 + assert forward_args.sparse_runtime_params is prediction + + +@pytest.mark.parametrize("has_block_sparse_inputs", [False, True]) +def test_forward_assigns_prepared_sparse_runtime_params( + has_block_sparse_inputs: bool, +) -> None: + attention = _make_sparse_prediction_forward_backend() + block_sparse_inputs = _make_block_sparse_inputs() if has_block_sparse_inputs else None + prediction = SparseRuntimeParams( + sparse_attn_kv_lens=torch.tensor([2]), + block_sparse_inputs=block_sparse_inputs, + ) + caller_params = SparseRuntimeParams(sparse_attn_kv_lens=torch.tensor([7])) + q = torch.empty((2, 4)) + k = torch.empty((2, 4)) + v = torch.empty((2, 4)) + metadata = _make_sparse_prediction_forward_metadata() + forward_args = AttentionForwardArgs( + output=torch.empty_like(q), + attention_input_type=AttentionInputType.context_only, + sparse_runtime_params=caller_params, + ) + + with patch.object( + trtllm_backend, "prepare_sparse_runtime_params", return_value=prediction + ) as prepare: + for _ in range(2): + with pytest.raises(_StopAfterShapeValidation): + attention.forward(q, k, v, metadata, forward_args) + + assert prepare.call_count == 2 + assert attention._ensure_rope_table_size.call_count == 2 + assert forward_args.sparse_runtime_params is prediction + + +@pytest.mark.parametrize("backend_name", ["FLASHINFER", "unknown"]) +def test_sparse_attention_backend_fallback_does_not_redispatch( + backend_name: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from tensorrt_llm._torch.attention.backends import utils as attention_backend_utils + from tensorrt_llm._torch.attention.backends.sparse.skip_softmax import SkipSoftmaxParams + + monkeypatch.setattr(attention_backend_utils, "IS_FLASHINFER_AVAILABLE", False) + + with patch.object( + attention_backend_utils, + "get_trtllm_sparse_attn_attention_backend", + ) as trtllm_sparse_resolver: + backend = attention_backend_utils.get_attention_backend( + backend_name, sparse_params=SkipSoftmaxParams() + ) + + assert backend is TrtllmAttention + trtllm_sparse_resolver.assert_not_called() + + +@pytest.mark.parametrize("backend_name", ["FLASHINFER", "unknown"]) +def test_trtllm_fallback_without_sparse_params_remains_dense( + backend_name: str, monkeypatch: pytest.MonkeyPatch +) -> None: + from tensorrt_llm._torch.attention.backends import utils as attention_backend_utils + + monkeypatch.setattr(attention_backend_utils, "IS_FLASHINFER_AVAILABLE", False) + + assert attention_backend_utils.get_attention_backend(backend_name) is TrtllmAttention diff --git a/tests/unittest/_torch/attention/test_attention_op_sync.py b/tests/unittest/_torch/attention/test_attention_op_sync.py index ba5c75f3d7b6..ccf11544be53 100644 --- a/tests/unittest/_torch/attention/test_attention_op_sync.py +++ b/tests/unittest/_torch/attention/test_attention_op_sync.py @@ -41,7 +41,7 @@ import textwrap import typing from dataclasses import fields -from types import SimpleNamespace +from types import SimpleNamespace, UnionType import pytest import torch @@ -382,7 +382,19 @@ def _dataclass_field_type(cls, name: str): return None if f is None: return None - return f.type if not isinstance(f.type, str) else None + if isinstance(f.type, str): + return None + return _unwrap_optional(f.type) + + +def _unwrap_optional(py_type): + """Return the payload type for ``Optional[T]`` annotations.""" + origin = typing.get_origin(py_type) + if origin in (typing.Union, UnionType): + args = [arg for arg in typing.get_args(py_type) if arg is not type(None)] + if len(args) == 1: + return args[0] + return py_type def _resolve_path(root_cls, path: tuple[str, ...]): @@ -404,7 +416,7 @@ def _python_category(py_type) -> str: confidently (the type check is then skipped for that kwarg).""" # Unwrap Optional[X] / Union[X, None]. origin = typing.get_origin(py_type) - if origin is typing.Union: + if origin in (typing.Union, UnionType): args = [a for a in typing.get_args(py_type) if a is not type(None)] if len(args) == 1: return _python_category(args[0]) @@ -558,7 +570,7 @@ def _verify_consumed(cls, chains: set[tuple[str, ...]], excluded=frozenset()): for f in fields(cls): if f.name in excluded: continue - ftype = f.type if not isinstance(f.type, str) else None + ftype = _dataclass_field_type(cls, f.name) if ftype is not None and dataclasses.is_dataclass(ftype): sub = {p[1:] for p in chains if len(p) >= 2 and p[0] == f.name} assert sub, ( @@ -566,7 +578,7 @@ def _verify_consumed(cls, chains: set[tuple[str, ...]], excluded=frozenset()): f"declared but `{f.name}.` is never read at the " f"call site." ) - _verify_consumed(ftype, sub) + _verify_consumed(ftype, sub, excluded=excluded) else: assert f.name in consumed, ( f"Field `{f.name}` on {cls.__name__} not consumed by the " @@ -609,7 +621,7 @@ def _all_forward_args_field_names() -> set[str]: def _walk(cls) -> None: for f in fields(cls): seen.add(f.name) - ftype = f.type if not isinstance(f.type, str) else None + ftype = _dataclass_field_type(cls, f.name) if ftype is not None and dataclasses.is_dataclass(ftype): _walk(ftype) diff --git a/tests/unittest/_torch/attention/test_fmha_manager.py b/tests/unittest/_torch/attention/test_fmha_manager.py index 17dd1beda5ce..5cf08f1bcd1a 100644 --- a/tests/unittest/_torch/attention/test_fmha_manager.py +++ b/tests/unittest/_torch/attention/test_fmha_manager.py @@ -31,6 +31,10 @@ AttentionInputType, PredefinedAttentionMask, ) +from tensorrt_llm._torch.attention.backends.sparse.params import ( + BlockSparseForwardInputs, + SparseRuntimeParams, +) from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttention from tensorrt_llm.models.modeling_utils import QuantConfig from tensorrt_llm.quantization.mode import QuantAlgo @@ -50,12 +54,14 @@ def _make_metadata( num_generations: int, num_ctx_tokens: int = 0, use_spec_decoding: bool = False, + num_sparse_topk: int = 0, ) -> SimpleNamespace: return SimpleNamespace( num_contexts=num_contexts, num_generations=num_generations, num_ctx_tokens=num_ctx_tokens, use_spec_decoding=use_spec_decoding, + num_sparse_topk=num_sparse_topk, ) @@ -538,6 +544,79 @@ def test_fmha_cache_tracks_attention_mask_data() -> None: assert len(manager._cache) == 2 +def test_block_sparse_requests_only_reach_libraries_that_declare_support() -> None: + events: list[tuple] = [] + attn, manager = _make_manager() + dense_fmha = FakeFmha(attn, "dense", events) + block_sparse_fmha = FakeFmha(attn, "block-sparse", events) + block_sparse_fmha.supports_block_sparse_inputs = True + manager.fmha_libs = [dense_fmha, block_sparse_fmha] + metadata = _make_metadata(num_contexts=1, num_generations=0, num_ctx_tokens=1) + forward_args = AttentionForwardArgs( + attention_input_type=AttentionInputType.context_only, + sparse_runtime_params=SparseRuntimeParams( + block_sparse_inputs=BlockSparseForwardInputs( + q_block_size=64, + kv_block_size=64, + exact_block_bits=torch.zeros((1, 1, 1, 1), dtype=torch.int32), + ) + ), + ) + + selected = manager.select(attn, torch.empty((1, 4)), None, None, metadata, forward_args) + + assert selected is block_sparse_fmha + assert [event[1] for event in events if event[0] == "support"] == ["block-sparse"] + + +@pytest.mark.parametrize("block_sparse_first", [False, True]) +def test_fmha_cache_separates_block_sparse_mode(block_sparse_first: bool) -> None: + events: list[tuple] = [] + attn, manager = _make_manager() + block_sparse_fmha = FakeFmha( + attn, + "block-sparse", + events, + support_predicate=lambda forward_args: ( + forward_args.sparse_runtime_params.block_sparse_inputs is not None + ), + ) + block_sparse_fmha.supports_block_sparse_inputs = True + dense_fmha = FakeFmha( + attn, + "dense", + events, + support_predicate=lambda forward_args: ( + forward_args.sparse_runtime_params.block_sparse_inputs is None + ), + ) + manager.fmha_libs = [block_sparse_fmha, dense_fmha] + metadata = _make_metadata(num_contexts=1, num_generations=0, num_ctx_tokens=1) + q = torch.empty((1, 4)) + by_mode = { + False: AttentionForwardArgs(attention_input_type=AttentionInputType.context_only), + True: AttentionForwardArgs( + attention_input_type=AttentionInputType.context_only, + sparse_runtime_params=SparseRuntimeParams( + block_sparse_inputs=BlockSparseForwardInputs( + q_block_size=64, + kv_block_size=64, + exact_block_bits=torch.zeros((1, 1), dtype=torch.uint32), + ), + ), + ), + } + order = (True, False) if block_sparse_first else (False, True) + + with patch.object(fmha_manager, "_is_fmha_cache_enabled", return_value=True): + selected = { + mode: manager.select(attn, q, None, None, metadata, by_mode[mode]) for mode in order + } + + assert selected == {False: dense_fmha, True: block_sparse_fmha} + assert len(manager._cache) == 2 + + @pytest.mark.parametrize("speculative_first", [False, True]) def test_fmha_cache_separates_speculative_decoding(speculative_first: bool) -> None: events: list[tuple] = [] diff --git a/tests/unittest/_torch/attention/test_fmha_registry.py b/tests/unittest/_torch/attention/test_fmha_registry.py index cd4cdb4fd7a2..bd6a65809afe 100644 --- a/tests/unittest/_torch/attention/test_fmha_registry.py +++ b/tests/unittest/_torch/attention/test_fmha_registry.py @@ -16,8 +16,10 @@ import pytest from tensorrt_llm._torch.attention.backends.fmha import registry +from tensorrt_llm._torch.attention.backends.fmha.interface import Fmha PRIMS_TS = "prims_ts" +PRIMS_TS_BLOCK_SPARSE = "prims_ts_block_sparse" def _canonical_names() -> tuple[str, ...]: @@ -39,10 +41,17 @@ def test_default_fmha_libs_exclude_prims_ts(monkeypatch: pytest.MonkeyPatch) -> monkeypatch.delenv("TLLM_FMHA_LIBS", raising=False) assert PRIMS_TS not in registry.DEFAULT_FMHA_LIBS + assert PRIMS_TS_BLOCK_SPARSE in registry.DEFAULT_FMHA_LIBS assert set(registry.DEFAULT_FMHA_LIBS) <= set(registry.FMHA_LIBS) assert _enabled_names() == registry.DEFAULT_FMHA_LIBS +def test_only_the_block_sparse_fmha_declares_block_sparse_support() -> None: + assert Fmha.supports_block_sparse_inputs is False + for name, fmha_cls in registry.FMHA_LIBS.items(): + assert fmha_cls.supports_block_sparse_inputs is (name == PRIMS_TS_BLOCK_SPARSE), name + + @pytest.mark.parametrize("value", ["", " ", ", ,"]) def test_empty_fmha_lib_env_uses_default( monkeypatch: pytest.MonkeyPatch, diff --git a/tests/unittest/_torch/attention/test_prims_ts_fmha.py b/tests/unittest/_torch/attention/test_prims_ts_fmha.py index 87dbd7364f4b..a18f41c00bf7 100644 --- a/tests/unittest/_torch/attention/test_prims_ts_fmha.py +++ b/tests/unittest/_torch/attention/test_prims_ts_fmha.py @@ -38,6 +38,7 @@ AttentionInputType, PredefinedAttentionMask, ) +from tensorrt_llm._torch.attention.backends.sparse.params import SparseRuntimeParams from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2 from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings import DataType @@ -185,7 +186,7 @@ def _support_result( is_fused_qkv=is_fused_qkv, ) if has_sparse_runtime_metadata: - forward_args.sparse_runtime_params.sparse_kv_indices = torch.empty(1) + forward_args.sparse_runtime_params = SparseRuntimeParams(sparse_kv_indices=torch.empty(1)) if attention_input_type == AttentionInputType.context_only: num_contexts, num_generations, num_ctx_tokens = 1, 0, 4 kv_lens = [4] diff --git a/tests/unittest/_torch/attention/test_skip_softmax_sm120.py b/tests/unittest/_torch/attention/test_skip_softmax_sm120.py index 979b0d3c6908..af8a067f7129 100644 --- a/tests/unittest/_torch/attention/test_skip_softmax_sm120.py +++ b/tests/unittest/_torch/attention/test_skip_softmax_sm120.py @@ -75,13 +75,16 @@ def _run_context( ) -> tuple: """Build a TRTLLM attention layer + no-cache context metadata and run a packed-QKV causal prefill. Mirrors ``test_attention_no_cache``.""" - AttentionCls = get_attention_backend("TRTLLM") + sparse_params = ( + sparse_attention_config.to_sparse_params() if sparse_attention_config is not None else None + ) + AttentionCls = get_attention_backend("TRTLLM", sparse_params=sparse_params) layer = AttentionCls( layer_idx=0, num_heads=num_heads, head_dim=head_dim, num_kv_heads=num_kv_heads, - sparse_attention_config=sparse_attention_config, + sparse_params=sparse_params, ) metadata = AttentionCls.Metadata( diff --git a/tests/unittest/_torch/visual_gen/test_fa4_cutlass_compatibility.py b/tests/unittest/_torch/visual_gen/test_fa4_cutlass_compatibility.py index eb9ea6941792..e1203891be26 100644 --- a/tests/unittest/_torch/visual_gen/test_fa4_cutlass_compatibility.py +++ b/tests/unittest/_torch/visual_gen/test_fa4_cutlass_compatibility.py @@ -1,6 +1,9 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 +import subprocess +import sys +import textwrap from importlib import import_module import pytest @@ -10,6 +13,7 @@ from tensorrt_llm._torch.visual_gen.attention_backend import flash_attn4, parallel # noqa: E402 from tensorrt_llm._torch.visual_gen.attention_backend.flash_attn4 import ( # noqa: E402 _install_cutlass_dsl_compatibility, + _install_flash_attn_tile_scheduler_compatibility, ) @@ -48,3 +52,54 @@ def test_cutlass_dsl_47_aliases_allow_fa4_interface_import() -> None: assert callable(interface.flash_attn_combine) assert callable(flash_attn4._flash_attn_fwd) assert callable(parallel._flash_attn_combine) + + +def test_fa4_work_tile_info_survives_cutlass_task_scheduling_import() -> None: + task_scheduling = pytest.importorskip("cutlass.experimental.task_scheduling") + tile_scheduler = pytest.importorskip("flash_attn.cute.tile_scheduler") + import cutlass + from cutlass.cutlass_dsl import Boolean + from cutlass.utils.static_persistent_tile_scheduler import WorkTileInfo as CutlassWorkTileInfo + + del task_scheduling + _install_flash_attn_tile_scheduler_compatibility() + tile_idx = (cutlass.Int32(1), cutlass.Int32(2), cutlass.Int32(3), cutlass.Int32(0)) + + # The task-scheduling import rewrote the shared CUTLASS class to three scalars. + with pytest.raises(ValueError, match="too many values to unpack"): + CutlassWorkTileInfo(tile_idx, Boolean(True)) + + fa4_tile = tile_scheduler.WorkTileInfo(tile_idx, Boolean(True)) + + assert "__init__" in vars(tile_scheduler.WorkTileInfo) + assert fa4_tile.tile_idx == tile_idx + assert bool(fa4_tile.is_valid_tile) + assert issubclass(tile_scheduler.WorkTileInfo, CutlassWorkTileInfo) + + +def test_fa4_work_tile_info_survives_task_scheduling_imported_first() -> None: + pytest.importorskip("cutlass.experimental.task_scheduling") + pytest.importorskip("flash_attn.cute.tile_scheduler") + script = textwrap.dedent( + """ + import cutlass + import cutlass.experimental.task_scheduling # noqa: F401 + from cutlass.cutlass_dsl import Boolean + + import tensorrt_llm._torch.visual_gen.attention_backend.flash_attn4 # noqa: F401 + from flash_attn.cute.tile_scheduler import WorkTileInfo + + tile_idx = (cutlass.Int32(1), cutlass.Int32(2), cutlass.Int32(3), cutlass.Int32(0)) + tile = WorkTileInfo(tile_idx, Boolean(True)) + assert tile.tile_idx == tile_idx + assert bool(tile.is_valid_tile) + print("fa4-work-tile-info-ok") + """ + ) + + result = subprocess.run( + [sys.executable, "-c", script], capture_output=True, text=True, timeout=900, check=False + ) + + assert result.returncode == 0, result.stderr[-4000:] + assert result.stdout.strip().endswith("fa4-work-tile-info-ok")