MiniMaxM3SparseAttentionConfig#

class tensorrt_llm.llmapi.MiniMaxM3SparseAttentionConfig(
*,
algorithm: Literal['minimax_m3'] = 'minimax_m3',
sparse_num_index_heads: Annotated[int, Gt(gt=0)] = 4,
sparse_index_dim: Annotated[int, Gt(gt=0)] = 128,
sparse_block_size: int = 128,
sparse_topk_blocks: int = 16,
sparse_init_blocks: int = 0,
sparse_local_blocks: int = 1,
sparse_score_type: Literal['max'] = 'max',
sparse_disable_index_value: bool = True,
indexer_kv_dtype: Literal['bf16', 'fp8'] = 'bf16',
fuse_qkv_index_projection: bool = False,
num_attention_heads: int | None = None,
num_key_value_heads: int | None = None,
implementation: Literal['triton', 'msa'] = 'triton',
)[source]#

Bases: BaseSparseAttentionConfig

Configuration for MiniMax-M3 block-sparse attention.

Drives the two-step sparse attention used by MiniMax-M3 layers 3..N:

  1. An index attention branch projects a per-head Q vector and a single replicated K vector, scores main K/V cache blocks, and selects the top-k blocks per (num_kv_heads, q_token) pair, with init_blocks forced at the head and local_blocks forced at the tail.

  2. A sparse GQA attention runs only over the selected blocks.

At runtime one of the MiniMax-M3 sparse attention backends under tensorrt_llm._torch.attention.backends.sparse.minimax_m3 is selected. The chosen backend runs on top of a MiniMaxM3KVCacheManagerV2 that allocates a paged side index-K cache of shape [num_slots, 1, sparse_index_dim] parallel to the main K/V cache. The M3 checkpoint sets disable_index_value=True on every sparse layer, so no index V cache is allocated.

field algorithm: Literal['minimax_m3'] = 'minimax_m3'#
field fuse_qkv_index_projection: bool = False#

Fuse Q/K/V and index-Q/index-K into one quantized projection. Index-Q is sharded with the KV heads and index-K is replicated. MSA batches also use a horizontal norm/RoPE/cache-insertion producer for prefill, mixed, and CUDA-graph decode execution. The MiniMax-M3-specific path requires the MSA implementation, indexer_kv_dtype=’fp8’, and an FP8 main KV cache.

field implementation: Literal['triton', 'msa'] = 'triton'#

Sparse attention implementation: ‘triton’ reference (default) or ‘msa’ (fmha_sm100 kernels). The ‘msa’ implementation requires an SM100 GPU, the fmha_sm100 package, and sparse_block_size == 128.

field indexer_kv_dtype: Literal['bf16', 'fp8'] = 'bf16'#

Storage and score-compute dtype for normalized index Q/K. ‘fp8’ uses unscaled E4M3 values with FP32 score accumulation and is supported only by the MSA implementation.

field num_attention_heads: int | None = None#

Global number of attention (query) heads. When unset, it falls back to pretrained_config.num_attention_heads.

field num_key_value_heads: int | None = None#

Global number of key/value heads. When unset, it falls back to pretrained_config.num_key_value_heads, then to num_attention_heads.

field sparse_block_size: int = 128#

Block size used by per-block scoring + top-k selection.

field sparse_disable_index_value: bool = True#

If True, skip the index V branch (M3 checkpoint default).

field sparse_index_dim: int = 128#

Per-head index Q/K dimension.

Constraints:
  • gt = 0

field sparse_init_blocks: int = 0#

Number of leading blocks forced into the top-k regardless of score.

field sparse_local_blocks: int = 1#

Number of trailing blocks forced into the top-k regardless of score.

field sparse_num_index_heads: Annotated[int, Gt(gt=0)] = 4#

Global checkpoint index-attention head count. Index heads shard with their KV-head groups in both separate and fused projections.

Constraints:
  • gt = 0

field sparse_score_type: Literal['max'] = 'max'#

Per-block score reduction; the M3 checkpoint sets ‘max’.

field sparse_topk_blocks: int = 16#

Number of top-k blocks per (kv_head, q_token).

__init__(**data: Any) → None#

Create a new model by parsing and validating input data from keyword arguments.

Raises [ValidationError][pydantic_core.ValidationError] if the input data cannot be validated to form a valid model.

self is explicitly positional-only to allow self as a field name.

get_indices_block_size() → int[source]#
supports_backend(backend: str) → bool[source]#

Override if the sparse attention algorithm does not support a subset of the possible backends.

to_sparse_metadata_params(**kwargs)[source]#

Lower into MiniMaxM3SparseMetadataParams for the attention metadata.

Head counts resolve as this config, then pretrained_config, then a default; num_key_value_heads falls back to num_attention_heads. Setting them on the config lets tests skip building a pretrained_config.

to_sparse_params(**kwargs)[source]#

Lower user-facing config into SparseParams.