From 9743cf02ea22729c487e098bda41949f94338533 Mon Sep 17 00:00:00 2001 From: Tracin <10434017+Tracin@users.noreply.github.com> Date: Mon, 25 May 2026 02:49:07 -0700 Subject: [PATCH 01/11] [None][feat] Add high-precision KV pool support for FP4 MLA Allocate a standalone BF16 tensor indexed by sequence slot for recent MLA tokens when FP4 KV cache is active. Add Triton kernels to store context and generation latent cache values into the high-precision pool before attention runs. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> --- .../_torch/attention_backend/trtllm.py | 257 ++++++++++++++++++ .../_torch/pyexecutor/model_engine.py | 32 +++ 2 files changed, 289 insertions(+) diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index 93e1a2bfe4c4..fa2e77b8201c 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -21,6 +21,8 @@ from typing import TYPE_CHECKING, List, Optional, Tuple import torch +import triton +import triton.language as tl if TYPE_CHECKING: from ..speculative.interface import SpecMetadata @@ -29,8 +31,11 @@ from tensorrt_llm._torch.attention_backend.fmha import ( Fmha, get_enabled_fmha_lib_classes) from tensorrt_llm._utils import get_sm_version, maybe_pin_memory, prefer_pinned +from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal import thop from tensorrt_llm.functional import AttentionMaskType +from tensorrt_llm.llmapi import SkipSoftmaxAttentionConfig +from tensorrt_llm.logger import logger from tensorrt_llm.models.modeling_utils import QuantConfig from ..utils import (compute_swizzled_sf_shape, get_global_attrs, @@ -43,6 +48,11 @@ from .sparse.params import SparseParams from .sparse.skip_softmax import SkipSoftmaxParams +# Circular buffer size for the high-precision BF16 KV pool (MLA FP4 models). +# Each sequence slot holds HP_BLOCK_SIZE token vectors; token at absolute +# position t maps to buffer slot t % HP_BLOCK_SIZE. +HP_BLOCK_SIZE: int = 16 + @functools.cache def generate_spec_decoding_position_offsets(max_num_requests: int, @@ -146,6 +156,24 @@ class TrtllmAttentionMetadata(AttentionMetadata): host_kv_cache_block_offsets: Optional[torch.Tensor] = None draft_kv_cache_block_offsets: Optional[torch.Tensor] = None + # Per-request seq_slot from SeqSlotManager (stable across steps). + # GPU tensor with stable address for CUDA graph compatibility. + # Shape: [max_num_sequences], int32. Values copied each step. + seq_slots: Optional[torch.Tensor] = None + seq_slots_cpu: Optional[torch.Tensor] = None + + # True during warmup forward passes (dummy requests, no real data). + is_warmup: bool = False + + # High-precision BF16 KV pool for MLA FP4 models, indexed by seq_slot. + # Shape: [max_num_sequences, num_local_layers, kv_factor, HP_BLOCK_SIZE * head_dim] + # Standalone tensor — not part of the block-based paged KV cache. + high_precision_kv_pool: Optional[torch.Tensor] = None + # Ownership tracking: maps seq_slot → request_id that last wrote it. + # Plain Python dict, updated during context phase, checked during decode. + # Debug only — runs outside CUDA graph. + hp_pool_owners: Optional[dict] = None + # Pre-computed FlashMLA tile-scheduler metadata and num_splits. # Computed once per forward pass in TrtllmAttention.forward() and reused across layers. flash_mla_tile_scheduler_metadata: Optional[torch.Tensor] = None @@ -419,6 +447,45 @@ def _post_init_with_buffers(self, buffers) -> None: pin_memory=prefer_pinned(), ) + # Allocate high-precision BF16 KV pool for MLA FP4 models. + # Standalone tensor indexed by seq_slot, not part of block-based paged KV cache. + # Each sequence gets a circular buffer of HP_BLOCK_SIZE=16 recent tokens at BF16. + if (self.kv_cache_manager is not None + and self.kv_cache_manager.kv_factor == 1 + and self.kv_cache_manager.dtype == DataType.NVFP4): + + self.seq_slots = self.get_empty( + buffers, + (self.max_num_sequences, ), + cache_name="seq_slots", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self.seq_slots_cpu = torch.empty( + self.max_num_sequences, + dtype=torch.int32, + device='cpu', + pin_memory=prefer_pinned(), + ) + self.hp_pool_owners = {} + num_local_layers = self.kv_cache_manager.num_local_layers + head_dim = self.kv_cache_manager.head_dim + kv_factor = self.kv_cache_manager.kv_factor + self.high_precision_kv_pool = self.get_empty( + buffers, + [ + self.max_num_sequences, num_local_layers, kv_factor, + HP_BLOCK_SIZE * head_dim + ], + cache_name="high_precision_kv_pool", + dtype=torch.bfloat16, + capture_graph=capture_graph, + ) + logger.info( + f"Allocated high-precision BF16 KV pool: shape=" + f"{list(self.high_precision_kv_pool.shape)}, " + f"size={self.high_precision_kv_pool.nbytes / (1 << 20):.1f} MB") + # Allocate static buffers for helix parallelism support. if self.enable_helix: self.helix_position_offsets = self.get_empty( @@ -1116,6 +1183,106 @@ def is_sm_version_trtllm_gen_kernel(self, sm): return not (sm < 100 or sm in [120, 121]) +# --------------------------------------------------------------------------- +# Triton kernels for storing latent cache into the high-precision KV pool +# --------------------------------------------------------------------------- + + +@triton.jit +def _hp_kv_store_context_kernel( + pool_ptr, + latent_cache_ptr, + seq_slots_ptr, # int32 [num_contexts] – seq_slot for each ctx seq + kv_lens_ptr, # int32 [num_contexts] – total KV length after prefill + token_offsets_ptr, # int32 [num_contexts] – excl. prefix-sum of prompt_lens + prompt_lens_ptr, # int32 [num_contexts] – number of new tokens per ctx seq + layer_idx, + pool_stride_seq, # pool.stride(0): elements between adjacent seq slots + pool_stride_layer, # pool.stride(1): elements between adjacent layers + lc_stride, # latent_cache.stride(0): elements between adjacent tokens + D: tl.constexpr, + BLOCK_D: tl.constexpr, + HP_BLOCK: tl.constexpr, # = HP_BLOCK_SIZE (16) +): + """Store the tail tokens of each context sequence into the HP KV pool. + + Grid: (num_contexts, HP_BLOCK_SIZE). + Only programs where buf_pos < kv_len % HP_BLOCK actually write. + + For a context sequence with total KV length L = num_cached + prompt_len: + - remainder = L % HP_BLOCK + - The last `remainder` new tokens (latent_cache positions + [offset + prompt_len - remainder, offset + prompt_len)) are stored + into pool slots [0, remainder), which correspond to the absolute token + positions [L - remainder, L) in the circular buffer. + """ + ctx_idx = tl.program_id(0) + buf_pos = tl.program_id(1) + + kv_len = tl.load(kv_lens_ptr + ctx_idx) + remainder = kv_len % HP_BLOCK + if buf_pos >= remainder: + return + + seq_slot = tl.load(seq_slots_ptr + ctx_idx) + prompt_len = tl.load(prompt_lens_ptr + ctx_idx) + tok_offset = tl.load(token_offsets_ptr + ctx_idx) + + # Index of this token within latent_cache: last `remainder` new tokens, + # buf_pos-th of them (0-indexed from the start of the tail). + token_idx = tok_offset + prompt_len - remainder + buf_pos + + offs_d = tl.arange(0, BLOCK_D) + mask_d = offs_d < D + src = tl.load(latent_cache_ptr + token_idx * lc_stride + offs_d, + mask=mask_d, + other=0.0) + + # Destination: pool[seq_slot, layer_idx, 0, buf_pos * D : (buf_pos+1) * D] + dst_base = (seq_slot * pool_stride_seq + layer_idx * pool_stride_layer + + buf_pos * D) + tl.store(pool_ptr + dst_base + offs_d, src, mask=mask_d) + + +@triton.jit +def _hp_kv_store_gen_kernel( + pool_ptr, + latent_cache_ptr, + seq_slots_ptr, # int32 [num_gen] – seq_slot for each gen seq + kv_lens_ptr, # int32 [num_gen] – total KV length after this decode step + gen_tok_start, # int – offset in latent_cache where gen tokens begin + layer_idx, + pool_stride_seq, + pool_stride_layer, + lc_stride, + D: tl.constexpr, + BLOCK_D: tl.constexpr, + HP_BLOCK: tl.constexpr, +): + """Store the current generation token into the HP KV pool. + + Grid: (num_gen_seqs,). + Each program stores one token into the circular buffer position + (kv_len - 1) % HP_BLOCK, overwriting the oldest entry. + """ + gen_idx = tl.program_id(0) + + seq_slot = tl.load(seq_slots_ptr + gen_idx) + kv_len = tl.load(kv_lens_ptr + gen_idx) + buf_pos = (kv_len - 1) % HP_BLOCK + + token_idx = gen_tok_start + gen_idx + offs_d = tl.arange(0, BLOCK_D) + mask_d = offs_d < D + src = tl.load(latent_cache_ptr + token_idx * lc_stride + offs_d, + mask=mask_d, + other=0.0) + + dst_base = (seq_slot * pool_stride_seq + layer_idx * pool_stride_layer + + buf_pos * D) + tl.store(pool_ptr + dst_base + offs_d, src, mask=mask_d) + + class TrtllmAttention(AttentionBackend[TrtllmAttentionMetadata]): Metadata = TrtllmAttentionMetadata @@ -1421,6 +1588,92 @@ def create_fmha_libs(self) -> None: if fmha_cls.is_available(self): self.fmha_libs.append(fmha_cls(self)) + def _update_high_precision_kv_for_fp4_mla( + self, + metadata: TrtllmAttentionMetadata, + latent_cache: Optional[torch.Tensor], + ) -> None: + """Store recent KV tokens at BF16 into the high-precision pool.""" + if metadata.hp_pool_owners is None: + return + num_contexts = metadata.num_contexts + num_seqs = metadata.num_seqs + local_layer = self.get_local_layer_idx(metadata) + + if local_layer == 0 and not metadata.is_cuda_graph and not metadata.is_warmup: + for batch_idx in range(num_contexts): + seq_slot = metadata.seq_slots_cpu[batch_idx].item() + request_id = metadata.request_ids[batch_idx] + metadata.hp_pool_owners[seq_slot] = request_id + + for batch_idx in range(num_contexts, num_seqs): + seq_slot = metadata.seq_slots_cpu[batch_idx].item() + request_id = metadata.request_ids[batch_idx] + owner = metadata.hp_pool_owners.get(seq_slot) + if owner != request_id: + raise RuntimeError( + f"HP KV pool ownership mismatch: seq_slot={seq_slot} " + f"is owned by request {owner} but request " + f"{request_id} is attempting to use it") + + if latent_cache is None: + return + + pool = metadata.high_precision_kv_pool + head_dim = latent_cache.shape[-1] + block_d = triton.next_power_of_2(head_dim) + pool_s0 = pool.stride(0) + pool_s1 = pool.stride(1) + lc_stride = latent_cache.stride(0) + + if num_contexts > 0: + prompt_lens_cpu = metadata.prompt_lens_cpu_runtime[:num_contexts] + token_offsets_cpu = torch.zeros(num_contexts, + dtype=torch.int32, + device='cpu') + if num_contexts > 1: + token_offsets_cpu[1:].copy_( + torch.cumsum(prompt_lens_cpu[:-1].to(torch.int32), dim=0)) + token_offsets_gpu = token_offsets_cpu.to(pool.device, + non_blocking=False) + prompt_lens_gpu = metadata.prompt_lens_cuda_runtime[:num_contexts] + + _hp_kv_store_context_kernel[(num_contexts, HP_BLOCK_SIZE)]( + pool, + latent_cache, + metadata.seq_slots, + metadata.kv_lens_cuda_runtime, + token_offsets_gpu, + prompt_lens_gpu, + local_layer, + pool_s0, + pool_s1, + lc_stride, + D=head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + + num_gen = num_seqs - num_contexts + if num_gen > 0: + ctx_tok_count = int( + metadata.prompt_lens_cpu_runtime[:num_contexts].sum().item()) + + _hp_kv_store_gen_kernel[(num_gen, )]( + pool, + latent_cache, + metadata.seq_slots[num_contexts:], + metadata.kv_lens_cuda_runtime[num_contexts:], + ctx_tok_count, + local_layer, + pool_s0, + pool_s1, + lc_stride, + D=head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + def forward( self, q: torch.Tensor, @@ -1638,6 +1891,10 @@ def forward( assert metadata.kv_cache_manager is None assert metadata.num_contexts == metadata.num_seqs + if metadata.high_precision_kv_pool is not None and self.is_mla_enable: + self._update_high_precision_kv_for_fp4_mla( + metadata, forward_args.latent_cache) + if not self.fmha_libs: self.create_fmha_libs() diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index ff1e5f8a5306..dc0c07a88803 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -2905,6 +2905,7 @@ def _prepare_tp_inputs( sequence_lengths = [] # per sequence prompt_lengths = [] # per sequence request_ids = [] # per request + seq_slots = [] # per request (stable seq_slot from SeqSlotManager) gather_ids = [] position_ids = [] # per sequence num_cached_tokens_per_seq = [] # per sequence @@ -2975,6 +2976,8 @@ def append_cross_attention_state(request: LlmRequest, for request in scheduled_requests.context_requests: request_ids.append(request.py_request_id) + seq_slots.append( + request.py_seq_slot if request.py_seq_slot is not None else 0) all_prompt_tokens = request.get_tokens(0) draft_lens.append(0) begin_compute = request.context_current_position @@ -3146,6 +3149,8 @@ def append_cross_attention_state(request: LlmRequest, previous_pos_indices = [] for request in extend_requests: request_ids.append(request.py_request_id) + seq_slots.append( + request.py_seq_slot if request.py_seq_slot is not None else 0) request_accepted_path[ request. py_request_id] = request.py_num_accepted_draft_tokens_indices @@ -3224,6 +3229,8 @@ def append_cross_attention_state(request: LlmRequest, for request in first_draft_requests: request_ids.append(request.py_request_id) + seq_slots.append( + request.py_seq_slot if request.py_seq_slot is not None else 0) all_prompt_tokens = request.get_tokens(0) draft_lens.append(0) begin_compute = len( @@ -3295,6 +3302,8 @@ def append_cross_attention_state(request: LlmRequest, # py_batch_idx is set after a request occupies a generation slot. # None means this is its first generation step on this worker. is_generation_admission = request.py_batch_idx is None + seq_slots.append(request.py_seq_slot if request. + py_seq_slot is not None else 0) # the request has no previous tensor: # (1) new_tokens_device is None, which means overlap scheduler is disabled; or # (2) a dummy request; or @@ -3751,6 +3760,14 @@ def previous_seq_slots_device(): attn_metadata.beam_width = 1 attn_metadata.request_ids = request_ids + attn_metadata.is_warmup = self.is_warmup + if hasattr(attn_metadata, + 'seq_slots') and attn_metadata.seq_slots is not None: + num_seqs = len(seq_slots) + attn_metadata.seq_slots_cpu[:num_seqs] = torch.tensor( + seq_slots, dtype=torch.int32) + attn_metadata.seq_slots[:num_seqs].copy_( + attn_metadata.seq_slots_cpu[:num_seqs], non_blocking=True) attn_metadata.prompt_lens = prompt_lengths attn_metadata.num_contexts = scheduled_requests.num_context_requests # Use num_chunked_ctx_requests to record the number of extend context requests, @@ -3909,6 +3926,7 @@ def _prepare_tp_inputs_no_cache( multi_modal_data = [] draft_lens = [] request_ids = [] + seq_slots = [] multimodal_params_list = [] for request in scheduled_requests.context_requests: @@ -3918,6 +3936,8 @@ def _prepare_tp_inputs_no_cache( context_start_idx = len(input_ids) input_ids.extend(prompt_tokens) request_ids.append(request.py_request_id) + seq_slots.append( + request.py_seq_slot if request.py_seq_slot is not None else 0) if request.position_ids is None: position_ids.extend(range(len(prompt_tokens))) else: @@ -4062,6 +4082,7 @@ def _prepare_star_attention_inputs( input_ids = [] prompt_lengths = [] request_ids = [] + seq_slots = [] gather_ids = [] position_ids = [] # for star attention, we need customized block ids @@ -4069,6 +4090,8 @@ def _prepare_star_attention_inputs( num_cached_tokens_per_seq = [] for request in scheduled_requests.context_requests: request_ids.append(request.py_request_id) + seq_slots.append( + request.py_seq_slot if request.py_seq_slot is not None else 0) prompt_lengths.append(request.py_prompt_len) ctx_iter = request.ctx_iters @@ -4180,6 +4203,8 @@ def _prepare_star_attention_inputs( for request in generation_requests: request_ids.append(request.py_request_id) + seq_slots.append( + request.py_seq_slot if request.py_seq_slot is not None else 0) prompt_lengths.append(request.py_prompt_len) input_token_id = request.get_token(0, request.get_num_tokens(0) - 1) @@ -4251,6 +4276,13 @@ def _prepare_star_attention_inputs( ) attn_metadata.request_ids = request_ids + if hasattr(attn_metadata, + 'seq_slots') and attn_metadata.seq_slots is not None: + num_seqs = len(seq_slots) + attn_metadata.seq_slots_cpu[:num_seqs] = torch.tensor( + seq_slots, dtype=torch.int32) + attn_metadata.seq_slots[:num_seqs].copy_( + attn_metadata.seq_slots_cpu[:num_seqs], non_blocking=True) attn_metadata.prompt_lens = prompt_lengths attn_metadata.num_contexts = num_contexts attn_metadata.num_queries = num_queries From 7eec65cb5ca4d4fdf68273977734688b4fb8e772 Mon Sep 17 00:00:00 2001 From: Tracin <10434017+Tracin@users.noreply.github.com> Date: Mon, 25 May 2026 02:50:30 -0700 Subject: [PATCH 02/11] [None][feat] Support NVFP4 MLA KV cache on FlashInfer Add shared FP4 MLA KV-cache helpers, V-scale storage, Triton no-dequant decode, and FlashInfer integration for NVFP4 MLA cache. Fold the offset, layout, memory, workspace, tail-scale, and scatter overflow fixes into the implementation, with focused unit coverage. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> --- .../batch_manager/kvCacheManager.h | 23 +- .../batch_manager/kvCacheManager.cpp | 47 +- .../nanobind/batch_manager/kvCacheManager.cpp | 17 +- .../_torch/attention_backend/flashinfer.py | 507 +++++- .../attention_backend/fp4_mla_kernels.py | 1287 ++++++++++++++ .../_torch/attention_backend/fp4_mla_kv.py | 1338 ++++++++++++++ .../_torch/attention_backend/trtllm.py | 210 +-- .../_torch/pyexecutor/py_executor_creator.py | 73 +- .../_torch/pyexecutor/resource_manager.py | 1543 ++++++++++++++++- .../attention/test_flashinfer_attention.py | 30 + .../_torch/attention/test_fp4_mla_kv.py | 1355 +++++++++++++++ .../executor/test_mla_tokens_per_block.py | 102 ++ 12 files changed, 6250 insertions(+), 282 deletions(-) create mode 100644 tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py create mode 100644 tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py create mode 100644 tests/unittest/_torch/attention/test_fp4_mla_kv.py create mode 100644 tests/unittest/_torch/executor/test_mla_tokens_per_block.py diff --git a/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h b/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h index c665f7a8df95..3930d56bc47b 100644 --- a/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h +++ b/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h @@ -2087,6 +2087,12 @@ class BaseKVCacheManager [[nodiscard]] virtual runtime::ITensor::SharedPtr getUniquePrimaryPool() const = 0; [[nodiscard]] virtual runtime::ITensor::SharedPtr getPrimaryPool(SizeType32 layer_idx) const = 0; [[nodiscard]] virtual runtime::ITensor::SharedPtr getIndexerKCachePool() const = 0; + + [[nodiscard]] virtual runtime::ITensor::SharedPtr getMlaVScalePool() const + { + return nullptr; + } + [[nodiscard]] virtual SizeType32 getPoolLayerIdx(SizeType32 layer_idx) const = 0; [[nodiscard]] virtual bool isPoolLayerFirst(SizeType32 layer_idx) const = 0; @@ -2250,7 +2256,7 @@ class KVCacheManager : public BaseKVCacheManager bool enableIndexerKCache = false, SizeType32 indexerKCacheQuantBlockSize = 128, SizeType32 indexerKCacheIndexHeadDim = 0, bool indexerKCacheUseFp4 = false, std::optional linearAttentionMetadata = std::nullopt, - std::vector const& poolConfigurations = {}); + std::vector const& poolConfigurations = {}, bool enableMlaVScalePool = false); KVCacheManager(std::vector const& numKvHeadsPerLayer, SizeType32 sizePerHead, SizeType32 tokensPerBlock, BlocksPerWindow const& blocksPerWindow, SizeType32 maxNumSequences, SizeType32 maxBeamWidth, @@ -2264,7 +2270,7 @@ class KVCacheManager : public BaseKVCacheManager bool enableIndexerKCache = false, SizeType32 indexerKCacheQuantBlockSize = 128, SizeType32 indexerKCacheIndexHeadDim = 0, bool indexerKCacheUseFp4 = false, std::optional linearAttentionMetadata = std::nullopt, - std::vector const& poolConfigurations = {}); + std::vector const& poolConfigurations = {}, bool enableMlaVScalePool = false); KVCacheManager(SizeType32 numLayers, SizeType32 numKvHeads, SizeType32 sizePerHead, SizeType32 tokensPerBlock, BlocksPerWindow const& blocksPerWindow, SizeType32 maxNumSequences, SizeType32 maxBeamWidth, @@ -2278,7 +2284,7 @@ class KVCacheManager : public BaseKVCacheManager bool enableIndexerKCache = false, SizeType32 indexerKCacheQuantBlockSize = 128, SizeType32 indexerKCacheIndexHeadDim = 0, bool indexerKCacheUseFp4 = false, std::optional linearAttentionMetadata = std::nullopt, - std::vector const& poolConfigurations = {}); + std::vector const& poolConfigurations = {}, bool enableMlaVScalePool = false); KVCacheManager(SizeType32 numLayers, SizeType32 numKvHeads, SizeType32 sizePerHead, SizeType32 tokensPerBlock, BlocksPerWindow const& blocksPerWindow, SizeType32 maxNumSequences, SizeType32 maxBeamWidth, @@ -2288,7 +2294,7 @@ class KVCacheManager : public BaseKVCacheManager bool enableIndexerKCache = false, SizeType32 indexerKCacheQuantBlockSize = 128, SizeType32 indexerKCacheIndexHeadDim = 0, bool indexerKCacheUseFp4 = false, std::optional linearAttentionMetadata = std::nullopt, - std::vector const& poolConfigurations = {}); + std::vector const& poolConfigurations = {}, bool enableMlaVScalePool = false); ~KVCacheManager() override = default; @@ -2564,6 +2570,11 @@ class KVCacheManager : public BaseKVCacheManager runtime::ITensor::SharedPtr getPrimaryPool(SizeType32 layer_idx) const override; runtime::ITensor::SharedPtr getIndexerKCachePool() const override; + runtime::ITensor::SharedPtr getMlaVScalePool() const override + { + return mMlaVScalePool; + } + SizeType32 getPoolLayerIdx(SizeType32 layer_idx) const override { return mBlockManager.getPoolLayerIdx(layer_idx); @@ -2633,6 +2644,8 @@ class KVCacheManager : public BaseKVCacheManager SizeType32 mMaxAttentionWindow; // Number of tokens per block SizeType32 mTokensPerBlock; + // Size of each attention head before FP4 packing. + SizeType32 mSizePerHead; // Number of tokens to fill up the sink tokens to a full block size SizeType32 mSinkBubbleLength; // Number of tokens in the sink blocks @@ -2645,6 +2658,7 @@ class KVCacheManager : public BaseKVCacheManager std::unordered_map mSequences; // Whether to cache KV pages for reuse bool mEnableBlockReuse; + bool mEnableMlaVScalePool; // Mutex to protect access to mSequences mutable std::mutex mSequencesMtx; // buffers for static tensors, will be created after allocating pools @@ -2652,6 +2666,7 @@ class KVCacheManager : public BaseKVCacheManager runtime::ITensor::SharedPtr mLayerToPoolMapping; runtime::ITensor::SharedPtr mBlockScalePoolPointers; runtime::ITensor::SharedPtr mIndexerKCachePoolPointers; + runtime::ITensor::SharedPtr mMlaVScalePool; // GPU bytes allocated for KV-cache std::size_t mAllocatedBytes{0}; }; diff --git a/cpp/tensorrt_llm/batch_manager/kvCacheManager.cpp b/cpp/tensorrt_llm/batch_manager/kvCacheManager.cpp index 0fb8af1527ae..14dc844cde31 100644 --- a/cpp/tensorrt_llm/batch_manager/kvCacheManager.cpp +++ b/cpp/tensorrt_llm/batch_manager/kvCacheManager.cpp @@ -87,6 +87,19 @@ std::vector getAllSequenceBlocks(BlockPtr lastBlock) return sequenceBlocks; } +SizeType32 getMlaVScaleElemsPerPage(SizeType32 sizePerHead, SizeType32 tokensPerBlock) +{ + // KVCacheManager only receives the full MLA latent head size. The Python V-scale view may consume fewer row + // groups when v_head_dim excludes RoPE channels, so this allocation intentionally keeps that headroom. + constexpr SizeType32 kFp4BlockSize = 16; + constexpr SizeType32 kScaleRowGroup = 128; + constexpr SizeType32 kScaleColGroup = 4; + auto const tokenScaleCols = tc::ceilDiv(tokensPerBlock, kFp4BlockSize); + auto const rowGroups = tc::ceilDiv(sizePerHead, kScaleRowGroup); + auto const colGroups = tc::ceilDiv(tokenScaleCols, kScaleColGroup); + return rowGroups * colGroups * 32 * 16; +} + // Compute maximum number of tokens that have been computed by prefill and generation. // Accounts for chunked prefill to avoid storing state that hasn't been written to KV cache yet. // We call LlmRequest::getContextRemainingLength to see how many tokens are still waiting to be computed in prefill. @@ -3173,13 +3186,13 @@ KVCacheManager::KVCacheManager(SizeType32 numLayers, SizeType32 numKvHeads, Size bool enableBlockReuse, CacheType cacheType, bool enablePartialReuse, bool copyOnPartialReuse, bool enableIndexerKCache, SizeType32 indexerKCacheQuantBlockSize, SizeType32 indexerKCacheIndexHeadDim, bool indexerKCacheUseFp4, std::optional linearAttentionMetadata, - std::vector const& poolConfigurations) + std::vector const& poolConfigurations, bool enableMlaVScalePool) : KVCacheManager(std::vector(numLayers, numKvHeads), sizePerHead, tokensPerBlock, blocksPerWindow, maxNumSequences, maxBeamWidth, maxAttentionWindowVec, dtype, sinkTokenLength, std::make_shared(reinterpret_cast(stream)), maxSequenceLength, chunkSize, enableBlockReuse, cacheType, std::nullopt, nullptr, enablePartialReuse, copyOnPartialReuse, nullptr, enableIndexerKCache, indexerKCacheQuantBlockSize, indexerKCacheIndexHeadDim, indexerKCacheUseFp4, - linearAttentionMetadata, poolConfigurations) + linearAttentionMetadata, poolConfigurations, enableMlaVScalePool) { } @@ -3192,13 +3205,13 @@ KVCacheManager::KVCacheManager(std::vector const& numKvHeadsPerLayer std::shared_ptr kvCacheConnectorManager, bool enableIndexerKCache, SizeType32 indexerKCacheQuantBlockSize, SizeType32 indexerKCacheIndexHeadDim, bool indexerKCacheUseFp4, std::optional linearAttentionMetadata, - std::vector const& poolConfigurations) + std::vector const& poolConfigurations, bool enableMlaVScalePool) : KVCacheManager(numKvHeadsPerLayer, sizePerHead, tokensPerBlock, blocksPerWindow, maxNumSequences, maxBeamWidth, maxAttentionWindowVec, dtype, sinkTokenLength, std::make_shared(reinterpret_cast(stream)), maxSequenceLength, chunkSize, enableBlockReuse, cacheType, secondaryOffloadMinPriority, eventManager, enablePartialReuse, copyOnPartialReuse, kvCacheConnectorManager, enableIndexerKCache, indexerKCacheQuantBlockSize, indexerKCacheIndexHeadDim, - indexerKCacheUseFp4, linearAttentionMetadata, poolConfigurations) + indexerKCacheUseFp4, linearAttentionMetadata, poolConfigurations, enableMlaVScalePool) { } @@ -3211,11 +3224,12 @@ KVCacheManager::KVCacheManager(std::vector const& numKvHeadsPerLayer std::shared_ptr kvCacheConnectorManager, bool enableIndexerKCache, SizeType32 indexerKCacheQuantBlockSize, SizeType32 indexerKCacheIndexHeadDim, bool indexerKCacheUseFp4, std::optional linearAttentionMetadata, - std::vector const& poolConfigurations) + std::vector const& poolConfigurations, bool enableMlaVScalePool) : mMaxBeamWidth(maxBeamWidth) , mDataType(dtype) , mMaxAttentionWindow(*std::max_element(maxAttentionWindowVec.begin(), maxAttentionWindowVec.end())) , mTokensPerBlock(tokensPerBlock) + , mSizePerHead(sizePerHead) , mSinkBubbleLength(BaseKVCacheManager::getSinkBubbleLength(sinkTokenLength, tokensPerBlock)) , mSinkBlockTokenLength(mSinkBubbleLength + sinkTokenLength) , mChunkSize(chunkSize) @@ -3227,6 +3241,7 @@ KVCacheManager::KVCacheManager(std::vector const& numKvHeadsPerLayer poolConfigurations) // disable block reuse for sink bubble since chopVectorIntoBlocks does not match KV cache blocks in this case , mEnableBlockReuse{mSinkBubbleLength > 0 ? false : enableBlockReuse} + , mEnableMlaVScalePool{enableMlaVScalePool} { // When num_layers < len(maxAttentionWindowVec), not all window sizes in the // repeating pattern are used. Update mMaxAttentionWindow to the actual @@ -3253,13 +3268,13 @@ KVCacheManager::KVCacheManager(SizeType32 numLayers, SizeType32 numKvHeads, Size std::shared_ptr kvCacheConnectorManager, bool enableIndexerKCache, SizeType32 indexerKCacheQuantBlockSize, SizeType32 indexerKCacheIndexHeadDim, bool indexerKCacheUseFp4, std::optional linearAttentionMetadata, - std::vector const& poolConfigurations) + std::vector const& poolConfigurations, bool enableMlaVScalePool) : KVCacheManager(std::vector(numLayers, numKvHeads), sizePerHead, tokensPerBlock, blocksPerWindow, maxNumSequences, maxBeamWidth, maxAttentionWindowVec, dtype, sinkTokenLength, std::move(stream), maxSequenceLength, chunkSize, enableBlockReuse, cacheType, secondaryOffloadMinPriority, std::move(eventManager), enablePartialReuse, copyOnPartialReuse, std::move(kvCacheConnectorManager), enableIndexerKCache, indexerKCacheQuantBlockSize, indexerKCacheIndexHeadDim, indexerKCacheUseFp4, linearAttentionMetadata, - poolConfigurations) + poolConfigurations, enableMlaVScalePool) { } @@ -3292,6 +3307,20 @@ void KVCacheManager::allocatePools(bool useUvm) cacheSizeBytes += (cacheVolume * 4) / 8; } } + if (mEnableMlaVScalePool) + { +#ifdef ENABLE_FP4 + TLLM_CHECK_WITH_INFO( + mDataType == nvinfer1::DataType::kFP4, "MLA V-scale pool is only supported for FP4 KV cache."); + auto const elemsPerPage = getMlaVScaleElemsPerPage(mSizePerHead, mTokensPerBlock); + auto const vScaleShape + = ITensor::makeShape({mBlockManager.getNumLayers(), mBlockManager.getNumPrimaryBlocks(), elemsPerPage}); + mMlaVScalePool = BufferManager::gpuSync(vScaleShape, nvinfer1::DataType::kFP8); + cacheSizeBytes += ITensor::volume(vScaleShape) * BufferDataType(nvinfer1::DataType::kFP8).getSize(); +#else + TLLM_THROW("MLA V-scale pool requires FP4 support."); +#endif + } // Save the total number of bytes allocated for the KV-cache for KvCacheStats mAllocatedBytes = cacheSizeBytes; if (tc::Logger::getLogger()->getLevel() <= tc::Logger::INFO) @@ -3350,6 +3379,10 @@ void KVCacheManager::allocatePools(bool useUvm) void KVCacheManager::releasePools() { mBlockManager.releasePools(); + if (mMlaVScalePool) + { + mMlaVScalePool->release(); + } } void KVCacheManager::startScheduling() diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp index 12b29d4981e2..99912321afe0 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp @@ -548,6 +548,19 @@ void tb::kv_cache_manager::KVCacheManagerBindings::initBindings(nb::module_& m) "get_indexer_k_cache_pool", [](tbk::BaseKVCacheManager& self) -> at::Tensor { return tr::Torch::tensor(self.getIndexerKCachePool()); }, nb::call_guard()) + .def( + "get_mla_v_scale_pool", + [](tbk::BaseKVCacheManager& self) + { + std::optional mla_v_scale_pool{std::nullopt}; + auto tensor = self.getMlaVScalePool(); + if (tensor) + { + mla_v_scale_pool = tr::Torch::tensor(tensor); + } + return mla_v_scale_pool; + }, + nb::call_guard()) .def( "get_unique_primary_pool", [](tbk::BaseKVCacheManager& self) { return self.getUniquePrimaryPool(); }, nb::call_guard()) @@ -655,7 +668,7 @@ void tb::kv_cache_manager::KVCacheManagerBindings::initBindings(nb::module_& m) tbk::CacheType, std::optional, std::shared_ptr, bool, bool, std::shared_ptr, bool, SizeType32, SizeType32, bool, std::optional, - std::vector const&>(), + std::vector const&, bool>(), nb::arg("num_kv_heads_per_layer"), nb::arg("size_per_head"), nb::arg("tokens_per_block"), nb::arg("blocks_per_window"), nb::arg("max_num_sequences"), nb::arg("max_beam_width"), nb::arg("max_attention_window_vec"), nb::arg("dtype"), nb::arg("sink_token_length"), nb::arg("stream"), @@ -667,7 +680,7 @@ void tb::kv_cache_manager::KVCacheManagerBindings::initBindings(nb::module_& m) nb::arg("indexer_k_cache_index_head_dim") = 0, nb::arg("indexer_k_cache_use_fp4") = false, nb::arg("linear_attention_metadata").none() = std::nullopt, nb::arg("pool_configurations") = std::vector{}, - nb::call_guard()) + nb::arg("enable_mla_v_scale_pool") = false, nb::call_guard()) .def( "scheduling_has_free_blocks", [](tbk::KVCacheManager& self, SizeType32 numRequired, SizeType32 windowSize) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index 7ced64dec19a..fff0a90b663d 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -10,18 +10,27 @@ from flashinfer.jit.core import check_cuda_arch from typing_extensions import Self +import tensorrt_llm.bindings from tensorrt_llm._torch.pyexecutor.sampling_utils import torch_multi_arange +from tensorrt_llm._utils import prefer_pinned from tensorrt_llm.functional import AttentionMaskType from tensorrt_llm.logger import logger from tensorrt_llm.models.modeling_utils import QuantConfig from ..metadata import KVCacheParams from ..utils import get_global_attrs, get_model_extra_attrs +from .fp4_mla_kv import (FP4_MLA_KV_GLOBAL_SCALE, HP_BLOCK_SIZE, + get_fp4_mla_decode_cache, + is_flashinfer_fp4_mla_attention_enabled, + run_fp4_mla_attention_decode, scatter_fp4_mla_kv_cache, + update_hp_kv_for_fp4_mla) from .interface import (AttentionBackend, AttentionForwardArgs, AttentionInputType, AttentionMetadata, CustomAttentionMask, MLAParams, PredefinedAttentionMask, merge_attention_forward_args) +_DataType = tensorrt_llm.bindings.DataType + try: check_cuda_arch() except RuntimeError: @@ -31,8 +40,6 @@ arch_list = f"{capability[0]}.{capability[1]}" os.environ["TORCH_CUDA_ARCH_LIST"] = arch_list -from tensorrt_llm._utils import prefer_pinned - _FORCE_RAGGED_FA2 = False """Used for testing.""" @@ -156,6 +163,61 @@ class FlashInferAttentionMetadata(AttentionMetadata): _mla_qo_indptr_buf: Optional[torch.Tensor] = field(init=False, default=None) _mla_kv_len_arr_buf: Optional[torch.Tensor] = field(init=False, default=None) + _fp4_mla_decode_kv_indices_buf: Optional[torch.Tensor] = field(init=False, + default=None) + _fp4_mla_decode_cache_buf: Optional[torch.Tensor] = field(init=False, + default=None) + _fp4_mla_global_scale: Optional[torch.Tensor] = field(init=False, + default=None) + + # --- MLA FP4 KV cache machinery (mirrors TrtllmAttentionMetadata). --- + # Per-request stable seq_slot from SeqSlotManager, for indexing into the + # high-precision BF16 KV pool. Populated by PyTorchModelEngine via + # hasattr(metadata, 'seq_slots'). + seq_slots: Optional[torch.Tensor] = field(init=False, default=None) + seq_slots_cpu: Optional[torch.Tensor] = field(init=False, default=None) + # BF16 circular buffer, shape + # [max_num_sequences, num_local_layers, kv_factor=1, HP_BLOCK_SIZE * head_dim]. + # Holds up to HP_BLOCK_SIZE most-recent latent vectors per seq. The + # dequant fallback overlays BF16 tail tokens from it; the no-dequant path + # uses it to requantize the active 16-token FP4 KV tile. + high_precision_kv_pool: Optional[torch.Tensor] = field(init=False, + default=None) + # Auxiliary FP4 MLA V-scale pool for the no-dequant PV path. The + # physical storage is flat per [local_layer, physical_page]; callers view + # it with get_fp4_mla_v_scale_pool_view(..., v_head_dim=kv_lora_rank). + fp4_mla_v_scale_pool: Optional[torch.Tensor] = field(init=False, + default=None) + _fp4_mla_attention_q_buf: Optional[torch.Tensor] = field(init=False, + default=None) + _fp4_mla_attention_p_buf: Optional[torch.Tensor] = field(init=False, + default=None) + _fp4_mla_attention_p_sf_buf: Optional[torch.Tensor] = field(init=False, + default=None) + _fp4_mla_attention_max_buf: Optional[torch.Tensor] = field(init=False, + default=None) + _fp4_mla_attention_denom_buf: Optional[torch.Tensor] = field(init=False, + default=None) + # Debug ownership map: seq_slot to last request_id that wrote it. + hp_pool_owners: Optional[dict] = field(init=False, default=None) + # True during warmup forward passes (dummy requests, no real data). + is_warmup: bool = field(init=False, default=False) + + # Runtime aliases consumed by the shared HP-pool update helper. + # Same naming as TrtllmAttentionMetadata so the helper can duck-type. + kv_lens_cuda_runtime: Optional[torch.Tensor] = field(init=False, + default=None) + prompt_lens_cuda_runtime: Optional[torch.Tensor] = field(init=False, + default=None) + prompt_lens_cpu_runtime: Optional[torch.Tensor] = field(init=False, + default=None) + # Stable backing buffers for the runtime slices above (allocated only when + # NVFP4 KV + MLA is active, to avoid memory cost on the default path). + _kv_lens_cuda_buf: Optional[torch.Tensor] = field(init=False, default=None) + _prompt_lens_cuda_buf: Optional[torch.Tensor] = field(init=False, + default=None) + _prompt_lens_cpu_buf: Optional[torch.Tensor] = field(init=False, + default=None) def needs_plan(self, plan_params: PlanParams) -> bool: if plan_params not in self._plan_params_to_wrappers: @@ -252,12 +314,15 @@ def plan_mla_decode( by prepare() so it runs outside of CUDA graph capture. """ if self._mla_decode_wrapper is None: + kv_indices_buf = (self._fp4_mla_decode_kv_indices_buf + if self.high_precision_kv_pool is not None else + self._paged_kv_indices) self._mla_decode_wrapper = flashinfer.mla.BatchMLAPagedAttentionWrapper( self.workspace_buffer, use_cuda_graph=self.is_cuda_graph, qo_indptr=self._mla_qo_indptr_buf, kv_indptr=self.paged_kv_indptr_decode, - kv_indices=self._paged_kv_indices, + kv_indices=kv_indices_buf, kv_len_arr=self._mla_kv_len_arr_buf, backend="auto", ) @@ -364,9 +429,19 @@ def _do_plan_mla_decode(self, plan_params: MLAPlanParams) -> None: """ num_gen = self.num_generations kv_indptr = self.paged_kv_indptr_decode[:num_gen + 1] - kv_indices = self._paged_kv_indices[self.num_context_blocks:self. - num_context_blocks + - self.num_generation_blocks] + if self.high_precision_kv_pool is not None: + kv_indices_buf = self._fp4_mla_decode_kv_indices_buf[:self. + num_generation_blocks] + compact_indices = torch.arange(self.num_generation_blocks, + dtype=torch.int32, + device=kv_indices_buf.device) + kv_indices_buf.copy_(compact_indices, non_blocking=True) + kv_indices = kv_indices_buf.clone() + else: + kv_indices_start = self.num_context_blocks + kv_indices_end = kv_indices_start + self.num_generation_blocks + kv_indices = self._paged_kv_indices[ + kv_indices_start:kv_indices_end].clone() kv_last_page = self._paged_kv_last_page_len[self.num_contexts:self. num_contexts + num_gen] @@ -596,6 +671,193 @@ def _post_init_with_buffers(self, buffers) -> None: self._mla_context_planned = False self._mla_decode_planned = False + # --- MLA FP4 KV cache buffers (only allocated when MLA + NVFP4). --- + if (self.kv_cache_manager is not None + and self.kv_cache_manager.kv_factor == 1 + and self.kv_cache_manager.dtype == _DataType.NVFP4): + self._allocate_fp4_mla_buffers(buffers, capture_graph) + + def _allocate_fp4_mla_buffers(self, buffers, capture_graph: bool) -> None: + """Allocate the HP BF16 KV pool, seq_slots, and runtime-alias backing + buffers used by ``fp4_mla_kv.update_hp_kv_for_fp4_mla``.""" + max_num_sequences = (self.max_num_sequences if self.max_num_sequences + is not None else self.max_num_requests) + + self.seq_slots = self.get_empty( + buffers, + (max_num_sequences, ), + cache_name="seq_slots", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self.seq_slots_cpu = torch.empty( + max_num_sequences, + dtype=torch.int32, + device='cpu', + pin_memory=prefer_pinned(), + ) + self.hp_pool_owners = {} + + num_local_layers = self.kv_cache_manager.num_local_layers + head_dim = self.kv_cache_manager.head_dim + kv_factor = self.kv_cache_manager.kv_factor + hp_pool_shape = [ + max_num_sequences, num_local_layers, kv_factor, + HP_BLOCK_SIZE * head_dim + ] + existing_hp_pool = self.high_precision_kv_pool + if (capture_graph and existing_hp_pool is not None + and existing_hp_pool.dtype == torch.bfloat16 + and existing_hp_pool.device.type == "cuda" + and len(existing_hp_pool.shape) == len(hp_pool_shape) + and all(existing_hp_pool.shape[idx] >= dim + for idx, dim in enumerate(hp_pool_shape))): + # HP KV is persistent seq-slot state. CUDA graph metadata must + # share it instead of reserving one full pool per captured graph. + self.high_precision_kv_pool = existing_hp_pool + else: + self.high_precision_kv_pool = self.get_empty( + buffers, + hp_pool_shape, + cache_name="high_precision_kv_pool", + dtype=torch.bfloat16, + capture_graph=capture_graph, + ) + + max_num_pages = self.kv_cache_manager.blocks_in_primary_pool + self._fp4_mla_decode_kv_indices_buf = self.get_empty( + buffers, + (max_num_pages, ), + cache_name="_fp4_mla_decode_kv_indices_buf", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self._fp4_mla_global_scale = self.get_empty( + buffers, + (1, ), + cache_name="_fp4_mla_global_scale", + dtype=torch.float32, + capture_graph=capture_graph, + ) + self._fp4_mla_global_scale.fill_(FP4_MLA_KV_GLOBAL_SCALE) + + if is_flashinfer_fp4_mla_attention_enabled(): + self.fp4_mla_v_scale_pool = self.kv_cache_manager.get_mla_v_scale_pool( + ) + if self.fp4_mla_v_scale_pool is None: + raise RuntimeError( + "FP4 MLA attention requires the C++ KV cache manager to " + "allocate the V-scale pool.") + + # Runtime-alias backing buffers: GPU for kv/prompt lens, CPU pinned for + # the helper's prompt_lens_cpu read. + self._kv_lens_cuda_buf = self.get_empty( + buffers, + (max_num_sequences, ), + cache_name="fp4_mla_kv_lens_cuda", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self._prompt_lens_cuda_buf = self.get_empty( + buffers, + (max_num_sequences, ), + cache_name="fp4_mla_prompt_lens_cuda", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self._prompt_lens_cpu_buf = torch.empty( + max_num_sequences, + dtype=torch.int32, + device='cpu', + pin_memory=prefer_pinned(), + ) + + logger.info( + f"FlashInfer MLA + NVFP4 KV: HP pool shape=" + f"{list(self.high_precision_kv_pool.shape)}, " + f"size={self.high_precision_kv_pool.nbytes / (1 << 20):.1f} MB") + if self.fp4_mla_v_scale_pool is not None: + logger.info( + f"FlashInfer MLA + NVFP4 KV: V scale pool shape=" + f"{list(self.fp4_mla_v_scale_pool.shape)}, " + f"size={self.fp4_mla_v_scale_pool.nbytes / (1 << 20):.1f} MB") + + def _populate_fp4_mla_runtime_aliases(self, kv_lens: torch.Tensor) -> None: + """Populate kv_lens / prompt_lens runtime-alias slices for use by the + shared HP-pool update helper. + + ``kv_lens`` is the CPU int tensor holding total KV length per sequence + after the current forward pass (see ``prepare`` where it's computed as + ``cached_token_lens + seq_lens_kv_cuda``). + """ + num_seqs = self.num_contexts + self.num_generations + kv_lens_int32 = kv_lens[:num_seqs].to(torch.int32) + self._kv_lens_cuda_buf[:num_seqs].copy_(kv_lens_int32, + non_blocking=True) + self.kv_lens_cuda_runtime = self._kv_lens_cuda_buf[:num_seqs] + + # prompt_lens = number of new tokens per sequence this forward pass. + # For self-attention this equals seq_lens_kv (= seq_lens). Use the + # int32 view stored on self.seq_lens_kv_cuda to keep the dtype stable. + prompt_lens = self.seq_lens_kv_cuda[:num_seqs].to(torch.int32) + self._prompt_lens_cuda_buf[:num_seqs].copy_(prompt_lens, + non_blocking=True) + self.prompt_lens_cuda_runtime = self._prompt_lens_cuda_buf[:num_seqs] + + # CPU pinned mirror (used by the helper to compute token offsets). + self._prompt_lens_cpu_buf[:num_seqs].copy_(prompt_lens.cpu(), + non_blocking=False) + self.prompt_lens_cpu_runtime = self._prompt_lens_cpu_buf[:num_seqs] + + def _populate_fp4_mla_batch_indices_positions(self) -> None: + """Populate append/scatter token metadata without FlashInfer helpers. + + FlashInfer's helper kernels are not reliable for the 128-token pages + required by the no-dequant FP4 MLA path. FP4 MLA only needs the generic + ragged append metadata: + batch_indices[token] = sequence index in this scheduled batch + positions[token] = absolute KV position written by that token + """ + num_seqs = self.num_contexts + self.num_generations + if num_seqs == 0 or self.num_tokens == 0: + return + + device = self._batch_indices.device + seq_lens = self.seq_lens_kv_cuda[:num_seqs].to(torch.int32) + seq_ids = torch.arange(num_seqs, dtype=torch.int32, device=device) + batch_indices = torch.repeat_interleave(seq_ids, + seq_lens, + output_size=self.num_tokens) + + # Per-sequence offsets of NEW KV tokens in the ragged batch. + # Compute this from seq_lens_kv_cuda instead of qo_indptr so the + # append positions stay correct if query lengths and KV lengths diverge. + kv_token_starts = torch.empty((num_seqs, ), + dtype=torch.int32, + device=device) + kv_token_starts[0].zero_() + if num_seqs > 1: + torch.cumsum(seq_lens[:-1], + dim=0, + dtype=torch.int32, + out=kv_token_starts[1:]) + kv_start = torch.repeat_interleave(kv_token_starts, + seq_lens, + output_size=self.num_tokens) + cached_start = torch.repeat_interleave( + self.cached_token_lens[:num_seqs].to(torch.int32), + seq_lens, + output_size=self.num_tokens, + ) + token_offsets = torch.arange(self.num_tokens, + dtype=torch.int32, + device=device) + positions = token_offsets - kv_start + cached_start + + self._batch_indices[:self.num_tokens].copy_(batch_indices, + non_blocking=True) + self._positions[:self.num_tokens].copy_(positions, non_blocking=True) + def create_cuda_graph_metadata(self, max_batch_size: int, sub_cross_metadata: bool = False, @@ -781,6 +1043,11 @@ def prepare(self) -> None: # number of tokens needed in the kv cache for each sequence after the next pass kv_lens = self.cached_token_lens + self.seq_lens_kv_cuda + # Populate runtime aliases consumed by the shared HP-pool update helper + # (fp4_mla_kv.update_hp_kv_for_fp4_mla). Only active when MLA + NVFP4. + if self.high_precision_kv_pool is not None: + self._populate_fp4_mla_runtime_aliases(kv_lens) + # start and end indices of each sequence in the ragged key and value # for self attention it's the same as qo_indptr so avoid computing twice. if self.is_cross: @@ -888,17 +1155,20 @@ def prepare(self) -> None: # For cross attention, num_tokens is 0 during decode, and we don't need to update kv cache. if self.num_tokens > 0: - batch_indices, positions = flashinfer.get_batch_indices_positions( - self.kv_indptr, - flashinfer.get_seq_lens(self.paged_kv_indptr, - self.paged_kv_last_page_len, - self.page_size), - self.num_tokens, - ) - self._batch_indices[:batch_indices.size(0)].copy_(batch_indices, - non_blocking=True) - self._positions[:positions.size(0)].copy_(positions, - non_blocking=True) + if self.high_precision_kv_pool is not None: + self._populate_fp4_mla_batch_indices_positions() + else: + batch_indices, positions = flashinfer.get_batch_indices_positions( + self.kv_indptr, + flashinfer.get_seq_lens(self.paged_kv_indptr, + self.paged_kv_last_page_len, + self.page_size), + self.num_tokens, + ) + self._batch_indices[:batch_indices.size(0)].copy_( + batch_indices, non_blocking=True) + self._positions[:positions.size(0)].copy_(positions, + non_blocking=True) # Multi-wrapper case (Gemma4 hybrid: different head_dim per layer) # shares one workspace_buffer; eager plan() would overwrite earlier @@ -1240,12 +1510,26 @@ def __init__( self.qk_nope_head_dim = mla_params.qk_nope_head_dim self.v_head_dim = mla_params.v_head_dim + # Phase 1 restriction: NVFP4 KV cache on the FlashInfer backend is + # only supported for MLA (DeepSeek-style) models. Non-MLA dense/GQA + # FlashInfer paths would need a separate scatter/dequant path that has + # not been implemented. + if getattr(self, "has_fp4_kv_cache", False) and not self.is_mla_enable: + raise NotImplementedError( + "NVFP4 KV cache on the FlashInfer attention backend is only " + "supported for MLA models. Set attn_backend='TRTLLM' for " + "non-MLA FP4 KV cache, or use BF16/FP8 KV cache with " + "FlashInfer.") + def update_quant_config(self, new_quant_config: Optional[QuantConfig]): self.quant_config = new_quant_config self.has_fp8_kv_cache = False + self.has_fp4_kv_cache = False if self.quant_config: self.has_fp8_kv_cache = self.quant_config.layer_quant_mode.has_fp8_kv_cache( ) + self.has_fp4_kv_cache = self.quant_config.layer_quant_mode.has_fp4_kv_cache( + ) @staticmethod def _process_multi_item_part_lens( @@ -1358,6 +1642,13 @@ def mla_rope_generation( # q_pe shape: [num_tokens, num_heads, qk_rope_head_dim] fused_q[..., self.kv_lora_rank:] = q_pe + def _local_layer_idx(self, metadata: "FlashInferAttentionMetadata") -> int: + """Layer index within the local pipeline-parallel slice. Mirrors + ``TrtllmAttention.get_local_layer_idx``.""" + if metadata.kv_cache_manager is None: + return self.layer_idx + return metadata.kv_cache_manager.layer_offsets[self.layer_idx] + def _get_mla_caches( self, metadata: "FlashInferAttentionMetadata", @@ -1397,7 +1688,13 @@ def _mla_forward_context( output: torch.Tensor, latent_cache: torch.Tensor, ) -> None: - """MLA context phase: append latent to MLA caches, run ragged prefill.""" + """MLA context phase: append latent to MLA caches, run ragged prefill. + + With NVFP4 KV cache, attention still runs in BF16 over the caller's + BF16 q/k/v (no history to read during context). The only side effect + of the cache write is a quantize-and-scatter of the new ckv/kpe into + the FP4 paged pool plus a BF16 mirror of the tail into the HP pool. + """ # 1. Append latent_cache to separate ckv/kpe paged caches. # latent_cache shape: [num_ctx_tokens, kv_lora_rank + qk_rope_head_dim] num_ctx_tokens = metadata.num_ctx_tokens @@ -1410,22 +1707,47 @@ def _mla_forward_context( append_ckv = append_ckv.to(kv_dtype) append_kpe = append_kpe.to(kv_dtype) - ckv_cache, kpe_cache = self._get_mla_caches(metadata) - - ctx_batch_indices = metadata.batch_indices[:num_ctx_tokens] - ctx_positions = metadata.positions[:num_ctx_tokens] - - flashinfer.page.append_paged_mla_kv_cache( - append_ckv, - append_kpe, - ctx_batch_indices, - ctx_positions, - ckv_cache, - kpe_cache, - metadata.paged_kv_indices, - metadata.paged_kv_indptr, - metadata.paged_kv_last_page_len, - ) + if self.has_fp4_kv_cache: + # MLA.forward_impl slices latent_cache to [:num_ctx_tokens] before + # dispatching to the context path. The FP4 scatter's grid is + # latent_cache.shape[0], so a regression that passes a full-batch + # latent here would over-write gen tokens. Fail loudly instead of + # silently corrupting the cache. + assert latent_cache.shape[0] == num_ctx_tokens, ( + f"FP4 MLA context scatter expected latent_cache of shape " + f"[num_ctx_tokens={num_ctx_tokens}, ...] but got " + f"{list(latent_cache.shape)}. Did MLA.forward_impl stop " + f"pre-slicing latent_cache per phase?") + scatter_fp4_mla_kv_cache( + metadata, + latent_cache, + self.layer_idx, + token_offset=0, + phase="context", + local_layer=self._local_layer_idx(metadata), + v_head_dim=self.kv_lora_rank, + ) + update_hp_kv_for_fp4_mla(metadata, + latent_cache, + self._local_layer_idx(metadata), + phase="context") + else: + ckv_cache, kpe_cache = self._get_mla_caches(metadata) + + ctx_batch_indices = metadata.batch_indices[:num_ctx_tokens] + ctx_positions = metadata.positions[:num_ctx_tokens] + + flashinfer.page.append_paged_mla_kv_cache( + append_ckv, + append_kpe, + ctx_batch_indices, + ctx_positions, + ckv_cache, + kpe_cache, + metadata.paged_kv_indices, + metadata.paged_kv_indptr, + metadata.paged_kv_last_page_len, + ) # 2. Run ragged prefill with expanded q, k, v num_contexts = metadata.num_contexts @@ -1469,32 +1791,83 @@ def _mla_forward_generation( kv_dtype = q.dtype if self.has_fp8_kv_cache: kv_dtype = torch.float8_e4m3fn - ckv_cache, kpe_cache = self._get_mla_caches(metadata) - assert latent_cache is not None, ( - "FlashInfer MLA generation requires latent_cache.") - # Append latent_cache to the paged MLA KV cache first. - # latent_cache shape: [num_tokens, kv_lora_rank + qk_rope_head_dim] - # RoPE must already be applied to the k_pe portion before calling this. - append_ckv = latent_cache[:, :self.kv_lora_rank] - append_kpe = latent_cache[:, self.kv_lora_rank:] - if self.has_fp8_kv_cache: - append_ckv = append_ckv.to(kv_dtype) - append_kpe = append_kpe.to(kv_dtype) - num_ctx_tokens = metadata.num_ctx_tokens - gen_batch_indices = metadata.batch_indices[num_ctx_tokens:] - gen_positions = metadata.positions[num_ctx_tokens:] - flashinfer.page.append_paged_mla_kv_cache( - append_ckv, - append_kpe, - gen_batch_indices, - gen_positions, - ckv_cache, - kpe_cache, - metadata.paged_kv_indices, - metadata.paged_kv_indptr, - metadata.paged_kv_last_page_len, - ) + if self.has_fp4_kv_cache: + use_fp4_attention = is_flashinfer_fp4_mla_attention_enabled() + if latent_cache is not None: + # MLA.forward_impl slices latent_cache to [num_ctx_tokens:] + # before dispatching to the generation path. The FP4 scatter's + # grid is latent_cache.shape[0], so a regression that passes a + # full-batch latent here would OOB-read batch_indices. Fail + # loudly instead of silently corrupting the cache. + num_gen_tokens = metadata.num_tokens - metadata.num_ctx_tokens + assert latent_cache.shape[0] == num_gen_tokens, ( + f"FP4 MLA generation scatter expected latent_cache of " + f"shape [num_gen_tokens={num_gen_tokens}, ...] but got " + f"{list(latent_cache.shape)}. Did MLA.forward_impl stop " + f"pre-slicing latent_cache per phase?") + if use_fp4_attention: + update_hp_kv_for_fp4_mla(metadata, + latent_cache, + self._local_layer_idx(metadata), + phase="generation") + scatter_fp4_mla_kv_cache( + metadata, + latent_cache, + self.layer_idx, + token_offset=metadata.num_ctx_tokens, + phase="generation", + local_layer=self._local_layer_idx(metadata), + v_head_dim=self.kv_lora_rank, + ) + else: + scatter_fp4_mla_kv_cache( + metadata, + latent_cache, + self.layer_idx, + token_offset=metadata.num_ctx_tokens, + ) + update_hp_kv_for_fp4_mla(metadata, + latent_cache, + self._local_layer_idx(metadata), + phase="generation") + if not use_fp4_attention: + combined_cache = get_fp4_mla_decode_cache( + metadata, + self.layer_idx, + self._local_layer_idx(metadata), + head_dim=self.kv_lora_rank + self.qk_rope_head_dim, + dtype=kv_dtype, + ) + ckv_cache = combined_cache[..., :self.kv_lora_rank] + kpe_cache = combined_cache[..., self.kv_lora_rank:] + else: + use_fp4_attention = False + ckv_cache, kpe_cache = self._get_mla_caches(metadata) + + # If latent_cache is provided, append it to the paged MLA KV cache first. + # latent_cache shape: [num_tokens, kv_lora_rank + qk_rope_head_dim] + # RoPE must already be applied to the k_pe portion before calling this. + if latent_cache is not None: + append_ckv = latent_cache[:, :self.kv_lora_rank] + append_kpe = latent_cache[:, self.kv_lora_rank:] + if self.has_fp8_kv_cache: + append_ckv = append_ckv.to(kv_dtype) + append_kpe = append_kpe.to(kv_dtype) + num_ctx_tokens = metadata.num_ctx_tokens + gen_batch_indices = metadata.batch_indices[num_ctx_tokens:] + gen_positions = metadata.positions[num_ctx_tokens:] + flashinfer.page.append_paged_mla_kv_cache( + append_ckv, + append_kpe, + gen_batch_indices, + gen_positions, + ckv_cache, + kpe_cache, + metadata.paged_kv_indices, + metadata.paged_kv_indptr, + metadata.paged_kv_last_page_len, + ) # fused_q layout: [num_tokens, num_heads * (kv_lora_rank + qk_rope_head_dim)] # Split into q_nope (absorbed) and q_pe (rope) @@ -1510,6 +1883,22 @@ def _mla_forward_generation( else: sm_scale = 1.0 / math.sqrt(qk_head_dim) + if use_fp4_attention: + out_view = output[:num_tokens].view(-1, self.num_heads, + self.kv_lora_rank) + run_fp4_mla_attention_decode( + metadata, + self.layer_idx, + self._local_layer_idx(metadata), + q_nope, + q_pe, + out_view, + sm_scale=sm_scale, + kv_lora_rank=self.kv_lora_rank, + qk_rope_head_dim=self.qk_rope_head_dim, + ) + return + plan_params = MLAPlanParams( num_heads=self.num_heads, kv_lora_rank=self.kv_lora_rank, diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py new file mode 100644 index 000000000000..18fa22e130c4 --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py @@ -0,0 +1,1287 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Triton kernels for MLA FP4 KV-cache helpers.""" + +import triton +import triton.language as tl + +# HP pool maintenance kernels + + +@triton.jit +def _hp_kv_store_context_kernel( + pool_ptr, + latent_cache_ptr, + seq_slots_ptr, # int32 [num_contexts], seq_slot for each ctx seq + kv_lens_ptr, # int32 [num_contexts], total KV length after prefill + token_offsets_ptr, # int32 [num_contexts], exclusive prompt prefix sum + prompt_lens_ptr, # int32 [num_contexts], new tokens per ctx seq + num_seq_slots, + num_layers, + layer_idx, + pool_stride_seq, # pool.stride(0): elements between adjacent seq slots + pool_stride_layer, # pool.stride(1): elements between adjacent layers + lc_stride, # latent_cache.stride(0): elements between adjacent tokens + D: tl.constexpr, # head_dim (runtime dimension, = latent_cache.shape[-1]) + POOL_HEAD_D: tl.constexpr, + BLOCK_D: tl.constexpr, # next_power_of_2(D), used for vectorised load/store + HP_BLOCK: tl.constexpr, # = HP_BLOCK_SIZE (16) +): + """Store the tail tokens of each context sequence into the HP KV pool. + + Grid: (num_contexts, HP_BLOCK_SIZE). + Only programs for tail tokens present in latent_cache actually write. + + For a context sequence with total KV length L = num_cached + prompt_len: + - remainder = L % HP_BLOCK + - The last ``remainder`` new tokens (latent_cache positions + [offset + prompt_len - remainder, offset + prompt_len)) are stored + into pool slots [0, remainder), which correspond to the absolute token + positions [L - remainder, L) in the circular buffer. + """ + ctx_idx = tl.program_id(0) + buf_pos = tl.program_id(1) + if (layer_idx < 0) | (layer_idx >= num_layers): + return + + kv_len = tl.load(kv_lens_ptr + ctx_idx) + remainder = kv_len % HP_BLOCK + prompt_len = tl.load(prompt_lens_ptr + ctx_idx) + store_count = tl.minimum(remainder, prompt_len) + first_buf_pos = remainder - store_count + if buf_pos < first_buf_pos: + return + if buf_pos >= remainder: + return + + seq_slot = tl.load(seq_slots_ptr + ctx_idx).to(tl.int64) + if (seq_slot < 0) | (seq_slot >= num_seq_slots): + return + tok_offset = tl.load(token_offsets_ptr + ctx_idx).to(tl.int64) + + # Index of this token within latent_cache: last `remainder` new tokens, + # buf_pos-th of them (0-indexed from the start of the tail). + token_idx = tok_offset + prompt_len.to(tl.int64) - remainder.to(tl.int64) + buf_pos.to(tl.int64) + + offs_d = tl.arange(0, BLOCK_D) + mask_d = offs_d < D + safe_offs_d = tl.where(mask_d, offs_d, 0) + src = tl.load( + latent_cache_ptr + token_idx * lc_stride + safe_offs_d, + mask=mask_d, + other=0.0, + ) + + # Destination: pool[seq_slot, layer_idx, 0, buf_pos, :D]. + dst_base = seq_slot * pool_stride_seq + layer_idx * pool_stride_layer + buf_pos * POOL_HEAD_D + tl.store(pool_ptr + dst_base + safe_offs_d, src, mask=mask_d) + + +@triton.jit +def _hp_kv_store_gen_kernel( + pool_ptr, + latent_cache_ptr, + seq_slots_ptr, # int32 [num_gen], seq_slot for each gen seq + kv_lens_ptr, # int32 [num_gen], total KV length after this decode step + gen_tok_start, # int, offset in latent_cache where gen tokens begin + num_seq_slots, + num_layers, + layer_idx, + pool_stride_seq, + pool_stride_layer, + lc_stride, + D: tl.constexpr, + POOL_HEAD_D: tl.constexpr, + BLOCK_D: tl.constexpr, + HP_BLOCK: tl.constexpr, +): + """Store the current generation token into the HP KV pool. + + Grid: (num_gen_seqs,). + Each program stores one token into the circular buffer position + (kv_len - 1) % HP_BLOCK, overwriting the oldest entry. + """ + gen_idx = tl.program_id(0) + if (layer_idx < 0) | (layer_idx >= num_layers): + return + + seq_slot = tl.load(seq_slots_ptr + gen_idx).to(tl.int64) + if (seq_slot < 0) | (seq_slot >= num_seq_slots): + return + kv_len = tl.load(kv_lens_ptr + gen_idx) + if kv_len <= 0: + return + buf_pos = (kv_len - 1) % HP_BLOCK + + token_idx = tl.cast(gen_tok_start, tl.int64) + gen_idx.to(tl.int64) + offs_d = tl.arange(0, BLOCK_D) + mask_d = offs_d < D + safe_offs_d = tl.where(mask_d, offs_d, 0) + src = tl.load( + latent_cache_ptr + token_idx * lc_stride + safe_offs_d, + mask=mask_d, + other=0.0, + ) + + dst_base = seq_slot * pool_stride_seq + layer_idx * pool_stride_layer + buf_pos * POOL_HEAD_D + tl.store(pool_ptr + dst_base + safe_offs_d, src, mask=mask_d) + + +@triton.jit +def _fp4_mla_swizzled_sf_offset( + row_idx, + col_idx, + SF_PER_TOKEN: tl.constexpr, +): + padded_cols = ((SF_PER_TOKEN + 3) // 4) * 4 + col_in_group = col_idx % 4 + col_group = col_idx // 4 + row_in_group0 = row_idx % 32 + row_in_group1 = (row_idx % 128) // 32 + row_group = row_idx // 128 + return ( + col_in_group + + col_group * (4 * 128) + + row_in_group0 * 16 + + row_in_group1 * 4 + + row_group * (128 * padded_cols) + ) + + +@triton.jit +def _fp4_mla_scatter_kernel( + kv_cache_ptr, + sf_cache_ptr, + q_fp4_ptr, + q_sf_ptr, + batch_indices_ptr, + positions_ptr, + paged_kv_indices_ptr, + paged_kv_indptr_ptr, + page_ids_len, + indptr_len, + num_pages, + token_offset, + page_size, + kv_s0, + kv_s1, + kv_s2, + kv_s3, + kv_s4, + sf_s0, + sf_s1, + sf_s2, + sf_s3, + sf_s4, + q_fp4_s0, + q_fp4_s1, + q_sf_s0, + q_sf_s1, + PACKED_D: tl.constexpr, + SF_PER_TOKEN: tl.constexpr, + BLOCK_PACKED_D: tl.constexpr, + BLOCK_SF: tl.constexpr, + USE_SWIZZLED_SF: tl.constexpr, +): + token_idx = tl.program_id(0) + metadata_token_idx = token_offset + token_idx + + # Keep page address math in int64; 128-token FP4 pages can overflow + # int32 offsets in large KV pools. + batch_idx = tl.load(batch_indices_ptr + metadata_token_idx).to(tl.int64) + position = tl.load(positions_ptr + metadata_token_idx).to(tl.int64) + if (batch_idx < 0) | (batch_idx + 1 >= indptr_len) | (position < 0): + return + + page_size_i64 = tl.cast(page_size, tl.int64) + page_idx = position // page_size_i64 + page_pos = position - page_idx * page_size_i64 + page_start = tl.load(paged_kv_indptr_ptr + batch_idx).to(tl.int64) + page_end = tl.load(paged_kv_indptr_ptr + batch_idx + 1).to(tl.int64) + page_table_offset = page_start + page_idx + if ( + (page_pos < 0) + | (page_pos >= page_size_i64) + | (page_table_offset < page_start) + | (page_table_offset >= page_end) + | (page_table_offset < 0) + | (page_table_offset >= page_ids_len) + ): + return + physical_page = tl.load(paged_kv_indices_ptr + page_table_offset).to(tl.int64) + if (physical_page < 0) | (physical_page >= num_pages): + return + + offs_packed = tl.arange(0, BLOCK_PACKED_D) + mask_packed = offs_packed < PACKED_D + safe_offs_packed = tl.where(mask_packed, offs_packed, 0) + q_vals = tl.load( + q_fp4_ptr + token_idx * q_fp4_s0 + safe_offs_packed * q_fp4_s1, + mask=mask_packed, + other=0, + ) + kv_dst = physical_page * kv_s0 + page_pos * kv_s2 + tl.store(kv_cache_ptr + kv_dst + safe_offs_packed * kv_s4, q_vals, mask=mask_packed) + + offs_sf = tl.arange(0, BLOCK_SF) + mask_sf = offs_sf < SF_PER_TOKEN + # Masked lanes are predicated off, but the address arithmetic still runs. + # In the swizzled layout, out-of-range cols land past the per-page stride; + # in the linear layout, they spill ~(BLOCK_SF - SF_PER_TOKEN) bytes past + # each page row. For the last physical page either case can fall outside + # the sf_cache allocation. Pin masked lanes to col 0 so all computed + # addresses stay in-bounds regardless of allocator slack. + safe_offs_sf = tl.where(mask_sf, offs_sf, 0) + sf_vals = tl.load( + q_sf_ptr + token_idx * q_sf_s0 + safe_offs_sf * q_sf_s1, + mask=mask_sf, + other=0, + ) + if USE_SWIZZLED_SF: + sf_offsets = _fp4_mla_swizzled_sf_offset(page_pos, safe_offs_sf, SF_PER_TOKEN) + sf_dst = physical_page * sf_s0 + tl.store(sf_cache_ptr + sf_dst + sf_offsets, sf_vals, mask=mask_sf) + else: + sf_dst = physical_page * sf_s0 + page_pos * sf_s2 + tl.store(sf_cache_ptr + sf_dst + safe_offs_sf * sf_s4, sf_vals, mask=mask_sf) + + +# FP4 conversion and cache kernels + + +@triton.jit +def _fp4_e2m1_to_f32(nibble): + magnitude = nibble & 0x7 + value = tl.where( + magnitude == 0, + 0.0, + tl.where( + magnitude == 1, + 0.5, + tl.where( + magnitude == 2, + 1.0, + tl.where( + magnitude == 3, + 1.5, + tl.where( + magnitude == 4, + 2.0, + tl.where(magnitude == 5, 3.0, tl.where(magnitude == 6, 4.0, 6.0)), + ), + ), + ), + ), + ) + sign = (nibble & 0x8) != 0 + return tl.where(sign, -value, value) + + +@triton.jit +def _fp4_e2m1_quantize(x): + abs_x = tl.abs(x) + magnitude = tl.where( + abs_x < 0.25, + 0, + tl.where( + abs_x < 0.75, + 1, + tl.where( + abs_x < 1.25, + 2, + tl.where( + abs_x < 1.75, + 3, + tl.where(abs_x < 2.5, 4, tl.where(abs_x < 3.5, 5, tl.where(abs_x < 5.0, 6, 7))), + ), + ), + ), + ) + sign = tl.where(x < 0.0, 8, 0) + return (magnitude | sign).to(tl.uint8) + + +@triton.jit +def _fp4_mla_v_scale_store_context_tokens_kernel( + kv_cache_ptr, + sf_cache_ptr, + v_sf_ptr, + latent_cache_ptr, + global_scale_ptr, + batch_indices_ptr, + positions_ptr, + paged_kv_indices_ptr, + paged_kv_indptr_ptr, + page_ids_len, + indptr_len, + metadata_num_tokens, + num_pages, + num_layers, + token_offset, + num_tokens, + local_layer, + page_size, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + lc_s0, + lc_s1, + vsf_s0, + vsf_s1, + HEAD_D: tl.constexpr, + V_HEAD_D: tl.constexpr, + HP_BLOCK: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, +): + token_idx = tl.program_id(0) + dim_block = tl.program_id(1) + if (local_layer < 0) | (local_layer >= num_layers): + return + if token_idx >= num_tokens: + return + + metadata_token_idx = token_offset + token_idx + if metadata_token_idx >= metadata_num_tokens: + return + + batch_idx = tl.load(batch_indices_ptr + metadata_token_idx).to(tl.int64) + position = tl.load(positions_ptr + metadata_token_idx).to(tl.int64) + if (batch_idx < 0) | (batch_idx + 1 >= indptr_len) | (position < 0): + return + if position % HP_BLOCK != 0: + return + + page_size_i64 = tl.cast(page_size, tl.int64) + page_idx = position // page_size_i64 + page_pos = position - page_idx * page_size_i64 + page_start = tl.load(paged_kv_indptr_ptr + batch_idx).to(tl.int64) + page_end = tl.load(paged_kv_indptr_ptr + batch_idx + 1).to(tl.int64) + page_table_offset = page_start + page_idx + if ( + (page_pos < 0) + | (page_pos >= page_size_i64) + | (page_table_offset < page_start) + | (page_table_offset >= page_end) + | (page_table_offset < 0) + | (page_table_offset >= page_ids_len) + ): + return + physical_page = tl.load(paged_kv_indices_ptr + page_table_offset).to(tl.int64) + if (physical_page < 0) | (physical_page >= num_pages): + return + + byte_offsets = tl.arange(0, HP_BLOCK // 2) + token_offsets = tl.arange(0, HP_BLOCK) + even_d = dim_block * FP4_BLOCK + byte_offsets * 2 + odd_d = even_d + 1 + all_d = dim_block * FP4_BLOCK + tl.arange(0, FP4_BLOCK) + mask_even_d = even_d < HEAD_D + mask_odd_d = odd_d < HEAD_D + mask_all_d = all_d < HEAD_D + safe_even_d = tl.where(mask_even_d, even_d, 0) + safe_odd_d = tl.where(mask_odd_d, odd_d, 0) + safe_all_d = tl.where(mask_all_d, all_d, 0) + + token_candidates = token_idx + token_offsets + valid_tokens = token_candidates < num_tokens + candidate_metadata = token_offset + token_candidates + valid_tokens = valid_tokens & (candidate_metadata < metadata_num_tokens) + safe_candidate_metadata = tl.where(valid_tokens, candidate_metadata, 0) + + candidate_batch = tl.load( + batch_indices_ptr + safe_candidate_metadata, mask=valid_tokens, other=-1 + ).to(tl.int64) + candidate_pos = tl.load( + positions_ptr + safe_candidate_metadata, mask=valid_tokens, other=-1 + ).to(tl.int64) + valid_tokens = valid_tokens & (candidate_batch == batch_idx) + valid_tokens = valid_tokens & (candidate_pos == position + token_offsets) + # int64 so safe_token_candidates * lc_s0 doesn't overflow when num_tokens * head_dim > 2^31. + safe_token_candidates = tl.where(valid_tokens, token_candidates, 0).to(tl.int64) + + even_values = tl.load( + latent_cache_ptr + safe_token_candidates[:, None] * lc_s0 + safe_even_d[None, :] * lc_s1, + mask=valid_tokens[:, None] & mask_even_d[None, :], + other=0.0, + ).to(tl.float32) + odd_values = tl.load( + latent_cache_ptr + safe_token_candidates[:, None] * lc_s0 + safe_odd_d[None, :] * lc_s1, + mask=valid_tokens[:, None] & mask_odd_d[None, :], + other=0.0, + ).to(tl.float32) + amax_per_token = tl.maximum( + tl.max(tl.abs(even_values), axis=1), + tl.max(tl.abs(odd_values), axis=1), + ) + tile_amax = tl.max(amax_per_token, axis=0) + global_scale = tl.load(global_scale_ptr) + # K consumes scales as [token, dim-block], while V consumes scales as + # [dim, token-block]. Only the compressed-KV prefix has both views, so + # tail K-only dims keep K's per-token scale. + shared_tile = dim_block * FP4_BLOCK < V_HEAD_D + tile_scale = tl.where(tile_amax > 0.0, tile_amax / 6.0, 1.0) + token_scale = tl.where(amax_per_token > 0.0, amax_per_token / 6.0, 1.0) + local_scale = tl.where(shared_tile, tile_scale, token_scale) + stored_scale = local_scale * global_scale + v_stored_scale = tile_scale * global_scale + + low = _fp4_e2m1_quantize(even_values / local_scale[:, None]) + high = _fp4_e2m1_quantize(odd_values / local_scale[:, None]) + packed = low | (high << 4) + + packed_cols = dim_block * (FP4_BLOCK // 2) + byte_offsets + page_positions = page_pos + token_offsets + kv_base = physical_page * kv_s0 + tl.store( + kv_cache_ptr + kv_base + page_positions[:, None] * kv_s2 + packed_cols[None, :] * kv_s4, + packed, + mask=valid_tokens[:, None] & mask_even_d[None, :], + ) + + k_sf_offsets = _fp4_mla_swizzled_sf_offset(page_positions, dim_block, SF_PER_TOKEN) + tl.store(sf_cache_ptr + physical_page * sf_s0 + k_sf_offsets, stored_scale, mask=valid_tokens) + + token_scale_col = page_pos // HP_BLOCK + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_all_d, token_scale_col, SF_PER_PAGE) + v_sf_base = tl.cast(local_layer, tl.int64) * tl.cast( + vsf_s0, tl.int64 + ) + physical_page * tl.cast(vsf_s1, tl.int64) + tl.store( + v_sf_ptr + v_sf_base + sf_offsets.to(tl.int64), + v_stored_scale, + mask=mask_all_d & (all_d < V_HEAD_D), + ) + + +@triton.jit +def _fp4_mla_v_scale_store_hp_tail_kernel( + kv_cache_ptr, + sf_cache_ptr, + v_sf_ptr, + hp_pool_ptr, + global_scale_ptr, + seq_slots_ptr, + kv_lens_ptr, + page_ids_ptr, + paged_kv_indptr_ptr, + page_ids_len, + indptr_len, + num_pages, + num_seq_slots, + num_layers, + local_layer, + page_size, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + pool_s0, + pool_s1, + vsf_s0, + vsf_s1, + HEAD_D: tl.constexpr, + V_HEAD_D: tl.constexpr, + HP_BLOCK: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, +): + seq_idx = tl.program_id(0) + dim_block = tl.program_id(1) + if (local_layer < 0) | (local_layer >= num_layers): + return + if seq_idx + 1 >= indptr_len: + return + + kv_len = tl.load(kv_lens_ptr + seq_idx) + remainder = kv_len % HP_BLOCK + if kv_len == 0: + return + + tail_count = tl.where(remainder == 0, HP_BLOCK, remainder) + block_base_pos = kv_len - tail_count + page_idx = block_base_pos // page_size + page_pos = block_base_pos - page_idx * page_size + page_start = tl.load(paged_kv_indptr_ptr + seq_idx).to(tl.int64) + page_end = tl.load(paged_kv_indptr_ptr + seq_idx + 1).to(tl.int64) + physical_page_offset = page_start + page_idx + if ( + (page_pos < 0) + | (page_pos >= page_size) + | (physical_page_offset < page_start) + | (physical_page_offset >= page_end) + | (physical_page_offset < 0) + | (physical_page_offset >= page_ids_len) + ): + return + physical_page = tl.load(page_ids_ptr + physical_page_offset).to(tl.int64) + if (physical_page < 0) | (physical_page >= num_pages): + return + seq_slot = tl.load(seq_slots_ptr + seq_idx).to(tl.int64) + if (seq_slot < 0) | (seq_slot >= num_seq_slots): + return + + byte_offsets = tl.arange(0, HP_BLOCK // 2) + token_offsets = tl.arange(0, HP_BLOCK) + even_d = dim_block * FP4_BLOCK + byte_offsets * 2 + odd_d = even_d + 1 + all_d = dim_block * FP4_BLOCK + tl.arange(0, FP4_BLOCK) + mask_even_d = even_d < HEAD_D + mask_odd_d = odd_d < HEAD_D + mask_all_d = all_d < HEAD_D + safe_even_d = tl.where(mask_even_d, even_d, 0) + safe_odd_d = tl.where(mask_odd_d, odd_d, 0) + safe_all_d = tl.where(mask_all_d, all_d, 0) + + hp_slots = (block_base_pos + token_offsets) % HP_BLOCK + valid_tokens = token_offsets < tail_count + even_values = tl.load( + hp_pool_ptr + + seq_slot * pool_s0 + + local_layer * pool_s1 + + hp_slots[:, None] * HEAD_D + + safe_even_d[None, :], + mask=valid_tokens[:, None] & mask_even_d[None, :], + other=0.0, + ).to(tl.float32) + odd_values = tl.load( + hp_pool_ptr + + seq_slot * pool_s0 + + local_layer * pool_s1 + + hp_slots[:, None] * HEAD_D + + safe_odd_d[None, :], + mask=valid_tokens[:, None] & mask_odd_d[None, :], + other=0.0, + ).to(tl.float32) + amax_per_token = tl.maximum( + tl.max(tl.abs(even_values), axis=1), + tl.max(tl.abs(odd_values), axis=1), + ) + tile_amax = tl.max(amax_per_token, axis=0) + global_scale = tl.load(global_scale_ptr) + # K consumes scales as [token, dim-block], while V consumes scales as + # [dim, token-block]. Only the compressed-KV prefix has both views, so + # tail K-only dims keep K's per-token scale. + shared_tile = dim_block * FP4_BLOCK < V_HEAD_D + tile_scale = tl.where(tile_amax > 0.0, tile_amax / 6.0, 1.0) + token_scale = tl.where(amax_per_token > 0.0, amax_per_token / 6.0, 1.0) + local_scale = tl.where(shared_tile, tile_scale, token_scale) + stored_scale = local_scale * global_scale + v_stored_scale = tile_scale * global_scale + + low = _fp4_e2m1_quantize(even_values / local_scale[:, None]) + high = _fp4_e2m1_quantize(odd_values / local_scale[:, None]) + packed = low | (high << 4) + + packed_cols = dim_block * (FP4_BLOCK // 2) + byte_offsets + page_positions = page_pos + token_offsets + kv_base = physical_page * kv_s0 + tl.store( + kv_cache_ptr + kv_base + page_positions[:, None] * kv_s2 + packed_cols[None, :] * kv_s4, + packed, + mask=valid_tokens[:, None] & mask_even_d[None, :], + ) + + k_sf_offsets = _fp4_mla_swizzled_sf_offset(page_positions, dim_block, SF_PER_TOKEN) + tl.store(sf_cache_ptr + physical_page * sf_s0 + k_sf_offsets, stored_scale, mask=valid_tokens) + + token_scale_col = page_pos // HP_BLOCK + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_all_d, token_scale_col, SF_PER_PAGE) + v_sf_base = tl.cast(local_layer, tl.int64) * tl.cast( + vsf_s0, tl.int64 + ) + physical_page * tl.cast(vsf_s1, tl.int64) + tl.store( + v_sf_ptr + v_sf_base + sf_offsets.to(tl.int64), + v_stored_scale, + mask=mask_all_d & (all_d < V_HEAD_D), + ) + + +@triton.jit +def _fp4_mla_load_values( + kv_cache_ptr, + sf_cache_ptr, + physical_page, + page_pos, + offs_d, + mask_d, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + sf_s2, + sf_s4, + D: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_TOKEN: tl.constexpr, + USE_SWIZZLED_SF: tl.constexpr, +): + packed_offsets = offs_d // 2 + packed = tl.load( + kv_cache_ptr + physical_page * kv_s0 + page_pos * kv_s2 + packed_offsets * kv_s4, + mask=mask_d, + other=0, + ) + low = packed & 0x0F + high = (packed >> 4) & 0x0F + nibble = tl.where((offs_d & 1) == 0, low, high) + + scale_offsets = offs_d // FP4_BLOCK + # See _fp4_mla_scatter_kernel for the rationale: pin masked lanes to col 0 + # so the address arithmetic stays inside the per-page stride for the last + # physical page, regardless of which SF layout is in use. + safe_scale_offsets = tl.where(mask_d, scale_offsets, 0) + if USE_SWIZZLED_SF: + swizzled_sf_offsets = _fp4_mla_swizzled_sf_offset( + page_pos, safe_scale_offsets, SF_PER_TOKEN + ) + scale = tl.load( + sf_cache_ptr + physical_page * sf_s0 + swizzled_sf_offsets, + mask=mask_d, + other=0.0, + ).to(tl.float32) + else: + scale = tl.load( + sf_cache_ptr + physical_page * sf_s0 + page_pos * sf_s2 + safe_scale_offsets * sf_s4, + mask=mask_d, + other=0.0, + ).to(tl.float32) + return _fp4_e2m1_to_f32(nibble) * scale + + +@triton.jit +def _fp4_mla_dequant_kernel( + out_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + page_ids_len, + num_pages, + kv_s0, + kv_s1, + kv_s2, + kv_s3, + kv_s4, + sf_s0, + sf_s1, + sf_s2, + sf_s3, + sf_s4, + out_s0, + out_s1, + out_s2, + D: tl.constexpr, + FP4_BLOCK: tl.constexpr, + BLOCK_D: tl.constexpr, + USE_SWIZZLED_SF: tl.constexpr, +): + compact_page = tl.program_id(0).to(tl.int64) + page_pos = tl.program_id(1).to(tl.int64) + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, + mask=valid_compact_page, + other=-1, + ).to(tl.int64) + valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + global_scale = tl.load(global_scale_ptr) + + offs_d = tl.arange(0, BLOCK_D) + mask_d = offs_d < D + safe_offs_d = tl.where(mask_d, offs_d, 0) + value = _fp4_mla_load_values( + kv_cache_ptr, + sf_cache_ptr, + safe_physical_page, + page_pos, + safe_offs_d, + mask_d & valid_physical_page, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + sf_s2, + sf_s4, + D, + FP4_BLOCK, + D // FP4_BLOCK, + USE_SWIZZLED_SF, + ) + + tl.store( + out_ptr + compact_page * out_s0 + page_pos * out_s1 + safe_offs_d * out_s2, + value / global_scale, + mask=mask_d, + ) + + +@triton.jit +def _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_offsets, + token_offsets, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + NUM_HEADS: tl.constexpr, +): + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + q_rows = q_row_base + head_offsets + mask_h = head_offsets < NUM_HEADS + safe_q_rows = tl.where(mask_h, q_rows, q_row_base) + scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) + + packed_k_offsets = tl.arange(0, BLOCK_K // 2) + scale_offsets = tl.arange(0, BLOCK_K // FP4_BLOCK) + # fp4_quantize_with_residual lays out the last Q groups as + # [main_0, residual_0, main_1, residual_1, ...]. The KV cache stores + # each tail K group once, so QK maps both logical Q tail groups to the + # same physical K group. + residual_groups = Q_RESIDUAL_D // FP4_BLOCK + non_residual_groups = K_HEAD_D // FP4_BLOCK - residual_groups + for q_start in tl.range(0, Q_HEAD_D, BLOCK_K): + q_elem_offsets = q_start + packed_k_offsets * 2 + q_group_offsets = q_elem_offsets // FP4_BLOCK + k_group_offsets = tl.where( + q_group_offsets < non_residual_groups, + q_group_offsets, + non_residual_groups + (q_group_offsets - non_residual_groups) // 2, + ) + byte_offsets_in_group = (q_elem_offsets % FP4_BLOCK) // 2 + packed_q_cols = q_start // 2 + packed_k_offsets + packed_k_cols = k_group_offsets * (FP4_BLOCK // 2) + byte_offsets_in_group + mask_k = q_elem_offsets < Q_HEAD_D + safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) + safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) + q_vals = tl.load( + q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + safe_packed_q_cols[None, :] * q_fp4_s1, + mask=mask_h[:, None] & mask_k[None, :], + other=0, + ) + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + safe_packed_k_cols[None, :] * kv_s4, + mask=valid_physical_page & mask_k[None, :], + other=0, + ) + + q_sf_cols = q_start // FP4_BLOCK + scale_offsets + k_sf_cols = tl.where( + q_sf_cols < non_residual_groups, + q_sf_cols, + non_residual_groups + (q_sf_cols - non_residual_groups) // 2, + ) + mask_sf = q_sf_cols < Q_SF_PER_TOKEN + safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) + safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) + q_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN + ) + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN + ) + q_scales = tl.load(q_sf_ptr + q_sf_offsets) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) + scores = tl.dot_scaled( + q_vals, + q_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + acc=scores, + fast_math=True, + ) + + return scores + + +@triton.jit +def _fp4_mla_attention_stats_kernel( + max_ptr, + denom_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + stats_s0, + sm_scale: tl.constexpr, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, BLOCK_T) + q_row_base = gen_idx * NUM_HEADS + kv_len = tl.load(kv_lens_ptr + gen_idx) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + + max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + denom = tl.zeros((BLOCK_H,), dtype=tl.float32) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + for page_rel in tl.range(0, MAX_PAGES): + page_start = page_rel * PAGE_SIZE + if page_start < kv_len: + compact_page = page_table_start + page_rel + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + offs_h, + offs_t, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + BLOCK_T, + BLOCK_K, + NUM_HEADS, + ) + valid_t = page_start + offs_t < kv_len + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + page_max = tl.max(scores, axis=1) + new_max = tl.maximum(max_score, page_max) + denom = denom * tl.exp(max_score - new_max) + tl.sum( + tl.exp(scores - new_max[:, None]), axis=1 + ) + max_score = new_max + + tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) + tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) + + +@triton.jit +def _fp4_mla_attention_prob_store_page_kernel( + probs_ptr, + max_ptr, + denom_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_rel, + page_ids_len, + num_pages, + probs_s0, + probs_s1, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + stats_s0, + sm_scale: tl.constexpr, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_K: tl.constexpr, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + + kv_len = tl.load(kv_lens_ptr + gen_idx) + page_start = page_rel * PAGE_SIZE + if page_start >= kv_len: + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, PAGE_SIZE) + valid_t = page_start + offs_t < kv_len + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if (compact_page < 0) | (compact_page >= page_ids_len): + return + q_row_base = gen_idx * NUM_HEADS + + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + offs_h, + offs_t, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + PAGE_SIZE, + BLOCK_K, + NUM_HEADS, + ) + max_score = tl.load(max_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + denom_valid = denom > 0.0 + safe_denom = tl.where(denom_valid, denom, 1.0) + safe_max = tl.where(denom_valid, max_score, 0.0) + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + probs = tl.exp(scores - safe_max[:, None]) / safe_denom[:, None] + probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) + + prob_rows = gen_idx * NUM_HEADS + offs_h + safe_prob_rows = tl.where(mask_h, prob_rows, gen_idx * NUM_HEADS) + tl.store( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + offs_t[None, :] * probs_s1, + probs, + mask=mask_h[:, None], + ) + + +@triton.jit +def _fp4_mla_attention_prob_pack_page_kernel( + p_fp4_ptr, + p_sf_ptr, + probs_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_rel, + page_ids_len, + p_s0, + p_s1, + probs_s0, + probs_s1, + NUM_HEADS: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, +): + gen_idx = tl.program_id(0) + token_group = tl.program_id(1) + head_block = tl.program_id(2) + + kv_len = tl.load(kv_lens_ptr + gen_idx) + page_start = page_rel * PAGE_SIZE + if page_start >= kv_len: + return + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if (compact_page < 0) | (compact_page >= page_ids_len): + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + mask_h = offs_h < NUM_HEADS + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + token_base = token_group * FP4_BLOCK + even_t = token_base + byte_offsets * 2 + odd_t = even_t + 1 + valid_even = page_start + even_t < kv_len + valid_odd = page_start + odd_t < kv_len + + prob_rows = gen_idx * NUM_HEADS + offs_h + safe_prob_rows = tl.where(mask_h, prob_rows, gen_idx * NUM_HEADS) + even_probs = tl.load( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + even_t[None, :] * probs_s1, + mask=mask_h[:, None] & valid_even[None, :], + other=0.0, + ) + odd_probs = tl.load( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + odd_t[None, :] * probs_s1, + mask=mask_h[:, None] & valid_odd[None, :], + other=0.0, + ) + amax = tl.maximum(tl.max(tl.abs(even_probs), axis=1), tl.max(tl.abs(odd_probs), axis=1)) + local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) + + p_rows = compact_page * NUM_HEADS + offs_h + safe_p_rows = tl.where(mask_h, p_rows, compact_page * NUM_HEADS) + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows, token_group, SF_PER_PAGE) + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h) + + even_quant = _fp4_e2m1_quantize(even_probs / local_scale[:, None]) + odd_quant = _fp4_e2m1_quantize(odd_probs / local_scale[:, None]) + packed = even_quant | (odd_quant << 4) + byte_cols = token_group * (FP4_BLOCK // 2) + byte_offsets + tl.store( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + byte_cols[None, :] * p_s1, + packed, + mask=mask_h[:, None], + ) + + +@triton.jit +def _fp4_mla_attention_pv_kernel( + out_ptr, + p_fp4_ptr, + p_sf_ptr, + kv_cache_ptr, + v_sf_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + out_s0, + out_s1, + out_s2, + p_s0, + p_s1, + kv_s0, + kv_s2, + kv_s4, + vsf_s0, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + dim_block = tl.program_id(2) + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + mask_h = offs_h < NUM_HEADS + mask_v = offs_v < V_HEAD_D + safe_offs_h = tl.where(mask_h, offs_h, 0) + safe_offs_v = tl.where(mask_v, offs_v, 0) + packed_t = tl.arange(0, PAGE_SIZE // 2) + scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) + even_t = packed_t * 2 + odd_t = even_t + 1 + v_packed_offsets = safe_offs_v // 2 + v_use_high_nibble = (safe_offs_v & 1) != 0 + + kv_len = tl.load(kv_lens_ptr + gen_idx) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES): + page_start = page_rel * PAGE_SIZE + if page_start < kv_len: + compact_page = page_table_start + page_rel + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + ) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + + p_rows = safe_compact_page * NUM_HEADS + offs_h + safe_p_rows = tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) + p_vals = tl.load( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, + mask=valid_compact_page & mask_h[:, None], + other=0, + ) + p_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + p_scales = tl.load(p_sf_ptr + p_sf_offsets) + + valid_even_t = page_start + even_t < kv_len + valid_odd_t = page_start + odd_t < kv_len + even_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + even_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), + other=0, + ) + odd_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + odd_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), + other=0, + ) + even_low = even_packed & 0x0F + even_high = (even_packed >> 4) & 0x0F + odd_low = odd_packed & 0x0F + odd_high = (odd_packed >> 4) & 0x0F + even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) + odd_nibble = tl.where(v_use_high_nibble[:, None], odd_high, odd_low) + v_vals = even_nibble | (odd_nibble << 4) + v_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) + acc = tl.dot_scaled( + p_vals, + p_scales, + "e2m1", + v_vals.T, + v_scales, + "e2m1", + acc=acc, + fast_math=True, + ) + + tl.store( + out_ptr + gen_idx * out_s0 + safe_offs_h[:, None] * out_s1 + safe_offs_v[None, :] * out_s2, + acc / (global_scale * P_GLOBAL_SCALE), + mask=mask_h[:, None] & mask_v[None, :], + ) + + +@triton.jit +def _fp4_mla_overlay_hp_tail_kernel( + out_ptr, + pool_ptr, + seq_slots_ptr, + kv_lens_ptr, + paged_kv_indptr_decode_ptr, + num_seq_slots, + num_layers, + num_pages, + layer_idx, + page_size, + out_s0, + out_s1, + out_s2, + pool_stride_seq, + pool_stride_layer, + D: tl.constexpr, + BLOCK_D: tl.constexpr, + HP_BLOCK: tl.constexpr, +): + gen_idx = tl.program_id(0) + tail_idx = tl.program_id(1) + if (layer_idx < 0) | (layer_idx >= num_layers): + return + + kv_len = tl.load(kv_lens_ptr + gen_idx) + tail_count = kv_len % HP_BLOCK + if tail_idx >= tail_count: + return + + abs_pos = (kv_len - tail_count + tail_idx).to(tl.int64) + page_size_i64 = tl.cast(page_size, tl.int64) + rel_page = abs_pos // page_size_i64 + page_pos = abs_pos - rel_page * page_size_i64 + compact_page = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + rel_page + if (compact_page < 0) | (compact_page >= num_pages): + return + hp_slot = abs_pos % HP_BLOCK + seq_slot = tl.load(seq_slots_ptr + gen_idx).to(tl.int64) + if (seq_slot < 0) | (seq_slot >= num_seq_slots): + return + + offs_d = tl.arange(0, BLOCK_D) + mask_d = offs_d < D + safe_offs_d = tl.where(mask_d, offs_d, 0) + src_base = seq_slot * pool_stride_seq + layer_idx * pool_stride_layer + hp_slot * D + value = tl.load(pool_ptr + src_base + safe_offs_d, mask=mask_d, other=0.0) + tl.store( + out_ptr + compact_page * out_s0 + page_pos * out_s1 + safe_offs_d * out_s2, + value, + mask=mask_d, + ) diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py new file mode 100644 index 000000000000..00664c41ed18 --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py @@ -0,0 +1,1338 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Shared MLA FP4 KV-cache helpers. + +The high-precision (HP) BF16 KV pool is a standalone circular buffer used +alongside the paged FP4 KV pool when MLA models run with NVFP4 KV cache. +Each sequence slot stores the ``HP_BLOCK_SIZE`` most-recent latent vectors at +BF16, so attention backends can consult BF16 values for the tail tokens that +do not yet fill a complete FP4 quant block of 16 elements along the sequence +dimension. + +Used by both ``TrtllmAttention`` (via an internal C++ attention op that reads +both pools) and ``FlashInferAttention`` (via either explicit Python-side +dequant into a BF16 workspace before calling FlashInfer MLA wrappers, or an +env-gated Triton attention path that reads packed FP4 Q, K, and V directly). +""" + +import os +from typing import Any, Literal, Optional + +import torch +import triton + +from .fp4_mla_kernels import ( + _fp4_mla_attention_prob_pack_page_kernel, + _fp4_mla_attention_prob_store_page_kernel, + _fp4_mla_attention_pv_kernel, + _fp4_mla_attention_stats_kernel, + _fp4_mla_dequant_kernel, + _fp4_mla_overlay_hp_tail_kernel, + _fp4_mla_scatter_kernel, + _fp4_mla_v_scale_store_context_tokens_kernel, + _fp4_mla_v_scale_store_hp_tail_kernel, + _hp_kv_store_context_kernel, + _hp_kv_store_gen_kernel, +) + +HP_BLOCK_SIZE: int = 16 +FP4_BLOCK_SIZE: int = 16 +FP4_MLA_TOKENS_PER_BLOCK: int = 128 +FP4_MLA_SCALE_ROW_GROUP: int = 128 +FP4_MLA_SCALE_COL_GROUP: int = 4 +FP4_MLA_KV_GLOBAL_SCALE: float = 448.0 * 6.0 / 448.0 * 6.0 +FP4_MLA_P_GLOBAL_SCALE: float = 448.0 * 6.0 +FP4_MLA_Q_RESIDUAL_DIM: int = 64 +FLASHINFER_FP4_MLA_ATTENTION_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION" +FLASHINFER_FP4_MLA_DEBUG_ENV = "TRTLLM_FLASHINFER_FP4_MLA_DEBUG" +_HPUpdatePhase = Literal["all", "context", "generation"] + + +# Environment and debug helpers + + +def _env_enabled(name: str) -> bool: + return os.getenv(name, "0").lower() in ( + "1", + "true", + "yes", + "on", + ) + + +def is_flashinfer_fp4_mla_attention_enabled() -> bool: + """Return whether FlashInfer MLA should allocate no-dequant FP4 attention buffers.""" + return _env_enabled(FLASHINFER_FP4_MLA_ATTENTION_ENV) + + +def _fp4_mla_debug_enabled() -> bool: + return _env_enabled(FLASHINFER_FP4_MLA_DEBUG_ENV) + + +def _fp4_mla_debug(message: str) -> None: + if _fp4_mla_debug_enabled(): + print(f"[fp4_mla_debug] {message}", flush=True) + + +def _tensor_layout(tensor: Optional[torch.Tensor]) -> str: + if tensor is None: + return "None" + return ( + f"shape={list(tensor.shape)} stride={list(tensor.stride())} " + f"dtype={tensor.dtype} device={tensor.device}" + ) + + +def _debug_tensor_range(name: str, tensor: Optional[torch.Tensor]) -> None: + if not _fp4_mla_debug_enabled(): + return + if tensor is None: + _fp4_mla_debug(f"{name}: None") + return + flat = tensor.detach().reshape(-1) + if flat.numel() == 0: + _fp4_mla_debug(f"{name}: empty {_tensor_layout(tensor)}") + return + try: + first = flat[: min(8, flat.numel())].cpu().tolist() + _fp4_mla_debug( + f"{name}: {_tensor_layout(tensor)} n={flat.numel()} " + f"min={flat.min().item()} max={flat.max().item()} first={first}" + ) + except RuntimeError as exc: + _fp4_mla_debug(f"{name}: failed to read range: {exc}") + + +def _debug_sync(label: str) -> None: + if not _fp4_mla_debug_enabled(): + return + if torch.cuda.is_current_stream_capturing(): + _fp4_mla_debug(f"{label}: skip sync during CUDA graph capture") + return + _fp4_mla_debug(f"{label}: synchronize") + torch.cuda.synchronize() + _fp4_mla_debug(f"{label}: sync complete") + + +def _ceil_div(lhs: int, rhs: int) -> int: + return (lhs + rhs - 1) // rhs + + +# FP4 MLA scale-layout helpers + + +def get_fp4_mla_v_scale_pool_size(v_head_dim: int, page_size: int) -> int: + """Return elements per page for the swizzled FP4 MLA V-scale pool. + + The PV matmul treats V as a RHS matrix shaped ``[v_head_dim, kv_tokens]``. + NVFP4 block scales therefore group along the token/K axis, not along the + latent dimension as the K-view cache does. The physical layout matches the + Triton block-scaled matmul scale layout: + ``[ceil(v_head_dim / 128), ceil(page_size / 16 / 4), 32, 16]``. + """ + token_scale_cols = _ceil_div(page_size, FP4_BLOCK_SIZE) + row_groups = _ceil_div(v_head_dim, FP4_MLA_SCALE_ROW_GROUP) + col_groups = _ceil_div(token_scale_cols, FP4_MLA_SCALE_COL_GROUP) + return row_groups * col_groups * 32 * 16 + + +def _get_fp4_mla_swizzled_scale_size(rows: int, cols: int) -> int: + scale_cols = _ceil_div(cols, FP4_BLOCK_SIZE) + row_groups = _ceil_div(rows, FP4_MLA_SCALE_ROW_GROUP) + col_groups = _ceil_div(scale_cols, FP4_MLA_SCALE_COL_GROUP) + return row_groups * col_groups * 32 * 16 + + +def _use_fp4_mla_swizzled_sf() -> bool: + return is_flashinfer_fp4_mla_attention_enabled() + + +def _get_fp4_mla_context_start_positions(metadata: Any, num_contexts: int) -> torch.Tensor: + kv_cache_params = getattr(metadata, "kv_cache_params", None) + cached_token_lens = getattr(kv_cache_params, "num_cached_tokens_per_seq", None) + if cached_token_lens is not None: + return torch.as_tensor(cached_token_lens[:num_contexts], dtype=torch.int64, device="cpu") + + return ( + ( + metadata.kv_lens_cuda_runtime[:num_contexts] + - metadata.prompt_lens_cuda_runtime[:num_contexts] + ) + .detach() + .cpu() + ) + + +def _validate_fp4_mla_context_start_alignment(metadata: Any, num_contexts: int) -> None: + context_start_positions = _get_fp4_mla_context_start_positions(metadata, num_contexts) + bad_start = (context_start_positions < 0) | ((context_start_positions % HP_BLOCK_SIZE) != 0) + if bool(torch.any(bad_start).item()): + starts = context_start_positions.detach().cpu().tolist() + raise ValueError( + "FP4 MLA shared-tile context update requires every context " + f"start position to be {HP_BLOCK_SIZE}-token aligned, got " + f"start positions {starts}." + ) + + +def get_fp4_mla_v_scale_pool_shape( + num_layers: int, + num_pages: int, + v_head_dim: int, + page_size: int, +) -> tuple[int, int, int, int, int, int]: + """Return the logical swizzled V-scale view shape. + + The leading dimensions are ``[layer, physical_page]``. The remaining + dimensions are the preshuffled ``[N // 128, K // 16 // 4, 32, 16]`` shape + consumed by Triton block-scaled matmul for the V/PV RHS operand. + """ + token_scale_cols = _ceil_div(page_size, FP4_BLOCK_SIZE) + return ( + num_layers, + num_pages, + _ceil_div(v_head_dim, FP4_MLA_SCALE_ROW_GROUP), + _ceil_div(token_scale_cols, FP4_MLA_SCALE_COL_GROUP), + 32, + 16, + ) + + +def get_fp4_mla_v_scale_pool_view( + metadata: Any, + *, + v_head_dim: int, +) -> torch.Tensor: + """View the auxiliary MLA V-scale pool in Triton's block-scaled layout.""" + pool = getattr(metadata, "fp4_mla_v_scale_pool", None) + if pool is None: + raise RuntimeError("FP4 MLA V scale pool is not allocated.") + + elems_per_page = get_fp4_mla_v_scale_pool_size(v_head_dim, metadata.page_size) + if pool.shape[-1] < elems_per_page: + raise RuntimeError( + f"FP4 MLA V scale pool page stride is too small: got " + f"{pool.shape[-1]}, need {elems_per_page}." + ) + + token_scale_cols = _ceil_div(metadata.page_size, FP4_BLOCK_SIZE) + col_groups = _ceil_div(token_scale_cols, FP4_MLA_SCALE_COL_GROUP) + shape = get_fp4_mla_v_scale_pool_shape( + pool.shape[0], pool.shape[1], v_head_dim, metadata.page_size + ) + strides = ( + pool.stride(0), + pool.stride(1), + col_groups * 32 * 16, + 32 * 16, + 16, + 1, + ) + return torch.as_strided(pool, size=shape, stride=strides) + + +# Python launch helpers + + +def _get_fp4_mla_global_scale(metadata: Any, device: torch.device) -> torch.Tensor: + global_scale = getattr(metadata, "_fp4_mla_global_scale", None) + if global_scale is None: + global_scale = torch.ones((1,), dtype=torch.float32, device=device) + return global_scale + + +def _get_fp4_mla_cache_tensors(metadata: Any, layer_idx: int) -> tuple[torch.Tensor, torch.Tensor]: + kv_cache = metadata.kv_cache_manager.get_buffers(layer_idx).view(torch.uint8) + sf_cache = metadata.kv_cache_manager.get_block_scale_buffers(layer_idx) + if sf_cache is None: + raise RuntimeError("NVFP4 KV cache scale pool is not available.") + return kv_cache, sf_cache + + +def _scatter_fp4_mla_kv_cache_2d_context( + metadata: Any, + latent_cache: torch.Tensor, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + v_sf: torch.Tensor, + global_scale: torch.Tensor, + *, + token_offset: int, + local_layer: int, + v_head_dim: int, + head_dim: int, + num_tokens: int, + num_dim_blocks: int, + sf_per_token: int, + sf_per_page: int, +) -> None: + num_contexts = metadata.num_contexts + if num_contexts > 0: + prompt_lens_cpu = metadata.prompt_lens_cpu_runtime[:num_contexts] + ctx_token_count = int(prompt_lens_cpu.sum().item()) + if num_tokens != ctx_token_count: + raise RuntimeError( + f"FP4 MLA 2D context scatter needs {ctx_token_count} context tokens, got " + f"{num_tokens}." + ) + _validate_fp4_mla_context_start_alignment(metadata, num_contexts) + + _fp4_mla_v_scale_store_context_tokens_kernel[ + ( + num_tokens, + num_dim_blocks, + ) + ]( + kv_cache, + sf_cache, + v_sf, + latent_cache, + global_scale, + metadata.batch_indices, + metadata.positions, + metadata.paged_kv_indices, + metadata.paged_kv_indptr, + metadata.paged_kv_indices.shape[0], + metadata.paged_kv_indptr.shape[0], + metadata.batch_indices.shape[0], + v_sf.shape[1], + v_sf.shape[0], + token_offset, + num_tokens, + local_layer, + metadata.page_size, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + sf_cache.stride(0), + latent_cache.stride(0), + latent_cache.stride(1), + v_sf.stride(0), + v_sf.stride(1), + HEAD_D=head_dim, + V_HEAD_D=v_head_dim, + HP_BLOCK=HP_BLOCK_SIZE, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_TOKEN=sf_per_token, + SF_PER_PAGE=sf_per_page, + ) + _debug_sync("scatter_fp4_mla_kv_cache_2d_context") + + +def _scatter_fp4_mla_kv_cache_2d_generation( + metadata: Any, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + v_sf: torch.Tensor, + global_scale: torch.Tensor, + *, + local_layer: int, + v_head_dim: int, + head_dim: int, + num_tokens: int, + num_dim_blocks: int, + sf_per_token: int, + sf_per_page: int, +) -> None: + num_contexts = metadata.num_contexts + num_seqs = metadata.num_seqs + num_gen = num_seqs - num_contexts + if num_gen <= 0: + return + if num_tokens != num_gen: + raise RuntimeError( + f"FP4 MLA 2D generation scatter expected {num_gen} generation tokens, got {num_tokens}." + ) + + pool = getattr(metadata, "high_precision_kv_pool", None) + if pool is None: + raise RuntimeError("FP4 MLA 2D generation scatter requires the HP KV pool.") + hp_head_dim = pool.shape[-1] // HP_BLOCK_SIZE + if hp_head_dim < head_dim: + raise RuntimeError( + f"FP4 MLA 2D generation scatter needs at least {head_dim} HP channels, got " + f"{hp_head_dim}." + ) + + page_ids = metadata.paged_kv_indices[metadata.num_context_blocks :] + _fp4_mla_v_scale_store_hp_tail_kernel[(num_gen, num_dim_blocks)]( + kv_cache, + sf_cache, + v_sf, + pool, + global_scale, + metadata.seq_slots[num_contexts:num_seqs], + metadata.kv_lens_cuda_runtime[num_contexts:num_seqs], + page_ids, + metadata.paged_kv_indptr_decode, + page_ids.shape[0], + metadata.paged_kv_indptr_decode.shape[0], + v_sf.shape[1], + pool.shape[0], + v_sf.shape[0], + local_layer, + metadata.page_size, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + sf_cache.stride(0), + pool.stride(0), + pool.stride(1), + v_sf.stride(0), + v_sf.stride(1), + HEAD_D=hp_head_dim, + V_HEAD_D=v_head_dim, + HP_BLOCK=HP_BLOCK_SIZE, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_TOKEN=sf_per_token, + SF_PER_PAGE=sf_per_page, + ) + _debug_sync("scatter_fp4_mla_kv_cache_2d_generation") + + +def _scatter_fp4_mla_kv_cache_1d( + metadata: Any, + latent_cache: torch.Tensor, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + global_scale: torch.Tensor, + *, + layer_idx: int, + token_offset: int, + num_tokens: int, + head_dim: int, + sf_per_token: int, + use_swizzled_sf: bool, +) -> None: + q_fp4, q_sf = torch.ops.trtllm.fp4_quantize( + latent_cache, global_scale, FP4_BLOCK_SIZE, False, False + ) + q_sf = q_sf.view(num_tokens, head_dim // FP4_BLOCK_SIZE) + + packed_dim = head_dim // 2 + block_packed_dim = triton.next_power_of_2(packed_dim) + block_sf = triton.next_power_of_2(sf_per_token) + + _fp4_mla_debug( + "scatter launch: " + f"num_tokens={num_tokens} token_offset={token_offset} " + f"page_size={metadata.page_size} layer_idx={layer_idx} " + f"head_dim={head_dim} packed_dim={packed_dim} " + f"sf_per_token={sf_per_token} use_swizzled_sf={use_swizzled_sf}" + ) + _fp4_mla_debug(f"scatter latent_cache: {_tensor_layout(latent_cache)}") + _fp4_mla_debug(f"scatter kv_cache: {_tensor_layout(kv_cache)}") + _fp4_mla_debug(f"scatter sf_cache: {_tensor_layout(sf_cache)}") + _debug_tensor_range( + "scatter batch_indices", + metadata.batch_indices[token_offset : token_offset + num_tokens], + ) + _debug_tensor_range( + "scatter positions", + metadata.positions[token_offset : token_offset + num_tokens], + ) + _debug_tensor_range("scatter paged_kv_indices", metadata.paged_kv_indices) + _debug_tensor_range("scatter paged_kv_indptr", metadata.paged_kv_indptr) + + _fp4_mla_scatter_kernel[(num_tokens,)]( + kv_cache, + sf_cache, + q_fp4, + q_sf, + metadata.batch_indices, + metadata.positions, + metadata.paged_kv_indices, + metadata.paged_kv_indptr, + metadata.paged_kv_indices.shape[0], + metadata.paged_kv_indptr.shape[0], + kv_cache.shape[0], + token_offset, + metadata.page_size, + kv_cache.stride(0), + kv_cache.stride(1), + kv_cache.stride(2), + kv_cache.stride(3), + kv_cache.stride(4), + sf_cache.stride(0), + sf_cache.stride(1), + sf_cache.stride(2), + sf_cache.stride(3), + sf_cache.stride(4), + q_fp4.stride(0), + q_fp4.stride(1), + q_sf.stride(0), + q_sf.stride(1), + PACKED_D=packed_dim, + SF_PER_TOKEN=sf_per_token, + BLOCK_PACKED_D=block_packed_dim, + BLOCK_SF=block_sf, + USE_SWIZZLED_SF=use_swizzled_sf, + ) + _debug_sync("scatter_fp4_mla_kv_cache") + + +# Public cache update and decode entry points + + +def scatter_fp4_mla_kv_cache( + metadata: Any, + latent_cache: Optional[torch.Tensor], + layer_idx: int, + *, + token_offset: int, + phase: Optional[_HPUpdatePhase] = None, + local_layer: Optional[int] = None, + v_head_dim: Optional[int] = None, +) -> None: + """Quantize MLA latent tokens and scatter them into the paged FP4 cache. + + Contract: this helper scatters exactly ``latent_cache.shape[0]`` tokens, + reading index metadata at ``batch_indices[token_offset : token_offset + N]`` + and ``positions[token_offset : token_offset + N]``. Callers must pass a + latent_cache pre-sliced to the current phase (context or generation) so + that ``shape[0]`` matches the number of index entries they intend to + consume. ``MLA.forward_impl`` (tensorrt_llm/_torch/modules/attention.py) + slices ``latent_cache[:num_ctx_tokens]`` for context and + ``latent_cache[num_ctx_tokens:]`` for generation before dispatching. + + When the no-dequant FP4 MLA attention path is enabled, callers should pass + ``phase``, ``local_layer``, and ``v_head_dim``. Context scatter then writes + the final FP4 tile representation directly: dimensions below + ``v_head_dim`` use one shared 16-token by 16-dim scale written into both + K's token-major scale layout and V's dim-major scale layout. Tail K-only + dimensions use K's per-token 1D scales. Generation scatter rewrites the + active 16-token tile from the HP pool, so the caller must update the HP pool + before invoking this helper. + """ + if latent_cache is None or latent_cache.numel() == 0: + return + + latent_cache = latent_cache.reshape(latent_cache.shape[0], -1).contiguous() + num_tokens = latent_cache.shape[0] + head_dim = latent_cache.shape[-1] + if head_dim % FP4_BLOCK_SIZE != 0: + raise ValueError( + f"FP4 MLA KV head_dim must be divisible by {FP4_BLOCK_SIZE}, got {head_dim}." + ) + indices_len = metadata.batch_indices.shape[0] + positions_len = metadata.positions.shape[0] + if token_offset + num_tokens > indices_len or token_offset + num_tokens > positions_len: + raise RuntimeError( + f"FP4 MLA scatter would read batch_indices[{token_offset}:" + f"{token_offset + num_tokens}] / positions[{token_offset}:" + f"{token_offset + num_tokens}], but only {indices_len} / " + f"{positions_len} entries are available. This indicates " + "latent_cache was not pre-sliced to the current phase's token " + "range (see MLA.forward_impl)." + ) + + use_swizzled_sf = _use_fp4_mla_swizzled_sf() + if use_swizzled_sf: + _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) + + global_scale = _get_fp4_mla_global_scale(metadata, latent_cache.device) + kv_cache, sf_cache = _get_fp4_mla_cache_tensors(metadata, layer_idx) + sf_per_token = head_dim // FP4_BLOCK_SIZE + + use_2d_scatter = ( + use_swizzled_sf + and phase in ("context", "generation") + and getattr(metadata, "fp4_mla_v_scale_pool", None) is not None + ) + if use_2d_scatter: + assert phase is not None + if local_layer is None or v_head_dim is None: + raise ValueError("Real FP4 MLA scatter requires local_layer and v_head_dim.") + if metadata.page_size % HP_BLOCK_SIZE != 0: + raise ValueError( + f"FP4 MLA scatter requires page_size divisible by " + f"{HP_BLOCK_SIZE}, got {metadata.page_size}." + ) + if v_head_dim > head_dim: + raise ValueError(f"FP4 MLA v_head_dim={v_head_dim} cannot exceed head_dim={head_dim}.") + if v_head_dim % FP4_BLOCK_SIZE != 0: + raise ValueError( + f"FP4 MLA v_head_dim must be divisible by {FP4_BLOCK_SIZE}, got {v_head_dim}." + ) + + sf_cache = sf_cache.view(torch.float8_e4m3fn) + v_sf = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=v_head_dim) + num_dim_blocks = triton.cdiv(head_dim, FP4_BLOCK_SIZE) + sf_per_page = metadata.page_size // HP_BLOCK_SIZE + _fp4_mla_debug( + "scatter 2d launch: " + f"phase={phase} num_tokens={num_tokens} " + f"token_offset={token_offset} layer_idx={layer_idx} " + f"local_layer={local_layer} head_dim={head_dim} " + f"v_head_dim={v_head_dim} num_dim_blocks={num_dim_blocks}" + ) + _fp4_mla_debug(f"scatter 2d kv_cache: {_tensor_layout(kv_cache)}") + _fp4_mla_debug(f"scatter 2d sf_cache: {_tensor_layout(sf_cache)}") + _fp4_mla_debug(f"scatter 2d v_sf: {_tensor_layout(v_sf)}") + + if phase == "context": + _scatter_fp4_mla_kv_cache_2d_context( + metadata, + latent_cache, + kv_cache, + sf_cache, + v_sf, + global_scale, + token_offset=token_offset, + local_layer=local_layer, + v_head_dim=v_head_dim, + head_dim=head_dim, + num_tokens=num_tokens, + num_dim_blocks=num_dim_blocks, + sf_per_token=sf_per_token, + sf_per_page=sf_per_page, + ) + else: + _scatter_fp4_mla_kv_cache_2d_generation( + metadata, + kv_cache, + sf_cache, + v_sf, + global_scale, + local_layer=local_layer, + v_head_dim=v_head_dim, + head_dim=head_dim, + num_tokens=num_tokens, + num_dim_blocks=num_dim_blocks, + sf_per_token=sf_per_token, + sf_per_page=sf_per_page, + ) + return + + _scatter_fp4_mla_kv_cache_1d( + metadata, + latent_cache, + kv_cache, + sf_cache, + global_scale, + layer_idx=layer_idx, + token_offset=token_offset, + num_tokens=num_tokens, + head_dim=head_dim, + sf_per_token=sf_per_token, + use_swizzled_sf=use_swizzled_sf, + ) + + +def _ensure_decode_workspace( + metadata: Any, + head_dim: int, + dtype: torch.dtype, +) -> torch.Tensor: + num_blocks = _get_decode_workspace_num_blocks(metadata) + workspace = getattr(metadata, "_fp4_mla_decode_cache_buf", None) + needs_alloc = ( + workspace is None + or workspace.shape[0] < max(num_blocks, 1) + or workspace.shape[1] != metadata.page_size + or workspace.shape[2] != head_dim + or workspace.dtype != dtype + ) + if needs_alloc: + if torch.cuda.is_current_stream_capturing(): + raise ValueError( + "Cannot allocate FlashInfer FP4 MLA decode workspace while " + "capturing a CUDA graph. Run a warmup prepare/forward first." + ) + workspace = torch.empty( + (max(num_blocks, 1), metadata.page_size, head_dim), + dtype=dtype, + device=metadata.paged_kv_indices.device, + ) + metadata._fp4_mla_decode_cache_buf = workspace + return workspace[:num_blocks] + + +def _get_decode_workspace_num_blocks(metadata: Any) -> int: + if metadata.is_cuda_graph: + max_blocks_per_seq = ( + metadata.kv_cache_manager.max_seq_len + metadata.page_size - 1 + ) // metadata.page_size + max_graph_blocks = metadata.max_num_requests * max_blocks_per_seq + return min( + metadata.kv_cache_manager.blocks_in_primary_pool, + max_graph_blocks, + ) + return metadata.num_generation_blocks + + +def _get_decode_src_page_ids(metadata: Any, num_blocks: int) -> torch.Tensor: + page_ids = ( + metadata._paged_kv_indices + if metadata.is_cuda_graph and hasattr(metadata, "_paged_kv_indices") + else metadata.paged_kv_indices + ) + src_page_ids = page_ids[metadata.num_context_blocks : metadata.num_context_blocks + num_blocks] + if src_page_ids.numel() != num_blocks: + raise RuntimeError( + f"FP4 MLA dequant needs {num_blocks} decode page ids from " + f"paged_kv_indices[{metadata.num_context_blocks}:" + f"{metadata.num_context_blocks + num_blocks}], got " + f"{src_page_ids.numel()}." + ) + return src_page_ids + + +def _validate_fp4_mla_cache_shape(page_size: int, head_dim: int) -> None: + if page_size != FP4_MLA_TOKENS_PER_BLOCK: + raise ValueError( + f"FP4 MLA KV cache requires tokens_per_block={FP4_MLA_TOKENS_PER_BLOCK} " + f"for swizzled block scales, got {page_size}." + ) + + sf_per_token = head_dim // FP4_BLOCK_SIZE + if head_dim % FP4_BLOCK_SIZE != 0 or sf_per_token % 4 != 0: + raise ValueError( + f"FP4 MLA KV head_dim must produce a scale column count divisible by 4; " + f"got head_dim={head_dim}, scale_columns={sf_per_token}." + ) + + +def _validate_fp4_mla_attention_q_shape(head_dim: int, q_residual_dim: int) -> None: + if q_residual_dim % FP4_BLOCK_SIZE != 0: + raise ValueError( + f"FP4 MLA Q residual_dim must be divisible by {FP4_BLOCK_SIZE}, got {q_residual_dim}." + ) + if q_residual_dim <= 0 or q_residual_dim > head_dim: + raise ValueError( + f"FP4 MLA Q residual_dim must be in (0, head_dim], got " + f"residual_dim={q_residual_dim}, head_dim={head_dim}." + ) + + q_head_dim = head_dim + q_residual_dim + q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE + if q_head_dim % FP4_BLOCK_SIZE != 0 or q_sf_per_token % FP4_MLA_SCALE_COL_GROUP != 0: + raise ValueError( + f"FP4 MLA residual Q must produce a scale column count divisible " + f"by {FP4_MLA_SCALE_COL_GROUP}; got q_head_dim={q_head_dim}, " + f"scale_columns={q_sf_per_token}." + ) + + +def get_fp4_mla_decode_cache( + metadata: Any, + layer_idx: int, + local_layer: int, + *, + head_dim: int, + dtype: torch.dtype, +) -> torch.Tensor: + """Build a compact dequantized MLA cache for FlashInfer decode.""" + # Must match scatter_fp4_mla_kv_cache: the env var picks the SF layout, + # not the page size. When the dequant fallback is the read path + # (env disabled), scatter wrote linear scales and we must read linear. + use_swizzled_sf = _use_fp4_mla_swizzled_sf() + if use_swizzled_sf: + _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) + combined = _ensure_decode_workspace(metadata, head_dim, dtype) + num_blocks = combined.shape[0] + if num_blocks == 0: + return combined + + kv_cache, sf_cache = _get_fp4_mla_cache_tensors(metadata, layer_idx) + sf_cache = sf_cache.view(torch.float8_e4m3fn) + global_scale = _get_fp4_mla_global_scale(metadata, combined.device) + src_page_ids = _get_decode_src_page_ids(metadata, num_blocks) + block_d = triton.next_power_of_2(head_dim) + + _fp4_mla_dequant_kernel[(num_blocks, metadata.page_size)]( + combined, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + src_page_ids.shape[0], + kv_cache.shape[0], + kv_cache.stride(0), + kv_cache.stride(1), + kv_cache.stride(2), + kv_cache.stride(3), + kv_cache.stride(4), + sf_cache.stride(0), + sf_cache.stride(1), + sf_cache.stride(2), + sf_cache.stride(3), + sf_cache.stride(4), + combined.stride(0), + combined.stride(1), + combined.stride(2), + D=head_dim, + FP4_BLOCK=FP4_BLOCK_SIZE, + BLOCK_D=block_d, + USE_SWIZZLED_SF=use_swizzled_sf, + ) + + num_gen = metadata.num_seqs - metadata.num_contexts + if num_gen > 0 and metadata.high_precision_kv_pool is not None: + pool = metadata.high_precision_kv_pool + _fp4_mla_overlay_hp_tail_kernel[(num_gen, HP_BLOCK_SIZE)]( + combined, + pool, + metadata.seq_slots[metadata.num_contexts : metadata.num_seqs], + metadata.kv_lens_cuda_runtime[metadata.num_contexts : metadata.num_seqs], + metadata.paged_kv_indptr_decode, + pool.shape[0], + pool.shape[1], + combined.shape[0], + local_layer, + metadata.page_size, + combined.stride(0), + combined.stride(1), + combined.stride(2), + pool.stride(0), + pool.stride(1), + D=head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + + return combined + + +def _ensure_workspace_tensor( + metadata: Any, + attr_name: str, + shape: tuple[int, ...], + *, + dtype: torch.dtype, + device: torch.device, +) -> torch.Tensor: + tensor = getattr(metadata, attr_name, None) + needs_alloc = ( + tensor is None + or tensor.dtype != dtype + or tensor.device != device + or len(tensor.shape) != len(shape) + or any(tensor.shape[idx] < dim for idx, dim in enumerate(shape)) + ) + if needs_alloc: + if torch.cuda.is_current_stream_capturing(): + raise ValueError( + f"Cannot allocate {attr_name} while capturing a CUDA graph. " + "Run a warmup prepare/forward first." + ) + tensor = torch.empty(shape, dtype=dtype, device=device) + setattr(metadata, attr_name, tensor) + + slices = tuple(slice(0, dim) for dim in shape) + return tensor[slices] + + +def _max_generation_pages(metadata: Any) -> int: + num_gen = metadata.num_seqs - metadata.num_contexts + if num_gen <= 0: + return 0 + num_blocks = getattr(metadata, "num_blocks", None) + if num_blocks is not None: + return max(num_blocks[metadata.num_contexts : metadata.num_seqs]) + return metadata.num_generation_blocks + + +def run_fp4_mla_attention_decode( + metadata: Any, + layer_idx: int, + local_layer: int, + q_nope: torch.Tensor, + q_pe: torch.Tensor, + output: torch.Tensor, + *, + sm_scale: float, + kv_lora_rank: int, + qk_rope_head_dim: int, +) -> None: + """Run MLA decode with FP4 QK and FP4 PV tensor-core matmuls. + + Q is quantized to FP4, QK reads the packed K-view cache with swizzled + block scales, softmax probabilities are quantized to FP4 per page, and PV + repacks V nibbles from the shared KV cache while reading the auxiliary + V-view scale pool. No BF16 dequantized KV workspace is materialized on + this path. + """ + if not is_flashinfer_fp4_mla_attention_enabled(): + raise RuntimeError( + f"FP4 MLA attention decode requires {FLASHINFER_FP4_MLA_ATTENTION_ENV}=1." + ) + + head_dim = kv_lora_rank + qk_rope_head_dim + _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) + if metadata.page_size != FP4_MLA_TOKENS_PER_BLOCK: + raise ValueError( + f"FP4 MLA attention decode requires page_size={FP4_MLA_TOKENS_PER_BLOCK}, " + f"got {metadata.page_size}." + ) + + num_gen = q_nope.shape[0] + if num_gen == 0: + return + num_gen_seqs = metadata.num_seqs - metadata.num_contexts + if num_gen != num_gen_seqs: + raise NotImplementedError( + "FP4 MLA attention decode currently supports one query token per " + f"generation sequence, got {num_gen} query tokens for " + f"{num_gen_seqs} sequences." + ) + + num_heads = q_nope.shape[1] + if q_pe.shape[:2] != (num_gen, num_heads): + raise ValueError("FP4 MLA attention q_nope/q_pe batch dimensions do not match.") + if output.shape[:2] != (num_gen, num_heads): + raise ValueError("FP4 MLA attention output batch dimensions do not match.") + if q_nope.shape[-1] != kv_lora_rank: + raise ValueError( + f"q_nope last dimension must be kv_lora_rank={kv_lora_rank}, got {q_nope.shape[-1]}." + ) + if q_pe.shape[-1] != qk_rope_head_dim: + raise ValueError( + f"q_pe last dimension must be qk_rope_head_dim={qk_rope_head_dim}, " + f"got {q_pe.shape[-1]}." + ) + + if getattr(metadata, "fp4_mla_v_scale_pool", None) is None: + raise RuntimeError( + "FP4 MLA attention decode requires the auxiliary V scale pool to be allocated." + ) + + global_scale = _get_fp4_mla_global_scale(metadata, q_nope.device) + q_residual_dim = FP4_MLA_Q_RESIDUAL_DIM + _validate_fp4_mla_attention_q_shape(head_dim, q_residual_dim) + + q_full = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_q_buf", + (num_gen, num_heads, head_dim), + dtype=q_nope.dtype, + device=q_nope.device, + ) + q_full[..., :kv_lora_rank].copy_(q_nope) + q_full[..., kv_lora_rank:].copy_(q_pe) + q_2d = q_full.reshape(num_gen * num_heads, head_dim) + if q_2d.dtype not in (torch.bfloat16, torch.float8_e4m3fn): + raise TypeError( + f"FP4 MLA residual Q quantization requires BF16 or FP8 Q; got {q_2d.dtype}." + ) + q_fp4, q_sf = torch.ops.trtllm.fp4_quantize_with_residual( + q_2d, + global_scale, + q_residual_dim, + is_act=True, + ) + q_sf = q_sf.view(torch.float8_e4m3fn) + + kv_cache, sf_cache = _get_fp4_mla_cache_tensors(metadata, layer_idx) + sf_cache = sf_cache.view(torch.float8_e4m3fn) + + num_gen_blocks = metadata.num_generation_blocks + total_p_rows = num_gen_blocks * num_heads + p_fp4 = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_p_buf", + (max(total_p_rows, 1), metadata.page_size // 2), + dtype=torch.uint8, + device=q_nope.device, + )[:total_p_rows] + p_sf = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_p_sf_buf", + (max(_get_fp4_mla_swizzled_scale_size(total_p_rows, metadata.page_size), 1),), + dtype=torch.float8_e4m3fn, + device=q_nope.device, + ) + p_probs = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_p_prob_buf", + (max(num_gen * num_heads, 1), metadata.page_size), + dtype=torch.float32, + device=q_nope.device, + )[: num_gen * num_heads] + stats_shape = (num_gen, num_heads) + max_scores = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_max_buf", + stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + denom = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_denom_buf", + stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + + v_sf = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=kv_lora_rank)[local_layer].view( + torch.float8_e4m3fn + ) + + src_page_ids = metadata.paged_kv_indices[ + metadata.num_context_blocks : metadata.num_context_blocks + num_gen_blocks + ] + kv_lens = metadata.kv_lens_cuda_runtime[metadata.num_contexts : metadata.num_seqs] + max_pages = _max_generation_pages(metadata) + if max_pages == 0: + return + + block_h = 128 + block_t = metadata.page_size + block_k = 256 + block_v = 128 + q_head_dim = head_dim + q_residual_dim + q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE + k_sf_per_token = head_dim // FP4_BLOCK_SIZE + sf_per_page = metadata.page_size // FP4_BLOCK_SIZE + num_head_blocks = triton.cdiv(num_heads, block_h) + + _fp4_mla_debug( + "attention decode: " + f"num_gen={num_gen} num_heads={num_heads} local_layer={local_layer} " + f"layer_idx={layer_idx} head_dim={head_dim} q_head_dim={q_head_dim} " + f"q_residual_dim={q_residual_dim} " + f"kv_lora_rank={kv_lora_rank} rope_dim={qk_rope_head_dim} " + f"num_gen_blocks={num_gen_blocks} max_pages={max_pages} " + f"num_head_blocks={num_head_blocks} sm_scale={sm_scale}" + ) + _fp4_mla_debug(f"attention q_nope: {_tensor_layout(q_nope)}") + _fp4_mla_debug(f"attention q_pe: {_tensor_layout(q_pe)}") + _fp4_mla_debug(f"attention output: {_tensor_layout(output)}") + _fp4_mla_debug(f"attention q_fp4: {_tensor_layout(q_fp4)}") + _fp4_mla_debug(f"attention q_sf: {_tensor_layout(q_sf)}") + _fp4_mla_debug(f"attention kv_cache: {_tensor_layout(kv_cache)}") + _fp4_mla_debug(f"attention sf_cache: {_tensor_layout(sf_cache)}") + _fp4_mla_debug(f"attention v_sf: {_tensor_layout(v_sf)}") + _fp4_mla_debug(f"attention p_fp4: {_tensor_layout(p_fp4)}") + _fp4_mla_debug(f"attention p_sf: {_tensor_layout(p_sf)}") + _fp4_mla_debug(f"attention p_probs: {_tensor_layout(p_probs)}") + _debug_tensor_range("attention src_page_ids", src_page_ids) + _debug_tensor_range("attention paged_kv_indptr_decode", metadata.paged_kv_indptr_decode) + _debug_tensor_range("attention kv_lens", kv_lens) + + _fp4_mla_debug( + "attention stats launch: " + f"grid=({num_gen}, {num_head_blocks}) " + f"block_h={block_h} block_t={block_t} block_k={block_k}" + ) + _fp4_mla_attention_stats_kernel[(num_gen, num_head_blocks)]( + max_scores, + denom, + q_fp4, + q_sf, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + kv_cache.shape[0], + q_fp4.stride(0), + q_fp4.stride(1), + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + sf_cache.stride(0), + max_scores.stride(0), + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + BLOCK_T=block_t, + BLOCK_K=block_k, + ) + _debug_sync("attention_stats") + + for page_rel in range(max_pages): + _fp4_mla_debug( + "attention prob page store launch: " + f"page_rel={page_rel} grid=({num_gen}, {num_head_blocks})" + ) + _fp4_mla_attention_prob_store_page_kernel[(num_gen, num_head_blocks)]( + p_probs, + max_scores, + denom, + q_fp4, + q_sf, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + page_rel, + src_page_ids.shape[0], + kv_cache.shape[0], + p_probs.stride(0), + p_probs.stride(1), + q_fp4.stride(0), + q_fp4.stride(1), + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + sf_cache.stride(0), + max_scores.stride(0), + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + BLOCK_H=block_h, + BLOCK_K=block_k, + ) + _debug_sync(f"attention_prob_page_store_{page_rel}") + + _fp4_mla_debug( + "attention prob page pack launch: " + f"page_rel={page_rel} grid=({num_gen}, {sf_per_page}, " + f"{num_head_blocks})" + ) + _fp4_mla_attention_prob_pack_page_kernel[ + ( + num_gen, + sf_per_page, + num_head_blocks, + ) + ]( + p_fp4, + p_sf, + p_probs, + metadata.paged_kv_indptr_decode, + kv_lens, + page_rel, + src_page_ids.shape[0], + p_fp4.stride(0), + p_fp4.stride(1), + p_probs.stride(0), + p_probs.stride(1), + NUM_HEADS=num_heads, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + BLOCK_H=block_h, + ) + _debug_sync(f"attention_prob_page_pack_{page_rel}") + + _fp4_mla_debug( + "attention pv launch: " + f"grid=({num_gen}, {num_head_blocks}, " + f"{triton.cdiv(kv_lora_rank, block_v)}) block_v={block_v}" + ) + _fp4_mla_attention_pv_kernel[ + ( + num_gen, + num_head_blocks, + triton.cdiv(kv_lora_rank, block_v), + ) + ]( + output, + p_fp4, + p_sf, + kv_cache, + v_sf, + global_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + kv_cache.shape[0], + output.stride(0), + output.stride(1), + output.stride(2), + p_fp4.stride(0), + p_fp4.stride(1), + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=kv_lora_rank, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + BLOCK_H=block_h, + BLOCK_V=block_v, + ) + _debug_sync("attention_pv") + + +def update_hp_kv_for_fp4_mla( + metadata: Any, + latent_cache: Optional[torch.Tensor], + local_layer: int, + *, + phase: _HPUpdatePhase = "all", +) -> None: + """Store recent KV tokens at BF16 into the high-precision pool. + + Called on every layer before the attention kernel. The pool acts as a + circular buffer of HP_BLOCK_SIZE slots per sequence: + + Context phase stores the last ``kv_len % HP_BLOCK_SIZE`` new tokens of + each request into buffer positions [0, remainder). These are the + tail tokens that do not fill a complete FP4 block of 16. + + Generation phase stores the single new token for each request into + position ``(kv_len - 1) % HP_BLOCK_SIZE``, overwriting the oldest + entry in the circular buffer. + + The Triton kernels use the GPU ``seq_slots`` tensor for scatter indexing + and are CUDA-graph-compatible for the generation phase. + + Args: + metadata: Attention metadata exposing ``num_contexts``, ``num_seqs``, + ``seq_slots`` / ``seq_slots_cpu``, ``request_ids``, + ``is_cuda_graph``, ``is_warmup``, ``hp_pool_owners``, + ``high_precision_kv_pool``, ``prompt_lens_cpu_runtime``, + ``prompt_lens_cuda_runtime``, ``kv_lens_cuda_runtime``. + latent_cache: MLA latent cache for the current tokens, shape + [num_tokens, head_dim]. When ``None``, only ownership tracking + runs (no data is written to the pool). + local_layer: Layer index within the local pipeline-parallel slice. + phase: Which portion of ``latent_cache`` is present. ``"all"`` means + context tokens followed by generation tokens, ``"context"`` means + only context tokens, and ``"generation"`` means only generation + tokens. + """ + if phase not in ("all", "context", "generation"): + raise ValueError(f"Unexpected FP4 MLA HP update phase: {phase}") + if metadata.hp_pool_owners is None: + return + num_contexts = metadata.num_contexts + num_seqs = metadata.num_seqs + update_context = phase in ("all", "context") + update_generation = phase in ("all", "generation") + + # ------------------------------------------------------------------ + # Ownership tracking (layer 0, eager mode only - debug guard). + # Context phase never uses CUDA graph; decode check is debug-only. + # ------------------------------------------------------------------ + if local_layer == 0 and not metadata.is_cuda_graph and not metadata.is_warmup: + # Context: register ownership of each seq_slot. + if update_context: + for batch_idx in range(num_contexts): + seq_slot = metadata.seq_slots_cpu[batch_idx].item() + request_id = metadata.request_ids[batch_idx] + metadata.hp_pool_owners[seq_slot] = request_id + + # Decode: verify that the expected request still owns each slot. + if update_generation: + for batch_idx in range(num_contexts, num_seqs): + seq_slot = metadata.seq_slots_cpu[batch_idx].item() + request_id = metadata.request_ids[batch_idx] + owner = metadata.hp_pool_owners.get(seq_slot) + if owner != request_id: + raise RuntimeError( + f"HP KV pool ownership mismatch: seq_slot={seq_slot} " + f"is owned by request {owner} but request " + f"{request_id} is attempting to use it" + ) + + if latent_cache is None: + return + + # ------------------------------------------------------------------ + # Triton kernel dispatch - runs on every layer, CUDA-graph-safe. + # ------------------------------------------------------------------ + pool = metadata.high_precision_kv_pool + head_dim = latent_cache.shape[-1] + pool_head_dim = pool.shape[-1] // HP_BLOCK_SIZE + if pool_head_dim < head_dim: + raise RuntimeError( + f"FP4 MLA HP pool head dimension is too small: got " + f"{pool_head_dim}, need at least {head_dim}." + ) + block_d = triton.next_power_of_2(head_dim) + pool_s0 = pool.stride(0) # stride across sequence slots + pool_s1 = pool.stride(1) # stride across layers + lc_stride = latent_cache.stride(0) + _fp4_mla_debug( + "hp update: " + f"phase={phase} local_layer={local_layer} num_contexts={num_contexts} " + f"num_seqs={num_seqs} head_dim={head_dim} " + f"pool_head_dim={pool_head_dim} block_d={block_d}" + ) + _fp4_mla_debug(f"hp latent_cache: {_tensor_layout(latent_cache)}") + _fp4_mla_debug(f"hp pool: {_tensor_layout(pool)}") + _debug_tensor_range("hp seq_slots", metadata.seq_slots[:num_seqs]) + _debug_tensor_range("hp kv_lens", metadata.kv_lens_cuda_runtime[:num_seqs]) + + # Context phase: store last (kv_len % HP_BLOCK_SIZE) new tokens. + if update_context and num_contexts > 0: + prompt_lens_cpu = metadata.prompt_lens_cpu_runtime[:num_contexts] + # Exclusive prefix sum: token offset in latent_cache for each ctx seq. + token_offsets_cpu = torch.zeros(num_contexts, dtype=torch.int32, device="cpu") + if num_contexts > 1: + token_offsets_cpu[1:].copy_(torch.cumsum(prompt_lens_cpu[:-1].to(torch.int32), dim=0)) + token_offsets_gpu = token_offsets_cpu.to(pool.device, non_blocking=False) + prompt_lens_gpu = metadata.prompt_lens_cuda_runtime[:num_contexts] + + _fp4_mla_debug( + "hp context launch: " + f"grid=({num_contexts}, {HP_BLOCK_SIZE}) " + f"token_offsets={token_offsets_cpu.tolist()}" + ) + _debug_tensor_range("hp context prompt_lens", prompt_lens_gpu) + _hp_kv_store_context_kernel[(num_contexts, HP_BLOCK_SIZE)]( + pool, + latent_cache, + metadata.seq_slots, + metadata.kv_lens_cuda_runtime, + token_offsets_gpu, + prompt_lens_gpu, + pool.shape[0], + pool.shape[1], + local_layer, + pool_s0, + pool_s1, + lc_stride, + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + _debug_sync("hp_context") + + # Generation phase: store current token at (kv_len - 1) % HP_BLOCK_SIZE. + num_gen = num_seqs - num_contexts + if update_generation and num_gen > 0: + gen_tok_start = 0 + if phase == "all": + # Scalar offset: number of context tokens packed before gen tokens. + gen_tok_start = int(metadata.prompt_lens_cpu_runtime[:num_contexts].sum().item()) + + _fp4_mla_debug(f"hp generation launch: grid=({num_gen},) gen_tok_start={gen_tok_start}") + _hp_kv_store_gen_kernel[(num_gen,)]( + pool, + latent_cache, + metadata.seq_slots[num_contexts:], + metadata.kv_lens_cuda_runtime[num_contexts:], + gen_tok_start, + pool.shape[0], + pool.shape[1], + local_layer, + pool_s0, + pool_s1, + lc_stride, + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + _debug_sync("hp_generation") diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index fa2e77b8201c..c34564c7fd1a 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -21,8 +21,6 @@ from typing import TYPE_CHECKING, List, Optional, Tuple import torch -import triton -import triton.language as tl if TYPE_CHECKING: from ..speculative.interface import SpecMetadata @@ -40,6 +38,7 @@ from ..utils import (compute_swizzled_sf_shape, get_global_attrs, get_model_extra_attrs) +from .fp4_mla_kv import HP_BLOCK_SIZE, update_hp_kv_for_fp4_mla from .interface import (AttentionBackend, AttentionForwardArgs, AttentionInputType, AttentionMask, AttentionMetadata, KVCacheParams, MLAParams, PositionalEmbeddingParams, @@ -48,11 +47,6 @@ from .sparse.params import SparseParams from .sparse.skip_softmax import SkipSoftmaxParams -# Circular buffer size for the high-precision BF16 KV pool (MLA FP4 models). -# Each sequence slot holds HP_BLOCK_SIZE token vectors; token at absolute -# position t maps to buffer slot t % HP_BLOCK_SIZE. -HP_BLOCK_SIZE: int = 16 - @functools.cache def generate_spec_decoding_position_offsets(max_num_requests: int, @@ -167,11 +161,11 @@ class TrtllmAttentionMetadata(AttentionMetadata): # High-precision BF16 KV pool for MLA FP4 models, indexed by seq_slot. # Shape: [max_num_sequences, num_local_layers, kv_factor, HP_BLOCK_SIZE * head_dim] - # Standalone tensor — not part of the block-based paged KV cache. + # Standalone tensor, not part of the block-based paged KV cache. high_precision_kv_pool: Optional[torch.Tensor] = None - # Ownership tracking: maps seq_slot → request_id that last wrote it. + # Ownership tracking: maps seq_slot to request_id that last wrote it. # Plain Python dict, updated during context phase, checked during decode. - # Debug only — runs outside CUDA graph. + # Debug only; runs outside CUDA graph. hp_pool_owners: Optional[dict] = None # Pre-computed FlashMLA tile-scheduler metadata and num_splits. @@ -1183,106 +1177,6 @@ def is_sm_version_trtllm_gen_kernel(self, sm): return not (sm < 100 or sm in [120, 121]) -# --------------------------------------------------------------------------- -# Triton kernels for storing latent cache into the high-precision KV pool -# --------------------------------------------------------------------------- - - -@triton.jit -def _hp_kv_store_context_kernel( - pool_ptr, - latent_cache_ptr, - seq_slots_ptr, # int32 [num_contexts] – seq_slot for each ctx seq - kv_lens_ptr, # int32 [num_contexts] – total KV length after prefill - token_offsets_ptr, # int32 [num_contexts] – excl. prefix-sum of prompt_lens - prompt_lens_ptr, # int32 [num_contexts] – number of new tokens per ctx seq - layer_idx, - pool_stride_seq, # pool.stride(0): elements between adjacent seq slots - pool_stride_layer, # pool.stride(1): elements between adjacent layers - lc_stride, # latent_cache.stride(0): elements between adjacent tokens - D: tl.constexpr, - BLOCK_D: tl.constexpr, - HP_BLOCK: tl.constexpr, # = HP_BLOCK_SIZE (16) -): - """Store the tail tokens of each context sequence into the HP KV pool. - - Grid: (num_contexts, HP_BLOCK_SIZE). - Only programs where buf_pos < kv_len % HP_BLOCK actually write. - - For a context sequence with total KV length L = num_cached + prompt_len: - - remainder = L % HP_BLOCK - - The last `remainder` new tokens (latent_cache positions - [offset + prompt_len - remainder, offset + prompt_len)) are stored - into pool slots [0, remainder), which correspond to the absolute token - positions [L - remainder, L) in the circular buffer. - """ - ctx_idx = tl.program_id(0) - buf_pos = tl.program_id(1) - - kv_len = tl.load(kv_lens_ptr + ctx_idx) - remainder = kv_len % HP_BLOCK - if buf_pos >= remainder: - return - - seq_slot = tl.load(seq_slots_ptr + ctx_idx) - prompt_len = tl.load(prompt_lens_ptr + ctx_idx) - tok_offset = tl.load(token_offsets_ptr + ctx_idx) - - # Index of this token within latent_cache: last `remainder` new tokens, - # buf_pos-th of them (0-indexed from the start of the tail). - token_idx = tok_offset + prompt_len - remainder + buf_pos - - offs_d = tl.arange(0, BLOCK_D) - mask_d = offs_d < D - src = tl.load(latent_cache_ptr + token_idx * lc_stride + offs_d, - mask=mask_d, - other=0.0) - - # Destination: pool[seq_slot, layer_idx, 0, buf_pos * D : (buf_pos+1) * D] - dst_base = (seq_slot * pool_stride_seq + layer_idx * pool_stride_layer + - buf_pos * D) - tl.store(pool_ptr + dst_base + offs_d, src, mask=mask_d) - - -@triton.jit -def _hp_kv_store_gen_kernel( - pool_ptr, - latent_cache_ptr, - seq_slots_ptr, # int32 [num_gen] – seq_slot for each gen seq - kv_lens_ptr, # int32 [num_gen] – total KV length after this decode step - gen_tok_start, # int – offset in latent_cache where gen tokens begin - layer_idx, - pool_stride_seq, - pool_stride_layer, - lc_stride, - D: tl.constexpr, - BLOCK_D: tl.constexpr, - HP_BLOCK: tl.constexpr, -): - """Store the current generation token into the HP KV pool. - - Grid: (num_gen_seqs,). - Each program stores one token into the circular buffer position - (kv_len - 1) % HP_BLOCK, overwriting the oldest entry. - """ - gen_idx = tl.program_id(0) - - seq_slot = tl.load(seq_slots_ptr + gen_idx) - kv_len = tl.load(kv_lens_ptr + gen_idx) - buf_pos = (kv_len - 1) % HP_BLOCK - - token_idx = gen_tok_start + gen_idx - offs_d = tl.arange(0, BLOCK_D) - mask_d = offs_d < D - src = tl.load(latent_cache_ptr + token_idx * lc_stride + offs_d, - mask=mask_d, - other=0.0) - - dst_base = (seq_slot * pool_stride_seq + layer_idx * pool_stride_layer + - buf_pos * D) - tl.store(pool_ptr + dst_base + offs_d, src, mask=mask_d) - - class TrtllmAttention(AttentionBackend[TrtllmAttentionMetadata]): Metadata = TrtllmAttentionMetadata @@ -1592,87 +1486,20 @@ def _update_high_precision_kv_for_fp4_mla( self, metadata: TrtllmAttentionMetadata, latent_cache: Optional[torch.Tensor], + attention_input_type: AttentionInputType = AttentionInputType.mixed, ) -> None: - """Store recent KV tokens at BF16 into the high-precision pool.""" - if metadata.hp_pool_owners is None: - return - num_contexts = metadata.num_contexts - num_seqs = metadata.num_seqs - local_layer = self.get_local_layer_idx(metadata) - - if local_layer == 0 and not metadata.is_cuda_graph and not metadata.is_warmup: - for batch_idx in range(num_contexts): - seq_slot = metadata.seq_slots_cpu[batch_idx].item() - request_id = metadata.request_ids[batch_idx] - metadata.hp_pool_owners[seq_slot] = request_id - - for batch_idx in range(num_contexts, num_seqs): - seq_slot = metadata.seq_slots_cpu[batch_idx].item() - request_id = metadata.request_ids[batch_idx] - owner = metadata.hp_pool_owners.get(seq_slot) - if owner != request_id: - raise RuntimeError( - f"HP KV pool ownership mismatch: seq_slot={seq_slot} " - f"is owned by request {owner} but request " - f"{request_id} is attempting to use it") - - if latent_cache is None: - return - - pool = metadata.high_precision_kv_pool - head_dim = latent_cache.shape[-1] - block_d = triton.next_power_of_2(head_dim) - pool_s0 = pool.stride(0) - pool_s1 = pool.stride(1) - lc_stride = latent_cache.stride(0) - - if num_contexts > 0: - prompt_lens_cpu = metadata.prompt_lens_cpu_runtime[:num_contexts] - token_offsets_cpu = torch.zeros(num_contexts, - dtype=torch.int32, - device='cpu') - if num_contexts > 1: - token_offsets_cpu[1:].copy_( - torch.cumsum(prompt_lens_cpu[:-1].to(torch.int32), dim=0)) - token_offsets_gpu = token_offsets_cpu.to(pool.device, - non_blocking=False) - prompt_lens_gpu = metadata.prompt_lens_cuda_runtime[:num_contexts] - - _hp_kv_store_context_kernel[(num_contexts, HP_BLOCK_SIZE)]( - pool, - latent_cache, - metadata.seq_slots, - metadata.kv_lens_cuda_runtime, - token_offsets_gpu, - prompt_lens_gpu, - local_layer, - pool_s0, - pool_s1, - lc_stride, - D=head_dim, - BLOCK_D=block_d, - HP_BLOCK=HP_BLOCK_SIZE, - ) - - num_gen = num_seqs - num_contexts - if num_gen > 0: - ctx_tok_count = int( - metadata.prompt_lens_cpu_runtime[:num_contexts].sum().item()) - - _hp_kv_store_gen_kernel[(num_gen, )]( - pool, - latent_cache, - metadata.seq_slots[num_contexts:], - metadata.kv_lens_cuda_runtime[num_contexts:], - ctx_tok_count, - local_layer, - pool_s0, - pool_s1, - lc_stride, - D=head_dim, - BLOCK_D=block_d, - HP_BLOCK=HP_BLOCK_SIZE, - ) + """Thin wrapper over the shared HP-pool update helper (see + ``fp4_mla_kv.update_hp_kv_for_fp4_mla`` for the full contract).""" + if attention_input_type == AttentionInputType.context_only: + phase = "context" + elif attention_input_type == AttentionInputType.generation_only: + phase = "generation" + else: + phase = "all" + update_hp_kv_for_fp4_mla(metadata, + latent_cache, + self.get_local_layer_idx(metadata), + phase=phase) def forward( self, @@ -1893,7 +1720,8 @@ def forward( if metadata.high_precision_kv_pool is not None and self.is_mla_enable: self._update_high_precision_kv_for_fp4_mla( - metadata, forward_args.latent_cache) + metadata, forward_args.latent_cache, + forward_args.attention_input_type) if not self.fmha_libs: self.create_fmha_libs() diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index 797f2fd48666..58e63d7a1184 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -44,6 +44,10 @@ from .model_loader import ModelLoader, _construct_checkpoint_loader from .py_executor import PyExecutor +FLASH_MLA_TOKENS_PER_BLOCK = 64 +FP4_MLA_TOKENS_PER_BLOCK = 128 +FLASHINFER_FP4_MLA_ATTENTION_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION" + class _ExecutorMemoryMonitor: """Currently this focuses on tracking memory usage and related errors.""" @@ -181,6 +185,55 @@ def _set_model_engines_cache_reuse(model_engines, cache_reuse: bool): engine.attn_runtime_features.cache_reuse = cache_reuse +def _has_fp4_kv_cache(model_config, kv_cache_config) -> bool: + kv_cache_quant_algo = getattr(getattr(model_config, "quant_config", None), + "kv_cache_quant_algo", None) + fp4_quant_values = { + QuantAlgo.NVFP4, + getattr(QuantAlgo.NVFP4, "value", None), + "NVFP4", + } + kv_cache_dtype = getattr(kv_cache_config, "dtype", None) + return ((isinstance(kv_cache_dtype, str) + and kv_cache_dtype.lower() == "nvfp4") + or (isinstance(kv_cache_quant_algo, str) + and kv_cache_quant_algo.upper() == "NVFP4") + or kv_cache_quant_algo in fp4_quant_values) + + +def _enable_fp4_mla_attention() -> bool: + return os.getenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "0").lower() in ( + "1", + "true", + "yes", + "on", + ) + + +def _select_mla_tokens_per_block(config, model_config, kv_cache_config, + tokens_per_block: int) -> int: + if not is_mla(config): + return tokens_per_block + + if (_has_fp4_kv_cache(model_config, kv_cache_config) + and _enable_fp4_mla_attention()): + tokens_per_block = FP4_MLA_TOKENS_PER_BLOCK + logger.info( + f"Change tokens_per_block to: {tokens_per_block} for using FP4 MLA attention" + ) + kv_cache_config.tokens_per_block = tokens_per_block + return tokens_per_block + + if model_config.enable_flash_mla: + tokens_per_block = FLASH_MLA_TOKENS_PER_BLOCK + logger.info( + f"Change tokens_per_block to: {tokens_per_block} for using FlashMLA" + ) + kv_cache_config.tokens_per_block = tokens_per_block + + return tokens_per_block + + def _get_mapping(_mapping: Mapping) -> Mapping: if _mapping is None: mapping = Mapping(world_size=tensorrt_llm.mpi_world_size(), @@ -647,24 +700,10 @@ def drafting_loop_wrapper(model): kv_cache_config.enable_block_reuse = False _set_model_engines_cache_reuse([model_engine, draft_model_engine], False) + tokens_per_block = _select_mla_tokens_per_block( + config, model_engine.model.model_config, kv_cache_config, + tokens_per_block) if is_mla(config): - if model_engine.model.model_config.enable_flash_mla: - tokens_per_block = 64 - # Propagate the override back to kv_cache_config so any consumer - # that later reads llm_args.kv_cache_config.tokens_per_block sees - # the effective value. KvCacheConnectorScheduler subclasses - # (LMCache, Dynamo KVBM) are instantiated further down via - # scheduler_cls(llm_args) and size their block pools from - # llm_args.kv_cache_config.tokens_per_block. Without this the - # connector's block size desynced from the KVCacheManager's - # actual tokens_per_block (user-set or default 32 vs. FlashMLA's - # forced 64), producing a frozen cache_block_ids view to the - # connector and silently-corrupted decode KV (#13320). - kv_cache_config.tokens_per_block = tokens_per_block - logger.info( - f"Change tokens_per_block to: {tokens_per_block} for using FlashMLA" - ) - sm_version = get_sm_version() if kv_cache_config.enable_block_reuse and sm_version not in [ 90, 100, 103, 120 diff --git a/tensorrt_llm/_torch/pyexecutor/resource_manager.py b/tensorrt_llm/_torch/pyexecutor/resource_manager.py index e4ad48050606..9a46a9164956 100644 --- a/tensorrt_llm/_torch/pyexecutor/resource_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/resource_manager.py @@ -58,6 +58,16 @@ BlocksPerWindow = Dict[int, Tuple[ int, int]] # window_size -> (blocks_in_primary_pool, blocks_in_secondary_pool) +FLASHINFER_FP4_MLA_ATTENTION_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION" + + +def _flashinfer_fp4_mla_attention_enabled() -> bool: + return os.getenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "0").lower() in ( + "1", + "true", + "yes", + "on", + ) @dataclass @@ -570,6 +580,7 @@ def append_to_kv_heads_per_layer(num_kv_heads_per_layer: List[int], # Forward the (possibly remapped) per-pool configurations. # window_size values are aligned with the post-clamp sizes. 'pool_configurations': pool_configurations_cpp, + 'enable_mla_v_scale_pool': self._enable_mla_v_scale_pool(), } if self.event_buffer_max_size > 0: @@ -1144,8 +1155,25 @@ def get_cache_bytes_per_token(self): cache_size_per_token, quant_vector_size=16, scaling_factor_dtype=DataType.FP8) + cache_size_bytes_per_token += self._get_mla_v_scale_bytes_per_token( + self.num_local_layers) return cache_size_bytes_per_token + def _enable_mla_v_scale_pool(self) -> bool: + return (self.dtype == DataType.NVFP4 + and self.kv_cache_type == CacheTypeCpp.SELFKONLY + and _flashinfer_fp4_mla_attention_enabled()) + + def _get_mla_v_scale_bytes_per_token(self, num_layers: int) -> int: + if not self._enable_mla_v_scale_pool(): + return 0 + token_scale_cols = math.ceil(self.tokens_per_block / 16) + # Match the C++ per-page allocation and spread it across page tokens + # for scheduler capacity accounting. + elems_per_page = (math.ceil(self.head_dim / 128) * + math.ceil(token_scale_cols / 4) * 32 * 16) + return num_layers * math.ceil(elems_per_page / self.tokens_per_block) + def calculate_max_num_blocks(self, kv_cache_config: KvCacheConfig, head_dim: int, @@ -1402,16 +1430,23 @@ def get_buffers(self, pool = self.get_pool_for_layer(layer_offset) layer_head_dim = pool.head_dim if pool else self.head_dim + layer_dtype = pool.dtype if pool else self.dtype assert kv_layout in ["NHD", "HND"], f"Unsupported kv_layout: {kv_layout}" + + element_per_container = 1 + if layer_dtype == DataType.NVFP4: + element_per_container = 2 + effective_head_dim = layer_head_dim // element_per_container + if kv_layout == "NHD": return result.reshape( result.shape[0], self.kv_factor, self.tokens_per_block, self.num_kv_heads_per_layer[layer_offset], - layer_head_dim, + effective_head_dim, ) else: return result.reshape( @@ -1419,9 +1454,86 @@ def get_buffers(self, self.kv_factor, self.num_kv_heads_per_layer[layer_offset], self.tokens_per_block, - layer_head_dim, + effective_head_dim, ) + def get_block_scale_buffers( + self, + layer_idx: int, + kv_layout: str = "NHD") -> Optional[torch.Tensor]: + '''Slice the NVFP4 block-scale tensor for a specified layer. + + The returned tensor uses raw uint8 storage for the FP8 E4M3 scale + bytes. The logical layouts mirror ``get_buffers()`` except the last + dimension is ``head_dim // 16`` scale values per token. + ''' + layer_offset = self.layer_offsets[layer_idx] + pool = self.get_pool_for_layer(layer_offset) + layer_head_dim = pool.head_dim if pool else self.head_dim + layer_dtype = pool.dtype if pool else self.dtype + if layer_dtype != DataType.NVFP4: + return None + + assert kv_layout in ["NHD", + "HND"], f"Unsupported kv_layout: {kv_layout}" + + pool_id = int(self.kv_cache_pool_mapping[layer_offset][0].item()) + pool_layer_idx = int(self.kv_cache_pool_mapping[layer_offset][1].item()) + scale_pool_ptr = int(self.kv_cache_pool_pointers[pool_id][0][1].item()) + if scale_pool_ptr == 0: + return None + + num_layers_in_pool = int( + (self.kv_cache_pool_mapping[:, 0] == pool_id).sum().item()) + num_kv_heads = self.num_kv_heads_per_layer[layer_offset] + scales_per_head = layer_head_dim // 16 + block_size = self.kv_factor * self.tokens_per_block * num_kv_heads * scales_per_head + layer_offset_elements = pool_layer_idx * block_size + + if kv_layout == "NHD": + shape = [ + self.blocks_in_primary_pool, + self.kv_factor, + self.tokens_per_block, + num_kv_heads, + scales_per_head, + ] + strides = [ + num_layers_in_pool * block_size, + self.tokens_per_block * num_kv_heads * scales_per_head, + num_kv_heads * scales_per_head, + scales_per_head, + 1, + ] + else: + shape = [ + self.blocks_in_primary_pool, + self.kv_factor, + num_kv_heads, + self.tokens_per_block, + scales_per_head, + ] + strides = [ + num_layers_in_pool * block_size, + self.tokens_per_block * num_kv_heads * scales_per_head, + scales_per_head, + num_kv_heads * scales_per_head, + 1, + ] + + return convert_to_torch_tensor( + TensorWrapper( + scale_pool_ptr + layer_offset_elements, + torch.uint8, + shape, + strides, + )) + + def get_mla_v_scale_pool(self) -> Optional[torch.Tensor]: + if not self._enable_mla_v_scale_pool(): + return None + return self.impl.get_mla_v_scale_pool() + def get_indexer_k_cache_pool_data(self, layer_idx: int) -> torch.Tensor: result = self.impl.get_indexer_k_cache_pool_data(layer_idx) return result.view(result.shape[0], -1) @@ -1563,6 +1675,7 @@ def _calculate_cache_bytes_per_token_for_layers( layer_elements, quant_vector_size=16, scaling_factor_dtype=DataType.FP8) + layer_bytes += self._get_mla_v_scale_bytes_per_token(1) total_bytes += layer_bytes return total_bytes @@ -2125,6 +2238,1432 @@ def reset_reuse_state(self): """Reset the reuse state of the KV cache manager.""" self.impl.reset_reuse_state() +class KVCacheManagerV2(BaseResourceManager): + + def __init__( + self, + kv_cache_config: KvCacheConfig, + kv_cache_type: CacheTypeCpp, + *, + num_layers: int, + num_kv_heads: Union[int, List[Optional[int]]], + head_dim: Union[int, List[int]], + tokens_per_block: int, + # Note that max_seq_len is not necessarily equal to kv_cache_config.num_tokens. + # It's derived from the model's BuildConfig for consistency with the C++ backend. + max_seq_len: int, + max_batch_size: int, + mapping: Mapping, + dtype: DataType = DataType.HALF, + spec_config=None, + layer_mask: Optional[List[bool]] = None, + vocab_size: int = None, + max_num_tokens: int = 8192, + model_config: Optional[ModelConfigCpp] = None, + max_beam_width: int = 1, + is_draft: bool = False, + kv_connector_manager: Optional[KvCacheConnectorManager] = None, + execution_stream: Optional[torch.cuda.Stream] = None, + is_disagg: bool = False, + **kwargs, + ) -> None: + self.mapping = mapping + self.dtype = dtype + self.is_disagg = is_disagg + + assert kv_connector_manager is None, "kv_connector_manager is not supported for KVCacheManagerV2" + assert max_beam_width == 1, "max_beam_width must be 1 for KVCacheManagerV2" + assert not (mapping.cp_config.get('cp_type') == CpType.STAR), \ + "Star attention is not supported for KVCacheManagerV2" + + self.kv_cache_type = kv_cache_type + self.pp_layers, self.num_layers = get_pp_layers( + num_layers, + mapping, + spec_config=spec_config, + layer_mask=layer_mask, + ) + self.is_draft = is_draft + self.num_local_layers = len(self.pp_layers) + self.layer_offsets = { + idx: offset + for offset, idx in enumerate(self.pp_layers) + } + self.max_beam_width = max_beam_width + + tp_size = mapping.tp_size + if mapping.enable_attention_dp: + tp_size = 1 + + self.num_kv_heads = num_kv_heads + self.head_dim = head_dim + self.tokens_per_block = tokens_per_block + self.max_seq_len = max_seq_len + self.max_batch_size = max_batch_size + self.kv_factor = 1 if kv_cache_type == CacheTypeCpp.SELFKONLY else 2 + from ..speculative import get_num_extra_kv_tokens + self.num_extra_kv_tokens = get_num_extra_kv_tokens(spec_config) + self.max_total_draft_tokens = spec_config.max_total_draft_tokens if spec_config is not None else 0 + + # Mirror V1's KV reserve sizing (see V1 __init__ for rationale). + self._kv_reserve_draft_tokens = self.max_total_draft_tokens + if (self.is_draft and spec_config is not None + and getattr(spec_config, 'use_dynamic_tree', False) + and getattr(spec_config, 'dynamic_tree_max_topK', 0) > 0): + draft_loop_tokens = spec_config.dynamic_tree_max_topK * spec_config.max_draft_len + self._kv_reserve_draft_tokens = max(self.max_total_draft_tokens, + draft_loop_tokens) + + self.event_buffer_max_size = kv_cache_config.event_buffer_max_size + + assert self.event_buffer_max_size == 0, "event_buffer_max_size must be 0" + + self._stream = execution_stream if execution_stream is not None else torch.cuda.current_stream( + ) + logger.info(f"[KVCacheManager] execution_stream: {self._stream}") + + # Determine max_attention_window_vec + if kv_cache_config.max_attention_window is not None: + + self.max_attention_window_vec = kv_cache_config.max_attention_window.copy( + ) # Make a copy to avoid modifying original + # Clamp all window sizes to max_seq_len before calculating the + # number of KV cache blocks. This prevents the KV cache pool from + # being skewed by the largest window values. + self.max_attention_window_vec = [ + min(max_seq_len, w) for w in self.max_attention_window_vec + ] + + self.max_attention_window_vec = [ + None if w == max_seq_len else w + for w in self.max_attention_window_vec + ] + + else: + self.max_attention_window_vec = [None] + + if isinstance(num_kv_heads, int): + self.num_kv_heads_per_layer = [ + (num_kv_heads + tp_size - 1) // tp_size + for _ in range(self.num_local_layers) + ] + self.total_num_kv_heads_per_layer = [ + (num_kv_heads + tp_size - 1) // tp_size + for _ in range(self.num_layers) + ] + else: + assert len(num_kv_heads) == self.num_layers + + def append_to_kv_heads_per_layer(num_kv_heads_per_layer: List[int], + kv_head: Optional[int]): + if kv_head is not None: + num_kv_heads_per_layer.append( + (kv_head + tp_size - 1) // tp_size) + else: + num_kv_heads_per_layer.append(0) + + self.num_kv_heads_per_layer = [] + if self.num_local_layers > 0: + for i in self.pp_layers: + kv_head = num_kv_heads[i] + append_to_kv_heads_per_layer(self.num_kv_heads_per_layer, + kv_head) + + self.total_num_kv_heads_per_layer = [] + for i in range(self.num_layers): + kv_head = num_kv_heads[i] + append_to_kv_heads_per_layer(self.total_num_kv_heads_per_layer, + kv_head) + + # Build per-layer head_dim (similar to num_kv_heads_per_layer) + if isinstance(head_dim, int): + self.head_dim_per_layer = [ + head_dim for _ in range(self.num_local_layers) + ] + else: + assert len(head_dim) == self.num_layers, \ + f"head_dim list length ({len(head_dim)}) must match num_layers ({self.num_layers})" + self.head_dim_per_layer = [] + if self.num_local_layers > 0: + for i in self.pp_layers: + self.head_dim_per_layer.append(head_dim[i]) + if len(set(self.head_dim_per_layer)) > 1: + logger.info( + f"Per-layer head_dim: {len(self.head_dim_per_layer)} layers, " + f"unique values={set(self.head_dim_per_layer)}") + + self.is_vswa = len(set(self.max_attention_window_vec)) > 1 + + quota = float('inf') + if kv_cache_config.max_gpu_total_bytes is not None and kv_cache_config.max_gpu_total_bytes > 0: + quota = int(kv_cache_config.max_gpu_total_bytes) + logger.info( + f"max_gpu_total_bytes is provided. New quota is {quota / (1 << 30)}GiB" + ) + if kv_cache_config.max_tokens is not None: + quota_from_max_tokens = int( + math.ceil( + self._get_quota_from_max_tokens(kv_cache_config.max_tokens) + / kv_cache_config.max_util_for_resume)) + quota = min(quota, quota_from_max_tokens) + logger.info( + f"max_tokens {kv_cache_config.max_tokens} is provided. Allowed quota from max_tokens is {quota_from_max_tokens / (1 << 30)}GiB. New quota is {quota / (1 << 30)}GiB" + ) + + assert quota != float( + 'inf' + ), "Quota not set. Check kv_cache_config.max_tokens or kv_cache_config.max_gpu_total_bytes" + + # Sync KV cache token capacity across ranks so all ranks allocate + # the same number of tokens and the scheduler produces identical + # batches. Normalize to token count before the allreduce because + # bytes_per_token varies across PP ranks (different local layers). + if mapping.world_size > 1: + dist = Distributed.get(mapping) + bytes_per_token = self.get_cache_bytes_per_token() + max_tokens = quota / bytes_per_token + max_tokens = dist.allreduce(max_tokens, op=ReduceOp.MIN) + quota = max_tokens * bytes_per_token + + logger.info( + f"KV cache manager v2 device quota set to {quota / (1 << 30)}GiB") + + cache_tiers: List[CacheTierConfig] = [GpuCacheTierConfig(quota=quota)] + if kv_cache_config.host_cache_size is not None and kv_cache_config.host_cache_size > 0: + host_quota = kv_cache_config.host_cache_size + else: + # The V2 MAX_UTILIZATION scheduler relies on suspend/resume to + # evict and later restore KV cache pages. Without a host tier, + # suspended pages have nowhere to be offloaded and resume() + # always fails, causing a scheduling deadlock where no + # generation request can ever make progress. + # + # Automatically provision a host tier matching the GPU quota so + # suspend/resume works out of the box. Cap at available host + # memory to avoid allocation failures. + try: + mem_available = os.sysconf('SC_PAGE_SIZE') * os.sysconf( + 'SC_AVPHYS_PAGES') + except (ValueError, OSError): + mem_available = float('inf') + host_quota = min(quota, int(mem_available * 0.5)) + if host_quota <= 0: + host_quota = quota + if host_quota > 0: + cache_tiers.append(HostCacheTierConfig(quota=host_quota)) + logger.info( + f"KV cache manager v2 host cache quota set to {host_quota / (1 << 30):.2f}GiB" + ) + disk_cache_size = kv_cache_config.disk_cache_size + if disk_cache_size is not None and disk_cache_size > 0: + disk_cache_path = kv_cache_config.disk_cache_path + assert disk_cache_path is not None + cache_tiers.append( + DiskCacheTierConfig(quota=disk_cache_size, + path=disk_cache_path)) + logger.info( + f"KV cache manager v2 disk cache quota set to {disk_cache_size / (1 << 30):.2f}GiB at {disk_cache_path}" + ) + + self.vocab_size = vocab_size + + config = self._build_cache_config( + kv_cache_config, + tokens_per_block=tokens_per_block, + vocab_size=vocab_size, + cache_tiers=cache_tiers, + ) + + self.kv_cache_manager_py_config = config + + self.impl = KVCacheManagerPy(config) + + self.num_pools = len(self.impl.layer_grouping) + + num_layers = len(config.layers) + self.layer_to_pool_mapping_dict: dict[int, int] = { + layer_id: self.impl.get_layer_group_id(layer_id) + for layer_id in typed_range(LayerId(num_layers)) + } + + (self.kv_cache_pool_pointers, + self.kv_cache_pool_mapping) = self._build_pool_mapping_tensors() + + self.kv_cache_map: dict[int, _KVCache] = {} + + # Tracks the draft length allocated by try_allocate_generation per + # request. Used by extend_capacity_for_tokens to compute the exact + # padding delta instead of blindly extending, which would cause + # unbounded capacity growth. + self._allocated_draft_lens: dict[int, int] = {} + + # Defensive cap for get_num_available_tokens: when host cache is + # enabled, clamp_max_seq_len_for_mem may return a value that spans + # both GPU and host tiers. Storing the explicit max_tokens (if set) + # lets us cap the result to GPU-only capacity so callers like CUDA + # graph warmup don't over-allocate beyond the GPU pool. + # None when max_tokens is not explicitly configured — other config + # paths (max_gpu_total_bytes, free_gpu_memory_fraction) are already + # bounded by the GPU quota passed to GpuCacheTierConfig. + self._gpu_max_tokens = kv_cache_config.max_tokens + + max_num_tokens = self.get_num_available_tokens( + token_num_upper_bound=max_seq_len) + + if max_seq_len > max_num_tokens: + logger.warning( + f"max_seq_len {max_seq_len} is greater than max_num_tokens {max_num_tokens} that can be allocated in kv cache manager, setting max_seq_len to {max_num_tokens}" + ) + # max_num_tokens is a float from clamp_max_seq_len_for_mem; cast + # so downstream int-only consumers (torch.randint size, range) + # stay int. + self.max_seq_len = int(max_num_tokens) + + # Pad max_blocks_per_seq to next multiple of 4 (copy_block_offsets kernel). + # Account for max single-sequence capacity = seq_len + extra KV tokens + + # _kv_reserve_draft_tokens (see __init__) + 1 base decode token. + max_seq_capacity = self.max_seq_len + self.num_extra_kv_tokens + self._kv_reserve_draft_tokens + 1 + self.max_blocks_per_seq = (max_seq_capacity + tokens_per_block - + 1) // tokens_per_block + if self.max_blocks_per_seq % 4 != 0: + self.max_blocks_per_seq = ((self.max_blocks_per_seq + 3) // 4) * 4 + + self.enable_block_reuse = kv_cache_config.enable_block_reuse + self.enable_partial_reuse = kv_cache_config.enable_partial_reuse + + # With pipeline parallelism, multiple microbatches can be in-flight + # simultaneously, so we need slots for all concurrent sequences. + # Plus 1 for cuda graph dummy request. + # In disaggregated mode, use a coefficient of 2: at any moment up to + # `max_num_sequences` requests can be actively generating while another + # up to `max_num_sequences` requests are still in KV transfer + # (TRANS_IN_PROGRESS) and continue to hold their index slots. The 2x + # capacity lets the next batch of active requests acquire slots without + # waiting for the previous batch's transfers to finish. + max_num_sequences = max_batch_size * mapping.pp_size + index_mapper_capacity = max_num_sequences * (2 if is_disagg else 1) + 1 + logger.info( + f"KVCacheManagerV2: IndexMapper capacity={index_mapper_capacity} " + f"(max_num_sequences={max_num_sequences}, is_disagg={is_disagg}, max_beam_width={max_beam_width})" + ) + self.index_mapper = IndexMapper(index_mapper_capacity, max_beam_width) + self._early_freed_index_requests: set[int] = set() + self.index_scales = torch.empty(self.num_pools, + dtype=torch.int32, + pin_memory=prefer_pinned(), + device='cpu') + self.kv_offset = torch.empty(self.num_pools, + dtype=torch.int32, + pin_memory=prefer_pinned(), + device='cpu') + for pool_id in range(self.num_pools): + layer_id = self.impl.layer_grouping[pool_id][0] + self.index_scales[pool_id] = self.impl.get_page_index_scale( + layer_id, Role.KEY) + if self.kv_cache_type != CacheTypeCpp.SELFKONLY: + self.kv_offset[pool_id] = exact_div( + self.impl.get_mem_pool_base_address(layer_id, Role.VALUE) - + self.impl.get_mem_pool_base_address(layer_id, Role.KEY), + self.impl.get_page_stride(layer_id, Role.KEY)) + else: + self.kv_offset[pool_id] = 0 + + # Keep unused block offsets as safe block index 0. + self.host_kv_cache_block_offsets = torch.zeros( + self.num_pools, + index_mapper_capacity * max_beam_width, + 2, # key and value + self.max_blocks_per_seq, + dtype=torch.int32, + pin_memory=prefer_pinned(), + device='cpu') + + def _get_quota_from_max_tokens(self, max_tokens: int) -> int: + return int(max_tokens * self.get_cache_bytes_per_token()) + + def _build_pool_mapping_tensors(self) -> Tuple[torch.Tensor, torch.Tensor]: + kv_cache_pool_pointers = torch.tensor([[ + self.impl.get_mem_pool_base_address( + self.impl.layer_grouping[pool_id][0], Role.KEY), 0 + ] for pool_id in range(self.num_pools)], + dtype=torch.int64, + device="cpu", + pin_memory=prefer_pinned()) + + if self.dtype == DataType.NVFP4: + kv_cache_pool_pointers = torch.stack([ + kv_cache_pool_pointers, + torch.tensor([[ + self.impl.get_mem_pool_base_address( + self.impl.layer_grouping[pool_id][0], + Role.KEY_BLOCK_SCALE), 0 + ] for pool_id in range(self.num_pools)], + dtype=torch.int64, + device="cpu", + pin_memory=prefer_pinned()) + ], + dim=-1) + + kv_cache_pool_mapping_list = [] + for layer_id in typed_range(LayerId(self.num_local_layers)): + layer_group_id = self.impl.get_layer_group_id(layer_id) + if self.dtype != DataType.NVFP4: + addr_offset = self.impl.get_mem_pool_base_address( + layer_id, Role.KEY) - int( + kv_cache_pool_pointers[layer_group_id][0]) + else: + addr_offset = self.impl.get_mem_pool_base_address( + layer_id, Role.KEY) - int( + kv_cache_pool_pointers[layer_group_id][0][0]) + block_scale_addr_offset = self.impl.get_mem_pool_base_address( + layer_id, Role.KEY_BLOCK_SCALE) - int( + kv_cache_pool_pointers[layer_group_id][0][1]) + block_scale_offset = exact_div( + block_scale_addr_offset, + self.get_layer_bytes_per_token( + layer_id, Role.KEY_BLOCK_SCALE) * self.kv_factor * + self.tokens_per_block) + offset = exact_div( + addr_offset, + self.get_layer_bytes_per_token(layer_id, Role.KEY) * + self.kv_factor * self.tokens_per_block) + + if self.dtype == DataType.NVFP4: + assert block_scale_offset == offset, "Block scale offset and offset should be the same" + + kv_cache_pool_mapping_list.append([layer_group_id, offset]) + + kv_cache_pool_mapping = torch.tensor(kv_cache_pool_mapping_list, + dtype=torch.int32, + device="cpu", + pin_memory=prefer_pinned()) + return kv_cache_pool_pointers, kv_cache_pool_mapping + + def _build_cache_config( + self, + kv_cache_config: KvCacheConfig, + *, + tokens_per_block: int, + vocab_size: Optional[int], + cache_tiers: List[CacheTierConfig], + ) -> KVCacheManagerConfigPy: + buffer_type = [Role.KEY] + if self.kv_cache_type != CacheTypeCpp.SELFKONLY: + buffer_type.append(Role.VALUE) + if kv_cache_config.dtype == "nvfp4": + for layer_idx, hd in enumerate(self.head_dim_per_layer): + assert hd % 2 == 0, \ + f"head_dim must be divisible by 2 for nvfp4 kv cache, but layer {layer_idx} has head_dim={hd}" + buffer_type.append(Role.KEY_BLOCK_SCALE) + if self.kv_cache_type != CacheTypeCpp.SELFKONLY: + buffer_type.append(Role.VALUE_BLOCK_SCALE) + + return KVCacheManagerConfigPy( + tokens_per_block=tokens_per_block, + vocab_size=vocab_size, + cache_tiers=cache_tiers, + max_util_for_resume=kv_cache_config.max_util_for_resume, + layers=[ + AttentionLayerConfig( + layer_id=layer_id, + buffers=[ + BufferConfig( + role=role, + size=self.get_layer_bytes_per_token( + local_layer_idx=layer_id, data_role=role) * + tokens_per_block, + ) for role in buffer_type + ], + sliding_window_size=self.max_attention_window_vec[ + self.pp_layers[layer_id] % + len(self.max_attention_window_vec)], + num_sink_tokens=None, + ) for layer_id in typed_range(LayerId(self.num_local_layers)) + ], + ) + + @property + def blocks_in_primary_pool(self) -> int: + """ + Get the number of blocks in the primary pool. + """ + return self.impl.get_page_index_upper_bound(0, Role.KEY) + + def get_buffers(self, + layer_idx: int, + kv_layout: str = "NHD") -> Optional[torch.Tensor]: + layer_offset = self.layer_offsets[layer_idx] + addr_key = self.impl.get_mem_pool_base_address(layer_offset, Role.KEY) + if self.kv_cache_type != CacheTypeCpp.SELFKONLY: + addr_value = self.impl.get_mem_pool_base_address( + layer_offset, Role.VALUE) + page_size_key = self.impl.get_page_stride(layer_offset, Role.KEY) + page_size_value = self.impl.get_page_stride(layer_offset, + Role.VALUE) + + assert addr_key + page_size_value == addr_value and page_size_key == page_size_value + + assert kv_layout in ["NHD", + "HND"], f"Unsupported kv_layout: {kv_layout}" + + element_per_container = 1 + dtype = self.dtype + if dtype == DataType.NVFP4: + element_per_container = 2 + dtype = torch.int8 + + layer_head_dim = self.head_dim_per_layer[layer_offset] + if kv_layout == "NHD": + shape = [ + self.impl.get_page_index_upper_bound(layer_offset, Role.KEY) // + self.kv_factor, + self.kv_factor, + self.tokens_per_block, + self.num_kv_heads_per_layer[layer_offset], + layer_head_dim // element_per_container, + ] + else: + shape = [ + self.impl.get_page_index_upper_bound(layer_offset, Role.KEY) // + self.kv_factor, + self.kv_factor, + self.num_kv_heads_per_layer[layer_offset], + self.tokens_per_block, + layer_head_dim // element_per_container, + ] + + return convert_to_torch_tensor(TensorWrapper( + addr_key, + dtype, + shape, + )) + + def get_block_scale_buffers( + self, + layer_idx: int, + kv_layout: str = "NHD") -> Optional[torch.Tensor]: + if self.dtype != DataType.NVFP4: + return None + + layer_offset = self.layer_offsets[layer_idx] + addr_key = self.impl.get_mem_pool_base_address(layer_offset, + Role.KEY_BLOCK_SCALE) + if self.kv_cache_type != CacheTypeCpp.SELFKONLY: + addr_value = self.impl.get_mem_pool_base_address( + layer_offset, Role.VALUE_BLOCK_SCALE) + page_size_key = self.impl.get_page_stride(layer_offset, + Role.KEY_BLOCK_SCALE) + page_size_value = self.impl.get_page_stride(layer_offset, + Role.VALUE_BLOCK_SCALE) + + assert addr_key + page_size_value == addr_value and page_size_key == page_size_value + + assert kv_layout in ["NHD", + "HND"], f"Unsupported kv_layout: {kv_layout}" + + scales_per_head = self.head_dim // 16 + if kv_layout == "NHD": + shape = [ + self.impl.get_page_index_upper_bound(layer_offset, Role.KEY) // + self.kv_factor, + self.kv_factor, + self.tokens_per_block, + self.num_kv_heads_per_layer[layer_offset], + scales_per_head, + ] + else: + shape = [ + self.impl.get_page_index_upper_bound(layer_offset, Role.KEY) // + self.kv_factor, + self.kv_factor, + self.num_kv_heads_per_layer[layer_offset], + self.tokens_per_block, + scales_per_head, + ] + + return convert_to_torch_tensor( + TensorWrapper( + addr_key, + torch.uint8, + shape, + )) + + def get_num_available_tokens(self, + *, + token_num_upper_bound: int, + batch_size: int = 1, + max_num_draft_tokens: int = 0) -> int: + extra_tokens = self.num_extra_kv_tokens + max_num_draft_tokens + # Token num upper bound is the maximum number of tokens that can be allocated in the kv cache manager. + # We need to add extra tokens to the token num upper bound to account for the extra tokens. + clamped = self.impl.clamp_max_seq_len_for_mem( + batch_size, token_num_upper_bound + extra_tokens) - extra_tokens + # clamp_max_seq_len_for_mem considers all tiers (GPU + host). When + # max_tokens is explicitly set, cap by GPU-only capacity so callers + # (e.g. CUDA graph warmup) don't exceed the GPU pool. + if self._gpu_max_tokens is not None: + clamped = min(clamped, self._gpu_max_tokens - extra_tokens) + return clamped + + def get_num_free_blocks(self) -> int: + # NOTE This method is used to get the number of blocks in the primary pool not the FREE blocks. + # However, since we only use this function when the kv cache manager is empty, so it is safe to do so. + assert len( + self.kv_cache_map + ) == 0, "get_num_free_blocks is only used when the kv cache manager is empty" + max_num_pages = max([ + self.impl.get_page_index_upper_bound(layer_id, Role.KEY) + for layer_id in typed_range(LayerId(self.num_local_layers)) + ]) + return max_num_pages // self.kv_factor + + # ---- Scheduling API (called by KVCacheV2Scheduler) ---- + + def is_request_active(self, request_id: int) -> bool: + """Return True if *request_id* has a live, non-suspended KV cache.""" + kv_cache = self.kv_cache_map.get(request_id) + return kv_cache is not None and kv_cache.is_active + + def _required_gen_capacity(self, req: LlmRequest, + current_capacity: int) -> int: + """Compute generation KV cache capacity for a request. + + Grows *current_capacity* by 1 + draft tokens. + """ + draft_len = get_draft_token_length(req) + return current_capacity + 1 + draft_len + + def try_allocate_generation(self, req: LlmRequest) -> bool: + """Try to allocate one additional KV cache slot for a generation request. + + Resumes from suspended state if needed, then resizes capacity by 1 (+ + draft tokens). Returns True on success, False if allocation failed. + """ + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None: + return False + + if not kv_cache.is_active: + if not kv_cache.resume(self._stream.cuda_stream): + return False + self._restore_page_index_bufs(req.py_request_id, kv_cache) + + draft_len = get_draft_token_length(req) + self._allocated_draft_lens[req.py_request_id] = draft_len + return kv_cache.resize( + self._required_gen_capacity(req, kv_cache.capacity)) + + def revert_allocate_generation(self, req: LlmRequest) -> None: + """Undo the capacity growth from try_allocate_generation. + + When attention DP causes can_queue=False after scheduling, the + forward pass is skipped but the scheduler already grew each + generation request's KV cache capacity by 1 (+draft tokens). + This method shrinks capacity back to undo that spurious growth + so it does not accumulate across iterations and overflow the + host page-index buffer. + """ + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None or not kv_cache.is_active: + return + draft_len = get_draft_token_length(req) + reverted_cap = kv_cache.capacity - 1 - draft_len + if reverted_cap < 0: + return + if not kv_cache.resize(reverted_cap): + raise RuntimeError( + f"Failed to revert KV cache capacity for request " + f"{req.py_request_id} from {kv_cache.capacity} to " + f"{reverted_cap}") + + def revert_allocate_context(self, req: LlmRequest) -> None: + """Undo the capacity growth from this iter's ``resize_context``. + + When delay batching (``_balance_adp_requests`` / + ``_waiting_requests``) defers a context request after V2 + scheduling, the forward pass is skipped for that request but the + scheduler already grew its KV cache capacity to cover the chunk. + This shrinks capacity back to the pre-resize value so the + freshly-allocated pages can be reused during the wait window — + important for long contexts where one deferred request can hold + GBs of KV. + """ + pre_cap = getattr(req, "py_ctx_pre_resize_cap", None) + if pre_cap is None: + return + # Mark as consumed even if the resize below is skipped, so a + # later iter does not see a stale snapshot. + req.py_ctx_pre_resize_cap = None + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None or not kv_cache.is_active: + return + if pre_cap >= kv_cache.capacity: + return + if not kv_cache.resize(pre_cap): + raise RuntimeError( + f"Failed to revert KV cache capacity for context " + f"request {req.py_request_id} from " + f"{kv_cache.capacity} to {pre_cap}") + if pre_cap > 0: + kv_cache.suspend() + + def _restore_page_index_bufs(self, request_id: int, kv_cache) -> None: + """Re-connect host page-index buffers after resume(). + + suspend() clears the base_page_index_buf pointers (sets them to + None) so the KV cache stops writing page indices to the host + buffer. After resume(), the KV cache has re-locked pages but + copy_batch_block_offsets still reads from the host buffer, so we + must re-connect the buffers to avoid stale/zero page indices that + would cause illegal memory accesses during the forward pass. + """ + index = self.index_mapper.get_index(request_id) + for i in range(self.max_beam_width): + for pool_idx in range(self.num_pools): + buffer: torch.Tensor = self.host_kv_cache_block_offsets[ + pool_idx, index * self.max_beam_width + i, 0] + kv_cache.set_base_page_index_buf(i, pool_idx, + memoryview(buffer.numpy())) + + def _resume_and_restore(self, req_id: int, kv_cache) -> bool: + """Resume a suspended KV cache and restore its page index buffers. + + Returns True if the cache is (or becomes) active, False on failure. + """ + if kv_cache.is_active: + return True + if not kv_cache.resume(self._stream.cuda_stream): + return False + self._restore_page_index_bufs(req_id, kv_cache) + return True + + def prepare_context(self, req: LlmRequest) -> bool: + """Create _KVCache, handle block reuse, and resume. Does NOT resize. + + For first chunk: creates _KVCache (with block reuse lookup if enabled), + sets context_current_position, and resumes from suspended state. + For subsequent chunks: verifies existing cache is active. + Returns True on success, False if preparation failed. + """ + if req.is_first_context_chunk: + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None: + # Last token cannot be recovered, so we don't include it in + # the input tokens to look up for the block that can be reused. + if self.enable_block_reuse: + all_tokens = req.get_tokens(DEFAULT_BEAM_INDEX) + tokens = self._augment_tokens_for_block_reuse( + all_tokens, req, end=len(all_tokens) - 1) + else: + tokens = None + kv_cache = self._create_kv_cache(req.py_request_id, + req.lora_task_id, + tokens, + cache_salt=req.cache_salt) + if kv_cache is None: + return False + kv_cache.cuda_stream = self._stream.cuda_stream + + if not self.enable_block_reuse: + kv_cache.stop_committing() + else: + req.context_current_position = kv_cache.num_committed_tokens + req.set_prepopulated_prompt_len(kv_cache.num_committed_tokens, + self.tokens_per_block) + + return self._resume_and_restore(req.py_request_id, kv_cache) + else: + # Subsequent chunk: cache must exist from first chunk. + # It may be suspended (e.g., evicted between chunks), so + # _resume_and_restore handles reactivation. + kv_cache = self.kv_cache_map.get(req.py_request_id) + assert kv_cache is not None, ( + f"KV cache missing for non-first context chunk, request {req.py_request_id}" + ) + return self._resume_and_restore(req.py_request_id, kv_cache) + + def resize_context(self, req: LlmRequest, num_tokens: int) -> bool: + """Resize KV cache to cover context_current_position + num_tokens. + + num_tokens is the number of tokens to be processed (i.e., + context_remaining_length or a chunk thereof). The target capacity is + computed as context_current_position + num_tokens so that block reuse + overlaps with existing capacity are handled correctly. + Returns True on success, False if resize failed (first chunk is + suspended on failure). + + Snapshots the pre-resize capacity on ``req.py_ctx_pre_resize_cap`` + when growth happens so ``revert_allocate_context`` can undo it if + delay batching defers the request. + """ + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None: + return False + + target = req.context_current_position + num_tokens + self.num_extra_kv_tokens + capacity = max(kv_cache.capacity, target) + pre_cap = kv_cache.capacity + + if not kv_cache.resize(capacity): + if req.is_first_context_chunk: + kv_cache.suspend() + return False + + # None means "no growth this iter, nothing to revert"; this also + # invalidates a stale snapshot from a prior iter on the same req. + req.py_ctx_pre_resize_cap = pre_cap if capacity > pre_cap else None + return True + + def extend_capacity_for_tokens(self, request: LlmRequest) -> None: + """Extend KV cache capacity for the CUDA-graph padding delta. + + ``try_allocate_generation`` allocated capacity for the schedule-reduced + draft length. After padding restores ``py_draft_tokens`` to the static + max, we must extend by exactly the difference so that the subsequent + rewind (which operates on the padded length) does not underflow. + + The delta is computed from ``_allocated_draft_lens`` (recorded by + ``try_allocate_generation``) vs the current draft length (post-padding). + """ + allocated = self._allocated_draft_lens.pop(request.py_request_id, None) + if allocated is None: + return + current_draft_len = get_draft_token_length(request) + delta = current_draft_len - allocated + if delta <= 0: + return + kv_cache = self.kv_cache_map[request.py_request_id] + new_capacity = kv_cache.capacity + delta + success = kv_cache.resize(new_capacity) + if not success: + raise ValueError( + f"Failed to extend capacity of KV cache for request " + f"{request.py_request_id} by {delta} tokens " + f"(target capacity {new_capacity})") + + def suspend_request(self, req: LlmRequest) -> None: + """Suspend a request's KV cache (move to host tier).""" + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is not None and kv_cache.is_active: + kv_cache.suspend() + + def resume_request(self, req: LlmRequest) -> bool: + """Resume a previously-suspended KV cache for *req*. + + Returns True if the cache is (or becomes) active on GPU, False if + resume was refused (e.g. GPU pressure above max_util_for_resume) + or no cache exists for the request. + """ + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None: + return False + return self._resume_and_restore(req.py_request_id, kv_cache) + + # ---- prepare_resources ---- + + @nvtx_range("prepare_resources_kv_cache_manager_v2") + def prepare_resources(self, scheduled_batch: ScheduledRequests): + if self.is_draft: + # Draft V2 manager: mirror the main manager by creating/resizing + # KV caches for scheduled requests (the main V2 scheduler does not + # know about the draft manager). + self._prepare_draft_resources(scheduled_batch) + return + + def _prepare_draft_resources(self, scheduled_batch: ScheduledRequests): + """Create/resize KV caches in the draft V2 manager for scheduled requests. + + The main V2 scheduler only manages the primary KV cache manager. + The draft manager must mirror context/generation allocations so that + its IndexMapper contains the correct request IDs for + copy_batch_block_offsets(). + """ + with request_context(True, scheduled_batch): + for req in scheduled_batch.context_requests: + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None: + kv_cache = self._create_kv_cache(req.py_request_id, + req.lora_task_id, + None, + cache_salt=req.cache_salt) + kv_cache.stop_committing() + if not self._resume_and_restore(req.py_request_id, kv_cache): + raise RuntimeError( + f"Failed to resume draft KV cache for request {req.py_request_id}" + ) + draft_len = get_draft_token_length(req) + capacity = (req.context_current_position + + req.context_chunk_size + draft_len + + self.num_extra_kv_tokens) + if not kv_cache.resize(capacity): + raise RuntimeError( + f"Draft KV cache context resize failed for request " + f"{req.py_request_id}: could not resize to {capacity} tokens" + ) + + for req in scheduled_batch.generation_requests: + kv_cache = self.kv_cache_map.get(req.py_request_id) + if kv_cache is None: + raise RuntimeError( + f"Missing draft KV cache for generation request {req.py_request_id}" + ) + if not self._resume_and_restore(req.py_request_id, kv_cache): + raise RuntimeError( + f"Failed to resume draft KV cache for request {req.py_request_id}" + ) + new_cap = self._required_gen_capacity(req, kv_cache.capacity) + # Pad the resize up to _kv_reserve_draft_tokens (see __init__); + # no-op when reserve == draft_token_length. + reserve_slack = (self._kv_reserve_draft_tokens - + get_draft_token_length(req)) + if reserve_slack > 0: + new_cap += reserve_slack + if not kv_cache.resize(new_cap): + raise RuntimeError( + f"Draft KV cache generation resize failed for request " + f"{req.py_request_id}: could not resize to {new_cap} tokens" + ) + + def _augment_tokens_for_block_reuse( + self, + tokens: Sequence[int], + req: LlmRequest, + start: int = 0, + end: int | None = None) -> Sequence[TokenIdExt]: + """Augment token sequence with multimodal content digests for block reuse. + + Multimodal placeholder tokens (e.g. image_token_id) share the same ID + regardless of the underlying content. This method replaces each + multimodal token region with TokenIdExt values produced by + gen_multimodal_cache_key_tokens(), embedding the content digest + (Blake3 hash) into the token sequence so that the radix tree can + distinguish blocks belonging to different images/videos. + + When *start*/*end* are given, they define the chunk bounds; only + `tokens[start:end]` is materialized and returned. This avoids + re-augmenting the full prompt on every chunk during chunked prefill. + + For text-only requests this is a no-op. + """ + if end is None: + end = len(tokens) + chunk_start = start + chunk_end = end + is_sliced = chunk_start != 0 or chunk_end != len(tokens) + + if (req.multimodal_hashes is None or req.multimodal_positions is None + or req.multimodal_lengths is None): + return tokens[chunk_start:chunk_end] if is_sliced else tokens + + result: list[TokenIdExt] = list(tokens[chunk_start:chunk_end]) + run_metadata = _resolve_multimodal_run_metadata(req) + if run_metadata is not None: + return _augment_tokens_with_mm_run_metadata(self.vocab_size, result, + req.multimodal_hashes, + run_metadata, + chunk_start, chunk_end) + + return _augment_tokens_with_contiguous_mm_metadata( + self.vocab_size, result, req.multimodal_hashes, + req.multimodal_positions, req.multimodal_lengths, chunk_start, + chunk_end) + + def get_kv_cache_stats(self): + kv_cache_stats = KvCacheStats() + kv_cache_stats.allocated_bytes = self.impl.get_quota(GPU_LEVEL) + + return kv_cache_stats + + def get_iteration_stats(self): + """V2 does not support per-iteration stats yet.""" + return None + + def get_block_ids_per_seq(self, request_ids: List[int]) -> torch.Tensor: + block_ids_per_seq = self.get_batch_cache_indices(request_ids) + block_ids_per_seq_tensors = [ + torch.tensor([ + i // self.num_local_layers if i != BAD_PAGE_INDEX else 0 + for i in sublist + ], + dtype=torch.int) for sublist in block_ids_per_seq + ] + padded_tensor = torch.nn.utils.rnn.pad_sequence( + block_ids_per_seq_tensors, batch_first=True, padding_value=0) + return padded_tensor + + def add_dummy_requests( + self, + request_ids: List[int], + # Note that token_nums should be past_kv_len + input_len (without + # spec decoding). The draft tokens will be added in this function, + # so we don't need to take care of it in the caller. When preparing + # token_nums, we should not take the draft tokens into account, so + # don't use the kv_cache_manager.max_seq_len, which includes both + # extra tokens and draft tokens. + token_nums: Optional[List[int]] = None, + is_gen: bool = False, + prepare_resource: bool = True, + max_num_draft_tokens: int = 0, + kv_reserve_draft_tokens: Optional[int] = None, + use_mrope: bool = False, + max_beam_width: int = 1, + num_extra_decoding_steps: int = 0, + draft_kv_cache_manager: Optional['BaseResourceManager'] = None): + _kv_draft = kv_reserve_draft_tokens if kv_reserve_draft_tokens is not None else max_num_draft_tokens + + beam_width = max_beam_width + requests = [] + + def release_resources(current_request: LlmRequest, + free_draft_resources: bool = False) -> None: + for req in requests: + self.free_resources(req) + self.free_resources(current_request) + if draft_kv_cache_manager is not None: + for req in requests: + draft_kv_cache_manager.free_resources(req) + if free_draft_resources: + draft_kv_cache_manager.free_resources(current_request) + + for i, req_id in enumerate(request_ids): + # exact choice of n can be ignored for dummy requests + sampling_params = SamplingParams(n=beam_width, + best_of=beam_width, + use_beam_search=beam_width > 1) + # Here 1+max_num_draft_tokens is used to extend the prompt length to + # a non-zero number to skip illegal memory access issue in MLA kernel + # during warmup. + token_num = token_nums[ + i] if token_nums is not None else 1 + max_num_draft_tokens + # token_num - 1 is the past history length in generation. + history_hint = max(0, token_num - 1) if is_gen else None + # TODO: support cross attention + encoder_input_tokens = None + # Using 1 instead of 0 prevents NaN during warmup in e.g. Deepseek + input_tokens = [1 for _ in range(token_num)] + req = LlmRequest(request_id=req_id, + max_new_tokens=1, + input_tokens=input_tokens, + sampling_config=SamplingConfig( + sampling_params._get_sampling_config()), + is_streaming=False, + encoder_input_tokens=encoder_input_tokens) + req.is_dummy_request = True + req.paged_kv_block_ids = [] + if prepare_resource: + # Dummy/warmup request. ``stop_committing()`` below blocks all + # writes to the radix tree, so the choice of branch does not + # affect committed state. ``cache_salt`` is left defaulted + # to None to avoid coupling synthetic data to any salted branch. + kv_cache = self._create_kv_cache(req.py_request_id, + req.lora_task_id, input_tokens) + assert kv_cache.num_committed_tokens == 0 + success = kv_cache.resume(self._stream.cuda_stream) + if not success: + release_resources(req) + return None + kv_cache.stop_committing() + dummy_capacity = token_num + self.num_extra_kv_tokens + num_extra_decoding_steps + # Need to hint the committed history to activate stale-block + # optimization and match the solver's pool budget. + success = kv_cache.resize(dummy_capacity, + history_length=history_hint) + if not success: + release_resources(req) + return None + draft_kv_cache = None + if draft_kv_cache_manager is not None: + draft_kv_cache = draft_kv_cache_manager._create_kv_cache( + req.py_request_id, req.lora_task_id, input_tokens) + # Dummy path: see comment above, no salt. + success = draft_kv_cache.resume( + draft_kv_cache_manager._stream.cuda_stream) + if not success: + release_resources(req, free_draft_resources=True) + return None + draft_kv_cache.stop_committing() + success = draft_kv_cache.resize(dummy_capacity) + if not success: + release_resources(req, free_draft_resources=True) + return None + + if is_gen: + req.state = LlmRequestState.GENERATION_IN_PROGRESS + req.prompt_len = token_num - 1 + req.py_prompt_len = req.prompt_len + req.py_draft_tokens = [1] * max_num_draft_tokens + if prepare_resource: + new_capacity = kv_cache.capacity + _kv_draft + 1 + success = kv_cache.resize(new_capacity, + history_length=history_hint) + if not success: + release_resources(req, + free_draft_resources=draft_kv_cache + is not None) + return None + if draft_kv_cache is not None: + success = draft_kv_cache.resize(new_capacity) + if not success: + release_resources(req, free_draft_resources=True) + return None + + if use_mrope: + _populate_dummy_mrope_config(req, token_num, is_gen) + requests.append(req) + + return requests + + def try_commit_blocks_for_reuse(self, request: LlmRequest, + kv_cache) -> None: + if (self.enable_block_reuse and not self.is_draft + and not request.is_dummy_request + and request.context_current_position + > kv_cache.num_committed_tokens): + tokens = self._augment_tokens_for_block_reuse( + request.get_tokens(DEFAULT_BEAM_INDEX), + request, + start=kv_cache.num_committed_tokens, + end=request.context_current_position) + kv_cache.commit(tokens) + kv_cache.stop_committing() + + def release_index_slot(self, request_id: int) -> None: + """Release IndexMapper slot early while keeping KV cache blocks allocated. + + After prefill completes on a context-only worker, the IndexMapper slot + (used for host_kv_cache_block_offsets during model forward) is no longer + needed. Releasing it early allows new requests to be scheduled while + the KV cache blocks are still being transferred via NIXL/UCX. + """ + self.index_mapper.remove_sequence(request_id) + self._early_freed_index_requests.add(request_id) + + def free_resources(self, request: LlmRequest, pin_on_release: bool = False): + self._allocated_draft_lens.pop(request.py_request_id, None) + kv_cache = self.kv_cache_map.pop(request.py_request_id, None) + if kv_cache is None: + return + self.try_commit_blocks_for_reuse(request, kv_cache) + kv_cache.close() + if request.py_request_id in self._early_freed_index_requests: + self._early_freed_index_requests.discard(request.py_request_id) + else: + self.index_mapper.remove_sequence(request.py_request_id) + + def get_batch_cache_indices( + self, + request_ids: List[int], + layer_idx: Optional[int] = None) -> List[List[int]]: + if layer_idx is None: + pool_id = 0 + else: + pool_id = self.layer_to_pool_mapping_dict[ + self.layer_offsets[layer_idx]] + return self._get_batch_cache_indices_by_pool_id(request_ids, + pool_id=pool_id, + is_kv_aggregate=True) + + def _get_batch_cache_indices_by_pool_id( + self, + request_ids: List[int], + *, + pool_id: int = 0, + is_kv_aggregate: bool = True) -> List[List[int]]: + + if is_kv_aggregate: + # Div by kv_factor to index kv cache with size [num_blocks, kv_factor, tokens_per_block, num_kv_heads, head_dim] + div_factor = self.kv_factor + else: + div_factor = 1 + + res = [] + + for req_id in request_ids: + idx_tensor = torch.as_tensor( + self.kv_cache_map[req_id].get_base_page_indices(pool_id)) + res.append((torch.where( + idx_tensor != BAD_PAGE_INDEX, + idx_tensor * self.index_scales[pool_id] // div_factor, + BAD_PAGE_INDEX)).tolist()) + + return res + + def get_cache_bytes_per_token(self) -> int: + data_roles = [Role.KEY] + if self.kv_cache_type != CacheTypeCpp.SELFKONLY: + data_roles.append(Role.VALUE) + if self.dtype == DataType.NVFP4: + data_roles.append(Role.KEY_BLOCK_SCALE) + if self.kv_cache_type != CacheTypeCpp.SELFKONLY: + data_roles.append(Role.VALUE_BLOCK_SCALE) + + return sum( + self.get_layer_bytes_per_token(local_layer_idx=local_layer_idx, + data_role=data_role) + for local_layer_idx in range(self.num_local_layers) + for data_role in data_roles) + + def get_layer_bytes_per_token(self, local_layer_idx: int, data_role: Role): + if self.dtype not in ( + DataType.FP8, + DataType.HALF, + DataType.BF16, + DataType.FLOAT, + DataType.NVFP4, + ): + raise ValueError(f"Cannot support {self.dtype} KV cache.") + + if data_role == Role.ALL: + kv_factor = self.kv_factor + elif data_role in [ + Role.KEY, Role.VALUE, Role.KEY_BLOCK_SCALE, + Role.VALUE_BLOCK_SCALE + ]: + if data_role in [Role.KEY_BLOCK_SCALE, Role.VALUE_BLOCK_SCALE]: + assert self.dtype == DataType.NVFP4, "NVFP4 is the only supported dtype for block quant data roles" + if data_role == Role.VALUE: + assert self.kv_cache_type != CacheTypeCpp.SELFKONLY, "VALUE data role is not supported for SELFKONLY cache type" + kv_factor = 1 + else: + raise ValueError(f"Invalid data role: {data_role}") + + cache_size_per_token = kv_factor * self.num_kv_heads_per_layer[ + local_layer_idx] * self.head_dim_per_layer[local_layer_idx] + + cache_size_bytes_per_token = get_size_in_bytes(cache_size_per_token, + self.dtype) + + if data_role in [Role.KEY, Role.VALUE]: + return cache_size_bytes_per_token + + quant_size_per_token = 0 + + if self.dtype == DataType.NVFP4: + quant_size_per_token = self.calculate_scaling_factor_size_bytes( + cache_size_per_token, + quant_vector_size=16, + scaling_factor_dtype=DataType.FP8, + ) + + if data_role in [Role.KEY_BLOCK_SCALE, Role.VALUE_BLOCK_SCALE]: + return quant_size_per_token + + # Role.ALL combines both + return cache_size_bytes_per_token + quant_size_per_token + + @staticmethod + def calculate_scaling_factor_size_bytes( + cache_size: int, quant_vector_size: int, + scaling_factor_dtype: DataType) -> int: + assert cache_size % quant_vector_size == 0, "NVFP4 cache size must be divisible by quant vector size" + return get_size_in_bytes(cache_size // quant_vector_size, + scaling_factor_dtype) + + def check_invalid_values_in_kv_cache(self, + fill_with_zero: bool = False) -> bool: + some_checks_unavailable = False + has_invalid_values = torch.tensor([False], + dtype=torch.bool, + device=torch.cuda.current_device()) + pool_handled = set() + + # Handle each layer from start to end to traverse the whole KV cache. + for layer_id, layer_offset in self.layer_offsets.items(): + pool_id = self.layer_to_pool_mapping_dict[layer_offset] + if pool_id in pool_handled: + continue + buffer = self.get_buffers(layer_id) + # process in chunks of 256 pages to avoid OoM + for i in range(0, buffer.shape[0], 256): + buffer_slice = buffer[i:i + 256] + try: + has_invalid_values.logical_or_( + torch.isnan(buffer_slice).any()) + has_invalid_values.logical_or_( + torch.isinf(buffer_slice).any()) + except NotImplementedError: + some_checks_unavailable = True + if fill_with_zero: + buffer.zero_() + pool_handled.add(pool_id) + torch.cuda.synchronize() + + if some_checks_unavailable: + logger.warning( + "`torch.isnan` or `torch.isinf` is not implemented for current kv cache dtype, related checks are skipped" + ) + return bool(has_invalid_values) + + def shutdown(self): + for kv_cache in self.kv_cache_map.values(): + kv_cache.close() + self.kv_cache_map.clear() + self.impl.shutdown() + + def get_max_resource_count(self) -> int: + # TODO: implement this + return 1 + + def get_needed_resource_to_completion(self, request: LlmRequest) -> int: + # TODO: implement this + # context_token_count = request.orig_prompt_len + # num_context_blocks = context_token_count // self.tokens_per_block + # remaining_tokens = context_token_count + request.max_new_tokens - num_context_blocks * self.tokens_per_block + # need_blocks = num_context_blocks + math.ceil( + # remaining_tokens / self.tokens_per_block) + # return need_blocks + return 0 + + # TODO: refactor get_cache_size_per_token and get_cache_bytes_per_token to use the same logic + @staticmethod + def get_cache_size_per_token(model_config: ModelConfigPython, + mapping: Mapping, + num_layers: Optional[int] = None, + **kwargs): + # get kv cache dtype bytes + mem_per_token = 2 + quant_config = model_config.quant_config + if quant_config is not None and quant_config.quant_mode.has_fp8_kv_cache( + ): + mem_per_token = 1 + + # get num key value heads + config = model_config.pretrained_config + num_key_value_heads = getattr(config, 'num_key_value_heads', + config.num_attention_heads) + if isinstance(num_key_value_heads, Iterable): + num_key_value_heads = sum(num_key_value_heads) / len( + num_key_value_heads) + + # get head dim + mla = hasattr(config, + "kv_lora_rank") and config.kv_lora_rank is not None + if mla: + head_dim = config.kv_lora_rank + config.qk_rope_head_dim + kv_factor = 1 + else: + tp_size = 1 if mapping.enable_attention_dp else mapping.tp_size + head_dim = getattr(config, "head_dim", None) + if not isinstance(head_dim, int): + head_dim = config.hidden_size // config.num_attention_heads + head_dim = head_dim * num_key_value_heads // tp_size + kv_factor = 2 + + num_attention_layers = KVCacheManager._resolve_num_attention_layers( + model_config, mapping, num_layers) + mem_per_token *= num_attention_layers * head_dim + + # K and V + mem_per_token *= kv_factor + return mem_per_token + + def update_resources(self, + scheduled_batch: ScheduledRequests, + attn_metadata: "AttentionMetadata" = None, + kv_cache_dtype_byte_size: float = None): + if not self.is_draft: + _update_kv_cache_draft_token_location(self, scheduled_batch, + attn_metadata, + kv_cache_dtype_byte_size) + for req in scheduled_batch.context_requests: + if req.py_request_id not in self.kv_cache_map: + continue + kv_cache = self.kv_cache_map[req.py_request_id] + # In the overlap scheduler, iteration N+1's eviction may + # suspend a ctx request's KV cache while iteration N's + # update_resources still needs to process it. Skip the + # resize — the request will be resumed by the scheduler + # on the next iteration. + if not kv_cache.is_active: + continue + if self.enable_block_reuse and not self.is_draft and not req.is_dummy_request: + if req.context_current_position > kv_cache.num_committed_tokens: + tokens = self._augment_tokens_for_block_reuse( + req.get_tokens(DEFAULT_BEAM_INDEX), + req, + start=kv_cache.num_committed_tokens, + end=req.context_current_position) + kv_cache.commit(tokens) + if req.context_remaining_length == 0: + kv_cache.stop_committing() + else: + success = kv_cache.resize(None, req.context_current_position) + if not success: + raise ValueError( + f"Failed to resize history length of KV cache for request {req.py_request_id} to {req.context_current_position} tokens at context update" + ) + + for req in scheduled_batch.generation_requests: + if req.py_request_id not in self.kv_cache_map: + continue + kv_cache = self.kv_cache_map[req.py_request_id] + # In the overlap scheduler, the scheduler for iteration N+1 + # may suspend a gen request's KV cache (via self-eviction or + # victim eviction) while iteration N's update_resources still + # needs to process it. Skip suspended caches — the request + # will be resumed by the scheduler on the next iteration. + if not kv_cache.is_active: + continue + new_capacity = None if req.state in ( + LlmRequestState.GENERATION_COMPLETE, + LlmRequestState.CONTEXT_INIT + ) else kv_cache.capacity - req.py_rewind_len + success = kv_cache.resize(new_capacity, req.max_beam_num_tokens - 1) + if not success: + raise ValueError( + f"Failed to resize KV cache for request {req.py_request_id} to capacity {new_capacity} and history length {req.max_beam_num_tokens - 1} tokens at generation update" + ) + + def copy_batch_block_offsets(self, dst_tensor: torch.Tensor, + request_ids: List[int], beam_width: int, + num_contexts: int, num_seqs: int): + assert beam_width == 1, "beam_width must be 1 for KVCacheManagerV2" + + copy_idx = self.index_mapper.get_copy_index(request_ids, num_contexts, + beam_width) + assert copy_idx.shape[0] == num_seqs + + copy_batch_block_offsets_to_device(self.host_kv_cache_block_offsets, + dst_tensor, copy_idx, + self.index_scales, self.kv_offset, + self._stream.cuda_stream) + + def _create_kv_cache(self, + request_id: int, + lora_task_id: int | None, + input_tokens: Sequence[TokenIdExt] | None, + cache_salt: str | None = None): + assert request_id not in self.kv_cache_map, f"KV cache for request {request_id} already exists" + if self.index_mapper.num_free_slots() == 0: + logger.warning( + "No free IndexMapper slots for request %s " + "(%d/%d slots in use, likely held by DISAGG_GENERATION_TRANS_IN_PROGRESS requests). " + "Skipping KV cache creation; request will retry next iteration.", + request_id, self.index_mapper.size(), self.index_mapper.size()) + return None + # ReuseScope.salt is int|None; derive a deterministic int from the + # cache_salt string so the same string yields the same reuse namespace + # across processes (matches C++ blockKey hashing on cacheSalt). + salt_int = (int.from_bytes( + hashlib.sha256(cache_salt.encode("utf-8")).digest()[:8], "little") + if cache_salt is not None else None) + kv_cache = self.impl.create_kv_cache( + ReuseScope(lora_id=lora_task_id, salt=salt_int), + input_tokens, + ) + self.kv_cache_map[request_id] = kv_cache + index = self.index_mapper.add_new_sequence(request_id) + for i in range(self.max_beam_width): + for pool_idx in range(self.num_pools): + buffer: torch.Tensor = self.host_kv_cache_block_offsets[ + pool_idx, index * self.max_beam_width + i, 0] + kv_cache.set_base_page_index_buf(i, pool_idx, + memoryview(buffer.numpy())) + return kv_cache + + def reset_reuse_state(self): + self.impl.clear_reusable_blocks() + class SlotManager: diff --git a/tests/unittest/_torch/attention/test_flashinfer_attention.py b/tests/unittest/_torch/attention/test_flashinfer_attention.py index 382e11f591ae..d91a45cd3bb0 100644 --- a/tests/unittest/_torch/attention/test_flashinfer_attention.py +++ b/tests/unittest/_torch/attention/test_flashinfer_attention.py @@ -666,3 +666,33 @@ def test_ragged_prefill_no_kv_cache_uses_cudnn_plan(self) -> None: "cudnn", msg="No-KV ragged prefill should use FlashInfer's cudnn backend", ) + + +class TestFlashInferFp4KvGuards(unittest.TestCase): + """Guards that FlashInfer + NVFP4 KV cache rejects unsupported configs at + init time with a clear NotImplementedError, rather than failing mid-forward. + + Phase 1 of NVFP4 KV cache support on the FlashInfer backend only covers + MLA; non-MLA FP4 should error out early pointing users to attn_backend + ``TRTLLM`` or a BF16/FP8 KV cache. + """ + + def test_non_mla_fp4_kv_raises_not_implemented(self): + from tensorrt_llm.models.modeling_utils import QuantConfig + from tensorrt_llm.quantization.mode import QuantAlgo + + quant_config = QuantConfig(kv_cache_quant_algo=QuantAlgo.NVFP4) + with self.assertRaises(NotImplementedError) as ctx: + FlashInferAttention( + layer_idx=0, + num_heads=8, + head_dim=64, + num_kv_heads=8, + quant_config=quant_config, + ) + # Message must point users at the supported alternatives so they can + # self-serve without reading the backend source. + msg = str(ctx.exception) + self.assertIn("NVFP4 KV cache", msg) + self.assertIn("MLA", msg) + self.assertIn("TRTLLM", msg) diff --git a/tests/unittest/_torch/attention/test_fp4_mla_kv.py b/tests/unittest/_torch/attention/test_fp4_mla_kv.py new file mode 100644 index 000000000000..09cb65ec7f63 --- /dev/null +++ b/tests/unittest/_torch/attention/test_fp4_mla_kv.py @@ -0,0 +1,1355 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Roundtrip tests for the FP4 MLA KV-cache kernels. + +Exercises ``scatter_fp4_mla_kv_cache`` and ``get_fp4_mla_decode_cache`` +(plus ``update_hp_kv_for_fp4_mla``) as a pair on a tiny V1 ``KVCacheManager``. +The goal is to catch stride / page-id / SF-layout bugs without standing up a +real model or FlashInfer wrapper. +""" + +import os +from types import SimpleNamespace + +import pytest +import torch + +import tensorrt_llm +from tensorrt_llm._torch.attention_backend.fp4_mla_kv import ( + FLASHINFER_FP4_MLA_ATTENTION_ENV, + FP4_BLOCK_SIZE, + FP4_MLA_KV_GLOBAL_SCALE, + FP4_MLA_P_GLOBAL_SCALE, + FP4_MLA_Q_RESIDUAL_DIM, + FP4_MLA_TOKENS_PER_BLOCK, + HP_BLOCK_SIZE, + get_fp4_mla_decode_cache, + get_fp4_mla_v_scale_pool_shape, + get_fp4_mla_v_scale_pool_size, + get_fp4_mla_v_scale_pool_view, + is_flashinfer_fp4_mla_attention_enabled, + run_fp4_mla_attention_decode, + scatter_fp4_mla_kv_cache, + update_hp_kv_for_fp4_mla, +) +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm.bindings.executor import KvCacheConfig +from tensorrt_llm.mapping import Mapping + +_DataType = tensorrt_llm.bindings.DataType +_CacheType = tensorrt_llm.bindings.internal.batch_manager.CacheType +_TEST_GLOBAL_SCALE = FP4_MLA_KV_GLOBAL_SCALE + + +def _swizzled_sf_offset(row_idx: int, col_idx: int, sf_per_token: int) -> int: + padded_cols = ((sf_per_token + 3) // 4) * 4 + return ( + col_idx % 4 + + (col_idx // 4) * (4 * 128) + + (row_idx % 32) * 16 + + ((row_idx % 128) // 32) * 4 + + (row_idx // 128) * (128 * padded_cols) + ) + + +def _is_pre_blackwell() -> bool: + return not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] < 10 + + +def _dequant_fp4_swizzled( + fp4_tensor: torch.Tensor, + sf_tensor: torch.Tensor, + *, + logical_dim: int, + sf_per_token: int, + global_scale: float, +) -> torch.Tensor: + fp4_values = torch.tensor( + [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], + dtype=torch.float32, + device=fp4_tensor.device, + ) + fp4_bytes = fp4_tensor.view(torch.uint8) + sf_flat = sf_tensor.view(torch.float8_e4m3fn).reshape(-1) + out = torch.empty( + (fp4_bytes.shape[0], logical_dim), + dtype=torch.float32, + device=fp4_tensor.device, + ) + + for row_idx in range(fp4_bytes.shape[0]): + for sf_col in range(sf_per_token): + start = sf_col * 16 + packed = fp4_bytes[row_idx, start // 2 : start // 2 + 8] + low = packed & 0x0F + high = (packed >> 4) & 0x0F + vals = torch.empty(16, dtype=torch.float32, device=fp4_tensor.device) + low_sign = torch.where( + (low & 0x08) != 0, + -torch.ones_like(low, dtype=torch.float32), + torch.ones_like(low, dtype=torch.float32), + ) + high_sign = torch.where( + (high & 0x08) != 0, + -torch.ones_like(high, dtype=torch.float32), + torch.ones_like(high, dtype=torch.float32), + ) + vals[0::2] = fp4_values[(low & 0x07).long()] * low_sign + vals[1::2] = fp4_values[(high & 0x07).long()] * high_sign + sf_offset = _swizzled_sf_offset(row_idx, sf_col, sf_per_token) + out[row_idx, start : start + 16] = vals * sf_flat[sf_offset].float() / global_scale + + return out + + +def _duplicate_tail_groups(tensor: torch.Tensor, residual_dim: int) -> torch.Tensor: + prefix = tensor[..., :-residual_dim] + tail = tensor[..., -residual_dim:].reshape(*tensor.shape[:-1], residual_dim // 16, 16) + duplicated_tail = tail.repeat_interleave(2, dim=-2).reshape( + *tensor.shape[:-1], + residual_dim * 2, + ) + return torch.cat((prefix, duplicated_tail), dim=-1) + + +def test_flashinfer_fp4_mla_attention_env(monkeypatch): + monkeypatch.delenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, raising=False) + assert not is_flashinfer_fp4_mla_attention_enabled() + + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "on") + assert is_flashinfer_fp4_mla_attention_enabled() + + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "0") + assert not is_flashinfer_fp4_mla_attention_enabled() + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_fp4_mla_v_scale_pool_view_shape(): + device = torch.device("cuda") + num_layers = 2 + num_pages = 3 + kv_lora_rank = 512 + qk_rope_head_dim = 64 + head_dim = kv_lora_rank + qk_rope_head_dim + page_size = FP4_MLA_TOKENS_PER_BLOCK + + allocated_page_elems = get_fp4_mla_v_scale_pool_size(head_dim, page_size) + pool = torch.empty( + (num_layers, num_pages, allocated_page_elems), + dtype=torch.float8_e4m3fn, + device=device, + ) + metadata = SimpleNamespace(page_size=page_size, fp4_mla_v_scale_pool=pool) + + view = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=kv_lora_rank) + + assert tuple(view.shape) == get_fp4_mla_v_scale_pool_shape( + num_layers, num_pages, kv_lora_rank, page_size + ) + assert view.data_ptr() == pool.data_ptr() + assert view.numel() == num_layers * num_pages * kv_lora_rank * (page_size // 16) + + +def _build_metadata(kv_cache_manager, *, num_tokens, page_size, num_layers): + """Build a minimal metadata namespace that satisfies the kernels' field + expectations. Single sequence, single layer slice, no draft tokens.""" + device = torch.device("cuda") + num_blocks = (num_tokens + page_size - 1) // page_size + + # Sequence 0 owns pages block_ids[:num_blocks] in the cache pool. + block_ids = kv_cache_manager.get_batch_cache_indices([0])[0][:num_blocks] + paged_kv_indices = torch.tensor(block_ids, dtype=torch.int32, device=device) + paged_kv_indptr = torch.tensor([0, num_blocks], dtype=torch.int32, device=device) + # Single-sequence decode: compact page range is [0, num_blocks). + paged_kv_indptr_decode = torch.tensor([0, num_blocks], dtype=torch.int32, device=device) + batch_indices = torch.zeros(num_tokens, dtype=torch.int32, device=device) + positions = torch.arange(num_tokens, dtype=torch.int32, device=device) + + # HP pool: one sequence slot, HP_BLOCK_SIZE * head_dim BF16 values per + # (seq_slot, layer) cell. + head_dim = kv_cache_manager.head_dim + hp_pool = torch.zeros( + (1, num_layers, 1, HP_BLOCK_SIZE * head_dim), + dtype=torch.bfloat16, + device=device, + ) + + seq_slots = torch.zeros(1, dtype=torch.int32, device=device) + seq_slots_cpu = torch.zeros(1, dtype=torch.int32, device="cpu") + kv_lens = torch.tensor([num_tokens], dtype=torch.int32, device=device) + prompt_lens_cuda = torch.tensor([num_tokens], dtype=torch.int32, device=device) + prompt_lens_cpu = torch.tensor([num_tokens], dtype=torch.int32) + global_scale = torch.tensor([_TEST_GLOBAL_SCALE], dtype=torch.float32, device=device) + + return SimpleNamespace( + kv_cache_manager=kv_cache_manager, + batch_indices=batch_indices, + positions=positions, + paged_kv_indices=paged_kv_indices, + paged_kv_indptr=paged_kv_indptr, + paged_kv_indptr_decode=paged_kv_indptr_decode, + page_size=page_size, + num_context_blocks=0, + num_generation_blocks=num_blocks, + num_contexts=1, + num_seqs=1, + high_precision_kv_pool=hp_pool, + fp4_mla_v_scale_pool=kv_cache_manager.get_mla_v_scale_pool(), + hp_pool_owners={}, + seq_slots=seq_slots, + seq_slots_cpu=seq_slots_cpu, + kv_lens_cuda_runtime=kv_lens, + prompt_lens_cuda_runtime=prompt_lens_cuda, + prompt_lens_cpu_runtime=prompt_lens_cpu, + _fp4_mla_global_scale=global_scale, + request_ids=[0], + is_cuda_graph=False, + is_warmup=False, + ) + + +def _build_multi_seq_metadata(kv_cache_manager, *, seq_lens, page_size, num_layers): + device = torch.device("cuda") + num_seqs = len(seq_lens) + request_ids = list(range(num_seqs)) + block_ids_per_seq = kv_cache_manager.get_batch_cache_indices(request_ids) + num_blocks = [(seq_len + page_size - 1) // page_size for seq_len in seq_lens] + + paged_kv_indices = torch.tensor( + [ + block_id + for seq_idx, seq_blocks in enumerate(block_ids_per_seq) + for block_id in seq_blocks[: num_blocks[seq_idx]] + ], + dtype=torch.int32, + device=device, + ) + indptr = [0] + for block_count in num_blocks: + indptr.append(indptr[-1] + block_count) + paged_kv_indptr = torch.tensor(indptr, dtype=torch.int32, device=device) + + batch_indices = torch.cat( + [ + torch.full((seq_len,), seq_idx, dtype=torch.int32, device=device) + for seq_idx, seq_len in enumerate(seq_lens) + ] + ) + positions = torch.cat( + [torch.arange(seq_len, dtype=torch.int32, device=device) for seq_len in seq_lens] + ) + + head_dim = kv_cache_manager.head_dim + hp_pool = torch.zeros( + (num_seqs, num_layers, 1, HP_BLOCK_SIZE * head_dim), + dtype=torch.bfloat16, + device=device, + ) + seq_slots = torch.arange(num_seqs, dtype=torch.int32, device=device) + seq_slots_cpu = torch.arange(num_seqs, dtype=torch.int32, device="cpu") + kv_lens = torch.tensor(seq_lens, dtype=torch.int32, device=device) + prompt_lens_cuda = torch.tensor(seq_lens, dtype=torch.int32, device=device) + prompt_lens_cpu = torch.tensor(seq_lens, dtype=torch.int32) + global_scale = torch.tensor([_TEST_GLOBAL_SCALE], dtype=torch.float32, device=device) + + return SimpleNamespace( + kv_cache_manager=kv_cache_manager, + batch_indices=batch_indices, + positions=positions, + paged_kv_indices=paged_kv_indices, + paged_kv_indptr=paged_kv_indptr, + paged_kv_indptr_decode=paged_kv_indptr.clone(), + page_size=page_size, + num_context_blocks=0, + num_generation_blocks=sum(num_blocks), + num_contexts=num_seqs, + num_seqs=num_seqs, + num_blocks=num_blocks, + high_precision_kv_pool=hp_pool, + fp4_mla_v_scale_pool=kv_cache_manager.get_mla_v_scale_pool(), + hp_pool_owners={}, + seq_slots=seq_slots, + seq_slots_cpu=seq_slots_cpu, + kv_lens_cuda_runtime=kv_lens, + prompt_lens_cuda_runtime=prompt_lens_cuda, + prompt_lens_cpu_runtime=prompt_lens_cpu, + _fp4_mla_global_scale=global_scale, + request_ids=request_ids, + is_cuda_graph=False, + is_warmup=False, + ) + + +def _build_fp4_mla_attention_decode_case(*, seq_lens, num_heads, seed): + torch.manual_seed(seed) + device = torch.device("cuda") + + kv_lora_rank = 512 + qk_rope_head_dim = 64 + head_dim = kv_lora_rank + qk_rope_head_dim + page_size = FP4_MLA_TOKENS_PER_BLOCK + num_layers = 1 + num_blocks = [(seq_len + page_size - 1) // page_size for seq_len in seq_lens] + max_seq_len = max(page_size, max(seq_lens)) + max_tokens = max(page_size, sum(num_blocks) * page_size) + + mapping = Mapping(world_size=1, tp_size=1, rank=0) + kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=max_tokens, enable_block_reuse=False), + _CacheType.SELFKONLY, + num_layers=num_layers, + num_kv_heads=1, + head_dim=head_dim, + tokens_per_block=page_size, + max_seq_len=max_seq_len, + max_batch_size=len(seq_lens), + mapping=mapping, + dtype=_DataType.NVFP4, + ) + kv_cache_manager.add_dummy_requests(list(range(len(seq_lens))), seq_lens) + kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() + kv_cache_manager.get_block_scale_buffers(0).zero_() + + metadata = _build_multi_seq_metadata( + kv_cache_manager, + seq_lens=seq_lens, + page_size=page_size, + num_layers=num_layers, + ) + assert metadata.fp4_mla_v_scale_pool is not None + + latent = ( + torch.randn(sum(seq_lens), head_dim, dtype=torch.bfloat16, device=device) * 0.25 + ).clamp_(-1.0, 1.0) + scatter_fp4_mla_kv_cache( + metadata, + latent, + layer_idx=0, + token_offset=0, + phase="context", + local_layer=0, + v_head_dim=kv_lora_rank, + ) + torch.cuda.synchronize() + + metadata.num_contexts = 0 + q_nope = ( + torch.randn(len(seq_lens), num_heads, kv_lora_rank, dtype=torch.bfloat16, device=device) + * 0.25 + ).clamp_(-1.0, 1.0) + q_pe = ( + torch.randn( + len(seq_lens), + num_heads, + qk_rope_head_dim, + dtype=torch.bfloat16, + device=device, + ) + * 0.25 + ).clamp_(-1.0, 1.0) + + return kv_cache_manager, metadata, q_nope, q_pe, kv_lora_rank, qk_rope_head_dim + + +def _fp4_mla_attention_decode_reference( + metadata, + q_nope, + q_pe, + *, + sm_scale, + kv_lora_rank, + qk_rope_head_dim, +): + head_dim = kv_lora_rank + qk_rope_head_dim + high_precision_kv_pool = metadata.high_precision_kv_pool + metadata.high_precision_kv_pool = None + try: + dequant_cache = get_fp4_mla_decode_cache( + metadata, + layer_idx=0, + local_layer=0, + head_dim=head_dim, + dtype=torch.bfloat16, + ) + finally: + metadata.high_precision_kv_pool = high_precision_kv_pool + + num_heads = q_nope.shape[1] + global_scale = metadata._fp4_mla_global_scale + q_full = torch.cat((q_nope, q_pe), dim=-1).reshape(-1, head_dim) + q_fp4, q_sf = torch.ops.trtllm.fp4_quantize_with_residual( + q_full, + global_scale, + FP4_MLA_Q_RESIDUAL_DIM, + is_act=True, + ) + q_logical_dim = head_dim + FP4_MLA_Q_RESIDUAL_DIM + q_dequant = _dequant_fp4_swizzled( + q_fp4, + q_sf.view(torch.float8_e4m3fn), + logical_dim=q_logical_dim, + sf_per_token=q_logical_dim // 16, + global_scale=_TEST_GLOBAL_SCALE, + ) + + p_dequant = _dequant_fp4_swizzled( + metadata._fp4_mla_attention_p_buf, + metadata._fp4_mla_attention_p_sf_buf, + logical_dim=metadata.page_size, + sf_per_token=metadata.page_size // 16, + global_scale=FP4_MLA_P_GLOBAL_SCALE, + ) + + indptr = metadata.paged_kv_indptr_decode.cpu().tolist() + kv_lens = metadata.kv_lens_cuda_runtime.cpu().tolist() + outputs = [] + exact_probs = [] + quantized_probs = [] + for seq_idx in range(metadata.num_seqs): + kv_len = kv_lens[seq_idx] + cache = dequant_cache[indptr[seq_idx] : indptr[seq_idx + 1]].reshape(-1, head_dim)[:kv_len] + logical_k = _duplicate_tail_groups(cache.float(), FP4_MLA_Q_RESIDUAL_DIM) + q_start = seq_idx * num_heads + q = q_dequant[q_start : q_start + num_heads] + probs = torch.softmax(torch.matmul(q, logical_k.transpose(0, 1)) * sm_scale, dim=-1) + + p_pages = [] + for page_rel in range(indptr[seq_idx + 1] - indptr[seq_idx]): + page_start = page_rel * metadata.page_size + valid_tokens = max(min(kv_len - page_start, metadata.page_size), 0) + if valid_tokens == 0: + continue + compact_page = indptr[seq_idx] + page_rel + p_start = compact_page * num_heads + p_pages.append(p_dequant[p_start : p_start + num_heads, :valid_tokens]) + p = torch.cat(p_pages, dim=-1) + + exact_probs.append(probs) + quantized_probs.append(p) + outputs.append(torch.matmul(probs, cache[:, :kv_lora_rank].float())) + return torch.stack(outputs, dim=0), exact_probs, quantized_probs + + +def _cuda_event_benchmark(fn, *, warmup_iters=10, iters=100): + for _ in range(warmup_iters): + fn() + torch.cuda.synchronize() + + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(iters): + fn() + end.record() + torch.cuda.synchronize() + return start.elapsed_time(end) / iters + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize( + ("num_tokens", "page_size"), + [(20, 16), (32, 16), (17, 16), (129, 128), (144, 128)], + ids=[ + "tail4_page16", + "aligned32_page16", + "tail1_page16", + "tail1_page128", + "aligned144_page128", + ], +) +def test_fp4_mla_scatter_gather_roundtrip(num_tokens: int, page_size: int, monkeypatch): + """Write BF16 latent through scatter + HP update, then read via the + dequant-gather + HP-overlay path and verify the two halves of the output: + + * Positions that fall in the FP4 region (before the last ``kv_len % 16`` + tokens) must match the input up to NVFP4 quant error. + * Positions covered by the HP overlay (the last ``kv_len % 16`` tokens) + must match the input exactly (BF16 roundtrip). + """ + monkeypatch.delenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, raising=False) + torch.manual_seed(0) + device = torch.device("cuda") + + # MLA shapes (DeepSeek-V3-Lite style, scaled down). + kv_lora_rank = 512 + qk_rope_head_dim = 64 + head_dim = kv_lora_rank + qk_rope_head_dim # 576, divisible by 16. + num_layers = 1 + max_seq_len = max(64, ((num_tokens + page_size - 1) // page_size) * page_size) + max_batch_size = 1 + + mapping = Mapping(world_size=1, tp_size=1, rank=0) + # max_tokens must cover at least ceil(max_seq_len / page_size) pages. + kv_cache_config = KvCacheConfig(max_tokens=max_seq_len, enable_block_reuse=False) + kv_cache_manager = KVCacheManager( + kv_cache_config, + _CacheType.SELFKONLY, + num_layers=num_layers, + num_kv_heads=1, + head_dim=head_dim, + tokens_per_block=page_size, + max_seq_len=max_seq_len, + max_batch_size=max_batch_size, + mapping=mapping, + dtype=_DataType.NVFP4, + ) + try: + kv_cache_manager.add_dummy_requests([0], [num_tokens]) + + # Zero the underlying data + scale pools so stale bytes can't mask bugs. + data_buf = kv_cache_manager.get_buffers(0).view(torch.uint8) + data_buf.zero_() + sf_buf = kv_cache_manager.get_block_scale_buffers(0) + assert sf_buf is not None, "V1 NVFP4 manager must expose block scales" + sf_buf.zero_() + + metadata = _build_metadata( + kv_cache_manager, num_tokens=num_tokens, page_size=page_size, num_layers=num_layers + ) + + # Stay inside the configured global-scale range so FP4 quantization + # does not saturate; narrower latents keep the FP4 tolerance reasonable. + latent = ( + torch.randn(num_tokens, head_dim, dtype=torch.bfloat16, device=device) * 1.5 + ).clamp_(-5.0, 5.0) + + # --- Write path ------------------------------------------------- + scatter_fp4_mla_kv_cache(metadata, latent, layer_idx=0, token_offset=0) + update_hp_kv_for_fp4_mla(metadata, latent, local_layer=0, phase="context") + + # Switch the metadata into "decode" shape for the read path. The + # decode kernels gather the entire 0..num_tokens range (as if every + # token were in the KV history for an upcoming decode step). + metadata.num_contexts = 0 + metadata.num_seqs = 1 + + # --- Read path -------------------------------------------------- + combined = get_fp4_mla_decode_cache( + metadata, + layer_idx=0, + local_layer=0, + head_dim=head_dim, + dtype=torch.bfloat16, + ) + # combined shape: [num_blocks, page_size, head_dim]. + flat = combined.reshape(-1, head_dim)[:num_tokens] + + # --- Assertions ------------------------------------------------- + tail = num_tokens % HP_BLOCK_SIZE + fp4_end = num_tokens - tail # exclusive + + # FP4-dequantized region: allow NVFP4 quant error. With unit global + # scale, worst-case absolute error is ~0.5 of the largest FP4 step + # within the value's block; 1.0 is a safe bound for the clamped + # latent range [-5, 5]. + if fp4_end > 0: + torch.testing.assert_close( + flat[:fp4_end].float(), + latent[:fp4_end].float(), + atol=1.0, + rtol=0.5, + msg=f"FP4 region mismatch for num_tokens={num_tokens}", + ) + + # HP overlay region: must be exact BF16 roundtrip. + if tail > 0: + torch.testing.assert_close( + flat[fp4_end:].float(), + latent[fp4_end:].float(), + atol=0.0, + rtol=0.0, + msg=f"HP overlay mismatch for num_tokens={num_tokens}", + ) + finally: + kv_cache_manager.shutdown() + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize( + ("page_size", "ctx_tokens"), + [(16, 32), (128, 128)], + ids=["page16", "page128"], +) +def test_fp4_mla_hp_overlay_generation_phase(page_size: int, ctx_tokens: int, monkeypatch): + """After an aligned context, perform one decode step and verify that the + decode token surfaces through the HP overlay at the first tail slot.""" + monkeypatch.delenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, raising=False) + torch.manual_seed(1) + device = torch.device("cuda") + + kv_lora_rank = 512 + qk_rope_head_dim = 64 + head_dim = kv_lora_rank + qk_rope_head_dim + num_layers = 1 + max_seq_len = max(64, page_size * 2) + # Aligned context -> no tail until the first decode token. + total_tokens = ctx_tokens + 1 + + mapping = Mapping(world_size=1, tp_size=1, rank=0) + kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=max_seq_len, enable_block_reuse=False), + _CacheType.SELFKONLY, + num_layers=num_layers, + num_kv_heads=1, + head_dim=head_dim, + tokens_per_block=page_size, + max_seq_len=max_seq_len, + max_batch_size=1, + mapping=mapping, + dtype=_DataType.NVFP4, + ) + try: + kv_cache_manager.add_dummy_requests([0], [total_tokens]) + + kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() + kv_cache_manager.get_block_scale_buffers(0).zero_() + + # Write the full context (32 tokens) in one scatter, then the single + # decode token separately, matching the production flow. + metadata = _build_metadata( + kv_cache_manager, num_tokens=total_tokens, page_size=page_size, num_layers=num_layers + ) + + ctx_latent = ( + torch.randn(ctx_tokens, head_dim, dtype=torch.bfloat16, device=device) * 1.5 + ).clamp_(-5.0, 5.0) + gen_latent = (torch.randn(1, head_dim, dtype=torch.bfloat16, device=device) * 1.5).clamp_( + -5.0, 5.0 + ) + + # Context scatter + HP update: kv_len temporarily = 32. + metadata.kv_lens_cuda_runtime = torch.tensor([ctx_tokens], dtype=torch.int32, device=device) + metadata.prompt_lens_cuda_runtime = torch.tensor( + [ctx_tokens], dtype=torch.int32, device=device + ) + metadata.prompt_lens_cpu_runtime = torch.tensor([ctx_tokens], dtype=torch.int32) + metadata.num_contexts = 1 + metadata.num_seqs = 1 + # Only ctx_tokens are visible to the scatter kernel this call. + metadata.positions = torch.arange(ctx_tokens, dtype=torch.int32, device=device) + metadata.batch_indices = torch.zeros(ctx_tokens, dtype=torch.int32, device=device) + scatter_fp4_mla_kv_cache(metadata, ctx_latent, layer_idx=0, token_offset=0) + update_hp_kv_for_fp4_mla(metadata, ctx_latent, local_layer=0, phase="context") + + # Decode scatter + HP update: append the single gen token at position 32. + metadata.kv_lens_cuda_runtime = torch.tensor( + [total_tokens], dtype=torch.int32, device=device + ) + metadata.positions = torch.tensor([ctx_tokens], dtype=torch.int32, device=device) + metadata.batch_indices = torch.zeros(1, dtype=torch.int32, device=device) + metadata.num_contexts = 0 + metadata.prompt_lens_cuda_runtime = torch.tensor([1], dtype=torch.int32, device=device) + metadata.prompt_lens_cpu_runtime = torch.tensor([1], dtype=torch.int32) + scatter_fp4_mla_kv_cache(metadata, gen_latent, layer_idx=0, token_offset=0) + update_hp_kv_for_fp4_mla(metadata, gen_latent, local_layer=0, phase="generation") + + # Now read the full 33-token history back. + num_blocks = (total_tokens + page_size - 1) // page_size + metadata.num_generation_blocks = num_blocks + metadata.paged_kv_indices = torch.tensor( + kv_cache_manager.get_batch_cache_indices([0])[0][:num_blocks], + dtype=torch.int32, + device=device, + ) + metadata.paged_kv_indptr_decode = torch.tensor( + [0, num_blocks], dtype=torch.int32, device=device + ) + + combined = get_fp4_mla_decode_cache( + metadata, + layer_idx=0, + local_layer=0, + head_dim=head_dim, + dtype=torch.bfloat16, + ) + flat = combined.reshape(-1, head_dim)[:total_tokens] + + # Position 32 is the lone tail token; HP overlay must return exactly + # gen_latent[0]. + torch.testing.assert_close( + flat[ctx_tokens].float(), + gen_latent[0].float(), + atol=0.0, + rtol=0.0, + msg="HP overlay did not restore the decode token", + ) + + # Positions 0..31 come from FP4 dequant; accept NVFP4 roundtrip noise. + torch.testing.assert_close( + flat[:ctx_tokens].float(), + ctx_latent.float(), + atol=1.0, + rtol=0.5, + msg="FP4 region mismatch on context tokens", + ) + finally: + kv_cache_manager.shutdown() + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("attention_env", ["0", "1"], ids=["linear_sf", "swizzled_sf"]) +def test_fp4_mla_scatter_last_page_no_oob(attention_env: str, monkeypatch): + """Scatter must not write past a page's scale region in either SF layout. + + With ``tokens_per_block=128`` the scatter's ``BLOCK_SF=64`` Triton block + has ``SF_PER_TOKEN=36`` valid lanes plus 28 masked-out lanes. For those + masked lanes the unconstrained offset (linear or swizzled) can exceed + the per-page stride. If masked-lane addresses are not pinned in-bounds, + the last physical page's masked stores fall past the sf_cache allocation, + which crashes with "illegal memory access" on Blackwell. + + The SF layout is gated by ``FLASHINFER_FP4_MLA_ATTENTION_ENV``; cover both + settings so a regression in either path is caught. The test writes tokens + that land exclusively on the LAST physical page and asserts: + 1. No bytes outside the target page are modified. + 2. The data round-trips correctly through the dequant path (so the + valid lanes still wrote the right values). + """ + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, attention_env) + torch.manual_seed(3) + device = torch.device("cuda") + + kv_lora_rank = 512 + qk_rope_head_dim = 64 + head_dim = kv_lora_rank + qk_rope_head_dim + page_size = FP4_MLA_TOKENS_PER_BLOCK + num_pages = 4 + max_seq_len = num_pages * page_size + num_layers = 1 + + mapping = Mapping(world_size=1, tp_size=1, rank=0) + kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=max_seq_len, enable_block_reuse=False), + _CacheType.SELFKONLY, + num_layers=num_layers, + num_kv_heads=1, + head_dim=head_dim, + tokens_per_block=page_size, + max_seq_len=max_seq_len, + max_batch_size=1, + mapping=mapping, + dtype=_DataType.NVFP4, + ) + try: + kv_cache_manager.add_dummy_requests([0], [max_seq_len]) + kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() + sf_buf = kv_cache_manager.get_block_scale_buffers(0) + assert sf_buf is not None + + block_ids = kv_cache_manager.get_batch_cache_indices([0])[0][:num_pages] + last_physical_page = block_ids[-1] + # Sentinel non-target pages so a stray cross-page write is observable. + sf_buf.view(torch.uint8).fill_(0xA5) + sf_buf[last_physical_page].zero_() + kv_cache_manager.get_buffers(0).view(torch.uint8).fill_(0xA5) + kv_cache_manager.get_buffers(0).view(torch.uint8)[last_physical_page].zero_() + snapshot_sf = sf_buf.view(torch.uint8).clone() + snapshot_kv = kv_cache_manager.get_buffers(0).view(torch.uint8).clone() + + # Write tokens that land on the last physical page only. + last_start = (num_pages - 1) * page_size + num_tokens = page_size + paged_kv_indices = torch.tensor(block_ids, dtype=torch.int32, device=device) + paged_kv_indptr = torch.tensor([0, num_pages], dtype=torch.int32, device=device) + batch_indices = torch.zeros(num_tokens, dtype=torch.int32, device=device) + positions = torch.arange( + last_start, last_start + num_tokens, dtype=torch.int32, device=device + ) + + # Single-sequence read-back metadata (one big seq covering all pages). + metadata = _build_metadata( + kv_cache_manager, + num_tokens=max_seq_len, + page_size=page_size, + num_layers=num_layers, + ) + # Override scatter-only fields to write just the last-page slice. + metadata.batch_indices = batch_indices + metadata.positions = positions + metadata.paged_kv_indices = paged_kv_indices + metadata.paged_kv_indptr = paged_kv_indptr + + latent = ( + torch.randn(num_tokens, head_dim, dtype=torch.bfloat16, device=device) * 1.5 + ).clamp_(-5.0, 5.0) + scatter_fp4_mla_kv_cache(metadata, latent, layer_idx=0, token_offset=0) + torch.cuda.synchronize() + + # Bytes for any non-target page must be unchanged. + sf_after = sf_buf.view(torch.uint8) + kv_after = kv_cache_manager.get_buffers(0).view(torch.uint8) + for pid in range(sf_after.shape[0]): + if pid == last_physical_page: + continue + torch.testing.assert_close( + sf_after[pid], + snapshot_sf[pid], + atol=0, + rtol=0, + msg=f"scatter wrote to sf of non-target page {pid}", + ) + torch.testing.assert_close( + kv_after[pid], + snapshot_kv[pid], + atol=0, + rtol=0, + msg=f"scatter wrote to kv of non-target page {pid}", + ) + + # Read back via dequant and verify round-trip correctness. + metadata.num_contexts = 0 + metadata.num_seqs = 1 + combined = get_fp4_mla_decode_cache( + metadata, + layer_idx=0, + local_layer=0, + head_dim=head_dim, + dtype=torch.bfloat16, + ).reshape(-1, head_dim) + recovered = combined[last_start : last_start + num_tokens] + torch.testing.assert_close( + recovered.float(), + latent.float(), + atol=1.0, + rtol=0.5, + msg="last-page FP4 round-trip mismatch", + ) + finally: + kv_cache_manager.shutdown() + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_fp4_mla_dequant_invalid_page_ids_no_oob(monkeypatch): + """Dequant should reject a short page-id slice and guard invalid physical + pages in-kernel instead of doing unchecked page-stride arithmetic.""" + monkeypatch.delenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, raising=False) + device = torch.device("cuda") + + kv_lora_rank = 512 + qk_rope_head_dim = 64 + head_dim = kv_lora_rank + qk_rope_head_dim + page_size = FP4_MLA_TOKENS_PER_BLOCK + num_layers = 1 + + mapping = Mapping(world_size=1, tp_size=1, rank=0) + kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=page_size, enable_block_reuse=False), + _CacheType.SELFKONLY, + num_layers=num_layers, + num_kv_heads=1, + head_dim=head_dim, + tokens_per_block=page_size, + max_seq_len=page_size, + max_batch_size=1, + mapping=mapping, + dtype=_DataType.NVFP4, + ) + try: + kv_cache_manager.add_dummy_requests([0], [page_size]) + kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() + kv_cache_manager.get_block_scale_buffers(0).zero_() + + metadata = _build_metadata( + kv_cache_manager, + num_tokens=page_size, + page_size=page_size, + num_layers=num_layers, + ) + metadata.num_contexts = 0 + metadata.num_seqs = 0 + metadata.num_context_blocks = 0 + metadata.num_generation_blocks = 1 + metadata.paged_kv_indices = torch.empty(0, dtype=torch.int32, device=device) + + with pytest.raises(RuntimeError, match="needs 1 decode page ids"): + get_fp4_mla_decode_cache( + metadata, + layer_idx=0, + local_layer=0, + head_dim=head_dim, + dtype=torch.bfloat16, + ) + + invalid_page = kv_cache_manager.get_buffers(0).shape[0] + metadata.paged_kv_indices = torch.tensor([invalid_page], dtype=torch.int32, device=device) + combined = get_fp4_mla_decode_cache( + metadata, + layer_idx=0, + local_layer=0, + head_dim=head_dim, + dtype=torch.bfloat16, + ) + torch.cuda.synchronize() + torch.testing.assert_close( + combined, + torch.zeros_like(combined), + atol=0, + rtol=0, + ) + finally: + kv_cache_manager.shutdown() + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_fp4_mla_real_scatter_writes_shared_2d_scales(monkeypatch): + """Real FP4 scatter shares 16x16 scales only where K and V overlap.""" + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") + torch.manual_seed(4) + device = torch.device("cuda") + + kv_lora_rank = 512 + qk_rope_head_dim = 64 + head_dim = kv_lora_rank + qk_rope_head_dim + page_size = FP4_MLA_TOKENS_PER_BLOCK + num_tokens = 32 + num_layers = 1 + + mapping = Mapping(world_size=1, tp_size=1, rank=0) + kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=page_size, enable_block_reuse=False), + _CacheType.SELFKONLY, + num_layers=num_layers, + num_kv_heads=1, + head_dim=head_dim, + tokens_per_block=page_size, + max_seq_len=page_size, + max_batch_size=1, + mapping=mapping, + dtype=_DataType.NVFP4, + ) + try: + kv_cache_manager.add_dummy_requests([0], [num_tokens]) + kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() + kv_cache_manager.get_block_scale_buffers(0).zero_() + + metadata = _build_metadata( + kv_cache_manager, + num_tokens=num_tokens, + page_size=page_size, + num_layers=num_layers, + ) + assert metadata.fp4_mla_v_scale_pool is not None + + row = ( + torch.arange(num_tokens, dtype=torch.float32, device=device) % HP_BLOCK_SIZE + 1.0 + ).view(num_tokens, 1) + col = ( + torch.arange(head_dim, dtype=torch.float32, device=device) % FP4_MLA_TOKENS_PER_BLOCK + + 1.0 + ).view(1, head_dim) + latent = (row * col / 512.0).to(torch.bfloat16) + + scatter_fp4_mla_kv_cache( + metadata, + latent, + layer_idx=0, + token_offset=0, + phase="context", + local_layer=0, + v_head_dim=kv_lora_rank, + ) + torch.cuda.synchronize() + + physical_page = kv_cache_manager.get_batch_cache_indices([0])[0][0] + sf_per_token = head_dim // 16 + sf_per_page = page_size // 16 + k_page = ( + kv_cache_manager.get_block_scale_buffers(0) + .view(torch.float8_e4m3fn)[physical_page] + .reshape(-1) + .view(torch.uint8) + ) + v_page = ( + get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=kv_lora_rank)[0, physical_page] + .reshape(-1) + .view(torch.uint8) + ) + + for token_block in range(num_tokens // HP_BLOCK_SIZE): + token_base = token_block * HP_BLOCK_SIZE + for dim_block in (0, 10, 31): + k_offsets = torch.tensor( + [ + _swizzled_sf_offset(row_idx, dim_block, sf_per_token) + for row_idx in range(token_base, token_base + HP_BLOCK_SIZE) + ], + dtype=torch.long, + device=device, + ) + k_bytes = k_page[k_offsets] + assert bool((k_bytes[0] != 0).item()) + torch.testing.assert_close( + k_bytes, + k_bytes[0].expand_as(k_bytes), + atol=0, + rtol=0, + msg=f"K scales are not shared for dim block {dim_block}", + ) + + v_offsets = torch.tensor( + [ + _swizzled_sf_offset(dim_block * 16 + row_idx, token_block, sf_per_page) + for row_idx in range(16) + ], + dtype=torch.long, + device=device, + ) + v_bytes = v_page[v_offsets] + torch.testing.assert_close( + v_bytes, + k_bytes[0].expand_as(v_bytes), + atol=0, + rtol=0, + msg=f"K/V scales disagree for dim block {dim_block}", + ) + + tail_dim_block = kv_lora_rank // FP4_BLOCK_SIZE + tail_offsets = torch.tensor( + [ + _swizzled_sf_offset(row_idx, tail_dim_block, sf_per_token) + for row_idx in range(token_base, token_base + HP_BLOCK_SIZE) + ], + dtype=torch.long, + device=device, + ) + tail_bytes = k_page[tail_offsets] + assert bool((tail_bytes[0] != 0).item()) + assert int(torch.unique(tail_bytes).numel()) > 1, "K-only tail scales must be per-token" + finally: + kv_cache_manager.shutdown() + + +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 tensor cores") +def test_fp4_mla_attention_decode_residual_qk_duplicates_k_tail(monkeypatch): + """Residual-Q QK must use the same cached K tail for main and residual groups.""" + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") + torch.manual_seed(6) + device = torch.device("cuda") + + kv_lora_rank = 512 + qk_rope_head_dim = 64 + head_dim = kv_lora_rank + qk_rope_head_dim + page_size = FP4_MLA_TOKENS_PER_BLOCK + num_tokens = page_size + num_layers = 1 + num_heads = 1 + + mapping = Mapping(world_size=1, tp_size=1, rank=0) + kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=page_size, enable_block_reuse=False), + _CacheType.SELFKONLY, + num_layers=num_layers, + num_kv_heads=1, + head_dim=head_dim, + tokens_per_block=page_size, + max_seq_len=page_size, + max_batch_size=1, + mapping=mapping, + dtype=_DataType.NVFP4, + ) + try: + kv_cache_manager.add_dummy_requests([0], [num_tokens]) + kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() + kv_cache_manager.get_block_scale_buffers(0).zero_() + + metadata = _build_metadata( + kv_cache_manager, + num_tokens=num_tokens, + page_size=page_size, + num_layers=num_layers, + ) + assert metadata.fp4_mla_v_scale_pool is not None + + token_pattern = torch.linspace( + -0.4, + 0.4, + num_tokens, + dtype=torch.float32, + device=device, + ).view(num_tokens, 1) + dim_pattern = torch.linspace( + -0.7, + 0.7, + FP4_MLA_Q_RESIDUAL_DIM, + dtype=torch.float32, + device=device, + ).view(1, FP4_MLA_Q_RESIDUAL_DIM) + latent = torch.zeros(num_tokens, head_dim, dtype=torch.bfloat16, device=device) + latent[:, -FP4_MLA_Q_RESIDUAL_DIM:] = (token_pattern + dim_pattern).to(torch.bfloat16) + + scatter_fp4_mla_kv_cache( + metadata, + latent, + layer_idx=0, + token_offset=0, + phase="context", + local_layer=0, + v_head_dim=kv_lora_rank, + ) + torch.cuda.synchronize() + + metadata.num_contexts = 0 + metadata.num_seqs = 1 + + q_nope = torch.zeros(1, num_heads, kv_lora_rank, dtype=torch.bfloat16, device=device) + q_pe = ( + torch.linspace( + -0.9, + 0.9, + qk_rope_head_dim, + dtype=torch.float32, + device=device, + ) + .view(1, num_heads, qk_rope_head_dim) + .to(torch.bfloat16) + ) + output = torch.empty(1, num_heads, kv_lora_rank, dtype=torch.bfloat16, device=device) + sm_scale = 0.1 + + run_fp4_mla_attention_decode( + metadata, + layer_idx=0, + local_layer=0, + q_nope=q_nope, + q_pe=q_pe, + output=output, + sm_scale=sm_scale, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + ) + torch.cuda.synchronize() + + q_full = torch.cat((q_nope, q_pe), dim=-1).reshape(num_heads, head_dim) + q_fp4, q_sf = torch.ops.trtllm.fp4_quantize_with_residual( + q_full, + metadata._fp4_mla_global_scale, + FP4_MLA_Q_RESIDUAL_DIM, + is_act=True, + ) + q_logical_dim = head_dim + FP4_MLA_Q_RESIDUAL_DIM + q_dequant = _dequant_fp4_swizzled( + q_fp4, + q_sf, + logical_dim=q_logical_dim, + sf_per_token=q_logical_dim // 16, + global_scale=_TEST_GLOBAL_SCALE, + ) + dequant_cache = get_fp4_mla_decode_cache( + metadata, + layer_idx=0, + local_layer=0, + head_dim=head_dim, + dtype=torch.bfloat16, + ).reshape(-1, head_dim)[:num_tokens] + logical_k = _duplicate_tail_groups(dequant_cache.float(), FP4_MLA_Q_RESIDUAL_DIM) + ref_scores = torch.matmul(q_dequant, logical_k.transpose(0, 1)) * sm_scale + ref_probs = torch.softmax(ref_scores, dim=-1) + probs = metadata._fp4_mla_attention_p_prob_buf[:num_heads, :num_tokens] + + torch.testing.assert_close( + probs, + ref_probs, + atol=2e-2, + rtol=2e-2, + msg="FP4 MLA residual-Q probabilities did not match duplicated K-tail reference", + ) + finally: + kv_cache_manager.shutdown() + + +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 tensor cores") +def test_fp4_mla_attention_decode_multi_seq_matches_reference(monkeypatch): + """Multiple decode sequences and heads must match a QK-softmax-PV reference.""" + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") + num_heads = 5 + seq_lens = [32, 128] + + ( + kv_cache_manager, + metadata, + q_nope, + q_pe, + kv_lora_rank, + qk_rope_head_dim, + ) = _build_fp4_mla_attention_decode_case( + seq_lens=seq_lens, + num_heads=num_heads, + seed=7, + ) + try: + output = torch.empty_like(q_nope) + sm_scale = 0.1 + run_fp4_mla_attention_decode( + metadata, + layer_idx=0, + local_layer=0, + q_nope=q_nope, + q_pe=q_pe, + output=output, + sm_scale=sm_scale, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + ) + torch.cuda.synchronize() + + ref_output, exact_probs, quantized_probs = _fp4_mla_attention_decode_reference( + metadata, + q_nope, + q_pe, + sm_scale=sm_scale, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + ) + for seq_idx, (exact_prob, quantized_prob) in enumerate(zip(exact_probs, quantized_probs)): + torch.testing.assert_close( + quantized_prob, + exact_prob, + atol=8e-2, + rtol=8e-2, + msg=f"FP4 MLA attention probabilities diverged for sequence {seq_idx}", + ) + torch.testing.assert_close( + output.float(), + ref_output, + atol=1e-1, + rtol=1e-1, + msg="FP4 MLA attention decode output diverged from the QK-softmax-PV reference", + ) + finally: + kv_cache_manager.shutdown() + + +@pytest.mark.skipif( + os.environ.get("TRTLLM_RUN_FP4_MLA_ATTENTION_BENCHMARK") != "1", + reason=("Manual perf benchmark; set TRTLLM_RUN_FP4_MLA_ATTENTION_BENCHMARK=1 to run"), +) +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 tensor cores") +@pytest.mark.parametrize("batch_size", [16, 32, 64, 128, 256], ids=lambda x: f"bs{x}") +@pytest.mark.parametrize("seq_len", [8192], ids=lambda x: f"seq{x}") +def test_fp4_mla_attention_decode_perf_benchmark( + batch_size, + seq_len, + monkeypatch, +): + """Opt-in microbenchmark for ``run_fp4_mla_attention_decode``. + + Run manually with: + ``TRTLLM_RUN_FP4_MLA_ATTENTION_BENCHMARK=1 pytest -s -k fp4_mla_attention_decode_perf``. + """ + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") + num_heads = 128 + ( + kv_cache_manager, + metadata, + q_nope, + q_pe, + kv_lora_rank, + qk_rope_head_dim, + ) = _build_fp4_mla_attention_decode_case( + seq_lens=[seq_len] * batch_size, + num_heads=num_heads, + seed=8, + ) + try: + output = torch.empty_like(q_nope) + + def run_decode(): + run_fp4_mla_attention_decode( + metadata, + layer_idx=0, + local_layer=0, + q_nope=q_nope, + q_pe=q_pe, + output=output, + sm_scale=0.1, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + ) + + run_decode() + torch.cuda.synchronize() + avg_ms = _cuda_event_benchmark(run_decode, warmup_iters=10, iters=50) + tokens = batch_size * seq_len + qk_dim = kv_lora_rank + qk_rope_head_dim + FP4_MLA_Q_RESIDUAL_DIM + pv_dim = kv_lora_rank + matmul_flops = 2 * batch_size * num_heads * seq_len * (qk_dim + pv_dim) + matmul_tflops = matmul_flops / avg_ms / 1e9 + print( + "\nrun_fp4_mla_attention_decode " + f"batch={batch_size} seq_len={seq_len} heads={num_heads}: " + f"{avg_ms:.4f} ms, {tokens * num_heads / avg_ms / 1e3:.2f} M token-head/s, " + f"{matmul_tflops:.2f} estimated matmul TFLOP/s" + ) + assert torch.isfinite(output.float()).all() + finally: + kv_cache_manager.shutdown() + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_fp4_mla_shared_tile_rejects_unaligned_context_start(monkeypatch): + """The no-dequant FP4 path requires context chunks to start on a 16-token boundary.""" + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") + torch.manual_seed(5) + device = torch.device("cuda") + + kv_lora_rank = 512 + qk_rope_head_dim = 64 + head_dim = kv_lora_rank + qk_rope_head_dim + page_size = FP4_MLA_TOKENS_PER_BLOCK + num_layers = 1 + cached_tokens = 8 + new_tokens = 16 + total_tokens = cached_tokens + new_tokens + + mapping = Mapping(world_size=1, tp_size=1, rank=0) + kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=page_size, enable_block_reuse=False), + _CacheType.SELFKONLY, + num_layers=num_layers, + num_kv_heads=1, + head_dim=head_dim, + tokens_per_block=page_size, + max_seq_len=page_size, + max_batch_size=1, + mapping=mapping, + dtype=_DataType.NVFP4, + ) + try: + kv_cache_manager.add_dummy_requests([0], [total_tokens]) + kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() + kv_cache_manager.get_block_scale_buffers(0).zero_() + + metadata = _build_metadata( + kv_cache_manager, + num_tokens=cached_tokens, + page_size=page_size, + num_layers=num_layers, + ) + assert metadata.fp4_mla_v_scale_pool is not None + + new_latent = ( + torch.randn(new_tokens, head_dim, dtype=torch.bfloat16, device=device) * 1.5 + ).clamp_(-5.0, 5.0) + metadata.kv_lens_cuda_runtime = torch.tensor( + [total_tokens], dtype=torch.int32, device=device + ) + metadata.prompt_lens_cuda_runtime = torch.tensor( + [new_tokens], dtype=torch.int32, device=device + ) + metadata.prompt_lens_cpu_runtime = torch.tensor([new_tokens], dtype=torch.int32) + + with pytest.raises( + ValueError, + match="start position.*16-token aligned", + ): + scatter_fp4_mla_kv_cache( + metadata, + new_latent, + layer_idx=0, + token_offset=0, + phase="context", + local_layer=0, + v_head_dim=kv_lora_rank, + ) + finally: + kv_cache_manager.shutdown() diff --git a/tests/unittest/_torch/executor/test_mla_tokens_per_block.py b/tests/unittest/_torch/executor/test_mla_tokens_per_block.py new file mode 100644 index 000000000000..5a42fe67cd01 --- /dev/null +++ b/tests/unittest/_torch/executor/test_mla_tokens_per_block.py @@ -0,0 +1,102 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace + +from tensorrt_llm._torch.pyexecutor.py_executor_creator import ( + FLASH_MLA_TOKENS_PER_BLOCK, + FLASHINFER_FP4_MLA_ATTENTION_ENV, + FP4_MLA_TOKENS_PER_BLOCK, + _select_mla_tokens_per_block, +) +from tensorrt_llm.quantization import QuantAlgo + + +def _mla_config(): + return SimpleNamespace(kv_lora_rank=512, qk_rope_head_dim=64) + + +def _non_mla_config(): + return SimpleNamespace() + + +def _model_config(kv_cache_quant_algo=None, enable_flash_mla=False): + quant_config = SimpleNamespace(kv_cache_quant_algo=kv_cache_quant_algo) + return SimpleNamespace(quant_config=quant_config, enable_flash_mla=enable_flash_mla) + + +def _kv_cache_config(dtype="auto", tokens_per_block=32): + return SimpleNamespace(dtype=dtype, tokens_per_block=tokens_per_block) + + +def test_non_mla_keeps_configured_tokens_per_block(): + kv_cache_config = _kv_cache_config(tokens_per_block=32) + + tokens_per_block = _select_mla_tokens_per_block( + _non_mla_config(), + _model_config(kv_cache_quant_algo=QuantAlgo.NVFP4, enable_flash_mla=True), + kv_cache_config, + kv_cache_config.tokens_per_block, + ) + + assert tokens_per_block == 32 + assert kv_cache_config.tokens_per_block == 32 + + +def test_flash_mla_non_fp4_uses_flash_mla_tokens_per_block(): + kv_cache_config = _kv_cache_config(tokens_per_block=32) + + tokens_per_block = _select_mla_tokens_per_block( + _mla_config(), + _model_config(enable_flash_mla=True), + kv_cache_config, + kv_cache_config.tokens_per_block, + ) + + assert tokens_per_block == FLASH_MLA_TOKENS_PER_BLOCK + assert kv_cache_config.tokens_per_block == FLASH_MLA_TOKENS_PER_BLOCK + + +def test_fp4_mla_dequant_flow_uses_flash_mla_tokens_per_block(monkeypatch): + monkeypatch.delenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, raising=False) + kv_cache_config = _kv_cache_config(tokens_per_block=32) + + tokens_per_block = _select_mla_tokens_per_block( + _mla_config(), + _model_config(kv_cache_quant_algo=QuantAlgo.NVFP4, enable_flash_mla=True), + kv_cache_config, + kv_cache_config.tokens_per_block, + ) + + assert tokens_per_block == FLASH_MLA_TOKENS_PER_BLOCK + assert kv_cache_config.tokens_per_block == FLASH_MLA_TOKENS_PER_BLOCK + + +def test_fp4_mla_attention_uses_128_tokens_per_block_from_quant_config(monkeypatch): + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") + kv_cache_config = _kv_cache_config(tokens_per_block=32) + + tokens_per_block = _select_mla_tokens_per_block( + _mla_config(), + _model_config(kv_cache_quant_algo=QuantAlgo.NVFP4, enable_flash_mla=True), + kv_cache_config, + kv_cache_config.tokens_per_block, + ) + + assert tokens_per_block == FP4_MLA_TOKENS_PER_BLOCK + assert kv_cache_config.tokens_per_block == FP4_MLA_TOKENS_PER_BLOCK + + +def test_fp4_mla_attention_uses_128_tokens_per_block_from_kv_cache_dtype(monkeypatch): + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") + kv_cache_config = _kv_cache_config(dtype="nvfp4", tokens_per_block=32) + + tokens_per_block = _select_mla_tokens_per_block( + _mla_config(), + _model_config(enable_flash_mla=True), + kv_cache_config, + kv_cache_config.tokens_per_block, + ) + + assert tokens_per_block == FP4_MLA_TOKENS_PER_BLOCK + assert kv_cache_config.tokens_per_block == FP4_MLA_TOKENS_PER_BLOCK From c5c2232e997ece03cb9991b092dca449baedc011 Mon Sep 17 00:00:00 2001 From: Tracin <10434017+Tracin@users.noreply.github.com> Date: Mon, 25 May 2026 02:52:02 -0700 Subject: [PATCH 03/11] [None][feat] Add CuTile and CuTe DSL FP4 MLA decode backends Add CuTile and CuTe DSL implementations for the FP4 MLA paged decode path and route them through the existing backend selector. Extend the opt-in FP4 MLA decode benchmark coverage to report the selected backend and estimated memory bandwidth. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> --- .../_torch/attention_backend/fp4_mla_cute.py | 803 ++++++ .../attention_backend/fp4_mla_cutile.py | 2475 +++++++++++++++++ .../_torch/attention_backend/fp4_mla_kv.py | 224 +- .../_torch/attention/test_fp4_mla_kv.py | 262 +- 4 files changed, 3676 insertions(+), 88 deletions(-) create mode 100644 tensorrt_llm/_torch/attention_backend/fp4_mla_cute.py create mode 100644 tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_cute.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_cute.py new file mode 100644 index 000000000000..fca764b1b5f8 --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_cute.py @@ -0,0 +1,803 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""CuTe DSL FP4 MLA decode backend. + +This module intentionally preserves the Python FP4 MLA decode contract from +``fp4_mla_kv.run_fp4_mla_attention_decode``. It consumes the same packed Q, +packed paged KV cache, swizzled scale tensors, page tables, and workspace +buffers as the Triton backend. +""" + +import math + +import torch + +from ...logger import logger +from ..cute_dsl_utils import IS_CUTLASS_DSL_AVAILABLE + +if IS_CUTLASS_DSL_AVAILABLE: + try: + from cuda.bindings import driver as cuda + except ImportError: + from cuda import cuda + + import cutlass + import cutlass.cute as cute + from cutlass._mlir.dialects import llvm + from cutlass.cute.runtime import from_dlpack + from cutlass.cutlass_dsl import T, dsl_user_op + + class _CUDAGraphCompatibleWrapper: + """Wrapper to make DLPack export safe during CUDA graph capture.""" + + def __init__(self, tensor: torch.Tensor) -> None: + self._tensor = tensor + + def __dlpack__(self, stream=None): + return self._tensor.__dlpack__(stream=-1) + + def __dlpack_device__(self): + return self._tensor.__dlpack_device__() + + def _to_cute(tensor: torch.Tensor) -> cute.Tensor: + return from_dlpack( + _CUDAGraphCompatibleWrapper(tensor.detach()), assumed_align=16 + ).mark_layout_dynamic() + + @cute.jit + def _swizzled_sf_offset(row_idx, col_idx, sf_per_token: cutlass.Constexpr): + padded_cols = ((sf_per_token + 3) // 4) * 4 + return ( + col_idx % 4 + + (col_idx // 4) * (4 * 128) + + (row_idx % 32) * 16 + + ((row_idx % 128) // 32) * 4 + + (row_idx // 128) * (128 * padded_cols) + ) + + @dsl_user_op + def _ptx_fp4_e2m1x2_to_f16x2(byte, *, loc=None, ip=None) -> cutlass.Uint32: + return cutlass.Uint32( + llvm.inline_asm( + T.i32(), + [cutlass.Uint32(byte).ir_value(loc=loc, ip=ip)], + """ + { + .reg .b8 in_8; + .reg .f16x2 out; + cvt.u8.u32 in_8, $1; + cvt.rn.f16x2.e2m1x2 out, in_8; + mov.b32 $0, out; + } + """, + "=r,r", + has_side_effects=False, + is_align_stack=False, + asm_dialect=llvm.AsmDialect.AD_ATT, + ) + ) + + @dsl_user_op + def _ptx_fp8_e4m3x2_to_f16x2(byte, *, loc=None, ip=None) -> cutlass.Uint32: + return cutlass.Uint32( + llvm.inline_asm( + T.i32(), + [cutlass.Uint32(byte).ir_value(loc=loc, ip=ip)], + """ + { + .reg .b16 in_16; + .reg .f16x2 out; + cvt.u16.u32 in_16, $1; + cvt.rn.f16x2.e4m3x2 out, in_16; + mov.b32 $0, out; + } + """, + "=r,r", + has_side_effects=False, + is_align_stack=False, + asm_dialect=llvm.AsmDialect.AD_ATT, + ) + ) + + @dsl_user_op + def _ptx_fp4_e2m1x2_from_f32(even, odd, *, loc=None, ip=None) -> cutlass.Uint32: + return cutlass.Uint32( + llvm.inline_asm( + T.i32(), + [ + cutlass.Float32(odd).ir_value(loc=loc, ip=ip), + cutlass.Float32(even).ir_value(loc=loc, ip=ip), + ], + """ + { + .reg .b8 out; + cvt.rn.satfinite.e2m1x2.f32 out, $1, $2; + mov.b32 $0, {out, out, out, out}; + } + """, + "=r,f,f", + has_side_effects=False, + is_align_stack=False, + asm_dialect=llvm.AsmDialect.AD_ATT, + ) + ) + + @dsl_user_op + def _ptx_fp8_e4m3x2_from_f32(low, high, *, loc=None, ip=None) -> cutlass.Uint32: + return cutlass.Uint32( + llvm.inline_asm( + T.i32(), + [ + cutlass.Float32(low).ir_value(loc=loc, ip=ip), + cutlass.Float32(high).ir_value(loc=loc, ip=ip), + ], + """ + { + .reg .b16 out; + cvt.rn.satfinite.e4m3x2.f32 out, $2, $1; + mov.b32 $0, {out, out}; + } + """, + "=r,f,f", + has_side_effects=False, + is_align_stack=False, + asm_dialect=llvm.AsmDialect.AD_ATT, + ) + ) + + @cute.jit + def _low_f16x2_lane_to_f32(bits): + half_bits = cutlass.Uint16(bits & cutlass.Uint32(0xFFFF)) + half_value = cutlass.Float16(llvm.bitcast(cutlass.Float16.mlir_type, half_bits.ir_value())) + return half_value.to(cutlass.Float32) + + @cute.jit + def _high_f16x2_lane_to_f32(bits): + half_bits = cutlass.Uint16((bits >> 16) & cutlass.Uint32(0xFFFF)) + half_value = cutlass.Float16(llvm.bitcast(cutlass.Float16.mlir_type, half_bits.ir_value())) + return half_value.to(cutlass.Float32) + + @cute.jit + def _fp4_e2m1_to_f32(nibble): + bits = _ptx_fp4_e2m1x2_to_f16x2(cute.Uint8(nibble & cute.Uint8(0x0F))) + return _low_f16x2_lane_to_f32(bits) + + @cute.jit + def _fp4_e2m1_quantize_packed(even, odd): + bits = _ptx_fp4_e2m1x2_from_f32(even, odd) + return cute.Uint8(bits & cutlass.Uint32(0xFF)) + + @cute.jit + def _load_fp4_value(packed_tensor, packed_offset, elem_idx): + packed = packed_tensor[packed_offset] + nibble = cute.Uint8(packed & 0x0F) + if (elem_idx & 1) != 0: + nibble = cute.Uint8((packed >> 4) & 0x0F) + return _fp4_e2m1_to_f32(nibble) + + @cute.jit + def _load_fp4_byte_pair(packed_tensor, packed_offset): + """Load one packed FP4 byte and return both nibbles as (low, high) f32. + + The PTX ``cvt.f16x2.e2m1x2`` instruction converts both nibbles in a + single op; the previous scalar path discarded the high half. + """ + packed = packed_tensor[packed_offset] + bits = _ptx_fp4_e2m1x2_to_f16x2(cute.Uint8(packed)) + return _low_f16x2_lane_to_f32(bits), _high_f16x2_lane_to_f32(bits) + + @cute.jit + def _fp8_e4m3fn_to_f32(byte): + bits = _ptx_fp8_e4m3x2_to_f16x2(cute.Uint8(byte)) + return _low_f16x2_lane_to_f32(bits) + + @cute.jit + def _fp8_e4m3fn_positive_from_f32(value): + bits = _ptx_fp8_e4m3x2_from_f32(value, cutlass.Float32(0.0)) + return cute.Uint8(bits & cutlass.Uint32(0xFF)) + + class _Fp4MlaDecodeCuteKernel: + def __init__( + self, + *, + num_heads: int, + kv_lora_rank: int, + qk_rope_head_dim: int, + q_residual_dim: int, + page_size: int, + max_pages: int, + q_fp4_stride0: int, + q_fp4_stride1: int, + kv_stride0: int, + kv_stride2: int, + kv_stride4: int, + sf_stride0: int, + v_sf_stride0: int, + p_stride0: int, + p_stride1: int, + page_stats_stride0: int, + page_stats_stride1: int, + page_stats_stride2: int, + output_dtype: torch.dtype, + ) -> None: + self.num_heads = num_heads + self.kv_lora_rank = kv_lora_rank + self.qk_rope_head_dim = qk_rope_head_dim + self.q_residual_dim = q_residual_dim + self.k_head_dim = kv_lora_rank + qk_rope_head_dim + self.q_head_dim = self.k_head_dim + q_residual_dim + self.page_size = page_size + self.max_pages = max_pages + self.q_fp4_stride0 = q_fp4_stride0 + self.q_fp4_stride1 = q_fp4_stride1 + self.kv_stride0 = kv_stride0 + self.kv_stride2 = kv_stride2 + self.kv_stride4 = kv_stride4 + self.sf_stride0 = sf_stride0 + self.v_sf_stride0 = v_sf_stride0 + self.p_stride0 = p_stride0 + self.p_stride1 = p_stride1 + self.page_stats_stride0 = page_stats_stride0 + self.page_stats_stride1 = page_stats_stride1 + self.page_stats_stride2 = page_stats_stride2 + self.output_dtype = output_dtype + self.fp4_block = 16 + self.k_sf_per_token = self.k_head_dim // self.fp4_block + self.q_sf_per_token = self.q_head_dim // self.fp4_block + self.p_sf_per_page = page_size // self.fp4_block + self.non_residual_groups = self.k_sf_per_token - q_residual_dim // self.fp4_block + self.log2_e = math.log2(math.e) + self.p_global_scale = 448.0 * 6.0 + self.head_tile = 128 + self.v_tile = 128 + + @cute.jit + def __call__( + self, + output, + max_scores, + denom, + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4, + q_sf, + kv_cache, + sf_cache, + v_sf, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + sm_scale: cutlass.Float32, + stream: cuda.CUstream, + ) -> None: + num_head_blocks = cute.ceil_div(self.num_heads, self.head_tile) + self._page_stats_pack_kernel( + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4, + q_sf, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + sm_scale, + ).launch( + grid=(output.shape[0], num_head_blocks, self.max_pages), + block=(self.head_tile, 1, 1), + stream=stream, + ) + + self._reduce_stats_kernel( + max_scores, + denom, + page_max, + page_sum, + ).launch( + grid=(output.shape[0], num_head_blocks, 1), + block=(self.head_tile, 1, 1), + stream=stream, + ) + + self._prob_scale_kernel( + max_scores, + denom, + page_max, + p_sf, + paged_kv_indptr_decode, + kv_lens, + ).launch( + grid=(output.shape[0], num_head_blocks, self.max_pages), + block=(self.head_tile, 1, 1), + stream=stream, + ) + + self._pv_kernel( + output, + p_fp4, + p_sf, + kv_cache, + v_sf, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + ).launch( + grid=( + output.shape[0], + self.num_heads, + cute.ceil_div(self.kv_lora_rank, self.v_tile), + ), + block=(self.v_tile, 1, 1), + stream=stream, + ) + + @cute.jit + def _qk_score( + self, + q_fp4, + q_sf, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + compact_page, + token_idx, + q_row, + sm_scale, + ): + physical_page = src_page_ids[compact_page] + score = cutlass.Float32(0.0) + bytes_per_group: cutlass.Constexpr = self.fp4_block // 2 + q_row_base = q_row * self.q_fp4_stride0 + kv_token_base = physical_page * self.kv_stride0 + token_idx * self.kv_stride2 + k_sf_page_base = physical_page * self.sf_stride0 + + # Keep q_group as a runtime loop (44 iters): unrolling it together + # with the inner byte-pair loop and the outer 128-token loop blows + # up compile time and instruction footprint. + for q_group in cutlass.range(self.q_sf_per_token, unroll=1): + k_group = q_group + if q_group >= self.non_residual_groups: + k_group = self.non_residual_groups + (q_group - self.non_residual_groups) // 2 + + # Scales only change at fp4_block boundaries — hoist out of the + # inner byte-pair loop instead of reloading per element. + q_scale = _fp8_e4m3fn_to_f32( + q_sf[_swizzled_sf_offset(q_row, q_group, self.q_sf_per_token)] + ) + k_scale = _fp8_e4m3fn_to_f32( + sf_cache[ + k_sf_page_base + + _swizzled_sf_offset(token_idx, k_group, self.k_sf_per_token) + ] + ) + qk_scale = q_scale * k_scale + + # Each FP4 byte holds 2 nibbles; process them together so the + # single cvt.f16x2.e2m1x2 produces 2 useful f32 values. + for byte_idx in cutlass.range_constexpr(bytes_per_group): + q_packed_col = q_group * bytes_per_group + byte_idx + k_packed_col = k_group * bytes_per_group + byte_idx + q_lo, q_hi = _load_fp4_byte_pair( + q_fp4, + q_row_base + q_packed_col * self.q_fp4_stride1, + ) + k_lo, k_hi = _load_fp4_byte_pair( + kv_cache, + kv_token_base + k_packed_col * self.kv_stride4, + ) + score += (q_lo * k_lo + q_hi * k_hi) * qk_scale + + scale = sm_scale / (global_scale[0] * global_scale[0]) + return score * scale + + @cute.jit + def _page_stats_offset(self, gen_idx, page_rel, head_idx): + return ( + gen_idx * self.page_stats_stride0 + + page_rel * self.page_stats_stride1 + + head_idx * self.page_stats_stride2 + ) + + @cute.kernel + def _page_stats_pack_kernel( + self, + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4, + q_sf, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + sm_scale: cutlass.Float32, + ) -> None: + tidx, _, _ = cute.arch.thread_idx() + gen_idx, head_block, page_rel = cute.arch.block_idx() + head_idx = head_block * self.head_tile + tidx + max_score = -cutlass.Float32.inf + sum_value = cutlass.Float32(0.0) + + if head_idx < self.num_heads: + q_row = gen_idx * self.num_heads + head_idx + kv_len = kv_lens[gen_idx] + page_start = page_rel * self.page_size + if page_start < kv_len: + page_table_start = paged_kv_indptr_decode[gen_idx] + compact_page = page_table_start + page_rel + scores = cute.make_fragment((self.page_size,), cutlass.Float32) + for token_offset in cutlass.range_constexpr(self.page_size): + token_abs = page_start + token_offset + score = -cutlass.Float32.inf + if token_abs < kv_len: + score = self._qk_score( + q_fp4, + q_sf, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + compact_page, + token_offset, + q_row, + sm_scale, + ) + if score > max_score: + max_score = score + scores[token_offset] = score + + p_row = compact_page * self.num_heads + head_idx + for token_group in cutlass.range_constexpr(self.p_sf_per_page): + local_max = cutlass.Float32(0.0) + probs = cute.make_fragment((self.fp4_block,), cutlass.Float32) + for idx in cutlass.range_constexpr(self.fp4_block): + token_offset = token_group * self.fp4_block + idx + token_abs = page_start + token_offset + prob = cutlass.Float32(0.0) + if token_abs < kv_len: + prob = cute.math.exp2( + (scores[token_offset] - max_score) * self.log2_e, + fastmath=True, + ) + sum_value += prob + probs[idx] = prob + if prob > local_max: + local_max = prob + + local_scale = cutlass.Float32(1.0) + stored_scale = cutlass.Float32(1.0) + if local_max > cutlass.Float32(0.0): + local_scale = local_max / cutlass.Float32(6.0) + stored_scale = local_scale * self.p_global_scale + if stored_scale > cutlass.Float32(448.0): + stored_scale = cutlass.Float32(448.0) + + p_sf[ + _swizzled_sf_offset( + p_row, + token_group, + self.p_sf_per_page, + ) + ] = _fp8_e4m3fn_positive_from_f32(stored_scale) + + for byte_idx in cutlass.range_constexpr(self.fp4_block // 2): + packed = _fp4_e2m1_quantize_packed( + probs[byte_idx * 2] / local_scale, + probs[byte_idx * 2 + 1] / local_scale, + ) + p_fp4[ + p_row * self.p_stride0 + + (token_group * (self.fp4_block // 2) + byte_idx) * self.p_stride1 + ] = packed + + stats_offset = self._page_stats_offset(gen_idx, page_rel, head_idx) + page_max[stats_offset] = max_score + page_sum[stats_offset] = sum_value + + @cute.kernel + def _reduce_stats_kernel(self, max_scores, denom, page_max, page_sum) -> None: + tidx, _, _ = cute.arch.thread_idx() + gen_idx, head_block, _ = cute.arch.block_idx() + head_idx = head_block * self.head_tile + tidx + if head_idx < self.num_heads: + max_score = -cutlass.Float32.inf + for page_rel in cutlass.range(self.max_pages, unroll=1): + stats_offset = self._page_stats_offset(gen_idx, page_rel, head_idx) + page_max_value = page_max[stats_offset] + if page_max_value > max_score: + max_score = page_max_value + + denom_value = cutlass.Float32(0.0) + for page_rel in cutlass.range(self.max_pages, unroll=1): + stats_offset = self._page_stats_offset(gen_idx, page_rel, head_idx) + page_sum_value = page_sum[stats_offset] + if page_sum_value > cutlass.Float32(0.0): + denom_value += page_sum_value * cute.math.exp2( + (page_max[stats_offset] - max_score) * self.log2_e, + fastmath=True, + ) + + q_row = gen_idx * self.num_heads + head_idx + max_scores[q_row] = max_score + denom[q_row] = denom_value + + @cute.kernel + def _prob_scale_kernel( + self, + max_scores, + denom, + page_max, + p_sf, + paged_kv_indptr_decode, + kv_lens, + ) -> None: + tidx, _, _ = cute.arch.thread_idx() + gen_idx, head_block, page_rel = cute.arch.block_idx() + head_idx = head_block * self.head_tile + tidx + if head_idx < self.num_heads: + kv_len = kv_lens[gen_idx] + page_start = page_rel * self.page_size + if page_start < kv_len: + stats_offset = self._page_stats_offset(gen_idx, page_rel, head_idx) + q_row = gen_idx * self.num_heads + head_idx + denom_value = denom[q_row] + factor = cutlass.Float32(0.0) + if denom_value > cutlass.Float32(0.0): + factor = ( + cute.math.exp2( + (page_max[stats_offset] - max_scores[q_row]) * self.log2_e, + fastmath=True, + ) + / denom_value + ) + + page_table_start = paged_kv_indptr_decode[gen_idx] + p_row = (page_table_start + page_rel) * self.num_heads + head_idx + for token_group in cutlass.range_constexpr(self.p_sf_per_page): + sf_offset = _swizzled_sf_offset( + p_row, + token_group, + self.p_sf_per_page, + ) + scaled = _fp8_e4m3fn_to_f32(p_sf[sf_offset]) * factor + p_sf[sf_offset] = _fp8_e4m3fn_positive_from_f32(scaled) + + @cute.kernel + def _pv_kernel( + self, + output, + p_fp4, + p_sf, + kv_cache, + v_sf, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + ) -> None: + tidx, _, _ = cute.arch.thread_idx() + gen_idx, head_idx, dim_block = cute.arch.block_idx() + v_dim = dim_block * self.v_tile + tidx + if v_dim < self.kv_lora_rank: + kv_len = kv_lens[gen_idx] + page_table_start = paged_kv_indptr_decode[gen_idx] + v_packed_col = v_dim // 2 # constant per thread + acc = cutlass.Float32(0.0) + bytes_per_group: cutlass.Constexpr = self.fp4_block // 2 + for page_rel in cutlass.range(self.max_pages, unroll=1): + page_start = page_rel * self.page_size + if page_start < kv_len: + compact_page = page_table_start + page_rel + physical_page = src_page_ids[compact_page] + p_row = compact_page * self.num_heads + head_idx + p_row_base = p_row * self.p_stride0 + v_page_base = ( + physical_page * self.kv_stride0 + v_packed_col * self.kv_stride4 + ) + v_sf_page_base = physical_page * self.v_sf_stride0 + # Process 16 tokens at a time — both scales are + # constant within each fp4_block, so hoist them. + for token_group in cutlass.range(self.p_sf_per_page, unroll=1): + group_start = token_group * self.fp4_block + if page_start + group_start < kv_len: + p_scale = _fp8_e4m3fn_to_f32( + p_sf[ + _swizzled_sf_offset( + p_row, + token_group, + self.p_sf_per_page, + ) + ] + ) + v_scale = _fp8_e4m3fn_to_f32( + v_sf[ + v_sf_page_base + + _swizzled_sf_offset( + v_dim, + token_group, + self.p_sf_per_page, + ) + ] + ) + pv_scale = p_scale * v_scale + # Each P byte holds 2 adjacent tokens; load + # once and use both nibbles. + for byte_idx in cutlass.range_constexpr(bytes_per_group): + token_a = group_start + byte_idx * 2 + token_b = token_a + 1 + token_abs_a = page_start + token_a + if token_abs_a < kv_len: + p_packed_col = token_group * bytes_per_group + byte_idx + p_lo, p_hi = _load_fp4_byte_pair( + p_fp4, + p_row_base + p_packed_col * self.p_stride1, + ) + v_a = _load_fp4_value( + kv_cache, + v_page_base + token_a * self.kv_stride2, + v_dim, + ) + acc += p_lo * v_a * pv_scale + if token_abs_a + 1 < kv_len: + v_b = _load_fp4_value( + kv_cache, + v_page_base + token_b * self.kv_stride2, + v_dim, + ) + acc += p_hi * v_b * pv_scale + + output[gen_idx, head_idx, v_dim] = ( + acc / (global_scale[0] * self.p_global_scale) + ).to(output.element_type) + + _COMPILE_CACHE: dict[tuple[int, ...], object] = {} + + def _storage_span(tensor: torch.Tensor) -> int: + if tensor.numel() == 0: + return 0 + return 1 + sum((size - 1) * stride for size, stride in zip(tensor.shape, tensor.stride())) + + def _flatten(tensor: torch.Tensor) -> torch.Tensor: + if tensor.is_contiguous(): + return tensor.reshape(-1) + return torch.as_strided( + tensor, + size=(_storage_span(tensor),), + stride=(1,), + storage_offset=tensor.storage_offset(), + ) + + def run_fp4_mla_attention_decode_cute( + *, + output: torch.Tensor, + max_scores: torch.Tensor, + denom: torch.Tensor, + page_max: torch.Tensor, + page_sum: torch.Tensor, + p_fp4: torch.Tensor, + p_sf: torch.Tensor, + q_fp4: torch.Tensor, + q_sf: torch.Tensor, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + v_sf: torch.Tensor, + global_scale: torch.Tensor, + src_page_ids: torch.Tensor, + paged_kv_indptr_decode: torch.Tensor, + kv_lens: torch.Tensor, + sm_scale: float, + kv_lora_rank: int, + qk_rope_head_dim: int, + q_residual_dim: int, + page_size: int, + max_pages: int, + ) -> None: + """Run FP4 MLA decode using the page-parallel CuTe DSL backend.""" + + q_fp4_flat = _flatten(q_fp4.view(torch.uint8)) + q_sf_flat = _flatten(q_sf.view(torch.uint8)) + kv_cache_flat = _flatten(kv_cache.view(torch.uint8)) + sf_cache_flat = _flatten(sf_cache.view(torch.uint8)) + v_sf_flat = _flatten(v_sf.view(torch.uint8)) + p_fp4_flat = _flatten(p_fp4.view(torch.uint8)) + p_sf_flat = _flatten(p_sf.view(torch.uint8)) + max_scores_flat = _flatten(max_scores) + denom_flat = _flatten(denom) + page_max_flat = _flatten(page_max) + page_sum_flat = _flatten(page_sum) + + stream = cuda.CUstream(torch.cuda.current_stream(output.device).cuda_stream) + cute_args = ( + _to_cute(output), + _to_cute(max_scores_flat), + _to_cute(denom_flat), + _to_cute(page_max_flat), + _to_cute(page_sum_flat), + _to_cute(p_fp4_flat), + _to_cute(p_sf_flat), + _to_cute(q_fp4_flat), + _to_cute(q_sf_flat), + _to_cute(kv_cache_flat), + _to_cute(sf_cache_flat), + _to_cute(v_sf_flat), + _to_cute(global_scale), + _to_cute(src_page_ids), + _to_cute(paged_kv_indptr_decode), + _to_cute(kv_lens), + cutlass.Float32(sm_scale), + stream, + ) + compile_key = ( + output.device.index or 0, + output.dtype is torch.bfloat16, + output.shape[0], + output.shape[1], + kv_lora_rank, + qk_rope_head_dim, + q_residual_dim, + page_size, + max_pages, + q_fp4.stride(0), + q_fp4.stride(1), + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + sf_cache.stride(0), + v_sf.stride(0), + p_fp4.stride(0), + p_fp4.stride(1), + page_max.stride(0), + page_max.stride(1), + page_max.stride(2), + ) + if compile_key not in _COMPILE_CACHE: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "CuTe FP4 MLA decode must be compiled before CUDA graph capture." + ) + logger.info( + "Compiling CuTe FP4 MLA decode kernel for " + f"num_gen={output.shape[0]}, num_heads={output.shape[1]}, " + f"kv_lora_rank={kv_lora_rank}, rope_dim={qk_rope_head_dim}, " + f"max_pages={max_pages}" + ) + kernel = _Fp4MlaDecodeCuteKernel( + num_heads=output.shape[1], + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + q_residual_dim=q_residual_dim, + page_size=page_size, + max_pages=max_pages, + q_fp4_stride0=q_fp4.stride(0), + q_fp4_stride1=q_fp4.stride(1), + kv_stride0=kv_cache.stride(0), + kv_stride2=kv_cache.stride(2), + kv_stride4=kv_cache.stride(4), + sf_stride0=sf_cache.stride(0), + v_sf_stride0=v_sf.stride(0), + p_stride0=p_fp4.stride(0), + p_stride1=p_fp4.stride(1), + page_stats_stride0=page_max.stride(0), + page_stats_stride1=page_max.stride(1), + page_stats_stride2=page_max.stride(2), + output_dtype=output.dtype, + ) + _COMPILE_CACHE[compile_key] = cute.compile(kernel, *cute_args) + + _COMPILE_CACHE[compile_key](*cute_args) + +else: + + def run_fp4_mla_attention_decode_cute(**_: object) -> None: + raise RuntimeError("CuTe DSL is not available in this environment.") diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py new file mode 100644 index 000000000000..337f017fe2c1 --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py @@ -0,0 +1,2475 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# SPDX-License-Identifier: Apache-2.0 + +"""FP4 MLA paged decode attention using Triton. + +The kernels are adapted from TensorRT-LLM's FP4 MLA decode path. This module +exposes the attention path for already-packed FP4 Q/K/V tensors and swizzled +FP8 block-scale tensors; quantization and KV-cache update helpers remain outside +this internal op. +""" + +from typing import Optional + +import torch +import triton +import triton.language as tl + +FP4_BLOCK_SIZE = 16 +FP4_MLA_P_GLOBAL_SCALE = 448.0 * 6.0 + + +def _ceil_div(lhs: int, rhs: int) -> int: + return (lhs + rhs - 1) // rhs + + +def _swizzled_scale_size(rows: int, logical_cols: int) -> int: + scale_cols = _ceil_div(logical_cols, FP4_BLOCK_SIZE) + padded_cols = _ceil_div(scale_cols, 4) * 4 + return _ceil_div(rows, 128) * 128 * padded_cols + + +def _get_kv_cache_strides(kv_cache: torch.Tensor) -> tuple[int, int, int, int, int, int]: + if kv_cache.dim() == 3: + num_pages, page_size, packed_dim = kv_cache.shape + return ( + num_pages, + page_size, + packed_dim, + kv_cache.stride(0), + kv_cache.stride(1), + kv_cache.stride(2), + ) + if kv_cache.dim() >= 5: + num_pages = kv_cache.shape[0] + page_size = kv_cache.shape[2] + packed_dim = kv_cache.shape[4] + return ( + num_pages, + page_size, + packed_dim, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + ) + raise ValueError( + "kv_cache must be shaped (num_pages, page_size, packed_dim) or (num_pages, ..., page_size, ..., packed_dim)." + ) + + +def _workspace_tensor( + workspace: Optional[torch.Tensor], + shape: tuple[int, ...], + *, + dtype: torch.dtype, + device: torch.device, + name: str, +) -> torch.Tensor: + if workspace is None: + if torch.cuda.is_current_stream_capturing(): + raise ValueError( + f"Cannot allocate {name} while capturing a CUDA graph. " + "Pass a preallocated workspace tensor." + ) + return torch.empty(shape, dtype=dtype, device=device) + + invalid = ( + workspace.dtype != dtype + or workspace.device != device + or len(workspace.shape) != len(shape) + or any(workspace.shape[idx] < dim for idx, dim in enumerate(shape)) + ) + if invalid: + raise ValueError( + f"{name} workspace must have shape at least {shape}, dtype={dtype}, " + f"and device={device}; got shape={tuple(workspace.shape)}, " + f"dtype={workspace.dtype}, device={workspace.device}." + ) + + slices = tuple(slice(0, dim) for dim in shape) + return workspace[slices] + + +@triton.jit +def _fp4_mla_swizzled_sf_offset(row_idx, col_idx, SF_PER_TOKEN: tl.constexpr): + padded_cols = ((SF_PER_TOKEN + 3) // 4) * 4 + col_in_group = col_idx % 4 + col_group = col_idx // 4 + row_in_group0 = row_idx % 32 + row_in_group1 = (row_idx % 128) // 32 + row_group = row_idx // 128 + return ( + col_in_group + + col_group * (4 * 128) + + row_in_group0 * 16 + + row_in_group1 * 4 + + row_group * (128 * padded_cols) + ) + + +@triton.jit +def _fp4_mla_swizzled_sf_offset_row_block( + row_group, row_offsets, col_idx, SF_PER_TOKEN: tl.constexpr +): + padded_cols = ((SF_PER_TOKEN + 3) // 4) * 4 + col_part = (col_idx % 4) + (col_idx // 4) * (4 * 128) + row_part = (row_offsets % 32) * 16 + ((row_offsets % 128) // 32) * 4 + return col_part + row_part + row_group * (128 * padded_cols) + + +@triton.jit +def _fp4_e2m1_quantize(x): + abs_x = tl.abs(x) + magnitude = tl.where( + abs_x < 0.25, + 0, + tl.where( + abs_x < 0.75, + 1, + tl.where( + abs_x < 1.25, + 2, + tl.where( + abs_x < 1.75, + 3, + tl.where(abs_x < 2.5, 4, tl.where(abs_x < 3.5, 5, tl.where(abs_x < 5.0, 6, 7))), + ), + ), + ), + ) + sign = tl.where(x < 0.0, 8, 0) + return (magnitude | sign).to(tl.uint8) + + +@triton.jit +def _fp4_e2m1_quantize_packed(even, odd): + return tl.inline_asm_elementwise( + """ + { + .reg .b8 r; + cvt.rn.satfinite.e2m1x2.f32 r, $1, $2; + mov.b32 $0, {r, r, r, r}; + } + """, + constraints="=r,f,f", + args=[odd.to(tl.float32), even.to(tl.float32)], + dtype=tl.uint8, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fp4_pack_low_nibbles(even_packed, odd_packed): + return tl.inline_asm_elementwise( + """ + { + .reg .b32 lo; + .reg .b32 hi; + and.b32 lo, $1, 15; + and.b32 hi, $2, 15; + shl.b32 hi, hi, 4; + or.b32 $0, lo, hi; + } + """, + constraints="=r,r,r", + args=[even_packed, odd_packed], + dtype=tl.uint8, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fp4_pack_high_nibbles(even_packed, odd_packed): + return tl.inline_asm_elementwise( + """ + { + .reg .b32 lo; + .reg .b32 hi; + shr.u32 lo, $1, 4; + and.b32 lo, lo, 15; + and.b32 hi, $2, 240; + or.b32 $0, lo, hi; + } + """, + constraints="=r,r,r", + args=[even_packed, odd_packed], + dtype=tl.uint8, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_start, + head_offsets, + token_offsets, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + NUM_HEADS: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, +): + if ASSUME_VALID_PAGES: + safe_compact_page = compact_page + physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) + safe_physical_page = physical_page + else: + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + ) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + q_rows = q_row_base + head_offsets + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_q_rows = q_rows + else: + mask_h = head_offsets < NUM_HEADS + safe_q_rows = tl.where(mask_h, q_rows, q_row_base) + scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) + + if ( + USE_TMA_DATA_LOAD + and ASSUME_FULL_HEADS + and ASSUME_VALID_PAGES + and Q_HEAD_D == 640 + and K_HEAD_D == 576 + and Q_RESIDUAL_D == 64 + and BLOCK_H == 128 + and BLOCK_T == 128 + and BLOCK_K == 512 + and FULL_BLOCK_END == 512 + and TAIL_BLOCK_K == 128 + ): + tl.assume(q_fp4_s0 % 8 == 0) + tl.assume(q_fp4_s1 == 1) + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + q_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 256], + ) + k_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 256], + ) + q_tail_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 64], + ) + k_tail_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 32], + ) + q_sf_full_view = tl.ext.make_view( + base=q_sf_ptr, + shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + k_sf_full_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + k_sf_tail_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 1, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + + full_q_vals = q_desc.load([(q_row_base + head_start).to(tl.int32), 0]) + full_k_vals = k_desc.load([safe_physical_page.to(tl.int32), 0, 0]) + full_k_vals = tl.reshape(full_k_vals, (BLOCK_T, 256)) + q_row_group = q_row_base // 128 + full_q_scales = tl.ext.load_view_tko(q_sf_full_view, [q_row_group.to(tl.int32), 0, 0, 0]) + full_q_scales = full_q_scales.reshape([1, 8, 32, 4, 4]).trans(0, 3, 2, 1, 4) + full_q_scales = full_q_scales.reshape([BLOCK_H, 32]) + full_k_scales = tl.ext.load_view_tko( + k_sf_full_view, [safe_physical_page.to(tl.int32), 0, 0, 0, 0] + ) + full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) + scores = tl.dot_scaled( + full_q_vals, + full_q_scales, + "e2m1", + full_k_vals.T, + full_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + tail_k_vals = k_tail_desc.load([safe_physical_page.to(tl.int32), 0, 256]) + tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) + tail_k_scales = tl.ext.load_view_tko( + k_sf_tail_view, [safe_physical_page.to(tl.int32), 0, 8, 0, 0] + ) + tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) + q_tail_vals = q_tail_desc.load([(q_row_base + head_start).to(tl.int32), 256]) + # Map Q tail groups [0, 1, ..., 7] onto K tail groups [0, 0, 1, 1, ..., 3, 3]. + q_tail_vals = q_tail_vals.reshape([BLOCK_H, 4, 2, 8]).trans(0, 1, 3, 2) + q_even_vals, q_odd_vals = tl.split(q_tail_vals) + q_even_vals = q_even_vals.reshape([BLOCK_H, 32]) + q_odd_vals = q_odd_vals.reshape([BLOCK_H, 32]) + + q_tail_sf_cols = 32 + tl.arange(0, 8) + q_tail_sf_offsets = _fp4_mla_swizzled_sf_offset( + q_rows[:, None], q_tail_sf_cols[None, :], Q_SF_PER_TOKEN + ) + q_tail_scales = tl.load(q_sf_ptr + q_tail_sf_offsets) + q_tail_scales = q_tail_scales.reshape([BLOCK_H, 4, 2]) + q_even_scales, q_odd_scales = tl.split(q_tail_scales) + scores = tl.dot_scaled( + q_even_vals, + q_even_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q_odd_vals, + q_odd_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + return scores + + packed_k_offsets = tl.arange(0, BLOCK_K // 2) + scale_offsets = tl.arange(0, BLOCK_K // FP4_BLOCK) + residual_groups = Q_RESIDUAL_D // FP4_BLOCK + non_residual_groups = K_HEAD_D // FP4_BLOCK - residual_groups + if USE_TMA_DATA_LOAD and FULL_BLOCK_END > 0: + tl.assume(q_fp4_s0 % 8 == 0) + tl.assume(q_fp4_s1 == 1) + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + q_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, BLOCK_K // 2], + ) + k_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, BLOCK_K // 2], + ) + if USE_TMA_DATA_LOAD and Q_RESIDUAL_D == 64 and TAIL_BLOCK_K == 128: + q_tail_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 64], + ) + k_tail_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 32], + ) + for q_start in tl.range(0, FULL_BLOCK_END, BLOCK_K): + q_elem_offsets = q_start + packed_k_offsets * 2 + q_group_offsets = q_elem_offsets // FP4_BLOCK + k_group_offsets = tl.where( + q_group_offsets < non_residual_groups, + q_group_offsets, + non_residual_groups + (q_group_offsets - non_residual_groups) // 2, + ) + byte_offsets_in_group = (q_elem_offsets % FP4_BLOCK) // 2 + packed_q_cols = q_start // 2 + packed_k_offsets + packed_k_cols = k_group_offsets * (FP4_BLOCK // 2) + byte_offsets_in_group + mask_k = q_elem_offsets < Q_HEAD_D + safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) + safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) + if ( + USE_TMA_DATA_LOAD + and FULL_BLOCK_END > 0 + and q_start + BLOCK_K <= non_residual_groups * FP4_BLOCK + ): + q_vals = q_desc.load([(q_row_base + head_start).to(tl.int32), q_start // 2]) + k_vals = k_desc.load([safe_physical_page.to(tl.int32), 0, q_start // 2]) + k_vals = tl.reshape(k_vals, (BLOCK_T, BLOCK_K // 2)) + if not ASSUME_VALID_PAGES: + k_vals = tl.where(valid_physical_page, k_vals, 0) + else: + q_vals = tl.load( + q_fp4_ptr + + safe_q_rows[:, None] * q_fp4_s0 + + safe_packed_q_cols[None, :] * q_fp4_s1, + mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], + other=0, + ) + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + safe_packed_k_cols[None, :] * kv_s4, + mask=mask_k[None, :] + if ASSUME_VALID_PAGES + else valid_physical_page & mask_k[None, :], + other=0, + ) + + q_sf_cols = q_start // FP4_BLOCK + scale_offsets + k_sf_cols = tl.where( + q_sf_cols < non_residual_groups, + q_sf_cols, + non_residual_groups + (q_sf_cols - non_residual_groups) // 2, + ) + mask_sf = q_sf_cols < Q_SF_PER_TOKEN + safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) + safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) + q_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN + ) + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN + ) + q_scales = tl.load(q_sf_ptr + q_sf_offsets) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) + scores = tl.dot_scaled( + q_vals, + q_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + if FULL_BLOCK_END < Q_HEAD_D: + q_start = FULL_BLOCK_END + if Q_RESIDUAL_D == 64 and TAIL_BLOCK_K == 128: + residual_packed_offsets = tl.arange(0, 32) + residual_scale_offsets = tl.arange(0, 4) + packed_k_cols = non_residual_groups * (FP4_BLOCK // 2) + residual_packed_offsets + if USE_TMA_DATA_LOAD: + k_vals = k_tail_desc.load( + [ + safe_physical_page.to(tl.int32), + 0, + (non_residual_groups * (FP4_BLOCK // 2)).to(tl.int32), + ] + ) + k_vals = tl.reshape(k_vals, (BLOCK_T, 32)) + if not ASSUME_VALID_PAGES: + k_vals = tl.where(valid_physical_page, k_vals, 0) + elif ASSUME_VALID_PAGES: + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + packed_k_cols[None, :] * kv_s4, + ) + else: + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + packed_k_cols[None, :] * kv_s4, + mask=valid_physical_page, + other=0, + ) + k_sf_cols = non_residual_groups + residual_scale_offsets + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], k_sf_cols[None, :], K_SF_PER_TOKEN + ) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) + + q_tail_cols = q_start // 2 + tl.arange(0, 64) + if USE_TMA_DATA_LOAD and ASSUME_FULL_HEADS: + q_tail_vals = q_tail_desc.load( + [(q_row_base + head_start).to(tl.int32), q_start // 2] + ) + elif ASSUME_FULL_HEADS: + q_tail_vals = tl.load( + q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1 + ) + else: + q_tail_vals = tl.load( + q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1, + mask=mask_h[:, None], + other=0, + ) + # Map Q tail groups [0, 1, ..., 7] onto K tail groups [0, 0, 1, 1, ..., 3, 3]. + q_tail_vals = q_tail_vals.reshape([BLOCK_H, 4, 2, 8]).trans(0, 1, 3, 2) + q_even_vals, q_odd_vals = tl.split(q_tail_vals) + q_even_vals = q_even_vals.reshape([BLOCK_H, 32]) + q_odd_vals = q_odd_vals.reshape([BLOCK_H, 32]) + + q_tail_sf_cols = q_start // FP4_BLOCK + tl.arange(0, 8) + q_tail_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], q_tail_sf_cols[None, :], Q_SF_PER_TOKEN + ) + q_tail_scales = tl.load(q_sf_ptr + q_tail_sf_offsets) + q_tail_scales = q_tail_scales.reshape([BLOCK_H, 4, 2]) + q_even_scales, q_odd_scales = tl.split(q_tail_scales) + scores = tl.dot_scaled( + q_even_vals, + q_even_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q_odd_vals, + q_odd_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + else: + tail_packed_offsets = tl.arange(0, TAIL_BLOCK_K // 2) + tail_scale_offsets = tl.arange(0, TAIL_BLOCK_K // FP4_BLOCK) + q_elem_offsets = q_start + tail_packed_offsets * 2 + q_group_offsets = q_elem_offsets // FP4_BLOCK + k_group_offsets = tl.where( + q_group_offsets < non_residual_groups, + q_group_offsets, + non_residual_groups + (q_group_offsets - non_residual_groups) // 2, + ) + byte_offsets_in_group = (q_elem_offsets % FP4_BLOCK) // 2 + packed_q_cols = q_start // 2 + tail_packed_offsets + packed_k_cols = k_group_offsets * (FP4_BLOCK // 2) + byte_offsets_in_group + mask_k = q_elem_offsets < Q_HEAD_D + safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) + safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) + q_vals = tl.load( + q_fp4_ptr + + safe_q_rows[:, None] * q_fp4_s0 + + safe_packed_q_cols[None, :] * q_fp4_s1, + mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], + other=0, + ) + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + safe_packed_k_cols[None, :] * kv_s4, + mask=mask_k[None, :] + if ASSUME_VALID_PAGES + else valid_physical_page & mask_k[None, :], + other=0, + ) + + q_sf_cols = q_start // FP4_BLOCK + tail_scale_offsets + k_sf_cols = tl.where( + q_sf_cols < non_residual_groups, + q_sf_cols, + non_residual_groups + (q_sf_cols - non_residual_groups) // 2, + ) + mask_sf = q_sf_cols < Q_SF_PER_TOKEN + safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) + safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) + q_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN + ) + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN + ) + q_scales = tl.load(q_sf_ptr + q_sf_offsets) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) + scores = tl.dot_scaled( + q_vals, + q_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + return scores + + +@triton.jit +def _fp4_mla_attention_stats_kernel( + max_ptr, + denom_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + stats_s0, + q_num_rows, + sm_scale, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, BLOCK_T) + q_row_base = gen_idx * NUM_HEADS + kv_len = tl.load(kv_lens_ptr + gen_idx) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + + max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + denom = tl.zeros((BLOCK_H,), dtype=tl.float32) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + for page_rel in tl.range(0, MAX_PAGES): + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES or page_start < kv_len: + compact_page = page_table_start + page_rel + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + BLOCK_T, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + if ASSUME_FULL_PAGES: + scores = tl.where(mask_h[:, None], scores * qk_scale, -float("inf")) + else: + valid_t = page_start + offs_t < kv_len + scores = tl.where( + mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf") + ) + page_max = tl.max(scores, axis=1) + new_max = tl.maximum(max_score, page_max) + denom = denom * tl.math.exp2((max_score - new_max) * 1.4426950408889634) + tl.sum( + tl.math.exp2((scores - new_max[:, None]) * 1.4426950408889634), axis=1 + ) + max_score = new_max + + tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) + tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) + + +@triton.jit +def _fp4_mla_attention_page_stats_kernel( + page_max_ptr, + page_sum_ptr, + p_fp4_ptr, + p_sf_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_stats_s0, + page_stats_s1, + p_s0, + p_s1, + p_num_rows, + q_num_rows, + sm_scale, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + PACK_PROBS: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_rel = tl.program_id(2) + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, BLOCK_T) + q_row_base = gen_idx * NUM_HEADS + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + gen_idx) + page_start = page_rel * PAGE_SIZE + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + out_offsets = gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h + + page_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + page_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) + if USE_TMA_DATA_LOAD and PACK_PROBS and ASSUME_FULL_HEADS and ASSUME_VALID_PAGES: + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + if ASSUME_FULL_PAGES or page_start < kv_len: + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + page_table_start + page_rel, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + BLOCK_T, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + if ASSUME_FULL_PAGES: + valid_t = tl.full([BLOCK_T], True, dtype=tl.int1) + else: + valid_t = page_start + offs_t < kv_len + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + if ASSUME_FULL_HEADS and ASSUME_FULL_PAGES: + scores = scores * qk_scale + page_max = tl.max(scores, axis=1) + exp_scores = tl.math.exp2((scores - page_max[:, None]) * 1.4426950408889634) + page_sum = tl.sum(exp_scores, axis=1) + else: + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + page_max = tl.max(scores, axis=1) + safe_page_max = tl.where(mask_h, page_max, 0.0) + exp_scores = tl.math.exp2((scores - safe_page_max[:, None]) * 1.4426950408889634) + exp_scores = tl.where(mask_h[:, None] & valid_t[None, :], exp_scores, 0.0) + page_sum = tl.sum(exp_scores, axis=1) + + if PACK_PROBS: + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + stored_scale = tl.where( + amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0 + ) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + if ASSUME_VALID_PAGES: + safe_compact_page = page_table_start + page_rel + else: + valid_compact_page = (page_table_start + page_rel >= 0) & ( + page_table_start + page_rel < page_ids_len + ) + safe_compact_page = tl.where(valid_compact_page, page_table_start + page_rel, 0) + p_rows = safe_compact_page * NUM_HEADS + offs_h + safe_p_rows = ( + p_rows + if ASSUME_FULL_HEADS + else tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) + ) + scale_cols = tl.arange(0, SF_PER_PAGE) + if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + safe_compact_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + if ASSUME_FULL_HEADS: + if ASSUME_VALID_PAGES: + tl.store(p_sf_ptr + sf_offsets, stored_scale) + else: + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_compact_page) + else: + tl.store( + p_sf_ptr + sf_offsets, + stored_scale, + mask=mask_h[:, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None], + ) + + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + if USE_TMA_DATA_LOAD and ASSUME_FULL_HEADS and ASSUME_VALID_PAGES: + p_desc.store( + [(safe_compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), + ) + elif ASSUME_FULL_HEADS: + if ASSUME_VALID_PAGES: + tl.store( + p_fp4_ptr + + safe_p_rows[:, None, None] * p_s0 + + byte_cols[None, :, :] * p_s1, + packed, + ) + else: + tl.store( + p_fp4_ptr + + safe_p_rows[:, None, None] * p_s0 + + byte_cols[None, :, :] * p_s1, + packed, + mask=valid_compact_page, + ) + else: + tl.store( + p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=mask_h[:, None, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None, None], + ) + + if ASSUME_FULL_HEADS: + tl.store(page_max_ptr + out_offsets, page_max) + tl.store(page_sum_ptr + out_offsets, page_sum) + else: + tl.store(page_max_ptr + out_offsets, page_max, mask=mask_h) + tl.store(page_sum_ptr + out_offsets, page_sum, mask=mask_h) + + +@triton.jit +def _fp4_mla_attention_reduce_stats_kernel( + max_ptr, + denom_ptr, + page_max_ptr, + page_sum_ptr, + stats_s0, + page_stats_s0, + page_stats_s1, + NUM_HEADS: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + + max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES): + page_max = tl.load( + page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.maximum(max_score, page_max) + + denom = tl.zeros((BLOCK_H,), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES): + page_max = tl.load( + page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + page_sum = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=0.0, + ) + denom += tl.where( + page_sum > 0.0, + page_sum * tl.math.exp2((page_max - max_score) * 1.4426950408889634), + 0.0, + ) + + tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) + tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) + + +@triton.jit +def _fp4_mla_attention_prob_scale_kernel( + p_sf_ptr, + max_ptr, + denom_ptr, + page_max_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + stats_s0, + page_stats_s0, + page_stats_s1, + NUM_HEADS: tl.constexpr, + PAGE_SIZE: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + BLOCK_H: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_rel = tl.program_id(2) + + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + gen_idx) + if page_start >= kv_len: + return + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if not ASSUME_VALID_PAGES: + if (compact_page < 0) | (compact_page >= page_ids_len): + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + page_max = tl.load( + page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.load(max_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + factor = tl.where( + denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0 + ) + + p_rows = compact_page * NUM_HEADS + offs_h + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, compact_page * NUM_HEADS) + ) + scale_cols = tl.arange(0, SF_PER_PAGE) + if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + compact_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) + tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) + + +@triton.jit +def _fp4_mla_attention_prob_store_page_kernel( + probs_ptr, + max_ptr, + denom_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_rel, + page_ids_len, + num_pages, + probs_s0, + probs_s1, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + stats_s0, + q_num_rows, + sm_scale, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + + kv_len = tl.load(kv_lens_ptr + gen_idx) + page_start = page_rel * PAGE_SIZE + if (not ASSUME_FULL_PAGES) and page_start >= kv_len: + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, PAGE_SIZE) + if ASSUME_FULL_PAGES: + valid_t = tl.full([PAGE_SIZE], True, dtype=tl.int1) + else: + valid_t = page_start + offs_t < kv_len + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if (compact_page < 0) | (compact_page >= page_ids_len): + return + q_row_base = gen_idx * NUM_HEADS + + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + PAGE_SIZE, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + max_score = tl.load(max_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + denom_valid = denom > 0.0 + safe_denom = tl.where(denom_valid, denom, 1.0) + safe_max = tl.where(denom_valid, max_score, 0.0) + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + probs = tl.math.exp2((scores - safe_max[:, None]) * 1.4426950408889634) / safe_denom[:, None] + probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) + + prob_rows = gen_idx * NUM_HEADS + offs_h + safe_prob_rows = tl.where(mask_h, prob_rows, gen_idx * NUM_HEADS) + tl.store( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + offs_t[None, :] * probs_s1, + probs, + mask=mask_h[:, None], + ) + + +@triton.jit +def _fp4_mla_attention_prob_pack_page_kernel( + p_fp4_ptr, + p_sf_ptr, + probs_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_rel, + page_ids_len, + p_s0, + p_s1, + probs_s0, + probs_s1, + NUM_HEADS: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + token_group = tl.program_id(1) + head_block = tl.program_id(2) + + kv_len = tl.load(kv_lens_ptr + gen_idx) + page_start = page_rel * PAGE_SIZE + if (not ASSUME_FULL_PAGES) and page_start >= kv_len: + return + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if (compact_page < 0) | (compact_page >= page_ids_len): + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + else: + mask_h = offs_h < NUM_HEADS + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + token_base = token_group * FP4_BLOCK + even_t = token_base + byte_offsets * 2 + odd_t = even_t + 1 + valid_even = page_start + even_t < kv_len + valid_odd = page_start + odd_t < kv_len + + prob_rows = gen_idx * NUM_HEADS + offs_h + safe_prob_rows = tl.where(mask_h, prob_rows, gen_idx * NUM_HEADS) + even_probs = tl.load( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + even_t[None, :] * probs_s1, + mask=mask_h[:, None] & valid_even[None, :], + other=0.0, + ) + odd_probs = tl.load( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + odd_t[None, :] * probs_s1, + mask=mask_h[:, None] & valid_odd[None, :], + other=0.0, + ) + amax = tl.maximum(tl.max(tl.abs(even_probs), axis=1), tl.max(tl.abs(odd_probs), axis=1)) + local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) + + p_rows = compact_page * NUM_HEADS + offs_h + safe_p_rows = tl.where(mask_h, p_rows, compact_page * NUM_HEADS) + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows, token_group, SF_PER_PAGE) + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h) + + even_quant = _fp4_e2m1_quantize(even_probs / local_scale[:, None]) + odd_quant = _fp4_e2m1_quantize(odd_probs / local_scale[:, None]) + packed = even_quant | (odd_quant << 4) + byte_cols = token_group * (FP4_BLOCK // 2) + byte_offsets + tl.store( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + byte_cols[None, :] * p_s1, + packed, + mask=mask_h[:, None], + ) + + +@triton.jit +def _fp4_mla_attention_prob_pack_page_fused_kernel( + p_fp4_ptr, + p_sf_ptr, + max_ptr, + denom_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_rel, + page_ids_len, + num_pages, + p_s0, + p_s1, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + stats_s0, + q_num_rows, + sm_scale, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + PAGE_REL_FROM_GRID: tl.constexpr = False, + ASSUME_FULL_HEADS: tl.constexpr = False, + ASSUME_FULL_PAGES: tl.constexpr = False, + ASSUME_VALID_PAGES: tl.constexpr = False, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + if PAGE_REL_FROM_GRID: + page_rel = tl.program_id(2) + + kv_len = tl.load(kv_lens_ptr + gen_idx) + page_start = page_rel * PAGE_SIZE + if (not ASSUME_FULL_PAGES) and page_start >= kv_len: + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, PAGE_SIZE) + if ASSUME_FULL_PAGES: + valid_t = tl.full([PAGE_SIZE], True, dtype=tl.int1) + else: + valid_t = page_start + offs_t < kv_len + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if (compact_page < 0) | (compact_page >= page_ids_len): + return + q_row_base = gen_idx * NUM_HEADS + + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + PAGE_SIZE, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + max_score = tl.load(max_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + denom_valid = denom > 0.0 + safe_denom = tl.where(denom_valid, denom, 1.0) + safe_max = tl.where(denom_valid, max_score, 0.0) + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + probs = tl.math.exp2((scores - safe_max[:, None]) * 1.4426950408889634) / safe_denom[:, None] + probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) + + grouped_probs = tl.reshape(probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(tl.abs(grouped_probs), axis=2) + local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) + scaled_probs = grouped_probs / tl.reshape(local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + p_rows = compact_page * NUM_HEADS + offs_h + safe_p_rows = tl.where(mask_h, p_rows, compact_page * NUM_HEADS) + scale_cols = tl.arange(0, SF_PER_PAGE) + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h[:, None]) + + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + tl.store( + p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=mask_h[:, None, None], + ) + + +@triton.jit +def _fp4_mla_attention_pv_kernel( + out_ptr, + p_fp4_ptr, + p_sf_ptr, + kv_cache_ptr, + v_sf_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + out_s0, + out_s1, + out_s2, + out_num_rows, + p_s0, + p_s1, + p_num_rows, + kv_s0, + kv_s2, + kv_s4, + vsf_s0, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + USE_TMA_P_LOAD: tl.constexpr, + USE_TMA_V_LOAD: tl.constexpr, + PV_LOOP_STAGES: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_FULL_V: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + dim_block = tl.program_id(2) + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + if ASSUME_FULL_V: + mask_v = tl.full([BLOCK_V], True, dtype=tl.int1) + safe_offs_v = offs_v + else: + mask_v = offs_v < V_HEAD_D + safe_offs_v = tl.where(mask_v, offs_v, 0) + packed_t = tl.arange(0, PAGE_SIZE // 2) + scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) + even_t = packed_t * 2 + odd_t = even_t + 1 + v_packed_offsets = safe_offs_v // 2 + v_use_high_nibble = (safe_offs_v & 1) != 0 + if ASSUME_FULL_V and BLOCK_V == 128: + v_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + v_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + if USE_TMA_P_LOAD: + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + if USE_TMA_V_LOAD and ASSUME_FULL_HEADS and ASSUME_FULL_V: + tl.assume(out_s1 % 8 == 0) + tl.assume(out_s2 == 1) + out_desc = tl.make_tensor_descriptor( + out_ptr, + shape=[out_num_rows, V_HEAD_D], + strides=[out_s1, out_s2], + block_shape=[BLOCK_H, BLOCK_V], + ) + if USE_TMA_V_LOAD: + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, PAGE_SIZE, BLOCK_V // 2], + ) + + if ( + USE_TMA_P_LOAD + and USE_TMA_V_LOAD + and ASSUME_FULL_HEADS + and ASSUME_FULL_PAGES + and ASSUME_FULL_V + and ASSUME_VALID_PAGES + and NUM_HEADS == 128 + and V_HEAD_D == 512 + and PAGE_SIZE == 128 + and BLOCK_H == 128 + and BLOCK_V == 128 + and SF_PER_PAGE == 8 + ): + p_view = tl.ext.make_view( + base=p_fp4_ptr, + shapes=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + tile_shape=[BLOCK_H, PAGE_SIZE // 2], + tile_dim_map=[0, 1], + ) + p_sf_view = tl.ext.make_view( + base=p_sf_ptr, + shapes=[p_num_rows // 128, SF_PER_PAGE // 4, 2, 256], + strides=[128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + v_sf_view = tl.ext.make_view( + base=v_sf_ptr, + shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], + strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + v_view = tl.ext.make_view( + base=kv_cache_ptr, + shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], + tile_dim_map=[0, 1, 2], + ) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): + compact_page = page_table_start + page_rel + physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) + + p_vals = tl.ext.load_view_tko( + p_view, + [(compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + ) + p_vals = p_vals.to(tl.uint8, bitcast=True) + p_scales = tl.ext.load_view_tko(p_sf_view, [compact_page.to(tl.int32), 0, 0, 0]) + p_scales = p_scales.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) + p_scales = p_scales.reshape([BLOCK_H, SF_PER_PAGE]) + + v_tile = tl.ext.load_view_tko( + v_view, + [ + physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ], + ) + v_tile = v_tile.to(tl.uint8, bitcast=True) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) + high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + v_scales = tl.ext.load_view_tko( + v_sf_view, + [ + physical_page.to(tl.int32), + dim_block, + 0, + 0, + 0, + ], + ) + v_scales = v_scales.reshape([1, 1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) + acc = tl.dot_scaled( + p_vals, + p_scales, + "e2m1", + v_vals.T, + v_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + out_vals = acc * out_scale + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + out_desc.store( + [ + (gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ], + out_vals, + ) + return + + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + gen_idx) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES or page_start < kv_len: + compact_page = page_table_start + page_rel + if ASSUME_VALID_PAGES: + safe_compact_page = compact_page + physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) + safe_physical_page = physical_page + else: + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + ) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + + p_rows = safe_compact_page * NUM_HEADS + offs_h + safe_p_rows = ( + p_rows + if ASSUME_FULL_HEADS + else tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) + ) + if USE_TMA_P_LOAD: + p_vals = p_desc.load( + [(safe_compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0] + ) + else: + p_vals = tl.load( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, + mask=mask_h[:, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None], + other=0, + ) + p_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + p_scales = tl.load(p_sf_ptr + p_sf_offsets) + + if ASSUME_FULL_PAGES: + valid_even_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + valid_odd_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + else: + valid_even_t = page_start + even_t < kv_len + valid_odd_t = page_start + odd_t < kv_len + if USE_TMA_V_LOAD: + v_tile = v_desc.load( + [ + safe_physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ] + ) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + if not ASSUME_VALID_PAGES: + v_tile = tl.where(valid_physical_page, v_tile, 0) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + if not ASSUME_FULL_PAGES: + even_packed = tl.where(valid_even_t[None, :], even_packed, 0) + odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) + low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) + high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + else: + even_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + even_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_even_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), + other=0, + ) + odd_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + odd_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_odd_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), + other=0, + ) + even_low = even_packed & 0x0F + even_high = (even_packed >> 4) & 0x0F + odd_low = odd_packed & 0x0F + even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) + odd_nibble = tl.where(v_use_high_nibble[:, None], odd_packed >> 4, odd_low) + v_vals = even_nibble | (odd_nibble << 4) + v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) + acc = tl.dot_scaled( + p_vals, + p_scales, + "e2m1", + v_vals.T, + v_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + if ASSUME_FULL_HEADS and ASSUME_FULL_V: + out_vals = acc * out_scale + if USE_TMA_V_LOAD: + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + out_desc.store( + [ + (gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ], + out_vals, + ) + else: + tl.store( + out_ptr + gen_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + ) + else: + tl.store( + out_ptr + + gen_idx * out_s0 + + safe_offs_h[:, None] * out_s1 + + safe_offs_v[None, :] * out_s2, + acc * out_scale, + mask=mask_h[:, None] & mask_v[None, :], + ) + + +def fp4_mla_paged_attention_internal( + q_fp4: torch.Tensor, + q_sf: torch.Tensor, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + v_sf: torch.Tensor, + global_scale: torch.Tensor, + src_page_ids: torch.Tensor, + paged_kv_indptr_decode: torch.Tensor, + kv_lens: torch.Tensor, + output: Optional[torch.Tensor] = None, + *, + sm_scale: float, + num_heads: Optional[int] = None, + v_head_dim: Optional[int] = None, + page_size: Optional[int] = None, + q_residual_dim: int = 0, + p_global_scale: float = FP4_MLA_P_GLOBAL_SCALE, + block_h: int = 128, + block_k: Optional[int] = None, + block_v: int = 128, + output_dtype: torch.dtype = torch.bfloat16, + max_pages: Optional[int] = None, + page_pipeline_streams: Optional[int] = None, + kernel_occupancy: Optional[int] = None, + kernel_num_ctas: Optional[int] = None, + kernel_num_stages: Optional[int] = None, + kernel_num_warps: Optional[int] = None, + pv_loop_stages: int = 1, + parallel_page_stats: Optional[bool] = None, + fused_prob_pack: Optional[bool] = None, + use_tma_data_load: Optional[bool] = None, + fused_prob_pack_single_launch: Optional[bool] = None, + pack_prob_in_page_stats: Optional[bool] = None, + assume_full_pages: Optional[bool] = None, + assume_valid_pages: Optional[bool] = None, + p_fp4_workspace: Optional[torch.Tensor] = None, + p_sf_workspace: Optional[torch.Tensor] = None, + p_probs_workspace: Optional[torch.Tensor] = None, + max_scores_workspace: Optional[torch.Tensor] = None, + denom_workspace: Optional[torch.Tensor] = None, + page_max_workspace: Optional[torch.Tensor] = None, + page_sum_workspace: Optional[torch.Tensor] = None, + **kwargs, +) -> torch.Tensor: + del kwargs + if not hasattr(tl, "dot_scaled"): + raise NotImplementedError( + "fp4_mla_paged_attention requires a Triton build with tl.dot_scaled." + ) + if not q_fp4.is_cuda: + raise ValueError("q_fp4 must be a CUDA tensor.") + if q_fp4.dtype != torch.uint8 or kv_cache.dtype != torch.uint8: + raise TypeError("q_fp4 and kv_cache must be packed FP4 tensors with dtype torch.uint8.") + if global_scale.numel() < 1: + raise ValueError("global_scale must contain at least one element.") + if q_fp4.dim() == 3: + inferred_num_gen, inferred_num_heads, packed_q_dim = q_fp4.shape + if num_heads is not None and num_heads != inferred_num_heads: + raise ValueError( + f"num_heads={num_heads} does not match q_fp4.shape[1]={inferred_num_heads}." + ) + num_gen = inferred_num_gen + num_heads = inferred_num_heads + q_fp4_2d = q_fp4.reshape(num_gen * num_heads, packed_q_dim) + elif q_fp4.dim() == 2: + if num_heads is None: + raise ValueError("num_heads is required when q_fp4 is 2D.") + if q_fp4.shape[0] % num_heads != 0: + raise ValueError("q_fp4.shape[0] must be divisible by num_heads.") + num_gen = q_fp4.shape[0] // num_heads + packed_q_dim = q_fp4.shape[1] + q_fp4_2d = q_fp4 + else: + raise ValueError("q_fp4 must be 2D or 3D.") + + num_pages, inferred_page_size, packed_k_dim, kv_s0, kv_s2, kv_s4 = _get_kv_cache_strides( + kv_cache + ) + if page_size is None: + page_size = inferred_page_size + if page_size != inferred_page_size: + raise ValueError( + f"page_size={page_size} does not match kv_cache page dimension {inferred_page_size}." + ) + if page_size % FP4_BLOCK_SIZE != 0: + raise ValueError(f"page_size must be divisible by {FP4_BLOCK_SIZE}.") + + q_head_dim = packed_q_dim * 2 + k_head_dim = packed_k_dim * 2 + if q_residual_dim < 0 or q_residual_dim % FP4_BLOCK_SIZE != 0: + raise ValueError(f"q_residual_dim must be a non-negative multiple of {FP4_BLOCK_SIZE}.") + if q_head_dim - q_residual_dim != k_head_dim: + raise ValueError( + f"q_head_dim - q_residual_dim must match K head dim: {q_head_dim} - {q_residual_dim} != {k_head_dim}." + ) + if q_head_dim % FP4_BLOCK_SIZE != 0 or k_head_dim % FP4_BLOCK_SIZE != 0: + raise ValueError(f"Q/K head dims must be divisible by {FP4_BLOCK_SIZE}.") + if v_head_dim is None: + v_head_dim = k_head_dim + if v_head_dim <= 0 or v_head_dim > k_head_dim: + raise ValueError(f"v_head_dim must be in (0, {k_head_dim}], got {v_head_dim}.") + q_sf_flat = q_sf.contiguous().view(-1) + if sf_cache.shape[0] < num_pages or v_sf.shape[0] < num_pages: + raise ValueError("sf_cache and v_sf must have a leading physical-page dimension.") + if q_sf_flat.numel() < _swizzled_scale_size(num_gen * num_heads, q_head_dim): + raise ValueError("q_sf is too small for the swizzled Q scale layout.") + if sf_cache.numel() < sf_cache.shape[0] * _swizzled_scale_size(page_size, k_head_dim): + raise ValueError("sf_cache is too small for the swizzled K scale layout.") + if v_sf.numel() < v_sf.shape[0] * _swizzled_scale_size(v_head_dim, page_size): + raise ValueError("v_sf is too small for the swizzled V scale layout.") + + if output is None: + output = torch.empty( + (num_gen, num_heads, v_head_dim), dtype=output_dtype, device=q_fp4.device + ) + elif output.shape != (num_gen, num_heads, v_head_dim): + raise ValueError( + f"output must have shape {(num_gen, num_heads, v_head_dim)}, got {tuple(output.shape)}." + ) + + if num_gen == 0: + return output + triton_backend = "nvt" + if block_k is None: + block_k = 512 if triton_backend == "nvt" else 256 + full_block_end = (q_head_dim // block_k) * block_k + tail_k = q_head_dim - full_block_end + tail_block_k = 1 << (tail_k - 1).bit_length() if tail_k > 0 else block_k + if max_pages is None: + if paged_kv_indptr_decode.numel() >= num_gen + 1: + page_counts = paged_kv_indptr_decode[1 : num_gen + 1] - paged_kv_indptr_decode[:num_gen] + max_pages = int(page_counts.max().item()) if page_counts.numel() > 0 else 0 + else: + max_pages = _ceil_div(int(kv_lens[:num_gen].max().item()), page_size) + if max_pages <= 0: + output.zero_() + return output + + q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE + k_sf_per_token = k_head_dim // FP4_BLOCK_SIZE + sf_per_page = page_size // FP4_BLOCK_SIZE + num_head_blocks = triton.cdiv(num_heads, block_h) + assume_full_heads = num_heads % block_h == 0 + assume_full_v = v_head_dim % block_v == 0 + if assume_full_pages is None: + assume_full_pages = False + assume_full_pages = bool(assume_full_pages) + if assume_valid_pages is None: + assume_valid_pages = False + assume_valid_pages = bool(assume_valid_pages) + total_p_rows = max(src_page_ids.numel() * num_heads, 1) + if page_pipeline_streams is None: + if triton_backend == "nvt" and max_pages >= 8 and num_gen >= 128: + page_pipeline_streams = 2 + else: + page_pipeline_streams = 1 + page_pipeline_streams = max(1, min(int(page_pipeline_streams), max_pages)) + launch_meta = {} + if kernel_occupancy is None and triton_backend == "nvt": + kernel_occupancy = 2 + if kernel_occupancy is not None: + launch_meta["occupancy"] = int(kernel_occupancy) + if kernel_num_ctas is not None: + launch_meta["num_ctas"] = int(kernel_num_ctas) + if kernel_num_stages is not None: + launch_meta["num_stages"] = int(kernel_num_stages) + if kernel_num_warps is not None: + launch_meta["num_warps"] = int(kernel_num_warps) + if fused_prob_pack is None: + fused_prob_pack = triton_backend == "nvt" + if fused_prob_pack_single_launch is None: + fused_prob_pack_single_launch = triton_backend == "nvt" and max_pages >= 8 + if use_tma_data_load is None: + use_tma_data_load = triton_backend == "nvt" + use_tma_data_load = bool(use_tma_data_load and hasattr(tl, "make_tensor_descriptor")) + if use_tma_data_load: + # Device-side descriptors may need Triton's allocator for descriptor scratch storage. + def alloc_fn(size: int, alignment: int, stream: Optional[int]): + return torch.empty(size, device=q_fp4.device, dtype=torch.int8) + + triton.set_allocator(alloc_fn) + + p_fp4 = _workspace_tensor( + p_fp4_workspace, + (total_p_rows, page_size // 2), + dtype=torch.uint8, + device=q_fp4.device, + name="p_fp4", + ) + p_sf = _workspace_tensor( + p_sf_workspace, + (_swizzled_scale_size(total_p_rows, page_size),), + dtype=q_sf.dtype, + device=q_fp4.device, + name="p_sf", + ) + if fused_prob_pack: + p_probs = None + else: + p_probs_shape = (max(num_gen * num_heads, 1), page_size) + if page_pipeline_streams > 1: + p_probs = _workspace_tensor( + p_probs_workspace, + (page_pipeline_streams, *p_probs_shape), + dtype=torch.float32, + device=q_fp4.device, + name="p_probs", + ) + else: + p_probs = _workspace_tensor( + p_probs_workspace, + p_probs_shape, + dtype=torch.float32, + device=q_fp4.device, + name="p_probs", + ) + max_scores = _workspace_tensor( + max_scores_workspace, + (num_gen, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="max_scores", + ) + denom = _workspace_tensor( + denom_workspace, + (num_gen, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="denom", + ) + + if parallel_page_stats is None: + parallel_page_stats = triton_backend == "nvt" and max_pages >= 8 + if pack_prob_in_page_stats is None: + pack_prob_in_page_stats = parallel_page_stats and fused_prob_pack + pack_prob_in_page_stats = bool( + pack_prob_in_page_stats and parallel_page_stats and fused_prob_pack + ) + if parallel_page_stats: + page_stats_shape = (num_gen, max_pages, num_heads) + page_max = _workspace_tensor( + page_max_workspace, + page_stats_shape, + dtype=torch.float32, + device=q_fp4.device, + name="page_max", + ) + page_sum = _workspace_tensor( + page_sum_workspace, + page_stats_shape, + dtype=torch.float32, + device=q_fp4.device, + name="page_sum", + ) + _fp4_mla_attention_page_stats_kernel[(num_gen, num_head_blocks, max_pages)]( + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + page_max.stride(0), + page_max.stride(1), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + BLOCK_T=page_size, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + PACK_PROBS=pack_prob_in_page_stats, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + _fp4_mla_attention_reduce_stats_kernel[(num_gen, num_head_blocks)]( + max_scores, + denom, + page_max, + page_sum, + max_scores.stride(0), + page_max.stride(0), + page_max.stride(1), + NUM_HEADS=num_heads, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + **launch_meta, + ) + if pack_prob_in_page_stats: + _fp4_mla_attention_prob_scale_kernel[(num_gen, num_head_blocks, max_pages)]( + p_sf, + max_scores, + denom, + page_max, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + max_scores.stride(0), + page_max.stride(0), + page_max.stride(1), + NUM_HEADS=num_heads, + PAGE_SIZE=page_size, + SF_PER_PAGE=sf_per_page, + BLOCK_H=block_h, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + else: + _fp4_mla_attention_stats_kernel[(num_gen, num_head_blocks)]( + max_scores, + denom, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + max_scores.stride(0), + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + BLOCK_T=page_size, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + + def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None): + if fused_prob_pack: + _fp4_mla_attention_prob_pack_page_fused_kernel[(num_gen, num_head_blocks)]( + p_fp4, + p_sf, + max_scores, + denom, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + page_rel, + src_page_ids.shape[0], + num_pages, + p_fp4.stride(0), + p_fp4.stride(1), + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + max_scores.stride(0), + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + return + assert p_probs_slot is not None + _fp4_mla_attention_prob_store_page_kernel[(num_gen, num_head_blocks)]( + p_probs_slot, + max_scores, + denom, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + page_rel, + src_page_ids.shape[0], + num_pages, + p_probs_slot.stride(0), + p_probs_slot.stride(1), + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + max_scores.stride(0), + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + BLOCK_H=block_h, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + _fp4_mla_attention_prob_pack_page_kernel[(num_gen, sf_per_page, num_head_blocks)]( + p_fp4, + p_sf, + p_probs_slot, + paged_kv_indptr_decode, + kv_lens, + page_rel, + src_page_ids.shape[0], + p_fp4.stride(0), + p_fp4.stride(1), + p_probs_slot.stride(0), + p_probs_slot.stride(1), + NUM_HEADS=num_heads, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + **launch_meta, + ) + + if pack_prob_in_page_stats: + pass + elif fused_prob_pack and fused_prob_pack_single_launch: + _fp4_mla_attention_prob_pack_page_fused_kernel[(num_gen, num_head_blocks, max_pages)]( + p_fp4, + p_sf, + max_scores, + denom, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + 0, + src_page_ids.shape[0], + num_pages, + p_fp4.stride(0), + p_fp4.stride(1), + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + max_scores.stride(0), + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + PAGE_REL_FROM_GRID=True, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + **launch_meta, + ) + elif page_pipeline_streams == 1: + for page_rel in range(max_pages): + _launch_prob_page(page_rel, p_probs) + else: + current_stream = torch.cuda.current_stream(q_fp4.device) + streams = [torch.cuda.Stream(device=q_fp4.device) for _ in range(page_pipeline_streams)] + for stream in streams: + stream.wait_stream(current_stream) + for page_rel in range(max_pages): + stream_idx = page_rel % page_pipeline_streams + with torch.cuda.stream(streams[stream_idx]): + if fused_prob_pack: + _launch_prob_page(page_rel) + else: + assert p_probs is not None + _launch_prob_page(page_rel, p_probs[stream_idx]) + for stream in streams: + current_stream.wait_stream(stream) + + num_dim_blocks = triton.cdiv(v_head_dim, block_v) + _fp4_mla_attention_pv_kernel[ + ( + num_gen, + num_head_blocks, + num_dim_blocks, + ) + ]( + output, + p_fp4, + p_sf, + kv_cache, + v_sf, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + output.stride(0), + output.stride(1), + output.stride(2), + output.shape[0] * output.shape[1], + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + kv_s0, + kv_s2, + kv_s4, + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + BLOCK_V=block_v, + USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, + USE_TMA_V_LOAD=use_tma_data_load and v_head_dim % block_v == 0, + PV_LOOP_STAGES=int(pv_loop_stages), + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_FULL_V=assume_full_v, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + return output + + +fp4_mla_paged_attention = fp4_mla_paged_attention_internal diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py index 00664c41ed18..7b32e3e53b47 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py @@ -44,7 +44,9 @@ FP4_MLA_P_GLOBAL_SCALE: float = 448.0 * 6.0 FP4_MLA_Q_RESIDUAL_DIM: int = 64 FLASHINFER_FP4_MLA_ATTENTION_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION" +FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION_BACKEND" FLASHINFER_FP4_MLA_DEBUG_ENV = "TRTLLM_FLASHINFER_FP4_MLA_DEBUG" +_FP4_MLA_CUTE_DSL_BACKEND = "cute_dsl" _HPUpdatePhase = Literal["all", "context", "generation"] @@ -65,6 +67,10 @@ def is_flashinfer_fp4_mla_attention_enabled() -> bool: return _env_enabled(FLASHINFER_FP4_MLA_ATTENTION_ENV) +def _fp4_mla_attention_backend() -> str: + return os.getenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, "triton").lower() + + def _fp4_mla_debug_enabled() -> bool: return _env_enabled(FLASHINFER_FP4_MLA_DEBUG_ENV) @@ -833,6 +839,46 @@ def _max_generation_pages(metadata: Any) -> int: return metadata.num_generation_blocks +def _host_int_list(value: Any, start: int, end: int) -> Optional[list[int]]: + if value is None: + return None + if isinstance(value, torch.Tensor): + if value.is_cuda: + return None + return [int(item) for item in value[start:end].tolist()] + try: + return [int(item) for item in value[start:end]] + except (TypeError, ValueError): + return None + + +def _infer_cutile_assume_full_pages(metadata: Any, max_pages: int, page_size: int) -> bool: + if getattr(metadata, "is_cuda_graph", False): + return False + + start = metadata.num_contexts + end = metadata.num_seqs + kv_cache_params = getattr(metadata, "kv_cache_params", None) + cached_token_lens = _host_int_list( + getattr(kv_cache_params, "num_cached_tokens_per_seq", None), + start, + end, + ) + seq_lens_kv = _host_int_list(getattr(metadata, "seq_lens_kv", None), start, end) + if cached_token_lens is not None and seq_lens_kv is not None: + if len(cached_token_lens) != len(seq_lens_kv): + return False + kv_lens = [ + cached_len + seq_len for cached_len, seq_len in zip(cached_token_lens, seq_lens_kv) + ] + elif kv_cache_params is None: + kv_lens = _host_int_list(getattr(metadata, "prompt_lens_cpu_runtime", None), start, end) + else: + return False + + return bool(kv_lens) and min(kv_lens) == max(kv_lens) == max_pages * page_size + + def run_fp4_mla_attention_decode( metadata: Any, layer_idx: int, @@ -927,6 +973,112 @@ def run_fp4_mla_attention_decode( sf_cache = sf_cache.view(torch.float8_e4m3fn) num_gen_blocks = metadata.num_generation_blocks + v_sf = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=kv_lora_rank)[local_layer].view( + torch.float8_e4m3fn + ) + + src_page_ids = metadata.paged_kv_indices[ + metadata.num_context_blocks : metadata.num_context_blocks + num_gen_blocks + ] + kv_lens = metadata.kv_lens_cuda_runtime[metadata.num_contexts : metadata.num_seqs] + max_pages = _max_generation_pages(metadata) + if max_pages == 0: + return + + backend = _fp4_mla_attention_backend() + if backend == "cutile": + from .fp4_mla_cutile import fp4_mla_paged_attention + + total_p_rows = max(src_page_ids.shape[0] * num_heads, 1) + p_fp4 = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_p_buf", + (total_p_rows, metadata.page_size // 2), + dtype=torch.uint8, + device=q_nope.device, + ) + p_sf = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_p_sf_buf", + (max(_get_fp4_mla_swizzled_scale_size(total_p_rows, metadata.page_size), 1),), + dtype=torch.float8_e4m3fn, + device=q_nope.device, + ) + stats_shape = (num_gen, num_heads) + max_scores = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_max_buf", + stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + denom = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_denom_buf", + stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + page_max = None + page_sum = None + if max_pages >= 8: + page_stats_shape = (num_gen, max_pages, num_heads) + page_max = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_page_max_buf", + page_stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + page_sum = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_page_sum_buf", + page_stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + assume_full_pages = _infer_cutile_assume_full_pages( + metadata, + max_pages, + metadata.page_size, + ) + assume_valid_pages = False + _fp4_mla_debug( + "attention decode cutile launch: " + f"num_gen={num_gen} num_heads={num_heads} local_layer={local_layer} " + f"layer_idx={layer_idx} head_dim={head_dim} kv_lora_rank={kv_lora_rank} " + f"rope_dim={qk_rope_head_dim} max_pages={max_pages} " + f"assume_full_pages={assume_full_pages} " + f"assume_valid_pages={assume_valid_pages}" + ) + fp4_mla_paged_attention( + q_fp4, + q_sf, + kv_cache, + sf_cache, + v_sf, + global_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + output, + sm_scale=float(sm_scale), + num_heads=num_heads, + v_head_dim=kv_lora_rank, + page_size=metadata.page_size, + q_residual_dim=q_residual_dim, + max_pages=max_pages, + assume_full_pages=assume_full_pages, + assume_valid_pages=assume_valid_pages, + p_fp4_workspace=p_fp4, + p_sf_workspace=p_sf, + max_scores_workspace=max_scores, + denom_workspace=denom, + page_max_workspace=page_max, + page_sum_workspace=page_sum, + ) + _debug_sync("attention_cutile") + return total_p_rows = num_gen_blocks * num_heads p_fp4 = _ensure_workspace_tensor( metadata, @@ -942,13 +1094,6 @@ def run_fp4_mla_attention_decode( dtype=torch.float8_e4m3fn, device=q_nope.device, ) - p_probs = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_p_prob_buf", - (max(num_gen * num_heads, 1), metadata.page_size), - dtype=torch.float32, - device=q_nope.device, - )[: num_gen * num_heads] stats_shape = (num_gen, num_heads) max_scores = _ensure_workspace_tensor( metadata, @@ -965,17 +1110,64 @@ def run_fp4_mla_attention_decode( device=q_nope.device, ) - v_sf = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=kv_lora_rank)[local_layer].view( - torch.float8_e4m3fn - ) + if backend == _FP4_MLA_CUTE_DSL_BACKEND: + from .fp4_mla_cute import run_fp4_mla_attention_decode_cute - src_page_ids = metadata.paged_kv_indices[ - metadata.num_context_blocks : metadata.num_context_blocks + num_gen_blocks - ] - kv_lens = metadata.kv_lens_cuda_runtime[metadata.num_contexts : metadata.num_seqs] - max_pages = _max_generation_pages(metadata) - if max_pages == 0: + page_stats_shape = (num_gen, max_pages, num_heads) + page_max = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_page_max_buf", + page_stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + page_sum = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_page_sum_buf", + page_stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + run_fp4_mla_attention_decode_cute( + output=output, + max_scores=max_scores, + denom=denom, + page_max=page_max, + page_sum=page_sum, + p_fp4=p_fp4, + p_sf=p_sf, + q_fp4=q_fp4, + q_sf=q_sf, + kv_cache=kv_cache, + sf_cache=sf_cache, + v_sf=v_sf, + global_scale=global_scale, + src_page_ids=src_page_ids, + paged_kv_indptr_decode=metadata.paged_kv_indptr_decode, + kv_lens=kv_lens, + sm_scale=float(sm_scale), + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + q_residual_dim=q_residual_dim, + page_size=metadata.page_size, + max_pages=max_pages, + ) + _debug_sync("attention_cute_dsl") return + if backend != "triton": + raise ValueError( + f"Unsupported FP4 MLA attention backend '{backend}'. " + f"Set {FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV} to 'triton', " + "'cutile', or 'cute_dsl'." + ) + + p_probs = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_p_prob_buf", + (max(num_gen * num_heads, 1), metadata.page_size), + dtype=torch.float32, + device=q_nope.device, + )[: num_gen * num_heads] block_h = 128 block_t = metadata.page_size diff --git a/tests/unittest/_torch/attention/test_fp4_mla_kv.py b/tests/unittest/_torch/attention/test_fp4_mla_kv.py index 09cb65ec7f63..252c5aced3a5 100644 --- a/tests/unittest/_torch/attention/test_fp4_mla_kv.py +++ b/tests/unittest/_torch/attention/test_fp4_mla_kv.py @@ -16,6 +16,7 @@ import tensorrt_llm from tensorrt_llm._torch.attention_backend.fp4_mla_kv import ( + FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, FLASHINFER_FP4_MLA_ATTENTION_ENV, FP4_BLOCK_SIZE, FP4_MLA_KV_GLOBAL_SCALE, @@ -32,6 +33,7 @@ scatter_fp4_mla_kv_cache, update_hp_kv_for_fp4_mla, ) +from tensorrt_llm._torch.cute_dsl_utils import IS_CUTLASS_DSL_AVAILABLE from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings.executor import KvCacheConfig from tensorrt_llm.mapping import Mapping @@ -392,13 +394,15 @@ def _fp4_mla_attention_decode_reference( global_scale=_TEST_GLOBAL_SCALE, ) - p_dequant = _dequant_fp4_swizzled( - metadata._fp4_mla_attention_p_buf, - metadata._fp4_mla_attention_p_sf_buf, - logical_dim=metadata.page_size, - sf_per_token=metadata.page_size // 16, - global_scale=FP4_MLA_P_GLOBAL_SCALE, - ) + p_dequant = None + if hasattr(metadata, "_fp4_mla_attention_p_buf"): + p_dequant = _dequant_fp4_swizzled( + metadata._fp4_mla_attention_p_buf, + metadata._fp4_mla_attention_p_sf_buf, + logical_dim=metadata.page_size, + sf_per_token=metadata.page_size // 16, + global_scale=FP4_MLA_P_GLOBAL_SCALE, + ) indptr = metadata.paged_kv_indptr_decode.cpu().tolist() kv_lens = metadata.kv_lens_cuda_runtime.cpu().tolist() @@ -413,16 +417,19 @@ def _fp4_mla_attention_decode_reference( q = q_dequant[q_start : q_start + num_heads] probs = torch.softmax(torch.matmul(q, logical_k.transpose(0, 1)) * sm_scale, dim=-1) - p_pages = [] - for page_rel in range(indptr[seq_idx + 1] - indptr[seq_idx]): - page_start = page_rel * metadata.page_size - valid_tokens = max(min(kv_len - page_start, metadata.page_size), 0) - if valid_tokens == 0: - continue - compact_page = indptr[seq_idx] + page_rel - p_start = compact_page * num_heads - p_pages.append(p_dequant[p_start : p_start + num_heads, :valid_tokens]) - p = torch.cat(p_pages, dim=-1) + if p_dequant is None: + p = probs + else: + p_pages = [] + for page_rel in range(indptr[seq_idx + 1] - indptr[seq_idx]): + page_start = page_rel * metadata.page_size + valid_tokens = max(min(kv_len - page_start, metadata.page_size), 0) + if valid_tokens == 0: + continue + compact_page = indptr[seq_idx] + page_rel + p_start = compact_page * num_heads + p_pages.append(p_dequant[p_start : p_start + num_heads, :valid_tokens]) + p = torch.cat(p_pages, dim=-1) exact_probs.append(probs) quantized_probs.append(p) @@ -445,6 +452,127 @@ def _cuda_event_benchmark(fn, *, warmup_iters=10, iters=100): return start.elapsed_time(end) / iters +def _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + *, + backend: str, + num_heads: int, + seq_lens: list[int], + seed: int, + check_probs: bool, +) -> None: + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, backend) + + ( + kv_cache_manager, + metadata, + q_nope, + q_pe, + kv_lora_rank, + qk_rope_head_dim, + ) = _build_fp4_mla_attention_decode_case( + seq_lens=seq_lens, + num_heads=num_heads, + seed=seed, + ) + try: + output = torch.empty_like(q_nope) + sm_scale = 0.1 + run_fp4_mla_attention_decode( + metadata, + layer_idx=0, + local_layer=0, + q_nope=q_nope, + q_pe=q_pe, + output=output, + sm_scale=sm_scale, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + ) + torch.cuda.synchronize() + + ref_output, exact_probs, quantized_probs = _fp4_mla_attention_decode_reference( + metadata, + q_nope, + q_pe, + sm_scale=sm_scale, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + ) + if check_probs: + for seq_idx, (exact_prob, quantized_prob) in enumerate( + zip(exact_probs, quantized_probs) + ): + torch.testing.assert_close( + quantized_prob, + exact_prob, + atol=8e-2, + rtol=8e-2, + msg=f"FP4 MLA attention probabilities diverged for sequence {seq_idx}", + ) + torch.testing.assert_close( + output.float(), + ref_output, + atol=1e-1, + rtol=1e-1, + msg=f"{backend} FP4 MLA attention decode output diverged from reference", + ) + finally: + kv_cache_manager.shutdown() + + +def _ceil_div(lhs: int, rhs: int) -> int: + return (lhs + rhs - 1) // rhs + + +def _estimate_fp4_mla_attention_decode_mbu_bytes( + *, + batch_size: int, + seq_len: int, + num_heads: int, + kv_lora_rank: int, + qk_rope_head_dim: int, + page_size: int, +) -> int: + """Estimate logical global-memory traffic for the FP4 MLA decode benchmark.""" + fp4_block_size = 16 + head_block_size = 128 + bf16_bytes = 2 + fp32_bytes = 4 + + q_input_dim = kv_lora_rank + qk_rope_head_dim + qk_dim = q_input_dim + FP4_MLA_Q_RESIDUAL_DIM + pages_per_seq = _ceil_div(seq_len, page_size) + padded_seq_len = pages_per_seq * page_size + head_blocks = _ceil_div(num_heads, head_block_size) + q_rows = batch_size * num_heads + + qk_fp4_bytes_per_token = qk_dim // 2 + _ceil_div(qk_dim, fp4_block_size) + v_fp4_bytes_per_token = kv_lora_rank // 2 + _ceil_div(kv_lora_rank, fp4_block_size) + p_bytes_per_seq = padded_seq_len // 2 + _ceil_div(padded_seq_len, fp4_block_size) + + q_setup_bytes = q_rows * q_input_dim * bf16_bytes * 3 + q_rows * qk_fp4_bytes_per_token + qk_cache_bytes = 2 * batch_size * head_blocks * padded_seq_len * qk_fp4_bytes_per_token + qk_q_bytes = 2 * batch_size * num_heads * pages_per_seq * qk_fp4_bytes_per_token + stats_bytes = batch_size * num_heads * 2 * fp32_bytes * (1 + pages_per_seq) + p_prob_bytes = batch_size * num_heads * padded_seq_len * fp32_bytes * 2 + p_quant_bytes = 2 * batch_size * num_heads * p_bytes_per_seq + pv_cache_bytes = batch_size * head_blocks * padded_seq_len * v_fp4_bytes_per_token + output_bytes = q_rows * kv_lora_rank * bf16_bytes + + return ( + q_setup_bytes + + qk_cache_bytes + + qk_q_bytes + + stats_bytes + + p_prob_bytes + + p_quant_bytes + + pv_cache_bytes + + output_bytes + ) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @pytest.mark.parametrize( ("num_tokens", "page_size"), @@ -1160,63 +1288,41 @@ def test_fp4_mla_attention_decode_residual_qk_duplicates_k_tail(monkeypatch): @pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 tensor cores") def test_fp4_mla_attention_decode_multi_seq_matches_reference(monkeypatch): """Multiple decode sequences and heads must match a QK-softmax-PV reference.""" - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") - num_heads = 5 - seq_lens = [32, 128] + _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + backend="triton", + num_heads=5, + seq_lens=[32, 128], + seed=7, + check_probs=True, + ) - ( - kv_cache_manager, - metadata, - q_nope, - q_pe, - kv_lora_rank, - qk_rope_head_dim, - ) = _build_fp4_mla_attention_decode_case( - seq_lens=seq_lens, - num_heads=num_heads, + +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +def test_fp4_mla_attention_decode_cutile_matches_reference(monkeypatch): + """CuTile decode backend must preserve the FP4 MLA residual-tail contract.""" + _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + backend="cutile", + num_heads=128, + seq_lens=[32, 128], seed=7, + check_probs=False, ) - try: - output = torch.empty_like(q_nope) - sm_scale = 0.1 - run_fp4_mla_attention_decode( - metadata, - layer_idx=0, - local_layer=0, - q_nope=q_nope, - q_pe=q_pe, - output=output, - sm_scale=sm_scale, - kv_lora_rank=kv_lora_rank, - qk_rope_head_dim=qk_rope_head_dim, - ) - torch.cuda.synchronize() - ref_output, exact_probs, quantized_probs = _fp4_mla_attention_decode_reference( - metadata, - q_nope, - q_pe, - sm_scale=sm_scale, - kv_lora_rank=kv_lora_rank, - qk_rope_head_dim=qk_rope_head_dim, - ) - for seq_idx, (exact_prob, quantized_prob) in enumerate(zip(exact_probs, quantized_probs)): - torch.testing.assert_close( - quantized_prob, - exact_prob, - atol=8e-2, - rtol=8e-2, - msg=f"FP4 MLA attention probabilities diverged for sequence {seq_idx}", - ) - torch.testing.assert_close( - output.float(), - ref_output, - atol=1e-1, - rtol=1e-1, - msg="FP4 MLA attention decode output diverged from the QK-softmax-PV reference", - ) - finally: - kv_cache_manager.shutdown() + +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +@pytest.mark.skipif(not IS_CUTLASS_DSL_AVAILABLE, reason="requires CuTe DSL") +def test_fp4_mla_attention_decode_cute_dsl_matches_reference(monkeypatch): + """CuTe DSL decode backend must preserve the FP4 MLA decode contract.""" + _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + backend="cute_dsl", + num_heads=2, + seq_lens=[32], + seed=10, + check_probs=False, + ) @pytest.mark.skipif( @@ -1237,6 +1343,8 @@ def test_fp4_mla_attention_decode_perf_benchmark( ``TRTLLM_RUN_FP4_MLA_ATTENTION_BENCHMARK=1 pytest -s -k fp4_mla_attention_decode_perf``. """ monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") + backend = os.environ.get(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, "triton") + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, backend) num_heads = 128 ( kv_cache_manager, @@ -1274,11 +1382,21 @@ def run_decode(): pv_dim = kv_lora_rank matmul_flops = 2 * batch_size * num_heads * seq_len * (qk_dim + pv_dim) matmul_tflops = matmul_flops / avg_ms / 1e9 + mbu_bytes = _estimate_fp4_mla_attention_decode_mbu_bytes( + batch_size=batch_size, + seq_len=seq_len, + num_heads=num_heads, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + page_size=metadata.page_size, + ) + mbu_tbps = mbu_bytes / avg_ms / 1e9 print( "\nrun_fp4_mla_attention_decode " - f"batch={batch_size} seq_len={seq_len} heads={num_heads}: " + f"backend={backend} batch={batch_size} seq_len={seq_len} heads={num_heads}: " f"{avg_ms:.4f} ms, {tokens * num_heads / avg_ms / 1e3:.2f} M token-head/s, " - f"{matmul_tflops:.2f} estimated matmul TFLOP/s" + f"{matmul_tflops:.2f} estimated matmul TFLOP/s, " + f"estimated MBU={mbu_tbps:.2f} TB/s" ) assert torch.isfinite(output.float()).all() finally: From c1006af150a3568547df7910cdb4b45a16272db1 Mon Sep 17 00:00:00 2001 From: Tracin <10434017+Tracin@users.noreply.github.com> Date: Mon, 25 May 2026 02:52:33 -0700 Subject: [PATCH 04/11] [None][test] Add FP4 MLA decode benchmark Add a standalone benchmark comparing the FP4 MLA decode backends against FlashInfer BF16 and trtllm-gen baselines. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> --- bench_fp4_mla_decode.py | 221 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 221 insertions(+) create mode 100644 bench_fp4_mla_decode.py diff --git a/bench_fp4_mla_decode.py b/bench_fp4_mla_decode.py new file mode 100644 index 000000000000..46cb8b6e064b --- /dev/null +++ b/bench_fp4_mla_decode.py @@ -0,0 +1,221 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Standalone benchmark for the FP4 MLA decode kernel. + +Compares the Triton, CuTile, and CuTe DSL FP4 backends against the +trtllm-gen bf16 baseline. +Run with: + python bench_fp4_mla_decode.py [--batch B] [--seq S] [--heads H] +""" + +import argparse +import os +import sys + +import torch + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "tests/unittest/_torch/attention")) + +os.environ.setdefault("TRTLLM_FLASHINFER_FP4_MLA_ATTENTION", "1") + +import flashinfer # noqa: E402 +from test_fp4_mla_kv import _build_fp4_mla_attention_decode_case # noqa: E402 + +from tensorrt_llm._torch.attention_backend.fp4_mla_kv import ( # noqa: E402 + FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, + FLASHINFER_FP4_MLA_ATTENTION_ENV, + FP4_MLA_Q_RESIDUAL_DIM, + run_fp4_mla_attention_decode, +) + +BACKEND_CHOICES = ("trtllm", "triton", "cutile", "cute_dsl") + + +def _bench(fn, warmup=10, iters=50): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(iters): + fn() + end.record() + torch.cuda.synchronize() + return start.elapsed_time(end) / iters + + +# B200 HBM3e peak bandwidth (~8 TB/s). Override via TRTLLM_HBM_GB_S= if needed. +HBM_PEAK_GB_S = float(os.environ.get("TRTLLM_HBM_GB_S", "8000")) +FP4_BLOCK_SIZE = 16 # E2M1 scale-group size used by the kernel +PAGE_SIZE = 128 + + +def _kernel_io_bytes(batch, seq, heads, kv_lora_rank, qk_rope_head_dim): + """Per-kernel-call HBM input + output bytes (FP4 = 0.5 B, scale = 1 B, output = 2 B).""" + q_head_dim = kv_lora_rank + qk_rope_head_dim + FP4_MLA_Q_RESIDUAL_DIM + k_head_dim = kv_lora_rank + qk_rope_head_dim + q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE + k_sf_per_token = k_head_dim // FP4_BLOCK_SIZE + pages = (seq + PAGE_SIZE - 1) // PAGE_SIZE + sf_per_page = PAGE_SIZE // FP4_BLOCK_SIZE + + q_fp4 = batch * heads * q_head_dim // 2 + q_sf = batch * heads * q_sf_per_token + kv_cache = batch * seq * k_head_dim // 2 + k_sf_cache = batch * seq * k_sf_per_token + v_sf_cache = batch * pages * kv_lora_rank * sf_per_page + out = batch * heads * kv_lora_rank * 2 # bf16/half + return q_fp4 + q_sf + kv_cache + k_sf_cache + v_sf_cache + out + + +def _bf16_mla_io_bytes(batch, seq, heads, kv_lora_rank, qk_rope_head_dim): + """Per-call HBM bytes for the bf16 MLA decode baseline (2 B/elem).""" + head_dim = kv_lora_rank + qk_rope_head_dim + q = batch * heads * head_dim * 2 + kv = batch * seq * head_dim * 2 # ckv + kpe paged caches + out = batch * heads * kv_lora_rank * 2 + return q + kv + out + + +def run_one_trtllm(batch, seq, heads): + """Bf16 baseline using the FlashInfer trtllm-gen MLA decode kernel.""" + device = torch.device("cuda") + kv_lora_rank = 512 + qk_rope_head_dim = 64 + qk_nope_head_dim = 128 # DeepSeek-V3 default; only used for fused scale convention. + head_dim_qk = kv_lora_rank + qk_rope_head_dim + page_size = 64 # trtllm-gen MLA decode only supports page_size of 32 or 64. + blocks_per_seq = (seq + page_size - 1) // page_size + total_pages = batch * blocks_per_seq + + torch.manual_seed(8) + # query layout: [batch, q_len=1, heads, kv_lora_rank + qk_rope_head_dim] + query = torch.randn(batch, 1, heads, head_dim_qk, dtype=torch.bfloat16, device=device) + # kv_cache layout: [num_pages, page_size, head_dim_ckv + head_dim_kpe] + kv_cache = torch.randn(total_pages, page_size, head_dim_qk, dtype=torch.bfloat16, device=device) + + block_tables = torch.arange(total_pages, dtype=torch.int32, device=device).view( + batch, blocks_per_seq + ) + seq_lens = torch.full((batch,), seq, dtype=torch.int32, device=device) + workspace = torch.zeros(128 * 1024 * 1024, dtype=torch.int8, device=device).view(-1, 4) + + output = torch.empty(batch, 1, heads, kv_lora_rank, dtype=torch.bfloat16, device=device) + + def run(): + flashinfer.mla.trtllm_batch_decode_with_kv_cache_mla( + query=query, + kv_cache=kv_cache, + workspace_buffer=workspace, + qk_nope_head_dim=qk_nope_head_dim, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + block_tables=block_tables, + seq_lens=seq_lens, + max_seq_len=seq, + out=output, + bmm1_scale=0.1, + bmm2_scale=1.0, + ) + + run() + torch.cuda.synchronize() + avg_ms = _bench(run) + qk_dim = kv_lora_rank + qk_rope_head_dim + pv_dim = kv_lora_rank + flops = 2 * batch * heads * seq * (qk_dim + pv_dim) + tflops = flops / avg_ms / 1e9 + bytes_per_call = _bf16_mla_io_bytes(batch, seq, heads, kv_lora_rank, qk_rope_head_dim) + gb_s = bytes_per_call / (avg_ms * 1e-3) / 1e9 + hbm_pct = 100.0 * gb_s / HBM_PEAK_GB_S + print( + f"backend={'trtllm':>12s} bs={batch:>3d} seq={seq:>5d} heads={heads:>3d}: " + f"{avg_ms:>7.3f} ms {tflops:>6.2f} TFLOP/s " + f"HBM {gb_s:>6.1f} GB/s ({hbm_pct:>4.1f}% of {HBM_PEAK_GB_S:.0f})", + flush=True, + ) + return avg_ms + + +def run_one(batch, seq, heads, backend): + os.environ[FLASHINFER_FP4_MLA_ATTENTION_ENV] = "1" + os.environ[FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV] = backend + ( + kv_cache_manager, + metadata, + q_nope, + q_pe, + kv_lora_rank, + qk_rope_head_dim, + ) = _build_fp4_mla_attention_decode_case(seq_lens=[seq] * batch, num_heads=heads, seed=8) + try: + output = torch.empty_like(q_nope) + + def run(): + run_fp4_mla_attention_decode( + metadata, + layer_idx=0, + local_layer=0, + q_nope=q_nope, + q_pe=q_pe, + output=output, + sm_scale=0.1, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + ) + + run() + torch.cuda.synchronize() + avg_ms = _bench(run) + qk_dim = kv_lora_rank + qk_rope_head_dim + FP4_MLA_Q_RESIDUAL_DIM + pv_dim = kv_lora_rank + flops = 2 * batch * heads * seq * (qk_dim + pv_dim) + tflops = flops / avg_ms / 1e9 + # HBM bandwidth = sum of global-mem reads (Q, Q-SF, KV, K-SF, V-SF) + output writes per + # call. Note: typically <1% of peak because L2 absorbs the KV reuse and the actual + # hot lane is SMEM (~85% of peak L1/SMEM throughput per ncu). HBM low ~= good caching. + bytes_per_call = _kernel_io_bytes(batch, seq, heads, kv_lora_rank, qk_rope_head_dim) + gb_s = bytes_per_call / (avg_ms * 1e-3) / 1e9 + hbm_pct = 100.0 * gb_s / HBM_PEAK_GB_S + print( + f"backend={backend:>12s} bs={batch:>3d} seq={seq:>5d} heads={heads:>3d}: " + f"{avg_ms:>7.3f} ms {tflops:>6.2f} TFLOP/s " + f"HBM {gb_s:>6.1f} GB/s ({hbm_pct:>4.1f}% of {HBM_PEAK_GB_S:.0f})", + flush=True, + ) + return avg_ms + finally: + kv_cache_manager.shutdown() + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--batch", type=int, default=None) + p.add_argument("--seq", type=int, default=32768) + p.add_argument("--heads", type=int, default=128) + p.add_argument( + "--backend", + default=None, + choices=BACKEND_CHOICES, + help=("Backend to benchmark; default runs all fast backends."), + ) + args = p.parse_args() + + batches = [args.batch] if args.batch else [16, 32, 64, 128, 256] + if args.backend: + backends = [args.backend] + else: + backends = list(BACKEND_CHOICES) + + for b in batches: + for be in backends: + if be == "trtllm": + run_one_trtllm(b, args.seq, args.heads) + else: + run_one(b, args.seq, args.heads, be) + + +if __name__ == "__main__": + main() From 5946c57a9ec68cb8bc5e27c01d7a5039429b2c5e Mon Sep 17 00:00:00 2001 From: Tracin <10434017+Tracin@users.noreply.github.com> Date: Mon, 25 May 2026 02:55:54 -0700 Subject: [PATCH 05/11] [None][fix] Force CUDA utility paths for FP4 MLA debugging Disable compile wrapping for the copy helper and force FLA utility device selection to CUDA for this debug path. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> --- tensorrt_llm/_torch/modules/fla/utils.py | 1 + tensorrt_llm/_torch/speculative/mtp.py | 8 ++++---- tensorrt_llm/_torch/speculative/one_model_sampler.py | 2 +- tensorrt_llm/_torch/utils.py | 2 +- 4 files changed, 7 insertions(+), 6 deletions(-) diff --git a/tensorrt_llm/_torch/modules/fla/utils.py b/tensorrt_llm/_torch/modules/fla/utils.py index 480051cbcc7d..3a5f9c1725d8 100644 --- a/tensorrt_llm/_torch/modules/fla/utils.py +++ b/tensorrt_llm/_torch/modules/fla/utils.py @@ -290,6 +290,7 @@ def _check_platform() -> Literal["nvidia", "amd", "intel", "musa"]: # However, the torch backend is 'cuda' for both Nvidia and AMD GPUs. # Therefore, we need to check the triton backend to determine the actual GPU vendor. device = get_available_device() if get_available_device() != "hip" else "cuda" +device = "cuda" device_torch_lib = getattr(torch, device) device_platform = _check_platform() diff --git a/tensorrt_llm/_torch/speculative/mtp.py b/tensorrt_llm/_torch/speculative/mtp.py index dda345844c19..8d00e5bd11f3 100644 --- a/tensorrt_llm/_torch/speculative/mtp.py +++ b/tensorrt_llm/_torch/speculative/mtp.py @@ -675,7 +675,7 @@ def unpack_sequence(packed_seq_cuda, seq_lens_cuda, seq_lens_cpu): mtp_past_hidden_states_pool.index_copy_(0, slot_ids, new_mtp_past_hidden_states) - @torch.compile(options={"max-autotune": True}) + # @torch.compile(options={"max-autotune": True}) def topk_kernel(self, gen_logprobs, num_gens, mtp_num_modules, spec_metadata): topk_value, topk_indices = torch.topk(gen_logprobs, @@ -689,7 +689,7 @@ def topk_kernel(self, gen_logprobs, num_gens, mtp_num_modules, num_gens, mtp_num_modules) return topk_value, topk_indices, draft_tokens - @torch.compile(options={"max-autotune": True}) + # @torch.compile(options={"max-autotune": True}) def process_generation_logits(self, logits, num_contexts): gen_logits = logits[num_contexts:] gen_logprobs = torch.softmax(gen_logits, dim=-1) @@ -1070,7 +1070,7 @@ def prepare_drafter_inputs( "attn_metadata": attn_metadata, } - @torch.compile(options={"max-autotune": True}) + # @torch.compile(options={"max-autotune": True}) def get_local_max_and_combined(self, logits, mapping_lm_tp=None): local_max_values, local_argmax = torch.max(logits, dim=-1, keepdim=True) # Adjust indices based on TP rank and size @@ -1089,7 +1089,7 @@ def get_local_max_and_combined(self, logits, mapping_lm_tp=None): dim=-1).flatten(-2) return combined - @torch.compile(options={"max-autotune": True}) + # @torch.compile(options={"max-autotune": True}) def get_draft_tokens_from_gathered(self, gathered): gathered_indices_float = gathered[..., 0::2] # Even positions: indices gathered_values_float = gathered[..., 1::2] # Odd positions: values diff --git a/tensorrt_llm/_torch/speculative/one_model_sampler.py b/tensorrt_llm/_torch/speculative/one_model_sampler.py index 6734b5e9f79a..2cdf238e0a8f 100644 --- a/tensorrt_llm/_torch/speculative/one_model_sampler.py +++ b/tensorrt_llm/_torch/speculative/one_model_sampler.py @@ -91,7 +91,7 @@ def apply_temperature( return logits.div_(temp.unsqueeze(dim=1)) -@torch.compile(options={"max-autotune": True}) +# @torch.compile(options={"max-autotune": True}) def sampling_batch_spec_dec_one_model( logits: torch.Tensor, temperatures: torch.Tensor, diff --git a/tensorrt_llm/_torch/utils.py b/tensorrt_llm/_torch/utils.py index 4e9c92c9ba76..30c90f8bce23 100644 --- a/tensorrt_llm/_torch/utils.py +++ b/tensorrt_llm/_torch/utils.py @@ -476,7 +476,7 @@ def tensor_to_str(x: torch.Tensor, num_elements: int = 10) -> str: ")") -@maybe_compile +# @maybe_compile def maybe_compiled_copy_(dst, src): dst.copy_(src) From 2cb6f174f4ec5fce8a741c1067232fd2a717c0b9 Mon Sep 17 00:00:00 2001 From: Tracin <10434017+Tracin@users.noreply.github.com> Date: Wed, 27 May 2026 00:57:34 -0700 Subject: [PATCH 06/11] Support FP4 MLA linear MTP Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> [None][fix] Avoid CuTe FP4 MLA large KV descriptor overflow Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> Support MPT. Update kernels. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> Support MPT. Update kernels. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> Support MPT. Update kernels. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> Support MPT. Update kernels. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> Support MPT. Update kernels. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> --- bench_fp4_mla_decode.py | 152 +- .../_torch/attention_backend/flashinfer.py | 234 +- .../_torch/attention_backend/fp4_mla.py | 2836 +++++++++++++++ .../_torch/attention_backend/fp4_mla_cute.py | 803 ---- .../attention_backend/fp4_mla_cutile.py | 1597 ++++++-- .../attention_backend/fp4_mla_cutile.py.bak | 3239 +++++++++++++++++ .../attention_backend/fp4_mla_kernels.py | 423 ++- .../_torch/attention_backend/fp4_mla_kv.py | 1530 -------- .../attention_backend/fp4_mla_triton.py | 1536 ++++++++ .../_torch/attention_backend/trtllm.py | 18 +- tensorrt_llm/_torch/modules/fla/utils.py | 7 +- .../_torch/pyexecutor/model_engine.py | 56 + .../_torch/pyexecutor/py_executor_creator.py | 33 +- tensorrt_llm/_torch/speculative/eagle3.py | 36 +- tensorrt_llm/_torch/speculative/mtp.py | 35 +- tensorrt_llm/_torch/speculative/utils.py | 4 + tensorrt_llm/llmapi/llm_args.py | 22 +- .../{test_fp4_mla_kv.py => test_fp4_mla.py} | 169 +- 18 files changed, 9922 insertions(+), 2808 deletions(-) create mode 100644 tensorrt_llm/_torch/attention_backend/fp4_mla.py delete mode 100644 tensorrt_llm/_torch/attention_backend/fp4_mla_cute.py create mode 100644 tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py.bak delete mode 100644 tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py create mode 100644 tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py rename tests/unittest/_torch/attention/{test_fp4_mla_kv.py => test_fp4_mla.py} (90%) diff --git a/bench_fp4_mla_decode.py b/bench_fp4_mla_decode.py index 46cb8b6e064b..109d134ddc15 100644 --- a/bench_fp4_mla_decode.py +++ b/bench_fp4_mla_decode.py @@ -3,10 +3,13 @@ """Standalone benchmark for the FP4 MLA decode kernel. -Compares the Triton, CuTile, and CuTe DSL FP4 backends against the -trtllm-gen bf16 baseline. +Compares the Triton and CuTile FP4 backends against the +trtllm-gen fp8 ("trtllm_fp8") baseline. Run with: - python bench_fp4_mla_decode.py [--batch B] [--seq S] [--heads H] + python bench_fp4_mla_decode.py [--batch B] [--seq S] [--heads H] [--q-len Q] + +--q-len (alias --mtp-len) sets the number of query tokens per sequence (>1 for +MTP / speculative decoding); it defaults to 1 (plain decode). """ import argparse @@ -20,19 +23,23 @@ os.environ.setdefault("TRTLLM_FLASHINFER_FP4_MLA_ATTENTION", "1") import flashinfer # noqa: E402 -from test_fp4_mla_kv import _build_fp4_mla_attention_decode_case # noqa: E402 +from test_fp4_mla import _build_fp4_mla_attention_decode_case # noqa: E402 -from tensorrt_llm._torch.attention_backend.fp4_mla_kv import ( # noqa: E402 +from tensorrt_llm._torch.attention_backend.fp4_mla import ( # noqa: E402 FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, FLASHINFER_FP4_MLA_ATTENTION_ENV, FP4_MLA_Q_RESIDUAL_DIM, run_fp4_mla_attention_decode, ) -BACKEND_CHOICES = ("trtllm", "triton", "cutile", "cute_dsl") +BACKEND_CHOICES = ( + "trtllm_fp8", + "triton", + "cutile", +) -def _bench(fn, warmup=10, iters=50): +def _bench(fn, warmup=0, iters=1): for _ in range(warmup): fn() torch.cuda.synchronize() @@ -52,58 +59,94 @@ def _bench(fn, warmup=10, iters=50): PAGE_SIZE = 128 -def _kernel_io_bytes(batch, seq, heads, kv_lora_rank, qk_rope_head_dim): +def _seq_lens_for_batch(batch, seq): + return [seq + 128 * (batch_idx // 10) for batch_idx in range(batch)] + + +def _seq_label(seq_lens): + return f"{seq_lens[0]}-{seq_lens[-1]}" if len(seq_lens) > 1 else str(seq_lens[0]) + + +def _kernel_io_bytes(seq_lens, heads, kv_lora_rank, qk_rope_head_dim, q_len=1): """Per-kernel-call HBM input + output bytes (FP4 = 0.5 B, scale = 1 B, output = 2 B).""" + batch = len(seq_lens) + total_seq = sum(seq_lens) q_head_dim = kv_lora_rank + qk_rope_head_dim + FP4_MLA_Q_RESIDUAL_DIM k_head_dim = kv_lora_rank + qk_rope_head_dim q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE k_sf_per_token = k_head_dim // FP4_BLOCK_SIZE - pages = (seq + PAGE_SIZE - 1) // PAGE_SIZE + pages = sum((seq_len + PAGE_SIZE - 1) // PAGE_SIZE for seq_len in seq_lens) sf_per_page = PAGE_SIZE // FP4_BLOCK_SIZE - q_fp4 = batch * heads * q_head_dim // 2 - q_sf = batch * heads * q_sf_per_token - kv_cache = batch * seq * k_head_dim // 2 - k_sf_cache = batch * seq * k_sf_per_token - v_sf_cache = batch * pages * kv_lora_rank * sf_per_page - out = batch * heads * kv_lora_rank * 2 # bf16/half + q_fp4 = batch * q_len * heads * q_head_dim // 2 + q_sf = batch * q_len * heads * q_sf_per_token + kv_cache = total_seq * k_head_dim // 2 + k_sf_cache = total_seq * k_sf_per_token + v_sf_cache = pages * kv_lora_rank * sf_per_page + out = batch * q_len * heads * kv_lora_rank * 2 # bf16/half return q_fp4 + q_sf + kv_cache + k_sf_cache + v_sf_cache + out -def _bf16_mla_io_bytes(batch, seq, heads, kv_lora_rank, qk_rope_head_dim): - """Per-call HBM bytes for the bf16 MLA decode baseline (2 B/elem).""" +def _trtllm_mla_io_bytes(seq_lens, heads, kv_lora_rank, qk_rope_head_dim, q_len=1, elem_bytes=2): + """Per-call HBM bytes for the trtllm-gen MLA decode baseline. + + ``elem_bytes`` is the byte width of the Q and KV-cache elements (2 for bf16, + 1 for fp8). The output is always written as bf16 (2 B/elem). + """ + batch = len(seq_lens) + total_seq = sum(seq_lens) head_dim = kv_lora_rank + qk_rope_head_dim - q = batch * heads * head_dim * 2 - kv = batch * seq * head_dim * 2 # ckv + kpe paged caches - out = batch * heads * kv_lora_rank * 2 + q = batch * q_len * heads * head_dim * elem_bytes + kv = total_seq * head_dim * elem_bytes # ckv + kpe paged caches + out = batch * q_len * heads * kv_lora_rank * 2 return q + kv + out -def run_one_trtllm(batch, seq, heads): - """Bf16 baseline using the FlashInfer trtllm-gen MLA decode kernel.""" +def run_one_trtllm(batch, seq, heads, q_len=1): + """Fp8 baseline using the FlashInfer trtllm-gen MLA decode kernel. + + Feeds fp8 (e4m3) Q and KV cache so the kernel uses fp8 tensor cores + (output stays bf16). + """ device = torch.device("cuda") + label = "trtllm_fp8" + io_dtype = torch.float8_e4m3fn + elem_bytes = 1 kv_lora_rank = 512 qk_rope_head_dim = 64 qk_nope_head_dim = 128 # DeepSeek-V3 default; only used for fused scale convention. head_dim_qk = kv_lora_rank + qk_rope_head_dim page_size = 64 # trtllm-gen MLA decode only supports page_size of 32 or 64. - blocks_per_seq = (seq + page_size - 1) // page_size - total_pages = batch * blocks_per_seq + seq_lens_list = _seq_lens_for_batch(batch, seq) + max_seq = max(seq_lens_list) + total_seq = sum(seq_lens_list) + blocks_per_seq = [(seq_len + page_size - 1) // page_size for seq_len in seq_lens_list] + max_blocks_per_seq = max(blocks_per_seq) + total_pages = sum(blocks_per_seq) torch.manual_seed(8) - # query layout: [batch, q_len=1, heads, kv_lora_rank + qk_rope_head_dim] - query = torch.randn(batch, 1, heads, head_dim_qk, dtype=torch.bfloat16, device=device) + # query layout: [batch, q_len, heads, kv_lora_rank + qk_rope_head_dim] + query = torch.randn(batch, q_len, heads, head_dim_qk, dtype=torch.bfloat16, device=device) # kv_cache layout: [num_pages, page_size, head_dim_ckv + head_dim_kpe] kv_cache = torch.randn(total_pages, page_size, head_dim_qk, dtype=torch.bfloat16, device=device) - - block_tables = torch.arange(total_pages, dtype=torch.int32, device=device).view( - batch, blocks_per_seq - ) - seq_lens = torch.full((batch,), seq, dtype=torch.int32, device=device) + # Quantize to fp8 e4m3 (randn ~ N(0,1) is well within e4m3 range). + query = query.to(io_dtype) + kv_cache = kv_cache.to(io_dtype) + + block_tables = torch.zeros((batch, max_blocks_per_seq), dtype=torch.int32, device=device) + page_start = 0 + for batch_idx, num_blocks in enumerate(blocks_per_seq): + block_tables[batch_idx, :num_blocks] = torch.arange( + page_start, page_start + num_blocks, dtype=torch.int32, device=device + ) + page_start += num_blocks + seq_lens = torch.tensor(seq_lens_list, dtype=torch.int32, device=device) workspace = torch.zeros(128 * 1024 * 1024, dtype=torch.int8, device=device).view(-1, 4) - output = torch.empty(batch, 1, heads, kv_lora_rank, dtype=torch.bfloat16, device=device) + output = torch.empty(batch, q_len, heads, kv_lora_rank, dtype=torch.bfloat16, device=device) + # bmm1_scale folds q_scale * k_scale * sm_scale / sqrt(head_dim_qk); a + # representative value (q_scale = k_scale = 1.0 for the fp8 unit-scale tensors). def run(): flashinfer.mla.trtllm_batch_decode_with_kv_cache_mla( query=query, @@ -114,10 +157,11 @@ def run(): qk_rope_head_dim=qk_rope_head_dim, block_tables=block_tables, seq_lens=seq_lens, - max_seq_len=seq, + max_seq_len=max_seq, out=output, bmm1_scale=0.1, bmm2_scale=1.0, + backend="trtllm-gen", ) run() @@ -125,13 +169,16 @@ def run(): avg_ms = _bench(run) qk_dim = kv_lora_rank + qk_rope_head_dim pv_dim = kv_lora_rank - flops = 2 * batch * heads * seq * (qk_dim + pv_dim) + flops = 2 * q_len * heads * total_seq * (qk_dim + pv_dim) tflops = flops / avg_ms / 1e9 - bytes_per_call = _bf16_mla_io_bytes(batch, seq, heads, kv_lora_rank, qk_rope_head_dim) + bytes_per_call = _trtllm_mla_io_bytes( + seq_lens_list, heads, kv_lora_rank, qk_rope_head_dim, q_len, elem_bytes + ) gb_s = bytes_per_call / (avg_ms * 1e-3) / 1e9 hbm_pct = 100.0 * gb_s / HBM_PEAK_GB_S print( - f"backend={'trtllm':>12s} bs={batch:>3d} seq={seq:>5d} heads={heads:>3d}: " + f"backend={label:>12s} bs={batch:>3d} seq={_seq_label(seq_lens_list):>11s} " + f"heads={heads:>3d} qlen={q_len:>2d}: " f"{avg_ms:>7.3f} ms {tflops:>6.2f} TFLOP/s " f"HBM {gb_s:>6.1f} GB/s ({hbm_pct:>4.1f}% of {HBM_PEAK_GB_S:.0f})", flush=True, @@ -139,9 +186,11 @@ def run(): return avg_ms -def run_one(batch, seq, heads, backend): +def run_one(batch, seq, heads, backend, q_len=1): os.environ[FLASHINFER_FP4_MLA_ATTENTION_ENV] = "1" os.environ[FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV] = backend + seq_lens = _seq_lens_for_batch(batch, seq) + total_seq = sum(seq_lens) ( kv_cache_manager, metadata, @@ -149,7 +198,9 @@ def run_one(batch, seq, heads, backend): q_pe, kv_lora_rank, qk_rope_head_dim, - ) = _build_fp4_mla_attention_decode_case(seq_lens=[seq] * batch, num_heads=heads, seed=8) + ) = _build_fp4_mla_attention_decode_case( + seq_lens=seq_lens, num_heads=heads, seed=8, query_len_per_seq=q_len + ) try: output = torch.empty_like(q_nope) @@ -171,16 +222,17 @@ def run(): avg_ms = _bench(run) qk_dim = kv_lora_rank + qk_rope_head_dim + FP4_MLA_Q_RESIDUAL_DIM pv_dim = kv_lora_rank - flops = 2 * batch * heads * seq * (qk_dim + pv_dim) + flops = 2 * q_len * heads * total_seq * (qk_dim + pv_dim) tflops = flops / avg_ms / 1e9 # HBM bandwidth = sum of global-mem reads (Q, Q-SF, KV, K-SF, V-SF) + output writes per # call. Note: typically <1% of peak because L2 absorbs the KV reuse and the actual # hot lane is SMEM (~85% of peak L1/SMEM throughput per ncu). HBM low ~= good caching. - bytes_per_call = _kernel_io_bytes(batch, seq, heads, kv_lora_rank, qk_rope_head_dim) + bytes_per_call = _kernel_io_bytes(seq_lens, heads, kv_lora_rank, qk_rope_head_dim, q_len) gb_s = bytes_per_call / (avg_ms * 1e-3) / 1e9 hbm_pct = 100.0 * gb_s / HBM_PEAK_GB_S print( - f"backend={backend:>12s} bs={batch:>3d} seq={seq:>5d} heads={heads:>3d}: " + f"backend={backend:>12s} bs={batch:>3d} seq={_seq_label(seq_lens):>11s} " + f"heads={heads:>3d} qlen={q_len:>2d}: " f"{avg_ms:>7.3f} ms {tflops:>6.2f} TFLOP/s " f"HBM {gb_s:>6.1f} GB/s ({hbm_pct:>4.1f}% of {HBM_PEAK_GB_S:.0f})", flush=True, @@ -193,8 +245,16 @@ def run(): def main(): p = argparse.ArgumentParser() p.add_argument("--batch", type=int, default=None) - p.add_argument("--seq", type=int, default=32768) + p.add_argument("--seq", type=int, default=30080) p.add_argument("--heads", type=int, default=128) + p.add_argument( + "--q-len", + "--mtp-len", + dest="q_len", + type=int, + default=1, + help="Query tokens per sequence (>1 for MTP / speculative decoding).", + ) p.add_argument( "--backend", default=None, @@ -203,7 +263,7 @@ def main(): ) args = p.parse_args() - batches = [args.batch] if args.batch else [16, 32, 64, 128, 256] + batches = [args.batch] if args.batch else [16, 30, 60, 120, 200, 300] if args.backend: backends = [args.backend] else: @@ -211,10 +271,10 @@ def main(): for b in batches: for be in backends: - if be == "trtllm": - run_one_trtllm(b, args.seq, args.heads) + if be == "trtllm_fp8": + run_one_trtllm(b, args.seq, args.heads, args.q_len) else: - run_one(b, args.seq, args.heads, be) + run_one(b, args.seq, args.heads, be, args.q_len) if __name__ == "__main__": diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index fff0a90b663d..a1508e69b87b 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -19,11 +19,12 @@ from ..metadata import KVCacheParams from ..utils import get_global_attrs, get_model_extra_attrs -from .fp4_mla_kv import (FP4_MLA_KV_GLOBAL_SCALE, HP_BLOCK_SIZE, - get_fp4_mla_decode_cache, - is_flashinfer_fp4_mla_attention_enabled, - run_fp4_mla_attention_decode, scatter_fp4_mla_kv_cache, - update_hp_kv_for_fp4_mla) +from .fp4_mla import (FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, + FP4_MLA_KV_GLOBAL_SCALE, HP_BLOCK_SIZE, + get_fp4_mla_decode_cache, + is_flashinfer_fp4_mla_attention_enabled, + run_fp4_mla_attention_decode, scatter_fp4_mla_kv_cache, + update_hp_kv_for_fp4_mla, update_page_stage_for_fp4_mla) from .interface import (AttentionBackend, AttentionForwardArgs, AttentionInputType, AttentionMetadata, CustomAttentionMask, MLAParams, PredefinedAttentionMask, @@ -183,6 +184,25 @@ class FlashInferAttentionMetadata(AttentionMetadata): # uses it to requantize the active 16-token FP4 KV tile. high_precision_kv_pool: Optional[torch.Tensor] = field(init=False, default=None) + fp4_mla_hp_snapshot_pool: Optional[torch.Tensor] = field(init=False, + default=None) + # BF16 staging buffer for the per-page dynamic-scale no-dequant path, + # shape [max_num_sequences, num_local_layers, kv_factor=1, + # page_size * head_dim]. Unlike high_precision_kv_pool (which buffers only + # the trailing 16-token NVFP4 V-block), this holds the whole in-progress + # page so the active page can be re-quantized to FP4 with the exact + # per-page amax / global scale each step. Allocated only when the + # FlashInfer no-dequant FP4 MLA path is active. + fp4_mla_page_stage_pool: Optional[torch.Tensor] = field(init=False, + default=None) + fp4_mla_page_stage_snapshot_pool: Optional[torch.Tensor] = field( + init=False, default=None) + # Per-page (K/V shared) FP4 global scale for the dynamic-scale path, shape + # [num_local_layers, num_physical_pages] fp32. page_gscale = 448*6/page_amax + # is baked into each page's K and V block scales at re-quant time and undone + # at read. Untouched pages stay 1.0 (== the previous static behaviour). + fp4_mla_page_scale_pool: Optional[torch.Tensor] = field(init=False, + default=None) # Auxiliary FP4 MLA V-scale pool for the no-dequant PV path. The # physical storage is flat per [local_layer, physical_page]; callers view # it with get_fp4_mla_v_scale_pool_view(..., v_head_dim=kv_lora_rank). @@ -679,7 +699,7 @@ def _post_init_with_buffers(self, buffers) -> None: def _allocate_fp4_mla_buffers(self, buffers, capture_graph: bool) -> None: """Allocate the HP BF16 KV pool, seq_slots, and runtime-alias backing - buffers used by ``fp4_mla_kv.update_hp_kv_for_fp4_mla``.""" + buffers used by ``fp4_mla.update_hp_kv_for_fp4_mla``.""" max_num_sequences = (self.max_num_sequences if self.max_num_sequences is not None else self.max_num_requests) @@ -723,6 +743,16 @@ def _allocate_fp4_mla_buffers(self, buffers, capture_graph: bool) -> None: dtype=torch.bfloat16, capture_graph=capture_graph, ) + if capture_graph: + self.fp4_mla_hp_snapshot_pool = self.get_empty( + buffers, + hp_pool_shape, + cache_name="fp4_mla_hp_snapshot_pool", + dtype=torch.bfloat16, + capture_graph=capture_graph, + ) + else: + self.fp4_mla_hp_snapshot_pool = None max_num_pages = self.kv_cache_manager.blocks_in_primary_pool self._fp4_mla_decode_kv_indices_buf = self.get_empty( @@ -749,6 +779,56 @@ def _allocate_fp4_mla_buffers(self, buffers, capture_graph: bool) -> None: "FP4 MLA attention requires the C++ KV cache manager to " "allocate the V-scale pool.") + # Per-page dynamic-scale staging buffer: holds the in-progress page + # in BF16 so it can be re-quantized to FP4 with the exact per-page + # amax. Sized to one full page (tokens_per_block) instead of the + # HP pool's 16-token NVFP4 V-block. + page_stage_slots = self.kv_cache_manager.tokens_per_block + stage_pool_shape = [ + max_num_sequences, num_local_layers, kv_factor, + page_stage_slots * head_dim + ] + existing_stage_pool = self.fp4_mla_page_stage_pool + if (capture_graph and existing_stage_pool is not None + and existing_stage_pool.dtype == torch.bfloat16 + and existing_stage_pool.device.type == "cuda" + and len(existing_stage_pool.shape) == len(stage_pool_shape) + and all(existing_stage_pool.shape[idx] >= dim + for idx, dim in enumerate(stage_pool_shape))): + # Persistent seq-slot state: CUDA graph metadata shares it + # rather than reserving a fresh pool per captured graph. + self.fp4_mla_page_stage_pool = existing_stage_pool + else: + self.fp4_mla_page_stage_pool = self.get_empty( + buffers, + stage_pool_shape, + cache_name="fp4_mla_page_stage_pool", + dtype=torch.bfloat16, + capture_graph=capture_graph, + ) + if capture_graph: + self.fp4_mla_page_stage_snapshot_pool = self.get_empty( + buffers, + stage_pool_shape, + cache_name="fp4_mla_page_stage_snapshot_pool", + dtype=torch.bfloat16, + capture_graph=capture_graph, + ) + else: + self.fp4_mla_page_stage_snapshot_pool = None + + # Per-page (K/V shared) FP4 global scale, one fp32 per physical + # page per local layer. Initialised to 1.0 so untouched pages and + # the read path match the previous static-scale behaviour. + self.fp4_mla_page_scale_pool = self.get_empty( + buffers, + [num_local_layers, max_num_pages], + cache_name="fp4_mla_page_scale_pool", + dtype=torch.float32, + capture_graph=capture_graph, + ) + self.fp4_mla_page_scale_pool.fill_(1.0) + # Runtime-alias backing buffers: GPU for kv/prompt lens, CPU pinned for # the helper's prompt_lens_cpu read. self._kv_lens_cuda_buf = self.get_empty( @@ -844,11 +924,14 @@ def _populate_fp4_mla_batch_indices_positions(self) -> None: kv_start = torch.repeat_interleave(kv_token_starts, seq_lens, output_size=self.num_tokens) - cached_start = torch.repeat_interleave( - self.cached_token_lens[:num_seqs].to(torch.int32), - seq_lens, - output_size=self.num_tokens, - ) + if self.kv_lens_cuda_runtime is not None: + cached_token_lens = self.kv_lens_cuda_runtime[:num_seqs] - seq_lens + else: + cached_token_lens = self.cached_token_lens[:num_seqs].to( + torch.int32) + cached_start = torch.repeat_interleave(cached_token_lens, + seq_lens, + output_size=self.num_tokens) token_offsets = torch.arange(self.num_tokens, dtype=torch.int32, device=device) @@ -858,6 +941,105 @@ def _populate_fp4_mla_batch_indices_positions(self) -> None: non_blocking=True) self._positions[:self.num_tokens].copy_(positions, non_blocking=True) + def repage_fp4_mla_decode_from_kv_lens(self) -> None: + """Rebuild the read-side decode paging from the corrected kv_lens. + + The overlap scheduler builds the generation metadata from the + all-draft-accepted over-estimate; ``_preprocess_inputs`` then corrects + ``kv_lens_cuda_runtime``. ``prepare()`` derived ``num_blocks`` / + ``num_generation_blocks`` / ``paged_kv_indptr_decode`` / + ``paged_kv_indices`` / ``paged_kv_last_page_len`` from the over-estimate. + This recomputes the generation slice of that paging from the corrected + kv_lens so the decode kernels that index the page table see lengths + consistent with what attention actually masks to. Opt-in + (``TRTLLM_FP4_MLA_OVERLAP_REPAGE``); host-syncs and is skipped under + CUDA-graph capture. + """ + if os.getenv("TRTLLM_FP4_MLA_OVERLAP_REPAGE", + "0").lower() not in ("1", "true", "yes", "on"): + return + if self.high_precision_kv_pool is None or self.kv_lens_cuda_runtime is None: + return + # num_blocks / num_generation_blocks are python scalars consumed when + # shaping kernel launches; mutating them cannot affect an already + # captured graph, so skip (and avoid the per-step host sync) under + # CUDA graphs and capture. + if getattr(self, "is_cuda_graph", False): + return + if torch.cuda.is_current_stream_capturing(): + return + num_contexts = self.num_contexts + num_seqs = self.num_contexts + self.num_generations + num_gen = num_seqs - num_contexts + if num_gen <= 0 or not self.num_blocks: + return + + page_size = self.page_size + kv_lens_gen = self.kv_lens_cuda_runtime[num_contexts:num_seqs].detach( + ).to("cpu", torch.int64).tolist() + new_blocks_gen = [(kv + page_size - 1) // page_size + for kv in kv_lens_gen] + old_blocks_gen = [ + int(b) for b in self.num_blocks[num_contexts:num_seqs] + ] + if new_blocks_gen == old_blocks_gen: + return # No page-boundary over-estimate this step. + + # Compact the generation page-id slice: per seq keep its first new_b + # block ids (corrected kv_len <= over-estimate => new_b <= old_b). + ctx_blocks = self.num_context_blocks + old_gen_total = sum(old_blocks_gen) + old_gen = self._paged_kv_indices[ctx_blocks:ctx_blocks + + old_gen_total].detach().to( + "cpu", torch.int64).tolist() + new_gen: list[int] = [] + off = 0 + for old_b, new_b in zip(old_blocks_gen, new_blocks_gen): + nb = min(new_b, old_b) + new_gen.extend(old_gen[off:off + nb]) + off += old_b + if new_gen: + self._paged_kv_indices[ctx_blocks:ctx_blocks + len(new_gen)].copy_( + torch.tensor(new_gen, dtype=self._paged_kv_indices.dtype), + non_blocking=False) + + for i, new_b in enumerate(new_blocks_gen): + self.num_blocks[num_contexts + i] = new_b + self.num_generation_blocks = sum(new_blocks_gen) + + indptr = [0] + for b in new_blocks_gen: + indptr.append(indptr[-1] + b) + self.paged_kv_indptr_decode[:len(indptr)].copy_(torch.tensor( + indptr, dtype=self.paged_kv_indptr_decode.dtype), + non_blocking=False) + + last_page = [ + kv - (b - 1) * page_size + for kv, b in zip(kv_lens_gen, new_blocks_gen) + ] + self._paged_kv_last_page_len[num_contexts:num_seqs].copy_( + torch.tensor(last_page, dtype=self._paged_kv_last_page_len.dtype), + non_blocking=False) + + def update_for_spec_dec(self) -> None: + if self.high_precision_kv_pool is None: + return + if self.kv_lens_cuda_runtime is None: + return + + num_seqs = self.num_seqs + prompt_lens = self.seq_lens_kv_cuda[:num_seqs].to(torch.int32) + self._prompt_lens_cuda_buf[:num_seqs].copy_(prompt_lens, + non_blocking=True) + self.prompt_lens_cuda_runtime = self._prompt_lens_cuda_buf[:num_seqs] + self._prompt_lens_cpu_buf[:num_seqs].copy_(prompt_lens.cpu(), + non_blocking=False) + self.prompt_lens_cpu_runtime = self._prompt_lens_cpu_buf[:num_seqs] + + if self.num_tokens > 0: + self._populate_fp4_mla_batch_indices_positions() + def create_cuda_graph_metadata(self, max_batch_size: int, sub_cross_metadata: bool = False, @@ -1044,7 +1226,7 @@ def prepare(self) -> None: kv_lens = self.cached_token_lens + self.seq_lens_kv_cuda # Populate runtime aliases consumed by the shared HP-pool update helper - # (fp4_mla_kv.update_hp_kv_for_fp4_mla). Only active when MLA + NVFP4. + # (fp4_mla.update_hp_kv_for_fp4_mla). Only active when MLA + NVFP4. if self.high_precision_kv_pool is not None: self._populate_fp4_mla_runtime_aliases(kv_lens) @@ -1731,6 +1913,10 @@ def _mla_forward_context( latent_cache, self._local_layer_idx(metadata), phase="context") + update_page_stage_for_fp4_mla(metadata, + latent_cache, + self._local_layer_idx(metadata), + phase="context") else: ckv_cache, kpe_cache = self._get_mla_caches(metadata) @@ -1806,11 +1992,13 @@ def _mla_forward_generation( f"shape [num_gen_tokens={num_gen_tokens}, ...] but got " f"{list(latent_cache.shape)}. Did MLA.forward_impl stop " f"pre-slicing latent_cache per phase?") + if (use_fp4_attention + and num_gen_tokens != metadata.num_generations + and os.getenv( + FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, + "triton").lower() not in ("triton", "cutile")): + use_fp4_attention = False if use_fp4_attention: - update_hp_kv_for_fp4_mla(metadata, - latent_cache, - self._local_layer_idx(metadata), - phase="generation") scatter_fp4_mla_kv_cache( metadata, latent_cache, @@ -1820,6 +2008,15 @@ def _mla_forward_generation( local_layer=self._local_layer_idx(metadata), v_head_dim=self.kv_lora_rank, ) + update_hp_kv_for_fp4_mla(metadata, + latent_cache, + self._local_layer_idx(metadata), + phase="generation") + update_page_stage_for_fp4_mla( + metadata, + latent_cache, + self._local_layer_idx(metadata), + phase="generation") else: scatter_fp4_mla_kv_cache( metadata, @@ -1831,6 +2028,11 @@ def _mla_forward_generation( latent_cache, self._local_layer_idx(metadata), phase="generation") + update_page_stage_for_fp4_mla( + metadata, + latent_cache, + self._local_layer_idx(metadata), + phase="generation") if not use_fp4_attention: combined_cache = get_fp4_mla_decode_cache( metadata, diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla.py b/tensorrt_llm/_torch/attention_backend/fp4_mla.py new file mode 100644 index 000000000000..4006a588fde1 --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla.py @@ -0,0 +1,2836 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Shared MLA FP4 KV-cache helpers. + +The high-precision (HP) BF16 KV pool is a standalone circular buffer used +alongside the paged FP4 KV pool when MLA models run with NVFP4 KV cache. +Each sequence slot stores the ``HP_BLOCK_SIZE`` most-recent latent vectors at +BF16, so attention backends can consult BF16 values for the tail tokens that +do not yet fill a complete FP4 quant block of 16 elements along the sequence +dimension. + +Used by both ``TrtllmAttention`` (via an internal C++ attention op that reads +both pools) and ``FlashInferAttention`` (via either explicit Python-side +dequant into a BF16 workspace before calling FlashInfer MLA wrappers, or an +env-gated Triton attention path that reads packed FP4 Q, K, and V directly). +""" + +import os +from typing import Any, Literal, Optional + +import torch +import triton +import triton.language as tl + +from tensorrt_llm.logger import logger + +from .fp4_mla_kernels import ( + _fp4_mla_dequant_kernel, + _fp4_mla_overlay_hp_tail_kernel, + _fp4_mla_scatter_kernel, + _fp4_mla_v_scale_store_context_tokens_kernel, + _fp4_mla_v_scale_store_generation_tiles_kernel, + _hp_kv_restore_rejected_from_pool_kernel, + _hp_kv_restore_rejected_from_values_kernel, + _hp_kv_store_context_kernel, + _hp_kv_store_gen_kernel, +) + +HP_BLOCK_SIZE: int = 16 +FP4_BLOCK_SIZE: int = 16 +FP4_MLA_TOKENS_PER_BLOCK: int = 128 +FP4_MLA_SCALE_ROW_GROUP: int = 128 +FP4_MLA_SCALE_COL_GROUP: int = 4 +FP4_MLA_KV_GLOBAL_SCALE: float = 448.0 * 6.0 / (448.0 * 6.0) +FP4_MLA_P_GLOBAL_SCALE: float = 448.0 * 6.0 +FP4_MLA_Q_RESIDUAL_DIM: int = 64 +FLASHINFER_FP4_MLA_ATTENTION_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION" +FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION_BACKEND" +FLASHINFER_FP4_MLA_DEBUG_ENV = "TRTLLM_FLASHINFER_FP4_MLA_DEBUG" +# Opt-in (triton only): re-quantize the active decode page each step +# from the BF16 staging buffer with an exact per-page (K/V shared) global scale, +# instead of the static FP4_MLA_KV_GLOBAL_SCALE. Off by default so the existing +# static-scale path is unchanged until validated. +FP4_MLA_PER_PAGE_SCALE_ENV = "TRTLLM_FP4_MLA_PER_PAGE_SCALE" +# Diagnostic (overlap + MTP debugging): assert that the read-side decode paging +# (num_blocks / num_generation_blocks / paged_kv_indptr_decode) and the gen +# token positions are consistent with the (corrected) kv_lens_cuda_runtime. +FP4_MLA_DEBUG_ASSERT_ENV = "TRTLLM_FP4_MLA_DEBUG_ASSERT" +# Fix (opt-in): rebuild the read-side decode paging from the corrected kv_lens +# in _preprocess_inputs after the overlap kv_lens correction. +FP4_MLA_OVERLAP_REPAGE_ENV = "TRTLLM_FP4_MLA_OVERLAP_REPAGE" +_HPUpdatePhase = Literal["all", "context", "generation"] +_FP4_MLA_MTP_HP_SNAPSHOTS = "_fp4_mla_mtp_hp_snapshots" +# Separate MTP snapshot store for the per-page dynamic-scale staging buffer +# (fp4_mla_page_stage_pool). Kept distinct from the HP-pool snapshots so the +# two rollback paths never alias. +_FP4_MLA_MTP_STAGE_SNAPSHOTS = "_fp4_mla_mtp_stage_snapshots" + + +# Environment and debug helpers + + +def _env_enabled(name: str) -> bool: + return os.getenv(name, "0").lower() in ( + "1", + "true", + "yes", + "on", + ) + + +def _env_int(name: str) -> Optional[int]: + value = os.environ.get(name) + if value is None or value == "": + return None + return int(value) + + +def is_flashinfer_fp4_mla_attention_enabled() -> bool: + """Return whether FlashInfer MLA should allocate no-dequant FP4 attention buffers.""" + return _env_enabled(FLASHINFER_FP4_MLA_ATTENTION_ENV) + + +def _fp4_mla_attention_backend() -> str: + return os.getenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, "triton").lower() + + +def fp4_mla_per_page_scale_enabled() -> bool: + """Return whether the per-page dynamic FP4 global-scale path is on. + + Only the ``triton`` backend implements the matching read side, so the + per-page store + dynamic Q scale are gated to it; on any other backend the + flag is ignored (the static FP4_MLA_KV_GLOBAL_SCALE path runs unchanged). + """ + return _env_enabled(FP4_MLA_PER_PAGE_SCALE_ENV) and _fp4_mla_attention_backend() == "triton" + + +def _fp4_mla_debug_enabled() -> bool: + return _env_enabled(FLASHINFER_FP4_MLA_DEBUG_ENV) + + +def _fp4_mla_debug(message: str) -> None: + if _fp4_mla_debug_enabled(): + print(f"[fp4_mla_debug] {message}", flush=True) + + +def _tensor_layout(tensor: Optional[torch.Tensor]) -> str: + if tensor is None: + return "None" + return ( + f"shape={list(tensor.shape)} stride={list(tensor.stride())} " + f"dtype={tensor.dtype} device={tensor.device}" + ) + + +def _debug_tensor_range(name: str, tensor: Optional[torch.Tensor]) -> None: + if not _fp4_mla_debug_enabled(): + return + if tensor is None: + _fp4_mla_debug(f"{name}: None") + return + flat = tensor.detach().reshape(-1) + if flat.numel() == 0: + _fp4_mla_debug(f"{name}: empty {_tensor_layout(tensor)}") + return + try: + first = flat[: min(8, flat.numel())].cpu().tolist() + _fp4_mla_debug( + f"{name}: {_tensor_layout(tensor)} n={flat.numel()} " + f"min={flat.min().item()} max={flat.max().item()} first={first}" + ) + except RuntimeError as exc: + _fp4_mla_debug(f"{name}: failed to read range: {exc}") + + +def _debug_sync(label: str) -> None: + if not _fp4_mla_debug_enabled(): + return + if torch.cuda.is_current_stream_capturing(): + _fp4_mla_debug(f"{label}: skip sync during CUDA graph capture") + return + _fp4_mla_debug(f"{label}: synchronize") + torch.cuda.synchronize() + _fp4_mla_debug(f"{label}: sync complete") + + +def _ceil_div(lhs: int, rhs: int) -> int: + return (lhs + rhs - 1) // rhs + + +def _host_int_list_during_forward(value: Any, start: int, end: int) -> Optional[list[int]]: + if torch.cuda.is_current_stream_capturing(): + return None + return _host_int_list(value, start, end) + + +# FP4 MLA scale-layout helpers + + +def get_fp4_mla_v_scale_pool_size(v_head_dim: int, page_size: int) -> int: + """Return elements per page for the swizzled FP4 MLA V-scale pool. + + The PV matmul treats V as a RHS matrix shaped ``[v_head_dim, kv_tokens]``. + NVFP4 block scales therefore group along the token/K axis, not along the + latent dimension as the K-view cache does. The physical layout matches the + Triton block-scaled matmul scale layout: + ``[ceil(v_head_dim / 128), ceil(page_size / 16 / 4), 32, 16]``. + """ + token_scale_cols = _ceil_div(page_size, FP4_BLOCK_SIZE) + row_groups = _ceil_div(v_head_dim, FP4_MLA_SCALE_ROW_GROUP) + col_groups = _ceil_div(token_scale_cols, FP4_MLA_SCALE_COL_GROUP) + return row_groups * col_groups * 32 * 16 + + +def _get_fp4_mla_swizzled_scale_size(rows: int, cols: int) -> int: + scale_cols = _ceil_div(cols, FP4_BLOCK_SIZE) + row_groups = _ceil_div(rows, FP4_MLA_SCALE_ROW_GROUP) + col_groups = _ceil_div(scale_cols, FP4_MLA_SCALE_COL_GROUP) + return row_groups * col_groups * 32 * 16 + + +def _use_fp4_mla_swizzled_sf() -> bool: + return is_flashinfer_fp4_mla_attention_enabled() + + +def _get_fp4_mla_context_start_positions(metadata: Any, num_contexts: int) -> torch.Tensor: + kv_cache_params = getattr(metadata, "kv_cache_params", None) + cached_token_lens = getattr(kv_cache_params, "num_cached_tokens_per_seq", None) + if cached_token_lens is not None: + return torch.as_tensor(cached_token_lens[:num_contexts], dtype=torch.int64, device="cpu") + + return ( + ( + metadata.kv_lens_cuda_runtime[:num_contexts] + - metadata.prompt_lens_cuda_runtime[:num_contexts] + ) + .detach() + .cpu() + ) + + +def _validate_fp4_mla_context_start_alignment(metadata: Any, num_contexts: int) -> None: + context_start_positions = _get_fp4_mla_context_start_positions(metadata, num_contexts) + bad_start = (context_start_positions < 0) | ((context_start_positions % HP_BLOCK_SIZE) != 0) + if bool(torch.any(bad_start).item()): + starts = context_start_positions.detach().cpu().tolist() + raise ValueError( + "FP4 MLA shared-tile context update requires every context " + f"start position to be {HP_BLOCK_SIZE}-token aligned, got " + f"start positions {starts}." + ) + + +def get_fp4_mla_v_scale_pool_shape( + num_layers: int, + num_pages: int, + v_head_dim: int, + page_size: int, +) -> tuple[int, int, int, int, int, int]: + """Return the logical swizzled V-scale view shape. + + The leading dimensions are ``[layer, physical_page]``. The remaining + dimensions are the preshuffled ``[N // 128, K // 16 // 4, 32, 16]`` shape + consumed by Triton block-scaled matmul for the V/PV RHS operand. + """ + token_scale_cols = _ceil_div(page_size, FP4_BLOCK_SIZE) + return ( + num_layers, + num_pages, + _ceil_div(v_head_dim, FP4_MLA_SCALE_ROW_GROUP), + _ceil_div(token_scale_cols, FP4_MLA_SCALE_COL_GROUP), + 32, + 16, + ) + + +def get_fp4_mla_v_scale_pool_view( + metadata: Any, + *, + v_head_dim: int, +) -> torch.Tensor: + """View the auxiliary MLA V-scale pool in Triton's block-scaled layout.""" + pool = getattr(metadata, "fp4_mla_v_scale_pool", None) + if pool is None: + raise RuntimeError("FP4 MLA V scale pool is not allocated.") + + elems_per_page = get_fp4_mla_v_scale_pool_size(v_head_dim, metadata.page_size) + if pool.shape[-1] < elems_per_page: + raise RuntimeError( + f"FP4 MLA V scale pool page stride is too small: got " + f"{pool.shape[-1]}, need {elems_per_page}." + ) + + token_scale_cols = _ceil_div(metadata.page_size, FP4_BLOCK_SIZE) + col_groups = _ceil_div(token_scale_cols, FP4_MLA_SCALE_COL_GROUP) + shape = get_fp4_mla_v_scale_pool_shape( + pool.shape[0], pool.shape[1], v_head_dim, metadata.page_size + ) + strides = ( + pool.stride(0), + pool.stride(1), + col_groups * 32 * 16, + 32 * 16, + 16, + 1, + ) + return torch.as_strided(pool, size=shape, stride=strides) + + +# Python launch helpers + + +def _get_fp4_mla_global_scale(metadata: Any, device: torch.device) -> torch.Tensor: + global_scale = getattr(metadata, "_fp4_mla_global_scale", None) + if global_scale is None: + global_scale = torch.ones((1,), dtype=torch.float32, device=device) + return global_scale + + +def _get_fp4_mla_kv_cache_tensors( + metadata: Any, layer_idx: int +) -> tuple[torch.Tensor, torch.Tensor]: + kv_cache = metadata.kv_cache_manager.get_buffers(layer_idx).view(torch.uint8) + sf_cache = metadata.kv_cache_manager.get_block_scale_buffers(layer_idx) + if sf_cache is None: + raise RuntimeError("NVFP4 KV cache scale pool is not available.") + return kv_cache, sf_cache + + +def _scatter_fp4_mla_kv_cache_2d_context( + metadata: Any, + latent_cache: torch.Tensor, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + v_sf: torch.Tensor, + global_scale: torch.Tensor, + *, + token_offset: int, + local_layer: int, + v_head_dim: int, + head_dim: int, + num_tokens: int, + num_dim_blocks: int, + sf_per_token: int, + sf_per_page: int, +) -> None: + num_contexts = metadata.num_contexts + if num_contexts > 0: + prompt_lens_cpu = metadata.prompt_lens_cpu_runtime[:num_contexts] + ctx_token_count = int(prompt_lens_cpu.sum().item()) + if num_tokens != ctx_token_count: + raise RuntimeError( + f"FP4 MLA 2D context scatter needs {ctx_token_count} context tokens, got " + f"{num_tokens}." + ) + _validate_fp4_mla_context_start_alignment(metadata, num_contexts) + + _fp4_mla_v_scale_store_context_tokens_kernel[ + ( + num_tokens, + num_dim_blocks, + ) + ]( + kv_cache, + sf_cache, + v_sf, + latent_cache, + global_scale, + metadata.batch_indices, + metadata.positions, + metadata.paged_kv_indices, + metadata.paged_kv_indptr, + metadata.paged_kv_indices.shape[0], + metadata.paged_kv_indptr.shape[0], + metadata.batch_indices.shape[0], + v_sf.shape[1], + v_sf.shape[0], + token_offset, + num_tokens, + local_layer, + metadata.page_size, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + sf_cache.stride(0), + latent_cache.stride(0), + latent_cache.stride(1), + v_sf.stride(0), + v_sf.stride(1), + HEAD_D=head_dim, + V_HEAD_D=v_head_dim, + HP_BLOCK=HP_BLOCK_SIZE, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_TOKEN=sf_per_token, + SF_PER_PAGE=sf_per_page, + ) + _debug_sync("scatter_fp4_mla_kv_cache_2d_context") + + +def _scatter_fp4_mla_kv_cache_2d_generation( + metadata: Any, + latent_cache: torch.Tensor, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + v_sf: torch.Tensor, + global_scale: torch.Tensor, + *, + local_layer: int, + v_head_dim: int, + head_dim: int, + num_tokens: int, + num_dim_blocks: int, + sf_per_token: int, + sf_per_page: int, +) -> None: + num_contexts = metadata.num_contexts + num_seqs = metadata.num_seqs + num_gen = num_seqs - num_contexts + if num_gen <= 0: + return + gen_token_lens = _host_int_list_during_forward( + getattr(metadata, "prompt_lens_cpu_runtime", None), num_contexts, num_seqs + ) + if gen_token_lens is not None: + expected_num_tokens = sum(gen_token_lens) + if num_tokens != expected_num_tokens: + raise RuntimeError( + "FP4 MLA 2D generation scatter token count mismatch: " + f"expected {expected_num_tokens} generation tokens from " + f"per-sequence lengths {gen_token_lens}, got {num_tokens}." + ) + if min(gen_token_lens) != max(gen_token_lens): + raise NotImplementedError( + "FP4 MLA no-dequant generation scatter currently supports " + f"uniform linear MTP lengths only, got {gen_token_lens}." + ) + elif num_tokens < num_gen: + raise RuntimeError( + f"FP4 MLA 2D generation scatter needs at least {num_gen} generation " + f"tokens, got {num_tokens}." + ) + elif num_tokens % num_gen != 0: + raise NotImplementedError( + "FP4 MLA no-dequant generation scatter requires a uniform " + f"generation length, got {num_tokens} tokens for {num_gen} sequences." + ) + + pool = getattr(metadata, "high_precision_kv_pool", None) + if pool is None: + raise RuntimeError("FP4 MLA 2D generation scatter requires the HP KV pool.") + hp_head_dim = pool.shape[-1] // HP_BLOCK_SIZE + if hp_head_dim < head_dim: + raise RuntimeError( + f"FP4 MLA 2D generation scatter needs at least {head_dim} HP channels, got " + f"{hp_head_dim}." + ) + + max_gen_len = max(gen_token_lens) if gen_token_lens is not None else num_tokens // num_gen + max_gen_tiles = _ceil_div(max_gen_len + HP_BLOCK_SIZE - 1, HP_BLOCK_SIZE) + page_ids = metadata.paged_kv_indices[metadata.num_context_blocks :] + _fp4_mla_v_scale_store_generation_tiles_kernel[ + ( + num_gen, + max(max_gen_tiles, 1), + num_dim_blocks, + ) + ]( + kv_cache, + sf_cache, + v_sf, + pool, + latent_cache, + global_scale, + metadata.seq_slots[num_contexts:num_seqs], + metadata.kv_lens_cuda_runtime[num_contexts:num_seqs], + metadata.prompt_lens_cuda_runtime[num_contexts:num_seqs], + page_ids, + metadata.paged_kv_indptr_decode, + page_ids.shape[0], + metadata.paged_kv_indptr_decode.shape[0], + pool.shape[0], + v_sf.shape[1], + v_sf.shape[0], + local_layer, + metadata.page_size, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + sf_cache.stride(0), + pool.stride(0), + pool.stride(1), + latent_cache.stride(0), + latent_cache.stride(1), + v_sf.stride(0), + v_sf.stride(1), + HEAD_D=hp_head_dim, + V_HEAD_D=v_head_dim, + HP_BLOCK=HP_BLOCK_SIZE, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_TOKEN=sf_per_token, + SF_PER_PAGE=sf_per_page, + ) + _debug_sync("scatter_fp4_mla_kv_cache_2d_generation") + + +def _check_fp4_mla_page_generation_single_page( + metadata: Any, + num_contexts: int, + num_seqs: int, + page_size: int, +) -> None: + """Raise if a step writes more than one generation token per sequence. + + The per-page re-quant v1 re-quantizes a single active page per step from the + staging buffer (which holds exactly one page). A linear-MTP draft of length + > 1 can straddle a page boundary, which would require the previous + (completing) page to be re-quantized too -- not yet supported. 1-token + decode never crosses, so only MTP drafts are rejected. Detectable only in + eager mode (host metadata available); under CUDA graph the captured shape + is assumed to be 1-token decode, so do not enable the per-page path together + with MTP + CUDA graph until the multi-page store lands. + """ + gen_token_lens = _host_int_list_during_forward( + getattr(metadata, "prompt_lens_cpu_runtime", None), num_contexts, num_seqs + ) + if gen_token_lens is None: + return + max_gen_len = max(gen_token_lens) if gen_token_lens else 1 + if max_gen_len > 1: + raise NotImplementedError( + "FP4 MLA per-page dynamic scale (v1) supports 1-token decode only; " + f"got a generation length of {max_gen_len} (linear MTP). Multi-page " + "re-quant for MTP drafts is a follow-up." + ) + + +def _store_fp4_mla_page_dynamic_generation( + metadata: Any, + latent_cache: torch.Tensor, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + v_sf: torch.Tensor, + page_scale_pool: torch.Tensor, + stage_pool: torch.Tensor, + *, + local_layer: int, + v_head_dim: int, + head_dim: int, + sf_per_token: int, + sf_per_page: int, +) -> None: + """Re-quantize the active decode page with an exact per-page global scale. + + Two passes (see ``fp4_mla_triton``): Pass A computes the shared K/V + page amax -> ``page_gscale`` into ``page_scale_pool``; Pass B re-quantizes + every FP4 tile of the active page from ``stage_pool`` (old tokens) + + ``latent_cache`` (new tokens), baking ``page_gscale`` into the K and V block + scales. Replaces the static 16-token tile scatter for the ``triton`` + path when the per-page scale is enabled. + """ + from .fp4_mla_triton import _fp4_mla_page_requant_gen_kernel, _fp4_mla_page_scale_gen_kernel + + num_contexts = metadata.num_contexts + num_seqs = metadata.num_seqs + num_gen = num_seqs - num_contexts + if num_gen <= 0: + return + + page_size = metadata.page_size + _check_fp4_mla_page_generation_single_page(metadata, num_contexts, num_seqs, page_size) + + pool_head_dim = stage_pool.shape[-1] // page_size + if pool_head_dim < head_dim: + raise RuntimeError( + f"FP4 MLA staging pool head dim {pool_head_dim} < latent head_dim {head_dim}." + ) + + page_ids = metadata.paged_kv_indices[metadata.num_context_blocks :] + seq_slots = metadata.seq_slots[num_contexts:num_seqs] + kv_lens = metadata.kv_lens_cuda_runtime[num_contexts:num_seqs] + gen_lens = metadata.prompt_lens_cuda_runtime[num_contexts:num_seqs] + indptr = metadata.paged_kv_indptr_decode + num_dim_blocks = triton.cdiv(head_dim, FP4_BLOCK_SIZE) + tiles_per_page = page_size // FP4_BLOCK_SIZE + block_d = triton.next_power_of_2(head_dim) + + # Pass A: per-page (K/V shared) amax -> page_gscale. + _fp4_mla_page_scale_gen_kernel[(num_gen,)]( + page_scale_pool, + stage_pool, + latent_cache, + seq_slots, + kv_lens, + gen_lens, + page_ids, + indptr, + page_ids.shape[0], + indptr.shape[0], + stage_pool.shape[0], + kv_cache.shape[0], + page_scale_pool.shape[0], + local_layer, + page_size, + page_scale_pool.stride(0), + stage_pool.stride(0), + stage_pool.stride(1), + latent_cache.stride(0), + latent_cache.stride(1), + HEAD_D=head_dim, + POOL_HEAD_D=pool_head_dim, + FP4_BLOCK=FP4_BLOCK_SIZE, + PAGE_SLOTS=page_size, + BLOCK_D=block_d, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + ) + _debug_sync("fp4_mla_page_scale_gen") + + # Pass B: re-quantize every tile of the active page with page_gscale. + _fp4_mla_page_requant_gen_kernel[(num_gen, tiles_per_page, num_dim_blocks)]( + kv_cache, + sf_cache, + v_sf, + stage_pool, + latent_cache, + page_scale_pool, + seq_slots, + kv_lens, + gen_lens, + page_ids, + indptr, + page_ids.shape[0], + indptr.shape[0], + stage_pool.shape[0], + kv_cache.shape[0], + page_scale_pool.shape[0], + local_layer, + page_size, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + sf_cache.stride(0), + stage_pool.stride(0), + stage_pool.stride(1), + latent_cache.stride(0), + latent_cache.stride(1), + v_sf.stride(0), + v_sf.stride(1), + page_scale_pool.stride(0), + HEAD_D=head_dim, + POOL_HEAD_D=pool_head_dim, + V_HEAD_D=v_head_dim, + PAGE_SLOTS=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_TOKEN=sf_per_token, + SF_PER_PAGE=sf_per_page, + ) + _debug_sync("fp4_mla_page_requant_gen") + + +def _scatter_fp4_mla_kv_cache_1d( + metadata: Any, + latent_cache: torch.Tensor, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + global_scale: torch.Tensor, + *, + layer_idx: int, + token_offset: int, + num_tokens: int, + head_dim: int, + sf_per_token: int, + use_swizzled_sf: bool, +) -> None: + q_fp4, q_sf = torch.ops.trtllm.fp4_quantize( + latent_cache, global_scale, FP4_BLOCK_SIZE, False, False + ) + q_sf = q_sf.view(num_tokens, head_dim // FP4_BLOCK_SIZE) + + packed_dim = head_dim // 2 + block_packed_dim = triton.next_power_of_2(packed_dim) + block_sf = triton.next_power_of_2(sf_per_token) + + _fp4_mla_debug( + "scatter launch: " + f"num_tokens={num_tokens} token_offset={token_offset} " + f"page_size={metadata.page_size} layer_idx={layer_idx} " + f"head_dim={head_dim} packed_dim={packed_dim} " + f"sf_per_token={sf_per_token} use_swizzled_sf={use_swizzled_sf}" + ) + _fp4_mla_debug(f"scatter latent_cache: {_tensor_layout(latent_cache)}") + _fp4_mla_debug(f"scatter kv_cache: {_tensor_layout(kv_cache)}") + _fp4_mla_debug(f"scatter sf_cache: {_tensor_layout(sf_cache)}") + _debug_tensor_range( + "scatter batch_indices", + metadata.batch_indices[token_offset : token_offset + num_tokens], + ) + _debug_tensor_range( + "scatter positions", + metadata.positions[token_offset : token_offset + num_tokens], + ) + _debug_tensor_range("scatter paged_kv_indices", metadata.paged_kv_indices) + _debug_tensor_range("scatter paged_kv_indptr", metadata.paged_kv_indptr) + + _fp4_mla_scatter_kernel[(num_tokens,)]( + kv_cache, + sf_cache, + q_fp4, + q_sf, + metadata.batch_indices, + metadata.positions, + metadata.paged_kv_indices, + metadata.paged_kv_indptr, + metadata.paged_kv_indices.shape[0], + metadata.paged_kv_indptr.shape[0], + kv_cache.shape[0], + token_offset, + metadata.page_size, + kv_cache.stride(0), + kv_cache.stride(1), + kv_cache.stride(2), + kv_cache.stride(3), + kv_cache.stride(4), + sf_cache.stride(0), + sf_cache.stride(1), + sf_cache.stride(2), + sf_cache.stride(3), + sf_cache.stride(4), + q_fp4.stride(0), + q_fp4.stride(1), + q_sf.stride(0), + q_sf.stride(1), + PACKED_D=packed_dim, + SF_PER_TOKEN=sf_per_token, + BLOCK_PACKED_D=block_packed_dim, + BLOCK_SF=block_sf, + USE_SWIZZLED_SF=use_swizzled_sf, + ) + _debug_sync("scatter_fp4_mla_kv_cache") + + +# Public cache update and decode entry points + + +def scatter_fp4_mla_kv_cache( + metadata: Any, + latent_cache: Optional[torch.Tensor], + layer_idx: int, + *, + token_offset: int, + phase: Optional[_HPUpdatePhase] = None, + local_layer: Optional[int] = None, + v_head_dim: Optional[int] = None, +) -> None: + """Quantize MLA latent tokens and scatter them into the paged FP4 cache. + + Contract: this helper scatters exactly ``latent_cache.shape[0]`` tokens, + reading index metadata at ``batch_indices[token_offset : token_offset + N]`` + and ``positions[token_offset : token_offset + N]``. Callers must pass a + latent_cache pre-sliced to the current phase (context or generation) so + that ``shape[0]`` matches the number of index entries they intend to + consume. ``MLA.forward_impl`` (tensorrt_llm/_torch/modules/attention.py) + slices ``latent_cache[:num_ctx_tokens]`` for context and + ``latent_cache[num_ctx_tokens:]`` for generation before dispatching. + + When the no-dequant FP4 MLA attention path is enabled, callers should pass + ``phase``, ``local_layer``, and ``v_head_dim``. Context scatter then writes + the final FP4 tile representation directly: dimensions below + ``v_head_dim`` use one shared 16-token by 16-dim scale written into both + K's token-major scale layout and V's dim-major scale layout. Tail K-only + dimensions use K's per-token 1D scales. Generation scatter rewrites each + touched 16-token tile by reading old tokens from the HP pool and new tokens + from ``latent_cache``; callers then update the HP pool after scatter. + """ + if latent_cache is None or latent_cache.numel() == 0: + return + + latent_cache = latent_cache.reshape(latent_cache.shape[0], -1).contiguous() + num_tokens = latent_cache.shape[0] + head_dim = latent_cache.shape[-1] + if head_dim % FP4_BLOCK_SIZE != 0: + raise ValueError( + f"FP4 MLA KV head_dim must be divisible by {FP4_BLOCK_SIZE}, got {head_dim}." + ) + indices_len = metadata.batch_indices.shape[0] + positions_len = metadata.positions.shape[0] + if token_offset + num_tokens > indices_len or token_offset + num_tokens > positions_len: + raise RuntimeError( + f"FP4 MLA scatter would read batch_indices[{token_offset}:" + f"{token_offset + num_tokens}] / positions[{token_offset}:" + f"{token_offset + num_tokens}], but only {indices_len} / " + f"{positions_len} entries are available. This indicates " + "latent_cache was not pre-sliced to the current phase's token " + "range (see MLA.forward_impl)." + ) + + use_swizzled_sf = _use_fp4_mla_swizzled_sf() + if use_swizzled_sf: + _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) + + global_scale = _get_fp4_mla_global_scale(metadata, latent_cache.device) + kv_cache, sf_cache = _get_fp4_mla_kv_cache_tensors(metadata, layer_idx) + sf_per_token = head_dim // FP4_BLOCK_SIZE + + use_2d_scatter = ( + use_swizzled_sf + and phase in ("context", "generation") + and getattr(metadata, "fp4_mla_v_scale_pool", None) is not None + ) + if use_2d_scatter: + assert phase is not None + if local_layer is None or v_head_dim is None: + raise ValueError("Real FP4 MLA scatter requires local_layer and v_head_dim.") + if metadata.page_size % HP_BLOCK_SIZE != 0: + raise ValueError( + f"FP4 MLA scatter requires page_size divisible by " + f"{HP_BLOCK_SIZE}, got {metadata.page_size}." + ) + if v_head_dim > head_dim: + raise ValueError(f"FP4 MLA v_head_dim={v_head_dim} cannot exceed head_dim={head_dim}.") + if v_head_dim % FP4_BLOCK_SIZE != 0: + raise ValueError( + f"FP4 MLA v_head_dim must be divisible by {FP4_BLOCK_SIZE}, got {v_head_dim}." + ) + + sf_cache = sf_cache.view(torch.float8_e4m3fn) + v_sf = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=v_head_dim) + num_dim_blocks = triton.cdiv(head_dim, FP4_BLOCK_SIZE) + sf_per_page = metadata.page_size // HP_BLOCK_SIZE + _fp4_mla_debug( + "scatter 2d launch: " + f"phase={phase} num_tokens={num_tokens} " + f"token_offset={token_offset} layer_idx={layer_idx} " + f"local_layer={local_layer} head_dim={head_dim} " + f"v_head_dim={v_head_dim} num_dim_blocks={num_dim_blocks}" + ) + _fp4_mla_debug(f"scatter 2d kv_cache: {_tensor_layout(kv_cache)}") + _fp4_mla_debug(f"scatter 2d sf_cache: {_tensor_layout(sf_cache)}") + _fp4_mla_debug(f"scatter 2d v_sf: {_tensor_layout(v_sf)}") + + if phase == "context": + _scatter_fp4_mla_kv_cache_2d_context( + metadata, + latent_cache, + kv_cache, + sf_cache, + v_sf, + global_scale, + token_offset=token_offset, + local_layer=local_layer, + v_head_dim=v_head_dim, + head_dim=head_dim, + num_tokens=num_tokens, + num_dim_blocks=num_dim_blocks, + sf_per_token=sf_per_token, + sf_per_page=sf_per_page, + ) + else: + page_scale_pool = getattr(metadata, "fp4_mla_page_scale_pool", None) + stage_pool = getattr(metadata, "fp4_mla_page_stage_pool", None) + if ( + fp4_mla_per_page_scale_enabled() + and page_scale_pool is not None + and stage_pool is not None + ): + # Per-page dynamic scale: re-quantize the active page from the + # BF16 staging buffer with its exact per-page global scale. + # NOTE: the staging buffer must hold this step's *pre-update* + # tokens, so update_page_stage_for_fp4_mla must run AFTER this. + _store_fp4_mla_page_dynamic_generation( + metadata, + latent_cache, + kv_cache, + sf_cache, + v_sf, + page_scale_pool, + stage_pool, + local_layer=local_layer, + v_head_dim=v_head_dim, + head_dim=head_dim, + sf_per_token=sf_per_token, + sf_per_page=sf_per_page, + ) + else: + _scatter_fp4_mla_kv_cache_2d_generation( + metadata, + latent_cache, + kv_cache, + sf_cache, + v_sf, + global_scale, + local_layer=local_layer, + v_head_dim=v_head_dim, + head_dim=head_dim, + num_tokens=num_tokens, + num_dim_blocks=num_dim_blocks, + sf_per_token=sf_per_token, + sf_per_page=sf_per_page, + ) + if phase == "context": + v_pack_page_ids = metadata.paged_kv_indices + else: + v_pack_page_ids = metadata.paged_kv_indices[metadata.num_context_blocks :] + _maybe_update_cutile_v_packed_cache( + metadata, + layer_idx, + kv_cache, + v_pack_page_ids, + v_head_dim=v_head_dim, + page_size=metadata.page_size, + ) + return + + _scatter_fp4_mla_kv_cache_1d( + metadata, + latent_cache, + kv_cache, + sf_cache, + global_scale, + layer_idx=layer_idx, + token_offset=token_offset, + num_tokens=num_tokens, + head_dim=head_dim, + sf_per_token=sf_per_token, + use_swizzled_sf=use_swizzled_sf, + ) + + +def _ensure_decode_workspace( + metadata: Any, + head_dim: int, + dtype: torch.dtype, +) -> torch.Tensor: + num_blocks = _get_decode_workspace_num_blocks(metadata) + workspace = getattr(metadata, "_fp4_mla_decode_cache_buf", None) + needs_alloc = ( + workspace is None + or workspace.shape[0] < max(num_blocks, 1) + or workspace.shape[1] != metadata.page_size + or workspace.shape[2] != head_dim + or workspace.dtype != dtype + ) + if needs_alloc: + if torch.cuda.is_current_stream_capturing(): + raise ValueError( + "Cannot allocate FlashInfer FP4 MLA decode workspace while " + "capturing a CUDA graph. Run a warmup prepare/forward first." + ) + workspace = torch.empty( + (max(num_blocks, 1), metadata.page_size, head_dim), + dtype=dtype, + device=metadata.paged_kv_indices.device, + ) + metadata._fp4_mla_decode_cache_buf = workspace + return workspace[:num_blocks] + + +def _get_decode_workspace_num_blocks(metadata: Any) -> int: + if metadata.is_cuda_graph: + max_blocks_per_seq = ( + metadata.kv_cache_manager.max_seq_len + metadata.page_size - 1 + ) // metadata.page_size + max_graph_blocks = metadata.max_num_requests * max_blocks_per_seq + return min( + metadata.kv_cache_manager.blocks_in_primary_pool, + max_graph_blocks, + ) + return metadata.num_generation_blocks + + +def _get_decode_src_page_ids(metadata: Any, num_blocks: int) -> torch.Tensor: + page_ids = ( + metadata._paged_kv_indices + if metadata.is_cuda_graph and hasattr(metadata, "_paged_kv_indices") + else metadata.paged_kv_indices + ) + src_page_ids = page_ids[metadata.num_context_blocks : metadata.num_context_blocks + num_blocks] + if src_page_ids.numel() != num_blocks: + raise RuntimeError( + f"FP4 MLA dequant needs {num_blocks} decode page ids from " + f"paged_kv_indices[{metadata.num_context_blocks}:" + f"{metadata.num_context_blocks + num_blocks}], got " + f"{src_page_ids.numel()}." + ) + return src_page_ids + + +def _assert_fp4_mla_decode_paging_consistent( + metadata: Any, + kv_lens: torch.Tensor, + num_gen_blocks: int, + query_len_per_seq: int, +) -> None: + """Flag read-side decode paging that diverges from the corrected kv_lens. + + With the overlap scheduler the generation metadata is first built from the + all-draft-accepted over-estimate; ``kv_lens_cuda_runtime`` is then corrected + in ``_preprocess_inputs`` and ``positions`` / ``batch_indices`` rebuilt. The + page-table side -- ``num_blocks`` / ``num_generation_blocks`` / + ``paged_kv_indptr_decode`` and the ``positions`` of the new tokens -- is what + the decode kernels index with. This check raises on the first decode where + any of those is inconsistent with the corrected kv_lens, so we can tell + whether stale paging (rather than masking) corrupts the read. Host-syncs; + gated by ``TRTLLM_FP4_MLA_DEBUG_ASSERT`` and skipped under CUDA-graph capture. + """ + if not _env_enabled(FP4_MLA_DEBUG_ASSERT_ENV): + return + if torch.cuda.is_current_stream_capturing(): + return + # Warmup builds synthetic metadata with dummy kv_lens but no real page-table + # allocation (num_generation_blocks==0), so the consistency check does not + # apply there. + if getattr(metadata, "is_warmup", False): + return + num_contexts = metadata.num_contexts + num_seqs = metadata.num_seqs + num_gen = num_seqs - num_contexts + if num_gen <= 0: + return + # Skip the benign context->generation reclassification state: the MTP draft + # loop resets num_contexts to 0 without rebuilding the decode page table + # (for FlashInfer the reorder is gated on enable_flash_mla), leaving + # num_context_blocks > 0 / stale num_generation_blocks. That forward's + # attention is degraded but rejected drafts fall back to the golden token, + # so it does NOT affect output accuracy (it fires with overlap off too). + # Only check pure steady-state generation, where a paging bug would + # genuinely corrupt the committed output. + if int(getattr(metadata, "num_context_blocks", 0)) != 0: + return + # Only the MTP *target* verification forward (query_len_per_seq > 1) + # determines the accepted/committed tokens, so it is the only forward whose + # paging staleness can change output accuracy. The draft-model forwards + # (query_len_per_seq == 1) read a separate KV layer and their bad output is + # rejected (and they mismatch in both overlap modes), so skip them here to + # isolate the accuracy-relevant path. + if query_len_per_seq <= 1: + return + + page_size = metadata.page_size + kv_lens_host = kv_lens.detach().to("cpu", torch.int64).tolist() + expected_blocks = [_ceil_div(kv, page_size) for kv in kv_lens_host] + expected_gen_blocks = sum(expected_blocks) + problems: list[str] = [] + + num_blocks = getattr(metadata, "num_blocks", None) + if num_blocks is not None: + actual_blocks = [int(b) for b in num_blocks[num_contexts:num_seqs]] + if actual_blocks != expected_blocks: + problems.append(f"per-seq num_blocks {actual_blocks} != expected {expected_blocks}") + if num_gen_blocks != expected_gen_blocks: + problems.append(f"num_generation_blocks={num_gen_blocks} != expected {expected_gen_blocks}") + + indptr = getattr(metadata, "paged_kv_indptr_decode", None) + if indptr is not None: + indptr_host = indptr[: num_gen + 1].detach().to("cpu", torch.int64).tolist() + expected_indptr = [0] + for b in expected_blocks: + expected_indptr.append(expected_indptr[-1] + b) + if indptr_host != expected_indptr: + problems.append(f"paged_kv_indptr_decode {indptr_host} != expected {expected_indptr}") + + positions = getattr(metadata, "positions", None) + prompt_lens = getattr(metadata, "prompt_lens_cuda_runtime", None) + if positions is not None and prompt_lens is not None: + num_ctx_tokens = int(getattr(metadata, "num_ctx_tokens", 0)) + gen_pos = positions[num_ctx_tokens:].detach().to("cpu", torch.int64).tolist() + pls = prompt_lens[num_contexts:num_seqs].detach().to("cpu", torch.int64).tolist() + off = 0 + for s, (kv_len, prompt_len) in enumerate(zip(kv_lens_host, pls)): + expected = list(range(kv_len - prompt_len, kv_len)) + got = gen_pos[off : off + prompt_len] + if got != expected: + problems.append( + f"seq{s} gen positions {got} != expected {expected} " + f"(kv_len={kv_len}, prompt_len={prompt_len})" + ) + off += prompt_len + + if problems: + ctx = ( + f"num_contexts={num_contexts}, num_seqs={num_seqs}, " + f"num_generations={getattr(metadata, 'num_generations', '?')}, " + f"query_len_per_seq={query_len_per_seq}, " + f"num_context_blocks={getattr(metadata, 'num_context_blocks', '?')}, " + f"num_blocks={getattr(metadata, 'num_blocks', '?')}, " + f"use_spec_decoding={getattr(metadata, 'use_spec_decoding', '?')}, " + f"is_spec_dec_mode={getattr(metadata, 'is_spec_dec_mode', '?')}, " + f"kv_lens={kv_lens_host}" + ) + msg = ( + "FP4 MLA decode paging inconsistent with corrected kv_lens " + f"({ctx}):\n " + "\n ".join(problems) + ) + if _env_enabled("TRTLLM_FP4_MLA_DEBUG_ASSERT_WARN"): + logger.warning(msg) + return + raise AssertionError(msg) + + +def _validate_fp4_mla_cache_shape(page_size: int, head_dim: int) -> None: + if page_size != FP4_MLA_TOKENS_PER_BLOCK: + raise ValueError( + f"FP4 MLA KV cache requires tokens_per_block={FP4_MLA_TOKENS_PER_BLOCK} " + f"for swizzled block scales, got {page_size}." + ) + + sf_per_token = head_dim // FP4_BLOCK_SIZE + if head_dim % FP4_BLOCK_SIZE != 0 or sf_per_token % 4 != 0: + raise ValueError( + f"FP4 MLA KV head_dim must produce a scale column count divisible by 4; " + f"got head_dim={head_dim}, scale_columns={sf_per_token}." + ) + + +def _validate_fp4_mla_attention_q_shape(head_dim: int, q_residual_dim: int) -> None: + if q_residual_dim % FP4_BLOCK_SIZE != 0: + raise ValueError( + f"FP4 MLA Q residual_dim must be divisible by {FP4_BLOCK_SIZE}, got {q_residual_dim}." + ) + if q_residual_dim <= 0 or q_residual_dim > head_dim: + raise ValueError( + f"FP4 MLA Q residual_dim must be in (0, head_dim], got " + f"residual_dim={q_residual_dim}, head_dim={head_dim}." + ) + + q_head_dim = head_dim + q_residual_dim + q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE + if q_head_dim % FP4_BLOCK_SIZE != 0 or q_sf_per_token % FP4_MLA_SCALE_COL_GROUP != 0: + raise ValueError( + f"FP4 MLA residual Q must produce a scale column count divisible " + f"by {FP4_MLA_SCALE_COL_GROUP}; got q_head_dim={q_head_dim}, " + f"scale_columns={q_sf_per_token}." + ) + + +def get_fp4_mla_decode_cache( + metadata: Any, + layer_idx: int, + local_layer: int, + *, + head_dim: int, + dtype: torch.dtype, +) -> torch.Tensor: + """Build a compact dequantized MLA cache for FlashInfer decode.""" + # Must match scatter_fp4_mla_kv_cache: the env var picks the SF layout, + # not the page size. When the dequant fallback is the read path + # (env disabled), scatter wrote linear scales and we must read linear. + use_swizzled_sf = _use_fp4_mla_swizzled_sf() + if use_swizzled_sf: + _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) + combined = _ensure_decode_workspace(metadata, head_dim, dtype) + num_blocks = combined.shape[0] + if num_blocks == 0: + return combined + + kv_cache, sf_cache = _get_fp4_mla_kv_cache_tensors(metadata, layer_idx) + sf_cache = sf_cache.view(torch.float8_e4m3fn) + global_scale = _get_fp4_mla_global_scale(metadata, combined.device) + src_page_ids = _get_decode_src_page_ids(metadata, num_blocks) + block_d = triton.next_power_of_2(head_dim) + + _fp4_mla_dequant_kernel[(num_blocks, metadata.page_size)]( + combined, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + src_page_ids.shape[0], + kv_cache.shape[0], + kv_cache.stride(0), + kv_cache.stride(1), + kv_cache.stride(2), + kv_cache.stride(3), + kv_cache.stride(4), + sf_cache.stride(0), + sf_cache.stride(1), + sf_cache.stride(2), + sf_cache.stride(3), + sf_cache.stride(4), + combined.stride(0), + combined.stride(1), + combined.stride(2), + D=head_dim, + FP4_BLOCK=FP4_BLOCK_SIZE, + BLOCK_D=block_d, + USE_SWIZZLED_SF=use_swizzled_sf, + ) + + num_gen = metadata.num_seqs - metadata.num_contexts + if num_gen > 0 and metadata.high_precision_kv_pool is not None: + pool = metadata.high_precision_kv_pool + _fp4_mla_overlay_hp_tail_kernel[(num_gen, HP_BLOCK_SIZE)]( + combined, + pool, + metadata.seq_slots[metadata.num_contexts : metadata.num_seqs], + metadata.kv_lens_cuda_runtime[metadata.num_contexts : metadata.num_seqs], + metadata.paged_kv_indptr_decode, + pool.shape[0], + pool.shape[1], + combined.shape[0], + local_layer, + metadata.page_size, + combined.stride(0), + combined.stride(1), + combined.stride(2), + pool.stride(0), + pool.stride(1), + D=head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + + return combined + + +def _ensure_workspace_tensor( + metadata: Any, + attr_name: str, + shape: tuple[int, ...], + *, + dtype: torch.dtype, + device: torch.device, +) -> torch.Tensor: + tensor = getattr(metadata, attr_name, None) + needs_alloc = ( + tensor is None + or tensor.dtype != dtype + or tensor.device != device + or len(tensor.shape) != len(shape) + or any(tensor.shape[idx] < dim for idx, dim in enumerate(shape)) + ) + if needs_alloc: + if torch.cuda.is_current_stream_capturing(): + raise ValueError( + f"Cannot allocate {attr_name} while capturing a CUDA graph. " + "Run a warmup prepare/forward first." + ) + tensor = torch.empty(shape, dtype=dtype, device=device) + setattr(metadata, attr_name, tensor) + + slices = tuple(slice(0, dim) for dim in shape) + return tensor[slices] + + +def _cutile_persistent_v_pack_enabled() -> bool: + if _fp4_mla_attention_backend() != "cutile": + return False + return os.getenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1").lower() not in ( + "0", + "false", + "no", + "off", + ) + + +def _cutile_v_packed_attr(layer_idx: int) -> str: + return f"_fp4_mla_attention_v_packed_buf_l{layer_idx}" + + +def _cutile_v_packed_valid_attr(layer_idx: int) -> str: + return f"_fp4_mla_attention_v_packed_valid_l{layer_idx}" + + +def _cutile_v_packed_shape( + kv_cache: torch.Tensor, + v_head_dim: int, + page_size: int, +) -> tuple[int, int]: + block_v = 128 + return (kv_cache.shape[0] * _ceil_div(v_head_dim, block_v) * block_v, page_size // 2) + + +def _maybe_update_cutile_v_packed_cache( + metadata: Any, + layer_idx: int, + kv_cache: torch.Tensor, + page_ids: torch.Tensor, + *, + v_head_dim: int, + page_size: int, +) -> None: + if not _cutile_persistent_v_pack_enabled(): + return + if v_head_dim % 128 != 0 or page_size != FP4_MLA_TOKENS_PER_BLOCK: + return + if page_ids.numel() == 0: + return + + from .fp4_mla_cutile import fp4_mla_repack_v_cache + + attr_name = _cutile_v_packed_attr(layer_idx) + v_packed = _ensure_workspace_tensor( + metadata, + attr_name, + _cutile_v_packed_shape(kv_cache, v_head_dim, page_size), + dtype=torch.uint8, + device=kv_cache.device, + ) + fp4_mla_repack_v_cache( + v_packed, + kv_cache, + page_ids, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=128, + ) + setattr(metadata, _cutile_v_packed_valid_attr(layer_idx), True) + + +def _get_cutile_v_packed_cache( + metadata: Any, + layer_idx: int, + kv_cache: torch.Tensor, + *, + v_head_dim: int, + page_size: int, +) -> Optional[torch.Tensor]: + if not _cutile_persistent_v_pack_enabled(): + return None + if not bool(getattr(metadata, _cutile_v_packed_valid_attr(layer_idx), False)): + return None + v_packed = getattr(metadata, _cutile_v_packed_attr(layer_idx), None) + expected_shape = _cutile_v_packed_shape(kv_cache, v_head_dim, page_size) + if ( + v_packed is None + or v_packed.dtype != torch.uint8 + or v_packed.device != kv_cache.device + or len(v_packed.shape) != 2 + or v_packed.shape[0] < expected_shape[0] + or v_packed.shape[1] < expected_shape[1] + ): + return None + return v_packed[: expected_shape[0], : expected_shape[1]] + + +def _max_generation_pages(metadata: Any) -> int: + num_gen = metadata.num_seqs - metadata.num_contexts + if num_gen <= 0: + return 0 + num_blocks = getattr(metadata, "num_blocks", None) + if num_blocks is not None: + return max(num_blocks[metadata.num_contexts : metadata.num_seqs]) + return metadata.num_generation_blocks + + +def _host_int_list(value: Any, start: int, end: int) -> Optional[list[int]]: + if value is None: + return None + if isinstance(value, torch.Tensor): + if value.is_cuda: + return None + return [int(item) for item in value[start:end].tolist()] + try: + return [int(item) for item in value[start:end]] + except (TypeError, ValueError): + return None + + +def _infer_cutile_assume_full_pages(metadata: Any, max_pages: int, page_size: int) -> bool: + if getattr(metadata, "is_cuda_graph", False): + return False + + start = metadata.num_contexts + end = metadata.num_seqs + kv_cache_params = getattr(metadata, "kv_cache_params", None) + cached_token_lens = _host_int_list( + getattr(kv_cache_params, "num_cached_tokens_per_seq", None), + start, + end, + ) + seq_lens_kv = _host_int_list(getattr(metadata, "seq_lens_kv", None), start, end) + if cached_token_lens is not None and seq_lens_kv is not None: + if len(cached_token_lens) != len(seq_lens_kv): + return False + kv_lens = [ + cached_len + seq_len for cached_len, seq_len in zip(cached_token_lens, seq_lens_kv) + ] + elif kv_cache_params is None: + kv_lens = _host_int_list(getattr(metadata, "prompt_lens_cpu_runtime", None), start, end) + else: + return False + + return bool(kv_lens) and min(kv_lens) == max(kv_lens) == max_pages * page_size + + +def _get_linear_mtp_query_len_per_seq( + metadata: Any, + *, + num_queries: int, + num_gen_seqs: int, +) -> int: + """Return the uniform generation query length required by linear MTP.""" + if num_gen_seqs <= 0: + return 1 + + start = metadata.num_contexts + end = metadata.num_seqs + query_lens = _host_int_list_during_forward( + getattr(metadata, "prompt_lens_cpu_runtime", None), start, end + ) + if query_lens is None: + query_lens = _host_int_list_during_forward(getattr(metadata, "seq_lens", None), start, end) + + if query_lens is None: + if num_queries % num_gen_seqs != 0: + raise NotImplementedError( + "FP4 MLA linear MTP requires a uniform generation query length; " + f"got {num_queries} query tokens for {num_gen_seqs} sequences." + ) + return num_queries // num_gen_seqs + + if sum(query_lens) != num_queries and num_queries == num_gen_seqs: + return 1 + if sum(query_lens) != num_queries: + raise RuntimeError( + "FP4 MLA generation query metadata does not match q shape: " + f"query_lens={query_lens}, total={sum(query_lens)}, " + f"q_tokens={num_queries}." + ) + if not query_lens: + return 1 + if min(query_lens) <= 0: + raise RuntimeError(f"FP4 MLA generation query lengths must be positive, got {query_lens}.") + if min(query_lens) != max(query_lens): + raise NotImplementedError( + "FP4 MLA no-dequant attention currently supports linear MTP with " + f"a uniform generation length per sequence, got {query_lens}." + ) + return query_lens[0] + + +def _run_triton_attention_decode( + *, + metadata: Any, + q_fp4: torch.Tensor, + q_sf: torch.Tensor, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + v_sf: torch.Tensor, + global_scale: torch.Tensor, + src_page_ids: torch.Tensor, + kv_lens: torch.Tensor, + p_fp4: torch.Tensor, + p_sf: torch.Tensor, + max_scores: torch.Tensor, + denom: torch.Tensor, + output: torch.Tensor, + num_queries: int, + num_heads: int, + head_dim: int, + kv_lora_rank: int, + q_residual_dim: int, + query_len_per_seq: int, + max_pages: int, + sm_scale: float, + q_global_scale: Optional[torch.Tensor] = None, + page_scale_pool: Optional[torch.Tensor] = None, + local_layer: int = 0, + use_per_page_scale: bool = False, +) -> None: + """Dispatch the ``triton`` FP4 MLA decode pipeline. + + Mirrors the four-stage layout used by ``fp4_mla_cutile.py`` + (page-stats with packed P -> reduce-stats -> prob-scale -> PV) but + routes through the self-contained kernels in + ``fp4_mla_triton.py``. Threads through the constexpr assume flags, + TMA descriptors, occupancy/num-warps launch meta, and pipelined PV loop. + """ + from .fp4_mla_triton import _fp4_mla_attention_page_stats_kernel as _attn_page_stats_kernel + from .fp4_mla_triton import _fp4_mla_attention_prob_scale_kernel as _attn_prob_scale_kernel + from .fp4_mla_triton import _fp4_mla_attention_pv_kernel as _attn_pv_kernel + from .fp4_mla_triton import _fp4_mla_attention_pv_reduce_kernel as _attn_pv_reduce_kernel + from .fp4_mla_triton import _fp4_mla_attention_reduce_stats_kernel as _attn_reduce_stats_kernel + + block_h = 128 + block_t = metadata.page_size + # Adaptive BLOCK_V: small batches need a finer V split to fill enough waves + # on B200 (~148 SMs). PV grid = num_queries * num_head_blocks(1) * + # (kv_lora_rank / BLOCK_V). We want >= ~2*num_SMs programs so that >1 CTA + # lands per SM and hides the L1TEX scoreboard stalls. Empirically (sweep): + # bs<=32 -> BLOCK_V=32; bs>=64 -> BLOCK_V=128. + # (BLOCK_V=16 is rejected by the V TMA descriptor min-stride requirement.) + block_v = 32 if num_queries <= 32 else 128 + q_head_dim = head_dim + q_residual_dim + # BLOCK_K = 512 matches cutile's "nvt" backend default and aligns the K-window + # with the residual-Q boundary (Q_HEAD_D = 640 = 512 + 128 tail). The + # residual-Q TMA tail path requires Q_RESIDUAL_D == 64 and TAIL_BLOCK_K == 128. + block_k = 512 + full_block_end = (q_head_dim // block_k) * block_k + tail_k = q_head_dim - full_block_end + tail_block_k = 1 << (tail_k - 1).bit_length() if tail_k > 0 else block_k + q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE + k_sf_per_token = head_dim // FP4_BLOCK_SIZE + sf_per_page = metadata.page_size // FP4_BLOCK_SIZE + num_head_blocks = triton.cdiv(num_heads, block_h) + + assume_full_heads = num_heads % block_h == 0 + assume_full_v = kv_lora_rank % block_v == 0 + # Match the cutile path: only mark pages "full" when we can prove every + # generation sequence has the same number of cached tokens AND + # query_len_per_seq == 1 (so the kv_len adjustment is a no-op). + assume_full_pages = ( + _infer_cutile_assume_full_pages(metadata, max_pages, metadata.page_size) + and query_len_per_seq == 1 + ) + # Leave validity checks on. Matches cutile's default and is correctness- + # safe. The perfect-shape PV fast path (tl.ext.make_view + load_view_tko) + # remains gated off — when measured on the TileIR backend (ENABLE_TILE=1) + # it was net-slower on the bench, so the cost of enabling it isn't worth + # the win on the FP4 MLA shapes we care about. + assume_valid_pages = False + # cutile checks only `make_tensor_descriptor`; on the nvt backend the + # presence of TMA descriptors implies `tl.ext.make_view` is available too. + use_tma_data_load = hasattr(triton.language, "make_tensor_descriptor") + + # Install the device-side scratch allocator on every call. Triton stores + # the allocator in a ContextVar (triton.runtime._allocation), so a single + # process-wide install is not visible from worker threads / asyncio tasks + # that run with a different Context — the kernel launch would then hit the + # default NullAllocator and raise. Matches the cutile path. + if use_tma_data_load: + + def _tma_alloc(size: int, alignment: int, stream): + return torch.empty(size, device=q_fp4.device, dtype=torch.int8) + + triton.set_allocator(_tma_alloc) + + # cutile-equivalent launch meta. occupancy=2 lets two CTAs land per SM + # which improves wave-tail efficiency at the bs=32 hot point. + # NOTE: num_stages=2 (instead of the Triton 3.6 default of 3) sidesteps + # the TritonGPUAutomaticWarpSpecialization + NVWSInsertTmemAref pass that + # ICEs on the page_stats kernel under Triton 3.6.0 / sm_100. + launch_meta = {"occupancy": 2} + # PV loop pipelining. With TMA loads, num_stages>=2 lets the next page's + # loads overlap with the current MMA via mbarrier. The PV report shows + # long_scoreboard=4.5 cycles avg on V loads at PV_LOOP_STAGES=2; bumping the + # depth pays off when the grid is small enough that occupancy can absorb + # the extra in-flight tile state — i.e. medium batch / large max_pages. + # Larger pipelines hurt at small batch (more live state, fewer dim blocks). + if num_queries <= 16 or max_pages <= 4: + pv_loop_stages = 2 + else: + pv_loop_stages = 3 + + # Page-stats kernel: per (query, head_block, page) program, does QK, + # softmax stats, and packs probs into FP4 with the per-page local-max + # scaling trick. The page-max correction is applied later by + # prob_scale_kernel via p_sf in-place rescaling. + page_stats_shape = (num_queries, max_pages, num_heads) + page_max = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_page_max_buf", + page_stats_shape, + dtype=torch.float32, + device=q_fp4.device, + ) + page_sum = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_page_sum_buf", + page_stats_shape, + dtype=torch.float32, + device=q_fp4.device, + ) + + pack_prob_in_page_stats = True + # Per-page scale pointers. When the per-page path is off, pass the static + # global_scale tensor as harmless dummies (the kernel never dereferences + # them under USE_PER_PAGE_SCALE=False) and a zero layer stride. + page_stats_q_gscale = q_global_scale if use_per_page_scale else global_scale + page_stats_page_scale = ( + page_scale_pool if (use_per_page_scale and page_scale_pool is not None) else global_scale + ) + page_stats_pscale_s0 = ( + page_scale_pool.stride(0) if (use_per_page_scale and page_scale_pool is not None) else 0 + ) + _attn_page_stats_kernel[(num_queries, num_head_blocks, max_pages)]( + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4, + q_sf, + kv_cache, + sf_cache, + global_scale, + page_stats_q_gscale, + page_stats_page_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + kv_cache.shape[0], + q_fp4.stride(0), + q_fp4.stride(1), + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + sf_cache.stride(0), + page_max.stride(0), + page_max.stride(1), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + q_fp4.shape[0], + sm_scale, + local_layer, + page_stats_pscale_s0, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + BLOCK_T=block_t, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + PACK_PROBS=pack_prob_in_page_stats, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + USE_PER_PAGE_SCALE=use_per_page_scale, + **launch_meta, + ) + + _attn_reduce_stats_kernel[(num_queries, num_head_blocks)]( + max_scores, + denom, + page_max, + page_sum, + max_pages, + max_scores.stride(0), + page_max.stride(0), + page_max.stride(1), + NUM_HEADS=num_heads, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + **launch_meta, + ) + + _attn_prob_scale_kernel[(num_queries, num_head_blocks, max_pages)]( + p_sf, + max_scores, + denom, + page_max, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + max_scores.stride(0), + page_max.stride(0), + page_max.stride(1), + NUM_HEADS=num_heads, + PAGE_SIZE=metadata.page_size, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + + num_dim_blocks = triton.cdiv(kv_lora_rank, block_v) + + # PV page split: partition the page range across additional programs and + # reduce in a follow-up kernel. ncu showed PV at waves/SM=0.49 for bs=32 — + # PV is L1-bandwidth bound, so raising in-flight CTAs is the lever. + # BLOCK_V is bounded below by the 16-byte TMA descriptor min-stride. + # PV page split: ncu shows that with the current shape (bs=32, max_pages=256) + # the PV kernel is L1-cache-throughput bound (long_scoreboard=4.5 cycles + # avg, L1 global LD hit-rate <40%). Increasing the program count via page + # splitting reduced waves/SM idle time but did NOT improve wall-time at + # current shapes — the per-CTA L1 thrash is the limit. Gate the split off + # by default; re-enable only for very small grids where occupancy is the + # bottleneck rather than per-CTA L1 pressure. + page_split = 1 + base_grid = num_queries * num_head_blocks * num_dim_blocks + if max_pages >= 16 and base_grid < 148: + for p in (8, 4, 2): + if max_pages % p == 0 and max_pages // p >= 16 and base_grid * p <= 148 * 4: + page_split = p + break + if page_split > 1: + pages_per_split = max_pages // page_split + partial_out = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_pv_partial_buf", + (num_queries, page_split, num_heads, kv_lora_rank), + dtype=torch.float32, + device=q_fp4.device, + ) + _attn_pv_kernel[(num_queries, num_head_blocks, num_dim_blocks * page_split)]( + output, + p_fp4, + p_sf, + kv_cache, + v_sf, + global_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + kv_cache.shape[0], + output.stride(0), + output.stride(1), + output.stride(2), + output.shape[0] * output.shape[1], + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=kv_lora_rank, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + BLOCK_H=block_h, + BLOCK_V=block_v, + USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, + USE_TMA_V_LOAD=use_tma_data_load and kv_lora_rank % block_v == 0, + PV_LOOP_STAGES=pv_loop_stages, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_FULL_V=assume_full_v, + ASSUME_VALID_PAGES=assume_valid_pages, + PAGE_SPLIT=page_split, + PAGES_PER_SPLIT=pages_per_split, + PARTIAL_OUT=True, + partial_out_ptr=partial_out, + partial_s0=partial_out.stride(0), + partial_s1=partial_out.stride(1), + partial_s2=partial_out.stride(2), + partial_s3=partial_out.stride(3), + **launch_meta, + ) + _attn_pv_reduce_kernel[(num_queries, num_head_blocks, num_dim_blocks)]( + output, + partial_out, + global_scale, + output.stride(0), + output.stride(1), + output.stride(2), + partial_out.stride(0), + partial_out.stride(1), + partial_out.stride(2), + partial_out.stride(3), + NUM_HEADS=num_heads, + V_HEAD_D=kv_lora_rank, + PAGE_SPLIT=page_split, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + BLOCK_H=block_h, + BLOCK_V=block_v, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_V=assume_full_v, + **launch_meta, + ) + else: + _attn_pv_kernel[(num_queries, num_head_blocks, num_dim_blocks)]( + output, + p_fp4, + p_sf, + kv_cache, + v_sf, + global_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + kv_cache.shape[0], + output.stride(0), + output.stride(1), + output.stride(2), + output.shape[0] * output.shape[1], + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=kv_lora_rank, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + BLOCK_H=block_h, + BLOCK_V=block_v, + USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, + USE_TMA_V_LOAD=use_tma_data_load and kv_lora_rank % block_v == 0, + PV_LOOP_STAGES=pv_loop_stages, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_FULL_V=assume_full_v, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + + +def run_fp4_mla_attention_decode( + metadata: Any, + layer_idx: int, + local_layer: int, + q_nope: torch.Tensor, + q_pe: torch.Tensor, + output: torch.Tensor, + *, + sm_scale: float, + kv_lora_rank: int, + qk_rope_head_dim: int, +) -> None: + """Run MLA decode with FP4 QK and FP4 PV tensor-core matmuls. + + Q is quantized to FP4, QK reads the packed K-view cache with swizzled + block scales, softmax probabilities are quantized to FP4 per page, and PV + repacks V nibbles from the shared KV cache while reading the auxiliary + V-view scale pool. No BF16 dequantized KV workspace is materialized on + this path. + """ + if not is_flashinfer_fp4_mla_attention_enabled(): + raise RuntimeError( + f"FP4 MLA attention decode requires {FLASHINFER_FP4_MLA_ATTENTION_ENV}=1." + ) + + head_dim = kv_lora_rank + qk_rope_head_dim + _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) + if metadata.page_size != FP4_MLA_TOKENS_PER_BLOCK: + raise ValueError( + f"FP4 MLA attention decode requires page_size={FP4_MLA_TOKENS_PER_BLOCK}, " + f"got {metadata.page_size}." + ) + + num_queries = q_nope.shape[0] + if num_queries == 0: + return + num_gen_seqs = metadata.num_seqs - metadata.num_contexts + query_len_per_seq = _get_linear_mtp_query_len_per_seq( + metadata, + num_queries=num_queries, + num_gen_seqs=num_gen_seqs, + ) + + num_heads = q_nope.shape[1] + if q_pe.shape[:2] != (num_queries, num_heads): + raise ValueError("FP4 MLA attention q_nope/q_pe batch dimensions do not match.") + if output.shape[:2] != (num_queries, num_heads): + raise ValueError("FP4 MLA attention output batch dimensions do not match.") + if q_nope.shape[-1] != kv_lora_rank: + raise ValueError( + f"q_nope last dimension must be kv_lora_rank={kv_lora_rank}, got {q_nope.shape[-1]}." + ) + if q_pe.shape[-1] != qk_rope_head_dim: + raise ValueError( + f"q_pe last dimension must be qk_rope_head_dim={qk_rope_head_dim}, " + f"got {q_pe.shape[-1]}." + ) + + if getattr(metadata, "fp4_mla_v_scale_pool", None) is None: + raise RuntimeError( + "FP4 MLA attention decode requires the auxiliary V scale pool to be allocated." + ) + + global_scale = _get_fp4_mla_global_scale(metadata, q_nope.device) + q_residual_dim = FP4_MLA_Q_RESIDUAL_DIM + _validate_fp4_mla_attention_q_shape(head_dim, q_residual_dim) + + q_full = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_q_buf", + (num_queries, num_heads, head_dim), + dtype=q_nope.dtype, + device=q_nope.device, + ) + q_full[..., :kv_lora_rank].copy_(q_nope) + q_full[..., kv_lora_rank:].copy_(q_pe) + q_2d = q_full.reshape(num_queries * num_heads, head_dim) + if q_2d.dtype not in (torch.bfloat16, torch.float8_e4m3fn): + raise TypeError( + f"FP4 MLA residual Q quantization requires BF16 or FP8 Q; got {q_2d.dtype}." + ) + # Per-page path: Q gets its own per-step dynamic global scale (independent of + # the KV per-page scale); the dev QK kernel divides by q_gscale * page_gscale. + page_scale_pool = getattr(metadata, "fp4_mla_page_scale_pool", None) + use_per_page_scale = fp4_mla_per_page_scale_enabled() and page_scale_pool is not None + if use_per_page_scale: + q_amax = q_2d.to(torch.float32).abs().amax().reshape(1) + q_global_scale = torch.where( + q_amax > 0.0, + torch.full_like(q_amax, FP4_MLA_P_GLOBAL_SCALE) / q_amax, + torch.ones_like(q_amax), + ) + else: + q_global_scale = global_scale + q_fp4, q_sf = torch.ops.trtllm.fp4_quantize_with_residual( + q_2d, + q_global_scale, + q_residual_dim, + is_act=True, + ) + q_sf = q_sf.view(torch.float8_e4m3fn) + + kv_cache, sf_cache = _get_fp4_mla_kv_cache_tensors(metadata, layer_idx) + sf_cache = sf_cache.view(torch.float8_e4m3fn) + + num_gen_blocks = metadata.num_generation_blocks + v_sf = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=kv_lora_rank)[local_layer].view( + torch.float8_e4m3fn + ) + + src_page_ids = metadata.paged_kv_indices[ + metadata.num_context_blocks : metadata.num_context_blocks + num_gen_blocks + ] + kv_lens = metadata.kv_lens_cuda_runtime[metadata.num_contexts : metadata.num_seqs] + max_pages = _max_generation_pages(metadata) + if max_pages == 0: + return + + _assert_fp4_mla_decode_paging_consistent(metadata, kv_lens, num_gen_blocks, query_len_per_seq) + + backend = _fp4_mla_attention_backend() + if backend == "cutile": + from .fp4_mla_cutile import fp4_mla_paged_attention + + total_p_rows = num_queries * max_pages * num_heads + p_fp4 = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_p_buf", + (max(total_p_rows, 1), metadata.page_size // 2), + dtype=torch.uint8, + device=q_nope.device, + )[:total_p_rows] + p_sf = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_p_sf_buf", + (max(_get_fp4_mla_swizzled_scale_size(total_p_rows, metadata.page_size), 1),), + dtype=torch.float8_e4m3fn, + device=q_nope.device, + ) + stats_shape = (num_queries, num_heads) + max_scores = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_max_buf", + stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + denom = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_denom_buf", + stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + page_max = None + page_sum = None + if max_pages >= 8: + page_stats_shape = (num_queries, max_pages, num_heads) + page_max = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_page_max_buf", + page_stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + page_sum = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_page_sum_buf", + page_stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + assume_full_pages = ( + _infer_cutile_assume_full_pages( + metadata, + max_pages, + metadata.page_size, + ) + and query_len_per_seq == 1 + ) + assume_valid_pages = False + cutile_block_h = _env_int("TRTLLM_FP4_MLA_BLOCK_H") or 128 + cutile_block_v = _env_int("TRTLLM_FP4_MLA_BLOCK_V") or 128 + cutile_num_gen_seqs = num_queries // query_len_per_seq + cutile_assume_valid_pages = assume_valid_pages or ( + assume_full_pages and src_page_ids.numel() == cutile_num_gen_seqs * max_pages + ) + cutile_auto_prepack_v = ( + hasattr(tl, "make_tensor_descriptor") + and num_heads % cutile_block_h == 0 + and assume_full_pages + and cutile_assume_valid_pages + and kv_lora_rank == 512 + and metadata.page_size == FP4_MLA_TOKENS_PER_BLOCK + and cutile_block_h in (64, 128) + and cutile_block_v == 128 + and metadata.page_size // FP4_BLOCK_SIZE == 8 + ) + cutile_prepack_v_env = os.environ.get("TRTLLM_FP4_MLA_PREPACK_V") + cutile_prepack_v_for_pv = ( + cutile_auto_prepack_v + if cutile_prepack_v_env is None + else cutile_prepack_v_env == "1" and cutile_auto_prepack_v + ) + v_packed = ( + _get_cutile_v_packed_cache( + metadata, + layer_idx, + kv_cache, + v_head_dim=kv_lora_rank, + page_size=metadata.page_size, + ) + if cutile_auto_prepack_v + else None + ) + use_cutile_v_packed_cache = v_packed is not None + if use_cutile_v_packed_cache: + cutile_prepack_v_for_pv = False + elif cutile_prepack_v_for_pv: + v_packed = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_v_packed_buf", + ( + kv_cache.shape[0] * triton.cdiv(kv_lora_rank, cutile_block_v) * cutile_block_v, + metadata.page_size // 2, + ), + dtype=torch.uint8, + device=q_nope.device, + ) + _fp4_mla_debug( + "attention decode cutile launch: " + f"num_queries={num_queries} query_len_per_seq={query_len_per_seq} " + f"num_heads={num_heads} local_layer={local_layer} " + f"layer_idx={layer_idx} head_dim={head_dim} kv_lora_rank={kv_lora_rank} " + f"rope_dim={qk_rope_head_dim} max_pages={max_pages} " + f"assume_full_pages={assume_full_pages} " + f"assume_valid_pages={assume_valid_pages} " + f"use_v_packed_cache={use_cutile_v_packed_cache}" + ) + fp4_mla_paged_attention( + q_fp4, + q_sf, + kv_cache, + sf_cache, + v_sf, + global_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + output, + sm_scale=float(sm_scale), + num_heads=num_heads, + v_head_dim=kv_lora_rank, + page_size=metadata.page_size, + q_residual_dim=q_residual_dim, + max_pages=max_pages, + query_len_per_seq=query_len_per_seq, + assume_full_pages=assume_full_pages, + assume_valid_pages=assume_valid_pages, + prepack_v_for_pv=cutile_prepack_v_for_pv, + use_prepacked_v_for_pv=use_cutile_v_packed_cache, + p_fp4_workspace=p_fp4, + p_sf_workspace=p_sf, + v_packed_workspace=v_packed, + max_scores_workspace=max_scores, + denom_workspace=denom, + page_max_workspace=page_max, + page_sum_workspace=page_sum, + ) + _debug_sync("attention_cutile") + return + + total_p_rows = num_queries * max_pages * num_heads + p_fp4 = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_p_buf", + (max(total_p_rows, 1), metadata.page_size // 2), + dtype=torch.uint8, + device=q_nope.device, + )[:total_p_rows] + p_sf = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_p_sf_buf", + (max(_get_fp4_mla_swizzled_scale_size(total_p_rows, metadata.page_size), 1),), + dtype=torch.float8_e4m3fn, + device=q_nope.device, + ) + stats_shape = (num_queries, num_heads) + max_scores = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_max_buf", + stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + denom = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_denom_buf", + stats_shape, + dtype=torch.float32, + device=q_nope.device, + ) + + if backend != "triton": + raise ValueError( + f"Unsupported FP4 MLA attention backend '{backend}'. " + f"Set {FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV} to " + "'triton' or 'cutile'." + ) + + # Self-contained triton path. Uses kernels copied from the cutile reference + # into fp4_mla_triton.py: TMA-loaded QK + fused page-stats pack, + # reduce-stats, prob-scale, and a specialized PV with tl.ext views. + if backend == "triton": + _run_triton_attention_decode( + metadata=metadata, + q_fp4=q_fp4, + q_sf=q_sf.contiguous().view(-1), + kv_cache=kv_cache, + sf_cache=sf_cache, + v_sf=v_sf, + global_scale=global_scale, + src_page_ids=src_page_ids, + kv_lens=kv_lens, + p_fp4=p_fp4, + p_sf=p_sf, + max_scores=max_scores, + denom=denom, + output=output, + num_queries=num_queries, + num_heads=num_heads, + head_dim=head_dim, + kv_lora_rank=kv_lora_rank, + q_residual_dim=q_residual_dim, + query_len_per_seq=query_len_per_seq, + max_pages=max_pages, + sm_scale=float(sm_scale), + q_global_scale=q_global_scale, + page_scale_pool=page_scale_pool, + local_layer=local_layer, + use_per_page_scale=use_per_page_scale, + ) + _debug_sync("attention_triton") + return + + +def _hp_pool_layer_view( + pool: torch.Tensor, + local_layer: int, + pool_head_dim: int, +) -> torch.Tensor: + return pool[:, local_layer, 0, :].view(pool.shape[0], HP_BLOCK_SIZE, pool_head_dim) + + +def _snapshot_hp_kv_for_mtp_generation( + metadata: Any, + pool: torch.Tensor, + local_layer: int, + *, + num_gen: int, + num_gen_tokens: int, + max_gen_len: int, + metadata_token_offset: int, + head_dim: int, + pool_head_dim: int, +) -> None: + if getattr(metadata, "is_warmup", False): + return + if num_gen_tokens <= num_gen: + return + if max_gen_len > HP_BLOCK_SIZE: + raise NotImplementedError( + "FP4 MLA HP-pool rollback for linear MTP supports at most " + f"{HP_BLOCK_SIZE} generation tokens per sequence, got {max_gen_len}." + ) + + end_token_offset = metadata_token_offset + num_gen_tokens + if ( + end_token_offset > metadata.batch_indices.shape[0] + or end_token_offset > metadata.positions.shape[0] + ): + raise RuntimeError( + "FP4 MLA HP-pool snapshot would read past generation metadata: " + f"token_offset={metadata_token_offset}, num_gen_tokens={num_gen_tokens}, " + f"batch_indices={metadata.batch_indices.shape[0]}, " + f"positions={metadata.positions.shape[0]}." + ) + + snapshots = getattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, None) + if snapshots is None: + snapshots = {} + setattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, snapshots) + + snapshot_pool = getattr(metadata, "fp4_mla_hp_snapshot_pool", None) + if snapshot_pool is not None: + snapshot_pool[:, local_layer, :, :].copy_(pool[:, local_layer, :, :]) + snapshots[int(local_layer)] = { + "mode": "pool", + "metadata_token_offset": metadata_token_offset, + "num_gen_tokens": num_gen_tokens, + "head_dim": head_dim, + "pool_head_dim": pool_head_dim, + } + return + + device = pool.device + token_indices = torch.arange( + metadata_token_offset, + end_token_offset, + dtype=torch.long, + device=device, + ) + batch_indices = metadata.batch_indices[token_indices].to(torch.long) + positions = metadata.positions[token_indices].to(torch.long) + seq_slots = metadata.seq_slots[batch_indices].to(torch.long) + hp_slots = torch.remainder(positions, HP_BLOCK_SIZE).to(torch.long) + first_new_positions = metadata.kv_lens_cuda_runtime[batch_indices].to( + torch.long + ) - metadata.prompt_lens_cuda_runtime[batch_indices].to(torch.long) + pool_view = _hp_pool_layer_view(pool, local_layer, pool_head_dim) + values = pool_view[seq_slots, hp_slots, :head_dim].clone() + + snapshots[int(local_layer)] = { + "mode": "values", + "batch_indices": batch_indices, + "seq_slots": seq_slots, + "hp_slots": hp_slots, + "positions": positions, + "first_new_positions": first_new_positions, + "values": values, + "head_dim": head_dim, + "pool_head_dim": pool_head_dim, + } + + +def repair_fp4_mla_hp_kv_for_mtp_rejection( + metadata: Any, + num_accepted_tokens: torch.Tensor, +) -> None: + """Restore HP-pool slots that belonged to rejected linear-MTP tokens. + + Packed FP4 pages past the accepted logical KV length are harmless because + later attention ignores them. The BF16 HP pool is a circular tail mirror, so + rejected speculative writes must be rolled back before the next tile rewrite. + """ + snapshots = getattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, None) + if not snapshots: + return + + keep_snapshots = False + try: + pool = getattr(metadata, "high_precision_kv_pool", None) + if pool is None: + return + accepted_tokens = num_accepted_tokens.to(device=pool.device) + for local_layer, snapshot in snapshots.items(): + head_dim = snapshot["head_dim"] + pool_head_dim = snapshot["pool_head_dim"] + block_d = triton.next_power_of_2(head_dim) + if snapshot.get("mode") == "pool": + keep_snapshots = True + _hp_kv_restore_rejected_from_pool_kernel[(snapshot["num_gen_tokens"],)]( + pool, + metadata.fp4_mla_hp_snapshot_pool, + metadata.batch_indices, + metadata.positions, + metadata.seq_slots, + metadata.kv_lens_cuda_runtime, + metadata.prompt_lens_cuda_runtime, + accepted_tokens, + snapshot["metadata_token_offset"], + snapshot["num_gen_tokens"], + metadata.batch_indices.shape[0], + metadata.num_seqs, + accepted_tokens.shape[0], + pool.shape[0], + pool.shape[1], + int(local_layer), + pool.stride(0), + pool.stride(1), + metadata.fp4_mla_hp_snapshot_pool.stride(0), + metadata.fp4_mla_hp_snapshot_pool.stride(1), + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + else: + positions = snapshot["positions"] + batch_indices = snapshot["batch_indices"] + seq_slots = snapshot["seq_slots"] + hp_slots = snapshot["hp_slots"] + first_new_positions = snapshot["first_new_positions"] + if positions.shape[0] == 0: + continue + _hp_kv_restore_rejected_from_values_kernel[(positions.shape[0],)]( + pool, + snapshot["values"], + batch_indices, + positions, + seq_slots, + hp_slots, + first_new_positions, + accepted_tokens, + positions.shape[0], + accepted_tokens.shape[0], + pool.shape[0], + pool.shape[1], + int(local_layer), + pool.stride(0), + pool.stride(1), + snapshot["values"].stride(0), + snapshot["values"].stride(1), + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + finally: + if not keep_snapshots: + setattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, None) + + +def update_hp_kv_for_fp4_mla( + metadata: Any, + latent_cache: Optional[torch.Tensor], + local_layer: int, + *, + phase: _HPUpdatePhase = "all", +) -> None: + """Store recent KV tokens at BF16 into the high-precision pool. + + Called on every layer before the attention kernel. The pool acts as a + circular buffer of HP_BLOCK_SIZE slots per sequence: + + Context phase stores the last ``kv_len % HP_BLOCK_SIZE`` new tokens of + each request into buffer positions [0, remainder). These are the + tail tokens that do not fill a complete FP4 block of 16. + + Generation phase stores every new token for each request into position + ``position % HP_BLOCK_SIZE``, overwriting the oldest entries in the + circular buffer. This supports linear MTP where a request contributes + more than one generation token in a forward pass. + + The Triton kernels use the GPU ``seq_slots`` tensor for scatter indexing + and are CUDA-graph-compatible for the generation phase. + + Args: + metadata: Attention metadata exposing ``num_contexts``, ``num_seqs``, + ``seq_slots`` / ``seq_slots_cpu``, ``request_ids``, + ``is_cuda_graph``, ``is_warmup``, ``hp_pool_owners``, + ``high_precision_kv_pool``, ``prompt_lens_cpu_runtime``, + ``prompt_lens_cuda_runtime``, ``kv_lens_cuda_runtime``. + latent_cache: MLA latent cache for the current tokens, shape + [num_tokens, head_dim]. When ``None``, only ownership tracking + runs (no data is written to the pool). + local_layer: Layer index within the local pipeline-parallel slice. + phase: Which portion of ``latent_cache`` is present. ``"all"`` means + context tokens followed by generation tokens, ``"context"`` means + only context tokens, and ``"generation"`` means only generation + tokens. + """ + if phase not in ("all", "context", "generation"): + raise ValueError(f"Unexpected FP4 MLA HP update phase: {phase}") + if metadata.hp_pool_owners is None: + return + num_contexts = metadata.num_contexts + num_seqs = metadata.num_seqs + update_context = phase in ("all", "context") + update_generation = phase in ("all", "generation") + + # ------------------------------------------------------------------ + # Ownership tracking (layer 0, eager mode only - debug guard). + # Context phase never uses CUDA graph; decode check is debug-only. + # ------------------------------------------------------------------ + if local_layer == 0 and not metadata.is_cuda_graph and not metadata.is_warmup: + # Context: register ownership of each seq_slot. + if update_context: + for batch_idx in range(num_contexts): + seq_slot = metadata.seq_slots_cpu[batch_idx].item() + request_id = metadata.request_ids[batch_idx] + metadata.hp_pool_owners[seq_slot] = request_id + + # Decode: verify that the expected request still owns each slot. + if update_generation: + for batch_idx in range(num_contexts, num_seqs): + seq_slot = metadata.seq_slots_cpu[batch_idx].item() + request_id = metadata.request_ids[batch_idx] + owner = metadata.hp_pool_owners.get(seq_slot) + if owner != request_id: + raise RuntimeError( + f"HP KV pool ownership mismatch: seq_slot={seq_slot} " + f"is owned by request {owner} but request " + f"{request_id} is attempting to use it" + ) + + if latent_cache is None: + return + + # ------------------------------------------------------------------ + # Triton kernel dispatch - runs on every layer, CUDA-graph-safe. + # ------------------------------------------------------------------ + pool = metadata.high_precision_kv_pool + head_dim = latent_cache.shape[-1] + pool_head_dim = pool.shape[-1] // HP_BLOCK_SIZE + if pool_head_dim < head_dim: + raise RuntimeError( + f"FP4 MLA HP pool head dimension is too small: got " + f"{pool_head_dim}, need at least {head_dim}." + ) + block_d = triton.next_power_of_2(head_dim) + pool_s0 = pool.stride(0) # stride across sequence slots + pool_s1 = pool.stride(1) # stride across layers + lc_stride = latent_cache.stride(0) + _fp4_mla_debug( + "hp update: " + f"phase={phase} local_layer={local_layer} num_contexts={num_contexts} " + f"num_seqs={num_seqs} head_dim={head_dim} " + f"pool_head_dim={pool_head_dim} block_d={block_d}" + ) + _fp4_mla_debug(f"hp latent_cache: {_tensor_layout(latent_cache)}") + _fp4_mla_debug(f"hp pool: {_tensor_layout(pool)}") + _debug_tensor_range("hp seq_slots", metadata.seq_slots[:num_seqs]) + _debug_tensor_range("hp kv_lens", metadata.kv_lens_cuda_runtime[:num_seqs]) + + # Context phase: store last (kv_len % HP_BLOCK_SIZE) new tokens. + if update_context and num_contexts > 0: + prompt_lens_cpu = metadata.prompt_lens_cpu_runtime[:num_contexts] + # Exclusive prefix sum: token offset in latent_cache for each ctx seq. + token_offsets_cpu = torch.zeros(num_contexts, dtype=torch.int32, device="cpu") + if num_contexts > 1: + token_offsets_cpu[1:].copy_(torch.cumsum(prompt_lens_cpu[:-1].to(torch.int32), dim=0)) + token_offsets_gpu = token_offsets_cpu.to(pool.device, non_blocking=False) + prompt_lens_gpu = metadata.prompt_lens_cuda_runtime[:num_contexts] + + _fp4_mla_debug( + "hp context launch: " + f"grid=({num_contexts}, {HP_BLOCK_SIZE}) " + f"token_offsets={token_offsets_cpu.tolist()}" + ) + _debug_tensor_range("hp context prompt_lens", prompt_lens_gpu) + _hp_kv_store_context_kernel[(num_contexts, HP_BLOCK_SIZE)]( + pool, + latent_cache, + metadata.seq_slots, + metadata.kv_lens_cuda_runtime, + token_offsets_gpu, + prompt_lens_gpu, + pool.shape[0], + pool.shape[1], + local_layer, + pool_s0, + pool_s1, + lc_stride, + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + _debug_sync("hp_context") + + # Generation phase: store current tokens at position % HP_BLOCK_SIZE. + num_gen = num_seqs - num_contexts + if update_generation and num_gen > 0: + gen_tok_start = 0 + metadata_token_offset = getattr(metadata, "num_ctx_tokens", 0) + if phase == "all": + # Scalar offset: number of context tokens packed before gen tokens. + gen_tok_start = int(metadata.prompt_lens_cpu_runtime[:num_contexts].sum().item()) + metadata_token_offset = gen_tok_start + num_gen_tokens = latent_cache.shape[0] - gen_tok_start + if num_gen_tokens < 0: + raise RuntimeError( + "FP4 MLA HP generation update received fewer latent tokens than " + f"the context prefix: latent_tokens={latent_cache.shape[0]}, " + f"context_tokens={gen_tok_start}." + ) + if num_gen_tokens == 0: + return + + gen_token_lens = _host_int_list_during_forward( + getattr(metadata, "prompt_lens_cpu_runtime", None), num_contexts, num_seqs + ) + if gen_token_lens is not None: + max_gen_len = max(gen_token_lens) + elif num_gen_tokens % num_gen == 0: + max_gen_len = num_gen_tokens // num_gen + else: + max_gen_len = num_gen_tokens + _snapshot_hp_kv_for_mtp_generation( + metadata, + pool, + local_layer, + num_gen=num_gen, + num_gen_tokens=num_gen_tokens, + max_gen_len=max_gen_len, + metadata_token_offset=metadata_token_offset, + head_dim=head_dim, + pool_head_dim=pool_head_dim, + ) + + _fp4_mla_debug( + "hp generation launch: " + f"grid=({num_gen_tokens},) gen_tok_start={gen_tok_start} " + f"metadata_token_offset={metadata_token_offset}" + ) + _hp_kv_store_gen_kernel[(num_gen_tokens,)]( + pool, + latent_cache, + metadata.seq_slots, + metadata.batch_indices, + metadata.positions, + gen_tok_start, + metadata_token_offset, + num_gen_tokens, + metadata.batch_indices.shape[0], + pool.shape[0], + pool.shape[1], + local_layer, + pool_s0, + pool_s1, + lc_stride, + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + _debug_sync("hp_generation") + + +def _stage_pool_layer_view( + pool: torch.Tensor, + local_layer: int, + pool_head_dim: int, + slots: int, +) -> torch.Tensor: + return pool[:, local_layer, 0, :].view(pool.shape[0], slots, pool_head_dim) + + +def _snapshot_page_stage_for_mtp_generation( + metadata: Any, + pool: torch.Tensor, + local_layer: int, + *, + slots: int, + num_gen: int, + num_gen_tokens: int, + max_gen_len: int, + metadata_token_offset: int, + head_dim: int, + pool_head_dim: int, +) -> None: + """Snapshot staging-pool slots before linear-MTP generation writes. + + Mirror of ``_snapshot_hp_kv_for_mtp_generation`` for the per-page staging + buffer, keyed under ``_FP4_MLA_MTP_STAGE_SNAPSHOTS`` so it never aliases the + HP-pool snapshots. ``slots`` is the page size (vs. the HP pool's + ``HP_BLOCK_SIZE``), so the per-sequence MTP draft length must not exceed it. + """ + if getattr(metadata, "is_warmup", False): + return + if num_gen_tokens <= num_gen: + return + if max_gen_len > slots: + raise NotImplementedError( + "FP4 MLA staging rollback for linear MTP supports at most " + f"{slots} generation tokens per sequence, got {max_gen_len}." + ) + + end_token_offset = metadata_token_offset + num_gen_tokens + if ( + end_token_offset > metadata.batch_indices.shape[0] + or end_token_offset > metadata.positions.shape[0] + ): + raise RuntimeError( + "FP4 MLA staging snapshot would read past generation metadata: " + f"token_offset={metadata_token_offset}, num_gen_tokens={num_gen_tokens}, " + f"batch_indices={metadata.batch_indices.shape[0]}, " + f"positions={metadata.positions.shape[0]}." + ) + + snapshots = getattr(metadata, _FP4_MLA_MTP_STAGE_SNAPSHOTS, None) + if snapshots is None: + snapshots = {} + setattr(metadata, _FP4_MLA_MTP_STAGE_SNAPSHOTS, snapshots) + + snapshot_pool = getattr(metadata, "fp4_mla_page_stage_snapshot_pool", None) + if snapshot_pool is not None: + snapshot_pool[:, local_layer, :, :].copy_(pool[:, local_layer, :, :]) + snapshots[int(local_layer)] = { + "mode": "pool", + "metadata_token_offset": metadata_token_offset, + "num_gen_tokens": num_gen_tokens, + "head_dim": head_dim, + "pool_head_dim": pool_head_dim, + "slots": slots, + } + return + + device = pool.device + token_indices = torch.arange( + metadata_token_offset, + end_token_offset, + dtype=torch.long, + device=device, + ) + batch_indices = metadata.batch_indices[token_indices].to(torch.long) + positions = metadata.positions[token_indices].to(torch.long) + seq_slots = metadata.seq_slots[batch_indices].to(torch.long) + stage_slots = torch.remainder(positions, slots).to(torch.long) + first_new_positions = metadata.kv_lens_cuda_runtime[batch_indices].to( + torch.long + ) - metadata.prompt_lens_cuda_runtime[batch_indices].to(torch.long) + pool_view = _stage_pool_layer_view(pool, local_layer, pool_head_dim, slots) + values = pool_view[seq_slots, stage_slots, :head_dim].clone() + + snapshots[int(local_layer)] = { + "mode": "values", + "batch_indices": batch_indices, + "seq_slots": seq_slots, + "hp_slots": stage_slots, + "positions": positions, + "first_new_positions": first_new_positions, + "values": values, + "head_dim": head_dim, + "pool_head_dim": pool_head_dim, + "slots": slots, + } + + +def update_page_stage_for_fp4_mla( + metadata: Any, + latent_cache: Optional[torch.Tensor], + local_layer: int, + *, + phase: _HPUpdatePhase = "all", +) -> None: + """Store recent KV tokens at BF16 into the per-page staging buffer. + + Separate twin of ``update_hp_kv_for_fp4_mla`` for the per-page + dynamic-scale path. Identical circular-buffer semantics, but the buffer + holds a whole page (``metadata.page_size`` slots) instead of the HP pool's + ``HP_BLOCK_SIZE`` slots, so the active page can be re-quantized to FP4 with + its exact per-page amax. Reuses the HP store kernels with ``HP_BLOCK`` set + to the page size. Ownership tracking is intentionally omitted -- the HP + update already validates it on this path. + """ + if phase not in ("all", "context", "generation"): + raise ValueError(f"Unexpected FP4 MLA staging update phase: {phase}") + pool = getattr(metadata, "fp4_mla_page_stage_pool", None) + if pool is None or latent_cache is None: + return + + num_contexts = metadata.num_contexts + num_seqs = metadata.num_seqs + update_context = phase in ("all", "context") + update_generation = phase in ("all", "generation") + + slots = metadata.page_size + head_dim = latent_cache.shape[-1] + pool_head_dim = pool.shape[-1] // slots + if pool_head_dim < head_dim: + raise RuntimeError( + f"FP4 MLA staging pool head dimension is too small: got " + f"{pool_head_dim}, need at least {head_dim}." + ) + block_d = triton.next_power_of_2(head_dim) + pool_s0 = pool.stride(0) + pool_s1 = pool.stride(1) + lc_stride = latent_cache.stride(0) + + # Context phase: store last (kv_len % page_size) new tokens. + if update_context and num_contexts > 0: + prompt_lens_cpu = metadata.prompt_lens_cpu_runtime[:num_contexts] + token_offsets_cpu = torch.zeros(num_contexts, dtype=torch.int32, device="cpu") + if num_contexts > 1: + token_offsets_cpu[1:].copy_(torch.cumsum(prompt_lens_cpu[:-1].to(torch.int32), dim=0)) + token_offsets_gpu = token_offsets_cpu.to(pool.device, non_blocking=False) + prompt_lens_gpu = metadata.prompt_lens_cuda_runtime[:num_contexts] + + _hp_kv_store_context_kernel[(num_contexts, slots)]( + pool, + latent_cache, + metadata.seq_slots, + metadata.kv_lens_cuda_runtime, + token_offsets_gpu, + prompt_lens_gpu, + pool.shape[0], + pool.shape[1], + local_layer, + pool_s0, + pool_s1, + lc_stride, + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=slots, + ) + _debug_sync("stage_context") + + # Generation phase: store current tokens at position % page_size. + num_gen = num_seqs - num_contexts + if update_generation and num_gen > 0: + gen_tok_start = 0 + metadata_token_offset = getattr(metadata, "num_ctx_tokens", 0) + if phase == "all": + gen_tok_start = int(metadata.prompt_lens_cpu_runtime[:num_contexts].sum().item()) + metadata_token_offset = gen_tok_start + num_gen_tokens = latent_cache.shape[0] - gen_tok_start + if num_gen_tokens < 0: + raise RuntimeError( + "FP4 MLA staging generation update received fewer latent tokens " + f"than the context prefix: latent_tokens={latent_cache.shape[0]}, " + f"context_tokens={gen_tok_start}." + ) + if num_gen_tokens == 0: + return + + gen_token_lens = _host_int_list_during_forward( + getattr(metadata, "prompt_lens_cpu_runtime", None), num_contexts, num_seqs + ) + if gen_token_lens is not None: + max_gen_len = max(gen_token_lens) + elif num_gen_tokens % num_gen == 0: + max_gen_len = num_gen_tokens // num_gen + else: + max_gen_len = num_gen_tokens + _snapshot_page_stage_for_mtp_generation( + metadata, + pool, + local_layer, + slots=slots, + num_gen=num_gen, + num_gen_tokens=num_gen_tokens, + max_gen_len=max_gen_len, + metadata_token_offset=metadata_token_offset, + head_dim=head_dim, + pool_head_dim=pool_head_dim, + ) + + _hp_kv_store_gen_kernel[(num_gen_tokens,)]( + pool, + latent_cache, + metadata.seq_slots, + metadata.batch_indices, + metadata.positions, + gen_tok_start, + metadata_token_offset, + num_gen_tokens, + metadata.batch_indices.shape[0], + pool.shape[0], + pool.shape[1], + local_layer, + pool_s0, + pool_s1, + lc_stride, + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=slots, + ) + _debug_sync("stage_generation") + + +def repair_fp4_mla_page_stage_for_mtp_rejection( + metadata: Any, + num_accepted_tokens: torch.Tensor, +) -> None: + """Restore staging-pool slots that belonged to rejected linear-MTP tokens. + + Separate twin of ``repair_fp4_mla_hp_kv_for_mtp_rejection`` for the per-page + staging buffer. Reuses the HP restore kernels with ``HP_BLOCK`` set to the + snapshot's page size. + """ + snapshots = getattr(metadata, _FP4_MLA_MTP_STAGE_SNAPSHOTS, None) + if not snapshots: + return + + keep_snapshots = False + try: + pool = getattr(metadata, "fp4_mla_page_stage_pool", None) + if pool is None: + return + accepted_tokens = num_accepted_tokens.to(device=pool.device) + for local_layer, snapshot in snapshots.items(): + head_dim = snapshot["head_dim"] + pool_head_dim = snapshot["pool_head_dim"] + slots = snapshot["slots"] + block_d = triton.next_power_of_2(head_dim) + if snapshot.get("mode") == "pool": + keep_snapshots = True + _hp_kv_restore_rejected_from_pool_kernel[(snapshot["num_gen_tokens"],)]( + pool, + metadata.fp4_mla_page_stage_snapshot_pool, + metadata.batch_indices, + metadata.positions, + metadata.seq_slots, + metadata.kv_lens_cuda_runtime, + metadata.prompt_lens_cuda_runtime, + accepted_tokens, + snapshot["metadata_token_offset"], + snapshot["num_gen_tokens"], + metadata.batch_indices.shape[0], + metadata.num_seqs, + accepted_tokens.shape[0], + pool.shape[0], + pool.shape[1], + int(local_layer), + pool.stride(0), + pool.stride(1), + metadata.fp4_mla_page_stage_snapshot_pool.stride(0), + metadata.fp4_mla_page_stage_snapshot_pool.stride(1), + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=slots, + ) + else: + positions = snapshot["positions"] + batch_indices = snapshot["batch_indices"] + seq_slots = snapshot["seq_slots"] + stage_slots = snapshot["hp_slots"] + first_new_positions = snapshot["first_new_positions"] + if positions.shape[0] == 0: + continue + _hp_kv_restore_rejected_from_values_kernel[(positions.shape[0],)]( + pool, + snapshot["values"], + batch_indices, + positions, + seq_slots, + stage_slots, + first_new_positions, + accepted_tokens, + positions.shape[0], + accepted_tokens.shape[0], + pool.shape[0], + pool.shape[1], + int(local_layer), + pool.stride(0), + pool.stride(1), + snapshot["values"].stride(0), + snapshot["values"].stride(1), + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=slots, + ) + finally: + if not keep_snapshots: + setattr(metadata, _FP4_MLA_MTP_STAGE_SNAPSHOTS, None) diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_cute.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_cute.py deleted file mode 100644 index fca764b1b5f8..000000000000 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_cute.py +++ /dev/null @@ -1,803 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -"""CuTe DSL FP4 MLA decode backend. - -This module intentionally preserves the Python FP4 MLA decode contract from -``fp4_mla_kv.run_fp4_mla_attention_decode``. It consumes the same packed Q, -packed paged KV cache, swizzled scale tensors, page tables, and workspace -buffers as the Triton backend. -""" - -import math - -import torch - -from ...logger import logger -from ..cute_dsl_utils import IS_CUTLASS_DSL_AVAILABLE - -if IS_CUTLASS_DSL_AVAILABLE: - try: - from cuda.bindings import driver as cuda - except ImportError: - from cuda import cuda - - import cutlass - import cutlass.cute as cute - from cutlass._mlir.dialects import llvm - from cutlass.cute.runtime import from_dlpack - from cutlass.cutlass_dsl import T, dsl_user_op - - class _CUDAGraphCompatibleWrapper: - """Wrapper to make DLPack export safe during CUDA graph capture.""" - - def __init__(self, tensor: torch.Tensor) -> None: - self._tensor = tensor - - def __dlpack__(self, stream=None): - return self._tensor.__dlpack__(stream=-1) - - def __dlpack_device__(self): - return self._tensor.__dlpack_device__() - - def _to_cute(tensor: torch.Tensor) -> cute.Tensor: - return from_dlpack( - _CUDAGraphCompatibleWrapper(tensor.detach()), assumed_align=16 - ).mark_layout_dynamic() - - @cute.jit - def _swizzled_sf_offset(row_idx, col_idx, sf_per_token: cutlass.Constexpr): - padded_cols = ((sf_per_token + 3) // 4) * 4 - return ( - col_idx % 4 - + (col_idx // 4) * (4 * 128) - + (row_idx % 32) * 16 - + ((row_idx % 128) // 32) * 4 - + (row_idx // 128) * (128 * padded_cols) - ) - - @dsl_user_op - def _ptx_fp4_e2m1x2_to_f16x2(byte, *, loc=None, ip=None) -> cutlass.Uint32: - return cutlass.Uint32( - llvm.inline_asm( - T.i32(), - [cutlass.Uint32(byte).ir_value(loc=loc, ip=ip)], - """ - { - .reg .b8 in_8; - .reg .f16x2 out; - cvt.u8.u32 in_8, $1; - cvt.rn.f16x2.e2m1x2 out, in_8; - mov.b32 $0, out; - } - """, - "=r,r", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - ) - ) - - @dsl_user_op - def _ptx_fp8_e4m3x2_to_f16x2(byte, *, loc=None, ip=None) -> cutlass.Uint32: - return cutlass.Uint32( - llvm.inline_asm( - T.i32(), - [cutlass.Uint32(byte).ir_value(loc=loc, ip=ip)], - """ - { - .reg .b16 in_16; - .reg .f16x2 out; - cvt.u16.u32 in_16, $1; - cvt.rn.f16x2.e4m3x2 out, in_16; - mov.b32 $0, out; - } - """, - "=r,r", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - ) - ) - - @dsl_user_op - def _ptx_fp4_e2m1x2_from_f32(even, odd, *, loc=None, ip=None) -> cutlass.Uint32: - return cutlass.Uint32( - llvm.inline_asm( - T.i32(), - [ - cutlass.Float32(odd).ir_value(loc=loc, ip=ip), - cutlass.Float32(even).ir_value(loc=loc, ip=ip), - ], - """ - { - .reg .b8 out; - cvt.rn.satfinite.e2m1x2.f32 out, $1, $2; - mov.b32 $0, {out, out, out, out}; - } - """, - "=r,f,f", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - ) - ) - - @dsl_user_op - def _ptx_fp8_e4m3x2_from_f32(low, high, *, loc=None, ip=None) -> cutlass.Uint32: - return cutlass.Uint32( - llvm.inline_asm( - T.i32(), - [ - cutlass.Float32(low).ir_value(loc=loc, ip=ip), - cutlass.Float32(high).ir_value(loc=loc, ip=ip), - ], - """ - { - .reg .b16 out; - cvt.rn.satfinite.e4m3x2.f32 out, $2, $1; - mov.b32 $0, {out, out}; - } - """, - "=r,f,f", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - ) - ) - - @cute.jit - def _low_f16x2_lane_to_f32(bits): - half_bits = cutlass.Uint16(bits & cutlass.Uint32(0xFFFF)) - half_value = cutlass.Float16(llvm.bitcast(cutlass.Float16.mlir_type, half_bits.ir_value())) - return half_value.to(cutlass.Float32) - - @cute.jit - def _high_f16x2_lane_to_f32(bits): - half_bits = cutlass.Uint16((bits >> 16) & cutlass.Uint32(0xFFFF)) - half_value = cutlass.Float16(llvm.bitcast(cutlass.Float16.mlir_type, half_bits.ir_value())) - return half_value.to(cutlass.Float32) - - @cute.jit - def _fp4_e2m1_to_f32(nibble): - bits = _ptx_fp4_e2m1x2_to_f16x2(cute.Uint8(nibble & cute.Uint8(0x0F))) - return _low_f16x2_lane_to_f32(bits) - - @cute.jit - def _fp4_e2m1_quantize_packed(even, odd): - bits = _ptx_fp4_e2m1x2_from_f32(even, odd) - return cute.Uint8(bits & cutlass.Uint32(0xFF)) - - @cute.jit - def _load_fp4_value(packed_tensor, packed_offset, elem_idx): - packed = packed_tensor[packed_offset] - nibble = cute.Uint8(packed & 0x0F) - if (elem_idx & 1) != 0: - nibble = cute.Uint8((packed >> 4) & 0x0F) - return _fp4_e2m1_to_f32(nibble) - - @cute.jit - def _load_fp4_byte_pair(packed_tensor, packed_offset): - """Load one packed FP4 byte and return both nibbles as (low, high) f32. - - The PTX ``cvt.f16x2.e2m1x2`` instruction converts both nibbles in a - single op; the previous scalar path discarded the high half. - """ - packed = packed_tensor[packed_offset] - bits = _ptx_fp4_e2m1x2_to_f16x2(cute.Uint8(packed)) - return _low_f16x2_lane_to_f32(bits), _high_f16x2_lane_to_f32(bits) - - @cute.jit - def _fp8_e4m3fn_to_f32(byte): - bits = _ptx_fp8_e4m3x2_to_f16x2(cute.Uint8(byte)) - return _low_f16x2_lane_to_f32(bits) - - @cute.jit - def _fp8_e4m3fn_positive_from_f32(value): - bits = _ptx_fp8_e4m3x2_from_f32(value, cutlass.Float32(0.0)) - return cute.Uint8(bits & cutlass.Uint32(0xFF)) - - class _Fp4MlaDecodeCuteKernel: - def __init__( - self, - *, - num_heads: int, - kv_lora_rank: int, - qk_rope_head_dim: int, - q_residual_dim: int, - page_size: int, - max_pages: int, - q_fp4_stride0: int, - q_fp4_stride1: int, - kv_stride0: int, - kv_stride2: int, - kv_stride4: int, - sf_stride0: int, - v_sf_stride0: int, - p_stride0: int, - p_stride1: int, - page_stats_stride0: int, - page_stats_stride1: int, - page_stats_stride2: int, - output_dtype: torch.dtype, - ) -> None: - self.num_heads = num_heads - self.kv_lora_rank = kv_lora_rank - self.qk_rope_head_dim = qk_rope_head_dim - self.q_residual_dim = q_residual_dim - self.k_head_dim = kv_lora_rank + qk_rope_head_dim - self.q_head_dim = self.k_head_dim + q_residual_dim - self.page_size = page_size - self.max_pages = max_pages - self.q_fp4_stride0 = q_fp4_stride0 - self.q_fp4_stride1 = q_fp4_stride1 - self.kv_stride0 = kv_stride0 - self.kv_stride2 = kv_stride2 - self.kv_stride4 = kv_stride4 - self.sf_stride0 = sf_stride0 - self.v_sf_stride0 = v_sf_stride0 - self.p_stride0 = p_stride0 - self.p_stride1 = p_stride1 - self.page_stats_stride0 = page_stats_stride0 - self.page_stats_stride1 = page_stats_stride1 - self.page_stats_stride2 = page_stats_stride2 - self.output_dtype = output_dtype - self.fp4_block = 16 - self.k_sf_per_token = self.k_head_dim // self.fp4_block - self.q_sf_per_token = self.q_head_dim // self.fp4_block - self.p_sf_per_page = page_size // self.fp4_block - self.non_residual_groups = self.k_sf_per_token - q_residual_dim // self.fp4_block - self.log2_e = math.log2(math.e) - self.p_global_scale = 448.0 * 6.0 - self.head_tile = 128 - self.v_tile = 128 - - @cute.jit - def __call__( - self, - output, - max_scores, - denom, - page_max, - page_sum, - p_fp4, - p_sf, - q_fp4, - q_sf, - kv_cache, - sf_cache, - v_sf, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - sm_scale: cutlass.Float32, - stream: cuda.CUstream, - ) -> None: - num_head_blocks = cute.ceil_div(self.num_heads, self.head_tile) - self._page_stats_pack_kernel( - page_max, - page_sum, - p_fp4, - p_sf, - q_fp4, - q_sf, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - sm_scale, - ).launch( - grid=(output.shape[0], num_head_blocks, self.max_pages), - block=(self.head_tile, 1, 1), - stream=stream, - ) - - self._reduce_stats_kernel( - max_scores, - denom, - page_max, - page_sum, - ).launch( - grid=(output.shape[0], num_head_blocks, 1), - block=(self.head_tile, 1, 1), - stream=stream, - ) - - self._prob_scale_kernel( - max_scores, - denom, - page_max, - p_sf, - paged_kv_indptr_decode, - kv_lens, - ).launch( - grid=(output.shape[0], num_head_blocks, self.max_pages), - block=(self.head_tile, 1, 1), - stream=stream, - ) - - self._pv_kernel( - output, - p_fp4, - p_sf, - kv_cache, - v_sf, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - ).launch( - grid=( - output.shape[0], - self.num_heads, - cute.ceil_div(self.kv_lora_rank, self.v_tile), - ), - block=(self.v_tile, 1, 1), - stream=stream, - ) - - @cute.jit - def _qk_score( - self, - q_fp4, - q_sf, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - compact_page, - token_idx, - q_row, - sm_scale, - ): - physical_page = src_page_ids[compact_page] - score = cutlass.Float32(0.0) - bytes_per_group: cutlass.Constexpr = self.fp4_block // 2 - q_row_base = q_row * self.q_fp4_stride0 - kv_token_base = physical_page * self.kv_stride0 + token_idx * self.kv_stride2 - k_sf_page_base = physical_page * self.sf_stride0 - - # Keep q_group as a runtime loop (44 iters): unrolling it together - # with the inner byte-pair loop and the outer 128-token loop blows - # up compile time and instruction footprint. - for q_group in cutlass.range(self.q_sf_per_token, unroll=1): - k_group = q_group - if q_group >= self.non_residual_groups: - k_group = self.non_residual_groups + (q_group - self.non_residual_groups) // 2 - - # Scales only change at fp4_block boundaries — hoist out of the - # inner byte-pair loop instead of reloading per element. - q_scale = _fp8_e4m3fn_to_f32( - q_sf[_swizzled_sf_offset(q_row, q_group, self.q_sf_per_token)] - ) - k_scale = _fp8_e4m3fn_to_f32( - sf_cache[ - k_sf_page_base - + _swizzled_sf_offset(token_idx, k_group, self.k_sf_per_token) - ] - ) - qk_scale = q_scale * k_scale - - # Each FP4 byte holds 2 nibbles; process them together so the - # single cvt.f16x2.e2m1x2 produces 2 useful f32 values. - for byte_idx in cutlass.range_constexpr(bytes_per_group): - q_packed_col = q_group * bytes_per_group + byte_idx - k_packed_col = k_group * bytes_per_group + byte_idx - q_lo, q_hi = _load_fp4_byte_pair( - q_fp4, - q_row_base + q_packed_col * self.q_fp4_stride1, - ) - k_lo, k_hi = _load_fp4_byte_pair( - kv_cache, - kv_token_base + k_packed_col * self.kv_stride4, - ) - score += (q_lo * k_lo + q_hi * k_hi) * qk_scale - - scale = sm_scale / (global_scale[0] * global_scale[0]) - return score * scale - - @cute.jit - def _page_stats_offset(self, gen_idx, page_rel, head_idx): - return ( - gen_idx * self.page_stats_stride0 - + page_rel * self.page_stats_stride1 - + head_idx * self.page_stats_stride2 - ) - - @cute.kernel - def _page_stats_pack_kernel( - self, - page_max, - page_sum, - p_fp4, - p_sf, - q_fp4, - q_sf, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - sm_scale: cutlass.Float32, - ) -> None: - tidx, _, _ = cute.arch.thread_idx() - gen_idx, head_block, page_rel = cute.arch.block_idx() - head_idx = head_block * self.head_tile + tidx - max_score = -cutlass.Float32.inf - sum_value = cutlass.Float32(0.0) - - if head_idx < self.num_heads: - q_row = gen_idx * self.num_heads + head_idx - kv_len = kv_lens[gen_idx] - page_start = page_rel * self.page_size - if page_start < kv_len: - page_table_start = paged_kv_indptr_decode[gen_idx] - compact_page = page_table_start + page_rel - scores = cute.make_fragment((self.page_size,), cutlass.Float32) - for token_offset in cutlass.range_constexpr(self.page_size): - token_abs = page_start + token_offset - score = -cutlass.Float32.inf - if token_abs < kv_len: - score = self._qk_score( - q_fp4, - q_sf, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - compact_page, - token_offset, - q_row, - sm_scale, - ) - if score > max_score: - max_score = score - scores[token_offset] = score - - p_row = compact_page * self.num_heads + head_idx - for token_group in cutlass.range_constexpr(self.p_sf_per_page): - local_max = cutlass.Float32(0.0) - probs = cute.make_fragment((self.fp4_block,), cutlass.Float32) - for idx in cutlass.range_constexpr(self.fp4_block): - token_offset = token_group * self.fp4_block + idx - token_abs = page_start + token_offset - prob = cutlass.Float32(0.0) - if token_abs < kv_len: - prob = cute.math.exp2( - (scores[token_offset] - max_score) * self.log2_e, - fastmath=True, - ) - sum_value += prob - probs[idx] = prob - if prob > local_max: - local_max = prob - - local_scale = cutlass.Float32(1.0) - stored_scale = cutlass.Float32(1.0) - if local_max > cutlass.Float32(0.0): - local_scale = local_max / cutlass.Float32(6.0) - stored_scale = local_scale * self.p_global_scale - if stored_scale > cutlass.Float32(448.0): - stored_scale = cutlass.Float32(448.0) - - p_sf[ - _swizzled_sf_offset( - p_row, - token_group, - self.p_sf_per_page, - ) - ] = _fp8_e4m3fn_positive_from_f32(stored_scale) - - for byte_idx in cutlass.range_constexpr(self.fp4_block // 2): - packed = _fp4_e2m1_quantize_packed( - probs[byte_idx * 2] / local_scale, - probs[byte_idx * 2 + 1] / local_scale, - ) - p_fp4[ - p_row * self.p_stride0 - + (token_group * (self.fp4_block // 2) + byte_idx) * self.p_stride1 - ] = packed - - stats_offset = self._page_stats_offset(gen_idx, page_rel, head_idx) - page_max[stats_offset] = max_score - page_sum[stats_offset] = sum_value - - @cute.kernel - def _reduce_stats_kernel(self, max_scores, denom, page_max, page_sum) -> None: - tidx, _, _ = cute.arch.thread_idx() - gen_idx, head_block, _ = cute.arch.block_idx() - head_idx = head_block * self.head_tile + tidx - if head_idx < self.num_heads: - max_score = -cutlass.Float32.inf - for page_rel in cutlass.range(self.max_pages, unroll=1): - stats_offset = self._page_stats_offset(gen_idx, page_rel, head_idx) - page_max_value = page_max[stats_offset] - if page_max_value > max_score: - max_score = page_max_value - - denom_value = cutlass.Float32(0.0) - for page_rel in cutlass.range(self.max_pages, unroll=1): - stats_offset = self._page_stats_offset(gen_idx, page_rel, head_idx) - page_sum_value = page_sum[stats_offset] - if page_sum_value > cutlass.Float32(0.0): - denom_value += page_sum_value * cute.math.exp2( - (page_max[stats_offset] - max_score) * self.log2_e, - fastmath=True, - ) - - q_row = gen_idx * self.num_heads + head_idx - max_scores[q_row] = max_score - denom[q_row] = denom_value - - @cute.kernel - def _prob_scale_kernel( - self, - max_scores, - denom, - page_max, - p_sf, - paged_kv_indptr_decode, - kv_lens, - ) -> None: - tidx, _, _ = cute.arch.thread_idx() - gen_idx, head_block, page_rel = cute.arch.block_idx() - head_idx = head_block * self.head_tile + tidx - if head_idx < self.num_heads: - kv_len = kv_lens[gen_idx] - page_start = page_rel * self.page_size - if page_start < kv_len: - stats_offset = self._page_stats_offset(gen_idx, page_rel, head_idx) - q_row = gen_idx * self.num_heads + head_idx - denom_value = denom[q_row] - factor = cutlass.Float32(0.0) - if denom_value > cutlass.Float32(0.0): - factor = ( - cute.math.exp2( - (page_max[stats_offset] - max_scores[q_row]) * self.log2_e, - fastmath=True, - ) - / denom_value - ) - - page_table_start = paged_kv_indptr_decode[gen_idx] - p_row = (page_table_start + page_rel) * self.num_heads + head_idx - for token_group in cutlass.range_constexpr(self.p_sf_per_page): - sf_offset = _swizzled_sf_offset( - p_row, - token_group, - self.p_sf_per_page, - ) - scaled = _fp8_e4m3fn_to_f32(p_sf[sf_offset]) * factor - p_sf[sf_offset] = _fp8_e4m3fn_positive_from_f32(scaled) - - @cute.kernel - def _pv_kernel( - self, - output, - p_fp4, - p_sf, - kv_cache, - v_sf, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - ) -> None: - tidx, _, _ = cute.arch.thread_idx() - gen_idx, head_idx, dim_block = cute.arch.block_idx() - v_dim = dim_block * self.v_tile + tidx - if v_dim < self.kv_lora_rank: - kv_len = kv_lens[gen_idx] - page_table_start = paged_kv_indptr_decode[gen_idx] - v_packed_col = v_dim // 2 # constant per thread - acc = cutlass.Float32(0.0) - bytes_per_group: cutlass.Constexpr = self.fp4_block // 2 - for page_rel in cutlass.range(self.max_pages, unroll=1): - page_start = page_rel * self.page_size - if page_start < kv_len: - compact_page = page_table_start + page_rel - physical_page = src_page_ids[compact_page] - p_row = compact_page * self.num_heads + head_idx - p_row_base = p_row * self.p_stride0 - v_page_base = ( - physical_page * self.kv_stride0 + v_packed_col * self.kv_stride4 - ) - v_sf_page_base = physical_page * self.v_sf_stride0 - # Process 16 tokens at a time — both scales are - # constant within each fp4_block, so hoist them. - for token_group in cutlass.range(self.p_sf_per_page, unroll=1): - group_start = token_group * self.fp4_block - if page_start + group_start < kv_len: - p_scale = _fp8_e4m3fn_to_f32( - p_sf[ - _swizzled_sf_offset( - p_row, - token_group, - self.p_sf_per_page, - ) - ] - ) - v_scale = _fp8_e4m3fn_to_f32( - v_sf[ - v_sf_page_base - + _swizzled_sf_offset( - v_dim, - token_group, - self.p_sf_per_page, - ) - ] - ) - pv_scale = p_scale * v_scale - # Each P byte holds 2 adjacent tokens; load - # once and use both nibbles. - for byte_idx in cutlass.range_constexpr(bytes_per_group): - token_a = group_start + byte_idx * 2 - token_b = token_a + 1 - token_abs_a = page_start + token_a - if token_abs_a < kv_len: - p_packed_col = token_group * bytes_per_group + byte_idx - p_lo, p_hi = _load_fp4_byte_pair( - p_fp4, - p_row_base + p_packed_col * self.p_stride1, - ) - v_a = _load_fp4_value( - kv_cache, - v_page_base + token_a * self.kv_stride2, - v_dim, - ) - acc += p_lo * v_a * pv_scale - if token_abs_a + 1 < kv_len: - v_b = _load_fp4_value( - kv_cache, - v_page_base + token_b * self.kv_stride2, - v_dim, - ) - acc += p_hi * v_b * pv_scale - - output[gen_idx, head_idx, v_dim] = ( - acc / (global_scale[0] * self.p_global_scale) - ).to(output.element_type) - - _COMPILE_CACHE: dict[tuple[int, ...], object] = {} - - def _storage_span(tensor: torch.Tensor) -> int: - if tensor.numel() == 0: - return 0 - return 1 + sum((size - 1) * stride for size, stride in zip(tensor.shape, tensor.stride())) - - def _flatten(tensor: torch.Tensor) -> torch.Tensor: - if tensor.is_contiguous(): - return tensor.reshape(-1) - return torch.as_strided( - tensor, - size=(_storage_span(tensor),), - stride=(1,), - storage_offset=tensor.storage_offset(), - ) - - def run_fp4_mla_attention_decode_cute( - *, - output: torch.Tensor, - max_scores: torch.Tensor, - denom: torch.Tensor, - page_max: torch.Tensor, - page_sum: torch.Tensor, - p_fp4: torch.Tensor, - p_sf: torch.Tensor, - q_fp4: torch.Tensor, - q_sf: torch.Tensor, - kv_cache: torch.Tensor, - sf_cache: torch.Tensor, - v_sf: torch.Tensor, - global_scale: torch.Tensor, - src_page_ids: torch.Tensor, - paged_kv_indptr_decode: torch.Tensor, - kv_lens: torch.Tensor, - sm_scale: float, - kv_lora_rank: int, - qk_rope_head_dim: int, - q_residual_dim: int, - page_size: int, - max_pages: int, - ) -> None: - """Run FP4 MLA decode using the page-parallel CuTe DSL backend.""" - - q_fp4_flat = _flatten(q_fp4.view(torch.uint8)) - q_sf_flat = _flatten(q_sf.view(torch.uint8)) - kv_cache_flat = _flatten(kv_cache.view(torch.uint8)) - sf_cache_flat = _flatten(sf_cache.view(torch.uint8)) - v_sf_flat = _flatten(v_sf.view(torch.uint8)) - p_fp4_flat = _flatten(p_fp4.view(torch.uint8)) - p_sf_flat = _flatten(p_sf.view(torch.uint8)) - max_scores_flat = _flatten(max_scores) - denom_flat = _flatten(denom) - page_max_flat = _flatten(page_max) - page_sum_flat = _flatten(page_sum) - - stream = cuda.CUstream(torch.cuda.current_stream(output.device).cuda_stream) - cute_args = ( - _to_cute(output), - _to_cute(max_scores_flat), - _to_cute(denom_flat), - _to_cute(page_max_flat), - _to_cute(page_sum_flat), - _to_cute(p_fp4_flat), - _to_cute(p_sf_flat), - _to_cute(q_fp4_flat), - _to_cute(q_sf_flat), - _to_cute(kv_cache_flat), - _to_cute(sf_cache_flat), - _to_cute(v_sf_flat), - _to_cute(global_scale), - _to_cute(src_page_ids), - _to_cute(paged_kv_indptr_decode), - _to_cute(kv_lens), - cutlass.Float32(sm_scale), - stream, - ) - compile_key = ( - output.device.index or 0, - output.dtype is torch.bfloat16, - output.shape[0], - output.shape[1], - kv_lora_rank, - qk_rope_head_dim, - q_residual_dim, - page_size, - max_pages, - q_fp4.stride(0), - q_fp4.stride(1), - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - sf_cache.stride(0), - v_sf.stride(0), - p_fp4.stride(0), - p_fp4.stride(1), - page_max.stride(0), - page_max.stride(1), - page_max.stride(2), - ) - if compile_key not in _COMPILE_CACHE: - if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "CuTe FP4 MLA decode must be compiled before CUDA graph capture." - ) - logger.info( - "Compiling CuTe FP4 MLA decode kernel for " - f"num_gen={output.shape[0]}, num_heads={output.shape[1]}, " - f"kv_lora_rank={kv_lora_rank}, rope_dim={qk_rope_head_dim}, " - f"max_pages={max_pages}" - ) - kernel = _Fp4MlaDecodeCuteKernel( - num_heads=output.shape[1], - kv_lora_rank=kv_lora_rank, - qk_rope_head_dim=qk_rope_head_dim, - q_residual_dim=q_residual_dim, - page_size=page_size, - max_pages=max_pages, - q_fp4_stride0=q_fp4.stride(0), - q_fp4_stride1=q_fp4.stride(1), - kv_stride0=kv_cache.stride(0), - kv_stride2=kv_cache.stride(2), - kv_stride4=kv_cache.stride(4), - sf_stride0=sf_cache.stride(0), - v_sf_stride0=v_sf.stride(0), - p_stride0=p_fp4.stride(0), - p_stride1=p_fp4.stride(1), - page_stats_stride0=page_max.stride(0), - page_stats_stride1=page_max.stride(1), - page_stats_stride2=page_max.stride(2), - output_dtype=output.dtype, - ) - _COMPILE_CACHE[compile_key] = cute.compile(kernel, *cute_args) - - _COMPILE_CACHE[compile_key](*cute_args) - -else: - - def run_fp4_mla_attention_decode_cute(**_: object) -> None: - raise RuntimeError("CuTe DSL is not available in this environment.") diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py index 337f017fe2c1..b821db911635 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py @@ -10,12 +10,14 @@ this internal op. """ +import os from typing import Optional import torch import triton import triton.language as tl + FP4_BLOCK_SIZE = 16 FP4_MLA_P_GLOBAL_SCALE = 448.0 * 6.0 @@ -24,6 +26,13 @@ def _ceil_div(lhs: int, rhs: int) -> int: return (lhs + rhs - 1) // rhs +def _env_int(name: str) -> Optional[int]: + value = os.environ.get(name) + if value is None or value == "": + return None + return int(value) + + def _swizzled_scale_size(rows: int, logical_cols: int) -> int: scale_cols = _ceil_div(logical_cols, FP4_BLOCK_SIZE) padded_cols = _ceil_div(scale_cols, 4) * 4 @@ -33,26 +42,12 @@ def _swizzled_scale_size(rows: int, logical_cols: int) -> int: def _get_kv_cache_strides(kv_cache: torch.Tensor) -> tuple[int, int, int, int, int, int]: if kv_cache.dim() == 3: num_pages, page_size, packed_dim = kv_cache.shape - return ( - num_pages, - page_size, - packed_dim, - kv_cache.stride(0), - kv_cache.stride(1), - kv_cache.stride(2), - ) + return num_pages, page_size, packed_dim, kv_cache.stride(0), kv_cache.stride(1), kv_cache.stride(2) if kv_cache.dim() >= 5: num_pages = kv_cache.shape[0] page_size = kv_cache.shape[2] packed_dim = kv_cache.shape[4] - return ( - num_pages, - page_size, - packed_dim, - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - ) + return num_pages, page_size, packed_dim, kv_cache.stride(0), kv_cache.stride(2), kv_cache.stride(4) raise ValueError( "kv_cache must be shaped (num_pages, page_size, packed_dim) or (num_pages, ..., page_size, ..., packed_dim)." ) @@ -69,8 +64,7 @@ def _workspace_tensor( if workspace is None: if torch.cuda.is_current_stream_capturing(): raise ValueError( - f"Cannot allocate {name} while capturing a CUDA graph. " - "Pass a preallocated workspace tensor." + f"Cannot allocate {name} while capturing a CUDA graph. Pass a preallocated workspace tensor." ) return torch.empty(shape, dtype=dtype, device=device) @@ -100,18 +94,12 @@ def _fp4_mla_swizzled_sf_offset(row_idx, col_idx, SF_PER_TOKEN: tl.constexpr): row_in_group1 = (row_idx % 128) // 32 row_group = row_idx // 128 return ( - col_in_group - + col_group * (4 * 128) - + row_in_group0 * 16 - + row_in_group1 * 4 - + row_group * (128 * padded_cols) + col_in_group + col_group * (4 * 128) + row_in_group0 * 16 + row_in_group1 * 4 + row_group * (128 * padded_cols) ) @triton.jit -def _fp4_mla_swizzled_sf_offset_row_block( - row_group, row_offsets, col_idx, SF_PER_TOKEN: tl.constexpr -): +def _fp4_mla_swizzled_sf_offset_row_block(row_group, row_offsets, col_idx, SF_PER_TOKEN: tl.constexpr): padded_cols = ((SF_PER_TOKEN + 3) // 4) * 4 col_part = (col_idx % 4) + (col_idx // 4) * (4 * 128) row_part = (row_offsets % 32) * 16 + ((row_offsets % 128) // 32) * 4 @@ -202,6 +190,198 @@ def _fp4_pack_high_nibbles(even_packed, odd_packed): ) +@triton.jit +def _fp4_pack_nibbles(even_packed, odd_packed): + return tl.inline_asm_elementwise( + """ + { + .reg .b32 lo; + .reg .b32 hi; + and.b32 lo, $2, 15; + and.b32 hi, $3, 15; + shl.b32 hi, hi, 4; + or.b32 $0, lo, hi; + + shr.u32 lo, $2, 4; + and.b32 lo, lo, 15; + and.b32 hi, $3, 240; + or.b32 $1, lo, hi; + } + """, + constraints="=r,=r,r,r", + args=[even_packed, odd_packed], + dtype=(tl.uint8, tl.uint8), + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fp4_mla_attention_v_repack_kernel( + v_packed_ptr, + kv_cache_ptr, + num_pages, + kv_s0, + kv_s2, + kv_s4, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BLOCK_V: tl.constexpr, + occupancy: tl.constexpr = 1, +): + page_idx = tl.program_id(0) + dim_block = tl.program_id(1) + + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_view = tl.ext.make_view( + base=kv_cache_ptr, + shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], + tile_dim_map=[0, 1, 2], + ) + v_tile = tl.ext.load_view_tko( + v_view, + [ + page_idx.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ], + ) + v_tile = v_tile.to(tl.uint8, bitcast=True) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + + out_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + row_base = (page_idx * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + out_desc.store([row_base.to(tl.int32), 0], v_vals) + + +@triton.jit +def _fp4_mla_attention_v_repack_pages_kernel( + v_packed_ptr, + kv_cache_ptr, + page_ids_ptr, + num_pages, + kv_s0, + kv_s2, + kv_s4, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BLOCK_V: tl.constexpr, + occupancy: tl.constexpr = 1, +): + page_list_idx = tl.program_id(0) + dim_block = tl.program_id(1) + page_idx = tl.load(page_ids_ptr + page_list_idx).to(tl.int64) + + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_view = tl.ext.make_view( + base=kv_cache_ptr, + shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], + tile_dim_map=[0, 1, 2], + ) + v_tile = tl.ext.load_view_tko( + v_view, + [ + page_idx.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ], + ) + v_tile = v_tile.to(tl.uint8, bitcast=True) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + + out_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + row_base = (page_idx * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + out_desc.store([row_base.to(tl.int32), 0], v_vals) + + +def fp4_mla_repack_v_cache( + v_packed: torch.Tensor, + kv_cache: torch.Tensor, + page_ids: Optional[torch.Tensor] = None, + *, + v_head_dim: int, + page_size: int, + block_v: int = 128, + kernel_occupancy: int = 8, + kernel_num_stages: int = 1, +) -> None: + """Populate the V-packed auxiliary cache consumed by the prepacked PV kernel.""" + if v_head_dim % block_v != 0: + raise ValueError(f"v_head_dim={v_head_dim} must be divisible by block_v={block_v}.") + if kv_cache.ndim < 5: + raise ValueError(f"kv_cache must expose the paged FP4 layout, got shape={tuple(kv_cache.shape)}.") + num_pages = kv_cache.shape[0] + num_dim_blocks = triton.cdiv(v_head_dim, block_v) + launch_meta = { + "occupancy": int(kernel_occupancy), + "num_stages": int(kernel_num_stages), + } + if page_ids is None: + if num_pages == 0: + return + _fp4_mla_attention_v_repack_kernel[(num_pages, num_dim_blocks)]( + v_packed, + kv_cache, + num_pages, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + BLOCK_V=block_v, + **launch_meta, + ) + return + + if page_ids.numel() == 0: + return + _fp4_mla_attention_v_repack_pages_kernel[(page_ids.numel(), num_dim_blocks)]( + v_packed, + kv_cache, + page_ids, + num_pages, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + BLOCK_V=block_v, + **launch_meta, + ) + + @triton.jit def _fp4_mla_qk_scores_tile( q_fp4_ptr, @@ -246,12 +426,8 @@ def _fp4_mla_qk_scores_tile( else: valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) safe_compact_page = tl.where(valid_compact_page, compact_page, 0) - physical_page = tl.load( - src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 - ).to(tl.int64) - valid_physical_page = ( - valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) - ) + physical_page = tl.load(src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1).to(tl.int64) + valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) safe_physical_page = tl.where(valid_physical_page, physical_page, 0) q_rows = q_row_base + head_offsets if ASSUME_FULL_HEADS: @@ -311,6 +487,13 @@ def _fp4_mla_qk_scores_tile( tile_shape=[1, 8, 2, 256], tile_dim_map=[0, 1, 2, 3], ) + q_sf_tail_view = tl.ext.make_view( + base=q_sf_ptr, + shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) k_sf_full_view = tl.ext.make_view( base=sf_cache_ptr, shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], @@ -333,9 +516,7 @@ def _fp4_mla_qk_scores_tile( full_q_scales = tl.ext.load_view_tko(q_sf_full_view, [q_row_group.to(tl.int32), 0, 0, 0]) full_q_scales = full_q_scales.reshape([1, 8, 32, 4, 4]).trans(0, 3, 2, 1, 4) full_q_scales = full_q_scales.reshape([BLOCK_H, 32]) - full_k_scales = tl.ext.load_view_tko( - k_sf_full_view, [safe_physical_page.to(tl.int32), 0, 0, 0, 0] - ) + full_k_scales = tl.ext.load_view_tko(k_sf_full_view, [safe_physical_page.to(tl.int32), 0, 0, 0, 0]) full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) scores = tl.dot_scaled( @@ -352,9 +533,7 @@ def _fp4_mla_qk_scores_tile( tail_k_vals = k_tail_desc.load([safe_physical_page.to(tl.int32), 0, 256]) tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) - tail_k_scales = tl.ext.load_view_tko( - k_sf_tail_view, [safe_physical_page.to(tl.int32), 0, 8, 0, 0] - ) + tail_k_scales = tl.ext.load_view_tko(k_sf_tail_view, [safe_physical_page.to(tl.int32), 0, 8, 0, 0]) tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) q_tail_vals = q_tail_desc.load([(q_row_base + head_start).to(tl.int32), 256]) @@ -365,9 +544,7 @@ def _fp4_mla_qk_scores_tile( q_odd_vals = q_odd_vals.reshape([BLOCK_H, 32]) q_tail_sf_cols = 32 + tl.arange(0, 8) - q_tail_sf_offsets = _fp4_mla_swizzled_sf_offset( - q_rows[:, None], q_tail_sf_cols[None, :], Q_SF_PER_TOKEN - ) + q_tail_sf_offsets = _fp4_mla_swizzled_sf_offset(q_rows[:, None], q_tail_sf_cols[None, :], Q_SF_PER_TOKEN) q_tail_scales = tl.load(q_sf_ptr + q_tail_sf_offsets) q_tail_scales = q_tail_scales.reshape([BLOCK_H, 4, 2]) q_even_scales, q_odd_scales = tl.split(q_tail_scales) @@ -444,11 +621,7 @@ def _fp4_mla_qk_scores_tile( mask_k = q_elem_offsets < Q_HEAD_D safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) - if ( - USE_TMA_DATA_LOAD - and FULL_BLOCK_END > 0 - and q_start + BLOCK_K <= non_residual_groups * FP4_BLOCK - ): + if USE_TMA_DATA_LOAD and FULL_BLOCK_END > 0 and q_start + BLOCK_K <= non_residual_groups * FP4_BLOCK: q_vals = q_desc.load([(q_row_base + head_start).to(tl.int32), q_start // 2]) k_vals = k_desc.load([safe_physical_page.to(tl.int32), 0, q_start // 2]) k_vals = tl.reshape(k_vals, (BLOCK_T, BLOCK_K // 2)) @@ -456,9 +629,7 @@ def _fp4_mla_qk_scores_tile( k_vals = tl.where(valid_physical_page, k_vals, 0) else: q_vals = tl.load( - q_fp4_ptr - + safe_q_rows[:, None] * q_fp4_s0 - + safe_packed_q_cols[None, :] * q_fp4_s1, + q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + safe_packed_q_cols[None, :] * q_fp4_s1, mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], other=0, ) @@ -467,9 +638,7 @@ def _fp4_mla_qk_scores_tile( + safe_physical_page * kv_s0 + token_offsets[:, None].to(tl.int64) * kv_s2 + safe_packed_k_cols[None, :] * kv_s4, - mask=mask_k[None, :] - if ASSUME_VALID_PAGES - else valid_physical_page & mask_k[None, :], + mask=mask_k[None, :] if ASSUME_VALID_PAGES else valid_physical_page & mask_k[None, :], other=0, ) @@ -482,12 +651,8 @@ def _fp4_mla_qk_scores_tile( mask_sf = q_sf_cols < Q_SF_PER_TOKEN safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) - q_sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN - ) - k_sf_offsets = _fp4_mla_swizzled_sf_offset( - token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN - ) + q_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN) + k_sf_offsets = _fp4_mla_swizzled_sf_offset(token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN) q_scales = tl.load(q_sf_ptr + q_sf_offsets) k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) scores = tl.dot_scaled( @@ -536,20 +701,14 @@ def _fp4_mla_qk_scores_tile( other=0, ) k_sf_cols = non_residual_groups + residual_scale_offsets - k_sf_offsets = _fp4_mla_swizzled_sf_offset( - token_offsets[:, None], k_sf_cols[None, :], K_SF_PER_TOKEN - ) + k_sf_offsets = _fp4_mla_swizzled_sf_offset(token_offsets[:, None], k_sf_cols[None, :], K_SF_PER_TOKEN) k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) q_tail_cols = q_start // 2 + tl.arange(0, 64) if USE_TMA_DATA_LOAD and ASSUME_FULL_HEADS: - q_tail_vals = q_tail_desc.load( - [(q_row_base + head_start).to(tl.int32), q_start // 2] - ) + q_tail_vals = q_tail_desc.load([(q_row_base + head_start).to(tl.int32), q_start // 2]) elif ASSUME_FULL_HEADS: - q_tail_vals = tl.load( - q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1 - ) + q_tail_vals = tl.load(q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1) else: q_tail_vals = tl.load( q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1, @@ -608,9 +767,7 @@ def _fp4_mla_qk_scores_tile( safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) q_vals = tl.load( - q_fp4_ptr - + safe_q_rows[:, None] * q_fp4_s0 - + safe_packed_q_cols[None, :] * q_fp4_s1, + q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + safe_packed_q_cols[None, :] * q_fp4_s1, mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], other=0, ) @@ -619,9 +776,7 @@ def _fp4_mla_qk_scores_tile( + safe_physical_page * kv_s0 + token_offsets[:, None].to(tl.int64) * kv_s2 + safe_packed_k_cols[None, :] * kv_s4, - mask=mask_k[None, :] - if ASSUME_VALID_PAGES - else valid_physical_page & mask_k[None, :], + mask=mask_k[None, :] if ASSUME_VALID_PAGES else valid_physical_page & mask_k[None, :], other=0, ) @@ -634,12 +789,8 @@ def _fp4_mla_qk_scores_tile( mask_sf = q_sf_cols < Q_SF_PER_TOKEN safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) - q_sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN - ) - k_sf_offsets = _fp4_mla_swizzled_sf_offset( - token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN - ) + q_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN) + k_sf_offsets = _fp4_mla_swizzled_sf_offset(token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN) q_scales = tl.load(q_sf_ptr + q_sf_offsets) k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) scores = tl.dot_scaled( @@ -763,9 +914,7 @@ def _fp4_mla_attention_stats_kernel( scores = tl.where(mask_h[:, None], scores * qk_scale, -float("inf")) else: valid_t = page_start + offs_t < kv_len - scores = tl.where( - mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf") - ) + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) page_max = tl.max(scores, axis=1) new_max = tl.maximum(max_score, page_max) denom = denom * tl.math.exp2((max_score - new_max) * 1.4426950408889634) + tl.sum( @@ -920,9 +1069,7 @@ def _fp4_mla_attention_page_stats_kernel( grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) amax = tl.max(grouped_probs, axis=2) inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) - stored_scale = tl.where( - amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0 - ) + stored_scale = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) even_probs, odd_probs = tl.split(pairs) @@ -931,25 +1078,17 @@ def _fp4_mla_attention_page_stats_kernel( if ASSUME_VALID_PAGES: safe_compact_page = page_table_start + page_rel else: - valid_compact_page = (page_table_start + page_rel >= 0) & ( - page_table_start + page_rel < page_ids_len - ) + valid_compact_page = (page_table_start + page_rel >= 0) & (page_table_start + page_rel < page_ids_len) safe_compact_page = tl.where(valid_compact_page, page_table_start + page_rel, 0) p_rows = safe_compact_page * NUM_HEADS + offs_h - safe_p_rows = ( - p_rows - if ASSUME_FULL_HEADS - else tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) - ) + safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) scale_cols = tl.arange(0, SF_PER_PAGE) if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( safe_compact_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) else: - sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE - ) + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) if ASSUME_FULL_HEADS: if ASSUME_VALID_PAGES: tl.store(p_sf_ptr + sf_offsets, stored_scale) @@ -959,9 +1098,7 @@ def _fp4_mla_attention_page_stats_kernel( tl.store( p_sf_ptr + sf_offsets, stored_scale, - mask=mask_h[:, None] - if ASSUME_VALID_PAGES - else valid_compact_page & mask_h[:, None], + mask=mask_h[:, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None], ) byte_offsets = tl.arange(0, FP4_BLOCK // 2) @@ -974,16 +1111,12 @@ def _fp4_mla_attention_page_stats_kernel( elif ASSUME_FULL_HEADS: if ASSUME_VALID_PAGES: tl.store( - p_fp4_ptr - + safe_p_rows[:, None, None] * p_s0 - + byte_cols[None, :, :] * p_s1, + p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, packed, ) else: tl.store( - p_fp4_ptr - + safe_p_rows[:, None, None] * p_s0 - + byte_cols[None, :, :] * p_s1, + p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, packed, mask=valid_compact_page, ) @@ -991,9 +1124,7 @@ def _fp4_mla_attention_page_stats_kernel( tl.store( p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, packed, - mask=mask_h[:, None, None] - if ASSUME_VALID_PAGES - else valid_compact_page & mask_h[:, None, None], + mask=mask_h[:, None, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None, None], ) if ASSUME_FULL_HEADS: @@ -1004,6 +1135,246 @@ def _fp4_mla_attention_page_stats_kernel( tl.store(page_sum_ptr + out_offsets, page_sum, mask=mask_h) +@triton.jit +def _fp4_mla_attention_page_stats_grouped_kernel( + page_max_ptr, + page_sum_ptr, + p_fp4_ptr, + p_sf_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len: tl.constexpr, + num_pages: tl.constexpr, + q_fp4_s0: tl.constexpr, + q_fp4_s1: tl.constexpr, + kv_s0: tl.constexpr, + kv_s2: tl.constexpr, + kv_s4: tl.constexpr, + sf_s0: tl.constexpr, + page_stats_s0: tl.constexpr, + page_stats_s1: tl.constexpr, + p_s0: tl.constexpr, + p_s1: tl.constexpr, + p_num_rows: tl.constexpr, + q_num_rows: tl.constexpr, + sm_scale: tl.constexpr, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + PACK_PROBS: tl.constexpr, + GROUP_REDUCE_STATS: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + GROUP_PAGES: tl.constexpr = 2, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_group = tl.program_id(2) + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_t = tl.arange(0, BLOCK_T) + scale_cols = tl.arange(0, SF_PER_PAGE) + q_row_base = gen_idx * NUM_HEADS + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + + tl.assume(q_fp4_s0 % 8 == 0) + tl.assume(q_fp4_s1 == 1) + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + + q_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 256], + ) + k_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 256], + ) + q_tail_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 32], + ) + k_tail_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 32], + ) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + q_sf_full_view = tl.ext.make_view( + base=q_sf_ptr, + shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + q_sf_tail_view = tl.ext.make_view( + base=q_sf_ptr, + shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + k_sf_full_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + k_sf_tail_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 1, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + + q_row_start = (q_row_base + head_block * BLOCK_H).to(tl.int32) + q_row_group = q_row_base // 128 + full_q_vals = q_desc.load([q_row_start, 0]) + full_q_scales = tl.ext.load_view_tko(q_sf_full_view, [q_row_group.to(tl.int32), 0, 0, 0]) + full_q_scales = full_q_scales.reshape([1, 8, 32, 4, 4]).trans(0, 3, 2, 1, 4) + full_q_scales = full_q_scales.reshape([BLOCK_H, 32]) + q0_vals = q_tail_desc.load([q_row_start, 256]) + q1_vals = q_tail_desc.load([q_row_start, 288]) + q0_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group.to(tl.int32), 8, 0, 0]) + q0_scales = q0_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) + q0_scales = q0_scales.reshape([BLOCK_H, 4]) + q1_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group.to(tl.int32), 9, 0, 0]) + q1_scales = q1_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) + q1_scales = q1_scales.reshape([BLOCK_H, 4]) + + group_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + group_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) + for page_group_off in tl.range(0, GROUP_PAGES): + page_rel = page_group * GROUP_PAGES + page_group_off + compact_page = page_table_start + page_rel + physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) + + scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) + full_k_vals = k_desc.load([physical_page.to(tl.int32), 0, 0]) + full_k_vals = tl.reshape(full_k_vals, (BLOCK_T, 256)) + full_k_scales = tl.ext.load_view_tko(k_sf_full_view, [physical_page.to(tl.int32), 0, 0, 0, 0]) + full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) + scores = tl.dot_scaled( + full_q_vals, + full_q_scales, + "e2m1", + full_k_vals.T, + full_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + tail_k_vals = k_tail_desc.load([physical_page.to(tl.int32), 0, 256]) + tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) + tail_k_scales = tl.ext.load_view_tko(k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0]) + tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) + scores = tl.dot_scaled( + q0_vals, + q0_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q1_vals, + q1_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + scores = scores * qk_scale + page_max = tl.max(scores, axis=1) + exp_scores = tl.math.exp2((scores - page_max[:, None]) * 1.4426950408889634) + page_sum = tl.sum(exp_scores, axis=1) + if GROUP_REDUCE_STATS: + next_group_max = tl.maximum(group_max, page_max) + group_sum = group_sum * tl.math.exp2((group_max - next_group_max) * 1.4426950408889634) + page_sum * tl.math.exp2( + (page_max - next_group_max) * 1.4426950408889634 + ) + group_max = next_group_max + + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + out_offsets = gen_idx * page_stats_s0 + page_rel * page_stats_s1 + offs_h + tl.store(page_max_ptr + out_offsets, page_max) + if not GROUP_REDUCE_STATS: + tl.store(page_sum_ptr + out_offsets, page_sum) + + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + compact_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + tl.store(p_sf_ptr + sf_offsets, stored_scale) + p_desc.store( + [(compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), + ) + if GROUP_REDUCE_STATS: + group_max_offsets = gen_idx * page_stats_s0 + (page_group * 2) * page_stats_s1 + offs_h + group_sum_offsets = group_max_offsets + page_stats_s1 + tl.store(page_sum_ptr + group_max_offsets, group_max) + tl.store(page_sum_ptr + group_sum_offsets, group_sum) + + @triton.jit def _fp4_mla_attention_reduce_stats_kernel( max_ptr, @@ -1016,6 +1387,8 @@ def _fp4_mla_attention_reduce_stats_kernel( NUM_HEADS: tl.constexpr, MAX_PAGES: tl.constexpr, BLOCK_H: tl.constexpr, + GROUP_REDUCE_STATS: tl.constexpr = False, + GROUP_PAGES: tl.constexpr = 1, occupancy: tl.constexpr = 1, ): gen_idx = tl.program_id(0) @@ -1025,31 +1398,50 @@ def _fp4_mla_attention_reduce_stats_kernel( safe_offs_h = tl.where(mask_h, offs_h, 0) max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) - for page_rel in tl.range(0, MAX_PAGES): - page_max = tl.load( - page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - max_score = tl.maximum(max_score, page_max) + if GROUP_REDUCE_STATS: + for group_rel in tl.range(0, MAX_PAGES // GROUP_PAGES): + group_max = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.maximum(max_score, group_max) + else: + for page_rel in tl.range(0, MAX_PAGES): + page_max = tl.load( + page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.maximum(max_score, page_max) denom = tl.zeros((BLOCK_H,), dtype=tl.float32) - for page_rel in tl.range(0, MAX_PAGES): - page_max = tl.load( - page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - page_sum = tl.load( - page_sum_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=0.0, - ) - denom += tl.where( - page_sum > 0.0, - page_sum * tl.math.exp2((page_max - max_score) * 1.4426950408889634), - 0.0, - ) + if GROUP_REDUCE_STATS: + for group_rel in tl.range(0, MAX_PAGES // GROUP_PAGES): + group_max = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + group_sum = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2 + 1) * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=0.0, + ) + denom += tl.where(group_sum > 0.0, group_sum * tl.math.exp2((group_max - max_score) * 1.4426950408889634), 0.0) + else: + for page_rel in tl.range(0, MAX_PAGES): + page_max = tl.load( + page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + page_sum = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=0.0, + ) + denom += tl.where(page_sum > 0.0, page_sum * tl.math.exp2((page_max - max_score) * 1.4426950408889634), 0.0) tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) @@ -1108,23 +1500,17 @@ def _fp4_mla_attention_prob_scale_kernel( ) max_score = tl.load(max_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) denom = tl.load(denom_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - factor = tl.where( - denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0 - ) + factor = tl.where(denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0) p_rows = compact_page * NUM_HEADS + offs_h - safe_p_rows = ( - p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, compact_page * NUM_HEADS) - ) + safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, compact_page * NUM_HEADS) scale_cols = tl.arange(0, SF_PER_PAGE) if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( compact_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) else: - sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE - ) + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) @@ -1556,9 +1942,7 @@ def _fp4_mla_attention_pv_kernel( dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE ) else: - v_sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE - ) + v_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE) if USE_TMA_P_LOAD: tl.assume(p_s0 % 8 == 0) tl.assume(p_s1 == 1) @@ -1633,7 +2017,7 @@ def _fp4_mla_attention_pv_kernel( page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) global_scale = tl.load(global_scale_ptr) out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) - acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + acc = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): compact_page = page_table_start + page_rel physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) @@ -1659,8 +2043,7 @@ def _fp4_mla_attention_pv_kernel( v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) even_packed, odd_packed = tl.split(v_pairs) - low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) - high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) v_vals = tl.reshape( tl.join(low_vals, high_vals).permute(0, 2, 1), (BLOCK_V, PAGE_SIZE // 2), @@ -1677,28 +2060,25 @@ def _fp4_mla_attention_pv_kernel( ) v_scales = v_scales.reshape([1, 1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) - acc = tl.dot_scaled( - p_vals, - p_scales, - "e2m1", - v_vals.T, + acc = tl.ext.dot_scaled( + v_vals, v_scales, "e2m1", + p_vals.T, + p_scales, + "e2m1", acc=acc, fast_math=True, rhs_k_pack=True, ) - out_vals = acc * out_scale + out_vals = acc.T * out_scale if out_ptr.dtype.element_ty == tl.bfloat16: out_vals = out_vals.to(tl.bfloat16) elif out_ptr.dtype.element_ty == tl.float16: out_vals = out_vals.to(tl.float16) out_desc.store( - [ - (gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), - (dim_block * BLOCK_V).to(tl.int32), - ], + [(gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], out_vals, ) return @@ -1722,35 +2102,23 @@ def _fp4_mla_attention_pv_kernel( else: valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) safe_compact_page = tl.where(valid_compact_page, compact_page, 0) - physical_page = tl.load( - src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 - ).to(tl.int64) - valid_physical_page = ( - valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + physical_page = tl.load(src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1).to( + tl.int64 ) + valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) safe_physical_page = tl.where(valid_physical_page, physical_page, 0) p_rows = safe_compact_page * NUM_HEADS + offs_h - safe_p_rows = ( - p_rows - if ASSUME_FULL_HEADS - else tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) - ) + safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) if USE_TMA_P_LOAD: - p_vals = p_desc.load( - [(safe_compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0] - ) + p_vals = p_desc.load([(safe_compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) else: p_vals = tl.load( p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, - mask=mask_h[:, None] - if ASSUME_VALID_PAGES - else valid_compact_page & mask_h[:, None], + mask=mask_h[:, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None], other=0, ) - p_sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE - ) + p_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) p_scales = tl.load(p_sf_ptr + p_sf_offsets) if ASSUME_FULL_PAGES: @@ -1759,7 +2127,10 @@ def _fp4_mla_attention_pv_kernel( else: valid_even_t = page_start + even_t < kv_len valid_odd_t = page_start + odd_t < kv_len - if USE_TMA_V_LOAD: + if USE_PREPACKED_V: + v_row = (safe_physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + elif USE_TMA_V_LOAD: v_tile = v_desc.load( [ safe_physical_page.to(tl.int32), @@ -1775,8 +2146,7 @@ def _fp4_mla_attention_pv_kernel( if not ASSUME_FULL_PAGES: even_packed = tl.where(valid_even_t[None, :], even_packed, 0) odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) - low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) - high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) v_vals = tl.reshape( tl.join(low_vals, high_vals).permute(0, 2, 1), (BLOCK_V, PAGE_SIZE // 2), @@ -1829,10 +2199,376 @@ def _fp4_mla_attention_pv_kernel( elif out_ptr.dtype.element_ty == tl.float16: out_vals = out_vals.to(tl.float16) out_desc.store( + [(gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], + out_vals, + ) + else: + tl.store( + out_ptr + gen_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + ) + else: + tl.store( + out_ptr + gen_idx * out_s0 + safe_offs_h[:, None] * out_s1 + safe_offs_v[None, :] * out_s2, + acc * out_scale, + mask=mask_h[:, None] & mask_v[None, :], + ) + + +@triton.jit +def _fp4_mla_attention_pv_prepacked_v_kernel( + out_ptr, + p_fp4_ptr, + p_sf_ptr, + kv_cache_ptr, + v_sf_ptr, + v_packed_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + out_s0, + out_s1, + out_s2, + out_num_rows, + p_s0, + p_s1, + p_num_rows, + kv_s0, + kv_s2, + kv_s4, + vsf_s0, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + USE_TMA_P_LOAD: tl.constexpr, + USE_TMA_V_LOAD: tl.constexpr, + USE_PREPACKED_V: tl.constexpr, + PV_M_PACKED_V: tl.constexpr, + PV_LOOP_STAGES: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_FULL_V: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + dim_block = tl.program_id(2) + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + if ASSUME_FULL_V: + mask_v = tl.full([BLOCK_V], True, dtype=tl.int1) + safe_offs_v = offs_v + else: + mask_v = offs_v < V_HEAD_D + safe_offs_v = tl.where(mask_v, offs_v, 0) + packed_t = tl.arange(0, PAGE_SIZE // 2) + scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) + even_t = packed_t * 2 + odd_t = even_t + 1 + v_packed_offsets = safe_offs_v // 2 + v_use_high_nibble = (safe_offs_v & 1) != 0 + if ASSUME_FULL_V and BLOCK_V == 128: + v_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + v_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE) + if USE_TMA_P_LOAD: + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + if USE_TMA_V_LOAD and ASSUME_FULL_HEADS and ASSUME_FULL_V: + tl.assume(out_s1 % 8 == 0) + tl.assume(out_s2 == 1) + out_desc = tl.make_tensor_descriptor( + out_ptr, + shape=[out_num_rows, V_HEAD_D], + strides=[out_s1, out_s2], + block_shape=[BLOCK_H, BLOCK_V], + ) + if USE_TMA_V_LOAD: + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, PAGE_SIZE, BLOCK_V // 2], + ) + if USE_PREPACKED_V: + v_packed_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + if ( + USE_TMA_P_LOAD + and USE_TMA_V_LOAD + and ASSUME_FULL_HEADS + and ASSUME_FULL_PAGES + and ASSUME_FULL_V + and ASSUME_VALID_PAGES + and NUM_HEADS == 128 + and V_HEAD_D == 512 + and PAGE_SIZE == 128 + and BLOCK_H == 128 + and BLOCK_V == 128 + and SF_PER_PAGE == 8 + ): + p_view = tl.ext.make_view( + base=p_fp4_ptr, + shapes=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + tile_shape=[BLOCK_H, PAGE_SIZE // 2], + tile_dim_map=[0, 1], + ) + p_sf_view = tl.ext.make_view( + base=p_sf_ptr, + shapes=[p_num_rows // 128, SF_PER_PAGE // 4, 2, 256], + strides=[128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + v_sf_view = tl.ext.make_view( + base=v_sf_ptr, + shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], + strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + v_view = tl.ext.make_view( + base=kv_cache_ptr, + shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], + tile_dim_map=[0, 1, 2], + ) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): + compact_page = page_table_start + page_rel + physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) + + p_vals = tl.ext.load_view_tko( + p_view, + [(compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + ) + p_vals = p_vals.to(tl.uint8, bitcast=True) + p_scales = tl.ext.load_view_tko(p_sf_view, [compact_page.to(tl.int32), 0, 0, 0]) + p_scales = p_scales.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) + p_scales = p_scales.reshape([BLOCK_H, SF_PER_PAGE]) + + if USE_PREPACKED_V: + v_row = (physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + else: + v_tile = tl.ext.load_view_tko( + v_view, + [ + physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ], + ) + v_tile = v_tile.to(tl.uint8, bitcast=True) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + if not PV_M_PACKED_V: + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + v_scales = tl.ext.load_view_tko( + v_sf_view, [ - (gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), - (dim_block * BLOCK_V).to(tl.int32), + physical_page.to(tl.int32), + dim_block, + 0, + 0, + 0, ], + ) + v_scales = v_scales.reshape([1, 1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) + if PV_M_PACKED_V and not USE_PREPACKED_V: + acc = tl.ext.dot_scaled( + v_tile.T, + v_scales, + "e2m1", + p_vals.T, + p_scales, + "e2m1", + acc=acc, + fast_math=True, + lhs_k_pack=False, + rhs_k_pack=True, + ) + else: + acc = tl.ext.dot_scaled( + v_vals, + v_scales, + "e2m1", + p_vals.T, + p_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + out_vals = acc.T * out_scale + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + out_desc.store( + [(gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], + out_vals, + ) + return + + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + gen_idx) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES or page_start < kv_len: + compact_page = page_table_start + page_rel + if ASSUME_VALID_PAGES: + safe_compact_page = compact_page + physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) + safe_physical_page = physical_page + else: + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load(src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1).to( + tl.int64 + ) + valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + + p_rows = safe_compact_page * NUM_HEADS + offs_h + safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) + if USE_TMA_P_LOAD: + p_vals = p_desc.load([(safe_compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) + else: + p_vals = tl.load( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, + mask=mask_h[:, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None], + other=0, + ) + p_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + p_scales = tl.load(p_sf_ptr + p_sf_offsets) + + if ASSUME_FULL_PAGES: + valid_even_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + valid_odd_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + else: + valid_even_t = page_start + even_t < kv_len + valid_odd_t = page_start + odd_t < kv_len + if USE_TMA_V_LOAD: + v_tile = v_desc.load( + [ + safe_physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ] + ) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + if not ASSUME_VALID_PAGES: + v_tile = tl.where(valid_physical_page, v_tile, 0) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + if not ASSUME_FULL_PAGES: + even_packed = tl.where(valid_even_t[None, :], even_packed, 0) + odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + else: + even_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + even_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_even_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), + other=0, + ) + odd_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + odd_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_odd_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), + other=0, + ) + even_low = even_packed & 0x0F + even_high = (even_packed >> 4) & 0x0F + odd_low = odd_packed & 0x0F + even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) + odd_nibble = tl.where(v_use_high_nibble[:, None], odd_packed >> 4, odd_low) + v_vals = even_nibble | (odd_nibble << 4) + v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) + acc = tl.dot_scaled( + p_vals, + p_scales, + "e2m1", + v_vals.T, + v_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + if ASSUME_FULL_HEADS and ASSUME_FULL_V: + out_vals = acc * out_scale + if USE_TMA_V_LOAD: + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + out_desc.store( + [(gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], out_vals, ) else: @@ -1842,10 +2578,7 @@ def _fp4_mla_attention_pv_kernel( ) else: tl.store( - out_ptr - + gen_idx * out_s0 - + safe_offs_h[:, None] * out_s1 - + safe_offs_v[None, :] * out_s2, + out_ptr + gen_idx * out_s0 + safe_offs_h[:, None] * out_s1 + safe_offs_v[None, :] * out_s2, acc * out_scale, mask=mask_h[:, None] & mask_v[None, :], ) @@ -1885,10 +2618,14 @@ def fp4_mla_paged_attention_internal( use_tma_data_load: Optional[bool] = None, fused_prob_pack_single_launch: Optional[bool] = None, pack_prob_in_page_stats: Optional[bool] = None, + page_stats_group_size: Optional[int] = None, assume_full_pages: Optional[bool] = None, assume_valid_pages: Optional[bool] = None, + prepack_v_for_pv: bool = False, + use_prepacked_v_for_pv: bool = False, p_fp4_workspace: Optional[torch.Tensor] = None, p_sf_workspace: Optional[torch.Tensor] = None, + v_packed_workspace: Optional[torch.Tensor] = None, p_probs_workspace: Optional[torch.Tensor] = None, max_scores_workspace: Optional[torch.Tensor] = None, denom_workspace: Optional[torch.Tensor] = None, @@ -1898,9 +2635,7 @@ def fp4_mla_paged_attention_internal( ) -> torch.Tensor: del kwargs if not hasattr(tl, "dot_scaled"): - raise NotImplementedError( - "fp4_mla_paged_attention requires a Triton build with tl.dot_scaled." - ) + raise NotImplementedError("fp4_mla_paged_attention requires a Triton build with tl.dot_scaled.") if not q_fp4.is_cuda: raise ValueError("q_fp4 must be a CUDA tensor.") if q_fp4.dtype != torch.uint8 or kv_cache.dtype != torch.uint8: @@ -1910,9 +2645,7 @@ def fp4_mla_paged_attention_internal( if q_fp4.dim() == 3: inferred_num_gen, inferred_num_heads, packed_q_dim = q_fp4.shape if num_heads is not None and num_heads != inferred_num_heads: - raise ValueError( - f"num_heads={num_heads} does not match q_fp4.shape[1]={inferred_num_heads}." - ) + raise ValueError(f"num_heads={num_heads} does not match q_fp4.shape[1]={inferred_num_heads}.") num_gen = inferred_num_gen num_heads = inferred_num_heads q_fp4_2d = q_fp4.reshape(num_gen * num_heads, packed_q_dim) @@ -1927,15 +2660,11 @@ def fp4_mla_paged_attention_internal( else: raise ValueError("q_fp4 must be 2D or 3D.") - num_pages, inferred_page_size, packed_k_dim, kv_s0, kv_s2, kv_s4 = _get_kv_cache_strides( - kv_cache - ) + num_pages, inferred_page_size, packed_k_dim, kv_s0, kv_s2, kv_s4 = _get_kv_cache_strides(kv_cache) if page_size is None: page_size = inferred_page_size if page_size != inferred_page_size: - raise ValueError( - f"page_size={page_size} does not match kv_cache page dimension {inferred_page_size}." - ) + raise ValueError(f"page_size={page_size} does not match kv_cache page dimension {inferred_page_size}.") if page_size % FP4_BLOCK_SIZE != 0: raise ValueError(f"page_size must be divisible by {FP4_BLOCK_SIZE}.") @@ -1964,19 +2693,33 @@ def fp4_mla_paged_attention_internal( raise ValueError("v_sf is too small for the swizzled V scale layout.") if output is None: - output = torch.empty( - (num_gen, num_heads, v_head_dim), dtype=output_dtype, device=q_fp4.device - ) + output = torch.empty((num_gen, num_heads, v_head_dim), dtype=output_dtype, device=q_fp4.device) elif output.shape != (num_gen, num_heads, v_head_dim): - raise ValueError( - f"output must have shape {(num_gen, num_heads, v_head_dim)}, got {tuple(output.shape)}." - ) + raise ValueError(f"output must have shape {(num_gen, num_heads, v_head_dim)}, got {tuple(output.shape)}.") if num_gen == 0: return output triton_backend = "nvt" + env_block_h = _env_int("TRTLLM_FP4_MLA_BLOCK_H") + env_block_k = _env_int("TRTLLM_FP4_MLA_BLOCK_K") + env_block_v = _env_int("TRTLLM_FP4_MLA_BLOCK_V") + env_pv_loop_stages = _env_int("TRTLLM_FP4_MLA_PV_LOOP_STAGES") + env_occupancy = _env_int("TRTLLM_FP4_MLA_OCCUPANCY") + env_num_warps = _env_int("TRTLLM_FP4_MLA_NUM_WARPS") + env_num_stages = _env_int("TRTLLM_FP4_MLA_NUM_STAGES") + env_group_pages = _env_int("TRTLLM_FP4_MLA_GROUP_PAGES") + env_group_reduce_stats = _env_int("TRTLLM_FP4_MLA_GROUP_REDUCE_STATS") + env_page_pipeline_streams = _env_int("TRTLLM_FP4_MLA_PAGE_PIPELINE_STREAMS") + if env_block_h is not None: + block_h = env_block_h if block_k is None: - block_k = 512 if triton_backend == "nvt" else 256 + block_k = env_block_k or (512 if triton_backend == "nvt" else 256) + elif env_block_k is not None: + block_k = env_block_k + if env_block_v is not None: + block_v = env_block_v + if env_pv_loop_stages is not None: + pv_loop_stages = env_pv_loop_stages full_block_end = (q_head_dim // block_k) * block_k tail_k = q_head_dim - full_block_end tail_block_k = 1 << (tail_k - 1).bit_length() if tail_k > 0 else block_k @@ -2002,7 +2745,18 @@ def fp4_mla_paged_attention_internal( if assume_valid_pages is None: assume_valid_pages = False assume_valid_pages = bool(assume_valid_pages) + if ( + not assume_valid_pages + and assume_full_pages + and src_page_ids.numel() == num_gen * max_pages + ): + # Full decode pages with an exactly-sized page table do not need the + # sentinel/physical-page validity masks. Keeping this inference inside + # the kernel wrapper lets the framework call path stay unchanged. + assume_valid_pages = True total_p_rows = max(src_page_ids.numel() * num_heads, 1) + if page_pipeline_streams is None and env_page_pipeline_streams is not None: + page_pipeline_streams = env_page_pipeline_streams if page_pipeline_streams is None: if triton_backend == "nvt" and max_pages >= 8 and num_gen >= 128: page_pipeline_streams = 2 @@ -2010,14 +2764,24 @@ def fp4_mla_paged_attention_internal( page_pipeline_streams = 1 page_pipeline_streams = max(1, min(int(page_pipeline_streams), max_pages)) launch_meta = {} - if kernel_occupancy is None and triton_backend == "nvt": - kernel_occupancy = 2 + if kernel_occupancy is None: + if env_occupancy is not None: + kernel_occupancy = env_occupancy + elif triton_backend == "nvt": + kernel_occupancy = 8 if kernel_occupancy is not None: launch_meta["occupancy"] = int(kernel_occupancy) if kernel_num_ctas is not None: launch_meta["num_ctas"] = int(kernel_num_ctas) + if kernel_num_stages is None: + if env_num_stages is not None: + kernel_num_stages = env_num_stages + elif triton_backend == "nvt": + kernel_num_stages = 1 if kernel_num_stages is not None: launch_meta["num_stages"] = int(kernel_num_stages) + if kernel_num_warps is None and env_num_warps is not None: + kernel_num_warps = env_num_warps if kernel_num_warps is not None: launch_meta["num_warps"] = int(kernel_num_warps) if fused_prob_pack is None: @@ -2048,6 +2812,54 @@ def alloc_fn(size: int, alignment: int, stream: Optional[int]): device=q_fp4.device, name="p_sf", ) + num_dim_blocks = triton.cdiv(v_head_dim, block_v) + v_repack_block_v = block_v + num_repack_dim_blocks = num_dim_blocks + auto_prepack_v_for_pv = ( + triton_backend == "nvt" + and use_tma_data_load + and assume_full_heads + and assume_full_pages + and assume_valid_pages + and v_head_dim == 512 + and page_size == 128 + and block_h in (64, 128) + and block_v == 128 + and sf_per_page == 8 + ) + if not prepack_v_for_pv and not use_prepacked_v_for_pv: + env_prepack_v = os.environ.get("TRTLLM_FP4_MLA_PREPACK_V") + prepack_v_for_pv = auto_prepack_v_for_pv if env_prepack_v is None else env_prepack_v == "1" + wants_prepacked_v_for_pv = bool(prepack_v_for_pv) or bool(use_prepacked_v_for_pv) + if use_prepacked_v_for_pv and v_packed_workspace is None: + raise ValueError("use_prepacked_v_for_pv requires v_packed_workspace to be provided.") + can_use_prepacked_v_for_pv = ( + wants_prepacked_v_for_pv + and use_tma_data_load + and assume_full_heads + and assume_full_pages + and assume_valid_pages + and v_head_dim == 512 + and page_size == 128 + and block_h in (64, 128) + and block_v == 128 + and sf_per_page == 8 + ) + if wants_prepacked_v_for_pv and not can_use_prepacked_v_for_pv: + raise ValueError( + "prepacked V PV path requires TMA, full heads/pages, valid pages, " + "v_head_dim=512, page_size=128, block_h in (64, 128), block_v=128, and sf_per_page=8." + ) + if can_use_prepacked_v_for_pv: + v_packed = _workspace_tensor( + v_packed_workspace, + (num_pages * num_dim_blocks * block_v, page_size // 2), + dtype=torch.uint8, + device=q_fp4.device, + name="v_packed", + ) + else: + v_packed = kv_cache if fused_prob_pack: p_probs = None else: @@ -2083,12 +2895,76 @@ def alloc_fn(size: int, alignment: int, stream: Optional[int]): name="denom", ) + v_repack_stream = None + if bool(prepack_v_for_pv) and can_use_prepacked_v_for_pv: + current_stream = torch.cuda.current_stream(q_fp4.device) + v_repack_stream = torch.cuda.Stream(device=q_fp4.device) + v_repack_stream.wait_stream(current_stream) + with torch.cuda.stream(v_repack_stream): + _fp4_mla_attention_v_repack_kernel[(num_pages, num_repack_dim_blocks)]( + v_packed, + kv_cache, + num_pages, + kv_s0, + kv_s2, + kv_s4, + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + BLOCK_V=v_repack_block_v, + **launch_meta, + ) + if parallel_page_stats is None: parallel_page_stats = triton_backend == "nvt" and max_pages >= 8 if pack_prob_in_page_stats is None: pack_prob_in_page_stats = parallel_page_stats and fused_prob_pack - pack_prob_in_page_stats = bool( - pack_prob_in_page_stats and parallel_page_stats and fused_prob_pack + pack_prob_in_page_stats = bool(pack_prob_in_page_stats and parallel_page_stats and fused_prob_pack) + page_stats_group_sizes = (2, 4, 8, 16, 32, 64, 128, 256, 512, 1024) + if page_stats_group_size is None and env_group_pages is not None: + page_stats_group_size = env_group_pages + if page_stats_group_size is None: + if num_gen <= 32: + target_group_pages = max(8, 8 * num_gen) + elif num_gen >= 256: + target_group_pages = 256 + else: + target_group_pages = max(8, min(512, 4 * num_gen)) + target_group_pages = min(max_pages, target_group_pages) + page_stats_group_size = next( + ( + group_size + for group_size in reversed(page_stats_group_sizes) + if group_size <= target_group_pages and max_pages % group_size == 0 + ), + 8, + ) + else: + page_stats_group_size = int(page_stats_group_size) + can_group_page_stats = ( + page_stats_group_size in page_stats_group_sizes + and triton_backend == "nvt" + and parallel_page_stats + and pack_prob_in_page_stats + and use_tma_data_load + and assume_full_heads + and assume_full_pages + and assume_valid_pages + and num_heads == 128 + and q_head_dim == 640 + and k_head_dim == 576 + and q_residual_dim == 64 + and page_size == 128 + and block_h == 128 + and block_k == 512 + and full_block_end == 512 + and tail_block_k == 128 + and sf_per_page == 8 + and max_pages % page_stats_group_size == 0 + ) + page_stats_group_size = page_stats_group_size if can_group_page_stats else 1 + group_reduce_stats = ( + (env_group_reduce_stats != 0 if env_group_reduce_stats is not None else triton_backend == "nvt") + and page_stats_group_size > 1 ) if parallel_page_stats: page_stats_shape = (num_gen, max_pages, num_heads) @@ -2106,56 +2982,112 @@ def alloc_fn(size: int, alignment: int, stream: Optional[int]): device=q_fp4.device, name="page_sum", ) - _fp4_mla_attention_page_stats_kernel[(num_gen, num_head_blocks, max_pages)]( - page_max, - page_sum, - p_fp4, - p_sf, - q_fp4_2d, - q_sf_flat, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], - num_pages, - q_fp4_2d.stride(0), - q_fp4_2d.stride(1), - kv_s0, - kv_s2, - kv_s4, - sf_cache.stride(0), - page_max.stride(0), - page_max.stride(1), - p_fp4.stride(0), - p_fp4.stride(1), - p_fp4.shape[0], - q_fp4_2d.shape[0], - sm_scale=sm_scale, - NUM_HEADS=num_heads, - Q_HEAD_D=q_head_dim, - K_HEAD_D=k_head_dim, - Q_RESIDUAL_D=q_residual_dim, - PAGE_SIZE=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - Q_SF_PER_TOKEN=q_sf_per_token, - K_SF_PER_TOKEN=k_sf_per_token, - SF_PER_PAGE=sf_per_page, - P_GLOBAL_SCALE=p_global_scale, - BLOCK_H=block_h, - BLOCK_T=page_size, - BLOCK_K=block_k, - FULL_BLOCK_END=full_block_end, - TAIL_BLOCK_K=tail_block_k, - USE_TMA_DATA_LOAD=use_tma_data_load, - PACK_PROBS=pack_prob_in_page_stats, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, - ) + if page_stats_group_size > 1: + _fp4_mla_attention_page_stats_grouped_kernel[ + (num_gen, num_head_blocks, max_pages // page_stats_group_size) + ]( + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + page_max.stride(0), + page_max.stride(1), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + BLOCK_T=page_size, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + PACK_PROBS=pack_prob_in_page_stats, + GROUP_REDUCE_STATS=group_reduce_stats, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + GROUP_PAGES=page_stats_group_size, + **launch_meta, + ) + else: + _fp4_mla_attention_page_stats_kernel[(num_gen, num_head_blocks, max_pages)]( + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + page_max.stride(0), + page_max.stride(1), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + BLOCK_T=page_size, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + PACK_PROBS=pack_prob_in_page_stats, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) _fp4_mla_attention_reduce_stats_kernel[(num_gen, num_head_blocks)]( max_scores, denom, @@ -2167,6 +3099,8 @@ def alloc_fn(size: int, alignment: int, stream: Optional[int]): NUM_HEADS=num_heads, MAX_PAGES=max_pages, BLOCK_H=block_h, + GROUP_REDUCE_STATS=group_reduce_stats, + GROUP_PAGES=page_stats_group_size, **launch_meta, ) if pack_prob_in_page_stats: @@ -2421,54 +3355,109 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None for stream in streams: current_stream.wait_stream(stream) - num_dim_blocks = triton.cdiv(v_head_dim, block_v) - _fp4_mla_attention_pv_kernel[ - ( - num_gen, - num_head_blocks, - num_dim_blocks, + if v_repack_stream is not None: + torch.cuda.current_stream(q_fp4.device).wait_stream(v_repack_stream) + + if can_use_prepacked_v_for_pv: + _fp4_mla_attention_pv_prepacked_v_kernel[ + ( + num_gen, + num_head_blocks, + num_dim_blocks, + ) + ]( + output, + p_fp4, + p_sf, + kv_cache, + v_sf, + v_packed, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + output.stride(0), + output.stride(1), + output.stride(2), + output.shape[0] * output.shape[1], + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + kv_s0, + kv_s2, + kv_s4, + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + BLOCK_V=block_v, + USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, + USE_TMA_V_LOAD=use_tma_data_load and v_head_dim % block_v == 0, + USE_PREPACKED_V=True, + PV_M_PACKED_V=False, + PV_LOOP_STAGES=int(pv_loop_stages), + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_FULL_V=assume_full_v, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + else: + num_dim_blocks = triton.cdiv(v_head_dim, block_v) + _fp4_mla_attention_pv_kernel[ + ( + num_gen, + num_head_blocks, + num_dim_blocks, + ) + ]( + output, + p_fp4, + p_sf, + kv_cache, + v_sf, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + output.stride(0), + output.stride(1), + output.stride(2), + output.shape[0] * output.shape[1], + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + kv_s0, + kv_s2, + kv_s4, + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + BLOCK_V=block_v, + USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, + USE_TMA_V_LOAD=use_tma_data_load and v_head_dim % block_v == 0, + PV_LOOP_STAGES=int(pv_loop_stages), + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_FULL_V=assume_full_v, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, ) - ]( - output, - p_fp4, - p_sf, - kv_cache, - v_sf, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], - num_pages, - output.stride(0), - output.stride(1), - output.stride(2), - output.shape[0] * output.shape[1], - p_fp4.stride(0), - p_fp4.stride(1), - p_fp4.shape[0], - kv_s0, - kv_s2, - kv_s4, - v_sf.stride(0), - NUM_HEADS=num_heads, - V_HEAD_D=v_head_dim, - PAGE_SIZE=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - SF_PER_PAGE=sf_per_page, - MAX_PAGES=max_pages, - P_GLOBAL_SCALE=p_global_scale, - BLOCK_H=block_h, - BLOCK_V=block_v, - USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, - USE_TMA_V_LOAD=use_tma_data_load and v_head_dim % block_v == 0, - PV_LOOP_STAGES=int(pv_loop_stages), - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_FULL_V=assume_full_v, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, - ) return output diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py.bak b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py.bak new file mode 100644 index 000000000000..3cd4f7bdb94c --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py.bak @@ -0,0 +1,3239 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# SPDX-License-Identifier: Apache-2.0 + +"""FP4 MLA paged decode attention using Triton. + +The kernels are adapted from TensorRT-LLM's FP4 MLA decode path. This module +exposes the attention path for already-packed FP4 Q/K/V tensors and swizzled +FP8 block-scale tensors; quantization and KV-cache update helpers remain outside +this internal op. +""" + +import os +from typing import Optional + +import torch +import triton +import triton.language as tl + +FP4_BLOCK_SIZE = 16 +FP4_MLA_P_GLOBAL_SCALE = 448.0 * 6.0 + + +def _ceil_div(lhs: int, rhs: int) -> int: + return (lhs + rhs - 1) // rhs + + +def _env_int(name: str) -> Optional[int]: + value = os.environ.get(name) + if value is None or value == "": + return None + return int(value) + + +def _swizzled_scale_size(rows: int, logical_cols: int) -> int: + scale_cols = _ceil_div(logical_cols, FP4_BLOCK_SIZE) + padded_cols = _ceil_div(scale_cols, 4) * 4 + return _ceil_div(rows, 128) * 128 * padded_cols + + +def _get_kv_cache_strides(kv_cache: torch.Tensor) -> tuple[int, int, int, int, int, int]: + if kv_cache.dim() == 3: + num_pages, page_size, packed_dim = kv_cache.shape + return ( + num_pages, + page_size, + packed_dim, + kv_cache.stride(0), + kv_cache.stride(1), + kv_cache.stride(2), + ) + if kv_cache.dim() >= 5: + num_pages = kv_cache.shape[0] + page_size = kv_cache.shape[2] + packed_dim = kv_cache.shape[4] + return ( + num_pages, + page_size, + packed_dim, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + ) + raise ValueError( + "kv_cache must be shaped (num_pages, page_size, packed_dim) or (num_pages, ..., page_size, ..., packed_dim)." + ) + + +def _workspace_tensor( + workspace: Optional[torch.Tensor], + shape: tuple[int, ...], + *, + dtype: torch.dtype, + device: torch.device, + name: str, +) -> torch.Tensor: + if workspace is None: + if torch.cuda.is_current_stream_capturing(): + raise ValueError( + f"Cannot allocate {name} while capturing a CUDA graph. " + "Pass a preallocated workspace tensor." + ) + return torch.empty(shape, dtype=dtype, device=device) + + invalid = ( + workspace.dtype != dtype + or workspace.device != device + or len(workspace.shape) != len(shape) + or any(workspace.shape[idx] < dim for idx, dim in enumerate(shape)) + ) + if invalid: + raise ValueError( + f"{name} workspace must have shape at least {shape}, dtype={dtype}, " + f"and device={device}; got shape={tuple(workspace.shape)}, " + f"dtype={workspace.dtype}, device={workspace.device}." + ) + + slices = tuple(slice(0, dim) for dim in shape) + return workspace[slices] + + +@triton.jit +def _fp4_mla_swizzled_sf_offset(row_idx, col_idx, SF_PER_TOKEN: tl.constexpr): + padded_cols = ((SF_PER_TOKEN + 3) // 4) * 4 + col_in_group = col_idx % 4 + col_group = col_idx // 4 + row_in_group0 = row_idx % 32 + row_in_group1 = (row_idx % 128) // 32 + row_group = row_idx // 128 + return ( + col_in_group + + col_group * (4 * 128) + + row_in_group0 * 16 + + row_in_group1 * 4 + + row_group * (128 * padded_cols) + ) + + +@triton.jit +def _fp4_mla_swizzled_sf_offset_row_block( + row_group, row_offsets, col_idx, SF_PER_TOKEN: tl.constexpr +): + padded_cols = ((SF_PER_TOKEN + 3) // 4) * 4 + col_part = (col_idx % 4) + (col_idx // 4) * (4 * 128) + row_part = (row_offsets % 32) * 16 + ((row_offsets % 128) // 32) * 4 + return col_part + row_part + row_group * (128 * padded_cols) + + +@triton.jit +def _fp4_e2m1_quantize(x): + abs_x = tl.abs(x) + magnitude = tl.where( + abs_x < 0.25, + 0, + tl.where( + abs_x < 0.75, + 1, + tl.where( + abs_x < 1.25, + 2, + tl.where( + abs_x < 1.75, + 3, + tl.where(abs_x < 2.5, 4, tl.where(abs_x < 3.5, 5, tl.where(abs_x < 5.0, 6, 7))), + ), + ), + ), + ) + sign = tl.where(x < 0.0, 8, 0) + return (magnitude | sign).to(tl.uint8) + + +@triton.jit +def _fp4_e2m1_quantize_packed(even, odd): + return tl.inline_asm_elementwise( + """ + { + .reg .b8 r; + cvt.rn.satfinite.e2m1x2.f32 r, $1, $2; + mov.b32 $0, {r, r, r, r}; + } + """, + constraints="=r,f,f", + args=[odd.to(tl.float32), even.to(tl.float32)], + dtype=tl.uint8, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fp4_pack_low_nibbles(even_packed, odd_packed): + return tl.inline_asm_elementwise( + """ + { + .reg .b32 lo; + .reg .b32 hi; + and.b32 lo, $1, 15; + and.b32 hi, $2, 15; + shl.b32 hi, hi, 4; + or.b32 $0, lo, hi; + } + """, + constraints="=r,r,r", + args=[even_packed, odd_packed], + dtype=tl.uint8, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fp4_pack_high_nibbles(even_packed, odd_packed): + return tl.inline_asm_elementwise( + """ + { + .reg .b32 lo; + .reg .b32 hi; + shr.u32 lo, $1, 4; + and.b32 lo, lo, 15; + and.b32 hi, $2, 240; + or.b32 $0, lo, hi; + } + """, + constraints="=r,r,r", + args=[even_packed, odd_packed], + dtype=tl.uint8, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fp4_pack_nibbles(even_packed, odd_packed): + return tl.inline_asm_elementwise( + """ + { + .reg .b32 lo; + .reg .b32 hi; + and.b32 lo, $2, 15; + and.b32 hi, $3, 15; + shl.b32 hi, hi, 4; + or.b32 $0, lo, hi; + + shr.u32 lo, $2, 4; + and.b32 lo, lo, 15; + and.b32 hi, $3, 240; + or.b32 $1, lo, hi; + } + """, + constraints="=r,=r,r,r", + args=[even_packed, odd_packed], + dtype=(tl.uint8, tl.uint8), + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fp4_mla_attention_v_repack_kernel( + v_packed_ptr, + kv_cache_ptr, + num_pages, + kv_s0, + kv_s2, + kv_s4, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BLOCK_V: tl.constexpr, + occupancy: tl.constexpr = 1, +): + page_idx = tl.program_id(0) + dim_block = tl.program_id(1) + + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_view = tl.ext.make_view( + base=kv_cache_ptr, + shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], + tile_dim_map=[0, 1, 2], + ) + v_tile = tl.ext.load_view_tko( + v_view, + [ + page_idx.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ], + ) + v_tile = v_tile.to(tl.uint8, bitcast=True) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + + out_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + row_base = (page_idx * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + out_desc.store([row_base.to(tl.int32), 0], v_vals) + + +@triton.jit +def _fp4_mla_attention_v_repack_pages_kernel( + v_packed_ptr, + kv_cache_ptr, + page_ids_ptr, + num_pages, + kv_s0, + kv_s2, + kv_s4, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BLOCK_V: tl.constexpr, + occupancy: tl.constexpr = 1, +): + page_list_idx = tl.program_id(0) + dim_block = tl.program_id(1) + page_idx = tl.load(page_ids_ptr + page_list_idx).to(tl.int64) + + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_view = tl.ext.make_view( + base=kv_cache_ptr, + shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], + tile_dim_map=[0, 1, 2], + ) + v_tile = tl.ext.load_view_tko( + v_view, + [ + page_idx.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ], + ) + v_tile = v_tile.to(tl.uint8, bitcast=True) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + + out_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + row_base = (page_idx * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + out_desc.store([row_base.to(tl.int32), 0], v_vals) + + +def fp4_mla_repack_v_cache( + v_packed: torch.Tensor, + kv_cache: torch.Tensor, + page_ids: Optional[torch.Tensor] = None, + *, + v_head_dim: int, + page_size: int, + block_v: int = 128, + kernel_occupancy: int = 2, + kernel_num_stages: int = 1, +) -> None: + """Populate the V-packed auxiliary cache consumed by the prepacked PV kernel.""" + if v_head_dim % block_v != 0: + raise ValueError(f"v_head_dim={v_head_dim} must be divisible by block_v={block_v}.") + if kv_cache.ndim < 5: + raise ValueError( + f"kv_cache must expose the paged FP4 layout, got shape={tuple(kv_cache.shape)}." + ) + num_pages = kv_cache.shape[0] + num_dim_blocks = triton.cdiv(v_head_dim, block_v) + launch_meta = { + "occupancy": int(kernel_occupancy), + "num_stages": int(kernel_num_stages), + } + if page_ids is None: + if num_pages == 0: + return + _fp4_mla_attention_v_repack_kernel[(num_pages, num_dim_blocks)]( + v_packed, + kv_cache, + num_pages, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + BLOCK_V=block_v, + **launch_meta, + ) + return + + if page_ids.numel() == 0: + return + _fp4_mla_attention_v_repack_pages_kernel[(page_ids.numel(), num_dim_blocks)]( + v_packed, + kv_cache, + page_ids, + num_pages, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + BLOCK_V=block_v, + **launch_meta, + ) + + +@triton.jit +def _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_start, + head_offsets, + token_offsets, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + NUM_HEADS: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, +): + if ASSUME_VALID_PAGES: + safe_compact_page = compact_page + physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) + safe_physical_page = physical_page + else: + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + ) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + q_rows = q_row_base + head_offsets + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_q_rows = q_rows + else: + mask_h = head_offsets < NUM_HEADS + safe_q_rows = tl.where(mask_h, q_rows, q_row_base) + scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) + + if ( + USE_TMA_DATA_LOAD + and ASSUME_FULL_HEADS + and ASSUME_VALID_PAGES + and Q_HEAD_D == 640 + and K_HEAD_D == 576 + and Q_RESIDUAL_D == 64 + and BLOCK_H == 128 + and BLOCK_T == 128 + and BLOCK_K == 512 + and FULL_BLOCK_END == 512 + and TAIL_BLOCK_K == 128 + ): + tl.assume(q_fp4_s0 % 8 == 0) + tl.assume(q_fp4_s1 == 1) + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + q_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 256], + ) + k_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 256], + ) + q_tail_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 64], + ) + k_tail_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 32], + ) + q_sf_full_view = tl.ext.make_view( + base=q_sf_ptr, + shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + k_sf_full_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + k_sf_tail_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 1, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + + full_q_vals = q_desc.load([(q_row_base + head_start).to(tl.int32), 0]) + full_k_vals = k_desc.load([safe_physical_page.to(tl.int32), 0, 0]) + full_k_vals = tl.reshape(full_k_vals, (BLOCK_T, 256)) + q_row_group = q_row_base // 128 + full_q_scales = tl.ext.load_view_tko(q_sf_full_view, [q_row_group.to(tl.int32), 0, 0, 0]) + full_q_scales = full_q_scales.reshape([1, 8, 32, 4, 4]).trans(0, 3, 2, 1, 4) + full_q_scales = full_q_scales.reshape([BLOCK_H, 32]) + full_k_scales = tl.ext.load_view_tko( + k_sf_full_view, [safe_physical_page.to(tl.int32), 0, 0, 0, 0] + ) + full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) + scores = tl.dot_scaled( + full_q_vals, + full_q_scales, + "e2m1", + full_k_vals.T, + full_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + tail_k_vals = k_tail_desc.load([safe_physical_page.to(tl.int32), 0, 256]) + tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) + tail_k_scales = tl.ext.load_view_tko( + k_sf_tail_view, [safe_physical_page.to(tl.int32), 0, 8, 0, 0] + ) + tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) + q_tail_vals = q_tail_desc.load([(q_row_base + head_start).to(tl.int32), 256]) + # Map Q tail groups [0, 1, ..., 7] onto K tail groups [0, 0, 1, 1, ..., 3, 3]. + q_tail_vals = q_tail_vals.reshape([BLOCK_H, 4, 2, 8]).trans(0, 1, 3, 2) + q_even_vals, q_odd_vals = tl.split(q_tail_vals) + q_even_vals = q_even_vals.reshape([BLOCK_H, 32]) + q_odd_vals = q_odd_vals.reshape([BLOCK_H, 32]) + + q_tail_sf_cols = 32 + tl.arange(0, 8) + q_tail_sf_offsets = _fp4_mla_swizzled_sf_offset( + q_rows[:, None], q_tail_sf_cols[None, :], Q_SF_PER_TOKEN + ) + q_tail_scales = tl.load(q_sf_ptr + q_tail_sf_offsets) + q_tail_scales = q_tail_scales.reshape([BLOCK_H, 4, 2]) + q_even_scales, q_odd_scales = tl.split(q_tail_scales) + scores = tl.dot_scaled( + q_even_vals, + q_even_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q_odd_vals, + q_odd_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + return scores + + packed_k_offsets = tl.arange(0, BLOCK_K // 2) + scale_offsets = tl.arange(0, BLOCK_K // FP4_BLOCK) + residual_groups = Q_RESIDUAL_D // FP4_BLOCK + non_residual_groups = K_HEAD_D // FP4_BLOCK - residual_groups + if USE_TMA_DATA_LOAD and FULL_BLOCK_END > 0: + tl.assume(q_fp4_s0 % 8 == 0) + tl.assume(q_fp4_s1 == 1) + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + q_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, BLOCK_K // 2], + ) + k_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, BLOCK_K // 2], + ) + if USE_TMA_DATA_LOAD and Q_RESIDUAL_D == 64 and TAIL_BLOCK_K == 128: + q_tail_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 64], + ) + k_tail_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 32], + ) + for q_start in tl.range(0, FULL_BLOCK_END, BLOCK_K): + q_elem_offsets = q_start + packed_k_offsets * 2 + q_group_offsets = q_elem_offsets // FP4_BLOCK + k_group_offsets = tl.where( + q_group_offsets < non_residual_groups, + q_group_offsets, + non_residual_groups + (q_group_offsets - non_residual_groups) // 2, + ) + byte_offsets_in_group = (q_elem_offsets % FP4_BLOCK) // 2 + packed_q_cols = q_start // 2 + packed_k_offsets + packed_k_cols = k_group_offsets * (FP4_BLOCK // 2) + byte_offsets_in_group + mask_k = q_elem_offsets < Q_HEAD_D + safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) + safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) + if ( + USE_TMA_DATA_LOAD + and FULL_BLOCK_END > 0 + and q_start + BLOCK_K <= non_residual_groups * FP4_BLOCK + ): + q_vals = q_desc.load([(q_row_base + head_start).to(tl.int32), q_start // 2]) + k_vals = k_desc.load([safe_physical_page.to(tl.int32), 0, q_start // 2]) + k_vals = tl.reshape(k_vals, (BLOCK_T, BLOCK_K // 2)) + if not ASSUME_VALID_PAGES: + k_vals = tl.where(valid_physical_page, k_vals, 0) + else: + q_vals = tl.load( + q_fp4_ptr + + safe_q_rows[:, None] * q_fp4_s0 + + safe_packed_q_cols[None, :] * q_fp4_s1, + mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], + other=0, + ) + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + safe_packed_k_cols[None, :] * kv_s4, + mask=mask_k[None, :] + if ASSUME_VALID_PAGES + else valid_physical_page & mask_k[None, :], + other=0, + ) + + q_sf_cols = q_start // FP4_BLOCK + scale_offsets + k_sf_cols = tl.where( + q_sf_cols < non_residual_groups, + q_sf_cols, + non_residual_groups + (q_sf_cols - non_residual_groups) // 2, + ) + mask_sf = q_sf_cols < Q_SF_PER_TOKEN + safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) + safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) + q_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN + ) + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN + ) + q_scales = tl.load(q_sf_ptr + q_sf_offsets) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) + scores = tl.dot_scaled( + q_vals, + q_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + if FULL_BLOCK_END < Q_HEAD_D: + q_start = FULL_BLOCK_END + if Q_RESIDUAL_D == 64 and TAIL_BLOCK_K == 128: + residual_packed_offsets = tl.arange(0, 32) + residual_scale_offsets = tl.arange(0, 4) + packed_k_cols = non_residual_groups * (FP4_BLOCK // 2) + residual_packed_offsets + if USE_TMA_DATA_LOAD: + k_vals = k_tail_desc.load( + [ + safe_physical_page.to(tl.int32), + 0, + (non_residual_groups * (FP4_BLOCK // 2)).to(tl.int32), + ] + ) + k_vals = tl.reshape(k_vals, (BLOCK_T, 32)) + if not ASSUME_VALID_PAGES: + k_vals = tl.where(valid_physical_page, k_vals, 0) + elif ASSUME_VALID_PAGES: + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + packed_k_cols[None, :] * kv_s4, + ) + else: + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + packed_k_cols[None, :] * kv_s4, + mask=valid_physical_page, + other=0, + ) + k_sf_cols = non_residual_groups + residual_scale_offsets + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], k_sf_cols[None, :], K_SF_PER_TOKEN + ) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) + + q_tail_cols = q_start // 2 + tl.arange(0, 64) + if USE_TMA_DATA_LOAD and ASSUME_FULL_HEADS: + q_tail_vals = q_tail_desc.load( + [(q_row_base + head_start).to(tl.int32), q_start // 2] + ) + elif ASSUME_FULL_HEADS: + q_tail_vals = tl.load( + q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1 + ) + else: + q_tail_vals = tl.load( + q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1, + mask=mask_h[:, None], + other=0, + ) + # Map Q tail groups [0, 1, ..., 7] onto K tail groups [0, 0, 1, 1, ..., 3, 3]. + q_tail_vals = q_tail_vals.reshape([BLOCK_H, 4, 2, 8]).trans(0, 1, 3, 2) + q_even_vals, q_odd_vals = tl.split(q_tail_vals) + q_even_vals = q_even_vals.reshape([BLOCK_H, 32]) + q_odd_vals = q_odd_vals.reshape([BLOCK_H, 32]) + + q_tail_sf_cols = q_start // FP4_BLOCK + tl.arange(0, 8) + q_tail_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], q_tail_sf_cols[None, :], Q_SF_PER_TOKEN + ) + q_tail_scales = tl.load(q_sf_ptr + q_tail_sf_offsets) + q_tail_scales = q_tail_scales.reshape([BLOCK_H, 4, 2]) + q_even_scales, q_odd_scales = tl.split(q_tail_scales) + scores = tl.dot_scaled( + q_even_vals, + q_even_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q_odd_vals, + q_odd_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + else: + tail_packed_offsets = tl.arange(0, TAIL_BLOCK_K // 2) + tail_scale_offsets = tl.arange(0, TAIL_BLOCK_K // FP4_BLOCK) + q_elem_offsets = q_start + tail_packed_offsets * 2 + q_group_offsets = q_elem_offsets // FP4_BLOCK + k_group_offsets = tl.where( + q_group_offsets < non_residual_groups, + q_group_offsets, + non_residual_groups + (q_group_offsets - non_residual_groups) // 2, + ) + byte_offsets_in_group = (q_elem_offsets % FP4_BLOCK) // 2 + packed_q_cols = q_start // 2 + tail_packed_offsets + packed_k_cols = k_group_offsets * (FP4_BLOCK // 2) + byte_offsets_in_group + mask_k = q_elem_offsets < Q_HEAD_D + safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) + safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) + q_vals = tl.load( + q_fp4_ptr + + safe_q_rows[:, None] * q_fp4_s0 + + safe_packed_q_cols[None, :] * q_fp4_s1, + mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], + other=0, + ) + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + safe_packed_k_cols[None, :] * kv_s4, + mask=mask_k[None, :] + if ASSUME_VALID_PAGES + else valid_physical_page & mask_k[None, :], + other=0, + ) + + q_sf_cols = q_start // FP4_BLOCK + tail_scale_offsets + k_sf_cols = tl.where( + q_sf_cols < non_residual_groups, + q_sf_cols, + non_residual_groups + (q_sf_cols - non_residual_groups) // 2, + ) + mask_sf = q_sf_cols < Q_SF_PER_TOKEN + safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) + safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) + q_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN + ) + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN + ) + q_scales = tl.load(q_sf_ptr + q_sf_offsets) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) + scores = tl.dot_scaled( + q_vals, + q_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + return scores + + +@triton.jit +def _fp4_mla_attention_stats_kernel( + max_ptr, + denom_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + stats_s0, + q_num_rows, + sm_scale, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, BLOCK_T) + q_row_base = query_idx * NUM_HEADS + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + + max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + denom = tl.zeros((BLOCK_H,), dtype=tl.float32) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + for page_rel in tl.range(0, MAX_PAGES): + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES or page_start < kv_len: + compact_page = page_table_start + page_rel + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + BLOCK_T, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + if ASSUME_FULL_PAGES: + scores = tl.where(mask_h[:, None], scores * qk_scale, -float("inf")) + else: + valid_t = page_start + offs_t < kv_len + scores = tl.where( + mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf") + ) + page_max = tl.max(scores, axis=1) + new_max = tl.maximum(max_score, page_max) + denom = denom * tl.math.exp2((max_score - new_max) * 1.4426950408889634) + tl.sum( + tl.math.exp2((scores - new_max[:, None]) * 1.4426950408889634), axis=1 + ) + max_score = new_max + + tl.store(max_ptr + query_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) + tl.store(denom_ptr + query_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) + + +@triton.jit +def _fp4_mla_attention_page_stats_kernel( + page_max_ptr, + page_sum_ptr, + p_fp4_ptr, + p_sf_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_stats_s0, + page_stats_s1, + p_s0, + p_s1, + p_num_rows, + q_num_rows, + sm_scale, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + PACK_PROBS: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_rel = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, BLOCK_T) + q_row_base = query_idx * NUM_HEADS + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_start = page_rel * PAGE_SIZE + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + out_offsets = query_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h + + page_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + page_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) + if USE_TMA_DATA_LOAD and PACK_PROBS and ASSUME_FULL_HEADS and ASSUME_VALID_PAGES: + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + if ASSUME_FULL_PAGES or page_start < kv_len: + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + page_table_start + page_rel, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + BLOCK_T, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + if ASSUME_FULL_PAGES: + valid_t = tl.full([BLOCK_T], True, dtype=tl.int1) + else: + valid_t = page_start + offs_t < kv_len + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + if ASSUME_FULL_HEADS and ASSUME_FULL_PAGES: + scores = scores * qk_scale + page_max = tl.max(scores, axis=1) + exp_scores = tl.math.exp2((scores - page_max[:, None]) * 1.4426950408889634) + page_sum = tl.sum(exp_scores, axis=1) + else: + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + page_max = tl.max(scores, axis=1) + safe_page_max = tl.where(mask_h, page_max, 0.0) + exp_scores = tl.math.exp2((scores - safe_page_max[:, None]) * 1.4426950408889634) + exp_scores = tl.where(mask_h[:, None] & valid_t[None, :], exp_scores, 0.0) + page_sum = tl.sum(exp_scores, axis=1) + + if PACK_PROBS: + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + stored_scale = tl.where( + amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0 + ) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + if not ASSUME_VALID_PAGES: + valid_compact_page = (page_table_start + page_rel >= 0) & ( + page_table_start + page_rel < page_ids_len + ) + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + ) + scale_cols = tl.arange(0, SF_PER_PAGE) + if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + if ASSUME_FULL_HEADS: + if ASSUME_VALID_PAGES: + tl.store(p_sf_ptr + sf_offsets, stored_scale) + else: + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_compact_page) + else: + tl.store( + p_sf_ptr + sf_offsets, + stored_scale, + mask=mask_h[:, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None], + ) + + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + if USE_TMA_DATA_LOAD and ASSUME_FULL_HEADS and ASSUME_VALID_PAGES: + p_desc.store( + [(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), + ) + elif ASSUME_FULL_HEADS: + if ASSUME_VALID_PAGES: + tl.store( + p_fp4_ptr + + safe_p_rows[:, None, None] * p_s0 + + byte_cols[None, :, :] * p_s1, + packed, + ) + else: + tl.store( + p_fp4_ptr + + safe_p_rows[:, None, None] * p_s0 + + byte_cols[None, :, :] * p_s1, + packed, + mask=valid_compact_page, + ) + else: + tl.store( + p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=mask_h[:, None, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None, None], + ) + + if ASSUME_FULL_HEADS: + tl.store(page_max_ptr + out_offsets, page_max) + tl.store(page_sum_ptr + out_offsets, page_sum) + else: + tl.store(page_max_ptr + out_offsets, page_max, mask=mask_h) + tl.store(page_sum_ptr + out_offsets, page_sum, mask=mask_h) + + +@triton.jit +def _fp4_mla_attention_page_stats_grouped_kernel( + page_max_ptr, + page_sum_ptr, + p_fp4_ptr, + p_sf_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len: tl.constexpr, + num_pages: tl.constexpr, + q_fp4_s0: tl.constexpr, + q_fp4_s1: tl.constexpr, + kv_s0: tl.constexpr, + kv_s2: tl.constexpr, + kv_s4: tl.constexpr, + sf_s0: tl.constexpr, + page_stats_s0: tl.constexpr, + page_stats_s1: tl.constexpr, + p_s0: tl.constexpr, + p_s1: tl.constexpr, + p_num_rows: tl.constexpr, + q_num_rows: tl.constexpr, + sm_scale: tl.constexpr, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + PACK_PROBS: tl.constexpr, + GROUP_REDUCE_STATS: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + GROUP_PAGES: tl.constexpr = 2, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_group = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + scale_cols = tl.arange(0, SF_PER_PAGE) + q_row_base = query_idx * NUM_HEADS + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + + tl.assume(q_fp4_s0 % 8 == 0) + tl.assume(q_fp4_s1 == 1) + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + + q_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 256], + ) + k_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 256], + ) + q_tail_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 32], + ) + k_tail_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 32], + ) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + q_sf_full_view = tl.ext.make_view( + base=q_sf_ptr, + shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + q_sf_tail_view = tl.ext.make_view( + base=q_sf_ptr, + shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + k_sf_full_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + k_sf_tail_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 1, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + + q_row_start = (q_row_base + head_block * BLOCK_H).to(tl.int32) + q_row_group = q_row_base // 128 + full_q_vals = q_desc.load([q_row_start, 0]) + full_q_scales = tl.ext.load_view_tko(q_sf_full_view, [q_row_group.to(tl.int32), 0, 0, 0]) + full_q_scales = full_q_scales.reshape([1, 8, 32, 4, 4]).trans(0, 3, 2, 1, 4) + full_q_scales = full_q_scales.reshape([BLOCK_H, 32]) + q0_vals = q_tail_desc.load([q_row_start, 256]) + q1_vals = q_tail_desc.load([q_row_start, 288]) + q0_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group.to(tl.int32), 8, 0, 0]) + q0_scales = q0_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) + q0_scales = q0_scales.reshape([BLOCK_H, 4]) + q1_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group.to(tl.int32), 9, 0, 0]) + q1_scales = q1_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) + q1_scales = q1_scales.reshape([BLOCK_H, 4]) + + group_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + group_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) + for page_group_off in tl.range(0, GROUP_PAGES): + page_rel = page_group * GROUP_PAGES + page_group_off + compact_page = page_table_start + page_rel + physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) + + scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) + full_k_vals = k_desc.load([physical_page.to(tl.int32), 0, 0]) + full_k_vals = tl.reshape(full_k_vals, (BLOCK_T, 256)) + full_k_scales = tl.ext.load_view_tko( + k_sf_full_view, [physical_page.to(tl.int32), 0, 0, 0, 0] + ) + full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) + scores = tl.dot_scaled( + full_q_vals, + full_q_scales, + "e2m1", + full_k_vals.T, + full_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + tail_k_vals = k_tail_desc.load([physical_page.to(tl.int32), 0, 256]) + tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) + tail_k_scales = tl.ext.load_view_tko( + k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0] + ) + tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) + scores = tl.dot_scaled( + q0_vals, + q0_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q1_vals, + q1_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + scores = scores * qk_scale + page_max = tl.max(scores, axis=1) + exp_scores = tl.math.exp2((scores - page_max[:, None]) * 1.4426950408889634) + page_sum = tl.sum(exp_scores, axis=1) + if GROUP_REDUCE_STATS: + next_group_max = tl.maximum(group_max, page_max) + group_sum = group_sum * tl.math.exp2( + (group_max - next_group_max) * 1.4426950408889634 + ) + page_sum * tl.math.exp2((page_max - next_group_max) * 1.4426950408889634) + group_max = next_group_max + + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + p_page = query_idx * MAX_PAGES + page_rel + out_offsets = query_idx * page_stats_s0 + page_rel * page_stats_s1 + offs_h + tl.store(page_max_ptr + out_offsets, page_max) + if not GROUP_REDUCE_STATS: + tl.store(page_sum_ptr + out_offsets, page_sum) + + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + tl.store(p_sf_ptr + sf_offsets, stored_scale) + p_desc.store( + [(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), + ) + if GROUP_REDUCE_STATS: + group_max_offsets = query_idx * page_stats_s0 + (page_group * 2) * page_stats_s1 + offs_h + group_sum_offsets = group_max_offsets + page_stats_s1 + tl.store(page_sum_ptr + group_max_offsets, group_max) + tl.store(page_sum_ptr + group_sum_offsets, group_sum) + + +@triton.jit +def _fp4_mla_attention_reduce_stats_kernel( + max_ptr, + denom_ptr, + page_max_ptr, + page_sum_ptr, + stats_s0, + page_stats_s0, + page_stats_s1, + NUM_HEADS: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + GROUP_REDUCE_STATS: tl.constexpr = False, + GROUP_PAGES: tl.constexpr = 1, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + + max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + if GROUP_REDUCE_STATS: + for group_rel in tl.range(0, MAX_PAGES // GROUP_PAGES): + group_max = tl.load( + page_sum_ptr + + gen_idx * page_stats_s0 + + (group_rel * 2) * page_stats_s1 + + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.maximum(max_score, group_max) + else: + for page_rel in tl.range(0, MAX_PAGES): + page_max = tl.load( + page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.maximum(max_score, page_max) + + denom = tl.zeros((BLOCK_H,), dtype=tl.float32) + if GROUP_REDUCE_STATS: + for group_rel in tl.range(0, MAX_PAGES // GROUP_PAGES): + group_max = tl.load( + page_sum_ptr + + gen_idx * page_stats_s0 + + (group_rel * 2) * page_stats_s1 + + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + group_sum = tl.load( + page_sum_ptr + + gen_idx * page_stats_s0 + + (group_rel * 2 + 1) * page_stats_s1 + + safe_offs_h, + mask=mask_h, + other=0.0, + ) + denom += tl.where( + group_sum > 0.0, + group_sum * tl.math.exp2((group_max - max_score) * 1.4426950408889634), + 0.0, + ) + else: + for page_rel in tl.range(0, MAX_PAGES): + page_max = tl.load( + page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + page_sum = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=0.0, + ) + denom += tl.where( + page_sum > 0.0, + page_sum * tl.math.exp2((page_max - max_score) * 1.4426950408889634), + 0.0, + ) + + tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) + tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) + + +@triton.jit +def _fp4_mla_attention_prob_scale_kernel( + p_sf_ptr, + max_ptr, + denom_ptr, + page_max_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + stats_s0, + page_stats_s0, + page_stats_s1, + NUM_HEADS: tl.constexpr, + PAGE_SIZE: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_rel = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + if page_start >= kv_len: + return + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if not ASSUME_VALID_PAGES: + if (compact_page < 0) | (compact_page >= page_ids_len): + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + page_max = tl.load( + page_max_ptr + query_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + factor = tl.where( + denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0 + ) + + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + scale_cols = tl.arange(0, SF_PER_PAGE) + if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) + tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) + + +@triton.jit +def _fp4_mla_attention_prob_store_page_kernel( + probs_ptr, + max_ptr, + denom_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_rel, + page_ids_len, + num_pages, + probs_s0, + probs_s1, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + stats_s0, + q_num_rows, + sm_scale, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_start = page_rel * PAGE_SIZE + if (not ASSUME_FULL_PAGES) and page_start >= kv_len: + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, PAGE_SIZE) + if ASSUME_FULL_PAGES: + valid_t = tl.full([PAGE_SIZE], True, dtype=tl.int1) + else: + valid_t = page_start + offs_t < kv_len + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if (compact_page < 0) | (compact_page >= page_ids_len): + return + q_row_base = query_idx * NUM_HEADS + + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + PAGE_SIZE, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + denom_valid = denom > 0.0 + safe_denom = tl.where(denom_valid, denom, 1.0) + safe_max = tl.where(denom_valid, max_score, 0.0) + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + probs = tl.math.exp2((scores - safe_max[:, None]) * 1.4426950408889634) / safe_denom[:, None] + probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) + + prob_rows = query_idx * NUM_HEADS + offs_h + safe_prob_rows = tl.where(mask_h, prob_rows, query_idx * NUM_HEADS) + tl.store( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + offs_t[None, :] * probs_s1, + probs, + mask=mask_h[:, None], + ) + + +@triton.jit +def _fp4_mla_attention_prob_pack_page_kernel( + p_fp4_ptr, + p_sf_ptr, + probs_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_rel, + page_ids_len, + p_s0, + p_s1, + probs_s0, + probs_s1, + NUM_HEADS: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + token_group = tl.program_id(1) + head_block = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_start = page_rel * PAGE_SIZE + if (not ASSUME_FULL_PAGES) and page_start >= kv_len: + return + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if (compact_page < 0) | (compact_page >= page_ids_len): + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + else: + mask_h = offs_h < NUM_HEADS + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + token_base = token_group * FP4_BLOCK + even_t = token_base + byte_offsets * 2 + odd_t = even_t + 1 + valid_even = page_start + even_t < kv_len + valid_odd = page_start + odd_t < kv_len + + prob_rows = query_idx * NUM_HEADS + offs_h + safe_prob_rows = tl.where(mask_h, prob_rows, query_idx * NUM_HEADS) + even_probs = tl.load( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + even_t[None, :] * probs_s1, + mask=mask_h[:, None] & valid_even[None, :], + other=0.0, + ) + odd_probs = tl.load( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + odd_t[None, :] * probs_s1, + mask=mask_h[:, None] & valid_odd[None, :], + other=0.0, + ) + amax = tl.maximum(tl.max(tl.abs(even_probs), axis=1), tl.max(tl.abs(odd_probs), axis=1)) + local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) + + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = tl.where(mask_h, p_rows, p_page * NUM_HEADS) + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows, token_group, SF_PER_PAGE) + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h) + + even_quant = _fp4_e2m1_quantize(even_probs / local_scale[:, None]) + odd_quant = _fp4_e2m1_quantize(odd_probs / local_scale[:, None]) + packed = even_quant | (odd_quant << 4) + byte_cols = token_group * (FP4_BLOCK // 2) + byte_offsets + tl.store( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + byte_cols[None, :] * p_s1, + packed, + mask=mask_h[:, None], + ) + + +@triton.jit +def _fp4_mla_attention_prob_pack_page_fused_kernel( + p_fp4_ptr, + p_sf_ptr, + max_ptr, + denom_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_rel, + page_ids_len, + num_pages, + p_s0, + p_s1, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + stats_s0, + q_num_rows, + sm_scale, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + PAGE_REL_FROM_GRID: tl.constexpr = False, + ASSUME_FULL_HEADS: tl.constexpr = False, + ASSUME_FULL_PAGES: tl.constexpr = False, + ASSUME_VALID_PAGES: tl.constexpr = False, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + if PAGE_REL_FROM_GRID: + page_rel = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_start = page_rel * PAGE_SIZE + if (not ASSUME_FULL_PAGES) and page_start >= kv_len: + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, PAGE_SIZE) + if ASSUME_FULL_PAGES: + valid_t = tl.full([PAGE_SIZE], True, dtype=tl.int1) + else: + valid_t = page_start + offs_t < kv_len + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if (compact_page < 0) | (compact_page >= page_ids_len): + return + q_row_base = query_idx * NUM_HEADS + + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + PAGE_SIZE, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + denom_valid = denom > 0.0 + safe_denom = tl.where(denom_valid, denom, 1.0) + safe_max = tl.where(denom_valid, max_score, 0.0) + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + probs = tl.math.exp2((scores - safe_max[:, None]) * 1.4426950408889634) / safe_denom[:, None] + probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) + + grouped_probs = tl.reshape(probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(tl.abs(grouped_probs), axis=2) + local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) + scaled_probs = grouped_probs / tl.reshape(local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = tl.where(mask_h, p_rows, p_page * NUM_HEADS) + scale_cols = tl.arange(0, SF_PER_PAGE) + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h[:, None]) + + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + tl.store( + p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=mask_h[:, None, None], + ) + + +@triton.jit +def _fp4_mla_attention_pv_kernel( + out_ptr, + p_fp4_ptr, + p_sf_ptr, + kv_cache_ptr, + v_sf_ptr, + v_packed_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + out_s0, + out_s1, + out_s2, + out_num_rows, + p_s0, + p_s1, + p_num_rows, + kv_s0, + kv_s2, + kv_s4, + vsf_s0, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + USE_TMA_P_LOAD: tl.constexpr, + USE_TMA_V_LOAD: tl.constexpr, + USE_PREPACKED_V: tl.constexpr, + PV_LOOP_STAGES: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_FULL_V: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + dim_block = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + if ASSUME_FULL_V: + mask_v = tl.full([BLOCK_V], True, dtype=tl.int1) + safe_offs_v = offs_v + else: + mask_v = offs_v < V_HEAD_D + safe_offs_v = tl.where(mask_v, offs_v, 0) + packed_t = tl.arange(0, PAGE_SIZE // 2) + scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) + even_t = packed_t * 2 + odd_t = even_t + 1 + v_packed_offsets = safe_offs_v // 2 + v_use_high_nibble = (safe_offs_v & 1) != 0 + if ASSUME_FULL_V and BLOCK_V == 128: + v_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + v_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + if USE_TMA_P_LOAD: + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + if USE_TMA_V_LOAD and ASSUME_FULL_HEADS and ASSUME_FULL_V: + tl.assume(out_s1 % 8 == 0) + tl.assume(out_s2 == 1) + out_desc = tl.make_tensor_descriptor( + out_ptr, + shape=[out_num_rows, V_HEAD_D], + strides=[out_s1, out_s2], + block_shape=[BLOCK_H, BLOCK_V], + ) + if USE_TMA_V_LOAD: + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, PAGE_SIZE, BLOCK_V // 2], + ) + if USE_PREPACKED_V: + v_packed_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + + if ( + USE_TMA_P_LOAD + and USE_TMA_V_LOAD + and ASSUME_FULL_HEADS + and ASSUME_FULL_PAGES + and ASSUME_FULL_V + and ASSUME_VALID_PAGES + and NUM_HEADS == 128 + and V_HEAD_D == 512 + and PAGE_SIZE == 128 + and BLOCK_H == 128 + and BLOCK_V == 128 + and SF_PER_PAGE == 8 + ): + p_view = tl.ext.make_view( + base=p_fp4_ptr, + shapes=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + tile_shape=[BLOCK_H, PAGE_SIZE // 2], + tile_dim_map=[0, 1], + ) + p_sf_view = tl.ext.make_view( + base=p_sf_ptr, + shapes=[p_num_rows // 128, SF_PER_PAGE // 4, 2, 256], + strides=[128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + v_sf_view = tl.ext.make_view( + base=v_sf_ptr, + shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], + strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + if not USE_PREPACKED_V: + v_view = tl.ext.make_view( + base=kv_cache_ptr, + shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], + tile_dim_map=[0, 1, 2], + ) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): + compact_page = page_table_start + page_rel + physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) + p_page = query_idx * MAX_PAGES + page_rel + + p_vals = tl.ext.load_view_tko( + p_view, + [(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + ) + p_vals = p_vals.to(tl.uint8, bitcast=True) + p_scales = tl.ext.load_view_tko(p_sf_view, [p_page.to(tl.int32), 0, 0, 0]) + p_scales = p_scales.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) + p_scales = p_scales.reshape([BLOCK_H, SF_PER_PAGE]) + + if USE_PREPACKED_V: + v_row = (physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + else: + v_tile = tl.ext.load_view_tko( + v_view, + [ + physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ], + ) + v_tile = v_tile.to(tl.uint8, bitcast=True) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) + high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + v_scales = tl.ext.load_view_tko( + v_sf_view, + [ + physical_page.to(tl.int32), + dim_block, + 0, + 0, + 0, + ], + ) + v_scales = v_scales.reshape([1, 1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) + acc = tl.dot_scaled( + p_vals, + p_scales, + "e2m1", + v_vals.T, + v_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + out_vals = acc * out_scale + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + out_desc.store( + [ + (query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ], + out_vals, + ) + return + + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES or page_start < kv_len: + compact_page = page_table_start + page_rel + if ASSUME_VALID_PAGES: + safe_compact_page = compact_page + physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) + safe_physical_page = physical_page + else: + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + ) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + ) + if USE_TMA_P_LOAD: + p_vals = p_desc.load([(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) + else: + p_vals = tl.load( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, + mask=mask_h[:, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None], + other=0, + ) + p_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + p_scales = tl.load(p_sf_ptr + p_sf_offsets) + + if ASSUME_FULL_PAGES: + valid_even_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + valid_odd_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + else: + valid_even_t = page_start + even_t < kv_len + valid_odd_t = page_start + odd_t < kv_len + if USE_PREPACKED_V: + v_row = (safe_physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + elif USE_TMA_V_LOAD: + v_tile = v_desc.load( + [ + safe_physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ] + ) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + if not ASSUME_VALID_PAGES: + v_tile = tl.where(valid_physical_page, v_tile, 0) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + if not ASSUME_FULL_PAGES: + even_packed = tl.where(valid_even_t[None, :], even_packed, 0) + odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) + low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) + high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + else: + even_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + even_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_even_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), + other=0, + ) + odd_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + odd_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_odd_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), + other=0, + ) + even_low = even_packed & 0x0F + even_high = (even_packed >> 4) & 0x0F + odd_low = odd_packed & 0x0F + even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) + odd_nibble = tl.where(v_use_high_nibble[:, None], odd_packed >> 4, odd_low) + v_vals = even_nibble | (odd_nibble << 4) + v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) + acc = tl.dot_scaled( + p_vals, + p_scales, + "e2m1", + v_vals.T, + v_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + if ASSUME_FULL_HEADS and ASSUME_FULL_V: + out_vals = acc * out_scale + if USE_TMA_V_LOAD: + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + out_desc.store( + [ + (query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ], + out_vals, + ) + else: + tl.store( + out_ptr + query_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + ) + else: + tl.store( + out_ptr + + query_idx * out_s0 + + safe_offs_h[:, None] * out_s1 + + safe_offs_v[None, :] * out_s2, + acc * out_scale, + mask=mask_h[:, None] & mask_v[None, :], + ) + + +def fp4_mla_paged_attention_internal( + q_fp4: torch.Tensor, + q_sf: torch.Tensor, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + v_sf: torch.Tensor, + global_scale: torch.Tensor, + src_page_ids: torch.Tensor, + paged_kv_indptr_decode: torch.Tensor, + kv_lens: torch.Tensor, + output: Optional[torch.Tensor] = None, + *, + sm_scale: float, + num_heads: Optional[int] = None, + v_head_dim: Optional[int] = None, + page_size: Optional[int] = None, + q_residual_dim: int = 0, + p_global_scale: float = FP4_MLA_P_GLOBAL_SCALE, + block_h: int = 128, + block_k: Optional[int] = None, + block_v: int = 128, + output_dtype: torch.dtype = torch.bfloat16, + max_pages: Optional[int] = None, + page_pipeline_streams: Optional[int] = None, + kernel_occupancy: Optional[int] = None, + kernel_num_ctas: Optional[int] = None, + kernel_num_stages: Optional[int] = None, + kernel_num_warps: Optional[int] = None, + pv_loop_stages: int = 1, + parallel_page_stats: Optional[bool] = None, + fused_prob_pack: Optional[bool] = None, + use_tma_data_load: Optional[bool] = None, + fused_prob_pack_single_launch: Optional[bool] = None, + pack_prob_in_page_stats: Optional[bool] = None, + page_stats_group_size: Optional[int] = None, + assume_full_pages: Optional[bool] = None, + assume_valid_pages: Optional[bool] = None, + query_len_per_seq: int = 1, + prepack_v_for_pv: bool = False, + use_prepacked_v_for_pv: bool = False, + p_fp4_workspace: Optional[torch.Tensor] = None, + p_sf_workspace: Optional[torch.Tensor] = None, + v_packed_workspace: Optional[torch.Tensor] = None, + p_probs_workspace: Optional[torch.Tensor] = None, + max_scores_workspace: Optional[torch.Tensor] = None, + denom_workspace: Optional[torch.Tensor] = None, + page_max_workspace: Optional[torch.Tensor] = None, + page_sum_workspace: Optional[torch.Tensor] = None, + **kwargs, +) -> torch.Tensor: + del kwargs + if not hasattr(tl, "dot_scaled"): + raise NotImplementedError( + "fp4_mla_paged_attention requires a Triton build with tl.dot_scaled." + ) + if not q_fp4.is_cuda: + raise ValueError("q_fp4 must be a CUDA tensor.") + if q_fp4.dtype != torch.uint8 or kv_cache.dtype != torch.uint8: + raise TypeError("q_fp4 and kv_cache must be packed FP4 tensors with dtype torch.uint8.") + if global_scale.numel() < 1: + raise ValueError("global_scale must contain at least one element.") + if q_fp4.dim() == 3: + inferred_num_queries, inferred_num_heads, packed_q_dim = q_fp4.shape + if num_heads is not None and num_heads != inferred_num_heads: + raise ValueError( + f"num_heads={num_heads} does not match q_fp4.shape[1]={inferred_num_heads}." + ) + num_queries = inferred_num_queries + num_heads = inferred_num_heads + q_fp4_2d = q_fp4.reshape(num_queries * num_heads, packed_q_dim) + elif q_fp4.dim() == 2: + if num_heads is None: + raise ValueError("num_heads is required when q_fp4 is 2D.") + if q_fp4.shape[0] % num_heads != 0: + raise ValueError("q_fp4.shape[0] must be divisible by num_heads.") + num_queries = q_fp4.shape[0] // num_heads + packed_q_dim = q_fp4.shape[1] + q_fp4_2d = q_fp4 + else: + raise ValueError("q_fp4 must be 2D or 3D.") + if query_len_per_seq <= 0: + raise ValueError(f"query_len_per_seq must be positive, got {query_len_per_seq}.") + if num_queries % query_len_per_seq != 0: + raise ValueError( + "q_fp4 query rows must be divisible by query_len_per_seq, got " + f"{num_queries} rows and query_len_per_seq={query_len_per_seq}." + ) + num_gen_seqs = num_queries // query_len_per_seq + + num_pages, inferred_page_size, packed_k_dim, kv_s0, kv_s2, kv_s4 = _get_kv_cache_strides( + kv_cache + ) + if page_size is None: + page_size = inferred_page_size + if page_size != inferred_page_size: + raise ValueError( + f"page_size={page_size} does not match kv_cache page dimension {inferred_page_size}." + ) + if page_size % FP4_BLOCK_SIZE != 0: + raise ValueError(f"page_size must be divisible by {FP4_BLOCK_SIZE}.") + + q_head_dim = packed_q_dim * 2 + k_head_dim = packed_k_dim * 2 + if q_residual_dim < 0 or q_residual_dim % FP4_BLOCK_SIZE != 0: + raise ValueError(f"q_residual_dim must be a non-negative multiple of {FP4_BLOCK_SIZE}.") + if q_head_dim - q_residual_dim != k_head_dim: + raise ValueError( + f"q_head_dim - q_residual_dim must match K head dim: {q_head_dim} - {q_residual_dim} != {k_head_dim}." + ) + if q_head_dim % FP4_BLOCK_SIZE != 0 or k_head_dim % FP4_BLOCK_SIZE != 0: + raise ValueError(f"Q/K head dims must be divisible by {FP4_BLOCK_SIZE}.") + if v_head_dim is None: + v_head_dim = k_head_dim + if v_head_dim <= 0 or v_head_dim > k_head_dim: + raise ValueError(f"v_head_dim must be in (0, {k_head_dim}], got {v_head_dim}.") + q_sf_flat = q_sf.contiguous().view(-1) + if sf_cache.shape[0] < num_pages or v_sf.shape[0] < num_pages: + raise ValueError("sf_cache and v_sf must have a leading physical-page dimension.") + if q_sf_flat.numel() < _swizzled_scale_size(num_queries * num_heads, q_head_dim): + raise ValueError("q_sf is too small for the swizzled Q scale layout.") + if sf_cache.numel() < sf_cache.shape[0] * _swizzled_scale_size(page_size, k_head_dim): + raise ValueError("sf_cache is too small for the swizzled K scale layout.") + if v_sf.numel() < v_sf.shape[0] * _swizzled_scale_size(v_head_dim, page_size): + raise ValueError("v_sf is too small for the swizzled V scale layout.") + + if output is None: + output = torch.empty( + (num_queries, num_heads, v_head_dim), dtype=output_dtype, device=q_fp4.device + ) + elif output.shape != (num_queries, num_heads, v_head_dim): + raise ValueError( + f"output must have shape {(num_queries, num_heads, v_head_dim)}, got {tuple(output.shape)}." + ) + + if num_queries == 0: + return output + triton_backend = "nvt" + env_block_h = _env_int("TRTLLM_FP4_MLA_BLOCK_H") + env_block_k = _env_int("TRTLLM_FP4_MLA_BLOCK_K") + env_block_v = _env_int("TRTLLM_FP4_MLA_BLOCK_V") + env_pv_loop_stages = _env_int("TRTLLM_FP4_MLA_PV_LOOP_STAGES") + env_occupancy = _env_int("TRTLLM_FP4_MLA_OCCUPANCY") + env_num_warps = _env_int("TRTLLM_FP4_MLA_NUM_WARPS") + env_num_stages = _env_int("TRTLLM_FP4_MLA_NUM_STAGES") + env_group_pages = _env_int("TRTLLM_FP4_MLA_GROUP_PAGES") + env_group_reduce_stats = _env_int("TRTLLM_FP4_MLA_GROUP_REDUCE_STATS") + env_page_pipeline_streams = _env_int("TRTLLM_FP4_MLA_PAGE_PIPELINE_STREAMS") + if env_block_h is not None: + block_h = env_block_h + if block_k is None: + block_k = env_block_k or (512 if triton_backend == "nvt" else 256) + elif env_block_k is not None: + block_k = env_block_k + if env_block_v is not None: + block_v = env_block_v + if env_pv_loop_stages is not None: + pv_loop_stages = env_pv_loop_stages + full_block_end = (q_head_dim // block_k) * block_k + tail_k = q_head_dim - full_block_end + tail_block_k = 1 << (tail_k - 1).bit_length() if tail_k > 0 else block_k + if max_pages is None: + if paged_kv_indptr_decode.numel() >= num_gen_seqs + 1: + page_counts = ( + paged_kv_indptr_decode[1 : num_gen_seqs + 1] - paged_kv_indptr_decode[:num_gen_seqs] + ) + max_pages = int(page_counts.max().item()) if page_counts.numel() > 0 else 0 + else: + max_pages = _ceil_div(int(kv_lens[:num_gen_seqs].max().item()), page_size) + if max_pages <= 0: + output.zero_() + return output + + q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE + k_sf_per_token = k_head_dim // FP4_BLOCK_SIZE + sf_per_page = page_size // FP4_BLOCK_SIZE + block_t = page_size + num_head_blocks = triton.cdiv(num_heads, block_h) + assume_full_heads = num_heads % block_h == 0 + assume_full_v = v_head_dim % block_v == 0 + if assume_full_pages is None: + assume_full_pages = False + assume_full_pages = bool(assume_full_pages) and query_len_per_seq == 1 + if assume_valid_pages is None: + assume_valid_pages = False + assume_valid_pages = bool(assume_valid_pages) + if ( + not assume_valid_pages + and assume_full_pages + and src_page_ids.numel() == num_gen_seqs * max_pages + ): + # Exactly-sized full-page decode tables can skip compact-page and + # physical-page validity masks. Keeping this inside the kernel wrapper + # lets call sites stay conservative. + assume_valid_pages = True + total_p_rows = max(num_queries * max_pages * num_heads, 1) + if page_pipeline_streams is None and env_page_pipeline_streams is not None: + page_pipeline_streams = env_page_pipeline_streams + if page_pipeline_streams is None: + if triton_backend == "nvt" and max_pages >= 8 and num_queries >= 128: + page_pipeline_streams = 2 + else: + page_pipeline_streams = 1 + page_pipeline_streams = max(1, min(int(page_pipeline_streams), max_pages)) + launch_meta = {} + if kernel_occupancy is None: + if env_occupancy is not None: + kernel_occupancy = env_occupancy + elif triton_backend == "nvt": + kernel_occupancy = 2 + if kernel_occupancy is not None: + launch_meta["occupancy"] = int(kernel_occupancy) + if kernel_num_ctas is not None: + launch_meta["num_ctas"] = int(kernel_num_ctas) + if kernel_num_stages is None and env_num_stages is not None: + kernel_num_stages = env_num_stages + if kernel_num_stages is not None: + launch_meta["num_stages"] = int(kernel_num_stages) + if kernel_num_warps is None and env_num_warps is not None: + kernel_num_warps = env_num_warps + if kernel_num_warps is not None: + launch_meta["num_warps"] = int(kernel_num_warps) + if fused_prob_pack is None: + fused_prob_pack = triton_backend == "nvt" + if fused_prob_pack_single_launch is None: + fused_prob_pack_single_launch = triton_backend == "nvt" and max_pages >= 8 + if use_tma_data_load is None: + use_tma_data_load = triton_backend == "nvt" + use_tma_data_load = bool(use_tma_data_load and hasattr(tl, "make_tensor_descriptor")) + if use_tma_data_load: + # Device-side descriptors may need Triton's allocator for descriptor scratch storage. + def alloc_fn(size: int, alignment: int, stream: Optional[int]): + return torch.empty(size, device=q_fp4.device, dtype=torch.int8) + + triton.set_allocator(alloc_fn) + + p_fp4 = _workspace_tensor( + p_fp4_workspace, + (total_p_rows, page_size // 2), + dtype=torch.uint8, + device=q_fp4.device, + name="p_fp4", + ) + p_sf = _workspace_tensor( + p_sf_workspace, + (_swizzled_scale_size(total_p_rows, page_size),), + dtype=q_sf.dtype, + device=q_fp4.device, + name="p_sf", + ) + num_dim_blocks = triton.cdiv(v_head_dim, block_v) + v_repack_block_v = block_v + num_repack_dim_blocks = num_dim_blocks + auto_prepack_v_for_pv = ( + triton_backend == "nvt" + and use_tma_data_load + and assume_full_heads + and assume_full_pages + and assume_valid_pages + and v_head_dim == 512 + and page_size == 128 + and block_h in (64, 128) + and block_v == 128 + and sf_per_page == 8 + ) + if not prepack_v_for_pv and not use_prepacked_v_for_pv: + env_prepack_v = os.environ.get("TRTLLM_FP4_MLA_PREPACK_V") + prepack_v_for_pv = auto_prepack_v_for_pv if env_prepack_v is None else env_prepack_v == "1" + wants_prepacked_v_for_pv = bool(prepack_v_for_pv) or bool(use_prepacked_v_for_pv) + if use_prepacked_v_for_pv and v_packed_workspace is None: + raise ValueError("use_prepacked_v_for_pv requires v_packed_workspace to be provided.") + can_use_prepacked_v_for_pv = ( + wants_prepacked_v_for_pv + and use_tma_data_load + and assume_full_heads + and assume_full_pages + and assume_valid_pages + and v_head_dim == 512 + and page_size == 128 + and block_h in (64, 128) + and block_v == 128 + and sf_per_page == 8 + ) + if wants_prepacked_v_for_pv and not can_use_prepacked_v_for_pv: + raise ValueError( + "prepacked V PV path requires TMA, full heads/pages, valid pages, " + "v_head_dim=512, page_size=128, block_h in (64, 128), " + "block_v=128, and sf_per_page=8." + ) + if can_use_prepacked_v_for_pv: + v_packed = _workspace_tensor( + v_packed_workspace, + (num_pages * num_dim_blocks * block_v, page_size // 2), + dtype=torch.uint8, + device=q_fp4.device, + name="v_packed", + ) + else: + v_packed = kv_cache + if fused_prob_pack: + p_probs = None + else: + p_probs_shape = (max(num_queries * num_heads, 1), page_size) + if page_pipeline_streams > 1: + p_probs = _workspace_tensor( + p_probs_workspace, + (page_pipeline_streams, *p_probs_shape), + dtype=torch.float32, + device=q_fp4.device, + name="p_probs", + ) + else: + p_probs = _workspace_tensor( + p_probs_workspace, + p_probs_shape, + dtype=torch.float32, + device=q_fp4.device, + name="p_probs", + ) + max_scores = _workspace_tensor( + max_scores_workspace, + (num_queries, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="max_scores", + ) + denom = _workspace_tensor( + denom_workspace, + (num_queries, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="denom", + ) + + v_repack_stream = None + if bool(prepack_v_for_pv) and can_use_prepacked_v_for_pv: + current_stream = torch.cuda.current_stream(q_fp4.device) + v_repack_stream = torch.cuda.Stream(device=q_fp4.device) + v_repack_stream.wait_stream(current_stream) + with torch.cuda.stream(v_repack_stream): + _fp4_mla_attention_v_repack_kernel[(num_pages, num_repack_dim_blocks)]( + v_packed, + kv_cache, + num_pages, + kv_s0, + kv_s2, + kv_s4, + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + BLOCK_V=v_repack_block_v, + **launch_meta, + ) + + if parallel_page_stats is None: + parallel_page_stats = triton_backend == "nvt" and max_pages >= 8 + if pack_prob_in_page_stats is None: + pack_prob_in_page_stats = parallel_page_stats and fused_prob_pack + pack_prob_in_page_stats = bool( + pack_prob_in_page_stats and parallel_page_stats and fused_prob_pack + ) + page_stats_group_sizes = (2, 4, 8, 16, 32, 64, 128, 256, 512, 1024) + if page_stats_group_size is None and env_group_pages is not None: + page_stats_group_size = env_group_pages + if page_stats_group_size is None: + page_stats_group_size = 1 + else: + page_stats_group_size = int(page_stats_group_size) + can_group_page_stats = ( + page_stats_group_size in page_stats_group_sizes + and triton_backend == "nvt" + and parallel_page_stats + and pack_prob_in_page_stats + and use_tma_data_load + and assume_full_heads + and assume_full_pages + and assume_valid_pages + and query_len_per_seq == 1 + and num_heads == 128 + and q_head_dim == 640 + and k_head_dim == 576 + and q_residual_dim == 64 + and page_size == 128 + and block_h == 128 + and block_k == 512 + and full_block_end == 512 + and tail_block_k == 128 + and sf_per_page == 8 + and max_pages % page_stats_group_size == 0 + ) + page_stats_group_size = page_stats_group_size if can_group_page_stats else 1 + group_reduce_stats = ( + env_group_reduce_stats != 0 + if env_group_reduce_stats is not None + else triton_backend == "nvt" + ) and page_stats_group_size > 1 + if parallel_page_stats: + page_stats_shape = (num_queries, max_pages, num_heads) + page_max = _workspace_tensor( + page_max_workspace, + page_stats_shape, + dtype=torch.float32, + device=q_fp4.device, + name="page_max", + ) + page_sum = _workspace_tensor( + page_sum_workspace, + page_stats_shape, + dtype=torch.float32, + device=q_fp4.device, + name="page_sum", + ) + if page_stats_group_size > 1: + _fp4_mla_attention_page_stats_grouped_kernel[ + (num_queries, num_head_blocks, max_pages // page_stats_group_size) + ]( + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + page_max.stride(0), + page_max.stride(1), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + BLOCK_T=block_t, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + PACK_PROBS=pack_prob_in_page_stats, + GROUP_REDUCE_STATS=group_reduce_stats, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + GROUP_PAGES=page_stats_group_size, + **launch_meta, + ) + else: + _fp4_mla_attention_page_stats_kernel[(num_queries, num_head_blocks, max_pages)]( + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + page_max.stride(0), + page_max.stride(1), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + BLOCK_T=block_t, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + PACK_PROBS=pack_prob_in_page_stats, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + _fp4_mla_attention_reduce_stats_kernel[(num_queries, num_head_blocks)]( + max_scores, + denom, + page_max, + page_sum, + max_scores.stride(0), + page_max.stride(0), + page_max.stride(1), + NUM_HEADS=num_heads, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + GROUP_REDUCE_STATS=group_reduce_stats, + GROUP_PAGES=page_stats_group_size, + **launch_meta, + ) + if pack_prob_in_page_stats: + _fp4_mla_attention_prob_scale_kernel[(num_queries, num_head_blocks, max_pages)]( + p_sf, + max_scores, + denom, + page_max, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + max_scores.stride(0), + page_max.stride(0), + page_max.stride(1), + NUM_HEADS=num_heads, + PAGE_SIZE=page_size, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + else: + _fp4_mla_attention_stats_kernel[(num_queries, num_head_blocks)]( + max_scores, + denom, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + max_scores.stride(0), + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + BLOCK_T=block_t, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + + def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None): + if fused_prob_pack: + _fp4_mla_attention_prob_pack_page_fused_kernel[(num_queries, num_head_blocks)]( + p_fp4, + p_sf, + max_scores, + denom, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + page_rel, + src_page_ids.shape[0], + num_pages, + p_fp4.stride(0), + p_fp4.stride(1), + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + max_scores.stride(0), + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + return + assert p_probs_slot is not None + _fp4_mla_attention_prob_store_page_kernel[(num_queries, num_head_blocks)]( + p_probs_slot, + max_scores, + denom, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + page_rel, + src_page_ids.shape[0], + num_pages, + p_probs_slot.stride(0), + p_probs_slot.stride(1), + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + max_scores.stride(0), + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + QUERY_LEN_PER_SEQ=query_len_per_seq, + BLOCK_H=block_h, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + _fp4_mla_attention_prob_pack_page_kernel[(num_queries, sf_per_page, num_head_blocks)]( + p_fp4, + p_sf, + p_probs_slot, + paged_kv_indptr_decode, + kv_lens, + page_rel, + src_page_ids.shape[0], + p_fp4.stride(0), + p_fp4.stride(1), + p_probs_slot.stride(0), + p_probs_slot.stride(1), + NUM_HEADS=num_heads, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + **launch_meta, + ) + + if pack_prob_in_page_stats: + pass + elif fused_prob_pack and fused_prob_pack_single_launch: + _fp4_mla_attention_prob_pack_page_fused_kernel[(num_queries, num_head_blocks, max_pages)]( + p_fp4, + p_sf, + max_scores, + denom, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + 0, + src_page_ids.shape[0], + num_pages, + p_fp4.stride(0), + p_fp4.stride(1), + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + max_scores.stride(0), + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + PAGE_REL_FROM_GRID=True, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + elif page_pipeline_streams == 1: + for page_rel in range(max_pages): + _launch_prob_page(page_rel, p_probs) + else: + current_stream = torch.cuda.current_stream(q_fp4.device) + streams = [torch.cuda.Stream(device=q_fp4.device) for _ in range(page_pipeline_streams)] + for stream in streams: + stream.wait_stream(current_stream) + for page_rel in range(max_pages): + stream_idx = page_rel % page_pipeline_streams + with torch.cuda.stream(streams[stream_idx]): + if fused_prob_pack: + _launch_prob_page(page_rel) + else: + assert p_probs is not None + _launch_prob_page(page_rel, p_probs[stream_idx]) + for stream in streams: + current_stream.wait_stream(stream) + + if v_repack_stream is not None: + torch.cuda.current_stream(q_fp4.device).wait_stream(v_repack_stream) + + _fp4_mla_attention_pv_kernel[ + ( + num_queries, + num_head_blocks, + num_dim_blocks, + ) + ]( + output, + p_fp4, + p_sf, + kv_cache, + v_sf, + v_packed, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + output.stride(0), + output.stride(1), + output.stride(2), + output.shape[0] * output.shape[1], + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + kv_s0, + kv_s2, + kv_s4, + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + BLOCK_V=block_v, + USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, + USE_TMA_V_LOAD=use_tma_data_load and v_head_dim % block_v == 0, + USE_PREPACKED_V=can_use_prepacked_v_for_pv, + PV_LOOP_STAGES=int(pv_loop_stages), + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_FULL_V=assume_full_v, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + return output + + +fp4_mla_paged_attention = fp4_mla_paged_attention_internal diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py index 18fa22e130c4..c1f853553d38 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py @@ -81,9 +81,13 @@ def _hp_kv_store_context_kernel( def _hp_kv_store_gen_kernel( pool_ptr, latent_cache_ptr, - seq_slots_ptr, # int32 [num_gen], seq_slot for each gen seq - kv_lens_ptr, # int32 [num_gen], total KV length after this decode step - gen_tok_start, # int, offset in latent_cache where gen tokens begin + seq_slots_ptr, # int32 [num_seqs], seq_slot for each sequence + batch_indices_ptr, # int32 [metadata_num_tokens], sequence index per token + positions_ptr, # int32 [metadata_num_tokens], absolute KV position per token + gen_tok_start, # int, offset in latent_cache where generation tokens begin + token_offset, # int, offset in metadata token arrays where generation tokens begin + num_tokens, # int, number of generation tokens in latent_cache + metadata_num_tokens, num_seq_slots, num_layers, layer_idx, @@ -95,25 +99,35 @@ def _hp_kv_store_gen_kernel( BLOCK_D: tl.constexpr, HP_BLOCK: tl.constexpr, ): - """Store the current generation token into the HP KV pool. + """Store generation tokens into the HP KV pool. - Grid: (num_gen_seqs,). + Grid: (num_generation_tokens,). Each program stores one token into the circular buffer position - (kv_len - 1) % HP_BLOCK, overwriting the oldest entry. + ``position % HP_BLOCK``, overwriting the oldest entry. The token metadata + is read from the same batch_indices/positions arrays used by the KV scatter + path, so this supports linear MTP where each sequence contributes multiple + generation tokens. """ - gen_idx = tl.program_id(0) + gen_token_idx = tl.program_id(0) if (layer_idx < 0) | (layer_idx >= num_layers): return + if gen_token_idx >= num_tokens: + return - seq_slot = tl.load(seq_slots_ptr + gen_idx).to(tl.int64) - if (seq_slot < 0) | (seq_slot >= num_seq_slots): + metadata_token_idx = token_offset + gen_token_idx + if metadata_token_idx >= metadata_num_tokens: return - kv_len = tl.load(kv_lens_ptr + gen_idx) - if kv_len <= 0: + batch_idx = tl.load(batch_indices_ptr + metadata_token_idx).to(tl.int64) + position = tl.load(positions_ptr + metadata_token_idx).to(tl.int64) + if (batch_idx < 0) | (position < 0): return - buf_pos = (kv_len - 1) % HP_BLOCK - token_idx = tl.cast(gen_tok_start, tl.int64) + gen_idx.to(tl.int64) + seq_slot = tl.load(seq_slots_ptr + batch_idx).to(tl.int64) + if (seq_slot < 0) | (seq_slot >= num_seq_slots): + return + buf_pos = position % HP_BLOCK + + token_idx = tl.cast(gen_tok_start, tl.int64) + gen_token_idx.to(tl.int64) offs_d = tl.arange(0, BLOCK_D) mask_d = offs_d < D safe_offs_d = tl.where(mask_d, offs_d, 0) @@ -127,6 +141,141 @@ def _hp_kv_store_gen_kernel( tl.store(pool_ptr + dst_base + safe_offs_d, src, mask=mask_d) +@triton.jit +def _hp_kv_restore_rejected_from_pool_kernel( + pool_ptr, + snapshot_pool_ptr, + batch_indices_ptr, + positions_ptr, + seq_slots_ptr, + kv_lens_ptr, + prompt_lens_ptr, + accepted_tokens_ptr, + token_offset, + num_tokens, + metadata_num_tokens, + num_seqs, + num_accepted_tokens, + num_seq_slots, + num_layers, + local_layer, + pool_stride_seq, + pool_stride_layer, + snapshot_stride_seq, + snapshot_stride_layer, + D: tl.constexpr, + POOL_HEAD_D: tl.constexpr, + BLOCK_D: tl.constexpr, + HP_BLOCK: tl.constexpr, +): + token_idx = tl.program_id(0) + metadata_token_idx = token_offset + token_idx + token_valid = ( + (token_idx < num_tokens) + & (metadata_token_idx >= 0) + & (metadata_token_idx < metadata_num_tokens) + & (local_layer >= 0) + & (local_layer < num_layers) + ) + + batch_idx = tl.load(batch_indices_ptr + metadata_token_idx, mask=token_valid, other=-1).to( + tl.int64 + ) + position = tl.load(positions_ptr + metadata_token_idx, mask=token_valid, other=-1).to(tl.int64) + batch_valid = ( + token_valid + & (batch_idx >= 0) + & (batch_idx < num_seqs) + & (batch_idx < num_accepted_tokens) + & (position >= 0) + ) + + seq_slot = tl.load(seq_slots_ptr + batch_idx, mask=batch_valid, other=-1).to(tl.int64) + kv_len = tl.load(kv_lens_ptr + batch_idx, mask=batch_valid, other=0).to(tl.int64) + prompt_len = tl.load(prompt_lens_ptr + batch_idx, mask=batch_valid, other=0).to(tl.int64) + accepted = tl.load(accepted_tokens_ptr + batch_idx, mask=batch_valid, other=0).to(tl.int64) + + hp_slot = position % HP_BLOCK + first_new_position = kv_len - prompt_len + should_restore = ( + batch_valid + & (seq_slot >= 0) + & (seq_slot < num_seq_slots) + & (hp_slot >= 0) + & (hp_slot < HP_BLOCK) + & (position >= first_new_position + accepted) + ) + + offs_d = tl.arange(0, BLOCK_D) + mask_d = offs_d < D + safe_offs_d = tl.where(mask_d, offs_d, 0) + src_base = ( + seq_slot * snapshot_stride_seq + local_layer * snapshot_stride_layer + hp_slot * POOL_HEAD_D + ) + dst_base = seq_slot * pool_stride_seq + local_layer * pool_stride_layer + hp_slot * POOL_HEAD_D + values = tl.load( + snapshot_pool_ptr + src_base + safe_offs_d, mask=should_restore & mask_d, other=0.0 + ) + tl.store(pool_ptr + dst_base + safe_offs_d, values, mask=should_restore & mask_d) + + +@triton.jit +def _hp_kv_restore_rejected_from_values_kernel( + pool_ptr, + values_ptr, + batch_indices_ptr, + positions_ptr, + seq_slots_ptr, + hp_slots_ptr, + first_new_positions_ptr, + accepted_tokens_ptr, + num_tokens, + num_accepted_tokens, + num_seq_slots, + num_layers, + local_layer, + pool_stride_seq, + pool_stride_layer, + values_stride_token, + values_stride_dim, + D: tl.constexpr, + POOL_HEAD_D: tl.constexpr, + BLOCK_D: tl.constexpr, + HP_BLOCK: tl.constexpr, +): + token_idx = tl.program_id(0) + token_valid = (token_idx < num_tokens) & (local_layer >= 0) & (local_layer < num_layers) + + batch_idx = tl.load(batch_indices_ptr + token_idx, mask=token_valid, other=-1).to(tl.int64) + position = tl.load(positions_ptr + token_idx, mask=token_valid, other=-1).to(tl.int64) + seq_slot = tl.load(seq_slots_ptr + token_idx, mask=token_valid, other=-1).to(tl.int64) + hp_slot = tl.load(hp_slots_ptr + token_idx, mask=token_valid, other=-1).to(tl.int64) + first_new_position = tl.load(first_new_positions_ptr + token_idx, mask=token_valid, other=0).to( + tl.int64 + ) + batch_valid = token_valid & (batch_idx >= 0) & (batch_idx < num_accepted_tokens) + accepted = tl.load(accepted_tokens_ptr + batch_idx, mask=batch_valid, other=0).to(tl.int64) + should_restore = ( + batch_valid + & (seq_slot >= 0) + & (seq_slot < num_seq_slots) + & (hp_slot >= 0) + & (hp_slot < HP_BLOCK) + & (position >= first_new_position + accepted) + ) + + offs_d = tl.arange(0, BLOCK_D) + mask_d = offs_d < D + safe_offs_d = tl.where(mask_d, offs_d, 0) + values = tl.load( + values_ptr + token_idx * values_stride_token + safe_offs_d * values_stride_dim, + mask=should_restore & mask_d, + other=0.0, + ) + dst_base = seq_slot * pool_stride_seq + local_layer * pool_stride_layer + hp_slot * POOL_HEAD_D + tl.store(pool_ptr + dst_base + safe_offs_d, values, mask=should_restore & mask_d) + + @triton.jit def _fp4_mla_swizzled_sf_offset( row_idx, @@ -600,6 +749,178 @@ def _fp4_mla_v_scale_store_hp_tail_kernel( ) +@triton.jit +def _fp4_mla_v_scale_store_generation_tiles_kernel( + kv_cache_ptr, + sf_cache_ptr, + v_sf_ptr, + hp_pool_ptr, + latent_cache_ptr, + global_scale_ptr, + seq_slots_ptr, + kv_lens_ptr, + prompt_lens_ptr, + page_ids_ptr, + paged_kv_indptr_ptr, + page_ids_len, + indptr_len, + num_seq_slots, + num_pages, + num_layers, + local_layer, + page_size, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + pool_s0, + pool_s1, + lc_s0, + lc_s1, + vsf_s0, + vsf_s1, + HEAD_D: tl.constexpr, + V_HEAD_D: tl.constexpr, + HP_BLOCK: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, +): + seq_idx = tl.program_id(0) + tile_idx = tl.program_id(1) + dim_block = tl.program_id(2) + if (local_layer < 0) | (local_layer >= num_layers): + return + if seq_idx + 1 >= indptr_len: + return + + kv_len = tl.load(kv_lens_ptr + seq_idx) + gen_len = tl.load(prompt_lens_ptr + seq_idx) + if gen_len <= 0: + return + first_new_pos = kv_len - gen_len + first_tile_pos = (first_new_pos // HP_BLOCK) * HP_BLOCK + block_base_pos = first_tile_pos + tile_idx * HP_BLOCK + if block_base_pos >= kv_len: + return + if block_base_pos + HP_BLOCK <= first_new_pos: + return + + page_idx = block_base_pos // page_size + page_pos = block_base_pos - page_idx * page_size + page_start = tl.load(paged_kv_indptr_ptr + seq_idx).to(tl.int64) + page_end = tl.load(paged_kv_indptr_ptr + seq_idx + 1).to(tl.int64) + physical_page_offset = page_start + page_idx + if ( + (page_pos < 0) + | (page_pos >= page_size) + | (physical_page_offset < page_start) + | (physical_page_offset >= page_end) + | (physical_page_offset < 0) + | (physical_page_offset >= page_ids_len) + ): + return + physical_page = tl.load(page_ids_ptr + physical_page_offset).to(tl.int64) + if (physical_page < 0) | (physical_page >= num_pages): + return + seq_slot = tl.load(seq_slots_ptr + seq_idx).to(tl.int64) + if (seq_slot < 0) | (seq_slot >= num_seq_slots): + return + + byte_offsets = tl.arange(0, HP_BLOCK // 2) + token_offsets = tl.arange(0, HP_BLOCK) + even_d = dim_block * FP4_BLOCK + byte_offsets * 2 + odd_d = even_d + 1 + all_d = dim_block * FP4_BLOCK + tl.arange(0, FP4_BLOCK) + mask_even_d = even_d < HEAD_D + mask_odd_d = odd_d < HEAD_D + mask_all_d = all_d < HEAD_D + safe_even_d = tl.where(mask_even_d, even_d, 0) + safe_odd_d = tl.where(mask_odd_d, odd_d, 0) + safe_all_d = tl.where(mask_all_d, all_d, 0) + + abs_positions = block_base_pos + token_offsets + valid_tokens = abs_positions < kv_len + from_latent = abs_positions >= first_new_pos + hp_slots = abs_positions % HP_BLOCK + new_token_offsets = abs_positions - first_new_pos + # Linear MTP uses a uniform generation length, so each sequence occupies a + # contiguous gen_len slice in latent_cache. + latent_tokens = seq_idx * gen_len + new_token_offsets + safe_latent_tokens = tl.where(valid_tokens & from_latent, latent_tokens, 0).to(tl.int64) + + hp_even = tl.load( + hp_pool_ptr + + seq_slot * pool_s0 + + local_layer * pool_s1 + + hp_slots[:, None] * HEAD_D + + safe_even_d[None, :], + mask=valid_tokens[:, None] & (~from_latent)[:, None] & mask_even_d[None, :], + other=0.0, + ).to(tl.float32) + hp_odd = tl.load( + hp_pool_ptr + + seq_slot * pool_s0 + + local_layer * pool_s1 + + hp_slots[:, None] * HEAD_D + + safe_odd_d[None, :], + mask=valid_tokens[:, None] & (~from_latent)[:, None] & mask_odd_d[None, :], + other=0.0, + ).to(tl.float32) + latent_even = tl.load( + latent_cache_ptr + safe_latent_tokens[:, None] * lc_s0 + safe_even_d[None, :] * lc_s1, + mask=valid_tokens[:, None] & from_latent[:, None] & mask_even_d[None, :], + other=0.0, + ).to(tl.float32) + latent_odd = tl.load( + latent_cache_ptr + safe_latent_tokens[:, None] * lc_s0 + safe_odd_d[None, :] * lc_s1, + mask=valid_tokens[:, None] & from_latent[:, None] & mask_odd_d[None, :], + other=0.0, + ).to(tl.float32) + even_values = hp_even + latent_even + odd_values = hp_odd + latent_odd + + amax_per_token = tl.maximum( + tl.max(tl.abs(even_values), axis=1), + tl.max(tl.abs(odd_values), axis=1), + ) + tile_amax = tl.max(amax_per_token, axis=0) + global_scale = tl.load(global_scale_ptr) + shared_tile = dim_block * FP4_BLOCK < V_HEAD_D + tile_scale = tl.where(tile_amax > 0.0, tile_amax / 6.0, 1.0) + token_scale = tl.where(amax_per_token > 0.0, amax_per_token / 6.0, 1.0) + local_scale = tl.where(shared_tile, tile_scale, token_scale) + stored_scale = local_scale * global_scale + v_stored_scale = tile_scale * global_scale + + low = _fp4_e2m1_quantize(even_values / local_scale[:, None]) + high = _fp4_e2m1_quantize(odd_values / local_scale[:, None]) + packed = low | (high << 4) + + packed_cols = dim_block * (FP4_BLOCK // 2) + byte_offsets + page_positions = page_pos + token_offsets + kv_base = physical_page * kv_s0 + tl.store( + kv_cache_ptr + kv_base + page_positions[:, None] * kv_s2 + packed_cols[None, :] * kv_s4, + packed, + mask=valid_tokens[:, None] & mask_even_d[None, :], + ) + + k_sf_offsets = _fp4_mla_swizzled_sf_offset(page_positions, dim_block, SF_PER_TOKEN) + tl.store(sf_cache_ptr + physical_page * sf_s0 + k_sf_offsets, stored_scale, mask=valid_tokens) + + token_scale_col = page_pos // HP_BLOCK + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_all_d, token_scale_col, SF_PER_PAGE) + v_sf_base = tl.cast(local_layer, tl.int64) * tl.cast( + vsf_s0, tl.int64 + ) + physical_page * tl.cast(vsf_s1, tl.int64) + tl.store( + v_sf_ptr + v_sf_base + sf_offsets.to(tl.int64), + v_stored_scale, + mask=mask_all_d & (all_d < V_HEAD_D), + ) + + @triton.jit def _fp4_mla_load_values( kv_cache_ptr, @@ -860,21 +1181,25 @@ def _fp4_mla_attention_stats_kernel( FP4_BLOCK: tl.constexpr, Q_SF_PER_TOKEN: tl.constexpr, K_SF_PER_TOKEN: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, MAX_PAGES: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_T: tl.constexpr, BLOCK_K: tl.constexpr, ): - gen_idx = tl.program_id(0) + query_idx = tl.program_id(0) head_block = tl.program_id(1) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) mask_h = offs_h < NUM_HEADS safe_offs_h = tl.where(mask_h, offs_h, 0) offs_t = tl.arange(0, BLOCK_T) - q_row_base = gen_idx * NUM_HEADS - kv_len = tl.load(kv_lens_ptr + gen_idx) - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + q_row_base = query_idx * NUM_HEADS + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) denom = tl.zeros((BLOCK_H,), dtype=tl.float32) @@ -922,8 +1247,8 @@ def _fp4_mla_attention_stats_kernel( ) max_score = new_max - tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) - tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) + tl.store(max_ptr + query_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) + tl.store(denom_ptr + query_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) @triton.jit @@ -960,13 +1285,17 @@ def _fp4_mla_attention_prob_store_page_kernel( FP4_BLOCK: tl.constexpr, Q_SF_PER_TOKEN: tl.constexpr, K_SF_PER_TOKEN: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_K: tl.constexpr, ): - gen_idx = tl.program_id(0) + query_idx = tl.program_id(0) head_block = tl.program_id(1) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ - kv_len = tl.load(kv_lens_ptr + gen_idx) + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) page_start = page_rel * PAGE_SIZE if page_start >= kv_len: return @@ -976,11 +1305,11 @@ def _fp4_mla_attention_prob_store_page_kernel( safe_offs_h = tl.where(mask_h, offs_h, 0) offs_t = tl.arange(0, PAGE_SIZE) valid_t = page_start + offs_t < kv_len - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) compact_page = page_table_start + page_rel if (compact_page < 0) | (compact_page >= page_ids_len): return - q_row_base = gen_idx * NUM_HEADS + q_row_base = query_idx * NUM_HEADS scores = _fp4_mla_qk_scores_tile( q_fp4_ptr, @@ -1011,8 +1340,8 @@ def _fp4_mla_attention_prob_store_page_kernel( BLOCK_K, NUM_HEADS, ) - max_score = tl.load(max_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - denom = tl.load(denom_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) + max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) global_scale = tl.load(global_scale_ptr) qk_scale = sm_scale / (global_scale * global_scale) denom_valid = denom > 0.0 @@ -1022,8 +1351,8 @@ def _fp4_mla_attention_prob_store_page_kernel( probs = tl.exp(scores - safe_max[:, None]) / safe_denom[:, None] probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) - prob_rows = gen_idx * NUM_HEADS + offs_h - safe_prob_rows = tl.where(mask_h, prob_rows, gen_idx * NUM_HEADS) + prob_rows = query_idx * NUM_HEADS + offs_h + safe_prob_rows = tl.where(mask_h, prob_rows, query_idx * NUM_HEADS) tl.store( probs_ptr + safe_prob_rows[:, None] * probs_s0 + offs_t[None, :] * probs_s1, probs, @@ -1049,18 +1378,23 @@ def _fp4_mla_attention_prob_pack_page_kernel( FP4_BLOCK: tl.constexpr, SF_PER_PAGE: tl.constexpr, P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, BLOCK_H: tl.constexpr, ): - gen_idx = tl.program_id(0) + query_idx = tl.program_id(0) token_group = tl.program_id(1) head_block = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ - kv_len = tl.load(kv_lens_ptr + gen_idx) + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) page_start = page_rel * PAGE_SIZE if page_start >= kv_len: return - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) compact_page = page_table_start + page_rel if (compact_page < 0) | (compact_page >= page_ids_len): return @@ -1074,8 +1408,8 @@ def _fp4_mla_attention_prob_pack_page_kernel( valid_even = page_start + even_t < kv_len valid_odd = page_start + odd_t < kv_len - prob_rows = gen_idx * NUM_HEADS + offs_h - safe_prob_rows = tl.where(mask_h, prob_rows, gen_idx * NUM_HEADS) + prob_rows = query_idx * NUM_HEADS + offs_h + safe_prob_rows = tl.where(mask_h, prob_rows, query_idx * NUM_HEADS) even_probs = tl.load( probs_ptr + safe_prob_rows[:, None] * probs_s0 + even_t[None, :] * probs_s1, mask=mask_h[:, None] & valid_even[None, :], @@ -1090,8 +1424,9 @@ def _fp4_mla_attention_prob_pack_page_kernel( local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) - p_rows = compact_page * NUM_HEADS + offs_h - safe_p_rows = tl.where(mask_h, p_rows, compact_page * NUM_HEADS) + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = tl.where(mask_h, p_rows, p_page * NUM_HEADS) sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows, token_group, SF_PER_PAGE) tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h) @@ -1133,14 +1468,17 @@ def _fp4_mla_attention_pv_kernel( PAGE_SIZE: tl.constexpr, FP4_BLOCK: tl.constexpr, SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, MAX_PAGES: tl.constexpr, P_GLOBAL_SCALE: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_V: tl.constexpr, ): - gen_idx = tl.program_id(0) + query_idx = tl.program_id(0) head_block = tl.program_id(1) dim_block = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) @@ -1155,8 +1493,9 @@ def _fp4_mla_attention_pv_kernel( v_packed_offsets = safe_offs_v // 2 v_use_high_nibble = (safe_offs_v & 1) != 0 - kv_len = tl.load(kv_lens_ptr + gen_idx) - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) global_scale = tl.load(global_scale_ptr) acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) for page_rel in tl.range(0, MAX_PAGES): @@ -1173,8 +1512,9 @@ def _fp4_mla_attention_pv_kernel( ) safe_physical_page = tl.where(valid_physical_page, physical_page, 0) - p_rows = safe_compact_page * NUM_HEADS + offs_h - safe_p_rows = tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = tl.where(mask_h, p_rows, p_page * NUM_HEADS) p_vals = tl.load( p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, mask=valid_compact_page & mask_h[:, None], @@ -1226,7 +1566,10 @@ def _fp4_mla_attention_pv_kernel( ) tl.store( - out_ptr + gen_idx * out_s0 + safe_offs_h[:, None] * out_s1 + safe_offs_v[None, :] * out_s2, + out_ptr + + query_idx * out_s0 + + safe_offs_h[:, None] * out_s1 + + safe_offs_v[None, :] * out_s2, acc / (global_scale * P_GLOBAL_SCALE), mask=mask_h[:, None] & mask_v[None, :], ) diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py deleted file mode 100644 index 7b32e3e53b47..000000000000 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_kv.py +++ /dev/null @@ -1,1530 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -"""Shared MLA FP4 KV-cache helpers. - -The high-precision (HP) BF16 KV pool is a standalone circular buffer used -alongside the paged FP4 KV pool when MLA models run with NVFP4 KV cache. -Each sequence slot stores the ``HP_BLOCK_SIZE`` most-recent latent vectors at -BF16, so attention backends can consult BF16 values for the tail tokens that -do not yet fill a complete FP4 quant block of 16 elements along the sequence -dimension. - -Used by both ``TrtllmAttention`` (via an internal C++ attention op that reads -both pools) and ``FlashInferAttention`` (via either explicit Python-side -dequant into a BF16 workspace before calling FlashInfer MLA wrappers, or an -env-gated Triton attention path that reads packed FP4 Q, K, and V directly). -""" - -import os -from typing import Any, Literal, Optional - -import torch -import triton - -from .fp4_mla_kernels import ( - _fp4_mla_attention_prob_pack_page_kernel, - _fp4_mla_attention_prob_store_page_kernel, - _fp4_mla_attention_pv_kernel, - _fp4_mla_attention_stats_kernel, - _fp4_mla_dequant_kernel, - _fp4_mla_overlay_hp_tail_kernel, - _fp4_mla_scatter_kernel, - _fp4_mla_v_scale_store_context_tokens_kernel, - _fp4_mla_v_scale_store_hp_tail_kernel, - _hp_kv_store_context_kernel, - _hp_kv_store_gen_kernel, -) - -HP_BLOCK_SIZE: int = 16 -FP4_BLOCK_SIZE: int = 16 -FP4_MLA_TOKENS_PER_BLOCK: int = 128 -FP4_MLA_SCALE_ROW_GROUP: int = 128 -FP4_MLA_SCALE_COL_GROUP: int = 4 -FP4_MLA_KV_GLOBAL_SCALE: float = 448.0 * 6.0 / 448.0 * 6.0 -FP4_MLA_P_GLOBAL_SCALE: float = 448.0 * 6.0 -FP4_MLA_Q_RESIDUAL_DIM: int = 64 -FLASHINFER_FP4_MLA_ATTENTION_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION" -FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION_BACKEND" -FLASHINFER_FP4_MLA_DEBUG_ENV = "TRTLLM_FLASHINFER_FP4_MLA_DEBUG" -_FP4_MLA_CUTE_DSL_BACKEND = "cute_dsl" -_HPUpdatePhase = Literal["all", "context", "generation"] - - -# Environment and debug helpers - - -def _env_enabled(name: str) -> bool: - return os.getenv(name, "0").lower() in ( - "1", - "true", - "yes", - "on", - ) - - -def is_flashinfer_fp4_mla_attention_enabled() -> bool: - """Return whether FlashInfer MLA should allocate no-dequant FP4 attention buffers.""" - return _env_enabled(FLASHINFER_FP4_MLA_ATTENTION_ENV) - - -def _fp4_mla_attention_backend() -> str: - return os.getenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, "triton").lower() - - -def _fp4_mla_debug_enabled() -> bool: - return _env_enabled(FLASHINFER_FP4_MLA_DEBUG_ENV) - - -def _fp4_mla_debug(message: str) -> None: - if _fp4_mla_debug_enabled(): - print(f"[fp4_mla_debug] {message}", flush=True) - - -def _tensor_layout(tensor: Optional[torch.Tensor]) -> str: - if tensor is None: - return "None" - return ( - f"shape={list(tensor.shape)} stride={list(tensor.stride())} " - f"dtype={tensor.dtype} device={tensor.device}" - ) - - -def _debug_tensor_range(name: str, tensor: Optional[torch.Tensor]) -> None: - if not _fp4_mla_debug_enabled(): - return - if tensor is None: - _fp4_mla_debug(f"{name}: None") - return - flat = tensor.detach().reshape(-1) - if flat.numel() == 0: - _fp4_mla_debug(f"{name}: empty {_tensor_layout(tensor)}") - return - try: - first = flat[: min(8, flat.numel())].cpu().tolist() - _fp4_mla_debug( - f"{name}: {_tensor_layout(tensor)} n={flat.numel()} " - f"min={flat.min().item()} max={flat.max().item()} first={first}" - ) - except RuntimeError as exc: - _fp4_mla_debug(f"{name}: failed to read range: {exc}") - - -def _debug_sync(label: str) -> None: - if not _fp4_mla_debug_enabled(): - return - if torch.cuda.is_current_stream_capturing(): - _fp4_mla_debug(f"{label}: skip sync during CUDA graph capture") - return - _fp4_mla_debug(f"{label}: synchronize") - torch.cuda.synchronize() - _fp4_mla_debug(f"{label}: sync complete") - - -def _ceil_div(lhs: int, rhs: int) -> int: - return (lhs + rhs - 1) // rhs - - -# FP4 MLA scale-layout helpers - - -def get_fp4_mla_v_scale_pool_size(v_head_dim: int, page_size: int) -> int: - """Return elements per page for the swizzled FP4 MLA V-scale pool. - - The PV matmul treats V as a RHS matrix shaped ``[v_head_dim, kv_tokens]``. - NVFP4 block scales therefore group along the token/K axis, not along the - latent dimension as the K-view cache does. The physical layout matches the - Triton block-scaled matmul scale layout: - ``[ceil(v_head_dim / 128), ceil(page_size / 16 / 4), 32, 16]``. - """ - token_scale_cols = _ceil_div(page_size, FP4_BLOCK_SIZE) - row_groups = _ceil_div(v_head_dim, FP4_MLA_SCALE_ROW_GROUP) - col_groups = _ceil_div(token_scale_cols, FP4_MLA_SCALE_COL_GROUP) - return row_groups * col_groups * 32 * 16 - - -def _get_fp4_mla_swizzled_scale_size(rows: int, cols: int) -> int: - scale_cols = _ceil_div(cols, FP4_BLOCK_SIZE) - row_groups = _ceil_div(rows, FP4_MLA_SCALE_ROW_GROUP) - col_groups = _ceil_div(scale_cols, FP4_MLA_SCALE_COL_GROUP) - return row_groups * col_groups * 32 * 16 - - -def _use_fp4_mla_swizzled_sf() -> bool: - return is_flashinfer_fp4_mla_attention_enabled() - - -def _get_fp4_mla_context_start_positions(metadata: Any, num_contexts: int) -> torch.Tensor: - kv_cache_params = getattr(metadata, "kv_cache_params", None) - cached_token_lens = getattr(kv_cache_params, "num_cached_tokens_per_seq", None) - if cached_token_lens is not None: - return torch.as_tensor(cached_token_lens[:num_contexts], dtype=torch.int64, device="cpu") - - return ( - ( - metadata.kv_lens_cuda_runtime[:num_contexts] - - metadata.prompt_lens_cuda_runtime[:num_contexts] - ) - .detach() - .cpu() - ) - - -def _validate_fp4_mla_context_start_alignment(metadata: Any, num_contexts: int) -> None: - context_start_positions = _get_fp4_mla_context_start_positions(metadata, num_contexts) - bad_start = (context_start_positions < 0) | ((context_start_positions % HP_BLOCK_SIZE) != 0) - if bool(torch.any(bad_start).item()): - starts = context_start_positions.detach().cpu().tolist() - raise ValueError( - "FP4 MLA shared-tile context update requires every context " - f"start position to be {HP_BLOCK_SIZE}-token aligned, got " - f"start positions {starts}." - ) - - -def get_fp4_mla_v_scale_pool_shape( - num_layers: int, - num_pages: int, - v_head_dim: int, - page_size: int, -) -> tuple[int, int, int, int, int, int]: - """Return the logical swizzled V-scale view shape. - - The leading dimensions are ``[layer, physical_page]``. The remaining - dimensions are the preshuffled ``[N // 128, K // 16 // 4, 32, 16]`` shape - consumed by Triton block-scaled matmul for the V/PV RHS operand. - """ - token_scale_cols = _ceil_div(page_size, FP4_BLOCK_SIZE) - return ( - num_layers, - num_pages, - _ceil_div(v_head_dim, FP4_MLA_SCALE_ROW_GROUP), - _ceil_div(token_scale_cols, FP4_MLA_SCALE_COL_GROUP), - 32, - 16, - ) - - -def get_fp4_mla_v_scale_pool_view( - metadata: Any, - *, - v_head_dim: int, -) -> torch.Tensor: - """View the auxiliary MLA V-scale pool in Triton's block-scaled layout.""" - pool = getattr(metadata, "fp4_mla_v_scale_pool", None) - if pool is None: - raise RuntimeError("FP4 MLA V scale pool is not allocated.") - - elems_per_page = get_fp4_mla_v_scale_pool_size(v_head_dim, metadata.page_size) - if pool.shape[-1] < elems_per_page: - raise RuntimeError( - f"FP4 MLA V scale pool page stride is too small: got " - f"{pool.shape[-1]}, need {elems_per_page}." - ) - - token_scale_cols = _ceil_div(metadata.page_size, FP4_BLOCK_SIZE) - col_groups = _ceil_div(token_scale_cols, FP4_MLA_SCALE_COL_GROUP) - shape = get_fp4_mla_v_scale_pool_shape( - pool.shape[0], pool.shape[1], v_head_dim, metadata.page_size - ) - strides = ( - pool.stride(0), - pool.stride(1), - col_groups * 32 * 16, - 32 * 16, - 16, - 1, - ) - return torch.as_strided(pool, size=shape, stride=strides) - - -# Python launch helpers - - -def _get_fp4_mla_global_scale(metadata: Any, device: torch.device) -> torch.Tensor: - global_scale = getattr(metadata, "_fp4_mla_global_scale", None) - if global_scale is None: - global_scale = torch.ones((1,), dtype=torch.float32, device=device) - return global_scale - - -def _get_fp4_mla_cache_tensors(metadata: Any, layer_idx: int) -> tuple[torch.Tensor, torch.Tensor]: - kv_cache = metadata.kv_cache_manager.get_buffers(layer_idx).view(torch.uint8) - sf_cache = metadata.kv_cache_manager.get_block_scale_buffers(layer_idx) - if sf_cache is None: - raise RuntimeError("NVFP4 KV cache scale pool is not available.") - return kv_cache, sf_cache - - -def _scatter_fp4_mla_kv_cache_2d_context( - metadata: Any, - latent_cache: torch.Tensor, - kv_cache: torch.Tensor, - sf_cache: torch.Tensor, - v_sf: torch.Tensor, - global_scale: torch.Tensor, - *, - token_offset: int, - local_layer: int, - v_head_dim: int, - head_dim: int, - num_tokens: int, - num_dim_blocks: int, - sf_per_token: int, - sf_per_page: int, -) -> None: - num_contexts = metadata.num_contexts - if num_contexts > 0: - prompt_lens_cpu = metadata.prompt_lens_cpu_runtime[:num_contexts] - ctx_token_count = int(prompt_lens_cpu.sum().item()) - if num_tokens != ctx_token_count: - raise RuntimeError( - f"FP4 MLA 2D context scatter needs {ctx_token_count} context tokens, got " - f"{num_tokens}." - ) - _validate_fp4_mla_context_start_alignment(metadata, num_contexts) - - _fp4_mla_v_scale_store_context_tokens_kernel[ - ( - num_tokens, - num_dim_blocks, - ) - ]( - kv_cache, - sf_cache, - v_sf, - latent_cache, - global_scale, - metadata.batch_indices, - metadata.positions, - metadata.paged_kv_indices, - metadata.paged_kv_indptr, - metadata.paged_kv_indices.shape[0], - metadata.paged_kv_indptr.shape[0], - metadata.batch_indices.shape[0], - v_sf.shape[1], - v_sf.shape[0], - token_offset, - num_tokens, - local_layer, - metadata.page_size, - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - sf_cache.stride(0), - latent_cache.stride(0), - latent_cache.stride(1), - v_sf.stride(0), - v_sf.stride(1), - HEAD_D=head_dim, - V_HEAD_D=v_head_dim, - HP_BLOCK=HP_BLOCK_SIZE, - FP4_BLOCK=FP4_BLOCK_SIZE, - SF_PER_TOKEN=sf_per_token, - SF_PER_PAGE=sf_per_page, - ) - _debug_sync("scatter_fp4_mla_kv_cache_2d_context") - - -def _scatter_fp4_mla_kv_cache_2d_generation( - metadata: Any, - kv_cache: torch.Tensor, - sf_cache: torch.Tensor, - v_sf: torch.Tensor, - global_scale: torch.Tensor, - *, - local_layer: int, - v_head_dim: int, - head_dim: int, - num_tokens: int, - num_dim_blocks: int, - sf_per_token: int, - sf_per_page: int, -) -> None: - num_contexts = metadata.num_contexts - num_seqs = metadata.num_seqs - num_gen = num_seqs - num_contexts - if num_gen <= 0: - return - if num_tokens != num_gen: - raise RuntimeError( - f"FP4 MLA 2D generation scatter expected {num_gen} generation tokens, got {num_tokens}." - ) - - pool = getattr(metadata, "high_precision_kv_pool", None) - if pool is None: - raise RuntimeError("FP4 MLA 2D generation scatter requires the HP KV pool.") - hp_head_dim = pool.shape[-1] // HP_BLOCK_SIZE - if hp_head_dim < head_dim: - raise RuntimeError( - f"FP4 MLA 2D generation scatter needs at least {head_dim} HP channels, got " - f"{hp_head_dim}." - ) - - page_ids = metadata.paged_kv_indices[metadata.num_context_blocks :] - _fp4_mla_v_scale_store_hp_tail_kernel[(num_gen, num_dim_blocks)]( - kv_cache, - sf_cache, - v_sf, - pool, - global_scale, - metadata.seq_slots[num_contexts:num_seqs], - metadata.kv_lens_cuda_runtime[num_contexts:num_seqs], - page_ids, - metadata.paged_kv_indptr_decode, - page_ids.shape[0], - metadata.paged_kv_indptr_decode.shape[0], - v_sf.shape[1], - pool.shape[0], - v_sf.shape[0], - local_layer, - metadata.page_size, - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - sf_cache.stride(0), - pool.stride(0), - pool.stride(1), - v_sf.stride(0), - v_sf.stride(1), - HEAD_D=hp_head_dim, - V_HEAD_D=v_head_dim, - HP_BLOCK=HP_BLOCK_SIZE, - FP4_BLOCK=FP4_BLOCK_SIZE, - SF_PER_TOKEN=sf_per_token, - SF_PER_PAGE=sf_per_page, - ) - _debug_sync("scatter_fp4_mla_kv_cache_2d_generation") - - -def _scatter_fp4_mla_kv_cache_1d( - metadata: Any, - latent_cache: torch.Tensor, - kv_cache: torch.Tensor, - sf_cache: torch.Tensor, - global_scale: torch.Tensor, - *, - layer_idx: int, - token_offset: int, - num_tokens: int, - head_dim: int, - sf_per_token: int, - use_swizzled_sf: bool, -) -> None: - q_fp4, q_sf = torch.ops.trtllm.fp4_quantize( - latent_cache, global_scale, FP4_BLOCK_SIZE, False, False - ) - q_sf = q_sf.view(num_tokens, head_dim // FP4_BLOCK_SIZE) - - packed_dim = head_dim // 2 - block_packed_dim = triton.next_power_of_2(packed_dim) - block_sf = triton.next_power_of_2(sf_per_token) - - _fp4_mla_debug( - "scatter launch: " - f"num_tokens={num_tokens} token_offset={token_offset} " - f"page_size={metadata.page_size} layer_idx={layer_idx} " - f"head_dim={head_dim} packed_dim={packed_dim} " - f"sf_per_token={sf_per_token} use_swizzled_sf={use_swizzled_sf}" - ) - _fp4_mla_debug(f"scatter latent_cache: {_tensor_layout(latent_cache)}") - _fp4_mla_debug(f"scatter kv_cache: {_tensor_layout(kv_cache)}") - _fp4_mla_debug(f"scatter sf_cache: {_tensor_layout(sf_cache)}") - _debug_tensor_range( - "scatter batch_indices", - metadata.batch_indices[token_offset : token_offset + num_tokens], - ) - _debug_tensor_range( - "scatter positions", - metadata.positions[token_offset : token_offset + num_tokens], - ) - _debug_tensor_range("scatter paged_kv_indices", metadata.paged_kv_indices) - _debug_tensor_range("scatter paged_kv_indptr", metadata.paged_kv_indptr) - - _fp4_mla_scatter_kernel[(num_tokens,)]( - kv_cache, - sf_cache, - q_fp4, - q_sf, - metadata.batch_indices, - metadata.positions, - metadata.paged_kv_indices, - metadata.paged_kv_indptr, - metadata.paged_kv_indices.shape[0], - metadata.paged_kv_indptr.shape[0], - kv_cache.shape[0], - token_offset, - metadata.page_size, - kv_cache.stride(0), - kv_cache.stride(1), - kv_cache.stride(2), - kv_cache.stride(3), - kv_cache.stride(4), - sf_cache.stride(0), - sf_cache.stride(1), - sf_cache.stride(2), - sf_cache.stride(3), - sf_cache.stride(4), - q_fp4.stride(0), - q_fp4.stride(1), - q_sf.stride(0), - q_sf.stride(1), - PACKED_D=packed_dim, - SF_PER_TOKEN=sf_per_token, - BLOCK_PACKED_D=block_packed_dim, - BLOCK_SF=block_sf, - USE_SWIZZLED_SF=use_swizzled_sf, - ) - _debug_sync("scatter_fp4_mla_kv_cache") - - -# Public cache update and decode entry points - - -def scatter_fp4_mla_kv_cache( - metadata: Any, - latent_cache: Optional[torch.Tensor], - layer_idx: int, - *, - token_offset: int, - phase: Optional[_HPUpdatePhase] = None, - local_layer: Optional[int] = None, - v_head_dim: Optional[int] = None, -) -> None: - """Quantize MLA latent tokens and scatter them into the paged FP4 cache. - - Contract: this helper scatters exactly ``latent_cache.shape[0]`` tokens, - reading index metadata at ``batch_indices[token_offset : token_offset + N]`` - and ``positions[token_offset : token_offset + N]``. Callers must pass a - latent_cache pre-sliced to the current phase (context or generation) so - that ``shape[0]`` matches the number of index entries they intend to - consume. ``MLA.forward_impl`` (tensorrt_llm/_torch/modules/attention.py) - slices ``latent_cache[:num_ctx_tokens]`` for context and - ``latent_cache[num_ctx_tokens:]`` for generation before dispatching. - - When the no-dequant FP4 MLA attention path is enabled, callers should pass - ``phase``, ``local_layer``, and ``v_head_dim``. Context scatter then writes - the final FP4 tile representation directly: dimensions below - ``v_head_dim`` use one shared 16-token by 16-dim scale written into both - K's token-major scale layout and V's dim-major scale layout. Tail K-only - dimensions use K's per-token 1D scales. Generation scatter rewrites the - active 16-token tile from the HP pool, so the caller must update the HP pool - before invoking this helper. - """ - if latent_cache is None or latent_cache.numel() == 0: - return - - latent_cache = latent_cache.reshape(latent_cache.shape[0], -1).contiguous() - num_tokens = latent_cache.shape[0] - head_dim = latent_cache.shape[-1] - if head_dim % FP4_BLOCK_SIZE != 0: - raise ValueError( - f"FP4 MLA KV head_dim must be divisible by {FP4_BLOCK_SIZE}, got {head_dim}." - ) - indices_len = metadata.batch_indices.shape[0] - positions_len = metadata.positions.shape[0] - if token_offset + num_tokens > indices_len or token_offset + num_tokens > positions_len: - raise RuntimeError( - f"FP4 MLA scatter would read batch_indices[{token_offset}:" - f"{token_offset + num_tokens}] / positions[{token_offset}:" - f"{token_offset + num_tokens}], but only {indices_len} / " - f"{positions_len} entries are available. This indicates " - "latent_cache was not pre-sliced to the current phase's token " - "range (see MLA.forward_impl)." - ) - - use_swizzled_sf = _use_fp4_mla_swizzled_sf() - if use_swizzled_sf: - _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) - - global_scale = _get_fp4_mla_global_scale(metadata, latent_cache.device) - kv_cache, sf_cache = _get_fp4_mla_cache_tensors(metadata, layer_idx) - sf_per_token = head_dim // FP4_BLOCK_SIZE - - use_2d_scatter = ( - use_swizzled_sf - and phase in ("context", "generation") - and getattr(metadata, "fp4_mla_v_scale_pool", None) is not None - ) - if use_2d_scatter: - assert phase is not None - if local_layer is None or v_head_dim is None: - raise ValueError("Real FP4 MLA scatter requires local_layer and v_head_dim.") - if metadata.page_size % HP_BLOCK_SIZE != 0: - raise ValueError( - f"FP4 MLA scatter requires page_size divisible by " - f"{HP_BLOCK_SIZE}, got {metadata.page_size}." - ) - if v_head_dim > head_dim: - raise ValueError(f"FP4 MLA v_head_dim={v_head_dim} cannot exceed head_dim={head_dim}.") - if v_head_dim % FP4_BLOCK_SIZE != 0: - raise ValueError( - f"FP4 MLA v_head_dim must be divisible by {FP4_BLOCK_SIZE}, got {v_head_dim}." - ) - - sf_cache = sf_cache.view(torch.float8_e4m3fn) - v_sf = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=v_head_dim) - num_dim_blocks = triton.cdiv(head_dim, FP4_BLOCK_SIZE) - sf_per_page = metadata.page_size // HP_BLOCK_SIZE - _fp4_mla_debug( - "scatter 2d launch: " - f"phase={phase} num_tokens={num_tokens} " - f"token_offset={token_offset} layer_idx={layer_idx} " - f"local_layer={local_layer} head_dim={head_dim} " - f"v_head_dim={v_head_dim} num_dim_blocks={num_dim_blocks}" - ) - _fp4_mla_debug(f"scatter 2d kv_cache: {_tensor_layout(kv_cache)}") - _fp4_mla_debug(f"scatter 2d sf_cache: {_tensor_layout(sf_cache)}") - _fp4_mla_debug(f"scatter 2d v_sf: {_tensor_layout(v_sf)}") - - if phase == "context": - _scatter_fp4_mla_kv_cache_2d_context( - metadata, - latent_cache, - kv_cache, - sf_cache, - v_sf, - global_scale, - token_offset=token_offset, - local_layer=local_layer, - v_head_dim=v_head_dim, - head_dim=head_dim, - num_tokens=num_tokens, - num_dim_blocks=num_dim_blocks, - sf_per_token=sf_per_token, - sf_per_page=sf_per_page, - ) - else: - _scatter_fp4_mla_kv_cache_2d_generation( - metadata, - kv_cache, - sf_cache, - v_sf, - global_scale, - local_layer=local_layer, - v_head_dim=v_head_dim, - head_dim=head_dim, - num_tokens=num_tokens, - num_dim_blocks=num_dim_blocks, - sf_per_token=sf_per_token, - sf_per_page=sf_per_page, - ) - return - - _scatter_fp4_mla_kv_cache_1d( - metadata, - latent_cache, - kv_cache, - sf_cache, - global_scale, - layer_idx=layer_idx, - token_offset=token_offset, - num_tokens=num_tokens, - head_dim=head_dim, - sf_per_token=sf_per_token, - use_swizzled_sf=use_swizzled_sf, - ) - - -def _ensure_decode_workspace( - metadata: Any, - head_dim: int, - dtype: torch.dtype, -) -> torch.Tensor: - num_blocks = _get_decode_workspace_num_blocks(metadata) - workspace = getattr(metadata, "_fp4_mla_decode_cache_buf", None) - needs_alloc = ( - workspace is None - or workspace.shape[0] < max(num_blocks, 1) - or workspace.shape[1] != metadata.page_size - or workspace.shape[2] != head_dim - or workspace.dtype != dtype - ) - if needs_alloc: - if torch.cuda.is_current_stream_capturing(): - raise ValueError( - "Cannot allocate FlashInfer FP4 MLA decode workspace while " - "capturing a CUDA graph. Run a warmup prepare/forward first." - ) - workspace = torch.empty( - (max(num_blocks, 1), metadata.page_size, head_dim), - dtype=dtype, - device=metadata.paged_kv_indices.device, - ) - metadata._fp4_mla_decode_cache_buf = workspace - return workspace[:num_blocks] - - -def _get_decode_workspace_num_blocks(metadata: Any) -> int: - if metadata.is_cuda_graph: - max_blocks_per_seq = ( - metadata.kv_cache_manager.max_seq_len + metadata.page_size - 1 - ) // metadata.page_size - max_graph_blocks = metadata.max_num_requests * max_blocks_per_seq - return min( - metadata.kv_cache_manager.blocks_in_primary_pool, - max_graph_blocks, - ) - return metadata.num_generation_blocks - - -def _get_decode_src_page_ids(metadata: Any, num_blocks: int) -> torch.Tensor: - page_ids = ( - metadata._paged_kv_indices - if metadata.is_cuda_graph and hasattr(metadata, "_paged_kv_indices") - else metadata.paged_kv_indices - ) - src_page_ids = page_ids[metadata.num_context_blocks : metadata.num_context_blocks + num_blocks] - if src_page_ids.numel() != num_blocks: - raise RuntimeError( - f"FP4 MLA dequant needs {num_blocks} decode page ids from " - f"paged_kv_indices[{metadata.num_context_blocks}:" - f"{metadata.num_context_blocks + num_blocks}], got " - f"{src_page_ids.numel()}." - ) - return src_page_ids - - -def _validate_fp4_mla_cache_shape(page_size: int, head_dim: int) -> None: - if page_size != FP4_MLA_TOKENS_PER_BLOCK: - raise ValueError( - f"FP4 MLA KV cache requires tokens_per_block={FP4_MLA_TOKENS_PER_BLOCK} " - f"for swizzled block scales, got {page_size}." - ) - - sf_per_token = head_dim // FP4_BLOCK_SIZE - if head_dim % FP4_BLOCK_SIZE != 0 or sf_per_token % 4 != 0: - raise ValueError( - f"FP4 MLA KV head_dim must produce a scale column count divisible by 4; " - f"got head_dim={head_dim}, scale_columns={sf_per_token}." - ) - - -def _validate_fp4_mla_attention_q_shape(head_dim: int, q_residual_dim: int) -> None: - if q_residual_dim % FP4_BLOCK_SIZE != 0: - raise ValueError( - f"FP4 MLA Q residual_dim must be divisible by {FP4_BLOCK_SIZE}, got {q_residual_dim}." - ) - if q_residual_dim <= 0 or q_residual_dim > head_dim: - raise ValueError( - f"FP4 MLA Q residual_dim must be in (0, head_dim], got " - f"residual_dim={q_residual_dim}, head_dim={head_dim}." - ) - - q_head_dim = head_dim + q_residual_dim - q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE - if q_head_dim % FP4_BLOCK_SIZE != 0 or q_sf_per_token % FP4_MLA_SCALE_COL_GROUP != 0: - raise ValueError( - f"FP4 MLA residual Q must produce a scale column count divisible " - f"by {FP4_MLA_SCALE_COL_GROUP}; got q_head_dim={q_head_dim}, " - f"scale_columns={q_sf_per_token}." - ) - - -def get_fp4_mla_decode_cache( - metadata: Any, - layer_idx: int, - local_layer: int, - *, - head_dim: int, - dtype: torch.dtype, -) -> torch.Tensor: - """Build a compact dequantized MLA cache for FlashInfer decode.""" - # Must match scatter_fp4_mla_kv_cache: the env var picks the SF layout, - # not the page size. When the dequant fallback is the read path - # (env disabled), scatter wrote linear scales and we must read linear. - use_swizzled_sf = _use_fp4_mla_swizzled_sf() - if use_swizzled_sf: - _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) - combined = _ensure_decode_workspace(metadata, head_dim, dtype) - num_blocks = combined.shape[0] - if num_blocks == 0: - return combined - - kv_cache, sf_cache = _get_fp4_mla_cache_tensors(metadata, layer_idx) - sf_cache = sf_cache.view(torch.float8_e4m3fn) - global_scale = _get_fp4_mla_global_scale(metadata, combined.device) - src_page_ids = _get_decode_src_page_ids(metadata, num_blocks) - block_d = triton.next_power_of_2(head_dim) - - _fp4_mla_dequant_kernel[(num_blocks, metadata.page_size)]( - combined, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - src_page_ids.shape[0], - kv_cache.shape[0], - kv_cache.stride(0), - kv_cache.stride(1), - kv_cache.stride(2), - kv_cache.stride(3), - kv_cache.stride(4), - sf_cache.stride(0), - sf_cache.stride(1), - sf_cache.stride(2), - sf_cache.stride(3), - sf_cache.stride(4), - combined.stride(0), - combined.stride(1), - combined.stride(2), - D=head_dim, - FP4_BLOCK=FP4_BLOCK_SIZE, - BLOCK_D=block_d, - USE_SWIZZLED_SF=use_swizzled_sf, - ) - - num_gen = metadata.num_seqs - metadata.num_contexts - if num_gen > 0 and metadata.high_precision_kv_pool is not None: - pool = metadata.high_precision_kv_pool - _fp4_mla_overlay_hp_tail_kernel[(num_gen, HP_BLOCK_SIZE)]( - combined, - pool, - metadata.seq_slots[metadata.num_contexts : metadata.num_seqs], - metadata.kv_lens_cuda_runtime[metadata.num_contexts : metadata.num_seqs], - metadata.paged_kv_indptr_decode, - pool.shape[0], - pool.shape[1], - combined.shape[0], - local_layer, - metadata.page_size, - combined.stride(0), - combined.stride(1), - combined.stride(2), - pool.stride(0), - pool.stride(1), - D=head_dim, - BLOCK_D=block_d, - HP_BLOCK=HP_BLOCK_SIZE, - ) - - return combined - - -def _ensure_workspace_tensor( - metadata: Any, - attr_name: str, - shape: tuple[int, ...], - *, - dtype: torch.dtype, - device: torch.device, -) -> torch.Tensor: - tensor = getattr(metadata, attr_name, None) - needs_alloc = ( - tensor is None - or tensor.dtype != dtype - or tensor.device != device - or len(tensor.shape) != len(shape) - or any(tensor.shape[idx] < dim for idx, dim in enumerate(shape)) - ) - if needs_alloc: - if torch.cuda.is_current_stream_capturing(): - raise ValueError( - f"Cannot allocate {attr_name} while capturing a CUDA graph. " - "Run a warmup prepare/forward first." - ) - tensor = torch.empty(shape, dtype=dtype, device=device) - setattr(metadata, attr_name, tensor) - - slices = tuple(slice(0, dim) for dim in shape) - return tensor[slices] - - -def _max_generation_pages(metadata: Any) -> int: - num_gen = metadata.num_seqs - metadata.num_contexts - if num_gen <= 0: - return 0 - num_blocks = getattr(metadata, "num_blocks", None) - if num_blocks is not None: - return max(num_blocks[metadata.num_contexts : metadata.num_seqs]) - return metadata.num_generation_blocks - - -def _host_int_list(value: Any, start: int, end: int) -> Optional[list[int]]: - if value is None: - return None - if isinstance(value, torch.Tensor): - if value.is_cuda: - return None - return [int(item) for item in value[start:end].tolist()] - try: - return [int(item) for item in value[start:end]] - except (TypeError, ValueError): - return None - - -def _infer_cutile_assume_full_pages(metadata: Any, max_pages: int, page_size: int) -> bool: - if getattr(metadata, "is_cuda_graph", False): - return False - - start = metadata.num_contexts - end = metadata.num_seqs - kv_cache_params = getattr(metadata, "kv_cache_params", None) - cached_token_lens = _host_int_list( - getattr(kv_cache_params, "num_cached_tokens_per_seq", None), - start, - end, - ) - seq_lens_kv = _host_int_list(getattr(metadata, "seq_lens_kv", None), start, end) - if cached_token_lens is not None and seq_lens_kv is not None: - if len(cached_token_lens) != len(seq_lens_kv): - return False - kv_lens = [ - cached_len + seq_len for cached_len, seq_len in zip(cached_token_lens, seq_lens_kv) - ] - elif kv_cache_params is None: - kv_lens = _host_int_list(getattr(metadata, "prompt_lens_cpu_runtime", None), start, end) - else: - return False - - return bool(kv_lens) and min(kv_lens) == max(kv_lens) == max_pages * page_size - - -def run_fp4_mla_attention_decode( - metadata: Any, - layer_idx: int, - local_layer: int, - q_nope: torch.Tensor, - q_pe: torch.Tensor, - output: torch.Tensor, - *, - sm_scale: float, - kv_lora_rank: int, - qk_rope_head_dim: int, -) -> None: - """Run MLA decode with FP4 QK and FP4 PV tensor-core matmuls. - - Q is quantized to FP4, QK reads the packed K-view cache with swizzled - block scales, softmax probabilities are quantized to FP4 per page, and PV - repacks V nibbles from the shared KV cache while reading the auxiliary - V-view scale pool. No BF16 dequantized KV workspace is materialized on - this path. - """ - if not is_flashinfer_fp4_mla_attention_enabled(): - raise RuntimeError( - f"FP4 MLA attention decode requires {FLASHINFER_FP4_MLA_ATTENTION_ENV}=1." - ) - - head_dim = kv_lora_rank + qk_rope_head_dim - _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) - if metadata.page_size != FP4_MLA_TOKENS_PER_BLOCK: - raise ValueError( - f"FP4 MLA attention decode requires page_size={FP4_MLA_TOKENS_PER_BLOCK}, " - f"got {metadata.page_size}." - ) - - num_gen = q_nope.shape[0] - if num_gen == 0: - return - num_gen_seqs = metadata.num_seqs - metadata.num_contexts - if num_gen != num_gen_seqs: - raise NotImplementedError( - "FP4 MLA attention decode currently supports one query token per " - f"generation sequence, got {num_gen} query tokens for " - f"{num_gen_seqs} sequences." - ) - - num_heads = q_nope.shape[1] - if q_pe.shape[:2] != (num_gen, num_heads): - raise ValueError("FP4 MLA attention q_nope/q_pe batch dimensions do not match.") - if output.shape[:2] != (num_gen, num_heads): - raise ValueError("FP4 MLA attention output batch dimensions do not match.") - if q_nope.shape[-1] != kv_lora_rank: - raise ValueError( - f"q_nope last dimension must be kv_lora_rank={kv_lora_rank}, got {q_nope.shape[-1]}." - ) - if q_pe.shape[-1] != qk_rope_head_dim: - raise ValueError( - f"q_pe last dimension must be qk_rope_head_dim={qk_rope_head_dim}, " - f"got {q_pe.shape[-1]}." - ) - - if getattr(metadata, "fp4_mla_v_scale_pool", None) is None: - raise RuntimeError( - "FP4 MLA attention decode requires the auxiliary V scale pool to be allocated." - ) - - global_scale = _get_fp4_mla_global_scale(metadata, q_nope.device) - q_residual_dim = FP4_MLA_Q_RESIDUAL_DIM - _validate_fp4_mla_attention_q_shape(head_dim, q_residual_dim) - - q_full = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_q_buf", - (num_gen, num_heads, head_dim), - dtype=q_nope.dtype, - device=q_nope.device, - ) - q_full[..., :kv_lora_rank].copy_(q_nope) - q_full[..., kv_lora_rank:].copy_(q_pe) - q_2d = q_full.reshape(num_gen * num_heads, head_dim) - if q_2d.dtype not in (torch.bfloat16, torch.float8_e4m3fn): - raise TypeError( - f"FP4 MLA residual Q quantization requires BF16 or FP8 Q; got {q_2d.dtype}." - ) - q_fp4, q_sf = torch.ops.trtllm.fp4_quantize_with_residual( - q_2d, - global_scale, - q_residual_dim, - is_act=True, - ) - q_sf = q_sf.view(torch.float8_e4m3fn) - - kv_cache, sf_cache = _get_fp4_mla_cache_tensors(metadata, layer_idx) - sf_cache = sf_cache.view(torch.float8_e4m3fn) - - num_gen_blocks = metadata.num_generation_blocks - v_sf = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=kv_lora_rank)[local_layer].view( - torch.float8_e4m3fn - ) - - src_page_ids = metadata.paged_kv_indices[ - metadata.num_context_blocks : metadata.num_context_blocks + num_gen_blocks - ] - kv_lens = metadata.kv_lens_cuda_runtime[metadata.num_contexts : metadata.num_seqs] - max_pages = _max_generation_pages(metadata) - if max_pages == 0: - return - - backend = _fp4_mla_attention_backend() - if backend == "cutile": - from .fp4_mla_cutile import fp4_mla_paged_attention - - total_p_rows = max(src_page_ids.shape[0] * num_heads, 1) - p_fp4 = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_p_buf", - (total_p_rows, metadata.page_size // 2), - dtype=torch.uint8, - device=q_nope.device, - ) - p_sf = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_p_sf_buf", - (max(_get_fp4_mla_swizzled_scale_size(total_p_rows, metadata.page_size), 1),), - dtype=torch.float8_e4m3fn, - device=q_nope.device, - ) - stats_shape = (num_gen, num_heads) - max_scores = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_max_buf", - stats_shape, - dtype=torch.float32, - device=q_nope.device, - ) - denom = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_denom_buf", - stats_shape, - dtype=torch.float32, - device=q_nope.device, - ) - page_max = None - page_sum = None - if max_pages >= 8: - page_stats_shape = (num_gen, max_pages, num_heads) - page_max = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_page_max_buf", - page_stats_shape, - dtype=torch.float32, - device=q_nope.device, - ) - page_sum = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_page_sum_buf", - page_stats_shape, - dtype=torch.float32, - device=q_nope.device, - ) - assume_full_pages = _infer_cutile_assume_full_pages( - metadata, - max_pages, - metadata.page_size, - ) - assume_valid_pages = False - _fp4_mla_debug( - "attention decode cutile launch: " - f"num_gen={num_gen} num_heads={num_heads} local_layer={local_layer} " - f"layer_idx={layer_idx} head_dim={head_dim} kv_lora_rank={kv_lora_rank} " - f"rope_dim={qk_rope_head_dim} max_pages={max_pages} " - f"assume_full_pages={assume_full_pages} " - f"assume_valid_pages={assume_valid_pages}" - ) - fp4_mla_paged_attention( - q_fp4, - q_sf, - kv_cache, - sf_cache, - v_sf, - global_scale, - src_page_ids, - metadata.paged_kv_indptr_decode, - kv_lens, - output, - sm_scale=float(sm_scale), - num_heads=num_heads, - v_head_dim=kv_lora_rank, - page_size=metadata.page_size, - q_residual_dim=q_residual_dim, - max_pages=max_pages, - assume_full_pages=assume_full_pages, - assume_valid_pages=assume_valid_pages, - p_fp4_workspace=p_fp4, - p_sf_workspace=p_sf, - max_scores_workspace=max_scores, - denom_workspace=denom, - page_max_workspace=page_max, - page_sum_workspace=page_sum, - ) - _debug_sync("attention_cutile") - return - total_p_rows = num_gen_blocks * num_heads - p_fp4 = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_p_buf", - (max(total_p_rows, 1), metadata.page_size // 2), - dtype=torch.uint8, - device=q_nope.device, - )[:total_p_rows] - p_sf = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_p_sf_buf", - (max(_get_fp4_mla_swizzled_scale_size(total_p_rows, metadata.page_size), 1),), - dtype=torch.float8_e4m3fn, - device=q_nope.device, - ) - stats_shape = (num_gen, num_heads) - max_scores = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_max_buf", - stats_shape, - dtype=torch.float32, - device=q_nope.device, - ) - denom = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_denom_buf", - stats_shape, - dtype=torch.float32, - device=q_nope.device, - ) - - if backend == _FP4_MLA_CUTE_DSL_BACKEND: - from .fp4_mla_cute import run_fp4_mla_attention_decode_cute - - page_stats_shape = (num_gen, max_pages, num_heads) - page_max = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_page_max_buf", - page_stats_shape, - dtype=torch.float32, - device=q_nope.device, - ) - page_sum = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_page_sum_buf", - page_stats_shape, - dtype=torch.float32, - device=q_nope.device, - ) - run_fp4_mla_attention_decode_cute( - output=output, - max_scores=max_scores, - denom=denom, - page_max=page_max, - page_sum=page_sum, - p_fp4=p_fp4, - p_sf=p_sf, - q_fp4=q_fp4, - q_sf=q_sf, - kv_cache=kv_cache, - sf_cache=sf_cache, - v_sf=v_sf, - global_scale=global_scale, - src_page_ids=src_page_ids, - paged_kv_indptr_decode=metadata.paged_kv_indptr_decode, - kv_lens=kv_lens, - sm_scale=float(sm_scale), - kv_lora_rank=kv_lora_rank, - qk_rope_head_dim=qk_rope_head_dim, - q_residual_dim=q_residual_dim, - page_size=metadata.page_size, - max_pages=max_pages, - ) - _debug_sync("attention_cute_dsl") - return - if backend != "triton": - raise ValueError( - f"Unsupported FP4 MLA attention backend '{backend}'. " - f"Set {FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV} to 'triton', " - "'cutile', or 'cute_dsl'." - ) - - p_probs = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_p_prob_buf", - (max(num_gen * num_heads, 1), metadata.page_size), - dtype=torch.float32, - device=q_nope.device, - )[: num_gen * num_heads] - - block_h = 128 - block_t = metadata.page_size - block_k = 256 - block_v = 128 - q_head_dim = head_dim + q_residual_dim - q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE - k_sf_per_token = head_dim // FP4_BLOCK_SIZE - sf_per_page = metadata.page_size // FP4_BLOCK_SIZE - num_head_blocks = triton.cdiv(num_heads, block_h) - - _fp4_mla_debug( - "attention decode: " - f"num_gen={num_gen} num_heads={num_heads} local_layer={local_layer} " - f"layer_idx={layer_idx} head_dim={head_dim} q_head_dim={q_head_dim} " - f"q_residual_dim={q_residual_dim} " - f"kv_lora_rank={kv_lora_rank} rope_dim={qk_rope_head_dim} " - f"num_gen_blocks={num_gen_blocks} max_pages={max_pages} " - f"num_head_blocks={num_head_blocks} sm_scale={sm_scale}" - ) - _fp4_mla_debug(f"attention q_nope: {_tensor_layout(q_nope)}") - _fp4_mla_debug(f"attention q_pe: {_tensor_layout(q_pe)}") - _fp4_mla_debug(f"attention output: {_tensor_layout(output)}") - _fp4_mla_debug(f"attention q_fp4: {_tensor_layout(q_fp4)}") - _fp4_mla_debug(f"attention q_sf: {_tensor_layout(q_sf)}") - _fp4_mla_debug(f"attention kv_cache: {_tensor_layout(kv_cache)}") - _fp4_mla_debug(f"attention sf_cache: {_tensor_layout(sf_cache)}") - _fp4_mla_debug(f"attention v_sf: {_tensor_layout(v_sf)}") - _fp4_mla_debug(f"attention p_fp4: {_tensor_layout(p_fp4)}") - _fp4_mla_debug(f"attention p_sf: {_tensor_layout(p_sf)}") - _fp4_mla_debug(f"attention p_probs: {_tensor_layout(p_probs)}") - _debug_tensor_range("attention src_page_ids", src_page_ids) - _debug_tensor_range("attention paged_kv_indptr_decode", metadata.paged_kv_indptr_decode) - _debug_tensor_range("attention kv_lens", kv_lens) - - _fp4_mla_debug( - "attention stats launch: " - f"grid=({num_gen}, {num_head_blocks}) " - f"block_h={block_h} block_t={block_t} block_k={block_k}" - ) - _fp4_mla_attention_stats_kernel[(num_gen, num_head_blocks)]( - max_scores, - denom, - q_fp4, - q_sf, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - metadata.paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], - kv_cache.shape[0], - q_fp4.stride(0), - q_fp4.stride(1), - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - sf_cache.stride(0), - max_scores.stride(0), - sm_scale=sm_scale, - NUM_HEADS=num_heads, - Q_HEAD_D=q_head_dim, - K_HEAD_D=head_dim, - Q_RESIDUAL_D=q_residual_dim, - PAGE_SIZE=metadata.page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - Q_SF_PER_TOKEN=q_sf_per_token, - K_SF_PER_TOKEN=k_sf_per_token, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - BLOCK_T=block_t, - BLOCK_K=block_k, - ) - _debug_sync("attention_stats") - - for page_rel in range(max_pages): - _fp4_mla_debug( - "attention prob page store launch: " - f"page_rel={page_rel} grid=({num_gen}, {num_head_blocks})" - ) - _fp4_mla_attention_prob_store_page_kernel[(num_gen, num_head_blocks)]( - p_probs, - max_scores, - denom, - q_fp4, - q_sf, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - metadata.paged_kv_indptr_decode, - kv_lens, - page_rel, - src_page_ids.shape[0], - kv_cache.shape[0], - p_probs.stride(0), - p_probs.stride(1), - q_fp4.stride(0), - q_fp4.stride(1), - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - sf_cache.stride(0), - max_scores.stride(0), - sm_scale=sm_scale, - NUM_HEADS=num_heads, - Q_HEAD_D=q_head_dim, - K_HEAD_D=head_dim, - Q_RESIDUAL_D=q_residual_dim, - PAGE_SIZE=metadata.page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - Q_SF_PER_TOKEN=q_sf_per_token, - K_SF_PER_TOKEN=k_sf_per_token, - BLOCK_H=block_h, - BLOCK_K=block_k, - ) - _debug_sync(f"attention_prob_page_store_{page_rel}") - - _fp4_mla_debug( - "attention prob page pack launch: " - f"page_rel={page_rel} grid=({num_gen}, {sf_per_page}, " - f"{num_head_blocks})" - ) - _fp4_mla_attention_prob_pack_page_kernel[ - ( - num_gen, - sf_per_page, - num_head_blocks, - ) - ]( - p_fp4, - p_sf, - p_probs, - metadata.paged_kv_indptr_decode, - kv_lens, - page_rel, - src_page_ids.shape[0], - p_fp4.stride(0), - p_fp4.stride(1), - p_probs.stride(0), - p_probs.stride(1), - NUM_HEADS=num_heads, - PAGE_SIZE=metadata.page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - SF_PER_PAGE=sf_per_page, - P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, - BLOCK_H=block_h, - ) - _debug_sync(f"attention_prob_page_pack_{page_rel}") - - _fp4_mla_debug( - "attention pv launch: " - f"grid=({num_gen}, {num_head_blocks}, " - f"{triton.cdiv(kv_lora_rank, block_v)}) block_v={block_v}" - ) - _fp4_mla_attention_pv_kernel[ - ( - num_gen, - num_head_blocks, - triton.cdiv(kv_lora_rank, block_v), - ) - ]( - output, - p_fp4, - p_sf, - kv_cache, - v_sf, - global_scale, - src_page_ids, - metadata.paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], - kv_cache.shape[0], - output.stride(0), - output.stride(1), - output.stride(2), - p_fp4.stride(0), - p_fp4.stride(1), - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - v_sf.stride(0), - NUM_HEADS=num_heads, - V_HEAD_D=kv_lora_rank, - PAGE_SIZE=metadata.page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - SF_PER_PAGE=sf_per_page, - MAX_PAGES=max_pages, - P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, - BLOCK_H=block_h, - BLOCK_V=block_v, - ) - _debug_sync("attention_pv") - - -def update_hp_kv_for_fp4_mla( - metadata: Any, - latent_cache: Optional[torch.Tensor], - local_layer: int, - *, - phase: _HPUpdatePhase = "all", -) -> None: - """Store recent KV tokens at BF16 into the high-precision pool. - - Called on every layer before the attention kernel. The pool acts as a - circular buffer of HP_BLOCK_SIZE slots per sequence: - - Context phase stores the last ``kv_len % HP_BLOCK_SIZE`` new tokens of - each request into buffer positions [0, remainder). These are the - tail tokens that do not fill a complete FP4 block of 16. - - Generation phase stores the single new token for each request into - position ``(kv_len - 1) % HP_BLOCK_SIZE``, overwriting the oldest - entry in the circular buffer. - - The Triton kernels use the GPU ``seq_slots`` tensor for scatter indexing - and are CUDA-graph-compatible for the generation phase. - - Args: - metadata: Attention metadata exposing ``num_contexts``, ``num_seqs``, - ``seq_slots`` / ``seq_slots_cpu``, ``request_ids``, - ``is_cuda_graph``, ``is_warmup``, ``hp_pool_owners``, - ``high_precision_kv_pool``, ``prompt_lens_cpu_runtime``, - ``prompt_lens_cuda_runtime``, ``kv_lens_cuda_runtime``. - latent_cache: MLA latent cache for the current tokens, shape - [num_tokens, head_dim]. When ``None``, only ownership tracking - runs (no data is written to the pool). - local_layer: Layer index within the local pipeline-parallel slice. - phase: Which portion of ``latent_cache`` is present. ``"all"`` means - context tokens followed by generation tokens, ``"context"`` means - only context tokens, and ``"generation"`` means only generation - tokens. - """ - if phase not in ("all", "context", "generation"): - raise ValueError(f"Unexpected FP4 MLA HP update phase: {phase}") - if metadata.hp_pool_owners is None: - return - num_contexts = metadata.num_contexts - num_seqs = metadata.num_seqs - update_context = phase in ("all", "context") - update_generation = phase in ("all", "generation") - - # ------------------------------------------------------------------ - # Ownership tracking (layer 0, eager mode only - debug guard). - # Context phase never uses CUDA graph; decode check is debug-only. - # ------------------------------------------------------------------ - if local_layer == 0 and not metadata.is_cuda_graph and not metadata.is_warmup: - # Context: register ownership of each seq_slot. - if update_context: - for batch_idx in range(num_contexts): - seq_slot = metadata.seq_slots_cpu[batch_idx].item() - request_id = metadata.request_ids[batch_idx] - metadata.hp_pool_owners[seq_slot] = request_id - - # Decode: verify that the expected request still owns each slot. - if update_generation: - for batch_idx in range(num_contexts, num_seqs): - seq_slot = metadata.seq_slots_cpu[batch_idx].item() - request_id = metadata.request_ids[batch_idx] - owner = metadata.hp_pool_owners.get(seq_slot) - if owner != request_id: - raise RuntimeError( - f"HP KV pool ownership mismatch: seq_slot={seq_slot} " - f"is owned by request {owner} but request " - f"{request_id} is attempting to use it" - ) - - if latent_cache is None: - return - - # ------------------------------------------------------------------ - # Triton kernel dispatch - runs on every layer, CUDA-graph-safe. - # ------------------------------------------------------------------ - pool = metadata.high_precision_kv_pool - head_dim = latent_cache.shape[-1] - pool_head_dim = pool.shape[-1] // HP_BLOCK_SIZE - if pool_head_dim < head_dim: - raise RuntimeError( - f"FP4 MLA HP pool head dimension is too small: got " - f"{pool_head_dim}, need at least {head_dim}." - ) - block_d = triton.next_power_of_2(head_dim) - pool_s0 = pool.stride(0) # stride across sequence slots - pool_s1 = pool.stride(1) # stride across layers - lc_stride = latent_cache.stride(0) - _fp4_mla_debug( - "hp update: " - f"phase={phase} local_layer={local_layer} num_contexts={num_contexts} " - f"num_seqs={num_seqs} head_dim={head_dim} " - f"pool_head_dim={pool_head_dim} block_d={block_d}" - ) - _fp4_mla_debug(f"hp latent_cache: {_tensor_layout(latent_cache)}") - _fp4_mla_debug(f"hp pool: {_tensor_layout(pool)}") - _debug_tensor_range("hp seq_slots", metadata.seq_slots[:num_seqs]) - _debug_tensor_range("hp kv_lens", metadata.kv_lens_cuda_runtime[:num_seqs]) - - # Context phase: store last (kv_len % HP_BLOCK_SIZE) new tokens. - if update_context and num_contexts > 0: - prompt_lens_cpu = metadata.prompt_lens_cpu_runtime[:num_contexts] - # Exclusive prefix sum: token offset in latent_cache for each ctx seq. - token_offsets_cpu = torch.zeros(num_contexts, dtype=torch.int32, device="cpu") - if num_contexts > 1: - token_offsets_cpu[1:].copy_(torch.cumsum(prompt_lens_cpu[:-1].to(torch.int32), dim=0)) - token_offsets_gpu = token_offsets_cpu.to(pool.device, non_blocking=False) - prompt_lens_gpu = metadata.prompt_lens_cuda_runtime[:num_contexts] - - _fp4_mla_debug( - "hp context launch: " - f"grid=({num_contexts}, {HP_BLOCK_SIZE}) " - f"token_offsets={token_offsets_cpu.tolist()}" - ) - _debug_tensor_range("hp context prompt_lens", prompt_lens_gpu) - _hp_kv_store_context_kernel[(num_contexts, HP_BLOCK_SIZE)]( - pool, - latent_cache, - metadata.seq_slots, - metadata.kv_lens_cuda_runtime, - token_offsets_gpu, - prompt_lens_gpu, - pool.shape[0], - pool.shape[1], - local_layer, - pool_s0, - pool_s1, - lc_stride, - D=head_dim, - POOL_HEAD_D=pool_head_dim, - BLOCK_D=block_d, - HP_BLOCK=HP_BLOCK_SIZE, - ) - _debug_sync("hp_context") - - # Generation phase: store current token at (kv_len - 1) % HP_BLOCK_SIZE. - num_gen = num_seqs - num_contexts - if update_generation and num_gen > 0: - gen_tok_start = 0 - if phase == "all": - # Scalar offset: number of context tokens packed before gen tokens. - gen_tok_start = int(metadata.prompt_lens_cpu_runtime[:num_contexts].sum().item()) - - _fp4_mla_debug(f"hp generation launch: grid=({num_gen},) gen_tok_start={gen_tok_start}") - _hp_kv_store_gen_kernel[(num_gen,)]( - pool, - latent_cache, - metadata.seq_slots[num_contexts:], - metadata.kv_lens_cuda_runtime[num_contexts:], - gen_tok_start, - pool.shape[0], - pool.shape[1], - local_layer, - pool_s0, - pool_s1, - lc_stride, - D=head_dim, - POOL_HEAD_D=pool_head_dim, - BLOCK_D=block_d, - HP_BLOCK=HP_BLOCK_SIZE, - ) - _debug_sync("hp_generation") diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py new file mode 100644 index 000000000000..6fd1920e19e4 --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py @@ -0,0 +1,1536 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Triton FP4 MLA decode-path kernels. + +This file owns the no-dequant ``triton`` attention backend selected via +``TRTLLM_FLASHINFER_FP4_MLA_ATTENTION_BACKEND=triton``. It is separate from +``fp4_mla_kernels.py``, which holds shared KV-cache scatter/dequant and HP-pool +helper kernels. + +Optimizations -- self-contained, public-Triton compatible (no ``tl.ext.*`` +or any private-Triton extension): + +* ``USE_TMA_DATA_LOAD`` path through ``_fp4_mla_qk_scores_tile`` -- builds + device-side TMA descriptors via ``tl.make_tensor_descriptor`` for the + full-K window and the Q residual tail, with the residual-Q permute/split + idiom that maps the Q tail's interleaved groups onto the K tail. +* ``ASSUME_FULL_HEADS`` / ``ASSUME_FULL_PAGES`` / ``ASSUME_VALID_PAGES`` + constexpr branches that drop ``tl.where`` masks on the hot decode path. +* ``PACK_PROBS`` fused page-stats kernel that quantizes P to FP4 in registers + before the page-max correction, eliminating the ``p_probs`` HBM round-trip. +* ``_fp4_mla_swizzled_sf_offset_row_block`` faster offset helper for the + perfect (NUM_HEADS=128, BLOCK_H=128) case used by page-stats and PV. +* ``tl.assume()`` stride hints in front of every ``make_tensor_descriptor`` + call -- helps Triton vectorize TMA loads. +* Pipelined PV loop via ``tl.range(..., num_stages=PV_LOOP_STAGES)``. +""" + +import triton +import triton.language as tl + +_LOG2_E = tl.constexpr(1.4426950408889634) + + +@triton.jit +def _fp4_mla_swizzled_sf_offset( + row_idx, + col_idx, + SF_PER_TOKEN: tl.constexpr, +): + padded_cols = ((SF_PER_TOKEN + 3) // 4) * 4 + col_in_group = col_idx % 4 + col_group = col_idx // 4 + row_in_group0 = row_idx % 32 + row_in_group1 = (row_idx % 128) // 32 + row_group = row_idx // 128 + return ( + col_in_group + + col_group * (4 * 128) + + row_in_group0 * 16 + + row_in_group1 * 4 + + row_group * (128 * padded_cols) + ) + + +@triton.jit +def _fp4_mla_swizzled_sf_offset_row_block( + row_group, row_offsets, col_idx, SF_PER_TOKEN: tl.constexpr +): + """Faster offset variant when the row group is a known constant. + + ``row_offsets`` ranges over [0, 128) within the row group; ``row_group`` + is the constant block index. Skips the divmod by 128 that the generic + helper has to do. + """ + padded_cols = ((SF_PER_TOKEN + 3) // 4) * 4 + col_part = (col_idx % 4) + (col_idx // 4) * (4 * 128) + row_part = (row_offsets % 32) * 16 + ((row_offsets % 128) // 32) * 4 + return col_part + row_part + row_group * (128 * padded_cols) + + +@triton.jit +def _fp4_e2m1_quantize(x): + abs_x = tl.abs(x) + magnitude = tl.where( + abs_x < 0.25, + 0, + tl.where( + abs_x < 0.75, + 1, + tl.where( + abs_x < 1.25, + 2, + tl.where( + abs_x < 1.75, + 3, + tl.where(abs_x < 2.5, 4, tl.where(abs_x < 3.5, 5, tl.where(abs_x < 5.0, 6, 7))), + ), + ), + ), + ) + sign = tl.where(x < 0.0, 8, 0) + return (magnitude | sign).to(tl.uint8) + + +@triton.jit +def _fp4_pack_low_nibbles(even_packed, odd_packed): + """PTX helper: pack the low nibbles of two bytes into one byte (low + high<<4).""" + return tl.inline_asm_elementwise( + """ + { + .reg .b32 lo; + .reg .b32 hi; + and.b32 lo, $1, 15; + and.b32 hi, $2, 15; + shl.b32 hi, hi, 4; + or.b32 $0, lo, hi; + } + """, + constraints="=r,r,r", + args=[even_packed, odd_packed], + dtype=tl.uint8, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fp4_pack_high_nibbles(even_packed, odd_packed): + """PTX helper: pack the high nibbles of two bytes into one byte (low + high<<4).""" + return tl.inline_asm_elementwise( + """ + { + .reg .b32 lo; + .reg .b32 hi; + shr.u32 lo, $1, 4; + and.b32 lo, lo, 15; + and.b32 hi, $2, 240; + or.b32 $0, lo, hi; + } + """, + constraints="=r,r,r", + args=[even_packed, odd_packed], + dtype=tl.uint8, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fp4_e2m1_quantize_packed(even, odd): + """Quantize two FP32 values into a single packed E2M1x2 byte via PTX.""" + return tl.inline_asm_elementwise( + """ + { + .reg .b8 r; + cvt.rn.satfinite.e2m1x2.f32 r, $1, $2; + mov.b32 $0, {r, r, r, r}; + } + """, + constraints="=r,f,f", + args=[odd.to(tl.float32), even.to(tl.float32)], + dtype=tl.uint8, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_start, + head_offsets, + token_offsets, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + NUM_HEADS: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, +): + if ASSUME_VALID_PAGES: + safe_compact_page = compact_page + physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) + safe_physical_page = physical_page + else: + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + ) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + q_rows = q_row_base + head_offsets + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_q_rows = q_rows + else: + mask_h = head_offsets < NUM_HEADS + safe_q_rows = tl.where(mask_h, q_rows, q_row_base) + scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) + + packed_k_offsets = tl.arange(0, BLOCK_K // 2) + scale_offsets = tl.arange(0, BLOCK_K // FP4_BLOCK) + residual_groups = Q_RESIDUAL_D // FP4_BLOCK + non_residual_groups = K_HEAD_D // FP4_BLOCK - residual_groups + if USE_TMA_DATA_LOAD and FULL_BLOCK_END > 0: + tl.assume(q_fp4_s0 % 8 == 0) + tl.assume(q_fp4_s1 == 1) + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + q_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, BLOCK_K // 2], + ) + k_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, BLOCK_K // 2], + ) + if USE_TMA_DATA_LOAD and Q_RESIDUAL_D == 64 and TAIL_BLOCK_K == 128: + q_tail_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 64], + ) + k_tail_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 32], + ) + # Static Python loop (not tl.range) — Triton unrolls it; this also keeps + # the Triton 3.6 AutomaticWarpSpecialization / NVWSInsertTmemAref passes + # from picking up the loop and ICEing on sm_100. + for q_start in range(0, FULL_BLOCK_END, BLOCK_K): + q_elem_offsets = q_start + packed_k_offsets * 2 + q_group_offsets = q_elem_offsets // FP4_BLOCK + k_group_offsets = tl.where( + q_group_offsets < non_residual_groups, + q_group_offsets, + non_residual_groups + (q_group_offsets - non_residual_groups) // 2, + ) + byte_offsets_in_group = (q_elem_offsets % FP4_BLOCK) // 2 + packed_q_cols = q_start // 2 + packed_k_offsets + packed_k_cols = k_group_offsets * (FP4_BLOCK // 2) + byte_offsets_in_group + mask_k = q_elem_offsets < Q_HEAD_D + safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) + safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) + if ( + USE_TMA_DATA_LOAD + and FULL_BLOCK_END > 0 + and q_start + BLOCK_K <= non_residual_groups * FP4_BLOCK + ): + q_vals = q_desc.load([(q_row_base + head_start).to(tl.int32), q_start // 2]) + k_vals = k_desc.load([safe_physical_page.to(tl.int32), 0, q_start // 2]) + k_vals = tl.reshape(k_vals, (BLOCK_T, BLOCK_K // 2)) + if not ASSUME_VALID_PAGES: + k_vals = tl.where(valid_physical_page, k_vals, 0) + else: + q_vals = tl.load( + q_fp4_ptr + + safe_q_rows[:, None] * q_fp4_s0 + + safe_packed_q_cols[None, :] * q_fp4_s1, + mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], + other=0, + ) + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + safe_packed_k_cols[None, :] * kv_s4, + mask=mask_k[None, :] + if ASSUME_VALID_PAGES + else valid_physical_page & mask_k[None, :], + other=0, + ) + + q_sf_cols = q_start // FP4_BLOCK + scale_offsets + k_sf_cols = tl.where( + q_sf_cols < non_residual_groups, + q_sf_cols, + non_residual_groups + (q_sf_cols - non_residual_groups) // 2, + ) + mask_sf = q_sf_cols < Q_SF_PER_TOKEN + safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) + safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) + q_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN + ) + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN + ) + q_scales = tl.load(q_sf_ptr + q_sf_offsets) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) + scores = tl.dot_scaled( + q_vals, + q_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + if FULL_BLOCK_END < Q_HEAD_D: + q_start = FULL_BLOCK_END + # Residual Q fast path disabled: it issues two chained dot_scaled + # calls into the same accumulator (one for even Q lane, one for odd), + # which lowers to a TMEM alloc with multiple uses and trips + # NVWSInsertTmemAref::hasOneUse() on Triton 3.6.0 / sm_100. + if False and Q_RESIDUAL_D == 64 and TAIL_BLOCK_K == 128: + residual_packed_offsets = tl.arange(0, 32) + residual_scale_offsets = tl.arange(0, 4) + packed_k_cols = non_residual_groups * (FP4_BLOCK // 2) + residual_packed_offsets + if USE_TMA_DATA_LOAD: + k_vals = k_tail_desc.load( + [ + safe_physical_page.to(tl.int32), + 0, + (non_residual_groups * (FP4_BLOCK // 2)).to(tl.int32), + ] + ) + k_vals = tl.reshape(k_vals, (BLOCK_T, 32)) + if not ASSUME_VALID_PAGES: + k_vals = tl.where(valid_physical_page, k_vals, 0) + elif ASSUME_VALID_PAGES: + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + packed_k_cols[None, :] * kv_s4, + ) + else: + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + packed_k_cols[None, :] * kv_s4, + mask=valid_physical_page, + other=0, + ) + k_sf_cols = non_residual_groups + residual_scale_offsets + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], k_sf_cols[None, :], K_SF_PER_TOKEN + ) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) + + q_tail_cols = q_start // 2 + tl.arange(0, 64) + if USE_TMA_DATA_LOAD and ASSUME_FULL_HEADS: + q_tail_vals = q_tail_desc.load( + [(q_row_base + head_start).to(tl.int32), q_start // 2] + ) + elif ASSUME_FULL_HEADS: + q_tail_vals = tl.load( + q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1 + ) + else: + q_tail_vals = tl.load( + q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1, + mask=mask_h[:, None], + other=0, + ) + # Map Q tail groups [0, 1, ..., 7] onto K tail groups [0, 0, 1, 1, ..., 3, 3]. + q_tail_vals = q_tail_vals.reshape([BLOCK_H, 4, 2, 8]).trans(0, 1, 3, 2) + q_even_vals, q_odd_vals = tl.split(q_tail_vals) + q_even_vals = q_even_vals.reshape([BLOCK_H, 32]) + q_odd_vals = q_odd_vals.reshape([BLOCK_H, 32]) + + q_tail_sf_cols = q_start // FP4_BLOCK + tl.arange(0, 8) + q_tail_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], q_tail_sf_cols[None, :], Q_SF_PER_TOKEN + ) + q_tail_scales = tl.load(q_sf_ptr + q_tail_sf_offsets) + q_tail_scales = q_tail_scales.reshape([BLOCK_H, 4, 2]) + q_even_scales, q_odd_scales = tl.split(q_tail_scales) + # Compute even/odd partial dots into fresh accumulators, then add + # back. Routing through `acc=scores` for both calls gives a single + # TMEM alloc with multiple uses, which trips + # NVWSInsertTmemAref::TmemAccessDag::build's hasOneUse() assertion + # on Triton 3.6.0 / sm_100. + tail_even = tl.dot_scaled( + q_even_vals, + q_even_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + fast_math=True, + rhs_k_pack=True, + ) + tail_odd = tl.dot_scaled( + q_odd_vals, + q_odd_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + fast_math=True, + rhs_k_pack=True, + ) + scores = scores + tail_even + tail_odd + else: + tail_packed_offsets = tl.arange(0, TAIL_BLOCK_K // 2) + tail_scale_offsets = tl.arange(0, TAIL_BLOCK_K // FP4_BLOCK) + q_elem_offsets = q_start + tail_packed_offsets * 2 + q_group_offsets = q_elem_offsets // FP4_BLOCK + k_group_offsets = tl.where( + q_group_offsets < non_residual_groups, + q_group_offsets, + non_residual_groups + (q_group_offsets - non_residual_groups) // 2, + ) + byte_offsets_in_group = (q_elem_offsets % FP4_BLOCK) // 2 + packed_q_cols = q_start // 2 + tail_packed_offsets + packed_k_cols = k_group_offsets * (FP4_BLOCK // 2) + byte_offsets_in_group + mask_k = q_elem_offsets < Q_HEAD_D + safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) + safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) + q_vals = tl.load( + q_fp4_ptr + + safe_q_rows[:, None] * q_fp4_s0 + + safe_packed_q_cols[None, :] * q_fp4_s1, + mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], + other=0, + ) + k_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + token_offsets[:, None].to(tl.int64) * kv_s2 + + safe_packed_k_cols[None, :] * kv_s4, + mask=mask_k[None, :] + if ASSUME_VALID_PAGES + else valid_physical_page & mask_k[None, :], + other=0, + ) + + q_sf_cols = q_start // FP4_BLOCK + tail_scale_offsets + k_sf_cols = tl.where( + q_sf_cols < non_residual_groups, + q_sf_cols, + non_residual_groups + (q_sf_cols - non_residual_groups) // 2, + ) + mask_sf = q_sf_cols < Q_SF_PER_TOKEN + safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) + safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) + q_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN + ) + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN + ) + q_scales = tl.load(q_sf_ptr + q_sf_offsets) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) + scores = tl.dot_scaled( + q_vals, + q_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + return scores + + +@triton.jit +def _fp4_mla_attention_page_stats_kernel( + page_max_ptr, + page_sum_ptr, + p_fp4_ptr, + p_sf_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + q_global_scale_ptr, + page_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_stats_s0, + page_stats_s1, + p_s0, + p_s1, + p_num_rows, + q_num_rows, + sm_scale, + local_layer, + pscale_s0, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + PACK_PROBS: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + USE_PER_PAGE_SCALE: tl.constexpr = False, + occupancy: tl.constexpr = 1, +): + """Page-stats fused QK + softmax-stats + (optional) FP4 P pack. + + Each program owns one (query, head_block, page). Probs are quantized + against the per-page max; the page-max correction + ``exp(page_max - global_max) / denom`` is folded into ``p_sf`` later by + ``_fp4_mla_attention_prob_scale_kernel``. + """ + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_rel = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, BLOCK_T) + q_row_base = query_idx * NUM_HEADS + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_start = page_rel * PAGE_SIZE + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + out_offsets = query_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h + + page_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + page_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) + if USE_TMA_DATA_LOAD and PACK_PROBS and ASSUME_FULL_HEADS and ASSUME_VALID_PAGES: + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + if ASSUME_FULL_PAGES or page_start < kv_len: + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + page_table_start + page_rel, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + BLOCK_T, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + if ASSUME_FULL_PAGES: + valid_t = tl.full([BLOCK_T], True, dtype=tl.int1) + else: + valid_t = page_start + offs_t < kv_len + global_scale = tl.load(global_scale_ptr) + if USE_PER_PAGE_SCALE: + # Independent dynamic Q scale and per-page (K/V shared) KV scale. + # page_gscale also folds into the stored P scale below so the per-page + # V scaling cancels inside the fused PV dot. + compact_page = page_table_start + page_rel + if ASSUME_VALID_PAGES: + phys_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) + else: + valid_cp = (compact_page >= 0) & (compact_page < page_ids_len) + phys_page = tl.load( + src_page_ids_ptr + tl.where(valid_cp, compact_page, 0), + mask=valid_cp, + other=0, + ).to(tl.int64) + phys_page = tl.where((phys_page >= 0) & (phys_page < num_pages), phys_page, 0) + page_gscale = tl.load(page_scale_ptr + local_layer * pscale_s0 + phys_page) + q_gscale = tl.load(q_global_scale_ptr) + qk_scale = sm_scale / (q_gscale * page_gscale) + else: + # Static scale: global_scale == 1.0, so page_gscale == 1.0 makes the + # stored-P fold below a no-op and qk_scale == sm_scale. + page_gscale = global_scale + qk_scale = sm_scale / (global_scale * global_scale) + if ASSUME_FULL_HEADS and ASSUME_FULL_PAGES: + scores = scores * qk_scale + page_max = tl.max(scores, axis=1) + exp_scores = tl.math.exp2((scores - page_max[:, None]) * _LOG2_E) + page_sum = tl.sum(exp_scores, axis=1) + else: + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + page_max = tl.max(scores, axis=1) + safe_page_max = tl.where(mask_h, page_max, 0.0) + exp_scores = tl.math.exp2((scores - safe_page_max[:, None]) * _LOG2_E) + exp_scores = tl.where(mask_h[:, None] & valid_t[None, :], exp_scores, 0.0) + page_sum = tl.sum(exp_scores, axis=1) + + if PACK_PROBS: + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + # Fold 1/page_gscale into the stored P block scale (page_gscale == 1.0 + # in the static path, so this is a no-op there). Combined with V's + # baked page_gscale, the page scale cancels in the PV dot and the end + # out_scale = 1/(global_scale*P_GLOBAL_SCALE) = 1/P_GLOBAL_SCALE stays. + stored_scale = tl.where( + amax > 0.0, + tl.minimum(amax * (P_GLOBAL_SCALE / 6.0) / page_gscale, 448.0), + 1.0, + ) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + if not ASSUME_VALID_PAGES: + valid_compact_page = (page_table_start + page_rel >= 0) & ( + page_table_start + page_rel < page_ids_len + ) + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + ).to(tl.int64) + scale_cols = tl.arange(0, SF_PER_PAGE) + if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + if ASSUME_FULL_HEADS: + if ASSUME_VALID_PAGES: + tl.store(p_sf_ptr + sf_offsets, stored_scale) + else: + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_compact_page) + else: + tl.store( + p_sf_ptr + sf_offsets, + stored_scale, + mask=mask_h[:, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None], + ) + + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + if USE_TMA_DATA_LOAD and ASSUME_FULL_HEADS and ASSUME_VALID_PAGES: + p_desc.store( + [(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), + ) + elif ASSUME_FULL_HEADS: + if ASSUME_VALID_PAGES: + tl.store( + p_fp4_ptr + + safe_p_rows[:, None, None] * p_s0 + + byte_cols[None, :, :] * p_s1, + packed, + ) + else: + tl.store( + p_fp4_ptr + + safe_p_rows[:, None, None] * p_s0 + + byte_cols[None, :, :] * p_s1, + packed, + mask=valid_compact_page, + ) + else: + tl.store( + p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=mask_h[:, None, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None, None], + ) + + if ASSUME_FULL_HEADS: + tl.store(page_max_ptr + out_offsets, page_max) + tl.store(page_sum_ptr + out_offsets, page_sum) + else: + tl.store(page_max_ptr + out_offsets, page_max, mask=mask_h) + tl.store(page_sum_ptr + out_offsets, page_sum, mask=mask_h) + + +@triton.jit +def _fp4_mla_attention_reduce_stats_kernel( + max_ptr, + denom_ptr, + page_max_ptr, + page_sum_ptr, + num_pages, + stats_s0, + page_stats_s0, + page_stats_s1, + NUM_HEADS: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + occupancy: tl.constexpr = 1, +): + """Combine per-page max/sum into a global max + denom per (query, head). + + Keep the page dimension in a loop instead of a 2D [pages, heads] vector. + Large decode batches can push max pages above 256, where the bulk vector + form becomes too large for a single Triton program. + """ + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + + max_score = tl.full((BLOCK_H,), -float("inf"), tl.float32) + for page_idx in tl.range(0, MAX_PAGES): + page_valid = page_idx < num_pages + page_offsets = gen_idx * page_stats_s0 + page_idx * page_stats_s1 + safe_offs_h + page_max = tl.load( + page_max_ptr + page_offsets, mask=mask_h & page_valid, other=-float("inf") + ) + max_score = tl.maximum(max_score, page_max) + + safe_max = tl.where(max_score > -float("inf"), max_score, 0.0) + denom = tl.zeros((BLOCK_H,), tl.float32) + for page_idx in tl.range(0, MAX_PAGES): + page_valid = page_idx < num_pages + page_offsets = gen_idx * page_stats_s0 + page_idx * page_stats_s1 + safe_offs_h + page_max = tl.load( + page_max_ptr + page_offsets, mask=mask_h & page_valid, other=-float("inf") + ) + page_sum = tl.load(page_sum_ptr + page_offsets, mask=mask_h & page_valid, other=0.0) + weights = tl.math.exp2((page_max - safe_max) * _LOG2_E) + denom += tl.where(page_sum > 0.0, page_sum * weights, 0.0) + + tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) + tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) + + +@triton.jit +def _fp4_mla_attention_prob_scale_kernel( + p_sf_ptr, + max_ptr, + denom_ptr, + page_max_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + stats_s0, + page_stats_s0, + page_stats_s1, + NUM_HEADS: tl.constexpr, + PAGE_SIZE: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + """Apply per-page softmax correction by scaling p_sf in place.""" + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_rel = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + if page_start >= kv_len: + return + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if not ASSUME_VALID_PAGES: + if (compact_page < 0) | (compact_page >= page_ids_len): + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + page_max = tl.load( + page_max_ptr + query_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + factor = tl.where(denom > 0.0, tl.math.exp2((page_max - max_score) * _LOG2_E) / denom, 0.0) + + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + ).to(tl.int64) + scale_cols = tl.arange(0, SF_PER_PAGE) + if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) + tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) + + +@triton.jit +def _fp4_mla_attention_pv_kernel( + out_ptr, + p_fp4_ptr, + p_sf_ptr, + kv_cache_ptr, + v_sf_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + out_s0, + out_s1, + out_s2, + out_num_rows, + p_s0, + p_s1, + p_num_rows, + kv_s0, + kv_s2, + kv_s4, + vsf_s0, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + USE_TMA_P_LOAD: tl.constexpr, + USE_TMA_V_LOAD: tl.constexpr, + PV_LOOP_STAGES: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_FULL_V: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + PAGE_SPLIT: tl.constexpr = 1, + PAGES_PER_SPLIT: tl.constexpr = 0, + PARTIAL_OUT: tl.constexpr = False, + partial_out_ptr=None, + partial_s0: tl.constexpr = 0, + partial_s1: tl.constexpr = 0, + partial_s2: tl.constexpr = 0, + partial_s3: tl.constexpr = 0, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + # When PAGE_SPLIT > 1 we encode (dim_block, split_idx) into program_id(2). + # The outer loop over pages is partitioned across split_idx programs so the + # grid grows by PAGE_SPLIT× — this lifts the bs<=32 PV grid out of the + # 0.5-wave-per-SM regime that the ncu report flagged. + prog2 = tl.program_id(2) + if PAGE_SPLIT > 1: + dim_block = prog2 // PAGE_SPLIT + split_idx = prog2 - dim_block * PAGE_SPLIT + else: + dim_block = prog2 + split_idx = 0 + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + if ASSUME_FULL_V: + mask_v = tl.full([BLOCK_V], True, dtype=tl.int1) + safe_offs_v = offs_v + else: + mask_v = offs_v < V_HEAD_D + safe_offs_v = tl.where(mask_v, offs_v, 0) + packed_t = tl.arange(0, PAGE_SIZE // 2) + scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) + even_t = packed_t * 2 + odd_t = even_t + 1 + v_packed_offsets = safe_offs_v // 2 + v_use_high_nibble = (safe_offs_v & 1) != 0 + if ASSUME_FULL_V and BLOCK_V == 128: + v_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + v_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + if USE_TMA_P_LOAD: + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + if USE_TMA_V_LOAD and ASSUME_FULL_HEADS and ASSUME_FULL_V: + tl.assume(out_s1 % 8 == 0) + tl.assume(out_s2 == 1) + out_desc = tl.make_tensor_descriptor( + out_ptr, + shape=[out_num_rows, V_HEAD_D], + strides=[out_s1, out_s2], + block_shape=[BLOCK_H, BLOCK_V], + ) + if USE_TMA_V_LOAD: + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, PAGE_SIZE, BLOCK_V // 2], + ) + + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + if PAGE_SPLIT > 1: + page_lo = split_idx * PAGES_PER_SPLIT + page_hi = tl.minimum(page_lo + PAGES_PER_SPLIT, MAX_PAGES) + else: + page_lo = 0 + page_hi = MAX_PAGES + for page_rel in tl.range(page_lo, page_hi, num_stages=PV_LOOP_STAGES): + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES or page_start < kv_len: + compact_page = page_table_start + page_rel + if ASSUME_VALID_PAGES: + safe_compact_page = compact_page + physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) + safe_physical_page = physical_page + else: + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + ) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + ).to(tl.int64) + if USE_TMA_P_LOAD: + p_vals = p_desc.load([(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) + else: + p_vals = tl.load( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, + mask=mask_h[:, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None], + other=0, + ) + p_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + p_scales = tl.load(p_sf_ptr + p_sf_offsets) + + if ASSUME_FULL_PAGES: + valid_even_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + valid_odd_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + else: + valid_even_t = page_start + even_t < kv_len + valid_odd_t = page_start + odd_t < kv_len + if USE_TMA_V_LOAD: + v_tile = v_desc.load( + [ + safe_physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ] + ) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + if not ASSUME_VALID_PAGES: + v_tile = tl.where(valid_physical_page, v_tile, 0) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + if not ASSUME_FULL_PAGES: + even_packed = tl.where(valid_even_t[None, :], even_packed, 0) + odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) + low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) + high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + else: + even_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + even_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_even_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), + other=0, + ) + odd_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + odd_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_odd_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), + other=0, + ) + even_low = even_packed & 0x0F + even_high = (even_packed >> 4) & 0x0F + odd_low = odd_packed & 0x0F + even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) + odd_nibble = tl.where(v_use_high_nibble[:, None], odd_packed >> 4, odd_low) + v_vals = even_nibble | (odd_nibble << 4) + v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) + acc = tl.dot_scaled( + p_vals, + p_scales, + "e2m1", + v_vals.T, + v_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + if PARTIAL_OUT: + # Write the unscaled partial accumulator to a float32 workspace; the + # reduce-PV kernel sums splits and applies out_scale + dtype cast. + # Layout: partial_out[query_idx, split_idx, head_offset, v_offset]. + base = ( + query_idx * partial_s0 + + split_idx * partial_s1 + + safe_offs_h[:, None] * partial_s2 + + safe_offs_v[None, :] * partial_s3 + ) + if ASSUME_FULL_HEADS and ASSUME_FULL_V: + tl.store(partial_out_ptr + base, acc) + else: + tl.store(partial_out_ptr + base, acc, mask=mask_h[:, None] & mask_v[None, :]) + elif ASSUME_FULL_HEADS and ASSUME_FULL_V: + out_vals = acc * out_scale + if USE_TMA_V_LOAD: + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + out_desc.store( + [ + (query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ], + out_vals, + ) + else: + tl.store( + out_ptr + query_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + ) + else: + tl.store( + out_ptr + + query_idx * out_s0 + + safe_offs_h[:, None] * out_s1 + + safe_offs_v[None, :] * out_s2, + acc * out_scale, + mask=mask_h[:, None] & mask_v[None, :], + ) + + +@triton.jit +def _fp4_mla_attention_pv_reduce_kernel( + out_ptr, + partial_ptr, + global_scale_ptr, + out_s0, + out_s1, + out_s2, + partial_s0, + partial_s1, + partial_s2, + partial_s3, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SPLIT: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_V: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + dim_block = tl.program_id(2) + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + if ASSUME_FULL_V: + mask_v = tl.full([BLOCK_V], True, dtype=tl.int1) + safe_offs_v = offs_v + else: + mask_v = offs_v < V_HEAD_D + safe_offs_v = tl.where(mask_v, offs_v, 0) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + for split_idx in tl.static_range(0, PAGE_SPLIT): + base = ( + query_idx * partial_s0 + + split_idx * partial_s1 + + safe_offs_h[:, None] * partial_s2 + + safe_offs_v[None, :] * partial_s3 + ) + acc += tl.load(partial_ptr + base, mask=mask_h[:, None] & mask_v[None, :], other=0.0) + out_vals = acc * out_scale + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + tl.store( + out_ptr + + query_idx * out_s0 + + safe_offs_h[:, None] * out_s1 + + safe_offs_v[None, :] * out_s2, + out_vals, + mask=mask_h[:, None] & mask_v[None, :], + ) + + +# --------------------------------------------------------------------------- +# Per-page dynamic-scale store path (triton only) +# +# Two passes per decode step over the *active* page (the page that holds the +# current step's new tokens): +# Pass A (_fp4_mla_page_scale_gen_kernel): compute the page amax over the +# shared K/V latent (old tokens read from the BF16 staging pool, new tokens +# from latent_cache) and write page_gscale = P_GLOBAL_SCALE / page_amax into +# the [num_layers, num_pages] fp32 page-scale pool. +# Pass B (_fp4_mla_page_requant_gen_kernel): re-quantize *every* FP4 tile of +# the active page from the same (staging + latent) source, baking page_gscale +# into the stored K and V block scales. Completed pages are never revisited, +# so their scale is frozen at the value computed on the step that filled them. +# +# Both passes assume the step's new tokens land in a single page (always true +# for 1-token decode; the Python dispatch guards the MTP boundary-cross case). +# --------------------------------------------------------------------------- + + +@triton.jit +def _fp4_mla_page_scale_gen_kernel( + page_scale_ptr, + stage_pool_ptr, + latent_cache_ptr, + seq_slots_ptr, + kv_lens_ptr, + prompt_lens_ptr, + page_ids_ptr, + paged_kv_indptr_ptr, + page_ids_len, + indptr_len, + num_seq_slots, + num_pages, + num_layers, + local_layer, + page_size, + pscale_s0, + pool_s0, + pool_s1, + lc_s0, + lc_s1, + HEAD_D: tl.constexpr, + POOL_HEAD_D: tl.constexpr, + FP4_BLOCK: tl.constexpr, + PAGE_SLOTS: tl.constexpr, + BLOCK_D: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, +): + seq_idx = tl.program_id(0) + if (local_layer < 0) | (local_layer >= num_layers): + return + if seq_idx + 1 >= indptr_len: + return + + kv_len = tl.load(kv_lens_ptr + seq_idx) + gen_len = tl.load(prompt_lens_ptr + seq_idx) + if (kv_len <= 0) | (gen_len <= 0): + return + first_new_pos = kv_len - gen_len + active_page = (kv_len - 1) // page_size + # New tokens must all land in the active page (boundary-cross guarded by + # the Python dispatch; bail defensively here too). + if first_new_pos // page_size != active_page: + return + page_pos_start = active_page * page_size + fill = kv_len - page_pos_start + + page_start = tl.load(paged_kv_indptr_ptr + seq_idx).to(tl.int64) + page_end = tl.load(paged_kv_indptr_ptr + seq_idx + 1).to(tl.int64) + physical_page_offset = page_start + active_page + if ( + (physical_page_offset < page_start) + | (physical_page_offset >= page_end) + | (physical_page_offset < 0) + | (physical_page_offset >= page_ids_len) + ): + return + physical_page = tl.load(page_ids_ptr + physical_page_offset).to(tl.int64) + if (physical_page < 0) | (physical_page >= num_pages): + return + seq_slot = tl.load(seq_slots_ptr + seq_idx).to(tl.int64) + if (seq_slot < 0) | (seq_slot >= num_seq_slots): + return + + offs_d = tl.arange(0, BLOCK_D) + mask_d = offs_d < HEAD_D + safe_d = tl.where(mask_d, offs_d, 0) + amax = 0.0 + for tile_idx in tl.range(0, PAGE_SLOTS // FP4_BLOCK): + token_offsets = tile_idx * FP4_BLOCK + tl.arange(0, FP4_BLOCK) + abs_pos = page_pos_start + token_offsets + valid = token_offsets < fill + from_latent = abs_pos >= first_new_pos + slot = abs_pos % PAGE_SLOTS + stage_vals = tl.load( + stage_pool_ptr + + seq_slot * pool_s0 + + local_layer * pool_s1 + + slot[:, None] * POOL_HEAD_D + + safe_d[None, :], + mask=valid[:, None] & (~from_latent)[:, None] & mask_d[None, :], + other=0.0, + ).to(tl.float32) + latent_tok = seq_idx * gen_len + (abs_pos - first_new_pos) + safe_latent = tl.where(valid & from_latent, latent_tok, 0).to(tl.int64) + latent_vals = tl.load( + latent_cache_ptr + safe_latent[:, None] * lc_s0 + safe_d[None, :] * lc_s1, + mask=valid[:, None] & from_latent[:, None] & mask_d[None, :], + other=0.0, + ).to(tl.float32) + vals = stage_vals + latent_vals + amax = tl.maximum(amax, tl.max(tl.abs(vals))) + + gscale = tl.where(amax > 0.0, P_GLOBAL_SCALE / amax, 1.0) + tl.store(page_scale_ptr + local_layer * pscale_s0 + physical_page, gscale) + + +@triton.jit +def _fp4_mla_page_requant_gen_kernel( + kv_cache_ptr, + sf_cache_ptr, + v_sf_ptr, + stage_pool_ptr, + latent_cache_ptr, + page_scale_ptr, + seq_slots_ptr, + kv_lens_ptr, + prompt_lens_ptr, + page_ids_ptr, + paged_kv_indptr_ptr, + page_ids_len, + indptr_len, + num_seq_slots, + num_pages, + num_layers, + local_layer, + page_size, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + pool_s0, + pool_s1, + lc_s0, + lc_s1, + vsf_s0, + vsf_s1, + pscale_s0, + HEAD_D: tl.constexpr, + POOL_HEAD_D: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SLOTS: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, +): + seq_idx = tl.program_id(0) + tile_idx = tl.program_id(1) + dim_block = tl.program_id(2) + if (local_layer < 0) | (local_layer >= num_layers): + return + if seq_idx + 1 >= indptr_len: + return + + kv_len = tl.load(kv_lens_ptr + seq_idx) + gen_len = tl.load(prompt_lens_ptr + seq_idx) + if gen_len <= 0: + return + first_new_pos = kv_len - gen_len + active_page = (kv_len - 1) // page_size + if first_new_pos // page_size != active_page: + return + # Re-quantize every tile of the active page (page_gscale changed), not just + # the tiles that the new tokens touch. + block_base_pos = active_page * page_size + tile_idx * FP4_BLOCK + if block_base_pos >= kv_len: + return + + page_pos = block_base_pos - active_page * page_size + page_start = tl.load(paged_kv_indptr_ptr + seq_idx).to(tl.int64) + page_end = tl.load(paged_kv_indptr_ptr + seq_idx + 1).to(tl.int64) + physical_page_offset = page_start + active_page + if ( + (page_pos < 0) + | (page_pos >= page_size) + | (physical_page_offset < page_start) + | (physical_page_offset >= page_end) + | (physical_page_offset < 0) + | (physical_page_offset >= page_ids_len) + ): + return + physical_page = tl.load(page_ids_ptr + physical_page_offset).to(tl.int64) + if (physical_page < 0) | (physical_page >= num_pages): + return + seq_slot = tl.load(seq_slots_ptr + seq_idx).to(tl.int64) + if (seq_slot < 0) | (seq_slot >= num_seq_slots): + return + + page_gscale = tl.load(page_scale_ptr + local_layer * pscale_s0 + physical_page) + + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + token_offsets = tl.arange(0, FP4_BLOCK) + even_d = dim_block * FP4_BLOCK + byte_offsets * 2 + odd_d = even_d + 1 + all_d = dim_block * FP4_BLOCK + tl.arange(0, FP4_BLOCK) + mask_even_d = even_d < HEAD_D + mask_odd_d = odd_d < HEAD_D + mask_all_d = all_d < HEAD_D + safe_even_d = tl.where(mask_even_d, even_d, 0) + safe_odd_d = tl.where(mask_odd_d, odd_d, 0) + safe_all_d = tl.where(mask_all_d, all_d, 0) + + abs_positions = block_base_pos + token_offsets + valid_tokens = abs_positions < kv_len + from_latent = abs_positions >= first_new_pos + slot = abs_positions % PAGE_SLOTS + new_token_offsets = abs_positions - first_new_pos + latent_tokens = seq_idx * gen_len + new_token_offsets + safe_latent_tokens = tl.where(valid_tokens & from_latent, latent_tokens, 0).to(tl.int64) + + stage_even = tl.load( + stage_pool_ptr + + seq_slot * pool_s0 + + local_layer * pool_s1 + + slot[:, None] * POOL_HEAD_D + + safe_even_d[None, :], + mask=valid_tokens[:, None] & (~from_latent)[:, None] & mask_even_d[None, :], + other=0.0, + ).to(tl.float32) + stage_odd = tl.load( + stage_pool_ptr + + seq_slot * pool_s0 + + local_layer * pool_s1 + + slot[:, None] * POOL_HEAD_D + + safe_odd_d[None, :], + mask=valid_tokens[:, None] & (~from_latent)[:, None] & mask_odd_d[None, :], + other=0.0, + ).to(tl.float32) + latent_even = tl.load( + latent_cache_ptr + safe_latent_tokens[:, None] * lc_s0 + safe_even_d[None, :] * lc_s1, + mask=valid_tokens[:, None] & from_latent[:, None] & mask_even_d[None, :], + other=0.0, + ).to(tl.float32) + latent_odd = tl.load( + latent_cache_ptr + safe_latent_tokens[:, None] * lc_s0 + safe_odd_d[None, :] * lc_s1, + mask=valid_tokens[:, None] & from_latent[:, None] & mask_odd_d[None, :], + other=0.0, + ).to(tl.float32) + even_values = stage_even + latent_even + odd_values = stage_odd + latent_odd + + amax_per_token = tl.maximum( + tl.max(tl.abs(even_values), axis=1), + tl.max(tl.abs(odd_values), axis=1), + ) + tile_amax = tl.max(amax_per_token, axis=0) + # K consumes scales as [token, dim-block], V as [dim, token-block]. Only the + # compressed-KV prefix has both views; tail K-only dims keep K's per-token + # scale. page_gscale is the shared per-page global scale. + shared_tile = dim_block * FP4_BLOCK < V_HEAD_D + tile_scale = tl.where(tile_amax > 0.0, tile_amax / 6.0, 1.0) + token_scale = tl.where(amax_per_token > 0.0, amax_per_token / 6.0, 1.0) + local_scale = tl.where(shared_tile, tile_scale, token_scale) + stored_scale = local_scale * page_gscale + v_stored_scale = tile_scale * page_gscale + + low = _fp4_e2m1_quantize(even_values / local_scale[:, None]) + high = _fp4_e2m1_quantize(odd_values / local_scale[:, None]) + packed = low | (high << 4) + + packed_cols = dim_block * (FP4_BLOCK // 2) + byte_offsets + page_positions = page_pos + token_offsets + kv_base = physical_page * kv_s0 + tl.store( + kv_cache_ptr + kv_base + page_positions[:, None] * kv_s2 + packed_cols[None, :] * kv_s4, + packed, + mask=valid_tokens[:, None] & mask_even_d[None, :], + ) + + k_sf_offsets = _fp4_mla_swizzled_sf_offset(page_positions, dim_block, SF_PER_TOKEN) + tl.store(sf_cache_ptr + physical_page * sf_s0 + k_sf_offsets, stored_scale, mask=valid_tokens) + + token_scale_col = page_pos // FP4_BLOCK + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_all_d, token_scale_col, SF_PER_PAGE) + v_sf_base = tl.cast(local_layer, tl.int64) * tl.cast( + vsf_s0, tl.int64 + ) + physical_page * tl.cast(vsf_s1, tl.int64) + tl.store( + v_sf_ptr + v_sf_base + sf_offsets.to(tl.int64), + v_stored_scale, + mask=mask_all_d & (all_d < V_HEAD_D), + ) diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index c34564c7fd1a..1b46fa5c09bc 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -38,7 +38,7 @@ from ..utils import (compute_swizzled_sf_shape, get_global_attrs, get_model_extra_attrs) -from .fp4_mla_kv import HP_BLOCK_SIZE, update_hp_kv_for_fp4_mla +from .fp4_mla import HP_BLOCK_SIZE, update_hp_kv_for_fp4_mla from .interface import (AttentionBackend, AttentionForwardArgs, AttentionInputType, AttentionMask, AttentionMetadata, KVCacheParams, MLAParams, PositionalEmbeddingParams, @@ -163,6 +163,7 @@ class TrtllmAttentionMetadata(AttentionMetadata): # Shape: [max_num_sequences, num_local_layers, kv_factor, HP_BLOCK_SIZE * head_dim] # Standalone tensor, not part of the block-based paged KV cache. high_precision_kv_pool: Optional[torch.Tensor] = None + fp4_mla_hp_snapshot_pool: Optional[torch.Tensor] = None # Ownership tracking: maps seq_slot to request_id that last wrote it. # Plain Python dict, updated during context phase, checked during decode. # Debug only; runs outside CUDA graph. @@ -475,6 +476,19 @@ def _post_init_with_buffers(self, buffers) -> None: dtype=torch.bfloat16, capture_graph=capture_graph, ) + if capture_graph: + self.fp4_mla_hp_snapshot_pool = self.get_empty( + buffers, + [ + self.max_num_sequences, num_local_layers, kv_factor, + HP_BLOCK_SIZE * head_dim + ], + cache_name="fp4_mla_hp_snapshot_pool", + dtype=torch.bfloat16, + capture_graph=capture_graph, + ) + else: + self.fp4_mla_hp_snapshot_pool = None logger.info( f"Allocated high-precision BF16 KV pool: shape=" f"{list(self.high_precision_kv_pool.shape)}, " @@ -1489,7 +1503,7 @@ def _update_high_precision_kv_for_fp4_mla( attention_input_type: AttentionInputType = AttentionInputType.mixed, ) -> None: """Thin wrapper over the shared HP-pool update helper (see - ``fp4_mla_kv.update_hp_kv_for_fp4_mla`` for the full contract).""" + ``fp4_mla.update_hp_kv_for_fp4_mla`` for the full contract).""" if attention_input_type == AttentionInputType.context_only: phase = "context" elif attention_input_type == AttentionInputType.generation_only: diff --git a/tensorrt_llm/_torch/modules/fla/utils.py b/tensorrt_llm/_torch/modules/fla/utils.py index 3a5f9c1725d8..3d1fc6f41ebe 100644 --- a/tensorrt_llm/_torch/modules/fla/utils.py +++ b/tensorrt_llm/_torch/modules/fla/utils.py @@ -276,7 +276,7 @@ def get_available_device() -> str: @lru_cache(maxsize=None) def _check_platform() -> Literal["nvidia", "amd", "intel", "musa"]: device = get_available_device() - if device == "cuda": + if device == "cuda" or device == "tileir": return "nvidia" elif device == "hip": return "amd" @@ -289,8 +289,9 @@ def _check_platform() -> Literal["nvidia", "amd", "intel", "musa"]: # For AMD GPUs, the triton backend is 'hip', while for Nvidia GPUs, the triton backend is 'cuda'. # However, the torch backend is 'cuda' for both Nvidia and AMD GPUs. # Therefore, we need to check the triton backend to determine the actual GPU vendor. -device = get_available_device() if get_available_device() != "hip" else "cuda" -device = "cuda" +device = get_available_device() +if device == "hip" or device == "tileir": + device = "cuda" device_torch_lib = getattr(torch, device) device_platform = _check_platform() diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index dc0c07a88803..4b95dd4095f1 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -2099,6 +2099,41 @@ def _preprocess_inputs(self, inputs: Dict[str, Any]): previous_kv_lens_offsets_cuda[:num_gen_requests] ) inputs['attn_metadata'].on_update_kv_lens() + elif getattr(inputs['attn_metadata'], 'kv_lens_cuda_runtime', + None) is not None: + # FlashInfer NVFP4 MLA does not expose kv_lens_cuda; it keeps + # the per-sequence KV length in kv_lens_cuda_runtime and + # precomputes the FP4 KV-cache write positions / batch indices + # from it during prepare(). The overlap scheduler builds the + # generation metadata from the all-draft-accepted estimate + # (num_cached_tokens_per_seq = past_seen + runtime_draft_len + + # 1), so without the same previous_kv_lens_offsets correction + # the FP4 KV scatter and the BF16 HP-pool overlay would write + # at over-estimated positions whenever the previous MTP step + # rejected a draft token, corrupting the KV cache. Apply the + # correction, then rebuild positions/batch indices from the + # corrected lengths so the captured graph replays them. + md = inputs['attn_metadata'] + if num_ctx_requests >= num_chunked_ctx_requests and num_chunked_ctx_requests > 0: + md.kv_lens_cuda_runtime[ + num_ctx_requests - + num_chunked_ctx_requests:num_ctx_requests] += ( + self. + previous_kv_lens_offsets_cuda[: + num_chunked_ctx_requests] + ) + else: + md.kv_lens_cuda_runtime[num_ctx_requests:num_seqs] += ( + self.previous_kv_lens_offsets_cuda[:num_gen_requests] + ) + md.on_update_kv_lens() + md._populate_fp4_mla_batch_indices_positions() + # Opt-in: also rebuild the read-side decode paging + # (num_blocks / paged_kv_indptr_decode / paged_kv_indices / + # last_page_len) from the corrected kv_lens, since prepare() + # derived those from the all-draft-accepted over-estimate. + if hasattr(md, "repage_fp4_mla_decode_from_kv_lens"): + md.repage_fp4_mla_decode_from_kv_lens() if self.guided_decoder is not None: self.guided_decoder.token_event.record() @@ -2146,6 +2181,27 @@ def _postprocess_inputs(self, inputs: Dict[str, Any]): self. previous_kv_lens_offsets_cuda[:num_gen_requests] ) + elif getattr(inputs['attn_metadata'], 'kv_lens_cuda_runtime', + None) is not None: + # Undo the FlashInfer NVFP4 MLA kv_lens correction applied in + # _preprocess_inputs so the captured graph re-applies it from + # the original (over-estimated) lengths on the post-capture + # replay. positions/batch indices are rebuilt by the captured + # _populate_fp4_mla_batch_indices_positions on every replay, so + # they do not need to be restored here. + md = inputs['attn_metadata'] + if num_ctx_requests >= num_chunked_ctx_requests and num_chunked_ctx_requests > 0: + md.kv_lens_cuda_runtime[ + num_ctx_requests - + num_chunked_ctx_requests:num_ctx_requests] -= ( + self. + previous_kv_lens_offsets_cuda[: + num_chunked_ctx_requests] + ) + else: + md.kv_lens_cuda_runtime[num_ctx_requests:num_seqs] -= ( + self.previous_kv_lens_offsets_cuda[:num_gen_requests] + ) def _get_all_rank_num_tokens(self, attn_metadata: AttentionMetadata): if self.enable_attention_dp: diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index 58e63d7a1184..7d3fc729cec4 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -210,6 +210,13 @@ def _enable_fp4_mla_attention() -> bool: ) +def _is_fp4_mla_flashinfer_attention_requested(llm_args: TorchLlmArgs, + kv_cache_config) -> bool: + return (llm_args.attn_backend == "FLASHINFER" + and _has_fp4_kv_cache(llm_args, kv_cache_config) + and _enable_fp4_mla_attention()) + + def _select_mla_tokens_per_block(config, model_config, kv_cache_config, tokens_per_block: int) -> int: if not is_mla(config): @@ -459,14 +466,36 @@ def create_py_executor( ) llm_args.disable_overlap_scheduler = True - # Check FLASHINFER compatibility with one-engine speculative decoding - if llm_args.attn_backend == "FLASHINFER": + if spec_config is not None and spec_config.spec_dec_mode.use_one_engine(): + if not spec_config.allow_advanced_sampling: + logger.warning( + f"Falling back to greedy decoding for {spec_config.decoding_type}. If you " + "want to use non-greedy sampling, please set allow_advanced_sampling=True." + ) + elif (spec_config.spec_dec_mode.is_mtp_eagle_one_model() + and not getattr(spec_config, "use_rejection_sampling", False)): + logger.warning( + "MTP-Eagle one-model advanced sampling is using strict " + "token-match acceptance. This can change the sampled output " + "distribution; set use_rejection_sampling=True to use exact " + "one-model speculative sampling.") + # Regular FlashInfer decode expects one query token per sequence. The + # FP4 MLA no-dequant path has its own linear-MTP handling, so allow that + # explicit configuration through. + fp4_mla_flashinfer = _is_fp4_mla_flashinfer_attention_requested( + llm_args, kv_cache_config) + if llm_args.attn_backend == "FLASHINFER" and not fp4_mla_flashinfer: raise ValueError( f"FLASHINFER attention backend is not supported with one-engine speculative " f"decoding mode '{spec_config.spec_dec_mode.name}'. The FLASHINFER backend's " f"decode path expects exactly 1 token per sequence, but one-engine speculative " f"decoding requires multiple tokens per sequence. Please use 'TRTLLM' attention " f"backend instead by setting attn_backend='TRTLLM'.") + if fp4_mla_flashinfer: + # FlashInfer metadata prepares page tables against one KV manager. + # Keep one-model MTP draft layers in the main manager so global + # draft layer ids are present in layer_offsets. + spec_config._allow_separate_draft_kv_cache = False if mm_encoder_only: llm_args.mm_encoder_only = True diff --git a/tensorrt_llm/_torch/speculative/eagle3.py b/tensorrt_llm/_torch/speculative/eagle3.py index 68dac4c5e9f0..32058e97afe5 100644 --- a/tensorrt_llm/_torch/speculative/eagle3.py +++ b/tensorrt_llm/_torch/speculative/eagle3.py @@ -18,7 +18,8 @@ from ..pyexecutor.sampler import TorchSampler from ..pyexecutor.scheduler import ScheduledRequests from .interface import SpecMetadata, SpecWorkerBase -from .mtp import MTPSampler, _select_mtp_position_ids +from .mtp import (MTPSampler, _repair_fp4_mla_hp_kv_after_mtp_acceptance, + _select_mtp_position_ids) from .sa_enhancer import SADraftEnhancer from .spec_tree_manager import SpecTreeManager @@ -675,6 +676,8 @@ def forward(self, # acceptance path (scans for thinking-phase tokens); ignored otherwise. accepted_tokens, num_accepted_tokens = self.sample_and_accept_draft_tokens( input_ids, logits, attn_metadata, spec_metadata) + _repair_fp4_mla_hp_kv_after_mtp_acceptance(attn_metadata, + num_accepted_tokens) # Mamba hybrid models need state updates after token acceptance because # the accepted token count affects which Mamba states are valid. The @@ -900,19 +903,32 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, has_kv_cache = inputs[ "attn_metadata"].kv_cache_manager is not None if has_kv_cache: - attn_metadata.host_request_types[:attn_metadata. - num_contexts].fill_(1) + host_request_types = getattr(attn_metadata, + "host_request_types", None) + if host_request_types is not None: + host_request_types[:attn_metadata. + num_contexts].fill_(1) attn_metadata.num_contexts = 0 + kv_lens_updated = False if hasattr(attn_metadata, 'kv_lens_cuda'): attn_metadata.kv_lens_cuda[num_contexts:batch_size] -= ( runtime_draft_len - num_accepted_tokens[num_contexts:]) attn_metadata.kv_lens_cuda[:num_contexts] += 1 + kv_lens_updated = True + elif getattr(attn_metadata, "kv_lens_cuda_runtime", + None) is not None: + attn_metadata.kv_lens_cuda_runtime[ + num_contexts:batch_size] -= ( + runtime_draft_len - + num_accepted_tokens[num_contexts:]) + attn_metadata.kv_lens_cuda_runtime[:num_contexts] += 1 + kv_lens_updated = True if has_kv_cache: self._prepare_flash_mla_generation_layout( attn_metadata, num_contexts, batch_size) - if hasattr(attn_metadata, 'kv_lens_cuda'): + if kv_lens_updated: attn_metadata.update_for_spec_dec() # Both Eagle3 and MTP Eagle drafters take ``draft_len + 1`` @@ -925,6 +941,10 @@ def _forward_linear_draft_loop(self, inputs, attn_metadata, spec_metadata, if hasattr(attn_metadata, 'kv_lens_cuda'): attn_metadata.kv_lens_cuda[:batch_size] += 1 attn_metadata.update_for_spec_dec() + elif getattr(attn_metadata, "kv_lens_cuda_runtime", + None) is not None: + attn_metadata.kv_lens_cuda_runtime[:batch_size] += 1 + attn_metadata.update_for_spec_dec() inputs = { "input_ids": new_draft_token, @@ -1009,7 +1029,7 @@ def _prepare_flash_mla_generation_layout(self, attn_metadata, num_contexts, attn_metadata.block_ids_per_seq[:batch_size, :].copy_( reorder_block_ids_per_seq, non_blocking=True) - @torch.compile(options={"max-autotune": True}) + # @torch.compile(options={"max-autotune": True}) def _get_local_max_and_combined(self, logits, mapping_lm_tp=None): local_max_values, local_argmax = torch.max(logits, dim=-1, keepdim=True) vocab_per_rank = logits.shape[-1] @@ -1024,7 +1044,7 @@ def _get_local_max_and_combined(self, logits, mapping_lm_tp=None): dim=-1).flatten(-2) return combined - @torch.compile(options={"max-autotune": True}) + # @torch.compile(options={"max-autotune": True}) def _get_draft_tokens_from_gathered(self, gathered): gathered_indices_float = gathered[..., 0::2] gathered_values_float = gathered[..., 1::2] @@ -1066,7 +1086,7 @@ def draft_sampler( else: return self._draft_sampler_greedy(logits) - @torch.compile(options={"max-autotune": True}) + # @torch.compile(options={"max-autotune": True}) def _topk_kernel(self, gen_logprobs, num_gens, mtp_num_modules, spec_metadata): topk_value, topk_indices = torch.topk(gen_logprobs, @@ -1080,7 +1100,7 @@ def _topk_kernel(self, gen_logprobs, num_gens, mtp_num_modules, num_gens, mtp_num_modules) return topk_value, topk_indices, draft_tokens - @torch.compile(options={"max-autotune": True}) + # @torch.compile(options={"max-autotune": True}) def _process_generation_logits(self, logits, num_contexts): gen_logits = logits[num_contexts:] gen_logprobs = torch.softmax(gen_logits, dim=-1) diff --git a/tensorrt_llm/_torch/speculative/mtp.py b/tensorrt_llm/_torch/speculative/mtp.py index 8d00e5bd11f3..6ec658280562 100644 --- a/tensorrt_llm/_torch/speculative/mtp.py +++ b/tensorrt_llm/_torch/speculative/mtp.py @@ -43,6 +43,22 @@ def _select_mtp_position_ids(position_ids: torch.Tensor, return position_ids[..., token_indices] +def _repair_fp4_mla_hp_kv_after_mtp_acceptance( + attn_metadata: Optional[AttentionMetadata], + num_accepted_tokens: torch.Tensor) -> None: + if attn_metadata is None or not getattr(attn_metadata, + "_fp4_mla_mtp_hp_snapshots", None): + return + + from ..attention_backend.fp4_mla import ( + repair_fp4_mla_hp_kv_for_mtp_rejection, + repair_fp4_mla_page_stage_for_mtp_rejection) + + repair_fp4_mla_hp_kv_for_mtp_rejection(attn_metadata, num_accepted_tokens) + repair_fp4_mla_page_stage_for_mtp_rejection(attn_metadata, + num_accepted_tokens) + + class MTPHiddenStatesManager(BaseResourceManager): def __init__(self, @@ -831,7 +847,10 @@ def sample_and_accept_draft_tokens( # Strict acceptance else: - if self.is_thop: + draft_tokens = spec_metadata.draft_tokens.reshape( + num_gens, mtp_num_modules) + if self.is_thop and not self._can_use_rejection_sampling( + spec_metadata, num_contexts): # Temporary buffer target_tokens_cache = torch.zeros(batch_size * (mtp_num_modules + 1), @@ -846,12 +865,7 @@ def sample_and_accept_draft_tokens( num_accepted_tokens, num_contexts, spec_metadata.runtime_draft_len) else: - # Reshape draft tokens for base implementation - draft_tokens = spec_metadata.draft_tokens.reshape( - num_gens, mtp_num_modules) - - # Use base implementation for strict acceptance - accepted_tokens, num_accepted_tokens = self._sample_and_accept_draft_tokens_base( + accepted_tokens, num_accepted_tokens = self._accept_draft_tokens( logits, draft_tokens, num_contexts, batch_size, spec_metadata) @@ -872,6 +886,8 @@ def sample_and_accept_draft_tokens( def change_attn_metadata(self, num_accepted_tokens: torch.Tensor, attn_metadata: AttentionMetadata): + _repair_fp4_mla_hp_kv_after_mtp_acceptance(attn_metadata, + num_accepted_tokens) self._prepare_attn_metadata_for_spec_dec(attn_metadata) batch_size = attn_metadata.num_seqs mtp_num_modules = self.spec_config.max_draft_len @@ -900,6 +916,11 @@ def change_attn_metadata(self, num_accepted_tokens: torch.Tensor, attn_metadata.kv_lens_cuda[num_contexts:batch_size].clamp_( min=mtp_num_modules) attn_metadata.on_update_kv_lens() + elif getattr(attn_metadata, "kv_lens_cuda_runtime", None) is not None: + attn_metadata.kv_lens_cuda_runtime[num_contexts:batch_size] -= ( + mtp_num_modules + 1 - + num_accepted_tokens[num_contexts:batch_size]) + attn_metadata.update_for_spec_dec() if attn_metadata.kv_cache_params is not None and not attn_metadata.is_cuda_graph: for i in range(num_contexts, batch_size): diff --git a/tensorrt_llm/_torch/speculative/utils.py b/tensorrt_llm/_torch/speculative/utils.py index 91f60243834c..12da82d69741 100644 --- a/tensorrt_llm/_torch/speculative/utils.py +++ b/tensorrt_llm/_torch/speculative/utils.py @@ -70,6 +70,10 @@ def get_spec_metadata(spec_config, mtp_num_modules=spec_config.max_draft_len, max_num_requests=max_num_requests, mtp_hidden_states_manager=spec_resource_manager, + allow_advanced_sampling=spec_config.allow_advanced_sampling, + use_rejection_sampling=use_rejection_sampling + and spec_config.spec_dec_mode.is_mtp_eagle_one_model(), + vocab_size=vocab_size, ) if spec_config.spec_dec_mode.is_mtp_eagle(): return Eagle3SpecMetadata( diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 609e11c5a27a..da7b58648bb9 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -4760,17 +4760,17 @@ def validate_speculative_config(self): exclude={"decoding_type"}) self.speculative_config = Eagle3DecodingConfig(**eagle_data) - if self.speculative_config.use_rejection_sampling and not isinstance( - self.speculative_config, Eagle3DecodingConfig): - # Rejection sampling is only wired up for Eagle3 one-model paths. - # Silently fall back for other spec types so the new default - # (True) does not break them. - # TODO: extend rejection sampling to the remaining speculative - # decoding paths (MTP / DraftTarget / PARD / DFlash / - # SaveHiddenStates / SA) and unify the dispatch in SpecMetadata - # so new spec algorithms get rejection sampling for free; once - # all paths are covered this whitelist guard can be removed. - self.speculative_config.use_rejection_sampling = False + if self.speculative_config.use_rejection_sampling: + is_supported_rejection_path = isinstance( + self.speculative_config, Eagle3DecodingConfig) or ( + isinstance(self.speculative_config, MTPDecodingConfig) + and self.speculative_config.spec_dec_mode. + is_mtp_eagle_one_model()) + if not is_supported_rejection_path: + raise ValueError( + "use_rejection_sampling is only supported for " + "PyTorch Eagle3 and MTP-Eagle one-model speculative " + "decoding paths.") if isinstance(self.speculative_config, PARDDecodingConfig): assert self.speculative_config.max_draft_len > 0, "PARD max_draft_len must be > 0" diff --git a/tests/unittest/_torch/attention/test_fp4_mla_kv.py b/tests/unittest/_torch/attention/test_fp4_mla.py similarity index 90% rename from tests/unittest/_torch/attention/test_fp4_mla_kv.py rename to tests/unittest/_torch/attention/test_fp4_mla.py index 252c5aced3a5..6732fe18eeda 100644 --- a/tests/unittest/_torch/attention/test_fp4_mla_kv.py +++ b/tests/unittest/_torch/attention/test_fp4_mla.py @@ -15,7 +15,7 @@ import torch import tensorrt_llm -from tensorrt_llm._torch.attention_backend.fp4_mla_kv import ( +from tensorrt_llm._torch.attention_backend.fp4_mla import ( FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, FLASHINFER_FP4_MLA_ATTENTION_ENV, FP4_BLOCK_SIZE, @@ -29,11 +29,11 @@ get_fp4_mla_v_scale_pool_size, get_fp4_mla_v_scale_pool_view, is_flashinfer_fp4_mla_attention_enabled, + repair_fp4_mla_hp_kv_for_mtp_rejection, run_fp4_mla_attention_decode, scatter_fp4_mla_kv_cache, update_hp_kv_for_fp4_mla, ) -from tensorrt_llm._torch.cute_dsl_utils import IS_CUTLASS_DSL_AVAILABLE from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings.executor import KvCacheConfig from tensorrt_llm.mapping import Mapping @@ -282,7 +282,7 @@ def _build_multi_seq_metadata(kv_cache_manager, *, seq_lens, page_size, num_laye ) -def _build_fp4_mla_attention_decode_case(*, seq_lens, num_heads, seed): +def _build_fp4_mla_attention_decode_case(*, seq_lens, num_heads, seed, query_len_per_seq=1): torch.manual_seed(seed) device = torch.device("cuda") @@ -335,13 +335,20 @@ def _build_fp4_mla_attention_decode_case(*, seq_lens, num_heads, seed): torch.cuda.synchronize() metadata.num_contexts = 0 + metadata.prompt_lens_cuda_runtime = torch.full( + (len(seq_lens),), query_len_per_seq, dtype=torch.int32, device=device + ) + metadata.prompt_lens_cpu_runtime = torch.full( + (len(seq_lens),), query_len_per_seq, dtype=torch.int32 + ) + num_queries = len(seq_lens) * query_len_per_seq q_nope = ( - torch.randn(len(seq_lens), num_heads, kv_lora_rank, dtype=torch.bfloat16, device=device) + torch.randn(num_queries, num_heads, kv_lora_rank, dtype=torch.bfloat16, device=device) * 0.25 ).clamp_(-1.0, 1.0) q_pe = ( torch.randn( - len(seq_lens), + num_queries, num_heads, qk_rope_head_dim, dtype=torch.bfloat16, @@ -406,34 +413,41 @@ def _fp4_mla_attention_decode_reference( indptr = metadata.paged_kv_indptr_decode.cpu().tolist() kv_lens = metadata.kv_lens_cuda_runtime.cpu().tolist() + num_seqs = metadata.num_seqs - metadata.num_contexts + query_len_per_seq = q_nope.shape[0] // num_seqs + max_pages = max(indptr[seq_idx + 1] - indptr[seq_idx] for seq_idx in range(num_seqs)) outputs = [] exact_probs = [] quantized_probs = [] - for seq_idx in range(metadata.num_seqs): + for seq_idx in range(num_seqs): kv_len = kv_lens[seq_idx] - cache = dequant_cache[indptr[seq_idx] : indptr[seq_idx + 1]].reshape(-1, head_dim)[:kv_len] - logical_k = _duplicate_tail_groups(cache.float(), FP4_MLA_Q_RESIDUAL_DIM) - q_start = seq_idx * num_heads - q = q_dequant[q_start : q_start + num_heads] - probs = torch.softmax(torch.matmul(q, logical_k.transpose(0, 1)) * sm_scale, dim=-1) - - if p_dequant is None: - p = probs - else: - p_pages = [] - for page_rel in range(indptr[seq_idx + 1] - indptr[seq_idx]): - page_start = page_rel * metadata.page_size - valid_tokens = max(min(kv_len - page_start, metadata.page_size), 0) - if valid_tokens == 0: - continue - compact_page = indptr[seq_idx] + page_rel - p_start = compact_page * num_heads - p_pages.append(p_dequant[p_start : p_start + num_heads, :valid_tokens]) - p = torch.cat(p_pages, dim=-1) - - exact_probs.append(probs) - quantized_probs.append(p) - outputs.append(torch.matmul(probs, cache[:, :kv_lora_rank].float())) + full_cache = dequant_cache[indptr[seq_idx] : indptr[seq_idx + 1]].reshape(-1, head_dim) + for query_offset in range(query_len_per_seq): + query_idx = seq_idx * query_len_per_seq + query_offset + effective_kv_len = kv_len - (query_len_per_seq - 1 - query_offset) + cache = full_cache[:effective_kv_len] + logical_k = _duplicate_tail_groups(cache.float(), FP4_MLA_Q_RESIDUAL_DIM) + q_start = query_idx * num_heads + q = q_dequant[q_start : q_start + num_heads] + probs = torch.softmax(torch.matmul(q, logical_k.transpose(0, 1)) * sm_scale, dim=-1) + + if p_dequant is None: + p = probs + else: + p_pages = [] + for page_rel in range(indptr[seq_idx + 1] - indptr[seq_idx]): + page_start = page_rel * metadata.page_size + valid_tokens = max(min(effective_kv_len - page_start, metadata.page_size), 0) + if valid_tokens == 0: + continue + p_page = query_idx * max_pages + page_rel + p_start = p_page * num_heads + p_pages.append(p_dequant[p_start : p_start + num_heads, :valid_tokens]) + p = torch.cat(p_pages, dim=-1) + + exact_probs.append(probs) + quantized_probs.append(p) + outputs.append(torch.matmul(probs, cache[:, :kv_lora_rank].float())) return torch.stack(outputs, dim=0), exact_probs, quantized_probs @@ -460,6 +474,7 @@ def _assert_fp4_mla_attention_decode_accuracy( seq_lens: list[int], seed: int, check_probs: bool, + query_len_per_seq: int = 1, ) -> None: monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, backend) @@ -475,6 +490,7 @@ def _assert_fp4_mla_attention_decode_accuracy( seq_lens=seq_lens, num_heads=num_heads, seed=seed, + query_len_per_seq=query_len_per_seq, ) try: output = torch.empty_like(q_nope) @@ -814,6 +830,73 @@ def test_fp4_mla_hp_overlay_generation_phase(page_size: int, ctx_tokens: int, mo kv_cache_manager.shutdown() +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("preallocated_snapshot", [False, True]) +def test_fp4_mla_hp_pool_restores_rejected_linear_mtp_tokens(preallocated_snapshot: bool): + device = torch.device("cuda") + head_dim = 8 + old_len = 30 + gen_len = 4 + accepted_len = 2 + + initial_hp = ( + torch.arange(HP_BLOCK_SIZE * head_dim, dtype=torch.float32, device=device) + .reshape(HP_BLOCK_SIZE, head_dim) + .to(torch.bfloat16) + ) + hp_pool = initial_hp.reshape(1, 1, 1, HP_BLOCK_SIZE * head_dim).clone() + gen_latent = ( + torch.arange(gen_len * head_dim, dtype=torch.float32, device=device) + .reshape(gen_len, head_dim) + .add_(1000.0) + .to(torch.bfloat16) + ) + metadata = SimpleNamespace( + high_precision_kv_pool=hp_pool, + hp_pool_owners={0: 0}, + seq_slots=torch.zeros(1, dtype=torch.int32, device=device), + seq_slots_cpu=torch.zeros(1, dtype=torch.int32, device="cpu"), + kv_lens_cuda_runtime=torch.tensor([old_len + gen_len], dtype=torch.int32, device=device), + prompt_lens_cuda_runtime=torch.tensor([gen_len], dtype=torch.int32, device=device), + prompt_lens_cpu_runtime=torch.tensor([gen_len], dtype=torch.int32), + batch_indices=torch.zeros(gen_len, dtype=torch.int32, device=device), + positions=torch.arange(old_len, old_len + gen_len, dtype=torch.int32, device=device), + num_contexts=0, + num_seqs=1, + request_ids=[0], + is_cuda_graph=False, + is_warmup=False, + ) + if preallocated_snapshot: + metadata.fp4_mla_hp_snapshot_pool = torch.empty_like(hp_pool) + + update_hp_kv_for_fp4_mla(metadata, gen_latent, local_layer=0, phase="generation") + repair_fp4_mla_hp_kv_for_mtp_rejection( + metadata, + torch.tensor([accepted_len], dtype=torch.int32, device=device), + ) + + hp_view = hp_pool.view(HP_BLOCK_SIZE, head_dim) + accepted_slots = [(old_len + idx) % HP_BLOCK_SIZE for idx in range(accepted_len)] + rejected_slots = [(old_len + idx) % HP_BLOCK_SIZE for idx in range(accepted_len, gen_len)] + torch.testing.assert_close( + hp_view[accepted_slots].float(), + gen_latent[:accepted_len].float(), + atol=0.0, + rtol=0.0, + ) + torch.testing.assert_close( + hp_view[rejected_slots].float(), + initial_hp[rejected_slots].float(), + atol=0.0, + rtol=0.0, + ) + if preallocated_snapshot: + assert getattr(metadata, "_fp4_mla_mtp_hp_snapshots") + else: + assert getattr(metadata, "_fp4_mla_mtp_hp_snapshots") is None + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @pytest.mark.parametrize("attention_env", ["0", "1"], ids=["linear_sf", "swizzled_sf"]) def test_fp4_mla_scatter_last_page_no_oob(attention_env: str, monkeypatch): @@ -1298,6 +1381,20 @@ def test_fp4_mla_attention_decode_multi_seq_matches_reference(monkeypatch): ) +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 tensor cores") +def test_fp4_mla_attention_decode_linear_mtp_matches_reference(monkeypatch): + """Linear MTP query rows must use per-query causal KV lengths.""" + _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + backend="triton", + num_heads=5, + seq_lens=[32, 128], + seed=11, + check_probs=True, + query_len_per_seq=3, + ) + + @pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") def test_fp4_mla_attention_decode_cutile_matches_reference(monkeypatch): """CuTile decode backend must preserve the FP4 MLA residual-tail contract.""" @@ -1312,16 +1409,16 @@ def test_fp4_mla_attention_decode_cutile_matches_reference(monkeypatch): @pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") -@pytest.mark.skipif(not IS_CUTLASS_DSL_AVAILABLE, reason="requires CuTe DSL") -def test_fp4_mla_attention_decode_cute_dsl_matches_reference(monkeypatch): - """CuTe DSL decode backend must preserve the FP4 MLA decode contract.""" +def test_fp4_mla_attention_decode_cutile_linear_mtp_matches_reference(monkeypatch): + """CuTile linear MTP rows must use per-query causal KV lengths.""" _assert_fp4_mla_attention_decode_accuracy( monkeypatch, - backend="cute_dsl", - num_heads=2, - seq_lens=[32], - seed=10, + backend="cutile", + num_heads=128, + seq_lens=[32, 128], + seed=13, check_probs=False, + query_len_per_seq=3, ) From f6d19cdc48337bea807d750189fd4b0b62bf9f44 Mon Sep 17 00:00:00 2001 From: Tracin <10434017+Tracin@users.noreply.github.com> Date: Tue, 2 Jun 2026 05:52:53 -0700 Subject: [PATCH 07/11] Apply FP4 MLA v2 integration fixes Reuse FP4 MLA V-packed storage across layers Handle partial FP4 MLA page-stat groups Tune FP4 MLA PV block size by batch Optimize FP4 MLA MTP decode path Add experimental FP4 MLA MTP fused QKPV path Optimize FP4 MLA MTP final-page path Fix cutile bug. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> --- bench_fp4_mla_decode.py | 20 +- .../_torch/attention_backend/fp4_mla.py | 280 +- .../attention_backend/fp4_mla_cutile.py | 6211 +++++++++++++---- .../unittest/_torch/attention/test_fp4_mla.py | 166 + 4 files changed, 5402 insertions(+), 1275 deletions(-) diff --git a/bench_fp4_mla_decode.py b/bench_fp4_mla_decode.py index 109d134ddc15..ffd4360362ee 100644 --- a/bench_fp4_mla_decode.py +++ b/bench_fp4_mla_decode.py @@ -60,7 +60,7 @@ def _bench(fn, warmup=0, iters=1): def _seq_lens_for_batch(batch, seq): - return [seq + 128 * (batch_idx // 10) for batch_idx in range(batch)] + return [seq] * batch def _seq_label(seq_lens): @@ -102,7 +102,7 @@ def _trtllm_mla_io_bytes(seq_lens, heads, kv_lora_rank, qk_rope_head_dim, q_len= return q + kv + out -def run_one_trtllm(batch, seq, heads, q_len=1): +def run_one_trtllm(batch, seq, heads, q_len=1, warmup=0, iters=1): """Fp8 baseline using the FlashInfer trtllm-gen MLA decode kernel. Feeds fp8 (e4m3) Q and KV cache so the kernel uses fp8 tensor cores @@ -166,7 +166,7 @@ def run(): run() torch.cuda.synchronize() - avg_ms = _bench(run) + avg_ms = _bench(run, warmup=warmup, iters=iters) qk_dim = kv_lora_rank + qk_rope_head_dim pv_dim = kv_lora_rank flops = 2 * q_len * heads * total_seq * (qk_dim + pv_dim) @@ -186,7 +186,7 @@ def run(): return avg_ms -def run_one(batch, seq, heads, backend, q_len=1): +def run_one(batch, seq, heads, backend, q_len=1, warmup=0, iters=1): os.environ[FLASHINFER_FP4_MLA_ATTENTION_ENV] = "1" os.environ[FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV] = backend seq_lens = _seq_lens_for_batch(batch, seq) @@ -219,7 +219,7 @@ def run(): run() torch.cuda.synchronize() - avg_ms = _bench(run) + avg_ms = _bench(run, warmup=warmup, iters=iters) qk_dim = kv_lora_rank + qk_rope_head_dim + FP4_MLA_Q_RESIDUAL_DIM pv_dim = kv_lora_rank flops = 2 * q_len * heads * total_seq * (qk_dim + pv_dim) @@ -244,7 +244,7 @@ def run(): def main(): p = argparse.ArgumentParser() - p.add_argument("--batch", type=int, default=None) + p.add_argument("--batch", type=int, nargs="+", default=None) p.add_argument("--seq", type=int, default=30080) p.add_argument("--heads", type=int, default=128) p.add_argument( @@ -261,9 +261,11 @@ def main(): choices=BACKEND_CHOICES, help=("Backend to benchmark; default runs all fast backends."), ) + p.add_argument("--warmup", type=int, default=0) + p.add_argument("--iters", type=int, default=1) args = p.parse_args() - batches = [args.batch] if args.batch else [16, 30, 60, 120, 200, 300] + batches = args.batch if args.batch else [16, 30, 60, 120, 200, 300] if args.backend: backends = [args.backend] else: @@ -272,9 +274,9 @@ def main(): for b in batches: for be in backends: if be == "trtllm_fp8": - run_one_trtllm(b, args.seq, args.heads, args.q_len) + run_one_trtllm(b, args.seq, args.heads, args.q_len, args.warmup, args.iters) else: - run_one(b, args.seq, args.heads, be, args.q_len) + run_one(b, args.seq, args.heads, be, args.q_len, args.warmup, args.iters) if __name__ == "__main__": diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla.py b/tensorrt_llm/_torch/attention_backend/fp4_mla.py index 4006a588fde1..4cdb7918e54c 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla.py @@ -869,7 +869,10 @@ def scatter_fp4_mla_kv_cache( if phase == "context": v_pack_page_ids = metadata.paged_kv_indices else: - v_pack_page_ids = metadata.paged_kv_indices[metadata.num_context_blocks :] + num_gen_blocks = metadata.num_generation_blocks + v_pack_page_ids = metadata.paged_kv_indices[ + metadata.num_context_blocks : metadata.num_context_blocks + num_gen_blocks + ] _maybe_update_cutile_v_packed_cache( metadata, layer_idx, @@ -877,6 +880,8 @@ def scatter_fp4_mla_kv_cache( v_pack_page_ids, v_head_dim=v_head_dim, page_size=metadata.page_size, + local_layer=local_layer, + v_sf=v_sf[local_layer], ) return @@ -1222,7 +1227,30 @@ def _cutile_persistent_v_pack_enabled() -> bool: ) +def _cutile_shared_v_pack_storage_enabled() -> bool: + return os.getenv("TRTLLM_FP4_MLA_SHARE_V_PACK_STORAGE", "1").lower() not in ( + "0", + "false", + "no", + "off", + ) + + +def _select_cutile_block_v(num_gen_seqs: int, query_len_per_seq: int = 1) -> int: + env_block_v = _env_int("TRTLLM_FP4_MLA_BLOCK_V") + if env_block_v is not None: + return env_block_v + threshold = _env_int("TRTLLM_FP4_MLA_BLOCK_V_AUTO_THRESHOLD") + if threshold is None: + threshold = 60 + if query_len_per_seq == 1 and num_gen_seqs >= threshold: + return 256 + return 128 + + def _cutile_v_packed_attr(layer_idx: int) -> str: + if _cutile_shared_v_pack_storage_enabled(): + return "_fp4_mla_attention_v_packed_buf" return f"_fp4_mla_attention_v_packed_buf_l{layer_idx}" @@ -1230,15 +1258,130 @@ def _cutile_v_packed_valid_attr(layer_idx: int) -> str: return f"_fp4_mla_attention_v_packed_valid_l{layer_idx}" +def _cutile_shared_v_packed_valid_attr() -> str: + return "_fp4_mla_attention_v_packed_valid_tag" + + def _cutile_v_packed_shape( kv_cache: torch.Tensor, v_head_dim: int, page_size: int, + block_v: int = 128, ) -> tuple[int, int]: - block_v = 128 return (kv_cache.shape[0] * _ceil_div(v_head_dim, block_v) * block_v, page_size // 2) +def _cutile_v_packed_cache_tag( + layer_idx: int, + kv_cache: torch.Tensor, + *, + v_head_dim: int, + page_size: int, + local_layer: Optional[int] = None, + v_sf: Optional[torch.Tensor] = None, + page_ids: Optional[torch.Tensor] = None, + block_v: int = 128, +) -> tuple[Any, ...]: + v_sf_tag = ( + None + if v_sf is None + else ( + int(v_sf.data_ptr()), + str(v_sf.device), + str(v_sf.dtype), + tuple(int(dim) for dim in v_sf.shape), + tuple(int(stride) for stride in v_sf.stride()), + ) + ) + page_ids_tag = ( + None + if page_ids is None + else ( + int(page_ids.data_ptr()), + str(page_ids.device), + str(page_ids.dtype), + tuple(int(dim) for dim in page_ids.shape), + tuple(int(stride) for stride in page_ids.stride()), + ) + ) + return ( + int(layer_idx), + None if local_layer is None else int(local_layer), + int(kv_cache.data_ptr()), + str(kv_cache.device), + str(kv_cache.dtype), + tuple(int(dim) for dim in kv_cache.shape), + tuple(int(stride) for stride in kv_cache.stride()), + int(v_head_dim), + int(page_size), + int(block_v), + v_sf_tag, + page_ids_tag, + ) + + +def _set_cutile_v_packed_cache_valid( + metadata: Any, + layer_idx: int, + kv_cache: torch.Tensor, + *, + v_head_dim: int, + page_size: int, + local_layer: Optional[int] = None, + v_sf: Optional[torch.Tensor] = None, + page_ids: Optional[torch.Tensor] = None, + block_v: int = 128, +) -> None: + valid_attr = ( + _cutile_shared_v_packed_valid_attr() + if _cutile_shared_v_pack_storage_enabled() + else _cutile_v_packed_valid_attr(layer_idx) + ) + setattr( + metadata, + valid_attr, + _cutile_v_packed_cache_tag( + layer_idx, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=page_ids, + ), + ) + + +def _is_cutile_v_packed_cache_valid( + metadata: Any, + layer_idx: int, + kv_cache: torch.Tensor, + *, + v_head_dim: int, + page_size: int, + local_layer: Optional[int] = None, + v_sf: Optional[torch.Tensor] = None, + page_ids: Optional[torch.Tensor] = None, + block_v: int = 128, +) -> bool: + valid_attr = ( + _cutile_shared_v_packed_valid_attr() + if _cutile_shared_v_pack_storage_enabled() + else _cutile_v_packed_valid_attr(layer_idx) + ) + return getattr(metadata, valid_attr, None) == _cutile_v_packed_cache_tag( + layer_idx, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=page_ids, + ) + + def _maybe_update_cutile_v_packed_cache( metadata: Any, layer_idx: int, @@ -1247,10 +1390,14 @@ def _maybe_update_cutile_v_packed_cache( *, v_head_dim: int, page_size: int, + local_layer: Optional[int] = None, + v_sf: Optional[torch.Tensor] = None, ) -> None: if not _cutile_persistent_v_pack_enabled(): return - if v_head_dim % 128 != 0 or page_size != FP4_MLA_TOKENS_PER_BLOCK: + num_gen_seqs = getattr(metadata, "num_seqs", 0) - getattr(metadata, "num_contexts", 0) + block_v = _select_cutile_block_v(num_gen_seqs) + if block_v not in (128, 256) or v_head_dim % block_v != 0 or page_size != FP4_MLA_TOKENS_PER_BLOCK: return if page_ids.numel() == 0: return @@ -1261,7 +1408,7 @@ def _maybe_update_cutile_v_packed_cache( v_packed = _ensure_workspace_tensor( metadata, attr_name, - _cutile_v_packed_shape(kv_cache, v_head_dim, page_size), + _cutile_v_packed_shape(kv_cache, v_head_dim, page_size, block_v), dtype=torch.uint8, device=kv_cache.device, ) @@ -1271,9 +1418,19 @@ def _maybe_update_cutile_v_packed_cache( page_ids, v_head_dim=v_head_dim, page_size=page_size, - block_v=128, + block_v=block_v, + ) + _set_cutile_v_packed_cache_valid( + metadata, + layer_idx, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=page_ids, ) - setattr(metadata, _cutile_v_packed_valid_attr(layer_idx), True) def _get_cutile_v_packed_cache( @@ -1283,13 +1440,27 @@ def _get_cutile_v_packed_cache( *, v_head_dim: int, page_size: int, + local_layer: Optional[int] = None, + v_sf: Optional[torch.Tensor] = None, + page_ids: Optional[torch.Tensor] = None, + block_v: int = 128, ) -> Optional[torch.Tensor]: if not _cutile_persistent_v_pack_enabled(): return None - if not bool(getattr(metadata, _cutile_v_packed_valid_attr(layer_idx), False)): + if not _is_cutile_v_packed_cache_valid( + metadata, + layer_idx, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=page_ids, + ): return None v_packed = getattr(metadata, _cutile_v_packed_attr(layer_idx), None) - expected_shape = _cutile_v_packed_shape(kv_cache, v_head_dim, page_size) + expected_shape = _cutile_v_packed_shape(kv_cache, v_head_dim, page_size, block_v) if ( v_packed is None or v_packed.dtype != torch.uint8 @@ -1331,6 +1502,30 @@ def _infer_cutile_assume_full_pages(metadata: Any, max_pages: int, page_size: in start = metadata.num_contexts end = metadata.num_seqs + block_counts = _host_int_list(getattr(metadata, "num_blocks", None), start, end) + if block_counts is not None and ( + not block_counts or min(block_counts) != max_pages or max(block_counts) != max_pages + ): + return False + + kv_lens_cuda = getattr(metadata, "kv_lens_cuda_runtime", None) + if isinstance(kv_lens_cuda, torch.Tensor): + cache_key = ( + start, + end, + max_pages, + page_size, + tuple(block_counts) if block_counts is not None else None, + kv_lens_cuda.data_ptr(), + ) + cache = getattr(metadata, "_fp4_mla_cutile_full_pages_cache", None) + if cache is not None and cache[0] == cache_key: + return bool(cache[1]) + kv_lens = [int(item) for item in kv_lens_cuda[start:end].detach().cpu().tolist()] + result = bool(kv_lens) and min(kv_lens) == max(kv_lens) == max_pages * page_size + setattr(metadata, "_fp4_mla_cutile_full_pages_cache", (cache_key, result)) + return result + kv_cache_params = getattr(metadata, "kv_cache_params", None) cached_token_lens = _host_int_list( getattr(kv_cache_params, "num_cached_tokens_per_seq", None), @@ -1955,33 +2150,48 @@ def run_fp4_mla_attention_decode( dtype=torch.float32, device=q_nope.device, ) - assume_full_pages = ( - _infer_cutile_assume_full_pages( - metadata, - max_pages, - metadata.page_size, - ) - and query_len_per_seq == 1 + cutile_storage_full_pages = _infer_cutile_assume_full_pages( + metadata, + max_pages, + metadata.page_size, ) + assume_full_pages = cutile_storage_full_pages and query_len_per_seq == 1 assume_valid_pages = False - cutile_block_h = _env_int("TRTLLM_FP4_MLA_BLOCK_H") or 128 - cutile_block_v = _env_int("TRTLLM_FP4_MLA_BLOCK_V") or 128 cutile_num_gen_seqs = num_queries // query_len_per_seq + cutile_block_h = _env_int("TRTLLM_FP4_MLA_BLOCK_H") or 128 + cutile_block_v = _select_cutile_block_v( + cutile_num_gen_seqs, + query_len_per_seq=query_len_per_seq, + ) + cutile_prepack_v_env = os.environ.get("TRTLLM_FP4_MLA_PREPACK_V") + cutile_storage_valid_pages = ( + assume_valid_pages + or ( + cutile_storage_full_pages + and src_page_ids.numel() == cutile_num_gen_seqs * max_pages + ) + ) + cutile_allow_qlen_prepack_v = query_len_per_seq > 1 and cutile_prepack_v_env != "0" cutile_assume_valid_pages = assume_valid_pages or ( assume_full_pages and src_page_ids.numel() == cutile_num_gen_seqs * max_pages ) cutile_auto_prepack_v = ( hasattr(tl, "make_tensor_descriptor") and num_heads % cutile_block_h == 0 - and assume_full_pages - and cutile_assume_valid_pages + and ( + assume_full_pages + or (cutile_allow_qlen_prepack_v and cutile_storage_full_pages) + ) + and ( + cutile_assume_valid_pages + or (cutile_allow_qlen_prepack_v and cutile_storage_valid_pages) + ) and kv_lora_rank == 512 and metadata.page_size == FP4_MLA_TOKENS_PER_BLOCK and cutile_block_h in (64, 128) - and cutile_block_v == 128 + and cutile_block_v in (128, 256) and metadata.page_size // FP4_BLOCK_SIZE == 8 ) - cutile_prepack_v_env = os.environ.get("TRTLLM_FP4_MLA_PREPACK_V") cutile_prepack_v_for_pv = ( cutile_auto_prepack_v if cutile_prepack_v_env is None @@ -1994,6 +2204,10 @@ def run_fp4_mla_attention_decode( kv_cache, v_head_dim=kv_lora_rank, page_size=metadata.page_size, + block_v=cutile_block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=src_page_ids, ) if cutile_auto_prepack_v else None @@ -2012,6 +2226,12 @@ def run_fp4_mla_attention_decode( dtype=torch.uint8, device=q_nope.device, ) + mark_cutile_v_packed_cache_valid = bool( + cutile_prepack_v_for_pv + and v_packed is not None + and _cutile_persistent_v_pack_enabled() + and _cutile_shared_v_pack_storage_enabled() + ) _fp4_mla_debug( "attention decode cutile launch: " f"num_queries={num_queries} query_len_per_seq={query_len_per_seq} " @@ -2040,7 +2260,13 @@ def run_fp4_mla_attention_decode( q_residual_dim=q_residual_dim, max_pages=max_pages, query_len_per_seq=query_len_per_seq, + block_v=cutile_block_v, assume_full_pages=assume_full_pages, + assume_full_pages_except_mtp_tail=( + cutile_storage_full_pages + and query_len_per_seq > 1 + and query_len_per_seq <= metadata.page_size + ), assume_valid_pages=assume_valid_pages, prepack_v_for_pv=cutile_prepack_v_for_pv, use_prepacked_v_for_pv=use_cutile_v_packed_cache, @@ -2052,6 +2278,18 @@ def run_fp4_mla_attention_decode( page_max_workspace=page_max, page_sum_workspace=page_sum, ) + if mark_cutile_v_packed_cache_valid: + _set_cutile_v_packed_cache_valid( + metadata, + layer_idx, + kv_cache, + v_head_dim=kv_lora_rank, + page_size=metadata.page_size, + block_v=cutile_block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=src_page_ids, + ) _debug_sync("attention_cutile") return diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py index b821db911635..f3c097ca3e12 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py @@ -106,6 +106,34 @@ def _fp4_mla_swizzled_sf_offset_row_block(row_group, row_offsets, col_idx, SF_PE return col_part + row_part + row_group * (128 * padded_cols) +@triton.jit +def _fp4_e2m1_to_f32(nibble): + magnitude = nibble & 0x7 + value = tl.where( + magnitude == 0, + 0.0, + tl.where( + magnitude == 1, + 0.5, + tl.where( + magnitude == 2, + 1.0, + tl.where( + magnitude == 3, + 1.5, + tl.where( + magnitude == 4, + 2.0, + tl.where(magnitude == 5, 3.0, tl.where(magnitude == 6, 4.0, 6.0)), + ), + ), + ), + ), + ) + sign = (nibble & 0x8) != 0 + return tl.where(sign, -value, value) + + @triton.jit def _fp4_e2m1_quantize(x): abs_x = tl.abs(x) @@ -576,6 +604,7 @@ def _fp4_mla_qk_scores_tile( scale_offsets = tl.arange(0, BLOCK_K // FP4_BLOCK) residual_groups = Q_RESIDUAL_D // FP4_BLOCK non_residual_groups = K_HEAD_D // FP4_BLOCK - residual_groups + token_start = tl.min(token_offsets, axis=0) if USE_TMA_DATA_LOAD and FULL_BLOCK_END > 0: tl.assume(q_fp4_s0 % 8 == 0) tl.assume(q_fp4_s1 == 1) @@ -623,7 +652,7 @@ def _fp4_mla_qk_scores_tile( safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) if USE_TMA_DATA_LOAD and FULL_BLOCK_END > 0 and q_start + BLOCK_K <= non_residual_groups * FP4_BLOCK: q_vals = q_desc.load([(q_row_base + head_start).to(tl.int32), q_start // 2]) - k_vals = k_desc.load([safe_physical_page.to(tl.int32), 0, q_start // 2]) + k_vals = k_desc.load([safe_physical_page.to(tl.int32), token_start.to(tl.int32), q_start // 2]) k_vals = tl.reshape(k_vals, (BLOCK_T, BLOCK_K // 2)) if not ASSUME_VALID_PAGES: k_vals = tl.where(valid_physical_page, k_vals, 0) @@ -677,7 +706,7 @@ def _fp4_mla_qk_scores_tile( k_vals = k_tail_desc.load( [ safe_physical_page.to(tl.int32), - 0, + token_start.to(tl.int32), (non_residual_groups * (FP4_BLOCK // 2)).to(tl.int32), ] ) @@ -839,6 +868,7 @@ def _fp4_mla_attention_stats_kernel( FP4_BLOCK: tl.constexpr, Q_SF_PER_TOKEN: tl.constexpr, K_SF_PER_TOKEN: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, MAX_PAGES: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_T: tl.constexpr, @@ -851,8 +881,10 @@ def _fp4_mla_attention_stats_kernel( ASSUME_VALID_PAGES: tl.constexpr, occupancy: tl.constexpr = 1, ): - gen_idx = tl.program_id(0) + query_idx = tl.program_id(0) head_block = tl.program_id(1) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) if ASSUME_FULL_HEADS: @@ -862,9 +894,10 @@ def _fp4_mla_attention_stats_kernel( mask_h = offs_h < NUM_HEADS safe_offs_h = tl.where(mask_h, offs_h, 0) offs_t = tl.arange(0, BLOCK_T) - q_row_base = gen_idx * NUM_HEADS - kv_len = tl.load(kv_lens_ptr + gen_idx) - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + q_row_base = query_idx * NUM_HEADS + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) denom = tl.zeros((BLOCK_H,), dtype=tl.float32) @@ -922,8 +955,8 @@ def _fp4_mla_attention_stats_kernel( ) max_score = new_max - tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) - tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) + tl.store(max_ptr + query_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) + tl.store(denom_ptr + query_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) @triton.jit @@ -965,6 +998,9 @@ def _fp4_mla_attention_page_stats_kernel( K_SF_PER_TOKEN: tl.constexpr, SF_PER_PAGE: tl.constexpr, P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_BY_QUERY: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_T: tl.constexpr, BLOCK_K: tl.constexpr, @@ -977,9 +1013,11 @@ def _fp4_mla_attention_page_stats_kernel( ASSUME_VALID_PAGES: tl.constexpr, occupancy: tl.constexpr = 1, ): - gen_idx = tl.program_id(0) + query_idx = tl.program_id(0) head_block = tl.program_id(1) page_rel = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) if ASSUME_FULL_HEADS: @@ -989,14 +1027,15 @@ def _fp4_mla_attention_page_stats_kernel( mask_h = offs_h < NUM_HEADS safe_offs_h = tl.where(mask_h, offs_h, 0) offs_t = tl.arange(0, BLOCK_T) - q_row_base = gen_idx * NUM_HEADS + q_row_base = query_idx * NUM_HEADS if ASSUME_FULL_PAGES: kv_len = 0 else: - kv_len = tl.load(kv_lens_ptr + gen_idx) + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) page_start = page_rel * PAGE_SIZE - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) - out_offsets = gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + out_offsets = query_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h page_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) page_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) @@ -1080,12 +1119,16 @@ def _fp4_mla_attention_page_stats_kernel( else: valid_compact_page = (page_table_start + page_rel >= 0) & (page_table_start + page_rel < page_ids_len) safe_compact_page = tl.where(valid_compact_page, page_table_start + page_rel, 0) - p_rows = safe_compact_page * NUM_HEADS + offs_h - safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) + if P_BY_QUERY: + p_page = query_idx * MAX_PAGES + page_rel + else: + p_page = safe_compact_page + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) scale_cols = tl.arange(0, SF_PER_PAGE) if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( - safe_compact_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) else: sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) @@ -1105,7 +1148,7 @@ def _fp4_mla_attention_page_stats_kernel( byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] if USE_TMA_DATA_LOAD and ASSUME_FULL_HEADS and ASSUME_VALID_PAGES: p_desc.store( - [(safe_compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + [(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), ) elif ASSUME_FULL_HEADS: @@ -1174,6 +1217,8 @@ def _fp4_mla_attention_page_stats_grouped_kernel( K_SF_PER_TOKEN: tl.constexpr, SF_PER_PAGE: tl.constexpr, P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + P_BY_QUERY: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_T: tl.constexpr, BLOCK_K: tl.constexpr, @@ -1184,19 +1229,34 @@ def _fp4_mla_attention_page_stats_grouped_kernel( GROUP_REDUCE_STATS: tl.constexpr, ASSUME_FULL_HEADS: tl.constexpr, ASSUME_FULL_PAGES: tl.constexpr, + MASK_MTP_FINAL_PAGE_ONLY: tl.constexpr, ASSUME_VALID_PAGES: tl.constexpr, + MAX_PAGES: tl.constexpr, GROUP_PAGES: tl.constexpr = 2, + ALLOW_PARTIAL_GROUPS: tl.constexpr = False, + PAGE_GROUP_OFFSET: tl.constexpr = 0, + DUPLICATE_TAIL_K: tl.constexpr = False, occupancy: tl.constexpr = 1, ): - gen_idx = tl.program_id(0) + query_idx = tl.program_id(0) head_block = tl.program_id(1) page_group = tl.program_id(2) + logical_page_group = page_group + PAGE_GROUP_OFFSET + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) offs_t = tl.arange(0, BLOCK_T) scale_cols = tl.arange(0, SF_PER_PAGE) - q_row_base = gen_idx * NUM_HEADS - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + q_row_base = query_idx * NUM_HEADS + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + if ASSUME_FULL_PAGES: + kv_len = 0 + elif MASK_MTP_FINAL_PAGE_ONLY: + kv_len = MAX_PAGES * PAGE_SIZE - (QUERY_LEN_PER_SEQ - 1 - query_offset) + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) global_scale = tl.load(global_scale_ptr) qk_scale = sm_scale / (global_scale * global_scale) @@ -1285,9 +1345,18 @@ def _fp4_mla_attention_page_stats_grouped_kernel( group_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) group_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) for page_group_off in tl.range(0, GROUP_PAGES): - page_rel = page_group * GROUP_PAGES + page_group_off + page_rel = logical_page_group * GROUP_PAGES + page_group_off + page_start = page_rel * PAGE_SIZE + valid_group_page = page_rel < MAX_PAGES + if not ASSUME_FULL_PAGES: + valid_group_page = valid_group_page & (page_start < kv_len) compact_page = page_table_start + page_rel - physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) + safe_compact_page = tl.where(valid_group_page, compact_page, page_table_start) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, + mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS), + other=0, + ).to(tl.int64) scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) full_k_vals = k_desc.load([physical_page.to(tl.int32), 0, 0]) @@ -1312,37 +1381,76 @@ def _fp4_mla_attention_page_stats_grouped_kernel( tail_k_scales = tl.ext.load_view_tko(k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0]) tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) - scores = tl.dot_scaled( - q0_vals, - q0_scales, - "e2m1", - tail_k_vals.T, - tail_k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) - scores = tl.dot_scaled( - q1_vals, - q1_scales, - "e2m1", - tail_k_vals.T, - tail_k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) + if DUPLICATE_TAIL_K: + q_tail_vals = tl.join(q0_vals, q1_vals).permute(0, 2, 1) + q_tail_vals = q_tail_vals.reshape([BLOCK_H, 64]) + q_tail_scales = tl.join(q0_scales, q1_scales).permute(0, 2, 1) + q_tail_scales = q_tail_scales.reshape([BLOCK_H, 8]) + tail_k_groups = tail_k_vals.reshape([BLOCK_T, 4, 8]) + tail_k_dup_vals = tl.join(tail_k_groups, tail_k_groups).permute(0, 1, 3, 2) + tail_k_dup_vals = tail_k_dup_vals.reshape([BLOCK_T, 64]) + tail_k_dup_scales = tl.join(tail_k_scales, tail_k_scales).permute(0, 2, 1) + tail_k_dup_scales = tail_k_dup_scales.reshape([BLOCK_T, 8]) + scores = tl.dot_scaled( + q_tail_vals, + q_tail_scales, + "e2m1", + tail_k_dup_vals.T, + tail_k_dup_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + else: + scores = tl.dot_scaled( + q0_vals, + q0_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q1_vals, + q1_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) scores = scores * qk_scale - page_max = tl.max(scores, axis=1) - exp_scores = tl.math.exp2((scores - page_max[:, None]) * 1.4426950408889634) + full_softmax_page = ASSUME_FULL_PAGES or (MASK_MTP_FINAL_PAGE_ONLY and page_rel < MAX_PAGES - 1) + if ASSUME_FULL_HEADS and ( + (ASSUME_FULL_PAGES and not ALLOW_PARTIAL_GROUPS) + or full_softmax_page + ): + page_max = tl.max(scores, axis=1) + exp_scores = tl.math.exp2((scores - page_max[:, None]) * 1.4426950408889634) + else: + if ASSUME_FULL_PAGES: + valid_t = tl.full([BLOCK_T], True, dtype=tl.int1) + else: + valid_t = page_start + offs_t < kv_len + scores = tl.where(valid_group_page & valid_t[None, :], scores, -float("inf")) + page_max = tl.max(scores, axis=1) + safe_page_max = tl.where(valid_group_page, page_max, 0.0) + exp_scores = tl.math.exp2((scores - safe_page_max[:, None]) * 1.4426950408889634) + exp_scores = tl.where(valid_group_page & valid_t[None, :], exp_scores, 0.0) page_sum = tl.sum(exp_scores, axis=1) if GROUP_REDUCE_STATS: next_group_max = tl.maximum(group_max, page_max) - group_sum = group_sum * tl.math.exp2((group_max - next_group_max) * 1.4426950408889634) + page_sum * tl.math.exp2( - (page_max - next_group_max) * 1.4426950408889634 + old_delta = tl.where(group_sum > 0.0, group_max - next_group_max, 0.0) + new_delta = tl.where(page_sum > 0.0, page_max - next_group_max, 0.0) + group_sum = group_sum * tl.math.exp2(old_delta * 1.4426950408889634) + page_sum * tl.math.exp2( + new_delta * 1.4426950408889634 ) group_max = next_group_max @@ -1355,171 +1463,450 @@ def _fp4_mla_attention_page_stats_grouped_kernel( even_probs, odd_probs = tl.split(pairs) packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) - out_offsets = gen_idx * page_stats_s0 + page_rel * page_stats_s1 + offs_h - tl.store(page_max_ptr + out_offsets, page_max) + out_offsets = query_idx * page_stats_s0 + page_rel * page_stats_s1 + offs_h + tl.store(page_max_ptr + out_offsets, page_max, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) if not GROUP_REDUCE_STATS: - tl.store(page_sum_ptr + out_offsets, page_sum) + tl.store(page_sum_ptr + out_offsets, page_sum, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) + if P_BY_QUERY: + p_page = query_idx * MAX_PAGES + page_rel + else: + p_page = safe_compact_page sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( - compact_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE - ) - tl.store(p_sf_ptr + sf_offsets, stored_scale) - p_desc.store( - [(compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], - tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) + if ALLOW_PARTIAL_GROUPS: + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + safe_p_rows = p_page * NUM_HEADS + offs_h + tl.store( + p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=valid_group_page, + ) + else: + p_desc.store( + [(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), + ) if GROUP_REDUCE_STATS: - group_max_offsets = gen_idx * page_stats_s0 + (page_group * 2) * page_stats_s1 + offs_h + group_max_offsets = query_idx * page_stats_s0 + (logical_page_group * 2) * page_stats_s1 + offs_h group_sum_offsets = group_max_offsets + page_stats_s1 tl.store(page_sum_ptr + group_max_offsets, group_max) tl.store(page_sum_ptr + group_sum_offsets, group_sum) @triton.jit -def _fp4_mla_attention_reduce_stats_kernel( - max_ptr, - denom_ptr, +def _fp4_mla_attention_page_stats_grouped_mtp_pair_kernel( page_max_ptr, page_sum_ptr, - stats_s0, - page_stats_s0, - page_stats_s1, - NUM_HEADS: tl.constexpr, - MAX_PAGES: tl.constexpr, - BLOCK_H: tl.constexpr, - GROUP_REDUCE_STATS: tl.constexpr = False, - GROUP_PAGES: tl.constexpr = 1, - occupancy: tl.constexpr = 1, -): - gen_idx = tl.program_id(0) - head_block = tl.program_id(1) - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - - max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) - if GROUP_REDUCE_STATS: - for group_rel in tl.range(0, MAX_PAGES // GROUP_PAGES): - group_max = tl.load( - page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - max_score = tl.maximum(max_score, group_max) - else: - for page_rel in tl.range(0, MAX_PAGES): - page_max = tl.load( - page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - max_score = tl.maximum(max_score, page_max) - - denom = tl.zeros((BLOCK_H,), dtype=tl.float32) - if GROUP_REDUCE_STATS: - for group_rel in tl.range(0, MAX_PAGES // GROUP_PAGES): - group_max = tl.load( - page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - group_sum = tl.load( - page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2 + 1) * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=0.0, - ) - denom += tl.where(group_sum > 0.0, group_sum * tl.math.exp2((group_max - max_score) * 1.4426950408889634), 0.0) - else: - for page_rel in tl.range(0, MAX_PAGES): - page_max = tl.load( - page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - page_sum = tl.load( - page_sum_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=0.0, - ) - denom += tl.where(page_sum > 0.0, page_sum * tl.math.exp2((page_max - max_score) * 1.4426950408889634), 0.0) - - tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) - tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) - - -@triton.jit -def _fp4_mla_attention_prob_scale_kernel( + p_fp4_ptr, p_sf_ptr, - max_ptr, - denom_ptr, - page_max_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, paged_kv_indptr_decode_ptr, kv_lens_ptr, - page_ids_len, - stats_s0, - page_stats_s0, - page_stats_s1, + page_ids_len: tl.constexpr, + num_pages: tl.constexpr, + q_fp4_s0: tl.constexpr, + q_fp4_s1: tl.constexpr, + kv_s0: tl.constexpr, + kv_s2: tl.constexpr, + kv_s4: tl.constexpr, + sf_s0: tl.constexpr, + page_stats_s0: tl.constexpr, + page_stats_s1: tl.constexpr, + p_s0: tl.constexpr, + p_s1: tl.constexpr, + p_num_rows: tl.constexpr, + q_num_rows: tl.constexpr, + sm_scale: tl.constexpr, NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + PACK_PROBS: tl.constexpr, + GROUP_REDUCE_STATS: tl.constexpr, ASSUME_FULL_HEADS: tl.constexpr, ASSUME_FULL_PAGES: tl.constexpr, ASSUME_VALID_PAGES: tl.constexpr, + MAX_PAGES: tl.constexpr, + GROUP_PAGES: tl.constexpr = 2, + ALLOW_PARTIAL_GROUPS: tl.constexpr = False, + Q_PER_GROUP: tl.constexpr = 2, + DUPLICATE_TAIL_K: tl.constexpr = False, occupancy: tl.constexpr = 1, ): - gen_idx = tl.program_id(0) + seq_idx = tl.program_id(0) head_block = tl.program_id(1) - page_rel = tl.program_id(2) - - page_start = page_rel * PAGE_SIZE - if ASSUME_FULL_PAGES: - kv_len = 0 - else: - kv_len = tl.load(kv_lens_ptr + gen_idx) - if page_start >= kv_len: - return - - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) - compact_page = page_table_start + page_rel - if not ASSUME_VALID_PAGES: - if (compact_page < 0) | (compact_page >= page_ids_len): - return + combo = tl.program_id(2) + num_query_groups = QUERY_LEN_PER_SEQ // Q_PER_GROUP + page_group = combo // num_query_groups + query_group = combo - page_group * num_query_groups + query_offset0 = query_group * Q_PER_GROUP + query_offset1 = query_offset0 + 1 + query_idx0 = seq_idx * QUERY_LEN_PER_SEQ + query_offset0 + query_idx1 = query_idx0 + 1 offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_offs_h = offs_h - else: - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - page_max = tl.load( - page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - max_score = tl.load(max_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - denom = tl.load(denom_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - factor = tl.where(denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0) - - p_rows = compact_page * NUM_HEADS + offs_h - safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, compact_page * NUM_HEADS) + offs_t = tl.arange(0, BLOCK_T) scale_cols = tl.arange(0, SF_PER_PAGE) - if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: - sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( - compact_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE - ) + q_row_base0 = query_idx0 * NUM_HEADS + q_row_base1 = query_idx1 * NUM_HEADS + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + if ASSUME_FULL_PAGES: + kv_len0 = 0 + kv_len1 = 0 else: - sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) - scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) - tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) + kv_len_base = tl.load(kv_lens_ptr + seq_idx) + kv_len0 = tl.maximum(kv_len_base - (QUERY_LEN_PER_SEQ - 1 - query_offset0), 0) + kv_len1 = tl.maximum(kv_len_base - (QUERY_LEN_PER_SEQ - 1 - query_offset1), 0) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + tl.assume(q_fp4_s0 % 8 == 0) + tl.assume(q_fp4_s1 == 1) + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) -@triton.jit -def _fp4_mla_attention_prob_store_page_kernel( - probs_ptr, - max_ptr, - denom_ptr, + q_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 256], + ) + k_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 256], + ) + q_tail_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 32], + ) + k_tail_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 32], + ) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + q_sf_full_view = tl.ext.make_view( + base=q_sf_ptr, + shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + q_sf_tail_view = tl.ext.make_view( + base=q_sf_ptr, + shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + k_sf_full_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + k_sf_tail_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 1, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + + q_row_start0 = (q_row_base0 + head_block * BLOCK_H).to(tl.int32) + q_row_group0 = q_row_base0 // 128 + full_q_vals0 = q_desc.load([q_row_start0, 0]) + full_q_scales0 = tl.ext.load_view_tko(q_sf_full_view, [q_row_group0.to(tl.int32), 0, 0, 0]) + full_q_scales0 = full_q_scales0.reshape([1, 8, 32, 4, 4]).trans(0, 3, 2, 1, 4) + full_q_scales0 = full_q_scales0.reshape([BLOCK_H, 32]) + q0_tail0_vals = q_tail_desc.load([q_row_start0, 256]) + q0_tail1_vals = q_tail_desc.load([q_row_start0, 288]) + q0_tail0_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group0.to(tl.int32), 8, 0, 0]) + q0_tail0_scales = q0_tail0_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) + q0_tail0_scales = q0_tail0_scales.reshape([BLOCK_H, 4]) + q0_tail1_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group0.to(tl.int32), 9, 0, 0]) + q0_tail1_scales = q0_tail1_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) + q0_tail1_scales = q0_tail1_scales.reshape([BLOCK_H, 4]) + + q_row_start1 = (q_row_base1 + head_block * BLOCK_H).to(tl.int32) + q_row_group1 = q_row_base1 // 128 + full_q_vals1 = q_desc.load([q_row_start1, 0]) + full_q_scales1 = tl.ext.load_view_tko(q_sf_full_view, [q_row_group1.to(tl.int32), 0, 0, 0]) + full_q_scales1 = full_q_scales1.reshape([1, 8, 32, 4, 4]).trans(0, 3, 2, 1, 4) + full_q_scales1 = full_q_scales1.reshape([BLOCK_H, 32]) + q1_tail0_vals = q_tail_desc.load([q_row_start1, 256]) + q1_tail1_vals = q_tail_desc.load([q_row_start1, 288]) + q1_tail0_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group1.to(tl.int32), 8, 0, 0]) + q1_tail0_scales = q1_tail0_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) + q1_tail0_scales = q1_tail0_scales.reshape([BLOCK_H, 4]) + q1_tail1_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group1.to(tl.int32), 9, 0, 0]) + q1_tail1_scales = q1_tail1_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) + q1_tail1_scales = q1_tail1_scales.reshape([BLOCK_H, 4]) + + group_max0 = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + group_sum0 = tl.zeros((BLOCK_H,), dtype=tl.float32) + group_max1 = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + group_sum1 = tl.zeros((BLOCK_H,), dtype=tl.float32) + for page_group_off in tl.range(0, GROUP_PAGES): + page_rel = page_group * GROUP_PAGES + page_group_off + page_start = page_rel * PAGE_SIZE + valid_group_page_base = page_rel < MAX_PAGES + compact_page = page_table_start + page_rel + safe_compact_page = tl.where(valid_group_page_base, compact_page, page_table_start) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, + mask=valid_group_page_base | (not ALLOW_PARTIAL_GROUPS), + other=0, + ).to(tl.int64) + + full_k_vals = k_desc.load([physical_page.to(tl.int32), 0, 0]) + full_k_vals = tl.reshape(full_k_vals, (BLOCK_T, 256)) + full_k_scales = tl.ext.load_view_tko(k_sf_full_view, [physical_page.to(tl.int32), 0, 0, 0, 0]) + full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) + tail_k_vals = k_tail_desc.load([physical_page.to(tl.int32), 0, 256]) + tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) + tail_k_scales = tl.ext.load_view_tko(k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0]) + tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) + + scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) + scores = tl.dot_scaled( + full_q_vals0, + full_q_scales0, + "e2m1", + full_k_vals.T, + full_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q0_tail0_vals, + q0_tail0_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q0_tail1_vals, + q0_tail1_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = scores * qk_scale + valid_group_page0 = valid_group_page_base + if not ASSUME_FULL_PAGES: + valid_group_page0 = valid_group_page0 & (page_start < kv_len0) + if ASSUME_FULL_PAGES: + valid_t0 = tl.full([BLOCK_T], True, dtype=tl.int1) + else: + valid_t0 = page_start + offs_t < kv_len0 + if ASSUME_FULL_HEADS and ASSUME_FULL_PAGES and not ALLOW_PARTIAL_GROUPS: + page_max0 = tl.max(scores, axis=1) + exp_scores = tl.math.exp2((scores - page_max0[:, None]) * 1.4426950408889634) + else: + scores = tl.where(valid_group_page0 & valid_t0[None, :], scores, -float("inf")) + page_max0 = tl.max(scores, axis=1) + safe_page_max0 = tl.where(valid_group_page0, page_max0, 0.0) + exp_scores = tl.math.exp2((scores - safe_page_max0[:, None]) * 1.4426950408889634) + exp_scores = tl.where(valid_group_page0 & valid_t0[None, :], exp_scores, 0.0) + page_sum0 = tl.sum(exp_scores, axis=1) + if GROUP_REDUCE_STATS: + next_group_max0 = tl.maximum(group_max0, page_max0) + old_delta0 = tl.where(group_sum0 > 0.0, group_max0 - next_group_max0, 0.0) + new_delta0 = tl.where(page_sum0 > 0.0, page_max0 - next_group_max0, 0.0) + group_sum0 = group_sum0 * tl.math.exp2(old_delta0 * 1.4426950408889634) + page_sum0 * tl.math.exp2( + new_delta0 * 1.4426950408889634 + ) + group_max0 = next_group_max0 + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + out_offsets0 = query_idx0 * page_stats_s0 + page_rel * page_stats_s1 + offs_h + tl.store(page_max_ptr + out_offsets0, page_max0, mask=valid_group_page0 | (not ALLOW_PARTIAL_GROUPS)) + if not GROUP_REDUCE_STATS: + tl.store(page_sum_ptr + out_offsets0, page_sum0, mask=valid_group_page0 | (not ALLOW_PARTIAL_GROUPS)) + p_page0 = query_idx0 * MAX_PAGES + page_rel + sf_offsets0 = _fp4_mla_swizzled_sf_offset_row_block( + p_page0, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + tl.store(p_sf_ptr + sf_offsets0, stored_scale, mask=valid_group_page0 | (not ALLOW_PARTIAL_GROUPS)) + if ALLOW_PARTIAL_GROUPS: + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + safe_p_rows0 = p_page0 * NUM_HEADS + offs_h + tl.store( + p_fp4_ptr + safe_p_rows0[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=valid_group_page0, + ) + else: + p_desc.store( + [(p_page0 * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), + ) + + scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) + scores = tl.dot_scaled( + full_q_vals1, + full_q_scales1, + "e2m1", + full_k_vals.T, + full_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q1_tail0_vals, + q1_tail0_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q1_tail1_vals, + q1_tail1_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = scores * qk_scale + valid_group_page1 = valid_group_page_base + if not ASSUME_FULL_PAGES: + valid_group_page1 = valid_group_page1 & (page_start < kv_len1) + if ASSUME_FULL_PAGES: + valid_t1 = tl.full([BLOCK_T], True, dtype=tl.int1) + else: + valid_t1 = page_start + offs_t < kv_len1 + if ASSUME_FULL_HEADS and ASSUME_FULL_PAGES and not ALLOW_PARTIAL_GROUPS: + page_max1 = tl.max(scores, axis=1) + exp_scores = tl.math.exp2((scores - page_max1[:, None]) * 1.4426950408889634) + else: + scores = tl.where(valid_group_page1 & valid_t1[None, :], scores, -float("inf")) + page_max1 = tl.max(scores, axis=1) + safe_page_max1 = tl.where(valid_group_page1, page_max1, 0.0) + exp_scores = tl.math.exp2((scores - safe_page_max1[:, None]) * 1.4426950408889634) + exp_scores = tl.where(valid_group_page1 & valid_t1[None, :], exp_scores, 0.0) + page_sum1 = tl.sum(exp_scores, axis=1) + if GROUP_REDUCE_STATS: + next_group_max1 = tl.maximum(group_max1, page_max1) + old_delta1 = tl.where(group_sum1 > 0.0, group_max1 - next_group_max1, 0.0) + new_delta1 = tl.where(page_sum1 > 0.0, page_max1 - next_group_max1, 0.0) + group_sum1 = group_sum1 * tl.math.exp2(old_delta1 * 1.4426950408889634) + page_sum1 * tl.math.exp2( + new_delta1 * 1.4426950408889634 + ) + group_max1 = next_group_max1 + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + out_offsets1 = query_idx1 * page_stats_s0 + page_rel * page_stats_s1 + offs_h + tl.store(page_max_ptr + out_offsets1, page_max1, mask=valid_group_page1 | (not ALLOW_PARTIAL_GROUPS)) + if not GROUP_REDUCE_STATS: + tl.store(page_sum_ptr + out_offsets1, page_sum1, mask=valid_group_page1 | (not ALLOW_PARTIAL_GROUPS)) + p_page1 = query_idx1 * MAX_PAGES + page_rel + sf_offsets1 = _fp4_mla_swizzled_sf_offset_row_block( + p_page1, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + tl.store(p_sf_ptr + sf_offsets1, stored_scale, mask=valid_group_page1 | (not ALLOW_PARTIAL_GROUPS)) + if ALLOW_PARTIAL_GROUPS: + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + safe_p_rows1 = p_page1 * NUM_HEADS + offs_h + tl.store( + p_fp4_ptr + safe_p_rows1[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=valid_group_page1, + ) + else: + p_desc.store( + [(p_page1 * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), + ) + + if GROUP_REDUCE_STATS: + group_max_offsets0 = query_idx0 * page_stats_s0 + (page_group * 2) * page_stats_s1 + offs_h + group_sum_offsets0 = group_max_offsets0 + page_stats_s1 + tl.store(page_sum_ptr + group_max_offsets0, group_max0) + tl.store(page_sum_ptr + group_sum_offsets0, group_sum0) + group_max_offsets1 = query_idx1 * page_stats_s0 + (page_group * 2) * page_stats_s1 + offs_h + group_sum_offsets1 = group_max_offsets1 + page_stats_s1 + tl.store(page_sum_ptr + group_max_offsets1, group_max1) + tl.store(page_sum_ptr + group_sum_offsets1, group_sum1) + + +@triton.jit +def _fp4_mla_attention_page_stats_grouped_generic_kernel( + page_max_ptr, + page_sum_ptr, + p_fp4_ptr, + p_sf_ptr, q_fp4_ptr, q_sf_ptr, kv_cache_ptr, @@ -1528,20 +1915,21 @@ def _fp4_mla_attention_prob_store_page_kernel( src_page_ids_ptr, paged_kv_indptr_decode_ptr, kv_lens_ptr, - page_rel, page_ids_len, num_pages, - probs_s0, - probs_s1, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - stats_s0, - q_num_rows, - sm_scale, + q_fp4_s0: tl.constexpr, + q_fp4_s1: tl.constexpr, + kv_s0: tl.constexpr, + kv_s2: tl.constexpr, + kv_s4: tl.constexpr, + sf_s0: tl.constexpr, + page_stats_s0: tl.constexpr, + page_stats_s1: tl.constexpr, + p_s0: tl.constexpr, + p_s1: tl.constexpr, + p_num_rows: tl.constexpr, + q_num_rows: tl.constexpr, + sm_scale: tl.constexpr, NUM_HEADS: tl.constexpr, Q_HEAD_D: tl.constexpr, K_HEAD_D: tl.constexpr, @@ -1550,185 +1938,229 @@ def _fp4_mla_attention_prob_store_page_kernel( FP4_BLOCK: tl.constexpr, Q_SF_PER_TOKEN: tl.constexpr, K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, BLOCK_K: tl.constexpr, FULL_BLOCK_END: tl.constexpr, TAIL_BLOCK_K: tl.constexpr, USE_TMA_DATA_LOAD: tl.constexpr, + PACK_PROBS: tl.constexpr, + GROUP_REDUCE_STATS: tl.constexpr, ASSUME_FULL_HEADS: tl.constexpr, ASSUME_FULL_PAGES: tl.constexpr, + MASK_MTP_FINAL_PAGE_ONLY: tl.constexpr, ASSUME_VALID_PAGES: tl.constexpr, + MAX_PAGES: tl.constexpr, + GROUP_PAGES: tl.constexpr = 2, + ALLOW_PARTIAL_GROUPS: tl.constexpr = False, + PAGE_GROUP_OFFSET: tl.constexpr = 0, + DUPLICATE_TAIL_K: tl.constexpr = False, occupancy: tl.constexpr = 1, ): gen_idx = tl.program_id(0) head_block = tl.program_id(1) - - kv_len = tl.load(kv_lens_ptr + gen_idx) - page_start = page_rel * PAGE_SIZE - if (not ASSUME_FULL_PAGES) and page_start >= kv_len: - return + page_group = tl.program_id(2) + logical_page_group = page_group + PAGE_GROUP_OFFSET offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_offs_h = offs_h - else: - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - offs_t = tl.arange(0, PAGE_SIZE) - if ASSUME_FULL_PAGES: - valid_t = tl.full([PAGE_SIZE], True, dtype=tl.int1) - else: - valid_t = page_start + offs_t < kv_len - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) - compact_page = page_table_start + page_rel - if (compact_page < 0) | (compact_page >= page_ids_len): - return + offs_t = tl.arange(0, BLOCK_T) + scale_cols = tl.arange(0, SF_PER_PAGE) q_row_base = gen_idx * NUM_HEADS + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) - scores = _fp4_mla_qk_scores_tile( - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - src_page_ids_ptr, - compact_page, - q_row_base, - head_block * BLOCK_H, - offs_h, - offs_t, - q_num_rows, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - page_ids_len, - num_pages, - Q_HEAD_D, - K_HEAD_D, - Q_RESIDUAL_D, - FP4_BLOCK, - Q_SF_PER_TOKEN, - K_SF_PER_TOKEN, - BLOCK_H, - PAGE_SIZE, - BLOCK_K, - FULL_BLOCK_END, - TAIL_BLOCK_K, - NUM_HEADS, - USE_TMA_DATA_LOAD, - ASSUME_FULL_HEADS, - ASSUME_VALID_PAGES, + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], ) - max_score = tl.load(max_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - denom = tl.load(denom_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) - global_scale = tl.load(global_scale_ptr) - qk_scale = sm_scale / (global_scale * global_scale) - denom_valid = denom > 0.0 - safe_denom = tl.where(denom_valid, denom, 1.0) - safe_max = tl.where(denom_valid, max_score, 0.0) - scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) - probs = tl.math.exp2((scores - safe_max[:, None]) * 1.4426950408889634) / safe_denom[:, None] - probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) - prob_rows = gen_idx * NUM_HEADS + offs_h - safe_prob_rows = tl.where(mask_h, prob_rows, gen_idx * NUM_HEADS) - tl.store( - probs_ptr + safe_prob_rows[:, None] * probs_s0 + offs_t[None, :] * probs_s1, - probs, - mask=mask_h[:, None], - ) + group_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + group_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) + for page_group_off in tl.range(0, GROUP_PAGES): + page_rel = logical_page_group * GROUP_PAGES + page_group_off + valid_group_page = page_rel < MAX_PAGES + compact_page = page_table_start + page_rel + safe_compact_page = tl.where(valid_group_page, compact_page, page_table_start) + + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + safe_compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + BLOCK_T, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + scores = scores * qk_scale + page_max = tl.max(scores, axis=1) + if ALLOW_PARTIAL_GROUPS: + page_max = tl.where(valid_group_page, page_max, -float("inf")) + safe_page_max = tl.where(valid_group_page, page_max, 0.0) + exp_scores = tl.math.exp2((scores - safe_page_max[:, None]) * 1.4426950408889634) + exp_scores = tl.where(valid_group_page, exp_scores, 0.0) + else: + exp_scores = tl.math.exp2((scores - page_max[:, None]) * 1.4426950408889634) + page_sum = tl.sum(exp_scores, axis=1) + if GROUP_REDUCE_STATS: + next_group_max = tl.maximum(group_max, page_max) + old_delta = tl.where(group_sum > 0.0, group_max - next_group_max, 0.0) + new_delta = tl.where(page_sum > 0.0, page_max - next_group_max, 0.0) + group_sum = group_sum * tl.math.exp2(old_delta * 1.4426950408889634) + page_sum * tl.math.exp2( + new_delta * 1.4426950408889634 + ) + group_max = next_group_max + + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + out_offsets = gen_idx * page_stats_s0 + page_rel * page_stats_s1 + offs_h + tl.store(page_max_ptr + out_offsets, page_max, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) + if not GROUP_REDUCE_STATS: + tl.store(page_sum_ptr + out_offsets, page_sum, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) + + safe_p_rows = safe_compact_page * NUM_HEADS + offs_h + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) + if ALLOW_PARTIAL_GROUPS: + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + tl.store( + p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=valid_group_page, + ) + else: + p_desc.store( + [(compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), + ) + + if GROUP_REDUCE_STATS: + group_max_offsets = gen_idx * page_stats_s0 + (logical_page_group * 2) * page_stats_s1 + offs_h + group_sum_offsets = group_max_offsets + page_stats_s1 + tl.store(page_sum_ptr + group_max_offsets, group_max) + tl.store(page_sum_ptr + group_sum_offsets, group_sum) @triton.jit -def _fp4_mla_attention_prob_pack_page_kernel( - p_fp4_ptr, - p_sf_ptr, - probs_ptr, - paged_kv_indptr_decode_ptr, - kv_lens_ptr, - page_rel, - page_ids_len, - p_s0, - p_s1, - probs_s0, - probs_s1, +def _fp4_mla_attention_reduce_stats_kernel( + max_ptr, + denom_ptr, + page_max_ptr, + page_sum_ptr, + stats_s0, + page_stats_s0, + page_stats_s1, NUM_HEADS: tl.constexpr, - PAGE_SIZE: tl.constexpr, - FP4_BLOCK: tl.constexpr, - SF_PER_PAGE: tl.constexpr, - P_GLOBAL_SCALE: tl.constexpr, + MAX_PAGES: tl.constexpr, BLOCK_H: tl.constexpr, - ASSUME_FULL_HEADS: tl.constexpr, - ASSUME_FULL_PAGES: tl.constexpr, + GROUP_REDUCE_STATS: tl.constexpr, + GROUP_PAGES: tl.constexpr, + NUM_PAGE_GROUPS: tl.constexpr, occupancy: tl.constexpr = 1, ): gen_idx = tl.program_id(0) - token_group = tl.program_id(1) - head_block = tl.program_id(2) - - kv_len = tl.load(kv_lens_ptr + gen_idx) - page_start = page_rel * PAGE_SIZE - if (not ASSUME_FULL_PAGES) and page_start >= kv_len: - return - - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) - compact_page = page_table_start + page_rel - if (compact_page < 0) | (compact_page >= page_ids_len): - return - + head_block = tl.program_id(1) offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - else: - mask_h = offs_h < NUM_HEADS - byte_offsets = tl.arange(0, FP4_BLOCK // 2) - token_base = token_group * FP4_BLOCK - even_t = token_base + byte_offsets * 2 - odd_t = even_t + 1 - valid_even = page_start + even_t < kv_len - valid_odd = page_start + odd_t < kv_len + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) - prob_rows = gen_idx * NUM_HEADS + offs_h - safe_prob_rows = tl.where(mask_h, prob_rows, gen_idx * NUM_HEADS) - even_probs = tl.load( - probs_ptr + safe_prob_rows[:, None] * probs_s0 + even_t[None, :] * probs_s1, - mask=mask_h[:, None] & valid_even[None, :], - other=0.0, - ) - odd_probs = tl.load( - probs_ptr + safe_prob_rows[:, None] * probs_s0 + odd_t[None, :] * probs_s1, - mask=mask_h[:, None] & valid_odd[None, :], - other=0.0, - ) - amax = tl.maximum(tl.max(tl.abs(even_probs), axis=1), tl.max(tl.abs(odd_probs), axis=1)) - local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) - stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) + max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + if GROUP_REDUCE_STATS: + for group_rel in tl.range(0, NUM_PAGE_GROUPS): + group_max = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.maximum(max_score, group_max) + else: + for page_rel in tl.range(0, MAX_PAGES): + page_max = tl.load( + page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.maximum(max_score, page_max) - p_rows = compact_page * NUM_HEADS + offs_h - safe_p_rows = tl.where(mask_h, p_rows, compact_page * NUM_HEADS) - sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows, token_group, SF_PER_PAGE) - tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h) + denom = tl.zeros((BLOCK_H,), dtype=tl.float32) + if GROUP_REDUCE_STATS: + for group_rel in tl.range(0, NUM_PAGE_GROUPS): + group_max = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + group_sum = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2 + 1) * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=0.0, + ) + denom += tl.where(group_sum > 0.0, group_sum * tl.math.exp2((group_max - max_score) * 1.4426950408889634), 0.0) + else: + for page_rel in tl.range(0, MAX_PAGES): + page_max = tl.load( + page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + page_sum = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=0.0, + ) + denom += tl.where(page_sum > 0.0, page_sum * tl.math.exp2((page_max - max_score) * 1.4426950408889634), 0.0) - even_quant = _fp4_e2m1_quantize(even_probs / local_scale[:, None]) - odd_quant = _fp4_e2m1_quantize(odd_probs / local_scale[:, None]) - packed = even_quant | (odd_quant << 4) - byte_cols = token_group * (FP4_BLOCK // 2) + byte_offsets - tl.store( - p_fp4_ptr + safe_p_rows[:, None] * p_s0 + byte_cols[None, :] * p_s1, - packed, - mask=mask_h[:, None], - ) + tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) + tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) @triton.jit -def _fp4_mla_attention_prob_pack_page_fused_kernel( +def _fp4_mla_attention_page_stats_half_grouped_kernel( + page_max_ptr, + page_sum_ptr, p_fp4_ptr, p_sf_ptr, - max_ptr, - denom_ptr, q_fp4_ptr, q_sf_ptr, kv_cache_ptr, @@ -1737,20 +2169,21 @@ def _fp4_mla_attention_prob_pack_page_fused_kernel( src_page_ids_ptr, paged_kv_indptr_decode_ptr, kv_lens_ptr, - page_rel, page_ids_len, num_pages, - p_s0, - p_s1, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - stats_s0, - q_num_rows, - sm_scale, + q_fp4_s0: tl.constexpr, + q_fp4_s1: tl.constexpr, + kv_s0: tl.constexpr, + kv_s2: tl.constexpr, + kv_s4: tl.constexpr, + sf_s0: tl.constexpr, + page_stats_s0: tl.constexpr, + page_stats_s1: tl.constexpr, + p_s0: tl.constexpr, + p_s1: tl.constexpr, + p_num_rows: tl.constexpr, + q_num_rows: tl.constexpr, + sm_scale: tl.constexpr, NUM_HEADS: tl.constexpr, Q_HEAD_D: tl.constexpr, K_HEAD_D: tl.constexpr, @@ -1762,826 +2195,3040 @@ def _fp4_mla_attention_prob_pack_page_fused_kernel( SF_PER_PAGE: tl.constexpr, P_GLOBAL_SCALE: tl.constexpr, BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, BLOCK_K: tl.constexpr, FULL_BLOCK_END: tl.constexpr, TAIL_BLOCK_K: tl.constexpr, USE_TMA_DATA_LOAD: tl.constexpr, - PAGE_REL_FROM_GRID: tl.constexpr = False, - ASSUME_FULL_HEADS: tl.constexpr = False, - ASSUME_FULL_PAGES: tl.constexpr = False, - ASSUME_VALID_PAGES: tl.constexpr = False, + GROUP_REDUCE_STATS: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + MAX_TILES: tl.constexpr, + GROUP_TILES: tl.constexpr, + TILES_PER_PAGE: tl.constexpr, + ALLOW_PARTIAL_GROUPS: tl.constexpr = False, occupancy: tl.constexpr = 1, ): gen_idx = tl.program_id(0) head_block = tl.program_id(1) - if PAGE_REL_FROM_GRID: - page_rel = tl.program_id(2) - - kv_len = tl.load(kv_lens_ptr + gen_idx) - page_start = page_rel * PAGE_SIZE - if (not ASSUME_FULL_PAGES) and page_start >= kv_len: - return + tile_group = tl.program_id(2) offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_offs_h = offs_h - else: - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - offs_t = tl.arange(0, PAGE_SIZE) - if ASSUME_FULL_PAGES: - valid_t = tl.full([PAGE_SIZE], True, dtype=tl.int1) - else: - valid_t = page_start + offs_t < kv_len - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) - compact_page = page_table_start + page_rel - if (compact_page < 0) | (compact_page >= page_ids_len): - return + local_t = tl.arange(0, BLOCK_T) + local_scale_cols = tl.arange(0, BLOCK_T // FP4_BLOCK) q_row_base = gen_idx * NUM_HEADS - - scores = _fp4_mla_qk_scores_tile( - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - src_page_ids_ptr, - compact_page, - q_row_base, - head_block * BLOCK_H, - offs_h, - offs_t, - q_num_rows, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - page_ids_len, - num_pages, - Q_HEAD_D, - K_HEAD_D, - Q_RESIDUAL_D, - FP4_BLOCK, - Q_SF_PER_TOKEN, - K_SF_PER_TOKEN, - BLOCK_H, - PAGE_SIZE, - BLOCK_K, - FULL_BLOCK_END, - TAIL_BLOCK_K, - NUM_HEADS, - USE_TMA_DATA_LOAD, - ASSUME_FULL_HEADS, - ASSUME_VALID_PAGES, - ) - max_score = tl.load(max_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - denom = tl.load(denom_ptr + gen_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) global_scale = tl.load(global_scale_ptr) qk_scale = sm_scale / (global_scale * global_scale) - denom_valid = denom > 0.0 - safe_denom = tl.where(denom_valid, denom, 1.0) - safe_max = tl.where(denom_valid, max_score, 0.0) - scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) - probs = tl.math.exp2((scores - safe_max[:, None]) * 1.4426950408889634) / safe_denom[:, None] - probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) - grouped_probs = tl.reshape(probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) - amax = tl.max(tl.abs(grouped_probs), axis=2) - local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) - stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) - scaled_probs = grouped_probs / tl.reshape(local_scale, (BLOCK_H, SF_PER_PAGE, 1)) - pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) - even_probs, odd_probs = tl.split(pairs) - packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) - p_rows = compact_page * NUM_HEADS + offs_h - safe_p_rows = tl.where(mask_h, p_rows, compact_page * NUM_HEADS) - scale_cols = tl.arange(0, SF_PER_PAGE) - sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) - tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h[:, None]) + group_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + group_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) + for group_off in tl.range(0, GROUP_TILES): + tile_rel = tile_group * GROUP_TILES + group_off + valid_tile = tile_rel < MAX_TILES + page_rel = tile_rel // TILES_PER_PAGE + tile_in_page = tile_rel - page_rel * TILES_PER_PAGE + token_base = tile_in_page * BLOCK_T + compact_page = page_table_start + page_rel + safe_compact_page = tl.where(valid_tile, compact_page, page_table_start) - byte_offsets = tl.arange(0, FP4_BLOCK // 2) - byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] - tl.store( - p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, - packed, - mask=mask_h[:, None, None], - ) + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + safe_compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + token_base + local_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + BLOCK_T, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + scores = scores * qk_scale + tile_max = tl.max(scores, axis=1) + tile_max = tl.where(valid_tile, tile_max, -float("inf")) + safe_tile_max = tl.where(valid_tile, tile_max, 0.0) + exp_scores = tl.math.exp2((scores - safe_tile_max[:, None]) * 1.4426950408889634) + exp_scores = tl.where(valid_tile, exp_scores, 0.0) + tile_sum = tl.sum(exp_scores, axis=1) + if GROUP_REDUCE_STATS: + next_group_max = tl.maximum(group_max, tile_max) + old_delta = tl.where(group_sum > 0.0, group_max - next_group_max, 0.0) + new_delta = tl.where(tile_sum > 0.0, tile_max - next_group_max, 0.0) + group_sum = group_sum * tl.math.exp2(old_delta * 1.4426950408889634) + tile_sum * tl.math.exp2( + new_delta * 1.4426950408889634 + ) + group_max = next_group_max + + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, BLOCK_T // FP4_BLOCK, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, BLOCK_T // FP4_BLOCK, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, BLOCK_T // FP4_BLOCK, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + out_offsets = gen_idx * page_stats_s0 + tile_rel * page_stats_s1 + offs_h + tl.store(page_max_ptr + out_offsets, tile_max, mask=valid_tile) + if not GROUP_REDUCE_STATS: + tl.store(page_sum_ptr + out_offsets, tile_sum, mask=valid_tile) + + sf_cols = tile_in_page * (BLOCK_T // FP4_BLOCK) + local_scale_cols + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + safe_compact_page, offs_h[:, None], sf_cols[None, :], SF_PER_PAGE + ) + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_tile) + + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = token_base // 2 + local_scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + p_rows = safe_compact_page * NUM_HEADS + offs_h + tl.store( + p_fp4_ptr + p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=valid_tile, + ) + + if GROUP_REDUCE_STATS: + group_max_offsets = gen_idx * page_stats_s0 + (tile_group * 2) * page_stats_s1 + offs_h + group_sum_offsets = group_max_offsets + page_stats_s1 + tl.store(page_sum_ptr + group_max_offsets, group_max) + tl.store(page_sum_ptr + group_sum_offsets, group_sum) @triton.jit -def _fp4_mla_attention_pv_kernel( - out_ptr, - p_fp4_ptr, +def _fp4_mla_attention_prob_scale_kernel( p_sf_ptr, - kv_cache_ptr, - v_sf_ptr, - global_scale_ptr, - src_page_ids_ptr, + max_ptr, + denom_ptr, + page_max_ptr, paged_kv_indptr_decode_ptr, kv_lens_ptr, page_ids_len, - num_pages, - out_s0, - out_s1, - out_s2, - out_num_rows, - p_s0, - p_s1, - p_num_rows, - kv_s0, - kv_s2, - kv_s4, - vsf_s0, + stats_s0, + page_stats_s0, + page_stats_s1, NUM_HEADS: tl.constexpr, - V_HEAD_D: tl.constexpr, PAGE_SIZE: tl.constexpr, - FP4_BLOCK: tl.constexpr, SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, MAX_PAGES: tl.constexpr, - P_GLOBAL_SCALE: tl.constexpr, + P_BY_QUERY: tl.constexpr, BLOCK_H: tl.constexpr, - BLOCK_V: tl.constexpr, - USE_TMA_P_LOAD: tl.constexpr, - USE_TMA_V_LOAD: tl.constexpr, - PV_LOOP_STAGES: tl.constexpr, ASSUME_FULL_HEADS: tl.constexpr, ASSUME_FULL_PAGES: tl.constexpr, - ASSUME_FULL_V: tl.constexpr, ASSUME_VALID_PAGES: tl.constexpr, occupancy: tl.constexpr = 1, ): - gen_idx = tl.program_id(0) + query_idx = tl.program_id(0) head_block = tl.program_id(1) - dim_block = tl.program_id(2) + page_rel = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + if page_start >= kv_len: + return + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if not ASSUME_VALID_PAGES: + if (compact_page < 0) | (compact_page >= page_ids_len): + return offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) if ASSUME_FULL_HEADS: mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) safe_offs_h = offs_h else: mask_h = offs_h < NUM_HEADS safe_offs_h = tl.where(mask_h, offs_h, 0) - if ASSUME_FULL_V: - mask_v = tl.full([BLOCK_V], True, dtype=tl.int1) - safe_offs_v = offs_v + page_max = tl.load( + page_max_ptr + query_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + factor = tl.where(denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0) + + if P_BY_QUERY: + p_page = query_idx * MAX_PAGES + page_rel else: - mask_v = offs_v < V_HEAD_D - safe_offs_v = tl.where(mask_v, offs_v, 0) - packed_t = tl.arange(0, PAGE_SIZE // 2) - scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) - even_t = packed_t * 2 - odd_t = even_t + 1 - v_packed_offsets = safe_offs_v // 2 - v_use_high_nibble = (safe_offs_v & 1) != 0 - if ASSUME_FULL_V and BLOCK_V == 128: - v_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( - dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + p_page = compact_page + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + scale_cols = tl.arange(0, SF_PER_PAGE) + if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) else: - v_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE) - if USE_TMA_P_LOAD: - tl.assume(p_s0 % 8 == 0) - tl.assume(p_s1 == 1) - p_desc = tl.make_tensor_descriptor( - p_fp4_ptr, - shape=[p_num_rows, PAGE_SIZE // 2], - strides=[p_s0, p_s1], - block_shape=[BLOCK_H, PAGE_SIZE // 2], - ) - if USE_TMA_V_LOAD and ASSUME_FULL_HEADS and ASSUME_FULL_V: - tl.assume(out_s1 % 8 == 0) - tl.assume(out_s2 == 1) - out_desc = tl.make_tensor_descriptor( - out_ptr, - shape=[out_num_rows, V_HEAD_D], - strides=[out_s1, out_s2], - block_shape=[BLOCK_H, BLOCK_V], + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) + tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) + + +@triton.jit +def _fp4_mla_attention_prob_scale_half_kernel( + p_sf_ptr, + max_ptr, + denom_ptr, + tile_max_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + stats_s0, + tile_stats_s0, + tile_stats_s1, + NUM_HEADS: tl.constexpr, + PAGE_SIZE: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_TILES: tl.constexpr, + TILES_PER_PAGE: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_H: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + tile_rel = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + page_rel = tile_rel // TILES_PER_PAGE + tile_in_page = tile_rel - page_rel * TILES_PER_PAGE + + page_start = page_rel * PAGE_SIZE + tile_in_page * BLOCK_T + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + if page_start >= kv_len: + return + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if not ASSUME_VALID_PAGES: + if (compact_page < 0) | (compact_page >= page_ids_len): + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + + tile_max = tl.load( + tile_max_ptr + query_idx * tile_stats_s0 + tile_rel * tile_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + factor = tl.where(denom > 0.0, tl.math.exp2((tile_max - max_score) * 1.4426950408889634) / denom, 0.0) + + scale_cols = tile_in_page * (BLOCK_T // 16) + tl.arange(0, BLOCK_T // 16) + if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + compact_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + p_rows = compact_page * NUM_HEADS + offs_h + safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, compact_page * NUM_HEADS) + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) + tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) + + +@triton.jit +def _fp4_mla_attention_prob_scale_from_group_stats_kernel( + p_sf_ptr, + page_max_ptr, + page_sum_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + page_stats_s0, + page_stats_s1, + NUM_HEADS: tl.constexpr, + PAGE_SIZE: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_BY_QUERY: tl.constexpr, + BLOCK_H: tl.constexpr, + NUM_PAGE_GROUPS: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + PAGE_REL_OFFSET: tl.constexpr = 0, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_rel = tl.program_id(2) + PAGE_REL_OFFSET + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + if page_start >= kv_len: + return + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if not ASSUME_VALID_PAGES: + if (compact_page < 0) | (compact_page >= page_ids_len): + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + + page_max = tl.load( + page_max_ptr + query_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + for group_rel in tl.range(0, NUM_PAGE_GROUPS): + group_max = tl.load( + page_sum_ptr + query_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.maximum(max_score, group_max) + + denom = tl.zeros((BLOCK_H,), dtype=tl.float32) + for group_rel in tl.range(0, NUM_PAGE_GROUPS): + group_max = tl.load( + page_sum_ptr + query_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + group_sum = tl.load( + page_sum_ptr + query_idx * page_stats_s0 + (group_rel * 2 + 1) * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=0.0, + ) + denom += tl.where( + group_sum > 0.0, + group_sum * tl.math.exp2((group_max - max_score) * 1.4426950408889634), + 0.0, + ) + + factor = tl.where(denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0) + + if P_BY_QUERY: + p_page = query_idx * MAX_PAGES + page_rel + else: + p_page = compact_page + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + scale_cols = tl.arange(0, SF_PER_PAGE) + if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) + tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) + + +@triton.jit +def _fp4_mla_attention_prob_store_page_kernel( + probs_ptr, + max_ptr, + denom_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_rel, + page_ids_len, + num_pages, + probs_s0, + probs_s1, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + stats_s0, + q_num_rows, + sm_scale, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_start = page_rel * PAGE_SIZE + if (not ASSUME_FULL_PAGES) and page_start >= kv_len: + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, PAGE_SIZE) + if ASSUME_FULL_PAGES: + valid_t = tl.full([PAGE_SIZE], True, dtype=tl.int1) + else: + valid_t = page_start + offs_t < kv_len + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if (compact_page < 0) | (compact_page >= page_ids_len): + return + q_row_base = query_idx * NUM_HEADS + + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + PAGE_SIZE, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + denom_valid = denom > 0.0 + safe_denom = tl.where(denom_valid, denom, 1.0) + safe_max = tl.where(denom_valid, max_score, 0.0) + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + probs = tl.math.exp2((scores - safe_max[:, None]) * 1.4426950408889634) / safe_denom[:, None] + probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) + + prob_rows = query_idx * NUM_HEADS + offs_h + safe_prob_rows = tl.where(mask_h, prob_rows, query_idx * NUM_HEADS) + tl.store( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + offs_t[None, :] * probs_s1, + probs, + mask=mask_h[:, None], + ) + + +@triton.jit +def _fp4_mla_attention_prob_pack_page_kernel( + p_fp4_ptr, + p_sf_ptr, + probs_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_rel, + page_ids_len, + p_s0, + p_s1, + probs_s0, + probs_s1, + NUM_HEADS: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_BY_QUERY: tl.constexpr, + BLOCK_H: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + token_group = tl.program_id(1) + head_block = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_start = page_rel * PAGE_SIZE + if (not ASSUME_FULL_PAGES) and page_start >= kv_len: + return + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if (compact_page < 0) | (compact_page >= page_ids_len): + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + else: + mask_h = offs_h < NUM_HEADS + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + token_base = token_group * FP4_BLOCK + even_t = token_base + byte_offsets * 2 + odd_t = even_t + 1 + valid_even = page_start + even_t < kv_len + valid_odd = page_start + odd_t < kv_len + + prob_rows = query_idx * NUM_HEADS + offs_h + safe_prob_rows = tl.where(mask_h, prob_rows, query_idx * NUM_HEADS) + even_probs = tl.load( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + even_t[None, :] * probs_s1, + mask=mask_h[:, None] & valid_even[None, :], + other=0.0, + ) + odd_probs = tl.load( + probs_ptr + safe_prob_rows[:, None] * probs_s0 + odd_t[None, :] * probs_s1, + mask=mask_h[:, None] & valid_odd[None, :], + other=0.0, + ) + amax = tl.maximum(tl.max(tl.abs(even_probs), axis=1), tl.max(tl.abs(odd_probs), axis=1)) + local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) + + if P_BY_QUERY: + p_page = query_idx * MAX_PAGES + page_rel + else: + p_page = compact_page + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = tl.where(mask_h, p_rows, p_page * NUM_HEADS) + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows, token_group, SF_PER_PAGE) + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h) + + even_quant = _fp4_e2m1_quantize(even_probs / local_scale[:, None]) + odd_quant = _fp4_e2m1_quantize(odd_probs / local_scale[:, None]) + packed = even_quant | (odd_quant << 4) + byte_cols = token_group * (FP4_BLOCK // 2) + byte_offsets + tl.store( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + byte_cols[None, :] * p_s1, + packed, + mask=mask_h[:, None], + ) + + +@triton.jit +def _fp4_mla_attention_prob_pack_page_fused_kernel( + p_fp4_ptr, + p_sf_ptr, + max_ptr, + denom_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_rel, + page_ids_len, + num_pages, + p_s0, + p_s1, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + stats_s0, + q_num_rows, + sm_scale, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_BY_QUERY: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + PAGE_REL_FROM_GRID: tl.constexpr = False, + ASSUME_FULL_HEADS: tl.constexpr = False, + ASSUME_FULL_PAGES: tl.constexpr = False, + ASSUME_VALID_PAGES: tl.constexpr = False, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + if PAGE_REL_FROM_GRID: + page_rel = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_start = page_rel * PAGE_SIZE + if (not ASSUME_FULL_PAGES) and page_start >= kv_len: + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + offs_t = tl.arange(0, PAGE_SIZE) + if ASSUME_FULL_PAGES: + valid_t = tl.full([PAGE_SIZE], True, dtype=tl.int1) + else: + valid_t = page_start + offs_t < kv_len + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if (compact_page < 0) | (compact_page >= page_ids_len): + return + q_row_base = query_idx * NUM_HEADS + + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + PAGE_SIZE, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, + ) + max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + denom_valid = denom > 0.0 + safe_denom = tl.where(denom_valid, denom, 1.0) + safe_max = tl.where(denom_valid, max_score, 0.0) + scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + probs = tl.math.exp2((scores - safe_max[:, None]) * 1.4426950408889634) / safe_denom[:, None] + probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) + + grouped_probs = tl.reshape(probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(tl.abs(grouped_probs), axis=2) + local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) + stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) + scaled_probs = grouped_probs / tl.reshape(local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + if P_BY_QUERY: + p_page = query_idx * MAX_PAGES + page_rel + else: + p_page = compact_page + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = tl.where(mask_h, p_rows, p_page * NUM_HEADS) + scale_cols = tl.arange(0, SF_PER_PAGE) + sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h[:, None]) + + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + tl.store( + p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=mask_h[:, None, None], + ) + + +@triton.jit +def _fp4_mla_attention_pv_kernel( + out_ptr, + p_fp4_ptr, + p_sf_ptr, + kv_cache_ptr, + v_sf_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + out_s0, + out_s1, + out_s2, + out_num_rows, + p_s0, + p_s1, + p_num_rows, + kv_s0, + kv_s2, + kv_s4, + vsf_s0, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_BY_QUERY: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + USE_TMA_P_LOAD: tl.constexpr, + USE_TMA_V_LOAD: tl.constexpr, + USE_TMA_OUT_STORE: tl.constexpr, + PV_LOOP_STAGES: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_FULL_V: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + dim_block = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + if ASSUME_FULL_V: + mask_v = tl.full([BLOCK_V], True, dtype=tl.int1) + safe_offs_v = offs_v + else: + mask_v = offs_v < V_HEAD_D + safe_offs_v = tl.where(mask_v, offs_v, 0) + packed_t = tl.arange(0, PAGE_SIZE // 2) + scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) + even_t = packed_t * 2 + odd_t = even_t + 1 + v_packed_offsets = safe_offs_v // 2 + v_use_high_nibble = (safe_offs_v & 1) != 0 + if ASSUME_FULL_V and BLOCK_V == 128: + v_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + v_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE) + if USE_TMA_P_LOAD: + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + if USE_TMA_OUT_STORE and ASSUME_FULL_HEADS and ASSUME_FULL_V: + tl.assume(out_s1 % 8 == 0) + tl.assume(out_s2 == 1) + out_desc = tl.make_tensor_descriptor( + out_ptr, + shape=[out_num_rows, V_HEAD_D], + strides=[out_s1, out_s2], + block_shape=[BLOCK_H, BLOCK_V], + ) + if USE_TMA_V_LOAD: + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, PAGE_SIZE, BLOCK_V // 2], + ) + + can_use_view_pv_fast_path = ( + USE_TMA_P_LOAD + and USE_TMA_V_LOAD + and USE_TMA_OUT_STORE + and ASSUME_FULL_HEADS + and ASSUME_FULL_PAGES + and ASSUME_FULL_V + and ASSUME_VALID_PAGES + and not P_BY_QUERY + and NUM_HEADS == 128 + and V_HEAD_D == 512 + and PAGE_SIZE == 128 + and BLOCK_H == 128 + and (BLOCK_V == 128 or BLOCK_V == 256) + and SF_PER_PAGE == 8 + ) + if can_use_view_pv_fast_path: + p_view = tl.ext.make_view( + base=p_fp4_ptr, + shapes=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + tile_shape=[BLOCK_H, PAGE_SIZE // 2], + tile_dim_map=[0, 1], + ) + p_sf_view = tl.ext.make_view( + base=p_sf_ptr, + shapes=[p_num_rows // 128, SF_PER_PAGE // 4, 2, 256], + strides=[128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + v_sf_view = tl.ext.make_view( + base=v_sf_ptr, + shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], + strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, BLOCK_V // 128, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + v_view = tl.ext.make_view( + base=kv_cache_ptr, + shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], + tile_dim_map=[0, 1, 2], + ) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): + compact_page = page_table_start + page_rel + physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) + + p_vals = tl.ext.load_view_tko( + p_view, + [(compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + ) + p_vals = p_vals.to(tl.uint8, bitcast=True) + p_scales = tl.ext.load_view_tko(p_sf_view, [compact_page.to(tl.int32), 0, 0, 0]) + p_scales = p_scales.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) + p_scales = p_scales.reshape([BLOCK_H, SF_PER_PAGE]) + + v_tile = tl.ext.load_view_tko( + v_view, + [ + physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ], + ) + v_tile = v_tile.to(tl.uint8, bitcast=True) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + v_scales = tl.ext.load_view_tko( + v_sf_view, + [ + physical_page.to(tl.int32), + dim_block * (BLOCK_V // 128), + 0, + 0, + 0, + ], + ) + v_scales = v_scales.reshape([1, BLOCK_V // 128, SF_PER_PAGE // 4, 32, 4, 4]).trans( + 0, 1, 4, 3, 2, 5 + ) + v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) + acc = tl.ext.dot_scaled( + v_vals, + v_scales, + "e2m1", + p_vals.T, + p_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + out_vals = acc.T * out_scale + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + if USE_TMA_OUT_STORE: + out_desc.store( + [(query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], + out_vals, + ) + else: + tl.store( + out_ptr + query_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + ) + return + + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES or page_start < kv_len: + compact_page = page_table_start + page_rel + if ASSUME_VALID_PAGES: + safe_compact_page = compact_page + physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) + safe_physical_page = physical_page + else: + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load(src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1).to( + tl.int64 + ) + valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + + if P_BY_QUERY: + p_page = query_idx * MAX_PAGES + page_rel + else: + p_page = safe_compact_page + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + if USE_TMA_P_LOAD: + p_vals = p_desc.load([(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) + else: + p_vals = tl.load( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, + mask=mask_h[:, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None], + other=0, + ) + p_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + p_scales = tl.load(p_sf_ptr + p_sf_offsets) + + if ASSUME_FULL_PAGES: + valid_even_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + valid_odd_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + else: + valid_even_t = page_start + even_t < kv_len + valid_odd_t = page_start + odd_t < kv_len + if USE_TMA_V_LOAD: + v_tile = v_desc.load( + [ + safe_physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ] + ) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + if not ASSUME_VALID_PAGES: + v_tile = tl.where(valid_physical_page, v_tile, 0) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + if not ASSUME_FULL_PAGES: + even_packed = tl.where(valid_even_t[None, :], even_packed, 0) + odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + else: + even_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + even_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_even_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), + other=0, + ) + odd_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + odd_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_odd_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), + other=0, + ) + even_low = even_packed & 0x0F + even_high = (even_packed >> 4) & 0x0F + odd_low = odd_packed & 0x0F + even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) + odd_nibble = tl.where(v_use_high_nibble[:, None], odd_packed >> 4, odd_low) + v_vals = even_nibble | (odd_nibble << 4) + v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) + acc = tl.dot_scaled( + p_vals, + p_scales, + "e2m1", + v_vals.T, + v_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + if ASSUME_FULL_HEADS and ASSUME_FULL_V: + out_vals = acc * out_scale + if USE_TMA_OUT_STORE: + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + out_desc.store( + [(query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], + out_vals, + ) + else: + tl.store( + out_ptr + query_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + ) + else: + tl.store( + out_ptr + query_idx * out_s0 + safe_offs_h[:, None] * out_s1 + safe_offs_v[None, :] * out_s2, + acc * out_scale, + mask=mask_h[:, None] & mask_v[None, :], + ) + + +@triton.jit +def _fp4_mla_attention_pv_prepacked_v_kernel( + out_ptr, + p_fp4_ptr, + p_sf_ptr, + kv_cache_ptr, + v_sf_ptr, + v_packed_ptr, + global_scale_ptr, + max_ptr, + denom_ptr, + page_max_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + out_s0, + out_s1, + out_s2, + out_num_rows, + p_s0, + p_s1, + p_num_rows, + stats_s0, + page_stats_s0, + page_stats_s1, + kv_s0, + kv_s2, + kv_s4, + vsf_s0, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_BY_QUERY: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + USE_TMA_P_LOAD: tl.constexpr, + USE_TMA_V_LOAD: tl.constexpr, + USE_TMA_OUT_STORE: tl.constexpr, + USE_PREPACKED_V: tl.constexpr, + PV_M_PACKED_V: tl.constexpr, + PV_LOOP_STAGES: tl.constexpr, + PV_APPLY_PROB_SCALE: tl.constexpr, + PV_SCALE_IN_SF: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_FULL_V: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + dim_block = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + if ASSUME_FULL_V: + mask_v = tl.full([BLOCK_V], True, dtype=tl.int1) + safe_offs_v = offs_v + else: + mask_v = offs_v < V_HEAD_D + safe_offs_v = tl.where(mask_v, offs_v, 0) + packed_t = tl.arange(0, PAGE_SIZE // 2) + scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) + even_t = packed_t * 2 + odd_t = even_t + 1 + v_packed_offsets = safe_offs_v // 2 + v_use_high_nibble = (safe_offs_v & 1) != 0 + if ASSUME_FULL_V and BLOCK_V == 128: + v_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + v_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE) + if USE_TMA_P_LOAD: + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + if USE_TMA_OUT_STORE and ASSUME_FULL_HEADS and ASSUME_FULL_V: + tl.assume(out_s1 % 8 == 0) + tl.assume(out_s2 == 1) + out_desc = tl.make_tensor_descriptor( + out_ptr, + shape=[out_num_rows, V_HEAD_D], + strides=[out_s1, out_s2], + block_shape=[BLOCK_H, BLOCK_V], + ) + if USE_TMA_V_LOAD: + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, PAGE_SIZE, BLOCK_V // 2], + ) + if USE_PREPACKED_V: + v_packed_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + can_use_view_pv_fast_path = ( + USE_TMA_P_LOAD + and USE_TMA_V_LOAD + and USE_TMA_OUT_STORE + and ASSUME_FULL_HEADS + and ASSUME_FULL_PAGES + and ASSUME_FULL_V + and ASSUME_VALID_PAGES + and not P_BY_QUERY + and NUM_HEADS == 128 + and V_HEAD_D == 512 + and PAGE_SIZE == 128 + and BLOCK_H == 128 + and (BLOCK_V == 128 or BLOCK_V == 256) + and SF_PER_PAGE == 8 + ) + if can_use_view_pv_fast_path: + p_view = tl.ext.make_view( + base=p_fp4_ptr, + shapes=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + tile_shape=[BLOCK_H, PAGE_SIZE // 2], + tile_dim_map=[0, 1], + ) + p_sf_view = tl.ext.make_view( + base=p_sf_ptr, + shapes=[p_num_rows // 128, SF_PER_PAGE // 4, 2, 256], + strides=[128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + v_sf_view = tl.ext.make_view( + base=v_sf_ptr, + shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], + strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, BLOCK_V // 128, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + v_view = tl.ext.make_view( + base=kv_cache_ptr, + shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], + tile_dim_map=[0, 1, 2], + ) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + if PV_APPLY_PROB_SCALE: + max_score = tl.load(max_ptr + query_idx * stats_s0 + offs_h) + denom = tl.load(denom_ptr + query_idx * stats_s0 + offs_h) + else: + max_score = tl.zeros((BLOCK_H,), dtype=tl.float32) + denom = tl.full((BLOCK_H,), 1.0, dtype=tl.float32) + acc = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): + compact_page = page_table_start + page_rel + physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) + if P_BY_QUERY: + p_page = query_idx * MAX_PAGES + page_rel + else: + p_page = compact_page + + p_vals = tl.ext.load_view_tko( + p_view, + [(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + ) + p_vals = p_vals.to(tl.uint8, bitcast=True) + p_scales = tl.ext.load_view_tko(p_sf_view, [p_page.to(tl.int32), 0, 0, 0]) + p_scales = p_scales.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) + p_scales = p_scales.reshape([BLOCK_H, SF_PER_PAGE]) + + if USE_PREPACKED_V: + v_row = (physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + else: + v_tile = tl.ext.load_view_tko( + v_view, + [ + physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ], + ) + v_tile = v_tile.to(tl.uint8, bitcast=True) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + if not PV_M_PACKED_V: + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + v_scales = tl.ext.load_view_tko( + v_sf_view, + [ + physical_page.to(tl.int32), + dim_block * (BLOCK_V // 128), + 0, + 0, + 0, + ], + ) + v_scales = v_scales.reshape([1, BLOCK_V // 128, SF_PER_PAGE // 4, 32, 4, 4]).trans( + 0, 1, 4, 3, 2, 5 + ) + v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) + if PV_M_PACKED_V and not USE_PREPACKED_V: + acc = tl.ext.dot_scaled( + v_tile.T, + v_scales, + "e2m1", + p_vals.T, + p_scales, + "e2m1", + acc=acc, + fast_math=True, + lhs_k_pack=False, + rhs_k_pack=True, + ) + else: + if PV_APPLY_PROB_SCALE: + page_max = tl.load( + page_max_ptr + query_idx * page_stats_s0 + page_rel * page_stats_s1 + offs_h + ) + factor = tl.where( + denom > 0.0, + tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, + 0.0, + ) + if PV_SCALE_IN_SF: + p_scales_with_factor = (p_scales.to(tl.float32) * factor[:, None]).to(tl.float8e4nv) + acc = tl.ext.dot_scaled( + v_vals, + v_scales, + "e2m1", + p_vals.T, + p_scales_with_factor, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + else: + page_acc = tl.ext.dot_scaled( + v_vals, + v_scales, + "e2m1", + p_vals.T, + p_scales, + "e2m1", + acc=tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32), + fast_math=True, + rhs_k_pack=True, + ) + acc += page_acc * factor[None, :] + else: + acc = tl.ext.dot_scaled( + v_vals, + v_scales, + "e2m1", + p_vals.T, + p_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + out_vals = acc.T * out_scale + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + if USE_TMA_OUT_STORE: + out_desc.store( + [(query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], + out_vals, + ) + else: + tl.store( + out_ptr + query_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + ) + return + + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES or page_start < kv_len: + compact_page = page_table_start + page_rel + if ASSUME_VALID_PAGES: + safe_compact_page = compact_page + physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) + safe_physical_page = physical_page + else: + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load(src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1).to( + tl.int64 + ) + valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + + if P_BY_QUERY: + p_page = query_idx * MAX_PAGES + page_rel + else: + p_page = safe_compact_page + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + if USE_TMA_P_LOAD: + p_vals = p_desc.load([(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) + else: + p_vals = tl.load( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, + mask=mask_h[:, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None], + other=0, + ) + p_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + p_scales = tl.load(p_sf_ptr + p_sf_offsets) + + if ASSUME_FULL_PAGES: + valid_even_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + valid_odd_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + else: + valid_even_t = page_start + even_t < kv_len + valid_odd_t = page_start + odd_t < kv_len + if USE_PREPACKED_V: + v_row = (safe_physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + if not ASSUME_VALID_PAGES: + v_vals = tl.where(valid_physical_page, v_vals, 0) + elif USE_TMA_V_LOAD: + v_tile = v_desc.load( + [ + safe_physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ] + ) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + if not ASSUME_VALID_PAGES: + v_tile = tl.where(valid_physical_page, v_tile, 0) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + if not ASSUME_FULL_PAGES: + even_packed = tl.where(valid_even_t[None, :], even_packed, 0) + odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) + low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + else: + even_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + even_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_even_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), + other=0, + ) + odd_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + odd_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_odd_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), + other=0, + ) + even_low = even_packed & 0x0F + even_high = (even_packed >> 4) & 0x0F + odd_low = odd_packed & 0x0F + even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) + odd_nibble = tl.where(v_use_high_nibble[:, None], odd_packed >> 4, odd_low) + v_vals = even_nibble | (odd_nibble << 4) + v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) + acc = tl.dot_scaled( + p_vals, + p_scales, + "e2m1", + v_vals.T, + v_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + if ASSUME_FULL_HEADS and ASSUME_FULL_V: + out_vals = acc * out_scale + if USE_TMA_OUT_STORE: + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + out_desc.store( + [(query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], + out_vals, + ) + else: + tl.store( + out_ptr + query_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + ) + else: + tl.store( + out_ptr + query_idx * out_s0 + safe_offs_h[:, None] * out_s1 + safe_offs_v[None, :] * out_s2, + acc * out_scale, + mask=mask_h[:, None] & mask_v[None, :], + ) + + +@triton.jit +def _fp4_mla_attention_pv_mtp_pair_prepacked_v_kernel( + out_ptr, + p_fp4_ptr, + p_sf_ptr, + v_sf_ptr, + v_packed_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + num_pages, + out_s0, + out_s1, + out_s2, + p_s0, + p_s1, + p_num_rows: tl.constexpr, + vsf_s0: tl.constexpr, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + NUM_DIM_BLOCKS: tl.constexpr, + Q_PER_GROUP: tl.constexpr, + PV_LOOP_STAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + seq_idx = tl.program_id(0) + head_block = tl.program_id(1) + combo = tl.program_id(2) + query_group = combo // NUM_DIM_BLOCKS + dim_block = combo - query_group * NUM_DIM_BLOCKS + query_base = seq_idx * QUERY_LEN_PER_SEQ + query_group * Q_PER_GROUP + query_idx0 = query_base + query_idx1 = query_base + 1 + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_view = tl.ext.make_view( + base=p_fp4_ptr, + shapes=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + tile_shape=[BLOCK_H, PAGE_SIZE // 2], + tile_dim_map=[0, 1], + ) + p_sf_view = tl.ext.make_view( + base=p_sf_ptr, + shapes=[p_num_rows // 128, SF_PER_PAGE // 4, 2, 256], + strides=[128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + v_sf_view = tl.ext.make_view( + base=v_sf_ptr, + shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], + strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, BLOCK_V // 128, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + v_packed_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + acc0 = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + acc1 = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): + compact_page = page_table_start + page_rel + physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) + + p_page0 = query_idx0 * MAX_PAGES + page_rel + p_vals0 = tl.ext.load_view_tko( + p_view, + [(p_page0 * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], ) - if USE_TMA_V_LOAD: - tl.assume(kv_s0 % 8 == 0) - tl.assume(kv_s2 % 8 == 0) - tl.assume(kv_s4 == 1) - v_desc = tl.make_tensor_descriptor( + p_vals0 = p_vals0.to(tl.uint8, bitcast=True) + p_scales0 = tl.ext.load_view_tko(p_sf_view, [p_page0.to(tl.int32), 0, 0, 0]) + p_scales0 = p_scales0.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) + p_scales0 = p_scales0.reshape([BLOCK_H, SF_PER_PAGE]) + + p_page1 = query_idx1 * MAX_PAGES + page_rel + p_vals1 = tl.ext.load_view_tko( + p_view, + [(p_page1 * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + ) + p_vals1 = p_vals1.to(tl.uint8, bitcast=True) + p_scales1 = tl.ext.load_view_tko(p_sf_view, [p_page1.to(tl.int32), 0, 0, 0]) + p_scales1 = p_scales1.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) + p_scales1 = p_scales1.reshape([BLOCK_H, SF_PER_PAGE]) + + v_row = (physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + v_scales = tl.ext.load_view_tko( + v_sf_view, + [ + physical_page.to(tl.int32), + dim_block * (BLOCK_V // 128), + 0, + 0, + 0, + ], + ) + v_scales = v_scales.reshape([1, BLOCK_V // 128, SF_PER_PAGE // 4, 32, 4, 4]).trans( + 0, 1, 4, 3, 2, 5 + ) + v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) + + acc0 = tl.ext.dot_scaled( + v_vals, + v_scales, + "e2m1", + p_vals0.T, + p_scales0, + "e2m1", + acc=acc0, + fast_math=True, + rhs_k_pack=True, + ) + acc1 = tl.ext.dot_scaled( + v_vals, + v_scales, + "e2m1", + p_vals1.T, + p_scales1, + "e2m1", + acc=acc1, + fast_math=True, + rhs_k_pack=True, + ) + + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + out_vals0 = acc0.T * out_scale + out_vals1 = acc1.T * out_scale + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals0 = out_vals0.to(tl.bfloat16) + out_vals1 = out_vals1.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals0 = out_vals0.to(tl.float16) + out_vals1 = out_vals1.to(tl.float16) + tl.store( + out_ptr + query_idx0 * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals0, + ) + tl.store( + out_ptr + query_idx1 * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals1, + ) + + +@triton.jit +def _fp4_mla_attention_online_qkpv_group_kernel( + partial_o_ptr, + partial_m_ptr, + partial_l_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + v_sf_ptr, + v_packed_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + num_pages: tl.constexpr, + q_fp4_s0: tl.constexpr, + q_fp4_s1: tl.constexpr, + kv_s0: tl.constexpr, + kv_s2: tl.constexpr, + kv_s4: tl.constexpr, + sf_s0: tl.constexpr, + vsf_s0: tl.constexpr, + po_s0: tl.constexpr, + po_s1: tl.constexpr, + po_s2: tl.constexpr, + po_s3: tl.constexpr, + pm_s0: tl.constexpr, + pm_s1: tl.constexpr, + q_num_rows: tl.constexpr, + sm_scale: tl.constexpr, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + BLOCK_V: tl.constexpr, + MAX_PAGES: tl.constexpr, + GROUP_PAGES: tl.constexpr, + NUM_DIM_BLOCKS: tl.constexpr, + ALLOW_PARTIAL_GROUPS: tl.constexpr, + FP4_PV: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + combo = tl.program_id(2) + page_group = combo // NUM_DIM_BLOCKS + dim_block = combo - page_group * NUM_DIM_BLOCKS + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_t = tl.arange(0, BLOCK_T) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + packed_t = tl.arange(0, PAGE_SIZE // 2) + scale_cols = tl.arange(0, SF_PER_PAGE) + q_row_base = gen_idx * NUM_HEADS + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + + tl.assume(q_fp4_s0 % 8 == 0) + tl.assume(q_fp4_s1 == 1) + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + q_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 256], + ) + k_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 256], + ) + q_tail_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, 32], + ) + k_tail_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, 32], + ) + v_packed_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + q_sf_full_view = tl.ext.make_view( + base=q_sf_ptr, + shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + q_sf_tail_view = tl.ext.make_view( + base=q_sf_ptr, + shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + k_sf_full_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 8, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + k_sf_tail_view = tl.ext.make_view( + base=sf_cache_ptr, + shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], + strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, 1, 1, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + q_row_start = (q_row_base + head_block * BLOCK_H).to(tl.int32) + q_row_group = q_row_base // 128 + full_q_vals = q_desc.load([q_row_start, 0]) + full_q_scales = tl.ext.load_view_tko(q_sf_full_view, [q_row_group.to(tl.int32), 0, 0, 0]) + full_q_scales = full_q_scales.reshape([1, 8, 32, 4, 4]).trans(0, 3, 2, 1, 4) + full_q_scales = full_q_scales.reshape([BLOCK_H, 32]) + q0_vals = q_tail_desc.load([q_row_start, 256]) + q1_vals = q_tail_desc.load([q_row_start, 288]) + q0_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group.to(tl.int32), 8, 0, 0]) + q0_scales = q0_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) + q0_scales = q0_scales.reshape([BLOCK_H, 4]) + q1_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group.to(tl.int32), 9, 0, 0]) + q1_scales = q1_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) + q1_scales = q1_scales.reshape([BLOCK_H, 4]) + + group_m = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + group_l = tl.zeros((BLOCK_H,), dtype=tl.float32) + group_o = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + for page_group_off in tl.range(0, GROUP_PAGES): + page_rel = page_group * GROUP_PAGES + page_group_off + valid_group_page = page_rel < MAX_PAGES + compact_page = page_table_start + page_rel + safe_compact_page = tl.where(valid_group_page, compact_page, page_table_start) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, + mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS), + other=0, + ).to(tl.int64) + + scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) + full_k_vals = k_desc.load([physical_page.to(tl.int32), 0, 0]) + full_k_vals = tl.reshape(full_k_vals, (BLOCK_T, 256)) + full_k_scales = tl.ext.load_view_tko(k_sf_full_view, [physical_page.to(tl.int32), 0, 0, 0, 0]) + full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) + scores = tl.dot_scaled( + full_q_vals, + full_q_scales, + "e2m1", + full_k_vals.T, + full_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + tail_k_vals = k_tail_desc.load([physical_page.to(tl.int32), 0, 256]) + tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) + tail_k_scales = tl.ext.load_view_tko(k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0]) + tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) + tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) + scores = tl.dot_scaled( + q0_vals, + q0_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + scores = tl.dot_scaled( + q1_vals, + q1_scales, + "e2m1", + tail_k_vals.T, + tail_k_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + scores = scores * qk_scale + page_m = tl.max(scores, axis=1) + if ALLOW_PARTIAL_GROUPS: + page_m = tl.where(valid_group_page, page_m, -float("inf")) + safe_page_m = tl.where(valid_group_page, page_m, 0.0) + exp_scores = tl.math.exp2((scores - safe_page_m[:, None]) * 1.4426950408889634) + exp_scores = tl.where(valid_group_page, exp_scores, 0.0) + else: + exp_scores = tl.math.exp2((scores - page_m[:, None]) * 1.4426950408889634) + page_l = tl.sum(exp_scores, axis=1) + + v_row = (physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + if FP4_PV: + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + p_scales = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) + p_scales = p_scales.to(tl.float8e4nv) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + p_vals = tl.reshape(_fp4_e2m1_quantize_packed(even_probs, odd_probs), (BLOCK_H, PAGE_SIZE // 2)) + v_scale_offsets = _fp4_mla_swizzled_sf_offset(offs_v[:, None], scale_cols[None, :], SF_PER_PAGE) + v_scales = tl.load(v_sf_ptr + physical_page * vsf_s0 + v_scale_offsets) + page_o = tl.ext.dot_scaled( + v_vals, + v_scales, + "e2m1", + p_vals.T, + p_scales, + "e2m1", + acc=tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32), + fast_math=True, + rhs_k_pack=True, + ) + else: + v_scale_cols = packed_t // (FP4_BLOCK // 2) + v_scale_offsets = _fp4_mla_swizzled_sf_offset(offs_v[:, None], v_scale_cols[None, :], SF_PER_PAGE) + v_scales = tl.load(v_sf_ptr + physical_page * vsf_s0 + v_scale_offsets).to(tl.float32) + v_low = _fp4_e2m1_to_f32(v_vals & 0x0F) * v_scales + v_high = _fp4_e2m1_to_f32((v_vals >> 4) & 0x0F) * v_scales + exp_pairs = tl.reshape(exp_scores, (BLOCK_H, PAGE_SIZE // 2, 2)) + exp_even, exp_odd = tl.split(exp_pairs) + page_o = tl.dot(v_low, exp_even.T, out_dtype=tl.float32) + tl.dot( + v_high, exp_odd.T, out_dtype=tl.float32 + ) + + next_m = tl.maximum(group_m, page_m) + old_delta = tl.where(group_l > 0.0, group_m - next_m, 0.0) + new_delta = tl.where(page_l > 0.0, page_m - next_m, 0.0) + old_scale = tl.math.exp2(old_delta * 1.4426950408889634) + new_scale = tl.math.exp2(new_delta * 1.4426950408889634) + group_o = group_o * old_scale[None, :] + page_o * new_scale[None, :] + group_l = group_l * old_scale + page_l * new_scale + group_m = next_m + + partial_o_offsets = ( + page_group * po_s0 + + gen_idx * po_s1 + + offs_h[:, None] * po_s2 + + offs_v[None, :] * po_s3 + ) + tl.store(partial_o_ptr + partial_o_offsets, group_o.T) + if dim_block == 0: + partial_ml_offsets = page_group * pm_s0 + gen_idx * pm_s1 + offs_h + tl.store(partial_m_ptr + partial_ml_offsets, group_m) + tl.store(partial_l_ptr + partial_ml_offsets, group_l) + + +@triton.jit +def _fp4_mla_pv_page_o_prepacked_raw( + v_sf_ptr, + v_packed_ptr, + physical_page, + dim_block: tl.constexpr, + p_vals, + p_scales, + num_pages: tl.constexpr, + vsf_s0: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, +): + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + packed_t = tl.arange(0, PAGE_SIZE // 2) + num_dim_blocks: tl.constexpr = V_HEAD_D // BLOCK_V + v_sf_view = tl.ext.make_view( + base=v_sf_ptr, + shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], + strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, BLOCK_V // 128, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + v_row = (physical_page * num_dim_blocks + dim_block) * BLOCK_V + v_vals = tl.load( + v_packed_ptr + (v_row + offs_v[:, None]) * (PAGE_SIZE // 2) + packed_t[None, :] + ) + v_scales = tl.ext.load_view_tko( + v_sf_view, + [ + physical_page.to(tl.int32), + dim_block * (BLOCK_V // 128), + 0, + 0, + 0, + ], + ) + v_scales = v_scales.reshape([1, BLOCK_V // 128, SF_PER_PAGE // 4, 32, 4, 4]).trans( + 0, 1, 4, 3, 2, 5 + ) + v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) + return tl.ext.dot_scaled( + v_vals, + v_scales, + "e2m1", + p_vals.T, + p_scales, + "e2m1", + acc=tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32), + fast_math=True, + rhs_k_pack=True, + ) + + +@triton.jit +def _fp4_mla_attention_gen_qkpv_group_kernel( + partial_o_ptr, + partial_m_ptr, + partial_l_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + v_sf_ptr, + v_packed_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len: tl.constexpr, + num_pages: tl.constexpr, + q_fp4_s0: tl.constexpr, + q_fp4_s1: tl.constexpr, + kv_s0: tl.constexpr, + kv_s2: tl.constexpr, + kv_s4: tl.constexpr, + sf_s0: tl.constexpr, + vsf_s0: tl.constexpr, + po_s0: tl.constexpr, + po_s1: tl.constexpr, + po_s2: tl.constexpr, + po_s3: tl.constexpr, + pm_s0: tl.constexpr, + pm_s1: tl.constexpr, + q_num_rows: tl.constexpr, + sm_scale: tl.constexpr, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + BLOCK_V: tl.constexpr, + MAX_PAGES: tl.constexpr, + GROUP_PAGES: tl.constexpr, + ALLOW_PARTIAL_GROUPS: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_group = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_t = tl.arange(0, BLOCK_T) + q_row_base = query_idx * NUM_HEADS + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + + group_m = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + group_l = tl.zeros((BLOCK_H,), dtype=tl.float32) + group_o0 = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + group_o1 = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + group_o2 = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + group_o3 = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + + for page_group_off in tl.range(0, GROUP_PAGES): + page_rel = page_group * GROUP_PAGES + page_group_off + page_start = page_rel * PAGE_SIZE + valid_group_page = page_rel < MAX_PAGES + valid_page_tokens = valid_group_page & (page_start < kv_len) + compact_page = page_table_start + page_rel + safe_compact_page = tl.where(valid_group_page, compact_page, page_table_start) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, + mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS), + other=0, + ).to(tl.int64) + + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, kv_cache_ptr, - shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - block_shape=[1, PAGE_SIZE, BLOCK_V // 2], + sf_cache_ptr, + src_page_ids_ptr, + safe_compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + BLOCK_T, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, ) + valid_t = page_start + offs_t < kv_len + scores = tl.where(valid_page_tokens & valid_t[None, :], scores * qk_scale, -float("inf")) + page_m = tl.max(scores, axis=1) + safe_page_m = tl.where(valid_page_tokens, page_m, 0.0) + exp_scores = tl.math.exp2((scores - safe_page_m[:, None]) * 1.4426950408889634) + exp_scores = tl.where(valid_page_tokens & valid_t[None, :], exp_scores, 0.0) + page_l = tl.sum(exp_scores, axis=1) - if ( - USE_TMA_P_LOAD - and USE_TMA_V_LOAD - and ASSUME_FULL_HEADS - and ASSUME_FULL_PAGES - and ASSUME_FULL_V - and ASSUME_VALID_PAGES - and NUM_HEADS == 128 - and V_HEAD_D == 512 - and PAGE_SIZE == 128 - and BLOCK_H == 128 - and BLOCK_V == 128 - and SF_PER_PAGE == 8 - ): - p_view = tl.ext.make_view( - base=p_fp4_ptr, - shapes=[p_num_rows, PAGE_SIZE // 2], - strides=[p_s0, p_s1], - tile_shape=[BLOCK_H, PAGE_SIZE // 2], - tile_dim_map=[0, 1], + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + p_scales = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) + p_scales = p_scales.to(tl.float8e4nv) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + p_vals = tl.reshape(_fp4_e2m1_quantize_packed(even_probs, odd_probs), (BLOCK_H, PAGE_SIZE // 2)) + + page_o0 = _fp4_mla_pv_page_o_prepacked_raw( + v_sf_ptr, + v_packed_ptr, + physical_page, + 0, + p_vals, + p_scales, + num_pages, + vsf_s0, + V_HEAD_D, + PAGE_SIZE, + SF_PER_PAGE, + BLOCK_H, + BLOCK_V, ) - p_sf_view = tl.ext.make_view( - base=p_sf_ptr, - shapes=[p_num_rows // 128, SF_PER_PAGE // 4, 2, 256], - strides=[128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, SF_PER_PAGE // 4, 2, 256], - tile_dim_map=[0, 1, 2, 3], + page_o1 = _fp4_mla_pv_page_o_prepacked_raw( + v_sf_ptr, + v_packed_ptr, + physical_page, + 1, + p_vals, + p_scales, + num_pages, + vsf_s0, + V_HEAD_D, + PAGE_SIZE, + SF_PER_PAGE, + BLOCK_H, + BLOCK_V, + ) + page_o2 = _fp4_mla_pv_page_o_prepacked_raw( + v_sf_ptr, + v_packed_ptr, + physical_page, + 2, + p_vals, + p_scales, + num_pages, + vsf_s0, + V_HEAD_D, + PAGE_SIZE, + SF_PER_PAGE, + BLOCK_H, + BLOCK_V, + ) + page_o3 = _fp4_mla_pv_page_o_prepacked_raw( + v_sf_ptr, + v_packed_ptr, + physical_page, + 3, + p_vals, + p_scales, + num_pages, + vsf_s0, + V_HEAD_D, + PAGE_SIZE, + SF_PER_PAGE, + BLOCK_H, + BLOCK_V, + ) + + next_m = tl.maximum(group_m, page_m) + old_delta = tl.where(group_l > 0.0, group_m - next_m, 0.0) + new_delta = tl.where(page_l > 0.0, page_m - next_m, 0.0) + old_scale = tl.math.exp2(old_delta * 1.4426950408889634) + new_scale = tl.math.exp2(new_delta * 1.4426950408889634) + group_o0 = group_o0 * old_scale[None, :] + page_o0 * new_scale[None, :] + group_o1 = group_o1 * old_scale[None, :] + page_o1 * new_scale[None, :] + group_o2 = group_o2 * old_scale[None, :] + page_o2 * new_scale[None, :] + group_o3 = group_o3 * old_scale[None, :] + page_o3 * new_scale[None, :] + group_l = group_l * old_scale + page_l * new_scale + group_m = next_m + + offs_v = tl.arange(0, BLOCK_V) + partial_base = page_group * po_s0 + query_idx * po_s1 + offs_h[:, None] * po_s2 + tl.store(partial_o_ptr + partial_base + offs_v[None, :] * po_s3, group_o0.T) + tl.store(partial_o_ptr + partial_base + (BLOCK_V + offs_v)[None, :] * po_s3, group_o1.T) + tl.store(partial_o_ptr + partial_base + (2 * BLOCK_V + offs_v)[None, :] * po_s3, group_o2.T) + tl.store(partial_o_ptr + partial_base + (3 * BLOCK_V + offs_v)[None, :] * po_s3, group_o3.T) + partial_ml_offsets = page_group * pm_s0 + query_idx * pm_s1 + offs_h + tl.store(partial_m_ptr + partial_ml_offsets, group_m) + tl.store(partial_l_ptr + partial_ml_offsets, group_l) + + +@triton.jit +def _fp4_mla_attention_mtp_fused_qkpv_group_kernel( + partial_o_ptr, + partial_m_ptr, + partial_l_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + v_sf_ptr, + v_packed_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len: tl.constexpr, + num_pages: tl.constexpr, + q_fp4_s0: tl.constexpr, + q_fp4_s1: tl.constexpr, + kv_s0: tl.constexpr, + kv_s2: tl.constexpr, + kv_s4: tl.constexpr, + sf_s0: tl.constexpr, + vsf_s0: tl.constexpr, + po_s0: tl.constexpr, + po_s1: tl.constexpr, + po_s2: tl.constexpr, + po_s3: tl.constexpr, + pm_s0: tl.constexpr, + pm_s1: tl.constexpr, + q_num_rows: tl.constexpr, + sm_scale: tl.constexpr, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + BLOCK_V: tl.constexpr, + MAX_PAGES: tl.constexpr, + GROUP_PAGES: tl.constexpr, + NUM_DIM_BLOCKS: tl.constexpr, + ALLOW_PARTIAL_GROUPS: tl.constexpr, + SPLIT_PV_K: tl.constexpr, + USE_TMA_DATA_LOAD: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + combo = tl.program_id(2) + page_group = combo // NUM_DIM_BLOCKS + dim_block = combo - page_group * NUM_DIM_BLOCKS + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_t = tl.arange(0, BLOCK_T) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + q_row_base = query_idx * NUM_HEADS + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + global_scale = tl.load(global_scale_ptr) + qk_scale = sm_scale / (global_scale * global_scale) + + v_packed_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * NUM_DIM_BLOCKS * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + v_sf_view = tl.ext.make_view( + base=v_sf_ptr, + shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], + strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, BLOCK_V // 128, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + + group_m = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + group_l = tl.zeros((BLOCK_H,), dtype=tl.float32) + group_o = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + for page_group_off in tl.range(0, GROUP_PAGES): + page_rel = page_group * GROUP_PAGES + page_group_off + page_start = page_rel * PAGE_SIZE + valid_group_page = page_rel < MAX_PAGES + valid_page_tokens = valid_group_page & (page_start < kv_len) + compact_page = page_table_start + page_rel + safe_compact_page = tl.where(valid_group_page, compact_page, page_table_start) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, + mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS), + other=0, + ).to(tl.int64) + + scores = _fp4_mla_qk_scores_tile( + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + src_page_ids_ptr, + safe_compact_page, + q_row_base, + head_block * BLOCK_H, + offs_h, + offs_t, + q_num_rows, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, + page_ids_len, + num_pages, + Q_HEAD_D, + K_HEAD_D, + Q_RESIDUAL_D, + FP4_BLOCK, + Q_SF_PER_TOKEN, + K_SF_PER_TOKEN, + BLOCK_H, + BLOCK_T, + BLOCK_K, + FULL_BLOCK_END, + TAIL_BLOCK_K, + NUM_HEADS, + USE_TMA_DATA_LOAD, + ASSUME_FULL_HEADS, + ASSUME_VALID_PAGES, ) - v_sf_view = tl.ext.make_view( - base=v_sf_ptr, - shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], - strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, 1, SF_PER_PAGE // 4, 2, 256], - tile_dim_map=[0, 1, 2, 3, 4], + valid_t = page_start + offs_t < kv_len + scores = tl.where(valid_page_tokens & valid_t[None, :], scores * qk_scale, -float("inf")) + page_m = tl.max(scores, axis=1) + safe_page_m = tl.where(valid_page_tokens, page_m, 0.0) + exp_scores = tl.math.exp2((scores - safe_page_m[:, None]) * 1.4426950408889634) + exp_scores = tl.where(valid_page_tokens & valid_t[None, :], exp_scores, 0.0) + page_l = tl.sum(exp_scores, axis=1) + + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + p_scales = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) + p_scales = p_scales.to(tl.float8e4nv) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + p_vals = tl.reshape(_fp4_e2m1_quantize_packed(even_probs, odd_probs), (BLOCK_H, PAGE_SIZE // 2)) + + v_row = (physical_page * NUM_DIM_BLOCKS + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + v_scales = tl.ext.load_view_tko( + v_sf_view, + [ + physical_page.to(tl.int32), + dim_block * (BLOCK_V // 128), + 0, + 0, + 0, + ], ) - v_view = tl.ext.make_view( - base=kv_cache_ptr, - shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], - tile_dim_map=[0, 1, 2], + v_scales = v_scales.reshape([1, BLOCK_V // 128, SF_PER_PAGE // 4, 32, 4, 4]).trans( + 0, 1, 4, 3, 2, 5 ) - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) - global_scale = tl.load(global_scale_ptr) - out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) - acc = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) - for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): - compact_page = page_table_start + page_rel - physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) - - p_vals = tl.ext.load_view_tko( - p_view, - [(compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], - ) - p_vals = p_vals.to(tl.uint8, bitcast=True) - p_scales = tl.ext.load_view_tko(p_sf_view, [compact_page.to(tl.int32), 0, 0, 0]) - p_scales = p_scales.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) - p_scales = p_scales.reshape([BLOCK_H, SF_PER_PAGE]) - - v_tile = tl.ext.load_view_tko( - v_view, - [ - physical_page.to(tl.int32), - 0, - (dim_block * (BLOCK_V // 2)).to(tl.int32), - ], - ) - v_tile = v_tile.to(tl.uint8, bitcast=True) - v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) - v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) - even_packed, odd_packed = tl.split(v_pairs) - low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) - v_vals = tl.reshape( - tl.join(low_vals, high_vals).permute(0, 2, 1), - (BLOCK_V, PAGE_SIZE // 2), + v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) + if SPLIT_PV_K: + v_val_halves = tl.reshape(v_vals, (BLOCK_V, 2, PAGE_SIZE // 4)).trans(0, 2, 1) + v_vals0, v_vals1 = tl.split(v_val_halves) + v_scale_halves = tl.reshape(v_scales, (BLOCK_V, 2, SF_PER_PAGE // 2)).trans(0, 2, 1) + v_scales0, v_scales1 = tl.split(v_scale_halves) + p_val_halves = tl.reshape(p_vals, (BLOCK_H, 2, PAGE_SIZE // 4)).trans(0, 2, 1) + p_vals0, p_vals1 = tl.split(p_val_halves) + p_scale_halves = tl.reshape(p_scales, (BLOCK_H, 2, SF_PER_PAGE // 2)).trans(0, 2, 1) + p_scales0, p_scales1 = tl.split(p_scale_halves) + page_o = tl.ext.dot_scaled( + v_vals0, + v_scales0, + "e2m1", + p_vals0.T, + p_scales0, + "e2m1", + acc=tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32), + fast_math=True, + rhs_k_pack=True, ) - v_scales = tl.ext.load_view_tko( - v_sf_view, - [ - physical_page.to(tl.int32), - dim_block, - 0, - 0, - 0, - ], + page_o = tl.ext.dot_scaled( + v_vals1, + v_scales1, + "e2m1", + p_vals1.T, + p_scales1, + "e2m1", + acc=page_o, + fast_math=True, + rhs_k_pack=True, ) - v_scales = v_scales.reshape([1, 1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) - v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) - acc = tl.ext.dot_scaled( + else: + page_o = tl.ext.dot_scaled( v_vals, v_scales, "e2m1", p_vals.T, p_scales, "e2m1", - acc=acc, + acc=tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32), fast_math=True, rhs_k_pack=True, ) - out_vals = acc.T * out_scale - if out_ptr.dtype.element_ty == tl.bfloat16: - out_vals = out_vals.to(tl.bfloat16) - elif out_ptr.dtype.element_ty == tl.float16: - out_vals = out_vals.to(tl.float16) - out_desc.store( - [(gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], - out_vals, + next_m = tl.maximum(group_m, page_m) + old_delta = tl.where(group_l > 0.0, group_m - next_m, 0.0) + new_delta = tl.where(page_l > 0.0, page_m - next_m, 0.0) + old_scale = tl.math.exp2(old_delta * 1.4426950408889634) + new_scale = tl.math.exp2(new_delta * 1.4426950408889634) + group_o = group_o * old_scale[None, :] + page_o * new_scale[None, :] + group_l = group_l * old_scale + page_l * new_scale + group_m = next_m + + partial_o_offsets = ( + page_group * po_s0 + + query_idx * po_s1 + + offs_h[:, None] * po_s2 + + offs_v[None, :] * po_s3 + ) + tl.store(partial_o_ptr + partial_o_offsets, group_o.T) + partial_ml_offsets = page_group * pm_s0 + query_idx * pm_s1 + offs_h + tl.store(partial_m_ptr + partial_ml_offsets, group_m) + tl.store(partial_l_ptr + partial_ml_offsets, group_l) + + +@triton.jit +def _fp4_mla_attention_online_qkpv_reduce_kernel( + out_ptr, + partial_o_ptr, + partial_m_ptr, + partial_l_ptr, + global_scale_ptr, + out_s0, + out_s1, + out_s2, + po_s0: tl.constexpr, + po_s1: tl.constexpr, + po_s2: tl.constexpr, + po_s3: tl.constexpr, + pm_s0: tl.constexpr, + pm_s1: tl.constexpr, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + NUM_PAGE_GROUPS: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + FP4_PV: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + dim_block = tl.program_id(2) + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + + global_m = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + for group_idx in tl.range(0, NUM_PAGE_GROUPS): + group_m = tl.load(partial_m_ptr + group_idx * pm_s0 + gen_idx * pm_s1 + offs_h) + global_m = tl.maximum(global_m, group_m) + + global_l = tl.zeros((BLOCK_H,), dtype=tl.float32) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + for group_idx in tl.range(0, NUM_PAGE_GROUPS): + group_m = tl.load(partial_m_ptr + group_idx * pm_s0 + gen_idx * pm_s1 + offs_h) + group_l = tl.load(partial_l_ptr + group_idx * pm_s0 + gen_idx * pm_s1 + offs_h) + scale = tl.where(group_l > 0.0, tl.math.exp2((group_m - global_m) * 1.4426950408889634), 0.0) + partial_o = tl.load( + partial_o_ptr + + group_idx * po_s0 + + gen_idx * po_s1 + + offs_h[:, None] * po_s2 + + offs_v[None, :] * po_s3 ) - return + acc += partial_o * scale[:, None] + global_l += group_l * scale - if ASSUME_FULL_PAGES: - kv_len = 0 - else: - kv_len = tl.load(kv_lens_ptr + gen_idx) - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) global_scale = tl.load(global_scale_ptr) - out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) - acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) - for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): - page_start = page_rel * PAGE_SIZE - if ASSUME_FULL_PAGES or page_start < kv_len: - compact_page = page_table_start + page_rel - if ASSUME_VALID_PAGES: - safe_compact_page = compact_page - physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) - safe_physical_page = physical_page - else: - valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) - safe_compact_page = tl.where(valid_compact_page, compact_page, 0) - physical_page = tl.load(src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1).to( - tl.int64 - ) - valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) - safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + out_scale = 1.0 / global_scale + if FP4_PV: + out_scale = out_scale / P_GLOBAL_SCALE + safe_l = tl.where(global_l > 0.0, global_l, 1.0) + out_vals = (acc / safe_l[:, None]) * out_scale + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + tl.store( + out_ptr + gen_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + ) - p_rows = safe_compact_page * NUM_HEADS + offs_h - safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) - if USE_TMA_P_LOAD: - p_vals = p_desc.load([(safe_compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) - else: - p_vals = tl.load( - p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, - mask=mask_h[:, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None], - other=0, - ) - p_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) - p_scales = tl.load(p_sf_ptr + p_sf_offsets) - if ASSUME_FULL_PAGES: - valid_even_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) - valid_odd_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) - else: - valid_even_t = page_start + even_t < kv_len - valid_odd_t = page_start + odd_t < kv_len - if USE_PREPACKED_V: - v_row = (safe_physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V - v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) - elif USE_TMA_V_LOAD: - v_tile = v_desc.load( - [ - safe_physical_page.to(tl.int32), - 0, - (dim_block * (BLOCK_V // 2)).to(tl.int32), - ] - ) - v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) - if not ASSUME_VALID_PAGES: - v_tile = tl.where(valid_physical_page, v_tile, 0) - v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) - even_packed, odd_packed = tl.split(v_pairs) - if not ASSUME_FULL_PAGES: - even_packed = tl.where(valid_even_t[None, :], even_packed, 0) - odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) - low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) - v_vals = tl.reshape( - tl.join(low_vals, high_vals).permute(0, 2, 1), - (BLOCK_V, PAGE_SIZE // 2), - ) - else: - even_packed = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + even_t[None, :].to(tl.int64) * kv_s2 - + v_packed_offsets[:, None] * kv_s4, - mask=(mask_v[:, None] & valid_even_t[None, :]) - if ASSUME_VALID_PAGES - else (valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), - other=0, - ) - odd_packed = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + odd_t[None, :].to(tl.int64) * kv_s2 - + v_packed_offsets[:, None] * kv_s4, - mask=(mask_v[:, None] & valid_odd_t[None, :]) - if ASSUME_VALID_PAGES - else (valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), - other=0, - ) - even_low = even_packed & 0x0F - even_high = (even_packed >> 4) & 0x0F - odd_low = odd_packed & 0x0F - even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) - odd_nibble = tl.where(v_use_high_nibble[:, None], odd_packed >> 4, odd_low) - v_vals = even_nibble | (odd_nibble << 4) - v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) - acc = tl.dot_scaled( - p_vals, +@triton.jit +def _fp4_mla_attention_pv_group_partial_prepacked_v_kernel( + partial_o_ptr, + p_fp4_ptr, + p_sf_ptr, + v_sf_ptr, + v_packed_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + num_pages, + po_s0: tl.constexpr, + po_s1: tl.constexpr, + po_s2: tl.constexpr, + po_s3: tl.constexpr, + po_num_rows: tl.constexpr, + p_s0: tl.constexpr, + p_s1: tl.constexpr, + p_num_rows: tl.constexpr, + vsf_s0: tl.constexpr, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + MAX_PAGES: tl.constexpr, + GROUP_PAGES: tl.constexpr, + NUM_DIM_BLOCKS: tl.constexpr, + ALLOW_PARTIAL_GROUPS: tl.constexpr, + PV_LOOP_STAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + combo = tl.program_id(2) + page_group = combo // NUM_DIM_BLOCKS + dim_block = combo - page_group * NUM_DIM_BLOCKS + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_view = tl.ext.make_view( + base=p_fp4_ptr, + shapes=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + tile_shape=[BLOCK_H, PAGE_SIZE // 2], + tile_dim_map=[0, 1], + ) + p_sf_view = tl.ext.make_view( + base=p_sf_ptr, + shapes=[p_num_rows // 128, SF_PER_PAGE // 4, 2, 256], + strides=[128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + v_sf_view = tl.ext.make_view( + base=v_sf_ptr, + shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], + strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, BLOCK_V // 128, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + v_packed_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + tl.assume(po_s2 % 8 == 0) + tl.assume(po_s3 == 1) + partial_o_desc = tl.make_tensor_descriptor( + partial_o_ptr, + shape=[po_num_rows, V_HEAD_D], + strides=[po_s2, po_s3], + block_shape=[BLOCK_H, BLOCK_V], + ) + + acc = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + for page_group_off in tl.range(0, GROUP_PAGES, num_stages=PV_LOOP_STAGES): + page_rel = page_group * GROUP_PAGES + page_group_off + valid_group_page = page_rel < MAX_PAGES + compact_page = page_table_start + page_rel + safe_compact_page = tl.where(valid_group_page, compact_page, page_table_start) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, + mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS), + other=0, + ).to(tl.int64) + + p_vals = tl.ext.load_view_tko( + p_view, + [(safe_compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + ) + p_vals = p_vals.to(tl.uint8, bitcast=True) + p_scales = tl.ext.load_view_tko(p_sf_view, [safe_compact_page.to(tl.int32), 0, 0, 0]) + p_scales = p_scales.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) + p_scales = p_scales.reshape([BLOCK_H, SF_PER_PAGE]) + + v_row = (physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + v_scales = tl.ext.load_view_tko( + v_sf_view, + [ + physical_page.to(tl.int32), + dim_block * (BLOCK_V // 128), + 0, + 0, + 0, + ], + ) + v_scales = v_scales.reshape([1, BLOCK_V // 128, SF_PER_PAGE // 4, 32, 4, 4]).trans( + 0, 1, 4, 3, 2, 5 + ) + v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) + if ALLOW_PARTIAL_GROUPS: + page_acc = tl.ext.dot_scaled( + v_vals, + v_scales, + "e2m1", + p_vals.T, p_scales, "e2m1", - v_vals.T, + acc=tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32), + fast_math=True, + rhs_k_pack=True, + ) + acc += tl.where(valid_group_page, page_acc, 0.0) + else: + acc = tl.ext.dot_scaled( + v_vals, v_scales, "e2m1", + p_vals.T, + p_scales, + "e2m1", acc=acc, fast_math=True, rhs_k_pack=True, ) - if ASSUME_FULL_HEADS and ASSUME_FULL_V: - out_vals = acc * out_scale - if USE_TMA_V_LOAD: - if out_ptr.dtype.element_ty == tl.bfloat16: - out_vals = out_vals.to(tl.bfloat16) - elif out_ptr.dtype.element_ty == tl.float16: - out_vals = out_vals.to(tl.float16) - out_desc.store( - [(gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], - out_vals, - ) - else: - tl.store( - out_ptr + gen_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, - out_vals, - ) - else: - tl.store( - out_ptr + gen_idx * out_s0 + safe_offs_h[:, None] * out_s1 + safe_offs_v[None, :] * out_s2, - acc * out_scale, - mask=mask_h[:, None] & mask_v[None, :], - ) + partial_o_desc.store( + [ + ( + page_group * (po_s0 // po_s2) + + gen_idx * (po_s1 // po_s2) + + head_block * BLOCK_H + ).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ], + acc.T, + ) @triton.jit -def _fp4_mla_attention_pv_prepacked_v_kernel( +def _fp4_mla_attention_pv_group_partial_reduce_kernel( out_ptr, - p_fp4_ptr, - p_sf_ptr, - kv_cache_ptr, - v_sf_ptr, - v_packed_ptr, + partial_o_ptr, + page_sum_ptr, global_scale_ptr, - src_page_ids_ptr, - paged_kv_indptr_decode_ptr, - kv_lens_ptr, - page_ids_len, - num_pages, out_s0, out_s1, out_s2, - out_num_rows, - p_s0, - p_s1, - p_num_rows, - kv_s0, - kv_s2, - kv_s4, - vsf_s0, + po_s0: tl.constexpr, + po_s1: tl.constexpr, + po_s2: tl.constexpr, + po_s3: tl.constexpr, + po_num_rows: tl.constexpr, + page_stats_s0: tl.constexpr, + page_stats_s1: tl.constexpr, NUM_HEADS: tl.constexpr, V_HEAD_D: tl.constexpr, - PAGE_SIZE: tl.constexpr, - FP4_BLOCK: tl.constexpr, - SF_PER_PAGE: tl.constexpr, - MAX_PAGES: tl.constexpr, + NUM_PAGE_GROUPS: tl.constexpr, P_GLOBAL_SCALE: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_V: tl.constexpr, - USE_TMA_P_LOAD: tl.constexpr, - USE_TMA_V_LOAD: tl.constexpr, - USE_PREPACKED_V: tl.constexpr, - PV_M_PACKED_V: tl.constexpr, - PV_LOOP_STAGES: tl.constexpr, - ASSUME_FULL_HEADS: tl.constexpr, - ASSUME_FULL_PAGES: tl.constexpr, - ASSUME_FULL_V: tl.constexpr, - ASSUME_VALID_PAGES: tl.constexpr, occupancy: tl.constexpr = 1, ): gen_idx = tl.program_id(0) head_block = tl.program_id(1) dim_block = tl.program_id(2) - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_offs_h = offs_h - else: - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - if ASSUME_FULL_V: - mask_v = tl.full([BLOCK_V], True, dtype=tl.int1) - safe_offs_v = offs_v - else: - mask_v = offs_v < V_HEAD_D - safe_offs_v = tl.where(mask_v, offs_v, 0) - packed_t = tl.arange(0, PAGE_SIZE // 2) - scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) - even_t = packed_t * 2 - odd_t = even_t + 1 - v_packed_offsets = safe_offs_v // 2 - v_use_high_nibble = (safe_offs_v & 1) != 0 - if ASSUME_FULL_V and BLOCK_V == 128: - v_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( - dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE - ) - else: - v_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE) - if USE_TMA_P_LOAD: - tl.assume(p_s0 % 8 == 0) - tl.assume(p_s1 == 1) - p_desc = tl.make_tensor_descriptor( - p_fp4_ptr, - shape=[p_num_rows, PAGE_SIZE // 2], - strides=[p_s0, p_s1], - block_shape=[BLOCK_H, PAGE_SIZE // 2], - ) - if USE_TMA_V_LOAD and ASSUME_FULL_HEADS and ASSUME_FULL_V: - tl.assume(out_s1 % 8 == 0) - tl.assume(out_s2 == 1) - out_desc = tl.make_tensor_descriptor( - out_ptr, - shape=[out_num_rows, V_HEAD_D], - strides=[out_s1, out_s2], - block_shape=[BLOCK_H, BLOCK_V], - ) - if USE_TMA_V_LOAD: - tl.assume(kv_s0 % 8 == 0) - tl.assume(kv_s2 % 8 == 0) - tl.assume(kv_s4 == 1) - v_desc = tl.make_tensor_descriptor( - kv_cache_ptr, - shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - block_shape=[1, PAGE_SIZE, BLOCK_V // 2], - ) - if USE_PREPACKED_V: - v_packed_desc = tl.make_tensor_descriptor( - v_packed_ptr, - shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], - strides=[PAGE_SIZE // 2, 1], - block_shape=[BLOCK_V, PAGE_SIZE // 2], - ) - if ( - USE_TMA_P_LOAD - and USE_TMA_V_LOAD - and ASSUME_FULL_HEADS - and ASSUME_FULL_PAGES - and ASSUME_FULL_V - and ASSUME_VALID_PAGES - and NUM_HEADS == 128 - and V_HEAD_D == 512 - and PAGE_SIZE == 128 - and BLOCK_H == 128 - and BLOCK_V == 128 - and SF_PER_PAGE == 8 - ): - p_view = tl.ext.make_view( - base=p_fp4_ptr, - shapes=[p_num_rows, PAGE_SIZE // 2], - strides=[p_s0, p_s1], - tile_shape=[BLOCK_H, PAGE_SIZE // 2], - tile_dim_map=[0, 1], - ) - p_sf_view = tl.ext.make_view( - base=p_sf_ptr, - shapes=[p_num_rows // 128, SF_PER_PAGE // 4, 2, 256], - strides=[128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, SF_PER_PAGE // 4, 2, 256], - tile_dim_map=[0, 1, 2, 3], - ) - v_sf_view = tl.ext.make_view( - base=v_sf_ptr, - shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], - strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, 1, SF_PER_PAGE // 4, 2, 256], - tile_dim_map=[0, 1, 2, 3, 4], - ) - v_view = tl.ext.make_view( - base=kv_cache_ptr, - shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], - tile_dim_map=[0, 1, 2], - ) - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) - global_scale = tl.load(global_scale_ptr) - out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) - acc = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) - for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): - compact_page = page_table_start + page_rel - physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) - - p_vals = tl.ext.load_view_tko( - p_view, - [(compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], - ) - p_vals = p_vals.to(tl.uint8, bitcast=True) - p_scales = tl.ext.load_view_tko(p_sf_view, [compact_page.to(tl.int32), 0, 0, 0]) - p_scales = p_scales.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) - p_scales = p_scales.reshape([BLOCK_H, SF_PER_PAGE]) - - if USE_PREPACKED_V: - v_row = (physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V - v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) - else: - v_tile = tl.ext.load_view_tko( - v_view, - [ - physical_page.to(tl.int32), - 0, - (dim_block * (BLOCK_V // 2)).to(tl.int32), - ], - ) - v_tile = v_tile.to(tl.uint8, bitcast=True) - v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) - if not PV_M_PACKED_V: - v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) - even_packed, odd_packed = tl.split(v_pairs) - low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) - v_vals = tl.reshape( - tl.join(low_vals, high_vals).permute(0, 2, 1), - (BLOCK_V, PAGE_SIZE // 2), - ) - v_scales = tl.ext.load_view_tko( - v_sf_view, - [ - physical_page.to(tl.int32), - dim_block, - 0, - 0, - 0, - ], - ) - v_scales = v_scales.reshape([1, 1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) - v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) - if PV_M_PACKED_V and not USE_PREPACKED_V: - acc = tl.ext.dot_scaled( - v_tile.T, - v_scales, - "e2m1", - p_vals.T, - p_scales, - "e2m1", - acc=acc, - fast_math=True, - lhs_k_pack=False, - rhs_k_pack=True, - ) - else: - acc = tl.ext.dot_scaled( - v_vals, - v_scales, - "e2m1", - p_vals.T, - p_scales, - "e2m1", - acc=acc, - fast_math=True, - rhs_k_pack=True, - ) + tl.assume(po_s2 % 8 == 0) + tl.assume(po_s3 == 1) + partial_o_desc = tl.make_tensor_descriptor( + partial_o_ptr, + shape=[po_num_rows, V_HEAD_D], + strides=[po_s2, po_s3], + block_shape=[BLOCK_H, BLOCK_V], + ) - out_vals = acc.T * out_scale - if out_ptr.dtype.element_ty == tl.bfloat16: - out_vals = out_vals.to(tl.bfloat16) - elif out_ptr.dtype.element_ty == tl.float16: - out_vals = out_vals.to(tl.float16) - out_desc.store( - [(gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], - out_vals, - ) - return + global_m = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + for group_idx in tl.range(0, NUM_PAGE_GROUPS): + group_m = tl.load(page_sum_ptr + gen_idx * page_stats_s0 + (group_idx * 2) * page_stats_s1 + offs_h) + global_m = tl.maximum(global_m, group_m) + + global_l = tl.zeros((BLOCK_H,), dtype=tl.float32) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + for group_idx in tl.range(0, NUM_PAGE_GROUPS): + group_m = tl.load(page_sum_ptr + gen_idx * page_stats_s0 + (group_idx * 2) * page_stats_s1 + offs_h) + group_l = tl.load(page_sum_ptr + gen_idx * page_stats_s0 + (group_idx * 2 + 1) * page_stats_s1 + offs_h) + scale = tl.where(group_l > 0.0, tl.math.exp2((group_m - global_m) * 1.4426950408889634), 0.0) + partial_o = partial_o_desc.load( + [ + ( + group_idx * (po_s0 // po_s2) + + gen_idx * (po_s1 // po_s2) + + head_block * BLOCK_H + ).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ] + ) + acc += partial_o * scale[:, None] + global_l += group_l * scale - if ASSUME_FULL_PAGES: - kv_len = 0 - else: - kv_len = tl.load(kv_lens_ptr + gen_idx) - page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) global_scale = tl.load(global_scale_ptr) out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) - acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) - for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): - page_start = page_rel * PAGE_SIZE - if ASSUME_FULL_PAGES or page_start < kv_len: - compact_page = page_table_start + page_rel - if ASSUME_VALID_PAGES: - safe_compact_page = compact_page - physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) - safe_physical_page = physical_page - else: - valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) - safe_compact_page = tl.where(valid_compact_page, compact_page, 0) - physical_page = tl.load(src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1).to( - tl.int64 - ) - valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) - safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + safe_l = tl.where(global_l > 0.0, global_l, 1.0) + out_vals = (acc / safe_l[:, None]) * out_scale + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + tl.store( + out_ptr + gen_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + ) - p_rows = safe_compact_page * NUM_HEADS + offs_h - safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, safe_compact_page * NUM_HEADS) - if USE_TMA_P_LOAD: - p_vals = p_desc.load([(safe_compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) - else: - p_vals = tl.load( - p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, - mask=mask_h[:, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None], - other=0, - ) - p_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) - p_scales = tl.load(p_sf_ptr + p_sf_offsets) - if ASSUME_FULL_PAGES: - valid_even_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) - valid_odd_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) - else: - valid_even_t = page_start + even_t < kv_len - valid_odd_t = page_start + odd_t < kv_len - if USE_TMA_V_LOAD: - v_tile = v_desc.load( - [ - safe_physical_page.to(tl.int32), - 0, - (dim_block * (BLOCK_V // 2)).to(tl.int32), - ] - ) - v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) - if not ASSUME_VALID_PAGES: - v_tile = tl.where(valid_physical_page, v_tile, 0) - v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) - even_packed, odd_packed = tl.split(v_pairs) - if not ASSUME_FULL_PAGES: - even_packed = tl.where(valid_even_t[None, :], even_packed, 0) - odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) - low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) - v_vals = tl.reshape( - tl.join(low_vals, high_vals).permute(0, 2, 1), - (BLOCK_V, PAGE_SIZE // 2), - ) - else: - even_packed = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + even_t[None, :].to(tl.int64) * kv_s2 - + v_packed_offsets[:, None] * kv_s4, - mask=(mask_v[:, None] & valid_even_t[None, :]) - if ASSUME_VALID_PAGES - else (valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), - other=0, - ) - odd_packed = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + odd_t[None, :].to(tl.int64) * kv_s2 - + v_packed_offsets[:, None] * kv_s4, - mask=(mask_v[:, None] & valid_odd_t[None, :]) - if ASSUME_VALID_PAGES - else (valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), - other=0, - ) - even_low = even_packed & 0x0F - even_high = (even_packed >> 4) & 0x0F - odd_low = odd_packed & 0x0F - even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) - odd_nibble = tl.where(v_use_high_nibble[:, None], odd_packed >> 4, odd_low) - v_vals = even_nibble | (odd_nibble << 4) - v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) - acc = tl.dot_scaled( - p_vals, +@triton.jit +def _fp4_mla_attention_pv_atomic_split_prepacked_v_kernel( + out_acc_ptr, + p_fp4_ptr, + p_sf_ptr, + v_sf_ptr, + v_packed_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + num_pages, + out_s0, + out_s1, + out_s2, + p_s0: tl.constexpr, + p_s1: tl.constexpr, + p_num_rows: tl.constexpr, + vsf_s0: tl.constexpr, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + MAX_PAGES: tl.constexpr, + GROUP_PAGES: tl.constexpr, + NUM_DIM_BLOCKS: tl.constexpr, + ALLOW_PARTIAL_GROUPS: tl.constexpr, + PV_LOOP_STAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + combo = tl.program_id(2) + page_group = combo // NUM_DIM_BLOCKS + dim_block = combo - page_group * NUM_DIM_BLOCKS + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_view = tl.ext.make_view( + base=p_fp4_ptr, + shapes=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + tile_shape=[BLOCK_H, PAGE_SIZE // 2], + tile_dim_map=[0, 1], + ) + p_sf_view = tl.ext.make_view( + base=p_sf_ptr, + shapes=[p_num_rows // 128, SF_PER_PAGE // 4, 2, 256], + strides=[128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3], + ) + v_sf_view = tl.ext.make_view( + base=v_sf_ptr, + shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], + strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], + tile_shape=[1, BLOCK_V // 128, SF_PER_PAGE // 4, 2, 256], + tile_dim_map=[0, 1, 2, 3, 4], + ) + v_packed_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + + acc = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) + for page_group_off in tl.range(0, GROUP_PAGES, num_stages=PV_LOOP_STAGES): + page_rel = page_group * GROUP_PAGES + page_group_off + valid_group_page = page_rel < MAX_PAGES + compact_page = page_table_start + page_rel + safe_compact_page = tl.where(valid_group_page, compact_page, page_table_start) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, + mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS), + other=0, + ).to(tl.int64) + + p_vals = tl.ext.load_view_tko( + p_view, + [(safe_compact_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], + ) + p_vals = p_vals.to(tl.uint8, bitcast=True) + p_scales = tl.ext.load_view_tko(p_sf_view, [safe_compact_page.to(tl.int32), 0, 0, 0]) + p_scales = p_scales.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) + p_scales = p_scales.reshape([BLOCK_H, SF_PER_PAGE]) + + v_row = (physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + v_scales = tl.ext.load_view_tko( + v_sf_view, + [ + physical_page.to(tl.int32), + dim_block * (BLOCK_V // 128), + 0, + 0, + 0, + ], + ) + v_scales = v_scales.reshape([1, BLOCK_V // 128, SF_PER_PAGE // 4, 32, 4, 4]).trans( + 0, 1, 4, 3, 2, 5 + ) + v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) + if ALLOW_PARTIAL_GROUPS: + page_acc = tl.ext.dot_scaled( + v_vals, + v_scales, + "e2m1", + p_vals.T, p_scales, "e2m1", - v_vals.T, + acc=tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32), + fast_math=True, + rhs_k_pack=True, + ) + acc += tl.where(valid_group_page, page_acc, 0.0) + else: + acc = tl.ext.dot_scaled( + v_vals, v_scales, "e2m1", + p_vals.T, + p_scales, + "e2m1", acc=acc, fast_math=True, rhs_k_pack=True, ) - if ASSUME_FULL_HEADS and ASSUME_FULL_V: - out_vals = acc * out_scale - if USE_TMA_V_LOAD: - if out_ptr.dtype.element_ty == tl.bfloat16: - out_vals = out_vals.to(tl.bfloat16) - elif out_ptr.dtype.element_ty == tl.float16: - out_vals = out_vals.to(tl.float16) - out_desc.store( - [(gen_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], - out_vals, - ) - else: - tl.store( - out_ptr + gen_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, - out_vals, - ) - else: - tl.store( - out_ptr + gen_idx * out_s0 + safe_offs_h[:, None] * out_s1 + safe_offs_v[None, :] * out_s2, - acc * out_scale, - mask=mask_h[:, None] & mask_v[None, :], - ) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + out_vals = acc.T * out_scale + tl.atomic_add( + out_acc_ptr + gen_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + sem="relaxed", + ) + + +@triton.jit +def _fp4_mla_attention_cast_acc_kernel( + out_ptr, + out_acc_ptr, + out_s0, + out_s1, + out_s2, + acc_s0, + acc_s1, + acc_s2, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + occupancy: tl.constexpr = 1, +): + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + dim_block = tl.program_id(2) + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + mask_h = offs_h < NUM_HEADS + mask_v = offs_v < V_HEAD_D + safe_h = tl.where(mask_h, offs_h, 0) + safe_v = tl.where(mask_v, offs_v, 0) + + vals = tl.load( + out_acc_ptr + gen_idx * acc_s0 + safe_h[:, None] * acc_s1 + safe_v[None, :] * acc_s2, + mask=mask_h[:, None] & mask_v[None, :], + other=0.0, + ) + if out_ptr.dtype.element_ty == tl.bfloat16: + vals = vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + vals = vals.to(tl.float16) + tl.store( + out_ptr + gen_idx * out_s0 + safe_h[:, None] * out_s1 + safe_v[None, :] * out_s2, + vals, + mask=mask_h[:, None] & mask_v[None, :], + ) def fp4_mla_paged_attention_internal( @@ -2620,7 +5267,9 @@ def fp4_mla_paged_attention_internal( pack_prob_in_page_stats: Optional[bool] = None, page_stats_group_size: Optional[int] = None, assume_full_pages: Optional[bool] = None, + assume_full_pages_except_mtp_tail: bool = False, assume_valid_pages: Optional[bool] = None, + query_len_per_seq: int = 1, prepack_v_for_pv: bool = False, use_prepacked_v_for_pv: bool = False, p_fp4_workspace: Optional[torch.Tensor] = None, @@ -2660,6 +5309,15 @@ def fp4_mla_paged_attention_internal( else: raise ValueError("q_fp4 must be 2D or 3D.") + if query_len_per_seq <= 0: + raise ValueError(f"query_len_per_seq must be positive, got {query_len_per_seq}.") + if num_gen % query_len_per_seq != 0: + raise ValueError( + "q_fp4 query rows must be divisible by query_len_per_seq, got " + f"{num_gen} rows and query_len_per_seq={query_len_per_seq}." + ) + num_gen_seqs = num_gen // query_len_per_seq + num_pages, inferred_page_size, packed_k_dim, kv_s0, kv_s2, kv_s4 = _get_kv_cache_strides(kv_cache) if page_size is None: page_size = inferred_page_size @@ -2703,13 +5361,56 @@ def fp4_mla_paged_attention_internal( env_block_h = _env_int("TRTLLM_FP4_MLA_BLOCK_H") env_block_k = _env_int("TRTLLM_FP4_MLA_BLOCK_K") env_block_v = _env_int("TRTLLM_FP4_MLA_BLOCK_V") + env_pv_block_h = _env_int("TRTLLM_FP4_MLA_PV_BLOCK_H") env_pv_loop_stages = _env_int("TRTLLM_FP4_MLA_PV_LOOP_STAGES") env_occupancy = _env_int("TRTLLM_FP4_MLA_OCCUPANCY") + env_num_ctas = _env_int("TRTLLM_FP4_MLA_NUM_CTAS") env_num_warps = _env_int("TRTLLM_FP4_MLA_NUM_WARPS") env_num_stages = _env_int("TRTLLM_FP4_MLA_NUM_STAGES") + env_page_stats_num_ctas = _env_int("TRTLLM_FP4_MLA_PAGE_STATS_NUM_CTAS") + env_page_stats_num_warps = _env_int("TRTLLM_FP4_MLA_PAGE_STATS_NUM_WARPS") + env_page_stats_num_stages = _env_int("TRTLLM_FP4_MLA_PAGE_STATS_NUM_STAGES") + env_pv_num_ctas = _env_int("TRTLLM_FP4_MLA_PV_NUM_CTAS") + env_pv_num_warps = _env_int("TRTLLM_FP4_MLA_PV_NUM_WARPS") + env_pv_num_stages = _env_int("TRTLLM_FP4_MLA_PV_NUM_STAGES") env_group_pages = _env_int("TRTLLM_FP4_MLA_GROUP_PAGES") env_group_reduce_stats = _env_int("TRTLLM_FP4_MLA_GROUP_REDUCE_STATS") env_page_pipeline_streams = _env_int("TRTLLM_FP4_MLA_PAGE_PIPELINE_STREAMS") + env_online_qkpv = _env_int("TRTLLM_FP4_MLA_ONLINE_QKPV") + env_online_qkpv_group_pages = _env_int("TRTLLM_FP4_MLA_ONLINE_QKPV_GROUP_PAGES") + env_online_qkpv_max_batch = _env_int("TRTLLM_FP4_MLA_ONLINE_QKPV_MAX_BATCH") + env_online_qkpv_fp4_pv = _env_int("TRTLLM_FP4_MLA_ONLINE_QKPV_FP4_PV") + env_gen_qkpv = _env_int("TRTLLM_FP4_MLA_GEN_QKPV") + env_gen_qkpv_group_pages = _env_int("TRTLLM_FP4_MLA_GEN_QKPV_GROUP_PAGES") + env_gen_qkpv_block_h = _env_int("TRTLLM_FP4_MLA_GEN_QKPV_BLOCK_H") + env_gen_qkpv_partial_dtype = os.environ.get("TRTLLM_FP4_MLA_GEN_QKPV_PARTIAL_DTYPE", "").lower() + env_mtp_fused_qkpv = _env_int("TRTLLM_FP4_MLA_MTP_FUSED_QKPV") + env_mtp_fused_qkpv_group_pages = _env_int("TRTLLM_FP4_MLA_MTP_FUSED_QKPV_GROUP_PAGES") + env_mtp_fused_qkpv_block_h = _env_int("TRTLLM_FP4_MLA_MTP_FUSED_QKPV_BLOCK_H") + env_mtp_fused_qkpv_block_v = _env_int("TRTLLM_FP4_MLA_MTP_FUSED_QKPV_BLOCK_V") + env_mtp_fused_qkpv_split_pv_k = _env_int("TRTLLM_FP4_MLA_MTP_FUSED_QKPV_SPLIT_PV_K") + env_mtp_fused_qkpv_occupancy = _env_int("TRTLLM_FP4_MLA_MTP_FUSED_QKPV_OCCUPANCY") + env_mtp_page_stats_pair = _env_int("TRTLLM_FP4_MLA_MTP_PAGE_STATS_PAIR") + env_mtp_pv_pair = _env_int("TRTLLM_FP4_MLA_MTP_PV_PAIR") + env_mtp_split_tail_group = _env_int("TRTLLM_FP4_MLA_MTP_SPLIT_TAIL_GROUP") + env_mtp_final_page_fast_path = _env_int("TRTLLM_FP4_MLA_MTP_FINAL_PAGE_FAST_PATH") + env_group_pv = _env_int("TRTLLM_FP4_MLA_GROUP_PV") + env_group_pv_max_batch = _env_int("TRTLLM_FP4_MLA_GROUP_PV_MAX_BATCH") + env_pv_atomic_split = _env_int("TRTLLM_FP4_MLA_PV_ATOMIC_SPLIT") + env_pv_atomic_group_pages = _env_int("TRTLLM_FP4_MLA_PV_ATOMIC_GROUP_PAGES") + env_pv_apply_prob_scale = _env_int("TRTLLM_FP4_MLA_PV_APPLY_PROB_SCALE") + env_pv_scale_in_sf = _env_int("TRTLLM_FP4_MLA_PV_SCALE_IN_SF") + env_scale_from_group_stats = _env_int("TRTLLM_FP4_MLA_SCALE_FROM_GROUP_STATS") + env_duplicate_tail_k = _env_int("TRTLLM_FP4_MLA_DUPLICATE_TAIL_K") + env_debug_page_stats_pack = _env_int("TRTLLM_FP4_MLA_DEBUG_PAGE_STATS_PACK") + env_debug_stop_after_page_stats = _env_int("TRTLLM_FP4_MLA_DEBUG_STOP_AFTER_PAGE_STATS") + env_mtp_page_stats_pair_group_pages = _env_int( + "TRTLLM_FP4_MLA_MTP_PAGE_STATS_PAIR_GROUP_PAGES" + ) + env_pv_tma = _env_int("TRTLLM_FP4_MLA_PV_TMA") + env_pv_p_tma = _env_int("TRTLLM_FP4_MLA_PV_P_TMA") + env_pv_v_tma = _env_int("TRTLLM_FP4_MLA_PV_V_TMA") + env_pv_out_tma = _env_int("TRTLLM_FP4_MLA_PV_OUT_TMA") if env_block_h is not None: block_h = env_block_h if block_k is None: @@ -2724,11 +5425,14 @@ def fp4_mla_paged_attention_internal( tail_k = q_head_dim - full_block_end tail_block_k = 1 << (tail_k - 1).bit_length() if tail_k > 0 else block_k if max_pages is None: - if paged_kv_indptr_decode.numel() >= num_gen + 1: - page_counts = paged_kv_indptr_decode[1 : num_gen + 1] - paged_kv_indptr_decode[:num_gen] + if paged_kv_indptr_decode.numel() >= num_gen_seqs + 1: + page_counts = ( + paged_kv_indptr_decode[1 : num_gen_seqs + 1] + - paged_kv_indptr_decode[:num_gen_seqs] + ) max_pages = int(page_counts.max().item()) if page_counts.numel() > 0 else 0 else: - max_pages = _ceil_div(int(kv_lens[:num_gen].max().item()), page_size) + max_pages = _ceil_div(int(kv_lens[:num_gen_seqs].max().item()), page_size) if max_pages <= 0: output.zero_() return output @@ -2738,23 +5442,38 @@ def fp4_mla_paged_attention_internal( sf_per_page = page_size // FP4_BLOCK_SIZE num_head_blocks = triton.cdiv(num_heads, block_h) assume_full_heads = num_heads % block_h == 0 + pv_block_h = env_pv_block_h if env_pv_block_h is not None else block_h + if pv_block_h <= 0: + raise ValueError(f"TRTLLM_FP4_MLA_PV_BLOCK_H must be positive, got {pv_block_h}.") + if pv_block_h not in (64, 128): + raise ValueError(f"TRTLLM_FP4_MLA_PV_BLOCK_H currently supports 64 or 128, got {pv_block_h}.") + num_pv_head_blocks = triton.cdiv(num_heads, pv_block_h) + assume_full_pv_heads = num_heads % pv_block_h == 0 assume_full_v = v_head_dim % block_v == 0 if assume_full_pages is None: assume_full_pages = False - assume_full_pages = bool(assume_full_pages) + assume_full_pages = bool(assume_full_pages) and query_len_per_seq == 1 + mask_mtp_final_page_only = ( + bool(assume_full_pages_except_mtp_tail) + and (env_mtp_final_page_fast_path != 0 if env_mtp_final_page_fast_path is not None else True) + and query_len_per_seq > 1 + and query_len_per_seq <= page_size + ) if assume_valid_pages is None: assume_valid_pages = False assume_valid_pages = bool(assume_valid_pages) if ( not assume_valid_pages and assume_full_pages - and src_page_ids.numel() == num_gen * max_pages + and src_page_ids.numel() == num_gen_seqs * max_pages ): # Full decode pages with an exactly-sized page table do not need the # sentinel/physical-page validity masks. Keeping this inference inside # the kernel wrapper lets the framework call path stay unchanged. assume_valid_pages = True - total_p_rows = max(src_page_ids.numel() * num_heads, 1) + p_by_query = query_len_per_seq != 1 + total_p_pages = num_gen * max_pages if p_by_query else src_page_ids.numel() + total_p_rows = max(total_p_pages * num_heads, 1) if page_pipeline_streams is None and env_page_pipeline_streams is not None: page_pipeline_streams = env_page_pipeline_streams if page_pipeline_streams is None: @@ -2764,6 +5483,7 @@ def fp4_mla_paged_attention_internal( page_pipeline_streams = 1 page_pipeline_streams = max(1, min(int(page_pipeline_streams), max_pages)) launch_meta = {} + explicit_kernel_occupancy = kernel_occupancy is not None or env_occupancy is not None if kernel_occupancy is None: if env_occupancy is not None: kernel_occupancy = env_occupancy @@ -2773,6 +5493,8 @@ def fp4_mla_paged_attention_internal( launch_meta["occupancy"] = int(kernel_occupancy) if kernel_num_ctas is not None: launch_meta["num_ctas"] = int(kernel_num_ctas) + elif env_num_ctas is not None: + launch_meta["num_ctas"] = int(env_num_ctas) if kernel_num_stages is None: if env_num_stages is not None: kernel_num_stages = env_num_stages @@ -2784,6 +5506,25 @@ def fp4_mla_paged_attention_internal( kernel_num_warps = env_num_warps if kernel_num_warps is not None: launch_meta["num_warps"] = int(kernel_num_warps) + page_stats_launch_meta = dict(launch_meta) + if env_page_stats_num_ctas is not None: + page_stats_launch_meta["num_ctas"] = int(env_page_stats_num_ctas) + if env_page_stats_num_stages is not None: + page_stats_launch_meta["num_stages"] = int(env_page_stats_num_stages) + if env_page_stats_num_warps is not None: + page_stats_launch_meta["num_warps"] = int(env_page_stats_num_warps) + elif kernel_num_warps is None and triton_backend == "nvt": + page_stats_launch_meta["num_warps"] = 8 + pv_launch_meta = dict(launch_meta) + if env_pv_num_ctas is not None: + pv_launch_meta["num_ctas"] = int(env_pv_num_ctas) + if env_pv_num_stages is not None: + pv_launch_meta["num_stages"] = int(env_pv_num_stages) + if env_pv_num_warps is not None: + pv_launch_meta["num_warps"] = int(env_pv_num_warps) + gen_qkpv_launch_meta = dict(page_stats_launch_meta) + if not explicit_kernel_occupancy: + gen_qkpv_launch_meta.pop("occupancy", None) if fused_prob_pack is None: fused_prob_pack = triton_backend == "nvt" if fused_prob_pack_single_launch is None: @@ -2797,133 +5538,618 @@ def alloc_fn(size: int, alignment: int, stream: Optional[int]): return torch.empty(size, device=q_fp4.device, dtype=torch.int8) triton.set_allocator(alloc_fn) + use_pv_tma_data_load = use_tma_data_load and (env_pv_tma is None or env_pv_tma != 0) + use_pv_p_tma_data_load = use_pv_tma_data_load and (env_pv_p_tma is None or env_pv_p_tma != 0) + use_pv_v_tma_data_load = use_pv_tma_data_load and (env_pv_v_tma is None or env_pv_v_tma != 0) + use_pv_out_tma_data_store = use_pv_tma_data_load and ( + env_pv_out_tma is None or env_pv_out_tma != 0 + ) p_fp4 = _workspace_tensor( p_fp4_workspace, (total_p_rows, page_size // 2), dtype=torch.uint8, device=q_fp4.device, - name="p_fp4", + name="p_fp4", + ) + p_sf = _workspace_tensor( + p_sf_workspace, + (_swizzled_scale_size(total_p_rows, page_size),), + dtype=q_sf.dtype, + device=q_fp4.device, + name="p_sf", + ) + num_dim_blocks = triton.cdiv(v_head_dim, block_v) + v_repack_block_v = block_v + num_repack_dim_blocks = num_dim_blocks + env_prepack_v = os.environ.get("TRTLLM_FP4_MLA_PREPACK_V") + can_address_all_compact_pages = src_page_ids.numel() == num_gen_seqs * max_pages + allow_partial_page_prepack_v_for_pv = ( + p_by_query + and can_address_all_compact_pages + and (bool(prepack_v_for_pv) or bool(use_prepacked_v_for_pv)) + ) + auto_prepack_v_for_pv = ( + triton_backend == "nvt" + and use_tma_data_load + and assume_full_heads + and (assume_full_pages or allow_partial_page_prepack_v_for_pv) + and (assume_valid_pages or allow_partial_page_prepack_v_for_pv) + and v_head_dim == 512 + and page_size == 128 + and block_h == 128 + and block_v in (128, 256) + and sf_per_page == 8 + ) + if not prepack_v_for_pv and not use_prepacked_v_for_pv: + prepack_v_for_pv = ( + auto_prepack_v_for_pv + if env_prepack_v is None + else env_prepack_v == "1" and auto_prepack_v_for_pv + ) + wants_prepacked_v_for_pv = bool(prepack_v_for_pv) or bool(use_prepacked_v_for_pv) + if use_prepacked_v_for_pv and v_packed_workspace is None: + raise ValueError("use_prepacked_v_for_pv requires v_packed_workspace to be provided.") + can_use_prepacked_v_for_pv = ( + wants_prepacked_v_for_pv + and use_tma_data_load + and assume_full_heads + and (assume_full_pages or allow_partial_page_prepack_v_for_pv) + and (assume_valid_pages or allow_partial_page_prepack_v_for_pv) + and v_head_dim == 512 + and page_size == 128 + and block_h in (64, 128) + and block_v in (128, 256) + and sf_per_page == 8 + ) + if wants_prepacked_v_for_pv and not can_use_prepacked_v_for_pv: + raise ValueError( + "prepacked V PV path requires TMA, full valid pages or explicit qlen>1 prepack, " + "v_head_dim=512, page_size=128, block_h in (64, 128), block_v in (128, 256), and sf_per_page=8." + ) + if can_use_prepacked_v_for_pv: + v_packed = _workspace_tensor( + v_packed_workspace, + (num_pages * num_dim_blocks * block_v, page_size // 2), + dtype=torch.uint8, + device=q_fp4.device, + name="v_packed", + ) + else: + v_packed = kv_cache + if fused_prob_pack: + p_probs = None + else: + p_probs_shape = (max(num_gen * num_heads, 1), page_size) + if page_pipeline_streams > 1: + p_probs = _workspace_tensor( + p_probs_workspace, + (page_pipeline_streams, *p_probs_shape), + dtype=torch.float32, + device=q_fp4.device, + name="p_probs", + ) + else: + p_probs = _workspace_tensor( + p_probs_workspace, + p_probs_shape, + dtype=torch.float32, + device=q_fp4.device, + name="p_probs", + ) + max_scores = _workspace_tensor( + max_scores_workspace, + (num_gen, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="max_scores", ) - p_sf = _workspace_tensor( - p_sf_workspace, - (_swizzled_scale_size(total_p_rows, page_size),), - dtype=q_sf.dtype, + denom = _workspace_tensor( + denom_workspace, + (num_gen, num_heads), + dtype=torch.float32, device=q_fp4.device, - name="p_sf", + name="denom", ) - num_dim_blocks = triton.cdiv(v_head_dim, block_v) - v_repack_block_v = block_v - num_repack_dim_blocks = num_dim_blocks - auto_prepack_v_for_pv = ( - triton_backend == "nvt" + v_repack_stream = None + debug_timing = os.environ.get("TRTLLM_FP4_MLA_DEBUG_TIMING") == "1" + debug_events = [] + + def _debug_mark(label: str) -> None: + if not debug_timing: + return + event = torch.cuda.Event(enable_timing=True) + event.record() + debug_events.append((label, event)) + + def _debug_report() -> None: + if not debug_timing or len(debug_events) < 2: + return + torch.cuda.synchronize(q_fp4.device) + parts = [] + for (start_label, start_event), (end_label, end_event) in zip( + debug_events, debug_events[1:] + ): + parts.append(f"{start_label}->{end_label}={start_event.elapsed_time(end_event):.3f}ms") + print("[fp4_mla_timing] " + " ".join(parts), flush=True) + + _debug_mark("start") + if bool(prepack_v_for_pv) and can_use_prepacked_v_for_pv: + current_stream = torch.cuda.current_stream(q_fp4.device) + v_repack_stream = torch.cuda.Stream(device=q_fp4.device) + v_repack_stream.wait_stream(current_stream) + with torch.cuda.stream(v_repack_stream): + _fp4_mla_attention_v_repack_kernel[(num_pages, num_repack_dim_blocks)]( + v_packed, + kv_cache, + num_pages, + kv_s0, + kv_s2, + kv_s4, + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + BLOCK_V=v_repack_block_v, + **launch_meta, + ) + + mtp_fused_block_h = env_mtp_fused_qkpv_block_h if env_mtp_fused_qkpv_block_h is not None else 32 + if env_mtp_fused_qkpv == 1 and mtp_fused_block_h not in (16, 32, 64, 128): + raise ValueError( + "TRTLLM_FP4_MLA_MTP_FUSED_QKPV_BLOCK_H currently supports 16, 32, 64, or 128, " + f"got {mtp_fused_block_h}." + ) + mtp_fused_block_v = env_mtp_fused_qkpv_block_v if env_mtp_fused_qkpv_block_v is not None else v_head_dim + if env_mtp_fused_qkpv == 1 and mtp_fused_block_v not in (128, 256, 512): + raise ValueError( + "TRTLLM_FP4_MLA_MTP_FUSED_QKPV_BLOCK_V currently supports 128, 256, or 512, " + f"got {mtp_fused_block_v}." + ) + if env_mtp_fused_qkpv == 1 and v_head_dim % mtp_fused_block_v != 0: + raise ValueError( + f"TRTLLM_FP4_MLA_MTP_FUSED_QKPV_BLOCK_V={mtp_fused_block_v} must divide v_head_dim={v_head_dim}." + ) + num_mtp_fused_dim_blocks = triton.cdiv(v_head_dim, mtp_fused_block_v) + mtp_fused_group_pages = env_mtp_fused_qkpv_group_pages if env_mtp_fused_qkpv_group_pages is not None else 128 + mtp_fused_group_pages = max(1, min(int(mtp_fused_group_pages), max_pages)) + mtp_fused_launch_meta = dict(page_stats_launch_meta) + mtp_fused_reduce_meta = dict(pv_launch_meta) + if env_mtp_fused_qkpv_occupancy is not None: + mtp_fused_launch_meta["occupancy"] = int(env_mtp_fused_qkpv_occupancy) + mtp_fused_reduce_meta["occupancy"] = int(env_mtp_fused_qkpv_occupancy) + elif env_occupancy is None: + mtp_fused_launch_meta["occupancy"] = 1 + mtp_fused_reduce_meta["occupancy"] = 1 + can_use_mtp_fused_qkpv = ( + env_mtp_fused_qkpv == 1 + and p_by_query + and query_len_per_seq == 4 + and can_use_prepacked_v_for_pv + and can_address_all_compact_pages and use_tma_data_load - and assume_full_heads - and assume_full_pages - and assume_valid_pages + and num_heads % mtp_fused_block_h == 0 + and num_heads == 128 + and q_head_dim == 640 + and k_head_dim == 576 + and q_residual_dim == 64 and v_head_dim == 512 and page_size == 128 - and block_h in (64, 128) + and block_k == 512 + and full_block_end == 512 + and tail_block_k == 128 + and block_v in (128, 256) + and sf_per_page == 8 + ) + if can_use_mtp_fused_qkpv: + if v_repack_stream is not None: + torch.cuda.current_stream(q_fp4.device).wait_stream(v_repack_stream) + v_repack_stream = None + num_mtp_fused_page_groups = _ceil_div(max_pages, mtp_fused_group_pages) + num_mtp_fused_head_blocks = triton.cdiv(num_heads, mtp_fused_block_h) + mtp_fused_partial_o = _workspace_tensor( + None, + (num_mtp_fused_page_groups, num_gen, num_heads, v_head_dim), + dtype=torch.float32, + device=q_fp4.device, + name="mtp_fused_partial_o", + ) + mtp_fused_partial_m = _workspace_tensor( + None, + (num_mtp_fused_page_groups, num_gen, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="mtp_fused_partial_m", + ) + mtp_fused_partial_l = _workspace_tensor( + None, + (num_mtp_fused_page_groups, num_gen, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="mtp_fused_partial_l", + ) + _fp4_mla_attention_mtp_fused_qkpv_group_kernel[ + (num_gen, num_mtp_fused_head_blocks, num_mtp_fused_page_groups * num_mtp_fused_dim_blocks) + ]( + mtp_fused_partial_o, + mtp_fused_partial_m, + mtp_fused_partial_l, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + v_sf, + v_packed, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + v_sf.stride(0), + mtp_fused_partial_o.stride(0), + mtp_fused_partial_o.stride(1), + mtp_fused_partial_o.stride(2), + mtp_fused_partial_o.stride(3), + mtp_fused_partial_m.stride(0), + mtp_fused_partial_m.stride(1), + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + BLOCK_H=mtp_fused_block_h, + BLOCK_T=page_size, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + BLOCK_V=mtp_fused_block_v, + MAX_PAGES=max_pages, + GROUP_PAGES=mtp_fused_group_pages, + NUM_DIM_BLOCKS=num_mtp_fused_dim_blocks, + ALLOW_PARTIAL_GROUPS=max_pages % mtp_fused_group_pages != 0, + SPLIT_PV_K=env_mtp_fused_qkpv_split_pv_k == 1, + USE_TMA_DATA_LOAD=use_tma_data_load, + ASSUME_FULL_HEADS=True, + ASSUME_VALID_PAGES=True, + **mtp_fused_launch_meta, + ) + _fp4_mla_attention_online_qkpv_reduce_kernel[ + (num_gen, num_mtp_fused_head_blocks, num_mtp_fused_dim_blocks) + ]( + output, + mtp_fused_partial_o, + mtp_fused_partial_m, + mtp_fused_partial_l, + global_scale, + output.stride(0), + output.stride(1), + output.stride(2), + mtp_fused_partial_o.stride(0), + mtp_fused_partial_o.stride(1), + mtp_fused_partial_o.stride(2), + mtp_fused_partial_o.stride(3), + mtp_fused_partial_m.stride(0), + mtp_fused_partial_m.stride(1), + NUM_HEADS=num_heads, + V_HEAD_D=v_head_dim, + NUM_PAGE_GROUPS=num_mtp_fused_page_groups, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=mtp_fused_block_h, + BLOCK_V=mtp_fused_block_v, + FP4_PV=True, + **mtp_fused_reduce_meta, + ) + return output + + gen_qkpv_block_h = env_gen_qkpv_block_h if env_gen_qkpv_block_h is not None else 32 + gen_qkpv_group_pages = env_gen_qkpv_group_pages if env_gen_qkpv_group_pages is not None else 128 + gen_qkpv_group_pages = max(1, min(int(gen_qkpv_group_pages), max_pages)) + can_use_gen_qkpv = ( + env_gen_qkpv == 1 + and p_by_query + and query_len_per_seq == 4 + and can_use_prepacked_v_for_pv + and can_address_all_compact_pages + and use_tma_data_load + and num_heads % gen_qkpv_block_h == 0 + and gen_qkpv_block_h in (16, 32, 64) + and num_heads == 128 + and q_head_dim == 640 + and k_head_dim == 576 + and q_residual_dim == 64 + and v_head_dim == 512 + and page_size == 128 + and block_k == 512 + and full_block_end == 512 + and tail_block_k == 128 and block_v == 128 and sf_per_page == 8 ) - if not prepack_v_for_pv and not use_prepacked_v_for_pv: - env_prepack_v = os.environ.get("TRTLLM_FP4_MLA_PREPACK_V") - prepack_v_for_pv = auto_prepack_v_for_pv if env_prepack_v is None else env_prepack_v == "1" - wants_prepacked_v_for_pv = bool(prepack_v_for_pv) or bool(use_prepacked_v_for_pv) - if use_prepacked_v_for_pv and v_packed_workspace is None: - raise ValueError("use_prepacked_v_for_pv requires v_packed_workspace to be provided.") - can_use_prepacked_v_for_pv = ( - wants_prepacked_v_for_pv + if env_gen_qkpv == 1 and not can_use_gen_qkpv: + raise ValueError( + "TRTLLM_FP4_MLA_GEN_QKPV=1 requires qlen=4, prepacked V, " + "num_heads=128, q/k/v dims 640/576/512, page_size=128, " + "BLOCK_K=512, BLOCK_V=128, and GEN_QKPV_BLOCK_H in (16, 32, 64)." + ) + if can_use_gen_qkpv: + if v_repack_stream is not None: + torch.cuda.current_stream(q_fp4.device).wait_stream(v_repack_stream) + v_repack_stream = None + num_gen_qkpv_page_groups = _ceil_div(max_pages, gen_qkpv_group_pages) + gen_qkpv_partial_dtype = torch.float32 + if env_gen_qkpv_partial_dtype in ("bf16", "bfloat16"): + gen_qkpv_partial_dtype = torch.bfloat16 + gen_qkpv_partial_o = _workspace_tensor( + None, + (num_gen_qkpv_page_groups, num_gen, num_heads, v_head_dim), + dtype=gen_qkpv_partial_dtype, + device=q_fp4.device, + name="gen_qkpv_partial_o", + ) + gen_qkpv_partial_m = _workspace_tensor( + None, + (num_gen_qkpv_page_groups, num_gen, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="gen_qkpv_partial_m", + ) + gen_qkpv_partial_l = _workspace_tensor( + None, + (num_gen_qkpv_page_groups, num_gen, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="gen_qkpv_partial_l", + ) + _fp4_mla_attention_gen_qkpv_group_kernel[ + (num_gen, triton.cdiv(num_heads, gen_qkpv_block_h), num_gen_qkpv_page_groups) + ]( + gen_qkpv_partial_o, + gen_qkpv_partial_m, + gen_qkpv_partial_l, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + v_sf, + v_packed, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + v_sf.stride(0), + gen_qkpv_partial_o.stride(0), + gen_qkpv_partial_o.stride(1), + gen_qkpv_partial_o.stride(2), + gen_qkpv_partial_o.stride(3), + gen_qkpv_partial_m.stride(0), + gen_qkpv_partial_m.stride(1), + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + BLOCK_H=gen_qkpv_block_h, + BLOCK_T=page_size, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + BLOCK_V=block_v, + MAX_PAGES=max_pages, + GROUP_PAGES=gen_qkpv_group_pages, + ALLOW_PARTIAL_GROUPS=max_pages % gen_qkpv_group_pages != 0, + USE_TMA_DATA_LOAD=use_tma_data_load, + ASSUME_FULL_HEADS=True, + ASSUME_VALID_PAGES=True, + **gen_qkpv_launch_meta, + ) + _fp4_mla_attention_online_qkpv_reduce_kernel[ + (num_gen, triton.cdiv(num_heads, gen_qkpv_block_h), num_dim_blocks) + ]( + output, + gen_qkpv_partial_o, + gen_qkpv_partial_m, + gen_qkpv_partial_l, + global_scale, + output.stride(0), + output.stride(1), + output.stride(2), + gen_qkpv_partial_o.stride(0), + gen_qkpv_partial_o.stride(1), + gen_qkpv_partial_o.stride(2), + gen_qkpv_partial_o.stride(3), + gen_qkpv_partial_m.stride(0), + gen_qkpv_partial_m.stride(1), + NUM_HEADS=num_heads, + V_HEAD_D=v_head_dim, + NUM_PAGE_GROUPS=num_gen_qkpv_page_groups, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=gen_qkpv_block_h, + BLOCK_V=block_v, + FP4_PV=True, + **pv_launch_meta, + ) + return output + + online_qkpv_max_batch = env_online_qkpv_max_batch if env_online_qkpv_max_batch is not None else 32 + can_use_online_qkpv = ( + env_online_qkpv == 1 + and can_use_prepacked_v_for_pv + and query_len_per_seq == 1 + and num_gen <= online_qkpv_max_batch and use_tma_data_load and assume_full_heads and assume_full_pages and assume_valid_pages + and num_heads == 128 + and q_head_dim == 640 + and k_head_dim == 576 + and q_residual_dim == 64 and v_head_dim == 512 and page_size == 128 and block_h in (64, 128) - and block_v == 128 + and block_k == 512 + and full_block_end == 512 + and tail_block_k == 128 + and block_v in (128, 256) and sf_per_page == 8 ) - if wants_prepacked_v_for_pv and not can_use_prepacked_v_for_pv: - raise ValueError( - "prepacked V PV path requires TMA, full heads/pages, valid pages, " - "v_head_dim=512, page_size=128, block_h in (64, 128), block_v=128, and sf_per_page=8." + if can_use_online_qkpv: + if v_repack_stream is not None: + torch.cuda.current_stream(q_fp4.device).wait_stream(v_repack_stream) + v_repack_stream = None + online_group_pages = env_online_qkpv_group_pages if env_online_qkpv_group_pages is not None else 128 + online_group_pages = max(1, min(int(online_group_pages), max_pages)) + num_online_page_groups = _ceil_div(max_pages, online_group_pages) + online_partial_o = _workspace_tensor( + None, + (num_online_page_groups, num_gen, num_heads, v_head_dim), + dtype=torch.float32, + device=q_fp4.device, + name="online_partial_o", + ) + online_partial_m = _workspace_tensor( + None, + (num_online_page_groups, num_gen, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="online_partial_m", + ) + online_partial_l = _workspace_tensor( + None, + (num_online_page_groups, num_gen, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="online_partial_l", + ) + _fp4_mla_attention_online_qkpv_group_kernel[ + (num_gen, num_head_blocks, num_dim_blocks * num_online_page_groups) + ]( + online_partial_o, + online_partial_m, + online_partial_l, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + v_sf, + v_packed, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + v_sf.stride(0), + online_partial_o.stride(0), + online_partial_o.stride(1), + online_partial_o.stride(2), + online_partial_o.stride(3), + online_partial_m.stride(0), + online_partial_m.stride(1), + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + BLOCK_T=page_size, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + BLOCK_V=block_v, + MAX_PAGES=max_pages, + GROUP_PAGES=online_group_pages, + NUM_DIM_BLOCKS=num_dim_blocks, + ALLOW_PARTIAL_GROUPS=max_pages % online_group_pages != 0, + FP4_PV=env_online_qkpv_fp4_pv == 1, + **page_stats_launch_meta, ) - if can_use_prepacked_v_for_pv: - v_packed = _workspace_tensor( - v_packed_workspace, - (num_pages * num_dim_blocks * block_v, page_size // 2), - dtype=torch.uint8, - device=q_fp4.device, - name="v_packed", + _fp4_mla_attention_online_qkpv_reduce_kernel[(num_gen, num_head_blocks, num_dim_blocks)]( + output, + online_partial_o, + online_partial_m, + online_partial_l, + global_scale, + output.stride(0), + output.stride(1), + output.stride(2), + online_partial_o.stride(0), + online_partial_o.stride(1), + online_partial_o.stride(2), + online_partial_o.stride(3), + online_partial_m.stride(0), + online_partial_m.stride(1), + NUM_HEADS=num_heads, + V_HEAD_D=v_head_dim, + NUM_PAGE_GROUPS=num_online_page_groups, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=block_h, + BLOCK_V=block_v, + FP4_PV=env_online_qkpv_fp4_pv == 1, + **pv_launch_meta, ) - else: - v_packed = kv_cache - if fused_prob_pack: - p_probs = None - else: - p_probs_shape = (max(num_gen * num_heads, 1), page_size) - if page_pipeline_streams > 1: - p_probs = _workspace_tensor( - p_probs_workspace, - (page_pipeline_streams, *p_probs_shape), - dtype=torch.float32, - device=q_fp4.device, - name="p_probs", - ) - else: - p_probs = _workspace_tensor( - p_probs_workspace, - p_probs_shape, - dtype=torch.float32, - device=q_fp4.device, - name="p_probs", - ) - max_scores = _workspace_tensor( - max_scores_workspace, - (num_gen, num_heads), - dtype=torch.float32, - device=q_fp4.device, - name="max_scores", - ) - denom = _workspace_tensor( - denom_workspace, - (num_gen, num_heads), - dtype=torch.float32, - device=q_fp4.device, - name="denom", - ) - - v_repack_stream = None - if bool(prepack_v_for_pv) and can_use_prepacked_v_for_pv: - current_stream = torch.cuda.current_stream(q_fp4.device) - v_repack_stream = torch.cuda.Stream(device=q_fp4.device) - v_repack_stream.wait_stream(current_stream) - with torch.cuda.stream(v_repack_stream): - _fp4_mla_attention_v_repack_kernel[(num_pages, num_repack_dim_blocks)]( - v_packed, - kv_cache, - num_pages, - kv_s0, - kv_s2, - kv_s4, - V_HEAD_D=v_head_dim, - PAGE_SIZE=page_size, - BLOCK_V=v_repack_block_v, - **launch_meta, - ) + return output if parallel_page_stats is None: parallel_page_stats = triton_backend == "nvt" and max_pages >= 8 if pack_prob_in_page_stats is None: pack_prob_in_page_stats = parallel_page_stats and fused_prob_pack pack_prob_in_page_stats = bool(pack_prob_in_page_stats and parallel_page_stats and fused_prob_pack) + debug_pack_prob_in_page_stats = ( + pack_prob_in_page_stats + if env_debug_page_stats_pack is None + else env_debug_page_stats_pack != 0 + ) page_stats_group_sizes = (2, 4, 8, 16, 32, 64, 128, 256, 512, 1024) if page_stats_group_size is None and env_group_pages is not None: page_stats_group_size = env_group_pages if page_stats_group_size is None: - if num_gen <= 32: + if p_by_query and query_len_per_seq > 1 and max_pages >= 512: + target_group_pages = 512 + elif num_gen <= 32: target_group_pages = max(8, 8 * num_gen) elif num_gen >= 256: target_group_pages = 256 @@ -2934,21 +6160,31 @@ def alloc_fn(size: int, alignment: int, stream: Optional[int]): ( group_size for group_size in reversed(page_stats_group_sizes) - if group_size <= target_group_pages and max_pages % group_size == 0 + if group_size <= target_group_pages ), 8, ) else: page_stats_group_size = int(page_stats_group_size) + if env_mtp_page_stats_pair == 1 and p_by_query and query_len_per_seq == 4: + if env_mtp_page_stats_pair_group_pages is not None: + page_stats_group_size = int(env_mtp_page_stats_pair_group_pages) + else: + # The two-query MTP page-stats kernel duplicates the QK/prob-pack + # work inside each page group. Keeping the group small avoids very + # large TileIR kernels while still sharing K loads across query pairs. + page_stats_group_size = min(page_stats_group_size, 2) + grouped_storage_full_pages = assume_full_pages or (p_by_query and can_address_all_compact_pages) + grouped_assume_valid_pages = assume_valid_pages or (p_by_query and can_address_all_compact_pages) can_group_page_stats = ( page_stats_group_size in page_stats_group_sizes and triton_backend == "nvt" and parallel_page_stats - and pack_prob_in_page_stats + and (pack_prob_in_page_stats or env_debug_stop_after_page_stats == 1) and use_tma_data_load and assume_full_heads - and assume_full_pages - and assume_valid_pages + and grouped_storage_full_pages + and grouped_assume_valid_pages and num_heads == 128 and q_head_dim == 640 and k_head_dim == 576 @@ -2959,15 +6195,21 @@ def alloc_fn(size: int, alignment: int, stream: Optional[int]): and full_block_end == 512 and tail_block_k == 128 and sf_per_page == 8 - and max_pages % page_stats_group_size == 0 ) page_stats_group_size = page_stats_group_size if can_group_page_stats else 1 + num_page_groups = _ceil_div(max_pages, page_stats_group_size) if page_stats_group_size > 1 else max_pages + allow_partial_page_groups = page_stats_group_size > 1 and max_pages % page_stats_group_size != 0 group_reduce_stats = ( (env_group_reduce_stats != 0 if env_group_reduce_stats is not None else triton_backend == "nvt") and page_stats_group_size > 1 ) + page_stats_max_entries = max_pages + page_stats_num_groups = num_page_groups + pv_apply_prob_scale = False + page_max = None + page_sum = None if parallel_page_stats: - page_stats_shape = (num_gen, max_pages, num_heads) + page_stats_shape = (num_gen, page_stats_max_entries, num_heads) page_max = _workspace_tensor( page_max_workspace, page_stats_shape, @@ -2982,10 +6224,198 @@ def alloc_fn(size: int, alignment: int, stream: Optional[int]): device=q_fp4.device, name="page_sum", ) - if page_stats_group_size > 1: - _fp4_mla_attention_page_stats_grouped_kernel[ - (num_gen, num_head_blocks, max_pages // page_stats_group_size) - ]( + can_use_mtp_page_stats_pair = ( + env_mtp_page_stats_pair == 1 + and p_by_query + and query_len_per_seq == 4 + and can_address_all_compact_pages + and block_h == 128 + and use_tma_data_load + and assume_full_heads + and grouped_storage_full_pages + and grouped_assume_valid_pages + and num_heads == 128 + and q_head_dim == 640 + and k_head_dim == 576 + and q_residual_dim == 64 + and page_size == 128 + and block_k == 512 + and full_block_end == 512 + and tail_block_k == 128 + and sf_per_page == 8 + ) + split_tail_group_enabled = ( + env_mtp_split_tail_group != 0 if env_mtp_split_tail_group is not None else False + ) + can_split_mtp_tail_group = ( + split_tail_group_enabled + and not can_use_mtp_page_stats_pair + and p_by_query + and query_len_per_seq > 1 + and page_stats_group_size > 1 + and num_page_groups > 1 + and (pack_prob_in_page_stats or env_debug_stop_after_page_stats == 1) + and group_reduce_stats + and can_address_all_compact_pages + and block_h == 128 + and use_tma_data_load + and assume_full_heads + and grouped_storage_full_pages + and grouped_assume_valid_pages + and num_heads == 128 + and q_head_dim == 640 + and k_head_dim == 576 + and q_residual_dim == 64 + and page_size == 128 + and block_k == 512 + and full_block_end == 512 + and tail_block_k == 128 + and sf_per_page == 8 + ) + if page_stats_group_size > 1 or can_use_mtp_page_stats_pair: + if can_use_mtp_page_stats_pair: + _fp4_mla_attention_page_stats_grouped_mtp_pair_kernel[ + ( + num_gen_seqs, + num_head_blocks, + num_page_groups * (query_len_per_seq // 2), + ) + ]( + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + page_max.stride(0), + page_max.stride(1), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + BLOCK_H=block_h, + BLOCK_T=page_size, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + PACK_PROBS=debug_pack_prob_in_page_stats, + GROUP_REDUCE_STATS=group_reduce_stats, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=grouped_assume_valid_pages, + MAX_PAGES=max_pages, + GROUP_PAGES=page_stats_group_size, + ALLOW_PARTIAL_GROUPS=allow_partial_page_groups, + Q_PER_GROUP=2, + DUPLICATE_TAIL_K=env_duplicate_tail_k == 1, + **page_stats_launch_meta, + ) + else: + page_stats_grouped_kernel = ( + _fp4_mla_attention_page_stats_grouped_kernel + if block_h == 128 + else _fp4_mla_attention_page_stats_grouped_generic_kernel + ) + if can_split_mtp_tail_group: + page_group_launches = ( + (num_page_groups - 1, True, False, 0), + (1, False, allow_partial_page_groups, num_page_groups - 1), + ) + else: + page_group_launches = ( + (num_page_groups, assume_full_pages, allow_partial_page_groups, 0), + ) + for grid_page_groups, launch_assume_full_pages, launch_allow_partial_groups, page_group_offset in page_group_launches: + if grid_page_groups <= 0: + continue + page_stats_grouped_kernel[(num_gen, num_head_blocks, grid_page_groups)]( + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4_2d, + q_sf_flat, + kv_cache, + sf_cache, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + num_pages, + q_fp4_2d.stride(0), + q_fp4_2d.stride(1), + kv_s0, + kv_s2, + kv_s4, + sf_cache.stride(0), + page_max.stride(0), + page_max.stride(1), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + q_fp4_2d.shape[0], + sm_scale=sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=k_head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + P_BY_QUERY=p_by_query, + BLOCK_H=block_h, + BLOCK_T=page_size, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + PACK_PROBS=debug_pack_prob_in_page_stats, + GROUP_REDUCE_STATS=group_reduce_stats, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=launch_assume_full_pages, + MASK_MTP_FINAL_PAGE_ONLY=mask_mtp_final_page_only, + ASSUME_VALID_PAGES=grouped_assume_valid_pages, + MAX_PAGES=max_pages, + GROUP_PAGES=page_stats_group_size, + ALLOW_PARTIAL_GROUPS=launch_allow_partial_groups, + PAGE_GROUP_OFFSET=page_group_offset, + DUPLICATE_TAIL_K=env_duplicate_tail_k == 1, + **page_stats_launch_meta, + ) + else: + _fp4_mla_attention_page_stats_kernel[(num_gen, num_head_blocks, max_pages)]( page_max, page_sum, p_fp4, @@ -3023,107 +6453,308 @@ def alloc_fn(size: int, alignment: int, stream: Optional[int]): K_SF_PER_TOKEN=k_sf_per_token, SF_PER_PAGE=sf_per_page, P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_BY_QUERY=p_by_query, BLOCK_H=block_h, BLOCK_T=page_size, BLOCK_K=block_k, FULL_BLOCK_END=full_block_end, TAIL_BLOCK_K=tail_block_k, USE_TMA_DATA_LOAD=use_tma_data_load, - PACK_PROBS=pack_prob_in_page_stats, - GROUP_REDUCE_STATS=group_reduce_stats, + PACK_PROBS=debug_pack_prob_in_page_stats, ASSUME_FULL_HEADS=assume_full_heads, ASSUME_FULL_PAGES=assume_full_pages, ASSUME_VALID_PAGES=assume_valid_pages, - GROUP_PAGES=page_stats_group_size, - **launch_meta, + **page_stats_launch_meta, ) - else: - _fp4_mla_attention_page_stats_kernel[(num_gen, num_head_blocks, max_pages)]( - page_max, - page_sum, + _debug_mark("page_stats") + if env_debug_stop_after_page_stats == 1: + _debug_report() + return output + group_pv_max_batch = env_group_pv_max_batch if env_group_pv_max_batch is not None else 64 + can_use_group_pv = ( + env_group_pv == 1 + and page_stats_group_size > 1 + and group_reduce_stats + and can_use_prepacked_v_for_pv + and query_len_per_seq == 1 + and not p_by_query + and num_gen <= group_pv_max_batch + and use_tma_data_load + and assume_full_heads + and assume_full_pages + and assume_valid_pages + and num_heads == 128 + and v_head_dim == 512 + and page_size == 128 + and block_h == 128 + and block_v in (128, 256) + and sf_per_page == 8 + ) + if can_use_group_pv: + if v_repack_stream is not None: + torch.cuda.current_stream(q_fp4.device).wait_stream(v_repack_stream) + v_repack_stream = None + group_pv_partial_dtype = torch.float32 + if os.environ.get("TRTLLM_FP4_MLA_GROUP_PV_PARTIAL_DTYPE", "").lower() in ("bf16", "bfloat16"): + group_pv_partial_dtype = torch.bfloat16 + group_partial_o = _workspace_tensor( + None, + (num_page_groups, num_gen, num_heads, v_head_dim), + dtype=group_pv_partial_dtype, + device=q_fp4.device, + name="group_partial_o", + ) + _fp4_mla_attention_pv_group_partial_prepacked_v_kernel[ + (num_gen, num_head_blocks, num_dim_blocks * num_page_groups) + ]( + group_partial_o, p_fp4, p_sf, - q_fp4_2d, - q_sf_flat, - kv_cache, - sf_cache, - global_scale, + v_sf, + v_packed, src_page_ids, paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], num_pages, - q_fp4_2d.stride(0), - q_fp4_2d.stride(1), - kv_s0, - kv_s2, - kv_s4, - sf_cache.stride(0), - page_max.stride(0), - page_max.stride(1), + group_partial_o.stride(0), + group_partial_o.stride(1), + group_partial_o.stride(2), + group_partial_o.stride(3), + group_partial_o.shape[0] * group_partial_o.shape[1] * group_partial_o.shape[2], p_fp4.stride(0), p_fp4.stride(1), p_fp4.shape[0], - q_fp4_2d.shape[0], - sm_scale=sm_scale, + v_sf.stride(0), NUM_HEADS=num_heads, - Q_HEAD_D=q_head_dim, - K_HEAD_D=k_head_dim, - Q_RESIDUAL_D=q_residual_dim, + V_HEAD_D=v_head_dim, PAGE_SIZE=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - Q_SF_PER_TOKEN=q_sf_per_token, - K_SF_PER_TOKEN=k_sf_per_token, SF_PER_PAGE=sf_per_page, + BLOCK_H=block_h, + BLOCK_V=block_v, + MAX_PAGES=max_pages, + GROUP_PAGES=page_stats_group_size, + NUM_DIM_BLOCKS=num_dim_blocks, + ALLOW_PARTIAL_GROUPS=allow_partial_page_groups, + PV_LOOP_STAGES=int(pv_loop_stages), + **pv_launch_meta, + ) + _fp4_mla_attention_pv_group_partial_reduce_kernel[ + (num_gen, num_head_blocks, num_dim_blocks) + ]( + output, + group_partial_o, + page_sum, + global_scale, + output.stride(0), + output.stride(1), + output.stride(2), + group_partial_o.stride(0), + group_partial_o.stride(1), + group_partial_o.stride(2), + group_partial_o.stride(3), + group_partial_o.shape[0] * group_partial_o.shape[1] * group_partial_o.shape[2], + page_sum.stride(0), + page_sum.stride(1), + NUM_HEADS=num_heads, + V_HEAD_D=v_head_dim, + NUM_PAGE_GROUPS=num_page_groups, P_GLOBAL_SCALE=p_global_scale, BLOCK_H=block_h, - BLOCK_T=page_size, - BLOCK_K=block_k, - FULL_BLOCK_END=full_block_end, - TAIL_BLOCK_K=tail_block_k, - USE_TMA_DATA_LOAD=use_tma_data_load, - PACK_PROBS=pack_prob_in_page_stats, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, + BLOCK_V=block_v, + **pv_launch_meta, ) - _fp4_mla_attention_reduce_stats_kernel[(num_gen, num_head_blocks)]( - max_scores, - denom, - page_max, - page_sum, - max_scores.stride(0), - page_max.stride(0), - page_max.stride(1), - NUM_HEADS=num_heads, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - GROUP_REDUCE_STATS=group_reduce_stats, - GROUP_PAGES=page_stats_group_size, - **launch_meta, + return output + scale_from_group_stats = ( + (env_scale_from_group_stats == 1 or (env_scale_from_group_stats is None and num_gen >= 64)) + and pack_prob_in_page_stats + and group_reduce_stats + and page_stats_group_size > 1 ) - if pack_prob_in_page_stats: - _fp4_mla_attention_prob_scale_kernel[(num_gen, num_head_blocks, max_pages)]( - p_sf, + if not scale_from_group_stats: + _fp4_mla_attention_reduce_stats_kernel[(num_gen, num_head_blocks)]( max_scores, denom, page_max, - paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], + page_sum, max_scores.stride(0), page_max.stride(0), page_max.stride(1), NUM_HEADS=num_heads, + MAX_PAGES=page_stats_max_entries, + BLOCK_H=block_h, + GROUP_REDUCE_STATS=group_reduce_stats, + GROUP_PAGES=page_stats_group_size, + NUM_PAGE_GROUPS=page_stats_num_groups, + **page_stats_launch_meta, + ) + pv_apply_prob_scale = ( + env_pv_apply_prob_scale == 1 + and pack_prob_in_page_stats + and not scale_from_group_stats + and can_use_prepacked_v_for_pv + and query_len_per_seq == 1 + and not p_by_query + and page_max is not None + ) + if pack_prob_in_page_stats and not pv_apply_prob_scale: + if scale_from_group_stats: + prob_scale_assume_valid_pages = ( + grouped_assume_valid_pages if mask_mtp_final_page_only else assume_valid_pages + ) + + def _launch_prob_scale_from_group_stats( + grid_pages: int, + launch_assume_full_pages: bool, + page_rel_offset: int, + ) -> None: + if grid_pages <= 0: + return + _fp4_mla_attention_prob_scale_from_group_stats_kernel[ + (num_gen, num_head_blocks, grid_pages) + ]( + p_sf, + page_max, + page_sum, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + page_max.stride(0), + page_max.stride(1), + NUM_HEADS=num_heads, + PAGE_SIZE=page_size, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_BY_QUERY=p_by_query, + BLOCK_H=block_h, + NUM_PAGE_GROUPS=num_page_groups, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=launch_assume_full_pages, + ASSUME_VALID_PAGES=prob_scale_assume_valid_pages, + PAGE_REL_OFFSET=page_rel_offset, + **page_stats_launch_meta, + ) + + if mask_mtp_final_page_only: + _launch_prob_scale_from_group_stats(max_pages - 1, True, 0) + _launch_prob_scale_from_group_stats(1, False, max_pages - 1) + else: + _launch_prob_scale_from_group_stats(max_pages, assume_full_pages, 0) + else: + _fp4_mla_attention_prob_scale_kernel[(num_gen, num_head_blocks, max_pages)]( + p_sf, + max_scores, + denom, + page_max, + paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + max_scores.stride(0), + page_max.stride(0), + page_max.stride(1), + NUM_HEADS=num_heads, + PAGE_SIZE=page_size, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_BY_QUERY=p_by_query, + BLOCK_H=block_h, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **page_stats_launch_meta, + ) + _debug_mark("prob_scale") + can_use_pv_atomic_split = ( + env_pv_atomic_split == 1 + and pack_prob_in_page_stats + and can_use_prepacked_v_for_pv + and query_len_per_seq == 1 + and not p_by_query + and use_tma_data_load + and assume_full_heads + and assume_full_pages + and assume_valid_pages + and assume_full_pv_heads + and pv_block_h == 128 + and num_heads == 128 + and v_head_dim == 512 + and page_size == 128 + and block_h == 128 + and block_v in (128, 256) + and sf_per_page == 8 + ) + if can_use_pv_atomic_split: + if v_repack_stream is not None: + torch.cuda.current_stream(q_fp4.device).wait_stream(v_repack_stream) + v_repack_stream = None + pv_atomic_group_pages = ( + env_pv_atomic_group_pages + if env_pv_atomic_group_pages is not None + else page_stats_group_size + ) + pv_atomic_group_pages = max(1, min(int(pv_atomic_group_pages), max_pages)) + num_pv_atomic_page_groups = _ceil_div(max_pages, pv_atomic_group_pages) + out_acc = _workspace_tensor( + None, + (num_gen, num_heads, v_head_dim), + dtype=torch.float32, + device=q_fp4.device, + name="pv_atomic_out_acc", + ) + out_acc.zero_() + _fp4_mla_attention_pv_atomic_split_prepacked_v_kernel[ + (num_gen, num_pv_head_blocks, num_dim_blocks * num_pv_atomic_page_groups) + ]( + out_acc, + p_fp4, + p_sf, + v_sf, + v_packed, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + num_pages, + out_acc.stride(0), + out_acc.stride(1), + out_acc.stride(2), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=v_head_dim, PAGE_SIZE=page_size, SF_PER_PAGE=sf_per_page, - BLOCK_H=block_h, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=pv_block_h, + BLOCK_V=block_v, + MAX_PAGES=max_pages, + GROUP_PAGES=pv_atomic_group_pages, + NUM_DIM_BLOCKS=num_dim_blocks, + ALLOW_PARTIAL_GROUPS=max_pages % pv_atomic_group_pages != 0, + PV_LOOP_STAGES=int(pv_loop_stages), + **pv_launch_meta, + ) + _fp4_mla_attention_cast_acc_kernel[ + (num_gen, num_pv_head_blocks, num_dim_blocks) + ]( + output, + out_acc, + output.stride(0), + output.stride(1), + output.stride(2), + out_acc.stride(0), + out_acc.stride(1), + out_acc.stride(2), + NUM_HEADS=num_heads, + V_HEAD_D=v_head_dim, + BLOCK_H=pv_block_h, + BLOCK_V=block_v, + **pv_launch_meta, ) + return output else: _fp4_mla_attention_stats_kernel[(num_gen, num_head_blocks)]( max_scores, @@ -3156,6 +6787,7 @@ def alloc_fn(size: int, alignment: int, stream: Optional[int]): Q_SF_PER_TOKEN=q_sf_per_token, K_SF_PER_TOKEN=k_sf_per_token, MAX_PAGES=max_pages, + QUERY_LEN_PER_SEQ=query_len_per_seq, BLOCK_H=block_h, BLOCK_T=page_size, BLOCK_K=block_k, @@ -3165,7 +6797,7 @@ def alloc_fn(size: int, alignment: int, stream: Optional[int]): ASSUME_FULL_HEADS=assume_full_heads, ASSUME_FULL_PAGES=assume_full_pages, ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, + **page_stats_launch_meta, ) def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None): @@ -3207,6 +6839,9 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None K_SF_PER_TOKEN=k_sf_per_token, SF_PER_PAGE=sf_per_page, P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_BY_QUERY=p_by_query, BLOCK_H=block_h, BLOCK_K=block_k, FULL_BLOCK_END=full_block_end, @@ -3215,7 +6850,7 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None ASSUME_FULL_HEADS=assume_full_heads, ASSUME_FULL_PAGES=assume_full_pages, ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, + **page_stats_launch_meta, ) return assert p_probs_slot is not None @@ -3253,6 +6888,7 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None FP4_BLOCK=FP4_BLOCK_SIZE, Q_SF_PER_TOKEN=q_sf_per_token, K_SF_PER_TOKEN=k_sf_per_token, + QUERY_LEN_PER_SEQ=query_len_per_seq, BLOCK_H=block_h, BLOCK_K=block_k, FULL_BLOCK_END=full_block_end, @@ -3261,7 +6897,7 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None ASSUME_FULL_HEADS=assume_full_heads, ASSUME_FULL_PAGES=assume_full_pages, ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, + **page_stats_launch_meta, ) _fp4_mla_attention_prob_pack_page_kernel[(num_gen, sf_per_page, num_head_blocks)]( p_fp4, @@ -3280,10 +6916,13 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None FP4_BLOCK=FP4_BLOCK_SIZE, SF_PER_PAGE=sf_per_page, P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_BY_QUERY=p_by_query, BLOCK_H=block_h, ASSUME_FULL_HEADS=assume_full_heads, ASSUME_FULL_PAGES=assume_full_pages, - **launch_meta, + **page_stats_launch_meta, ) if pack_prob_in_page_stats: @@ -3326,6 +6965,9 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None K_SF_PER_TOKEN=k_sf_per_token, SF_PER_PAGE=sf_per_page, P_GLOBAL_SCALE=p_global_scale, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_BY_QUERY=p_by_query, BLOCK_H=block_h, BLOCK_K=block_k, FULL_BLOCK_END=full_block_end, @@ -3334,7 +6976,7 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None PAGE_REL_FROM_GRID=True, ASSUME_FULL_HEADS=assume_full_heads, ASSUME_FULL_PAGES=assume_full_pages, - **launch_meta, + **page_stats_launch_meta, ) elif page_pipeline_streams == 1: for page_rel in range(max_pages): @@ -3359,10 +7001,66 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None torch.cuda.current_stream(q_fp4.device).wait_stream(v_repack_stream) if can_use_prepacked_v_for_pv: + mtp_pv_pair_enabled = env_mtp_pv_pair != 0 if env_mtp_pv_pair is not None else True + can_use_mtp_pv_pair = ( + mtp_pv_pair_enabled + and p_by_query + and query_len_per_seq == 4 + and can_address_all_compact_pages + and num_heads == 128 + and v_head_dim == 512 + and page_size == 128 + and pv_block_h == 128 + and block_v in (128, 256) + and sf_per_page == 8 + and assume_full_pv_heads + and assume_full_v + ) + if can_use_mtp_pv_pair: + _fp4_mla_attention_pv_mtp_pair_prepacked_v_kernel[ + ( + num_gen_seqs, + num_pv_head_blocks, + num_dim_blocks * (query_len_per_seq // 2), + ) + ]( + output, + p_fp4, + p_sf, + v_sf, + v_packed, + global_scale, + src_page_ids, + paged_kv_indptr_decode, + num_pages, + output.stride(0), + output.stride(1), + output.stride(2), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=pv_block_h, + BLOCK_V=block_v, + NUM_DIM_BLOCKS=num_dim_blocks, + Q_PER_GROUP=2, + PV_LOOP_STAGES=int(pv_loop_stages), + **pv_launch_meta, + ) + _debug_mark("pv") + _debug_report() + return output _fp4_mla_attention_pv_prepacked_v_kernel[ ( num_gen, - num_head_blocks, + num_pv_head_blocks, num_dim_blocks, ) ]( @@ -3373,6 +7071,9 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None v_sf, v_packed, global_scale, + max_scores, + denom, + page_max if page_max is not None else max_scores, src_page_ids, paged_kv_indptr_decode, kv_lens, @@ -3385,6 +7086,9 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None p_fp4.stride(0), p_fp4.stride(1), p_fp4.shape[0], + max_scores.stride(0), + page_max.stride(0) if page_max is not None else max_scores.stride(0), + page_max.stride(1) if page_max is not None else max_scores.stride(0), kv_s0, kv_s2, kv_s4, @@ -3394,27 +7098,36 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None PAGE_SIZE=page_size, FP4_BLOCK=FP4_BLOCK_SIZE, SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, MAX_PAGES=max_pages, + P_BY_QUERY=p_by_query, P_GLOBAL_SCALE=p_global_scale, - BLOCK_H=block_h, + BLOCK_H=pv_block_h, BLOCK_V=block_v, - USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, - USE_TMA_V_LOAD=use_tma_data_load and v_head_dim % block_v == 0, + USE_TMA_P_LOAD=use_pv_p_tma_data_load and assume_full_pv_heads and assume_valid_pages, + USE_TMA_V_LOAD=use_pv_v_tma_data_load and v_head_dim % block_v == 0, + USE_TMA_OUT_STORE=use_pv_out_tma_data_store + and (not p_by_query or env_pv_out_tma == 1) + and assume_full_pv_heads + and assume_full_v, USE_PREPACKED_V=True, PV_M_PACKED_V=False, PV_LOOP_STAGES=int(pv_loop_stages), - ASSUME_FULL_HEADS=assume_full_heads, + PV_APPLY_PROB_SCALE=pv_apply_prob_scale, + PV_SCALE_IN_SF=env_pv_scale_in_sf == 1, + ASSUME_FULL_HEADS=assume_full_pv_heads, ASSUME_FULL_PAGES=assume_full_pages, ASSUME_FULL_V=assume_full_v, ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, + **pv_launch_meta, ) + _debug_mark("pv") else: num_dim_blocks = triton.cdiv(v_head_dim, block_v) _fp4_mla_attention_pv_kernel[ ( num_gen, - num_head_blocks, + num_pv_head_blocks, num_dim_blocks, ) ]( @@ -3445,19 +7158,27 @@ def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None PAGE_SIZE=page_size, FP4_BLOCK=FP4_BLOCK_SIZE, SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, MAX_PAGES=max_pages, + P_BY_QUERY=p_by_query, P_GLOBAL_SCALE=p_global_scale, - BLOCK_H=block_h, + BLOCK_H=pv_block_h, BLOCK_V=block_v, - USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, - USE_TMA_V_LOAD=use_tma_data_load and v_head_dim % block_v == 0, + USE_TMA_P_LOAD=use_pv_p_tma_data_load and assume_full_pv_heads and assume_valid_pages, + USE_TMA_V_LOAD=use_pv_v_tma_data_load and v_head_dim % block_v == 0, + USE_TMA_OUT_STORE=use_pv_out_tma_data_store + and (not p_by_query or env_pv_out_tma == 1) + and assume_full_pv_heads + and assume_full_v, PV_LOOP_STAGES=int(pv_loop_stages), - ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_HEADS=assume_full_pv_heads, ASSUME_FULL_PAGES=assume_full_pages, ASSUME_FULL_V=assume_full_v, ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, + **pv_launch_meta, ) + _debug_mark("pv") + _debug_report() return output diff --git a/tests/unittest/_torch/attention/test_fp4_mla.py b/tests/unittest/_torch/attention/test_fp4_mla.py index 6732fe18eeda..913f305f0f2a 100644 --- a/tests/unittest/_torch/attention/test_fp4_mla.py +++ b/tests/unittest/_torch/attention/test_fp4_mla.py @@ -33,6 +33,8 @@ run_fp4_mla_attention_decode, scatter_fp4_mla_kv_cache, update_hp_kv_for_fp4_mla, + _get_cutile_v_packed_cache, + _maybe_update_cutile_v_packed_cache, ) from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings.executor import KvCacheConfig @@ -1408,6 +1410,153 @@ def test_fp4_mla_attention_decode_cutile_matches_reference(monkeypatch): ) +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +def test_fp4_mla_attention_decode_cutile_shared_v_pack_matches_reference(monkeypatch): + """Shared V-packed storage must preserve the prepacked PV fast path.""" + monkeypatch.setenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1") + monkeypatch.setenv("TRTLLM_FP4_MLA_SHARE_V_PACK_STORAGE", "1") + _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + backend="cutile", + num_heads=128, + seq_lens=[128, 128], + seed=17, + check_probs=False, + ) + + +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +def test_fp4_mla_attention_decode_cutile_grouped_tail_matches_reference(monkeypatch): + """Grouped page-stats must handle a partial final page group.""" + monkeypatch.setenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1") + monkeypatch.setenv("TRTLLM_FP4_MLA_SHARE_V_PACK_STORAGE", "1") + monkeypatch.setenv("TRTLLM_FP4_MLA_GROUP_PAGES", "8") + _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + backend="cutile", + num_heads=128, + seq_lens=[9 * FP4_MLA_TOKENS_PER_BLOCK, 9 * FP4_MLA_TOKENS_PER_BLOCK], + seed=19, + check_probs=False, + ) + + +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +def test_fp4_mla_cutile_shared_v_pack_storage_is_layer_tagged(monkeypatch): + """Layer ownership is metadata state; storage is reused across layers.""" + monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, "cutile") + monkeypatch.setenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1") + monkeypatch.setenv("TRTLLM_FP4_MLA_SHARE_V_PACK_STORAGE", "1") + + device = torch.device("cuda") + metadata = SimpleNamespace() + num_pages = 4 + page_size = FP4_MLA_TOKENS_PER_BLOCK + v_head_dim = 512 + head_dim = v_head_dim + 64 + kv_cache = torch.randint( + 0, + 256, + (num_pages, 1, page_size, 1, head_dim // 2), + dtype=torch.uint8, + device=device, + ) + page_ids = torch.arange(num_pages, dtype=torch.int32, device=device) + v_sf = torch.empty( + (2, num_pages, v_head_dim, page_size // FP4_BLOCK_SIZE), + dtype=torch.float8_e4m3fn, + device=device, + ) + + _maybe_update_cutile_v_packed_cache( + metadata, + 0, + kv_cache, + page_ids, + v_head_dim=v_head_dim, + page_size=page_size, + local_layer=0, + v_sf=v_sf[0], + ) + torch.cuda.synchronize() + shared = metadata._fp4_mla_attention_v_packed_buf + shared_ptr = shared.data_ptr() + assert not hasattr(metadata, "_fp4_mla_attention_v_packed_buf_l0") + assert _get_cutile_v_packed_cache( + metadata, + 0, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + local_layer=0, + v_sf=v_sf[0], + page_ids=page_ids, + ) is not None + assert ( + _get_cutile_v_packed_cache( + metadata, + 1, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + local_layer=1, + v_sf=v_sf[1], + page_ids=page_ids, + ) + is None + ) + + _maybe_update_cutile_v_packed_cache( + metadata, + 1, + kv_cache, + page_ids, + v_head_dim=v_head_dim, + page_size=page_size, + local_layer=1, + v_sf=v_sf[1], + ) + torch.cuda.synchronize() + assert metadata._fp4_mla_attention_v_packed_buf.data_ptr() == shared_ptr + assert not hasattr(metadata, "_fp4_mla_attention_v_packed_buf_l1") + assert ( + _get_cutile_v_packed_cache( + metadata, + 0, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + local_layer=0, + v_sf=v_sf[0], + page_ids=page_ids, + ) + is None + ) + assert ( + _get_cutile_v_packed_cache( + metadata, + 1, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + local_layer=0, + v_sf=v_sf[0], + page_ids=page_ids, + ) + is None + ) + assert _get_cutile_v_packed_cache( + metadata, + 1, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + local_layer=1, + v_sf=v_sf[1], + page_ids=page_ids, + ) is not None + + @pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") def test_fp4_mla_attention_decode_cutile_linear_mtp_matches_reference(monkeypatch): """CuTile linear MTP rows must use per-query causal KV lengths.""" @@ -1422,6 +1571,23 @@ def test_fp4_mla_attention_decode_cutile_linear_mtp_matches_reference(monkeypatc ) +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +def test_fp4_mla_attention_decode_cutile_grouped_mtp_matches_reference(monkeypatch): + """CuTile grouped page-stats must mask MTP future tokens on the final page.""" + monkeypatch.setenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1") + monkeypatch.setenv("TRTLLM_FP4_MLA_SHARE_V_PACK_STORAGE", "1") + monkeypatch.setenv("TRTLLM_FP4_MLA_GROUP_PAGES", "8") + _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + backend="cutile", + num_heads=128, + seq_lens=[9 * FP4_MLA_TOKENS_PER_BLOCK, 9 * FP4_MLA_TOKENS_PER_BLOCK], + seed=23, + check_probs=False, + query_len_per_seq=4, + ) + + @pytest.mark.skipif( os.environ.get("TRTLLM_RUN_FP4_MLA_ATTENTION_BENCHMARK") != "1", reason=("Manual perf benchmark; set TRTLLM_RUN_FP4_MLA_ATTENTION_BENCHMARK=1 to run"), From ebd1347dc1b6009fef6add46c264606b90bd850b Mon Sep 17 00:00:00 2001 From: Tracin <10434017+Tracin@users.noreply.github.com> Date: Tue, 9 Jun 2026 19:01:25 -0700 Subject: [PATCH 08/11] Remove debug codes. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> Triton kernel prepack V. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> --- bench_fp4_mla_decode.py | 6 +- .../_torch/attention_backend/flashinfer.py | 167 +- .../_torch/attention_backend/fp4_mla.py | 2784 +++++++------- .../attention_backend/fp4_mla_cutile.py.bak | 3239 ----------------- .../attention_backend/fp4_mla_kernels.py | 9 +- .../attention_backend/fp4_mla_triton.py | 1728 ++++++--- .../_torch/attention_backend/trtllm.py | 6 +- .../_torch/pyexecutor/model_engine.py | 14 +- tensorrt_llm/_torch/speculative/mtp.py | 7 +- 9 files changed, 2554 insertions(+), 5406 deletions(-) delete mode 100644 tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py.bak diff --git a/bench_fp4_mla_decode.py b/bench_fp4_mla_decode.py index ffd4360362ee..99c18642bc79 100644 --- a/bench_fp4_mla_decode.py +++ b/bench_fp4_mla_decode.py @@ -39,7 +39,7 @@ ) -def _bench(fn, warmup=0, iters=1): +def _bench(fn, warmup=10, iters=50): for _ in range(warmup): fn() torch.cuda.synchronize() @@ -261,8 +261,8 @@ def main(): choices=BACKEND_CHOICES, help=("Backend to benchmark; default runs all fast backends."), ) - p.add_argument("--warmup", type=int, default=0) - p.add_argument("--iters", type=int, default=1) + p.add_argument("--warmup", type=int, default=10) + p.add_argument("--iters", type=int, default=50) args = p.parse_args() batches = args.batch if args.batch else [16, 30, 60, 120, 200, 300] diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index a1508e69b87b..891880e9c717 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -24,7 +24,7 @@ get_fp4_mla_decode_cache, is_flashinfer_fp4_mla_attention_enabled, run_fp4_mla_attention_decode, scatter_fp4_mla_kv_cache, - update_hp_kv_for_fp4_mla, update_page_stage_for_fp4_mla) + update_hp_kv_for_fp4_mla) from .interface import (AttentionBackend, AttentionForwardArgs, AttentionInputType, AttentionMetadata, CustomAttentionMask, MLAParams, PredefinedAttentionMask, @@ -186,23 +186,6 @@ class FlashInferAttentionMetadata(AttentionMetadata): default=None) fp4_mla_hp_snapshot_pool: Optional[torch.Tensor] = field(init=False, default=None) - # BF16 staging buffer for the per-page dynamic-scale no-dequant path, - # shape [max_num_sequences, num_local_layers, kv_factor=1, - # page_size * head_dim]. Unlike high_precision_kv_pool (which buffers only - # the trailing 16-token NVFP4 V-block), this holds the whole in-progress - # page so the active page can be re-quantized to FP4 with the exact - # per-page amax / global scale each step. Allocated only when the - # FlashInfer no-dequant FP4 MLA path is active. - fp4_mla_page_stage_pool: Optional[torch.Tensor] = field(init=False, - default=None) - fp4_mla_page_stage_snapshot_pool: Optional[torch.Tensor] = field( - init=False, default=None) - # Per-page (K/V shared) FP4 global scale for the dynamic-scale path, shape - # [num_local_layers, num_physical_pages] fp32. page_gscale = 448*6/page_amax - # is baked into each page's K and V block scales at re-quant time and undone - # at read. Untouched pages stay 1.0 (== the previous static behaviour). - fp4_mla_page_scale_pool: Optional[torch.Tensor] = field(init=False, - default=None) # Auxiliary FP4 MLA V-scale pool for the no-dequant PV path. The # physical storage is flat per [local_layer, physical_page]; callers view # it with get_fp4_mla_v_scale_pool_view(..., v_head_dim=kv_lora_rank). @@ -218,8 +201,6 @@ class FlashInferAttentionMetadata(AttentionMetadata): default=None) _fp4_mla_attention_denom_buf: Optional[torch.Tensor] = field(init=False, default=None) - # Debug ownership map: seq_slot to last request_id that wrote it. - hp_pool_owners: Optional[dict] = field(init=False, default=None) # True during warmup forward passes (dummy requests, no real data). is_warmup: bool = field(init=False, default=False) @@ -716,7 +697,6 @@ def _allocate_fp4_mla_buffers(self, buffers, capture_graph: bool) -> None: device='cpu', pin_memory=prefer_pinned(), ) - self.hp_pool_owners = {} num_local_layers = self.kv_cache_manager.num_local_layers head_dim = self.kv_cache_manager.head_dim @@ -779,56 +759,6 @@ def _allocate_fp4_mla_buffers(self, buffers, capture_graph: bool) -> None: "FP4 MLA attention requires the C++ KV cache manager to " "allocate the V-scale pool.") - # Per-page dynamic-scale staging buffer: holds the in-progress page - # in BF16 so it can be re-quantized to FP4 with the exact per-page - # amax. Sized to one full page (tokens_per_block) instead of the - # HP pool's 16-token NVFP4 V-block. - page_stage_slots = self.kv_cache_manager.tokens_per_block - stage_pool_shape = [ - max_num_sequences, num_local_layers, kv_factor, - page_stage_slots * head_dim - ] - existing_stage_pool = self.fp4_mla_page_stage_pool - if (capture_graph and existing_stage_pool is not None - and existing_stage_pool.dtype == torch.bfloat16 - and existing_stage_pool.device.type == "cuda" - and len(existing_stage_pool.shape) == len(stage_pool_shape) - and all(existing_stage_pool.shape[idx] >= dim - for idx, dim in enumerate(stage_pool_shape))): - # Persistent seq-slot state: CUDA graph metadata shares it - # rather than reserving a fresh pool per captured graph. - self.fp4_mla_page_stage_pool = existing_stage_pool - else: - self.fp4_mla_page_stage_pool = self.get_empty( - buffers, - stage_pool_shape, - cache_name="fp4_mla_page_stage_pool", - dtype=torch.bfloat16, - capture_graph=capture_graph, - ) - if capture_graph: - self.fp4_mla_page_stage_snapshot_pool = self.get_empty( - buffers, - stage_pool_shape, - cache_name="fp4_mla_page_stage_snapshot_pool", - dtype=torch.bfloat16, - capture_graph=capture_graph, - ) - else: - self.fp4_mla_page_stage_snapshot_pool = None - - # Per-page (K/V shared) FP4 global scale, one fp32 per physical - # page per local layer. Initialised to 1.0 so untouched pages and - # the read path match the previous static-scale behaviour. - self.fp4_mla_page_scale_pool = self.get_empty( - buffers, - [num_local_layers, max_num_pages], - cache_name="fp4_mla_page_scale_pool", - dtype=torch.float32, - capture_graph=capture_graph, - ) - self.fp4_mla_page_scale_pool.fill_(1.0) - # Runtime-alias backing buffers: GPU for kv/prompt lens, CPU pinned for # the helper's prompt_lens_cpu read. self._kv_lens_cuda_buf = self.get_empty( @@ -941,87 +871,6 @@ def _populate_fp4_mla_batch_indices_positions(self) -> None: non_blocking=True) self._positions[:self.num_tokens].copy_(positions, non_blocking=True) - def repage_fp4_mla_decode_from_kv_lens(self) -> None: - """Rebuild the read-side decode paging from the corrected kv_lens. - - The overlap scheduler builds the generation metadata from the - all-draft-accepted over-estimate; ``_preprocess_inputs`` then corrects - ``kv_lens_cuda_runtime``. ``prepare()`` derived ``num_blocks`` / - ``num_generation_blocks`` / ``paged_kv_indptr_decode`` / - ``paged_kv_indices`` / ``paged_kv_last_page_len`` from the over-estimate. - This recomputes the generation slice of that paging from the corrected - kv_lens so the decode kernels that index the page table see lengths - consistent with what attention actually masks to. Opt-in - (``TRTLLM_FP4_MLA_OVERLAP_REPAGE``); host-syncs and is skipped under - CUDA-graph capture. - """ - if os.getenv("TRTLLM_FP4_MLA_OVERLAP_REPAGE", - "0").lower() not in ("1", "true", "yes", "on"): - return - if self.high_precision_kv_pool is None or self.kv_lens_cuda_runtime is None: - return - # num_blocks / num_generation_blocks are python scalars consumed when - # shaping kernel launches; mutating them cannot affect an already - # captured graph, so skip (and avoid the per-step host sync) under - # CUDA graphs and capture. - if getattr(self, "is_cuda_graph", False): - return - if torch.cuda.is_current_stream_capturing(): - return - num_contexts = self.num_contexts - num_seqs = self.num_contexts + self.num_generations - num_gen = num_seqs - num_contexts - if num_gen <= 0 or not self.num_blocks: - return - - page_size = self.page_size - kv_lens_gen = self.kv_lens_cuda_runtime[num_contexts:num_seqs].detach( - ).to("cpu", torch.int64).tolist() - new_blocks_gen = [(kv + page_size - 1) // page_size - for kv in kv_lens_gen] - old_blocks_gen = [ - int(b) for b in self.num_blocks[num_contexts:num_seqs] - ] - if new_blocks_gen == old_blocks_gen: - return # No page-boundary over-estimate this step. - - # Compact the generation page-id slice: per seq keep its first new_b - # block ids (corrected kv_len <= over-estimate => new_b <= old_b). - ctx_blocks = self.num_context_blocks - old_gen_total = sum(old_blocks_gen) - old_gen = self._paged_kv_indices[ctx_blocks:ctx_blocks + - old_gen_total].detach().to( - "cpu", torch.int64).tolist() - new_gen: list[int] = [] - off = 0 - for old_b, new_b in zip(old_blocks_gen, new_blocks_gen): - nb = min(new_b, old_b) - new_gen.extend(old_gen[off:off + nb]) - off += old_b - if new_gen: - self._paged_kv_indices[ctx_blocks:ctx_blocks + len(new_gen)].copy_( - torch.tensor(new_gen, dtype=self._paged_kv_indices.dtype), - non_blocking=False) - - for i, new_b in enumerate(new_blocks_gen): - self.num_blocks[num_contexts + i] = new_b - self.num_generation_blocks = sum(new_blocks_gen) - - indptr = [0] - for b in new_blocks_gen: - indptr.append(indptr[-1] + b) - self.paged_kv_indptr_decode[:len(indptr)].copy_(torch.tensor( - indptr, dtype=self.paged_kv_indptr_decode.dtype), - non_blocking=False) - - last_page = [ - kv - (b - 1) * page_size - for kv, b in zip(kv_lens_gen, new_blocks_gen) - ] - self._paged_kv_last_page_len[num_contexts:num_seqs].copy_( - torch.tensor(last_page, dtype=self._paged_kv_last_page_len.dtype), - non_blocking=False) - def update_for_spec_dec(self) -> None: if self.high_precision_kv_pool is None: return @@ -1913,10 +1762,6 @@ def _mla_forward_context( latent_cache, self._local_layer_idx(metadata), phase="context") - update_page_stage_for_fp4_mla(metadata, - latent_cache, - self._local_layer_idx(metadata), - phase="context") else: ckv_cache, kpe_cache = self._get_mla_caches(metadata) @@ -2012,11 +1857,6 @@ def _mla_forward_generation( latent_cache, self._local_layer_idx(metadata), phase="generation") - update_page_stage_for_fp4_mla( - metadata, - latent_cache, - self._local_layer_idx(metadata), - phase="generation") else: scatter_fp4_mla_kv_cache( metadata, @@ -2028,11 +1868,6 @@ def _mla_forward_generation( latent_cache, self._local_layer_idx(metadata), phase="generation") - update_page_stage_for_fp4_mla( - metadata, - latent_cache, - self._local_layer_idx(metadata), - phase="generation") if not use_fp4_attention: combined_cache = get_fp4_mla_decode_cache( metadata, diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla.py b/tensorrt_llm/_torch/attention_backend/fp4_mla.py index 4cdb7918e54c..fd4453d16959 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla.py @@ -22,8 +22,6 @@ import triton import triton.language as tl -from tensorrt_llm.logger import logger - from .fp4_mla_kernels import ( _fp4_mla_dequant_kernel, _fp4_mla_overlay_hp_tail_kernel, @@ -41,33 +39,18 @@ FP4_MLA_TOKENS_PER_BLOCK: int = 128 FP4_MLA_SCALE_ROW_GROUP: int = 128 FP4_MLA_SCALE_COL_GROUP: int = 4 -FP4_MLA_KV_GLOBAL_SCALE: float = 448.0 * 6.0 / (448.0 * 6.0) +FP4_MLA_KV_GLOBAL_SCALE: float = 448.0 * 6.0 / 448 * 6.0 FP4_MLA_P_GLOBAL_SCALE: float = 448.0 * 6.0 +# Max finite e4m3 magnitude for FP4 MLA block-scale clamping. +FP4_MLA_E4M3_MAX: float = 448.0 FP4_MLA_Q_RESIDUAL_DIM: int = 64 FLASHINFER_FP4_MLA_ATTENTION_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION" FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION_BACKEND" -FLASHINFER_FP4_MLA_DEBUG_ENV = "TRTLLM_FLASHINFER_FP4_MLA_DEBUG" -# Opt-in (triton only): re-quantize the active decode page each step -# from the BF16 staging buffer with an exact per-page (K/V shared) global scale, -# instead of the static FP4_MLA_KV_GLOBAL_SCALE. Off by default so the existing -# static-scale path is unchanged until validated. -FP4_MLA_PER_PAGE_SCALE_ENV = "TRTLLM_FP4_MLA_PER_PAGE_SCALE" -# Diagnostic (overlap + MTP debugging): assert that the read-side decode paging -# (num_blocks / num_generation_blocks / paged_kv_indptr_decode) and the gen -# token positions are consistent with the (corrected) kv_lens_cuda_runtime. -FP4_MLA_DEBUG_ASSERT_ENV = "TRTLLM_FP4_MLA_DEBUG_ASSERT" -# Fix (opt-in): rebuild the read-side decode paging from the corrected kv_lens -# in _preprocess_inputs after the overlap kv_lens correction. -FP4_MLA_OVERLAP_REPAGE_ENV = "TRTLLM_FP4_MLA_OVERLAP_REPAGE" _HPUpdatePhase = Literal["all", "context", "generation"] _FP4_MLA_MTP_HP_SNAPSHOTS = "_fp4_mla_mtp_hp_snapshots" -# Separate MTP snapshot store for the per-page dynamic-scale staging buffer -# (fp4_mla_page_stage_pool). Kept distinct from the HP-pool snapshots so the -# two rollback paths never alias. -_FP4_MLA_MTP_STAGE_SNAPSHOTS = "_fp4_mla_mtp_stage_snapshots" -# Environment and debug helpers +# Environment helpers def _env_enabled(name: str) -> bool: @@ -79,6 +62,18 @@ def _env_enabled(name: str) -> bool: ) +def _env_enabled_default(name: str, default: bool) -> bool: + value = os.getenv(name) + if value is None or value == "": + return default + return value.lower() in ( + "1", + "true", + "yes", + "on", + ) + + def _env_int(name: str) -> Optional[int]: value = os.environ.get(name) if value is None or value == "": @@ -95,67 +90,21 @@ def _fp4_mla_attention_backend() -> str: return os.getenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, "triton").lower() -def fp4_mla_per_page_scale_enabled() -> bool: - """Return whether the per-page dynamic FP4 global-scale path is on. - - Only the ``triton`` backend implements the matching read side, so the - per-page store + dynamic Q scale are gated to it; on any other backend the - flag is ignored (the static FP4_MLA_KV_GLOBAL_SCALE path runs unchanged). - """ - return _env_enabled(FP4_MLA_PER_PAGE_SCALE_ENV) and _fp4_mla_attention_backend() == "triton" - - -def _fp4_mla_debug_enabled() -> bool: - return _env_enabled(FLASHINFER_FP4_MLA_DEBUG_ENV) - - -def _fp4_mla_debug(message: str) -> None: - if _fp4_mla_debug_enabled(): - print(f"[fp4_mla_debug] {message}", flush=True) - - -def _tensor_layout(tensor: Optional[torch.Tensor]) -> str: - if tensor is None: - return "None" - return ( - f"shape={list(tensor.shape)} stride={list(tensor.stride())} " - f"dtype={tensor.dtype} device={tensor.device}" - ) - - -def _debug_tensor_range(name: str, tensor: Optional[torch.Tensor]) -> None: - if not _fp4_mla_debug_enabled(): - return - if tensor is None: - _fp4_mla_debug(f"{name}: None") - return - flat = tensor.detach().reshape(-1) - if flat.numel() == 0: - _fp4_mla_debug(f"{name}: empty {_tensor_layout(tensor)}") - return - try: - first = flat[: min(8, flat.numel())].cpu().tolist() - _fp4_mla_debug( - f"{name}: {_tensor_layout(tensor)} n={flat.numel()} " - f"min={flat.min().item()} max={flat.max().item()} first={first}" - ) - except RuntimeError as exc: - _fp4_mla_debug(f"{name}: failed to read range: {exc}") +def _ceil_div(lhs: int, rhs: int) -> int: + return (lhs + rhs - 1) // rhs -def _debug_sync(label: str) -> None: - if not _fp4_mla_debug_enabled(): - return - if torch.cuda.is_current_stream_capturing(): - _fp4_mla_debug(f"{label}: skip sync during CUDA graph capture") - return - _fp4_mla_debug(f"{label}: synchronize") - torch.cuda.synchronize() - _fp4_mla_debug(f"{label}: sync complete") +_SM_COUNT_CACHE: dict[int, int] = {} -def _ceil_div(lhs: int, rhs: int) -> int: - return (lhs + rhs - 1) // rhs +def _get_sm_count(device: torch.device) -> int: + """Return the SM (multiprocessor) count for ``device``, cached per index.""" + index = device.index if device.index is not None else torch.cuda.current_device() + count = _SM_COUNT_CACHE.get(index) + if count is None: + count = torch.cuda.get_device_properties(index).multi_processor_count + _SM_COUNT_CACHE[index] = count + return count def _host_int_list_during_forward(value: Any, start: int, end: int) -> Optional[list[int]]: @@ -364,7 +313,6 @@ def _scatter_fp4_mla_kv_cache_2d_context( SF_PER_TOKEN=sf_per_token, SF_PER_PAGE=sf_per_page, ) - _debug_sync("scatter_fp4_mla_kv_cache_2d_context") def _scatter_fp4_mla_kv_cache_2d_generation( @@ -470,161 +418,6 @@ def _scatter_fp4_mla_kv_cache_2d_generation( SF_PER_TOKEN=sf_per_token, SF_PER_PAGE=sf_per_page, ) - _debug_sync("scatter_fp4_mla_kv_cache_2d_generation") - - -def _check_fp4_mla_page_generation_single_page( - metadata: Any, - num_contexts: int, - num_seqs: int, - page_size: int, -) -> None: - """Raise if a step writes more than one generation token per sequence. - - The per-page re-quant v1 re-quantizes a single active page per step from the - staging buffer (which holds exactly one page). A linear-MTP draft of length - > 1 can straddle a page boundary, which would require the previous - (completing) page to be re-quantized too -- not yet supported. 1-token - decode never crosses, so only MTP drafts are rejected. Detectable only in - eager mode (host metadata available); under CUDA graph the captured shape - is assumed to be 1-token decode, so do not enable the per-page path together - with MTP + CUDA graph until the multi-page store lands. - """ - gen_token_lens = _host_int_list_during_forward( - getattr(metadata, "prompt_lens_cpu_runtime", None), num_contexts, num_seqs - ) - if gen_token_lens is None: - return - max_gen_len = max(gen_token_lens) if gen_token_lens else 1 - if max_gen_len > 1: - raise NotImplementedError( - "FP4 MLA per-page dynamic scale (v1) supports 1-token decode only; " - f"got a generation length of {max_gen_len} (linear MTP). Multi-page " - "re-quant for MTP drafts is a follow-up." - ) - - -def _store_fp4_mla_page_dynamic_generation( - metadata: Any, - latent_cache: torch.Tensor, - kv_cache: torch.Tensor, - sf_cache: torch.Tensor, - v_sf: torch.Tensor, - page_scale_pool: torch.Tensor, - stage_pool: torch.Tensor, - *, - local_layer: int, - v_head_dim: int, - head_dim: int, - sf_per_token: int, - sf_per_page: int, -) -> None: - """Re-quantize the active decode page with an exact per-page global scale. - - Two passes (see ``fp4_mla_triton``): Pass A computes the shared K/V - page amax -> ``page_gscale`` into ``page_scale_pool``; Pass B re-quantizes - every FP4 tile of the active page from ``stage_pool`` (old tokens) + - ``latent_cache`` (new tokens), baking ``page_gscale`` into the K and V block - scales. Replaces the static 16-token tile scatter for the ``triton`` - path when the per-page scale is enabled. - """ - from .fp4_mla_triton import _fp4_mla_page_requant_gen_kernel, _fp4_mla_page_scale_gen_kernel - - num_contexts = metadata.num_contexts - num_seqs = metadata.num_seqs - num_gen = num_seqs - num_contexts - if num_gen <= 0: - return - - page_size = metadata.page_size - _check_fp4_mla_page_generation_single_page(metadata, num_contexts, num_seqs, page_size) - - pool_head_dim = stage_pool.shape[-1] // page_size - if pool_head_dim < head_dim: - raise RuntimeError( - f"FP4 MLA staging pool head dim {pool_head_dim} < latent head_dim {head_dim}." - ) - - page_ids = metadata.paged_kv_indices[metadata.num_context_blocks :] - seq_slots = metadata.seq_slots[num_contexts:num_seqs] - kv_lens = metadata.kv_lens_cuda_runtime[num_contexts:num_seqs] - gen_lens = metadata.prompt_lens_cuda_runtime[num_contexts:num_seqs] - indptr = metadata.paged_kv_indptr_decode - num_dim_blocks = triton.cdiv(head_dim, FP4_BLOCK_SIZE) - tiles_per_page = page_size // FP4_BLOCK_SIZE - block_d = triton.next_power_of_2(head_dim) - - # Pass A: per-page (K/V shared) amax -> page_gscale. - _fp4_mla_page_scale_gen_kernel[(num_gen,)]( - page_scale_pool, - stage_pool, - latent_cache, - seq_slots, - kv_lens, - gen_lens, - page_ids, - indptr, - page_ids.shape[0], - indptr.shape[0], - stage_pool.shape[0], - kv_cache.shape[0], - page_scale_pool.shape[0], - local_layer, - page_size, - page_scale_pool.stride(0), - stage_pool.stride(0), - stage_pool.stride(1), - latent_cache.stride(0), - latent_cache.stride(1), - HEAD_D=head_dim, - POOL_HEAD_D=pool_head_dim, - FP4_BLOCK=FP4_BLOCK_SIZE, - PAGE_SLOTS=page_size, - BLOCK_D=block_d, - P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, - ) - _debug_sync("fp4_mla_page_scale_gen") - - # Pass B: re-quantize every tile of the active page with page_gscale. - _fp4_mla_page_requant_gen_kernel[(num_gen, tiles_per_page, num_dim_blocks)]( - kv_cache, - sf_cache, - v_sf, - stage_pool, - latent_cache, - page_scale_pool, - seq_slots, - kv_lens, - gen_lens, - page_ids, - indptr, - page_ids.shape[0], - indptr.shape[0], - stage_pool.shape[0], - kv_cache.shape[0], - page_scale_pool.shape[0], - local_layer, - page_size, - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - sf_cache.stride(0), - stage_pool.stride(0), - stage_pool.stride(1), - latent_cache.stride(0), - latent_cache.stride(1), - v_sf.stride(0), - v_sf.stride(1), - page_scale_pool.stride(0), - HEAD_D=head_dim, - POOL_HEAD_D=pool_head_dim, - V_HEAD_D=v_head_dim, - PAGE_SLOTS=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - SF_PER_TOKEN=sf_per_token, - SF_PER_PAGE=sf_per_page, - ) - _debug_sync("fp4_mla_page_requant_gen") def _scatter_fp4_mla_kv_cache_1d( @@ -650,27 +443,6 @@ def _scatter_fp4_mla_kv_cache_1d( block_packed_dim = triton.next_power_of_2(packed_dim) block_sf = triton.next_power_of_2(sf_per_token) - _fp4_mla_debug( - "scatter launch: " - f"num_tokens={num_tokens} token_offset={token_offset} " - f"page_size={metadata.page_size} layer_idx={layer_idx} " - f"head_dim={head_dim} packed_dim={packed_dim} " - f"sf_per_token={sf_per_token} use_swizzled_sf={use_swizzled_sf}" - ) - _fp4_mla_debug(f"scatter latent_cache: {_tensor_layout(latent_cache)}") - _fp4_mla_debug(f"scatter kv_cache: {_tensor_layout(kv_cache)}") - _fp4_mla_debug(f"scatter sf_cache: {_tensor_layout(sf_cache)}") - _debug_tensor_range( - "scatter batch_indices", - metadata.batch_indices[token_offset : token_offset + num_tokens], - ) - _debug_tensor_range( - "scatter positions", - metadata.positions[token_offset : token_offset + num_tokens], - ) - _debug_tensor_range("scatter paged_kv_indices", metadata.paged_kv_indices) - _debug_tensor_range("scatter paged_kv_indptr", metadata.paged_kv_indptr) - _fp4_mla_scatter_kernel[(num_tokens,)]( kv_cache, sf_cache, @@ -705,7 +477,6 @@ def _scatter_fp4_mla_kv_cache_1d( BLOCK_SF=block_sf, USE_SWIZZLED_SF=use_swizzled_sf, ) - _debug_sync("scatter_fp4_mla_kv_cache") # Public cache update and decode entry points @@ -796,16 +567,6 @@ def scatter_fp4_mla_kv_cache( v_sf = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=v_head_dim) num_dim_blocks = triton.cdiv(head_dim, FP4_BLOCK_SIZE) sf_per_page = metadata.page_size // HP_BLOCK_SIZE - _fp4_mla_debug( - "scatter 2d launch: " - f"phase={phase} num_tokens={num_tokens} " - f"token_offset={token_offset} layer_idx={layer_idx} " - f"local_layer={local_layer} head_dim={head_dim} " - f"v_head_dim={v_head_dim} num_dim_blocks={num_dim_blocks}" - ) - _fp4_mla_debug(f"scatter 2d kv_cache: {_tensor_layout(kv_cache)}") - _fp4_mla_debug(f"scatter 2d sf_cache: {_tensor_layout(sf_cache)}") - _fp4_mla_debug(f"scatter 2d v_sf: {_tensor_layout(v_sf)}") if phase == "context": _scatter_fp4_mla_kv_cache_2d_context( @@ -825,47 +586,21 @@ def scatter_fp4_mla_kv_cache( sf_per_page=sf_per_page, ) else: - page_scale_pool = getattr(metadata, "fp4_mla_page_scale_pool", None) - stage_pool = getattr(metadata, "fp4_mla_page_stage_pool", None) - if ( - fp4_mla_per_page_scale_enabled() - and page_scale_pool is not None - and stage_pool is not None - ): - # Per-page dynamic scale: re-quantize the active page from the - # BF16 staging buffer with its exact per-page global scale. - # NOTE: the staging buffer must hold this step's *pre-update* - # tokens, so update_page_stage_for_fp4_mla must run AFTER this. - _store_fp4_mla_page_dynamic_generation( - metadata, - latent_cache, - kv_cache, - sf_cache, - v_sf, - page_scale_pool, - stage_pool, - local_layer=local_layer, - v_head_dim=v_head_dim, - head_dim=head_dim, - sf_per_token=sf_per_token, - sf_per_page=sf_per_page, - ) - else: - _scatter_fp4_mla_kv_cache_2d_generation( - metadata, - latent_cache, - kv_cache, - sf_cache, - v_sf, - global_scale, - local_layer=local_layer, - v_head_dim=v_head_dim, - head_dim=head_dim, - num_tokens=num_tokens, - num_dim_blocks=num_dim_blocks, - sf_per_token=sf_per_token, - sf_per_page=sf_per_page, - ) + _scatter_fp4_mla_kv_cache_2d_generation( + metadata, + latent_cache, + kv_cache, + sf_cache, + v_sf, + global_scale, + local_layer=local_layer, + v_head_dim=v_head_dim, + head_dim=head_dim, + num_tokens=num_tokens, + num_dim_blocks=num_dim_blocks, + sf_per_token=sf_per_token, + sf_per_page=sf_per_page, + ) if phase == "context": v_pack_page_ids = metadata.paged_kv_indices else: @@ -883,6 +618,17 @@ def scatter_fp4_mla_kv_cache( local_layer=local_layer, v_sf=v_sf[local_layer], ) + _maybe_update_triton_v_packed_cache( + metadata, + layer_idx, + kv_cache, + v_pack_page_ids, + num_queries=num_tokens, + v_head_dim=v_head_dim, + page_size=metadata.page_size, + local_layer=local_layer, + v_sf=v_sf[local_layer], + ) return _scatter_fp4_mla_kv_cache_1d( @@ -959,118 +705,6 @@ def _get_decode_src_page_ids(metadata: Any, num_blocks: int) -> torch.Tensor: return src_page_ids -def _assert_fp4_mla_decode_paging_consistent( - metadata: Any, - kv_lens: torch.Tensor, - num_gen_blocks: int, - query_len_per_seq: int, -) -> None: - """Flag read-side decode paging that diverges from the corrected kv_lens. - - With the overlap scheduler the generation metadata is first built from the - all-draft-accepted over-estimate; ``kv_lens_cuda_runtime`` is then corrected - in ``_preprocess_inputs`` and ``positions`` / ``batch_indices`` rebuilt. The - page-table side -- ``num_blocks`` / ``num_generation_blocks`` / - ``paged_kv_indptr_decode`` and the ``positions`` of the new tokens -- is what - the decode kernels index with. This check raises on the first decode where - any of those is inconsistent with the corrected kv_lens, so we can tell - whether stale paging (rather than masking) corrupts the read. Host-syncs; - gated by ``TRTLLM_FP4_MLA_DEBUG_ASSERT`` and skipped under CUDA-graph capture. - """ - if not _env_enabled(FP4_MLA_DEBUG_ASSERT_ENV): - return - if torch.cuda.is_current_stream_capturing(): - return - # Warmup builds synthetic metadata with dummy kv_lens but no real page-table - # allocation (num_generation_blocks==0), so the consistency check does not - # apply there. - if getattr(metadata, "is_warmup", False): - return - num_contexts = metadata.num_contexts - num_seqs = metadata.num_seqs - num_gen = num_seqs - num_contexts - if num_gen <= 0: - return - # Skip the benign context->generation reclassification state: the MTP draft - # loop resets num_contexts to 0 without rebuilding the decode page table - # (for FlashInfer the reorder is gated on enable_flash_mla), leaving - # num_context_blocks > 0 / stale num_generation_blocks. That forward's - # attention is degraded but rejected drafts fall back to the golden token, - # so it does NOT affect output accuracy (it fires with overlap off too). - # Only check pure steady-state generation, where a paging bug would - # genuinely corrupt the committed output. - if int(getattr(metadata, "num_context_blocks", 0)) != 0: - return - # Only the MTP *target* verification forward (query_len_per_seq > 1) - # determines the accepted/committed tokens, so it is the only forward whose - # paging staleness can change output accuracy. The draft-model forwards - # (query_len_per_seq == 1) read a separate KV layer and their bad output is - # rejected (and they mismatch in both overlap modes), so skip them here to - # isolate the accuracy-relevant path. - if query_len_per_seq <= 1: - return - - page_size = metadata.page_size - kv_lens_host = kv_lens.detach().to("cpu", torch.int64).tolist() - expected_blocks = [_ceil_div(kv, page_size) for kv in kv_lens_host] - expected_gen_blocks = sum(expected_blocks) - problems: list[str] = [] - - num_blocks = getattr(metadata, "num_blocks", None) - if num_blocks is not None: - actual_blocks = [int(b) for b in num_blocks[num_contexts:num_seqs]] - if actual_blocks != expected_blocks: - problems.append(f"per-seq num_blocks {actual_blocks} != expected {expected_blocks}") - if num_gen_blocks != expected_gen_blocks: - problems.append(f"num_generation_blocks={num_gen_blocks} != expected {expected_gen_blocks}") - - indptr = getattr(metadata, "paged_kv_indptr_decode", None) - if indptr is not None: - indptr_host = indptr[: num_gen + 1].detach().to("cpu", torch.int64).tolist() - expected_indptr = [0] - for b in expected_blocks: - expected_indptr.append(expected_indptr[-1] + b) - if indptr_host != expected_indptr: - problems.append(f"paged_kv_indptr_decode {indptr_host} != expected {expected_indptr}") - - positions = getattr(metadata, "positions", None) - prompt_lens = getattr(metadata, "prompt_lens_cuda_runtime", None) - if positions is not None and prompt_lens is not None: - num_ctx_tokens = int(getattr(metadata, "num_ctx_tokens", 0)) - gen_pos = positions[num_ctx_tokens:].detach().to("cpu", torch.int64).tolist() - pls = prompt_lens[num_contexts:num_seqs].detach().to("cpu", torch.int64).tolist() - off = 0 - for s, (kv_len, prompt_len) in enumerate(zip(kv_lens_host, pls)): - expected = list(range(kv_len - prompt_len, kv_len)) - got = gen_pos[off : off + prompt_len] - if got != expected: - problems.append( - f"seq{s} gen positions {got} != expected {expected} " - f"(kv_len={kv_len}, prompt_len={prompt_len})" - ) - off += prompt_len - - if problems: - ctx = ( - f"num_contexts={num_contexts}, num_seqs={num_seqs}, " - f"num_generations={getattr(metadata, 'num_generations', '?')}, " - f"query_len_per_seq={query_len_per_seq}, " - f"num_context_blocks={getattr(metadata, 'num_context_blocks', '?')}, " - f"num_blocks={getattr(metadata, 'num_blocks', '?')}, " - f"use_spec_decoding={getattr(metadata, 'use_spec_decoding', '?')}, " - f"is_spec_dec_mode={getattr(metadata, 'is_spec_dec_mode', '?')}, " - f"kv_lens={kv_lens_host}" - ) - msg = ( - "FP4 MLA decode paging inconsistent with corrected kv_lens " - f"({ctx}):\n " + "\n ".join(problems) - ) - if _env_enabled("TRTLLM_FP4_MLA_DEBUG_ASSERT_WARN"): - logger.warning(msg) - return - raise AssertionError(msg) - - def _validate_fp4_mla_cache_shape(page_size: int, head_dim: int) -> None: if page_size != FP4_MLA_TOKENS_PER_BLOCK: raise ValueError( @@ -1183,7 +817,6 @@ def get_fp4_mla_decode_cache( BLOCK_D=block_d, HP_BLOCK=HP_BLOCK_SIZE, ) - return combined @@ -1248,6 +881,24 @@ def _select_cutile_block_v(num_gen_seqs: int, query_len_per_seq: int = 1) -> int return 128 +def _select_triton_block_v(num_queries: int, *, prefer_prepacked_v: bool = False) -> int: + env_block_v = _env_int("TRTLLM_FP4_MLA_BLOCK_V") + if env_block_v is not None: + return env_block_v + if prefer_prepacked_v: + return 128 + return 32 if num_queries <= 32 else 128 + + +def _v_packed_shape( + kv_cache: torch.Tensor, + v_head_dim: int, + page_size: int, + block_v: int, +) -> tuple[int, int]: + return (kv_cache.shape[0] * _ceil_div(v_head_dim, block_v) * block_v, page_size // 2) + + def _cutile_v_packed_attr(layer_idx: int) -> str: if _cutile_shared_v_pack_storage_enabled(): return "_fp4_mla_attention_v_packed_buf" @@ -1268,7 +919,7 @@ def _cutile_v_packed_shape( page_size: int, block_v: int = 128, ) -> tuple[int, int]: - return (kv_cache.shape[0] * _ceil_div(v_head_dim, block_v) * block_v, page_size // 2) + return _v_packed_shape(kv_cache, v_head_dim, page_size, block_v) def _cutile_v_packed_cache_tag( @@ -1397,7 +1048,11 @@ def _maybe_update_cutile_v_packed_cache( return num_gen_seqs = getattr(metadata, "num_seqs", 0) - getattr(metadata, "num_contexts", 0) block_v = _select_cutile_block_v(num_gen_seqs) - if block_v not in (128, 256) or v_head_dim % block_v != 0 or page_size != FP4_MLA_TOKENS_PER_BLOCK: + if ( + block_v not in (128, 256) + or v_head_dim % block_v != 0 + or page_size != FP4_MLA_TOKENS_PER_BLOCK + ): return if page_ids.numel() == 0: return @@ -1473,522 +1128,1159 @@ def _get_cutile_v_packed_cache( return v_packed[: expected_shape[0], : expected_shape[1]] -def _max_generation_pages(metadata: Any) -> int: - num_gen = metadata.num_seqs - metadata.num_contexts - if num_gen <= 0: - return 0 - num_blocks = getattr(metadata, "num_blocks", None) - if num_blocks is not None: - return max(num_blocks[metadata.num_contexts : metadata.num_seqs]) - return metadata.num_generation_blocks - - -def _host_int_list(value: Any, start: int, end: int) -> Optional[list[int]]: - if value is None: - return None - if isinstance(value, torch.Tensor): - if value.is_cuda: - return None - return [int(item) for item in value[start:end].tolist()] - try: - return [int(item) for item in value[start:end]] - except (TypeError, ValueError): - return None - - -def _infer_cutile_assume_full_pages(metadata: Any, max_pages: int, page_size: int) -> bool: - if getattr(metadata, "is_cuda_graph", False): +def _triton_prepack_v_enabled() -> bool: + if _fp4_mla_attention_backend() != "triton": return False + default = _env_enabled_default("TRTLLM_FP4_MLA_PREPACK_V", True) + return _env_enabled_default("TRTLLM_FP4_MLA_TRITON_PREPACK_V", default) - start = metadata.num_contexts - end = metadata.num_seqs - block_counts = _host_int_list(getattr(metadata, "num_blocks", None), start, end) - if block_counts is not None and ( - not block_counts or min(block_counts) != max_pages or max(block_counts) != max_pages - ): + +def _triton_persistent_v_pack_enabled() -> bool: + if not _triton_prepack_v_enabled(): return False + return os.getenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1").lower() not in ( + "0", + "false", + "no", + "off", + ) - kv_lens_cuda = getattr(metadata, "kv_lens_cuda_runtime", None) - if isinstance(kv_lens_cuda, torch.Tensor): - cache_key = ( - start, - end, - max_pages, - page_size, - tuple(block_counts) if block_counts is not None else None, - kv_lens_cuda.data_ptr(), - ) - cache = getattr(metadata, "_fp4_mla_cutile_full_pages_cache", None) - if cache is not None and cache[0] == cache_key: - return bool(cache[1]) - kv_lens = [int(item) for item in kv_lens_cuda[start:end].detach().cpu().tolist()] - result = bool(kv_lens) and min(kv_lens) == max(kv_lens) == max_pages * page_size - setattr(metadata, "_fp4_mla_cutile_full_pages_cache", (cache_key, result)) - return result - kv_cache_params = getattr(metadata, "kv_cache_params", None) - cached_token_lens = _host_int_list( - getattr(kv_cache_params, "num_cached_tokens_per_seq", None), - start, - end, +def _triton_can_prepack_v(v_head_dim: int, page_size: int, block_v: int) -> bool: + return ( + _triton_persistent_v_pack_enabled() + and hasattr(tl, "make_tensor_descriptor") + and block_v in (32, 128) + and v_head_dim % block_v == 0 + and page_size == FP4_MLA_TOKENS_PER_BLOCK ) - seq_lens_kv = _host_int_list(getattr(metadata, "seq_lens_kv", None), start, end) - if cached_token_lens is not None and seq_lens_kv is not None: - if len(cached_token_lens) != len(seq_lens_kv): - return False - kv_lens = [ - cached_len + seq_len for cached_len, seq_len in zip(cached_token_lens, seq_lens_kv) - ] - elif kv_cache_params is None: - kv_lens = _host_int_list(getattr(metadata, "prompt_lens_cpu_runtime", None), start, end) - else: - return False - return bool(kv_lens) and min(kv_lens) == max(kv_lens) == max_pages * page_size +def _triton_v_packed_attr(layer_idx: int) -> str: + if _cutile_shared_v_pack_storage_enabled(): + return "_fp4_mla_triton_attention_v_packed_buf" + return f"_fp4_mla_triton_attention_v_packed_buf_l{layer_idx}" -def _get_linear_mtp_query_len_per_seq( - metadata: Any, - *, - num_queries: int, - num_gen_seqs: int, -) -> int: - """Return the uniform generation query length required by linear MTP.""" - if num_gen_seqs <= 0: - return 1 - start = metadata.num_contexts - end = metadata.num_seqs - query_lens = _host_int_list_during_forward( - getattr(metadata, "prompt_lens_cpu_runtime", None), start, end - ) - if query_lens is None: - query_lens = _host_int_list_during_forward(getattr(metadata, "seq_lens", None), start, end) +def _triton_v_packed_valid_attr(layer_idx: int) -> str: + return f"_fp4_mla_triton_attention_v_packed_valid_l{layer_idx}" - if query_lens is None: - if num_queries % num_gen_seqs != 0: - raise NotImplementedError( - "FP4 MLA linear MTP requires a uniform generation query length; " - f"got {num_queries} query tokens for {num_gen_seqs} sequences." - ) - return num_queries // num_gen_seqs - if sum(query_lens) != num_queries and num_queries == num_gen_seqs: - return 1 - if sum(query_lens) != num_queries: - raise RuntimeError( - "FP4 MLA generation query metadata does not match q shape: " - f"query_lens={query_lens}, total={sum(query_lens)}, " - f"q_tokens={num_queries}." - ) - if not query_lens: - return 1 - if min(query_lens) <= 0: - raise RuntimeError(f"FP4 MLA generation query lengths must be positive, got {query_lens}.") - if min(query_lens) != max(query_lens): - raise NotImplementedError( - "FP4 MLA no-dequant attention currently supports linear MTP with " - f"a uniform generation length per sequence, got {query_lens}." - ) - return query_lens[0] +def _triton_shared_v_packed_valid_attr() -> str: + return "_fp4_mla_triton_attention_v_packed_valid_tag" -def _run_triton_attention_decode( +def _triton_v_packed_cache_tag( + layer_idx: int, + kv_cache: torch.Tensor, *, + v_head_dim: int, + page_size: int, + local_layer: Optional[int] = None, + v_sf: Optional[torch.Tensor] = None, + page_ids: Optional[torch.Tensor] = None, + block_v: int = 128, +) -> tuple[Any, ...]: + return ( + "triton", + _cutile_v_packed_cache_tag( + layer_idx, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=page_ids, + ), + ) + + +def _set_triton_v_packed_cache_valid( metadata: Any, - q_fp4: torch.Tensor, - q_sf: torch.Tensor, + layer_idx: int, kv_cache: torch.Tensor, - sf_cache: torch.Tensor, - v_sf: torch.Tensor, - global_scale: torch.Tensor, - src_page_ids: torch.Tensor, - kv_lens: torch.Tensor, - p_fp4: torch.Tensor, - p_sf: torch.Tensor, - max_scores: torch.Tensor, - denom: torch.Tensor, - output: torch.Tensor, - num_queries: int, - num_heads: int, - head_dim: int, - kv_lora_rank: int, - q_residual_dim: int, - query_len_per_seq: int, - max_pages: int, - sm_scale: float, - q_global_scale: Optional[torch.Tensor] = None, - page_scale_pool: Optional[torch.Tensor] = None, - local_layer: int = 0, - use_per_page_scale: bool = False, + *, + v_head_dim: int, + page_size: int, + local_layer: Optional[int] = None, + v_sf: Optional[torch.Tensor] = None, + page_ids: Optional[torch.Tensor] = None, + block_v: int = 128, ) -> None: - """Dispatch the ``triton`` FP4 MLA decode pipeline. - - Mirrors the four-stage layout used by ``fp4_mla_cutile.py`` - (page-stats with packed P -> reduce-stats -> prob-scale -> PV) but - routes through the self-contained kernels in - ``fp4_mla_triton.py``. Threads through the constexpr assume flags, - TMA descriptors, occupancy/num-warps launch meta, and pipelined PV loop. - """ - from .fp4_mla_triton import _fp4_mla_attention_page_stats_kernel as _attn_page_stats_kernel - from .fp4_mla_triton import _fp4_mla_attention_prob_scale_kernel as _attn_prob_scale_kernel - from .fp4_mla_triton import _fp4_mla_attention_pv_kernel as _attn_pv_kernel - from .fp4_mla_triton import _fp4_mla_attention_pv_reduce_kernel as _attn_pv_reduce_kernel - from .fp4_mla_triton import _fp4_mla_attention_reduce_stats_kernel as _attn_reduce_stats_kernel + valid_attr = ( + _triton_shared_v_packed_valid_attr() + if _cutile_shared_v_pack_storage_enabled() + else _triton_v_packed_valid_attr(layer_idx) + ) + setattr( + metadata, + valid_attr, + _triton_v_packed_cache_tag( + layer_idx, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=page_ids, + ), + ) - block_h = 128 - block_t = metadata.page_size - # Adaptive BLOCK_V: small batches need a finer V split to fill enough waves - # on B200 (~148 SMs). PV grid = num_queries * num_head_blocks(1) * - # (kv_lora_rank / BLOCK_V). We want >= ~2*num_SMs programs so that >1 CTA - # lands per SM and hides the L1TEX scoreboard stalls. Empirically (sweep): - # bs<=32 -> BLOCK_V=32; bs>=64 -> BLOCK_V=128. - # (BLOCK_V=16 is rejected by the V TMA descriptor min-stride requirement.) - block_v = 32 if num_queries <= 32 else 128 - q_head_dim = head_dim + q_residual_dim - # BLOCK_K = 512 matches cutile's "nvt" backend default and aligns the K-window - # with the residual-Q boundary (Q_HEAD_D = 640 = 512 + 128 tail). The - # residual-Q TMA tail path requires Q_RESIDUAL_D == 64 and TAIL_BLOCK_K == 128. - block_k = 512 - full_block_end = (q_head_dim // block_k) * block_k - tail_k = q_head_dim - full_block_end - tail_block_k = 1 << (tail_k - 1).bit_length() if tail_k > 0 else block_k - q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE - k_sf_per_token = head_dim // FP4_BLOCK_SIZE - sf_per_page = metadata.page_size // FP4_BLOCK_SIZE - num_head_blocks = triton.cdiv(num_heads, block_h) - assume_full_heads = num_heads % block_h == 0 - assume_full_v = kv_lora_rank % block_v == 0 - # Match the cutile path: only mark pages "full" when we can prove every - # generation sequence has the same number of cached tokens AND - # query_len_per_seq == 1 (so the kv_len adjustment is a no-op). - assume_full_pages = ( - _infer_cutile_assume_full_pages(metadata, max_pages, metadata.page_size) - and query_len_per_seq == 1 +def _is_triton_v_packed_cache_valid( + metadata: Any, + layer_idx: int, + kv_cache: torch.Tensor, + *, + v_head_dim: int, + page_size: int, + local_layer: Optional[int] = None, + v_sf: Optional[torch.Tensor] = None, + page_ids: Optional[torch.Tensor] = None, + block_v: int = 128, +) -> bool: + valid_attr = ( + _triton_shared_v_packed_valid_attr() + if _cutile_shared_v_pack_storage_enabled() + else _triton_v_packed_valid_attr(layer_idx) + ) + return getattr(metadata, valid_attr, None) == _triton_v_packed_cache_tag( + layer_idx, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=page_ids, ) - # Leave validity checks on. Matches cutile's default and is correctness- - # safe. The perfect-shape PV fast path (tl.ext.make_view + load_view_tko) - # remains gated off — when measured on the TileIR backend (ENABLE_TILE=1) - # it was net-slower on the bench, so the cost of enabling it isn't worth - # the win on the FP4 MLA shapes we care about. - assume_valid_pages = False - # cutile checks only `make_tensor_descriptor`; on the nvt backend the - # presence of TMA descriptors implies `tl.ext.make_view` is available too. - use_tma_data_load = hasattr(triton.language, "make_tensor_descriptor") - # Install the device-side scratch allocator on every call. Triton stores - # the allocator in a ContextVar (triton.runtime._allocation), so a single - # process-wide install is not visible from worker threads / asyncio tasks - # that run with a different Context — the kernel launch would then hit the - # default NullAllocator and raise. Matches the cutile path. - if use_tma_data_load: - def _tma_alloc(size: int, alignment: int, stream): - return torch.empty(size, device=q_fp4.device, dtype=torch.int8) +def _get_triton_v_packed_cache( + metadata: Any, + layer_idx: int, + kv_cache: torch.Tensor, + *, + v_head_dim: int, + page_size: int, + local_layer: Optional[int] = None, + v_sf: Optional[torch.Tensor] = None, + page_ids: Optional[torch.Tensor] = None, + block_v: int = 128, +) -> Optional[torch.Tensor]: + if not _triton_can_prepack_v(v_head_dim, page_size, block_v): + return None + if not _is_triton_v_packed_cache_valid( + metadata, + layer_idx, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=page_ids, + ): + return None + v_packed = getattr(metadata, _triton_v_packed_attr(layer_idx), None) + expected_shape = _v_packed_shape(kv_cache, v_head_dim, page_size, block_v) + if ( + v_packed is None + or v_packed.dtype != torch.uint8 + or v_packed.device != kv_cache.device + or len(v_packed.shape) != 2 + or v_packed.shape[0] < expected_shape[0] + or v_packed.shape[1] < expected_shape[1] + ): + return None + return v_packed[: expected_shape[0], : expected_shape[1]] - triton.set_allocator(_tma_alloc) - # cutile-equivalent launch meta. occupancy=2 lets two CTAs land per SM - # which improves wave-tail efficiency at the bs=32 hot point. - # NOTE: num_stages=2 (instead of the Triton 3.6 default of 3) sidesteps - # the TritonGPUAutomaticWarpSpecialization + NVWSInsertTmemAref pass that - # ICEs on the page_stats kernel under Triton 3.6.0 / sm_100. - launch_meta = {"occupancy": 2} - # PV loop pipelining. With TMA loads, num_stages>=2 lets the next page's - # loads overlap with the current MMA via mbarrier. The PV report shows - # long_scoreboard=4.5 cycles avg on V loads at PV_LOOP_STAGES=2; bumping the - # depth pays off when the grid is small enough that occupancy can absorb - # the extra in-flight tile state — i.e. medium batch / large max_pages. - # Larger pipelines hurt at small batch (more live state, fewer dim blocks). - if num_queries <= 16 or max_pages <= 4: - pv_loop_stages = 2 - else: - pv_loop_stages = 3 +def _update_triton_v_packed_cache( + metadata: Any, + layer_idx: int, + kv_cache: torch.Tensor, + page_ids: torch.Tensor, + *, + v_head_dim: int, + page_size: int, + block_v: int, + local_layer: Optional[int] = None, + v_sf: Optional[torch.Tensor] = None, +) -> Optional[torch.Tensor]: + if not _triton_can_prepack_v(v_head_dim, page_size, block_v): + return None + if page_ids.numel() == 0: + return None + from .fp4_mla_triton import fp4_mla_repack_v_cache - # Page-stats kernel: per (query, head_block, page) program, does QK, - # softmax stats, and packs probs into FP4 with the per-page local-max - # scaling trick. The page-max correction is applied later by - # prob_scale_kernel via p_sf in-place rescaling. - page_stats_shape = (num_queries, max_pages, num_heads) - page_max = _ensure_workspace_tensor( + def _tma_alloc(size: int, alignment: int, stream): + return torch.empty(size, device=kv_cache.device, dtype=torch.int8) + + triton.set_allocator(_tma_alloc) + attr_name = _triton_v_packed_attr(layer_idx) + v_packed = _ensure_workspace_tensor( metadata, - "_fp4_mla_attention_page_max_buf", - page_stats_shape, - dtype=torch.float32, - device=q_fp4.device, + attr_name, + _v_packed_shape(kv_cache, v_head_dim, page_size, block_v), + dtype=torch.uint8, + device=kv_cache.device, ) - page_sum = _ensure_workspace_tensor( + fp4_mla_repack_v_cache( + v_packed, + kv_cache, + page_ids, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + ) + _set_triton_v_packed_cache_valid( metadata, - "_fp4_mla_attention_page_sum_buf", - page_stats_shape, - dtype=torch.float32, - device=q_fp4.device, + layer_idx, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=page_ids, ) + return v_packed - pack_prob_in_page_stats = True - # Per-page scale pointers. When the per-page path is off, pass the static - # global_scale tensor as harmless dummies (the kernel never dereferences - # them under USE_PER_PAGE_SCALE=False) and a zero layer stride. - page_stats_q_gscale = q_global_scale if use_per_page_scale else global_scale - page_stats_page_scale = ( - page_scale_pool if (use_per_page_scale and page_scale_pool is not None) else global_scale - ) - page_stats_pscale_s0 = ( - page_scale_pool.stride(0) if (use_per_page_scale and page_scale_pool is not None) else 0 - ) - _attn_page_stats_kernel[(num_queries, num_head_blocks, max_pages)]( - page_max, - page_sum, - p_fp4, - p_sf, - q_fp4, - q_sf, + +def _maybe_update_triton_v_packed_cache( + metadata: Any, + layer_idx: int, + kv_cache: torch.Tensor, + page_ids: torch.Tensor, + *, + num_queries: int, + v_head_dim: int, + page_size: int, + local_layer: Optional[int] = None, + v_sf: Optional[torch.Tensor] = None, +) -> None: + block_v = _select_triton_block_v(num_queries, prefer_prepacked_v=_triton_prepack_v_enabled()) + _update_triton_v_packed_cache( + metadata, + layer_idx, kv_cache, - sf_cache, - global_scale, - page_stats_q_gscale, - page_stats_page_scale, - src_page_ids, - metadata.paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], - kv_cache.shape[0], - q_fp4.stride(0), - q_fp4.stride(1), - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - sf_cache.stride(0), - page_max.stride(0), - page_max.stride(1), - p_fp4.stride(0), - p_fp4.stride(1), - p_fp4.shape[0], - q_fp4.shape[0], - sm_scale, - local_layer, - page_stats_pscale_s0, - NUM_HEADS=num_heads, - Q_HEAD_D=q_head_dim, - K_HEAD_D=head_dim, - Q_RESIDUAL_D=q_residual_dim, - PAGE_SIZE=metadata.page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - Q_SF_PER_TOKEN=q_sf_per_token, - K_SF_PER_TOKEN=k_sf_per_token, - SF_PER_PAGE=sf_per_page, - P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, - QUERY_LEN_PER_SEQ=query_len_per_seq, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - BLOCK_T=block_t, - BLOCK_K=block_k, - FULL_BLOCK_END=full_block_end, - TAIL_BLOCK_K=tail_block_k, - USE_TMA_DATA_LOAD=use_tma_data_load, - PACK_PROBS=pack_prob_in_page_stats, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - USE_PER_PAGE_SCALE=use_per_page_scale, - **launch_meta, + page_ids, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, ) - _attn_reduce_stats_kernel[(num_queries, num_head_blocks)]( - max_scores, - denom, - page_max, - page_sum, - max_pages, - max_scores.stride(0), - page_max.stride(0), - page_max.stride(1), - NUM_HEADS=num_heads, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - **launch_meta, + +def _max_generation_pages(metadata: Any) -> int: + num_gen = metadata.num_seqs - metadata.num_contexts + if num_gen <= 0: + return 0 + num_blocks = getattr(metadata, "num_blocks", None) + if num_blocks is not None: + return max(num_blocks[metadata.num_contexts : metadata.num_seqs]) + return metadata.num_generation_blocks + + +def _host_int_list(value: Any, start: int, end: int) -> Optional[list[int]]: + if value is None: + return None + if isinstance(value, torch.Tensor): + if value.is_cuda: + return None + return [int(item) for item in value[start:end].tolist()] + try: + return [int(item) for item in value[start:end]] + except (TypeError, ValueError): + return None + + +def _infer_cutile_assume_full_pages(metadata: Any, max_pages: int, page_size: int) -> bool: + if getattr(metadata, "is_cuda_graph", False): + return False + + start = metadata.num_contexts + end = metadata.num_seqs + block_counts = _host_int_list(getattr(metadata, "num_blocks", None), start, end) + if block_counts is not None and ( + not block_counts or min(block_counts) != max_pages or max(block_counts) != max_pages + ): + return False + + kv_lens_cuda = getattr(metadata, "kv_lens_cuda_runtime", None) + if isinstance(kv_lens_cuda, torch.Tensor): + cache_key = ( + start, + end, + max_pages, + page_size, + tuple(block_counts) if block_counts is not None else None, + kv_lens_cuda.data_ptr(), + ) + cache = getattr(metadata, "_fp4_mla_cutile_full_pages_cache", None) + if cache is not None and cache[0] == cache_key: + return bool(cache[1]) + kv_lens = [int(item) for item in kv_lens_cuda[start:end].detach().cpu().tolist()] + result = bool(kv_lens) and min(kv_lens) == max(kv_lens) == max_pages * page_size + setattr(metadata, "_fp4_mla_cutile_full_pages_cache", (cache_key, result)) + return result + + kv_cache_params = getattr(metadata, "kv_cache_params", None) + cached_token_lens = _host_int_list( + getattr(kv_cache_params, "num_cached_tokens_per_seq", None), + start, + end, ) + seq_lens_kv = _host_int_list(getattr(metadata, "seq_lens_kv", None), start, end) + if cached_token_lens is not None and seq_lens_kv is not None: + if len(cached_token_lens) != len(seq_lens_kv): + return False + kv_lens = [ + cached_len + seq_len for cached_len, seq_len in zip(cached_token_lens, seq_lens_kv) + ] + elif kv_cache_params is None: + kv_lens = _host_int_list(getattr(metadata, "prompt_lens_cpu_runtime", None), start, end) + else: + return False - _attn_prob_scale_kernel[(num_queries, num_head_blocks, max_pages)]( - p_sf, - max_scores, - denom, - page_max, - metadata.paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], - max_scores.stride(0), - page_max.stride(0), - page_max.stride(1), - NUM_HEADS=num_heads, - PAGE_SIZE=metadata.page_size, - SF_PER_PAGE=sf_per_page, - QUERY_LEN_PER_SEQ=query_len_per_seq, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, + return bool(kv_lens) and min(kv_lens) == max(kv_lens) == max_pages * page_size + + +def _get_linear_mtp_query_len_per_seq( + metadata: Any, + *, + num_queries: int, + num_gen_seqs: int, +) -> int: + """Return the uniform generation query length required by linear MTP.""" + if num_gen_seqs <= 0: + return 1 + + start = metadata.num_contexts + end = metadata.num_seqs + query_lens = _host_int_list_during_forward( + getattr(metadata, "prompt_lens_cpu_runtime", None), start, end ) + if query_lens is None: + query_lens = _host_int_list_during_forward(getattr(metadata, "seq_lens", None), start, end) - num_dim_blocks = triton.cdiv(kv_lora_rank, block_v) + if query_lens is None: + if num_queries % num_gen_seqs != 0: + raise NotImplementedError( + "FP4 MLA linear MTP requires a uniform generation query length; " + f"got {num_queries} query tokens for {num_gen_seqs} sequences." + ) + return num_queries // num_gen_seqs - # PV page split: partition the page range across additional programs and - # reduce in a follow-up kernel. ncu showed PV at waves/SM=0.49 for bs=32 — - # PV is L1-bandwidth bound, so raising in-flight CTAs is the lever. - # BLOCK_V is bounded below by the 16-byte TMA descriptor min-stride. - # PV page split: ncu shows that with the current shape (bs=32, max_pages=256) - # the PV kernel is L1-cache-throughput bound (long_scoreboard=4.5 cycles - # avg, L1 global LD hit-rate <40%). Increasing the program count via page - # splitting reduced waves/SM idle time but did NOT improve wall-time at - # current shapes — the per-CTA L1 thrash is the limit. Gate the split off - # by default; re-enable only for very small grids where occupancy is the - # bottleneck rather than per-CTA L1 pressure. - page_split = 1 - base_grid = num_queries * num_head_blocks * num_dim_blocks - if max_pages >= 16 and base_grid < 148: - for p in (8, 4, 2): - if max_pages % p == 0 and max_pages // p >= 16 and base_grid * p <= 148 * 4: - page_split = p - break - if page_split > 1: - pages_per_split = max_pages // page_split - partial_out = _ensure_workspace_tensor( - metadata, - "_fp4_mla_attention_pv_partial_buf", - (num_queries, page_split, num_heads, kv_lora_rank), - dtype=torch.float32, - device=q_fp4.device, + if sum(query_lens) != num_queries and num_queries == num_gen_seqs: + return 1 + if sum(query_lens) != num_queries: + raise RuntimeError( + "FP4 MLA generation query metadata does not match q shape: " + f"query_lens={query_lens}, total={sum(query_lens)}, " + f"q_tokens={num_queries}." ) - _attn_pv_kernel[(num_queries, num_head_blocks, num_dim_blocks * page_split)]( - output, - p_fp4, - p_sf, - kv_cache, - v_sf, - global_scale, - src_page_ids, - metadata.paged_kv_indptr_decode, + if not query_lens: + return 1 + if min(query_lens) <= 0: + raise RuntimeError(f"FP4 MLA generation query lengths must be positive, got {query_lens}.") + if min(query_lens) != max(query_lens): + raise NotImplementedError( + "FP4 MLA no-dequant attention currently supports linear MTP with " + f"a uniform generation length per sequence, got {query_lens}." + ) + return query_lens[0] + + +def _run_triton_attention_decode( + *, + metadata: Any, + layer_idx: int, + local_layer: int, + q_fp4: torch.Tensor, + q_sf: torch.Tensor, + kv_cache: torch.Tensor, + sf_cache: torch.Tensor, + v_sf: torch.Tensor, + global_scale: torch.Tensor, + src_page_ids: torch.Tensor, + kv_lens: torch.Tensor, + p_fp4: torch.Tensor, + p_sf: torch.Tensor, + max_scores: torch.Tensor, + denom: torch.Tensor, + output: torch.Tensor, + num_queries: int, + num_heads: int, + head_dim: int, + kv_lora_rank: int, + q_residual_dim: int, + query_len_per_seq: int, + max_pages: int, + sm_scale: float, + q_global_scale: Optional[torch.Tensor] = None, +) -> None: + """Dispatch the ``triton`` FP4 MLA decode pipeline. + + Mirrors the four-stage layout used by ``fp4_mla_cutile.py`` + (page-stats with packed P -> reduce-stats -> prob-scale -> PV) but + routes through the self-contained kernels in + ``fp4_mla_triton.py``. Threads through the constexpr assume flags, + TMA descriptors, occupancy/num-warps launch meta, and pipelined PV loop. + """ + from .fp4_mla_triton import ( + _fp4_mla_attention_group_reduce_stats_kernel as _attn_group_reduce_stats_kernel, + ) + from .fp4_mla_triton import ( + _fp4_mla_attention_page_stats_grouped_kernel as _attn_page_stats_grouped_kernel, + ) + from .fp4_mla_triton import _fp4_mla_attention_page_stats_kernel as _attn_page_stats_kernel + from .fp4_mla_triton import ( + _fp4_mla_attention_page_stats_mtp_kernel as _attn_page_stats_mtp_kernel, + ) + from .fp4_mla_triton import _fp4_mla_attention_prob_scale_kernel as _attn_prob_scale_kernel + from .fp4_mla_triton import _fp4_mla_attention_pv_kernel as _attn_pv_kernel + from .fp4_mla_triton import ( + _fp4_mla_attention_pv_prepacked_v_kernel as _attn_pv_prepacked_v_kernel, + ) + from .fp4_mla_triton import _fp4_mla_attention_pv_reduce_kernel as _attn_pv_reduce_kernel + from .fp4_mla_triton import _fp4_mla_attention_reduce_stats_kernel as _attn_reduce_stats_kernel + + block_h = 128 + block_t = metadata.page_size + # Adaptive BLOCK_V: the fallback PV path uses a finer V split at small batch + # on B200 (~148 SMs). PV grid = num_queries * num_head_blocks(1) * + # (kv_lora_rank / BLOCK_V). We want >= ~2*num_SMs programs so that >1 CTA + # lands per SM and hides the L1TEX scoreboard stalls. Empirically (sweep): + # bs<=32 -> BLOCK_V=32; bs>=64 -> BLOCK_V=128. + # (BLOCK_V=16 is rejected by the V TMA descriptor min-stride requirement.) + # With prepacked V, BLOCK_V=128 avoids reloading the same P tile four times + # and matches the cutile prepacked-V tile shape. + block_v = _select_triton_block_v(num_queries, prefer_prepacked_v=_triton_prepack_v_enabled()) + q_head_dim = head_dim + q_residual_dim + # BLOCK_K = 512 matches cutile's "nvt" backend default and aligns the K-window + # with the residual-Q boundary (Q_HEAD_D = 640 = 512 + 128 tail). The + # residual-Q TMA tail path requires Q_RESIDUAL_D == 64 and TAIL_BLOCK_K == 128. + block_k = 512 + full_block_end = (q_head_dim // block_k) * block_k + tail_k = q_head_dim - full_block_end + tail_block_k = 1 << (tail_k - 1).bit_length() if tail_k > 0 else block_k + q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE + k_sf_per_token = head_dim // FP4_BLOCK_SIZE + sf_per_page = metadata.page_size // FP4_BLOCK_SIZE + num_head_blocks = triton.cdiv(num_heads, block_h) + + assume_full_heads = num_heads % block_h == 0 + assume_full_v = kv_lora_rank % block_v == 0 + # Match the cutile path: only mark pages "full" when we can prove every + # generation sequence has the same number of cached tokens AND + # query_len_per_seq == 1 (so the kv_len adjustment is a no-op). + assume_full_pages = ( + _infer_cutile_assume_full_pages(metadata, max_pages, metadata.page_size) + and query_len_per_seq == 1 + ) + # Leave validity checks on. Matches cutile's default and is correctness- + # safe. The perfect-shape PV fast path (tl.ext.make_view + load_view_tko) + # remains gated off — when measured on the TileIR backend (ENABLE_TILE=1) + # it was net-slower on the bench, so the cost of enabling it isn't worth + # the win on the FP4 MLA shapes we care about. + assume_valid_pages = False + num_gen_seqs = num_queries // query_len_per_seq + if ( + not assume_valid_pages + and assume_full_pages + and src_page_ids.numel() == num_gen_seqs * max_pages + ): + assume_valid_pages = True + # cutile checks only `make_tensor_descriptor`; on the nvt backend the + # presence of TMA descriptors implies `tl.ext.make_view` is available too. + use_tma_data_load = hasattr(triton.language, "make_tensor_descriptor") + + # Install the device-side scratch allocator on every call. Triton stores + # the allocator in a ContextVar (triton.runtime._allocation), so a single + # process-wide install is not visible from worker threads / asyncio tasks + # that run with a different Context — the kernel launch would then hit the + # default NullAllocator and raise. Matches the cutile path. + if use_tma_data_load: + + def _tma_alloc(size: int, alignment: int, stream): + return torch.empty(size, device=q_fp4.device, dtype=torch.int8) + + triton.set_allocator(_tma_alloc) + + # cutile-equivalent launch meta. occupancy=2 lets two CTAs land per SM + # which improves wave-tail efficiency at the bs=32 hot point. + # NOTE: num_stages=2 (instead of the Triton 3.6 default of 3) sidesteps + # the TritonGPUAutomaticWarpSpecialization + NVWSInsertTmemAref pass that + # ICEs on the page_stats kernel under Triton 3.6.0 / sm_100. + launch_meta = {"occupancy": 2} + # The matmul kernels (page-stats QK and PV) are register-limited: at the + # Triton default of num_warps=4 the [BLOCK_H, BLOCK_T] epilogue spills the + # register file down to ~2 CTAs/SM (12.5% occupancy), so there are too few + # warps to hide the QK/PV load latency (ncu: ~0.3 eligible warps/scheduler). + # Spreading the tile epilogue over num_warps=8 halves the per-thread + # register need and roughly doubles resident warps. Matches the cutile + # ("nvt") backend, which launches page-stats at num_warps=8. Both are + # overridable for tuning. + sm_count = _get_sm_count(q_fp4.device) + # page-stats num_warps: the full-pages fast path (uniform q_len==1 decode) + # benefits from num_warps=8 (more warps hide the QK load latency); the + # masked path (q_len>1 / ragged lengths) carries extra per-thread state and + # measured markedly faster at num_warps=4 (e.g. bs256 q_len4: 131->95ms). + page_stats_num_warps = _env_int("TRTLLM_FP4_MLA_PAGE_STATS_NUM_WARPS") + if page_stats_num_warps is None: + page_stats_num_warps = 8 if assume_full_pages else 4 + page_stats_launch_meta = {"occupancy": 2, "num_warps": page_stats_num_warps} + # PV benefits from num_warps=8 across shapes measured. + pv_num_warps = _env_int("TRTLLM_FP4_MLA_PV_NUM_WARPS") or 8 + pv_launch_meta = {"occupancy": 2, "num_warps": pv_num_warps} + # The MTP-fused page-stats kernel holds K live across the q_len row loop and + # so carries more state than the masked one-page kernel; it wants + # num_warps=8 (measured bs256 q_len4: nw4 108ms -> nw8 85ms). + mtp_num_warps = _env_int("TRTLLM_FP4_MLA_MTP_NUM_WARPS") or 8 + mtp_launch_meta = {"occupancy": 2, "num_warps": mtp_num_warps} + # PV loop pipelining. With TMA loads, num_stages>=2 lets the next page's + # loads overlap with the current MMA via mbarrier. The PV report shows + # long_scoreboard=4.5 cycles avg on V loads at PV_LOOP_STAGES=2; bumping the + # depth pays off when the grid is small enough that occupancy can absorb + # the extra in-flight tile state — i.e. medium batch / large max_pages. + # Larger pipelines hurt at small batch (more live state, fewer dim blocks). + if num_queries <= 16 or max_pages <= 4: + pv_loop_stages = 2 + else: + pv_loop_stages = 3 + + # Page-stats kernel: per (query, head_block, page) program, does QK, + # softmax stats, and packs probs into FP4 with the per-page local-max + # scaling trick. The page-max correction is applied later by + # prob_scale_kernel via p_sf in-place rescaling. + page_stats_shape = (num_queries, max_pages, num_heads) + page_max = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_page_max_buf", + page_stats_shape, + dtype=torch.float32, + device=q_fp4.device, + ) + page_sum = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_page_sum_buf", + page_stats_shape, + dtype=torch.float32, + device=q_fp4.device, + ) + + pack_prob_in_page_stats = True + # Q and KV share the same static global scale, so this reduces to the + # global_scale^2 correction. + page_stats_q_gscale = q_global_scale if q_global_scale is not None else global_scale + + # Grouped page-stats: walk multiple pages per CTA so Q (and the TMA + # descriptors) load once and amortize across the group. The one-page-per-CTA + # kernel is work-bound at long context -- it reloads Q for every page and + # pays a per-CTA prologue 16k times -- and ncu shows raising its occupancy + # does not help (no extra warps to fill, the work itself is the cost). + # Grouping cuts both. Outputs stay per-page so every downstream stage is + # unchanged. Restricted to the perfect decode shape the grouped kernel was + # written for; everything else keeps the one-page kernel. + # NOTE: grouped page-stats is OFF by default. It walks multiple pages per + # CTA with Q held live for reuse, but once Q is indexed correctly per query + # the held per-query Q tiles add enough register pressure to drop occupancy, + # and it measured net-slower than the one-page kernel on every shape tested + # (the earlier apparent win came from a since-fixed bug that loaded a + # constant, cacheable Q slice). Kept behind an opt-in flag for future work. + # Shape gate shared by the grouped and MTP-fused page-stats kernels. + standard_page_stats_shape = ( + use_tma_data_load + and pack_prob_in_page_stats + and assume_full_heads + and num_heads == block_h + and num_head_blocks == 1 + and block_k == full_block_end + and block_t == metadata.page_size + and q_residual_dim == FP4_MLA_Q_RESIDUAL_DIM + and q_head_dim - full_block_end == tail_block_k + and tail_block_k == 2 * q_residual_dim + and full_block_end + == (head_dim // FP4_BLOCK_SIZE - q_residual_dim // FP4_BLOCK_SIZE) * FP4_BLOCK_SIZE + and metadata.page_size == FP4_MLA_TOKENS_PER_BLOCK + ) + # MTP-fused page-stats: for linear MTP (query_len_per_seq > 1) the q_len + # query rows of a sequence share the same K, so one CTA per (seq, page) + # loads K once and feeds all q_len QK matmuls -- cutting K reloads q_len-fold + # on the load-latency-bound decode QK. Only the masked path applies here + # (q_len>1 forces assume_full_pages/valid_pages False). + mtp_page_stats_enabled = _env_enabled_default("TRTLLM_FP4_MLA_TRITON_MTP_PAGE_STATS", True) + can_mtp_page_stats = ( + mtp_page_stats_enabled + and query_len_per_seq > 1 + and num_gen_seqs * query_len_per_seq == num_queries + and not assume_full_pages + and not assume_valid_pages + and standard_page_stats_shape + ) + group_page_stats_enabled = _env_enabled_default("TRTLLM_FP4_MLA_TRITON_GROUP_PAGE_STATS", False) + can_group_page_stats = ( + group_page_stats_enabled and not can_mtp_page_stats and standard_page_stats_shape + ) + if can_mtp_page_stats: + _attn_page_stats_mtp_kernel[(num_gen_seqs, num_head_blocks, max_pages)]( + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4, + q_sf, + kv_cache, + sf_cache, + global_scale, + page_stats_q_gscale, + src_page_ids, + metadata.paged_kv_indptr_decode, kv_lens, src_page_ids.shape[0], kv_cache.shape[0], - output.stride(0), - output.stride(1), - output.stride(2), - output.shape[0] * output.shape[1], + q_fp4.stride(0), + q_fp4.stride(1), + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + sf_cache.stride(0), + page_max.stride(0), + page_max.stride(1), p_fp4.stride(0), p_fp4.stride(1), p_fp4.shape[0], + q_fp4.shape[0], + sm_scale, + NUM_HEADS=num_heads, + Q_HEAD_D=q_head_dim, + K_HEAD_D=head_dim, + Q_RESIDUAL_D=q_residual_dim, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, + SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + BLOCK_T=block_t, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + **mtp_launch_meta, + ) + elif can_group_page_stats: + # Each grouped CTA reloads Q once, so total Q reloads == CTA count; we + # want the fewest CTAs that still fill ~one wave, with each CTA walking + # as many pages as possible. Mirrors the cutile decode heuristic + # (group_pages ~ 8 * num_gen, total CTAs ~ max_pages / 8 ~ one wave). + # Over-splitting into more, lighter CTAs both adds Q reloads and risks a + # second, mostly-empty occupancy wave -- measured net-slower. + ps_group_pages_env = _env_int("TRTLLM_FP4_MLA_TRITON_PAGE_STATS_GROUP_PAGES") + ps_group_pages_cap = _env_int("TRTLLM_FP4_MLA_TRITON_PAGE_STATS_GROUP_PAGES_CAP") or 128 + if ps_group_pages_env is not None: + ps_group_pages = max(1, ps_group_pages_env) + else: + # 8 * num_gen mirrors the cutile decode heuristic, but cap it: a + # heavy serial page loop limits this kernel's occupancy, so very + # large groups (e.g. one group spanning every page at big batch) + # collapse to a few mega-CTAs and run several-fold slower. The cap + # keeps per-CTA work bounded and CTA count scaling with batch. + ps_group_pages = min(max(8, 8 * max(num_gen_seqs, 1)), ps_group_pages_cap, max_pages) + ps_num_groups = _ceil_div(max_pages, ps_group_pages) + ps_loop_stages = _env_int("TRTLLM_FP4_MLA_TRITON_PAGE_STATS_STAGES") or 2 + _attn_page_stats_grouped_kernel[(num_queries, num_head_blocks, ps_num_groups)]( + page_max, + page_sum, + p_fp4, + p_sf, + q_fp4, + q_sf, + kv_cache, + sf_cache, + global_scale, + page_stats_q_gscale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + kv_cache.shape[0], + q_fp4.stride(0), + q_fp4.stride(1), kv_cache.stride(0), kv_cache.stride(2), kv_cache.stride(4), - v_sf.stride(0), + sf_cache.stride(0), + page_max.stride(0), + page_max.stride(1), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + q_fp4.shape[0], + sm_scale, NUM_HEADS=num_heads, - V_HEAD_D=kv_lora_rank, + Q_HEAD_D=q_head_dim, + K_HEAD_D=head_dim, + Q_RESIDUAL_D=q_residual_dim, PAGE_SIZE=metadata.page_size, FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, QUERY_LEN_PER_SEQ=query_len_per_seq, MAX_PAGES=max_pages, - P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, BLOCK_H=block_h, - BLOCK_V=block_v, - USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, - USE_TMA_V_LOAD=use_tma_data_load and kv_lora_rank % block_v == 0, - PV_LOOP_STAGES=pv_loop_stages, - ASSUME_FULL_HEADS=assume_full_heads, + BLOCK_T=block_t, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + GROUP_PAGES=ps_group_pages, ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_FULL_V=assume_full_v, ASSUME_VALID_PAGES=assume_valid_pages, - PAGE_SPLIT=page_split, - PAGES_PER_SPLIT=pages_per_split, - PARTIAL_OUT=True, - partial_out_ptr=partial_out, - partial_s0=partial_out.stride(0), - partial_s1=partial_out.stride(1), - partial_s2=partial_out.stride(2), - partial_s3=partial_out.stride(3), - **launch_meta, - ) - _attn_pv_reduce_kernel[(num_queries, num_head_blocks, num_dim_blocks)]( - output, - partial_out, - global_scale, - output.stride(0), - output.stride(1), - output.stride(2), - partial_out.stride(0), - partial_out.stride(1), - partial_out.stride(2), - partial_out.stride(3), - NUM_HEADS=num_heads, - V_HEAD_D=kv_lora_rank, - PAGE_SPLIT=page_split, - P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, - BLOCK_H=block_h, - BLOCK_V=block_v, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_V=assume_full_v, - **launch_meta, + PAGE_LOOP_STAGES=ps_loop_stages, + **page_stats_launch_meta, ) else: - _attn_pv_kernel[(num_queries, num_head_blocks, num_dim_blocks)]( - output, + _attn_page_stats_kernel[(num_queries, num_head_blocks, max_pages)]( + page_max, + page_sum, p_fp4, p_sf, + q_fp4, + q_sf, kv_cache, - v_sf, + sf_cache, global_scale, + page_stats_q_gscale, src_page_ids, metadata.paged_kv_indptr_decode, kv_lens, src_page_ids.shape[0], kv_cache.shape[0], - output.stride(0), - output.stride(1), - output.stride(2), - output.shape[0] * output.shape[1], - p_fp4.stride(0), - p_fp4.stride(1), - p_fp4.shape[0], + q_fp4.stride(0), + q_fp4.stride(1), kv_cache.stride(0), kv_cache.stride(2), kv_cache.stride(4), - v_sf.stride(0), + sf_cache.stride(0), + page_max.stride(0), + page_max.stride(1), + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + q_fp4.shape[0], + sm_scale, NUM_HEADS=num_heads, - V_HEAD_D=kv_lora_rank, + Q_HEAD_D=q_head_dim, + K_HEAD_D=head_dim, + Q_RESIDUAL_D=q_residual_dim, PAGE_SIZE=metadata.page_size, FP4_BLOCK=FP4_BLOCK_SIZE, + Q_SF_PER_TOKEN=q_sf_per_token, + K_SF_PER_TOKEN=k_sf_per_token, SF_PER_PAGE=sf_per_page, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, QUERY_LEN_PER_SEQ=query_len_per_seq, MAX_PAGES=max_pages, - P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, BLOCK_H=block_h, - BLOCK_V=block_v, - USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, - USE_TMA_V_LOAD=use_tma_data_load and kv_lora_rank % block_v == 0, - PV_LOOP_STAGES=pv_loop_stages, + BLOCK_T=block_t, + BLOCK_K=block_k, + FULL_BLOCK_END=full_block_end, + TAIL_BLOCK_K=tail_block_k, + USE_TMA_DATA_LOAD=use_tma_data_load, + PACK_PROBS=pack_prob_in_page_stats, ASSUME_FULL_HEADS=assume_full_heads, ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_FULL_V=assume_full_v, ASSUME_VALID_PAGES=assume_valid_pages, + **page_stats_launch_meta, + ) + # Two-level softmax-stats reduction. The single-level reduce launched only + # (num_queries * num_head_blocks) CTAs, each serially walking all max_pages + # twice -- at small batch that handful of CTAs left the GPU almost idle and + # the reduce cost more than the QK matmul. Level 1 parallelizes the page + # reduction across a page-group axis (online-softmax partials, pipelined); + # level 2 reuses the existing reduce kernel to fold the few groups into the + # global (max, denom). When the (query, head) grid already fills the GPU the + # group count collapses to 1 and this degenerates to the original reduce. + seqhead_ctas = num_queries * num_head_blocks + # Aim for ~3 waves of level-1 CTAs so page loads have enough memory-level + # parallelism to hide latency, while keeping the group count small enough + # that the level-2 combine loop stays short. + target_l1_ctas = 3 * sm_count + num_reduce_groups = _ceil_div(target_l1_ctas, max(seqhead_ctas, 1)) + num_reduce_groups = max(1, min(num_reduce_groups, max_pages, 64)) + if num_reduce_groups <= 1: + _attn_reduce_stats_kernel[(num_queries, num_head_blocks)]( + max_scores, + denom, + page_max, + page_sum, + max_pages, + max_scores.stride(0), + page_max.stride(0), + page_max.stride(1), + NUM_HEADS=num_heads, + MAX_PAGES=max_pages, + BLOCK_H=block_h, **launch_meta, ) - - -def run_fp4_mla_attention_decode( - metadata: Any, - layer_idx: int, - local_layer: int, - q_nope: torch.Tensor, - q_pe: torch.Tensor, - output: torch.Tensor, - *, - sm_scale: float, - kv_lora_rank: int, - qk_rope_head_dim: int, + else: + group_pages = _ceil_div(max_pages, num_reduce_groups) + num_reduce_groups = _ceil_div(max_pages, group_pages) + group_max = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_group_max_buf", + (num_queries, num_reduce_groups, num_heads), + dtype=torch.float32, + device=q_fp4.device, + ) + group_sum = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_group_sum_buf", + (num_queries, num_reduce_groups, num_heads), + dtype=torch.float32, + device=q_fp4.device, + ) + _attn_group_reduce_stats_kernel[(num_queries, num_head_blocks, num_reduce_groups)]( + group_max, + group_sum, + page_max, + page_sum, + max_pages, + group_max.stride(0), + group_max.stride(1), + page_max.stride(0), + page_max.stride(1), + NUM_HEADS=num_heads, + GROUP_PAGES=group_pages, + BLOCK_H=block_h, + PIPELINE_STAGES=min(group_pages, 4), + **launch_meta, + ) + _attn_reduce_stats_kernel[(num_queries, num_head_blocks)]( + max_scores, + denom, + group_max, + group_sum, + num_reduce_groups, + max_scores.stride(0), + group_max.stride(0), + group_max.stride(1), + NUM_HEADS=num_heads, + MAX_PAGES=num_reduce_groups, + BLOCK_H=block_h, + **launch_meta, + ) + _attn_prob_scale_kernel[(num_queries, num_head_blocks, max_pages)]( + p_sf, + max_scores, + denom, + page_max, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + max_scores.stride(0), + page_max.stride(0), + page_max.stride(1), + NUM_HEADS=num_heads, + PAGE_SIZE=metadata.page_size, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + BLOCK_H=block_h, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **launch_meta, + ) + num_dim_blocks = triton.cdiv(kv_lora_rank, block_v) + v_packed = _get_triton_v_packed_cache( + metadata, + layer_idx, + kv_cache, + v_head_dim=kv_lora_rank, + page_size=metadata.page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + page_ids=src_page_ids, + ) + if ( + v_packed is None + and _triton_can_prepack_v(kv_lora_rank, metadata.page_size, block_v) + and not torch.cuda.is_current_stream_capturing() + ): + v_packed = _update_triton_v_packed_cache( + metadata, + layer_idx, + kv_cache, + src_page_ids, + v_head_dim=kv_lora_rank, + page_size=metadata.page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + ) + use_triton_v_packed_cache = v_packed is not None + + # PV page split: partition the page range across additional programs and + # reduce in a follow-up kernel. ncu showed PV at waves/SM=0.49 for bs=32 — + # PV is L1-bandwidth bound, so raising in-flight CTAs is the lever. + # BLOCK_V is bounded below by the 16-byte TMA descriptor min-stride. + # PV page split: ncu shows that with the current shape (bs=32, max_pages=256) + # the PV kernel is L1-cache-throughput bound (long_scoreboard=4.5 cycles + # avg, L1 global LD hit-rate <40%). Increasing the program count via page + # splitting reduced waves/SM idle time but did NOT improve wall-time at + # current shapes — the per-CTA L1 thrash is the limit. Gate the split off + # by default; re-enable only for very small grids where occupancy is the + # bottleneck rather than per-CTA L1 pressure. + page_split = 1 + base_grid = num_queries * num_head_blocks * num_dim_blocks + if max_pages >= 16 and base_grid < 148: + for p in (8, 4, 2): + if max_pages % p == 0 and max_pages // p >= 16 and base_grid * p <= 148 * 4: + page_split = p + break + if page_split > 1: + pages_per_split = max_pages // page_split + partial_out = _ensure_workspace_tensor( + metadata, + "_fp4_mla_attention_pv_partial_buf", + (num_queries, page_split, num_heads, kv_lora_rank), + dtype=torch.float32, + device=q_fp4.device, + ) + if use_triton_v_packed_cache: + _attn_pv_prepacked_v_kernel[ + (num_queries, num_head_blocks, num_dim_blocks * page_split) + ]( + output, + p_fp4, + p_sf, + v_packed, + v_sf, + global_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + kv_cache.shape[0], + output.stride(0), + output.stride(1), + output.stride(2), + output.shape[0] * output.shape[1], + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=kv_lora_rank, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + BLOCK_H=block_h, + BLOCK_V=block_v, + USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, + USE_TMA_OUT_STORE=use_tma_data_load and assume_full_heads and assume_full_v, + PV_LOOP_STAGES=pv_loop_stages, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_FULL_V=assume_full_v, + ASSUME_VALID_PAGES=assume_valid_pages, + PAGE_SPLIT=page_split, + PAGES_PER_SPLIT=pages_per_split, + PARTIAL_OUT=True, + partial_out_ptr=partial_out, + partial_s0=partial_out.stride(0), + partial_s1=partial_out.stride(1), + partial_s2=partial_out.stride(2), + partial_s3=partial_out.stride(3), + **pv_launch_meta, + ) + else: + _attn_pv_kernel[(num_queries, num_head_blocks, num_dim_blocks * page_split)]( + output, + p_fp4, + p_sf, + kv_cache, + kv_cache, + v_sf, + global_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + kv_cache.shape[0], + output.stride(0), + output.stride(1), + output.stride(2), + output.shape[0] * output.shape[1], + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=kv_lora_rank, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + BLOCK_H=block_h, + BLOCK_V=block_v, + USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, + USE_TMA_V_LOAD=use_tma_data_load and kv_lora_rank % block_v == 0, + USE_PREPACKED_V=False, + PV_LOOP_STAGES=pv_loop_stages, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_FULL_V=assume_full_v, + ASSUME_VALID_PAGES=assume_valid_pages, + PAGE_SPLIT=page_split, + PAGES_PER_SPLIT=pages_per_split, + PARTIAL_OUT=True, + partial_out_ptr=partial_out, + partial_s0=partial_out.stride(0), + partial_s1=partial_out.stride(1), + partial_s2=partial_out.stride(2), + partial_s3=partial_out.stride(3), + **pv_launch_meta, + ) + _attn_pv_reduce_kernel[(num_queries, num_head_blocks, num_dim_blocks)]( + output, + partial_out, + global_scale, + output.stride(0), + output.stride(1), + output.stride(2), + partial_out.stride(0), + partial_out.stride(1), + partial_out.stride(2), + partial_out.stride(3), + NUM_HEADS=num_heads, + V_HEAD_D=kv_lora_rank, + PAGE_SPLIT=page_split, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + BLOCK_H=block_h, + BLOCK_V=block_v, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_V=assume_full_v, + **launch_meta, + ) + else: + if use_triton_v_packed_cache: + _attn_pv_prepacked_v_kernel[(num_queries, num_head_blocks, num_dim_blocks)]( + output, + p_fp4, + p_sf, + v_packed, + v_sf, + global_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + kv_cache.shape[0], + output.stride(0), + output.stride(1), + output.stride(2), + output.shape[0] * output.shape[1], + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=kv_lora_rank, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + BLOCK_H=block_h, + BLOCK_V=block_v, + USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, + USE_TMA_OUT_STORE=use_tma_data_load and assume_full_heads and assume_full_v, + PV_LOOP_STAGES=pv_loop_stages, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_FULL_V=assume_full_v, + ASSUME_VALID_PAGES=assume_valid_pages, + **pv_launch_meta, + ) + else: + _attn_pv_kernel[(num_queries, num_head_blocks, num_dim_blocks)]( + output, + p_fp4, + p_sf, + kv_cache, + kv_cache, + v_sf, + global_scale, + src_page_ids, + metadata.paged_kv_indptr_decode, + kv_lens, + src_page_ids.shape[0], + kv_cache.shape[0], + output.stride(0), + output.stride(1), + output.stride(2), + output.shape[0] * output.shape[1], + p_fp4.stride(0), + p_fp4.stride(1), + p_fp4.shape[0], + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + v_sf.stride(0), + NUM_HEADS=num_heads, + V_HEAD_D=kv_lora_rank, + PAGE_SIZE=metadata.page_size, + FP4_BLOCK=FP4_BLOCK_SIZE, + SF_PER_PAGE=sf_per_page, + QUERY_LEN_PER_SEQ=query_len_per_seq, + MAX_PAGES=max_pages, + P_GLOBAL_SCALE=FP4_MLA_P_GLOBAL_SCALE, + BLOCK_H=block_h, + BLOCK_V=block_v, + USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, + USE_TMA_V_LOAD=use_tma_data_load and kv_lora_rank % block_v == 0, + USE_PREPACKED_V=False, + PV_LOOP_STAGES=pv_loop_stages, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_FULL_V=assume_full_v, + ASSUME_VALID_PAGES=assume_valid_pages, + **pv_launch_meta, + ) + + +def run_fp4_mla_attention_decode( + metadata: Any, + layer_idx: int, + local_layer: int, + q_nope: torch.Tensor, + q_pe: torch.Tensor, + output: torch.Tensor, + *, + sm_scale: float, + kv_lora_rank: int, + qk_rope_head_dim: int, ) -> None: """Run MLA decode with FP4 QK and FP4 PV tensor-core matmuls. @@ -2059,19 +2351,8 @@ def run_fp4_mla_attention_decode( raise TypeError( f"FP4 MLA residual Q quantization requires BF16 or FP8 Q; got {q_2d.dtype}." ) - # Per-page path: Q gets its own per-step dynamic global scale (independent of - # the KV per-page scale); the dev QK kernel divides by q_gscale * page_gscale. - page_scale_pool = getattr(metadata, "fp4_mla_page_scale_pool", None) - use_per_page_scale = fp4_mla_per_page_scale_enabled() and page_scale_pool is not None - if use_per_page_scale: - q_amax = q_2d.to(torch.float32).abs().amax().reshape(1) - q_global_scale = torch.where( - q_amax > 0.0, - torch.full_like(q_amax, FP4_MLA_P_GLOBAL_SCALE) / q_amax, - torch.ones_like(q_amax), - ) - else: - q_global_scale = global_scale + backend = _fp4_mla_attention_backend() + q_global_scale = global_scale q_fp4, q_sf = torch.ops.trtllm.fp4_quantize_with_residual( q_2d, q_global_scale, @@ -2096,9 +2377,6 @@ def run_fp4_mla_attention_decode( if max_pages == 0: return - _assert_fp4_mla_decode_paging_consistent(metadata, kv_lens, num_gen_blocks, query_len_per_seq) - - backend = _fp4_mla_attention_backend() if backend == "cutile": from .fp4_mla_cutile import fp4_mla_paged_attention @@ -2164,12 +2442,8 @@ def run_fp4_mla_attention_decode( query_len_per_seq=query_len_per_seq, ) cutile_prepack_v_env = os.environ.get("TRTLLM_FP4_MLA_PREPACK_V") - cutile_storage_valid_pages = ( - assume_valid_pages - or ( - cutile_storage_full_pages - and src_page_ids.numel() == cutile_num_gen_seqs * max_pages - ) + cutile_storage_valid_pages = assume_valid_pages or ( + cutile_storage_full_pages and src_page_ids.numel() == cutile_num_gen_seqs * max_pages ) cutile_allow_qlen_prepack_v = query_len_per_seq > 1 and cutile_prepack_v_env != "0" cutile_assume_valid_pages = assume_valid_pages or ( @@ -2178,10 +2452,7 @@ def run_fp4_mla_attention_decode( cutile_auto_prepack_v = ( hasattr(tl, "make_tensor_descriptor") and num_heads % cutile_block_h == 0 - and ( - assume_full_pages - or (cutile_allow_qlen_prepack_v and cutile_storage_full_pages) - ) + and (assume_full_pages or (cutile_allow_qlen_prepack_v and cutile_storage_full_pages)) and ( cutile_assume_valid_pages or (cutile_allow_qlen_prepack_v and cutile_storage_valid_pages) @@ -2232,16 +2503,6 @@ def run_fp4_mla_attention_decode( and _cutile_persistent_v_pack_enabled() and _cutile_shared_v_pack_storage_enabled() ) - _fp4_mla_debug( - "attention decode cutile launch: " - f"num_queries={num_queries} query_len_per_seq={query_len_per_seq} " - f"num_heads={num_heads} local_layer={local_layer} " - f"layer_idx={layer_idx} head_dim={head_dim} kv_lora_rank={kv_lora_rank} " - f"rope_dim={qk_rope_head_dim} max_pages={max_pages} " - f"assume_full_pages={assume_full_pages} " - f"assume_valid_pages={assume_valid_pages} " - f"use_v_packed_cache={use_cutile_v_packed_cache}" - ) fp4_mla_paged_attention( q_fp4, q_sf, @@ -2290,7 +2551,6 @@ def run_fp4_mla_attention_decode( v_sf=v_sf, page_ids=src_page_ids, ) - _debug_sync("attention_cutile") return total_p_rows = num_queries * max_pages * num_heads @@ -2331,441 +2591,52 @@ def run_fp4_mla_attention_decode( "'triton' or 'cutile'." ) - # Self-contained triton path. Uses kernels copied from the cutile reference - # into fp4_mla_triton.py: TMA-loaded QK + fused page-stats pack, - # reduce-stats, prob-scale, and a specialized PV with tl.ext views. + # Self-contained public-Triton path: TMA-loaded QK + fused page-stats pack, + # reduce-stats, prob-scale, and PV with an optional prepacked V cache. if backend == "triton": _run_triton_attention_decode( metadata=metadata, - q_fp4=q_fp4, - q_sf=q_sf.contiguous().view(-1), - kv_cache=kv_cache, - sf_cache=sf_cache, - v_sf=v_sf, - global_scale=global_scale, - src_page_ids=src_page_ids, - kv_lens=kv_lens, - p_fp4=p_fp4, - p_sf=p_sf, - max_scores=max_scores, - denom=denom, - output=output, - num_queries=num_queries, - num_heads=num_heads, - head_dim=head_dim, - kv_lora_rank=kv_lora_rank, - q_residual_dim=q_residual_dim, - query_len_per_seq=query_len_per_seq, - max_pages=max_pages, - sm_scale=float(sm_scale), - q_global_scale=q_global_scale, - page_scale_pool=page_scale_pool, + layer_idx=layer_idx, local_layer=local_layer, - use_per_page_scale=use_per_page_scale, - ) - _debug_sync("attention_triton") - return - - -def _hp_pool_layer_view( - pool: torch.Tensor, - local_layer: int, - pool_head_dim: int, -) -> torch.Tensor: - return pool[:, local_layer, 0, :].view(pool.shape[0], HP_BLOCK_SIZE, pool_head_dim) - - -def _snapshot_hp_kv_for_mtp_generation( - metadata: Any, - pool: torch.Tensor, - local_layer: int, - *, - num_gen: int, - num_gen_tokens: int, - max_gen_len: int, - metadata_token_offset: int, - head_dim: int, - pool_head_dim: int, -) -> None: - if getattr(metadata, "is_warmup", False): - return - if num_gen_tokens <= num_gen: - return - if max_gen_len > HP_BLOCK_SIZE: - raise NotImplementedError( - "FP4 MLA HP-pool rollback for linear MTP supports at most " - f"{HP_BLOCK_SIZE} generation tokens per sequence, got {max_gen_len}." - ) - - end_token_offset = metadata_token_offset + num_gen_tokens - if ( - end_token_offset > metadata.batch_indices.shape[0] - or end_token_offset > metadata.positions.shape[0] - ): - raise RuntimeError( - "FP4 MLA HP-pool snapshot would read past generation metadata: " - f"token_offset={metadata_token_offset}, num_gen_tokens={num_gen_tokens}, " - f"batch_indices={metadata.batch_indices.shape[0]}, " - f"positions={metadata.positions.shape[0]}." - ) - - snapshots = getattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, None) - if snapshots is None: - snapshots = {} - setattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, snapshots) - - snapshot_pool = getattr(metadata, "fp4_mla_hp_snapshot_pool", None) - if snapshot_pool is not None: - snapshot_pool[:, local_layer, :, :].copy_(pool[:, local_layer, :, :]) - snapshots[int(local_layer)] = { - "mode": "pool", - "metadata_token_offset": metadata_token_offset, - "num_gen_tokens": num_gen_tokens, - "head_dim": head_dim, - "pool_head_dim": pool_head_dim, - } - return - - device = pool.device - token_indices = torch.arange( - metadata_token_offset, - end_token_offset, - dtype=torch.long, - device=device, - ) - batch_indices = metadata.batch_indices[token_indices].to(torch.long) - positions = metadata.positions[token_indices].to(torch.long) - seq_slots = metadata.seq_slots[batch_indices].to(torch.long) - hp_slots = torch.remainder(positions, HP_BLOCK_SIZE).to(torch.long) - first_new_positions = metadata.kv_lens_cuda_runtime[batch_indices].to( - torch.long - ) - metadata.prompt_lens_cuda_runtime[batch_indices].to(torch.long) - pool_view = _hp_pool_layer_view(pool, local_layer, pool_head_dim) - values = pool_view[seq_slots, hp_slots, :head_dim].clone() - - snapshots[int(local_layer)] = { - "mode": "values", - "batch_indices": batch_indices, - "seq_slots": seq_slots, - "hp_slots": hp_slots, - "positions": positions, - "first_new_positions": first_new_positions, - "values": values, - "head_dim": head_dim, - "pool_head_dim": pool_head_dim, - } - - -def repair_fp4_mla_hp_kv_for_mtp_rejection( - metadata: Any, - num_accepted_tokens: torch.Tensor, -) -> None: - """Restore HP-pool slots that belonged to rejected linear-MTP tokens. - - Packed FP4 pages past the accepted logical KV length are harmless because - later attention ignores them. The BF16 HP pool is a circular tail mirror, so - rejected speculative writes must be rolled back before the next tile rewrite. - """ - snapshots = getattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, None) - if not snapshots: - return - - keep_snapshots = False - try: - pool = getattr(metadata, "high_precision_kv_pool", None) - if pool is None: - return - accepted_tokens = num_accepted_tokens.to(device=pool.device) - for local_layer, snapshot in snapshots.items(): - head_dim = snapshot["head_dim"] - pool_head_dim = snapshot["pool_head_dim"] - block_d = triton.next_power_of_2(head_dim) - if snapshot.get("mode") == "pool": - keep_snapshots = True - _hp_kv_restore_rejected_from_pool_kernel[(snapshot["num_gen_tokens"],)]( - pool, - metadata.fp4_mla_hp_snapshot_pool, - metadata.batch_indices, - metadata.positions, - metadata.seq_slots, - metadata.kv_lens_cuda_runtime, - metadata.prompt_lens_cuda_runtime, - accepted_tokens, - snapshot["metadata_token_offset"], - snapshot["num_gen_tokens"], - metadata.batch_indices.shape[0], - metadata.num_seqs, - accepted_tokens.shape[0], - pool.shape[0], - pool.shape[1], - int(local_layer), - pool.stride(0), - pool.stride(1), - metadata.fp4_mla_hp_snapshot_pool.stride(0), - metadata.fp4_mla_hp_snapshot_pool.stride(1), - D=head_dim, - POOL_HEAD_D=pool_head_dim, - BLOCK_D=block_d, - HP_BLOCK=HP_BLOCK_SIZE, - ) - else: - positions = snapshot["positions"] - batch_indices = snapshot["batch_indices"] - seq_slots = snapshot["seq_slots"] - hp_slots = snapshot["hp_slots"] - first_new_positions = snapshot["first_new_positions"] - if positions.shape[0] == 0: - continue - _hp_kv_restore_rejected_from_values_kernel[(positions.shape[0],)]( - pool, - snapshot["values"], - batch_indices, - positions, - seq_slots, - hp_slots, - first_new_positions, - accepted_tokens, - positions.shape[0], - accepted_tokens.shape[0], - pool.shape[0], - pool.shape[1], - int(local_layer), - pool.stride(0), - pool.stride(1), - snapshot["values"].stride(0), - snapshot["values"].stride(1), - D=head_dim, - POOL_HEAD_D=pool_head_dim, - BLOCK_D=block_d, - HP_BLOCK=HP_BLOCK_SIZE, - ) - finally: - if not keep_snapshots: - setattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, None) - - -def update_hp_kv_for_fp4_mla( - metadata: Any, - latent_cache: Optional[torch.Tensor], - local_layer: int, - *, - phase: _HPUpdatePhase = "all", -) -> None: - """Store recent KV tokens at BF16 into the high-precision pool. - - Called on every layer before the attention kernel. The pool acts as a - circular buffer of HP_BLOCK_SIZE slots per sequence: - - Context phase stores the last ``kv_len % HP_BLOCK_SIZE`` new tokens of - each request into buffer positions [0, remainder). These are the - tail tokens that do not fill a complete FP4 block of 16. - - Generation phase stores every new token for each request into position - ``position % HP_BLOCK_SIZE``, overwriting the oldest entries in the - circular buffer. This supports linear MTP where a request contributes - more than one generation token in a forward pass. - - The Triton kernels use the GPU ``seq_slots`` tensor for scatter indexing - and are CUDA-graph-compatible for the generation phase. - - Args: - metadata: Attention metadata exposing ``num_contexts``, ``num_seqs``, - ``seq_slots`` / ``seq_slots_cpu``, ``request_ids``, - ``is_cuda_graph``, ``is_warmup``, ``hp_pool_owners``, - ``high_precision_kv_pool``, ``prompt_lens_cpu_runtime``, - ``prompt_lens_cuda_runtime``, ``kv_lens_cuda_runtime``. - latent_cache: MLA latent cache for the current tokens, shape - [num_tokens, head_dim]. When ``None``, only ownership tracking - runs (no data is written to the pool). - local_layer: Layer index within the local pipeline-parallel slice. - phase: Which portion of ``latent_cache`` is present. ``"all"`` means - context tokens followed by generation tokens, ``"context"`` means - only context tokens, and ``"generation"`` means only generation - tokens. - """ - if phase not in ("all", "context", "generation"): - raise ValueError(f"Unexpected FP4 MLA HP update phase: {phase}") - if metadata.hp_pool_owners is None: - return - num_contexts = metadata.num_contexts - num_seqs = metadata.num_seqs - update_context = phase in ("all", "context") - update_generation = phase in ("all", "generation") - - # ------------------------------------------------------------------ - # Ownership tracking (layer 0, eager mode only - debug guard). - # Context phase never uses CUDA graph; decode check is debug-only. - # ------------------------------------------------------------------ - if local_layer == 0 and not metadata.is_cuda_graph and not metadata.is_warmup: - # Context: register ownership of each seq_slot. - if update_context: - for batch_idx in range(num_contexts): - seq_slot = metadata.seq_slots_cpu[batch_idx].item() - request_id = metadata.request_ids[batch_idx] - metadata.hp_pool_owners[seq_slot] = request_id - - # Decode: verify that the expected request still owns each slot. - if update_generation: - for batch_idx in range(num_contexts, num_seqs): - seq_slot = metadata.seq_slots_cpu[batch_idx].item() - request_id = metadata.request_ids[batch_idx] - owner = metadata.hp_pool_owners.get(seq_slot) - if owner != request_id: - raise RuntimeError( - f"HP KV pool ownership mismatch: seq_slot={seq_slot} " - f"is owned by request {owner} but request " - f"{request_id} is attempting to use it" - ) - - if latent_cache is None: - return - - # ------------------------------------------------------------------ - # Triton kernel dispatch - runs on every layer, CUDA-graph-safe. - # ------------------------------------------------------------------ - pool = metadata.high_precision_kv_pool - head_dim = latent_cache.shape[-1] - pool_head_dim = pool.shape[-1] // HP_BLOCK_SIZE - if pool_head_dim < head_dim: - raise RuntimeError( - f"FP4 MLA HP pool head dimension is too small: got " - f"{pool_head_dim}, need at least {head_dim}." - ) - block_d = triton.next_power_of_2(head_dim) - pool_s0 = pool.stride(0) # stride across sequence slots - pool_s1 = pool.stride(1) # stride across layers - lc_stride = latent_cache.stride(0) - _fp4_mla_debug( - "hp update: " - f"phase={phase} local_layer={local_layer} num_contexts={num_contexts} " - f"num_seqs={num_seqs} head_dim={head_dim} " - f"pool_head_dim={pool_head_dim} block_d={block_d}" - ) - _fp4_mla_debug(f"hp latent_cache: {_tensor_layout(latent_cache)}") - _fp4_mla_debug(f"hp pool: {_tensor_layout(pool)}") - _debug_tensor_range("hp seq_slots", metadata.seq_slots[:num_seqs]) - _debug_tensor_range("hp kv_lens", metadata.kv_lens_cuda_runtime[:num_seqs]) - - # Context phase: store last (kv_len % HP_BLOCK_SIZE) new tokens. - if update_context and num_contexts > 0: - prompt_lens_cpu = metadata.prompt_lens_cpu_runtime[:num_contexts] - # Exclusive prefix sum: token offset in latent_cache for each ctx seq. - token_offsets_cpu = torch.zeros(num_contexts, dtype=torch.int32, device="cpu") - if num_contexts > 1: - token_offsets_cpu[1:].copy_(torch.cumsum(prompt_lens_cpu[:-1].to(torch.int32), dim=0)) - token_offsets_gpu = token_offsets_cpu.to(pool.device, non_blocking=False) - prompt_lens_gpu = metadata.prompt_lens_cuda_runtime[:num_contexts] - - _fp4_mla_debug( - "hp context launch: " - f"grid=({num_contexts}, {HP_BLOCK_SIZE}) " - f"token_offsets={token_offsets_cpu.tolist()}" - ) - _debug_tensor_range("hp context prompt_lens", prompt_lens_gpu) - _hp_kv_store_context_kernel[(num_contexts, HP_BLOCK_SIZE)]( - pool, - latent_cache, - metadata.seq_slots, - metadata.kv_lens_cuda_runtime, - token_offsets_gpu, - prompt_lens_gpu, - pool.shape[0], - pool.shape[1], - local_layer, - pool_s0, - pool_s1, - lc_stride, - D=head_dim, - POOL_HEAD_D=pool_head_dim, - BLOCK_D=block_d, - HP_BLOCK=HP_BLOCK_SIZE, - ) - _debug_sync("hp_context") - - # Generation phase: store current tokens at position % HP_BLOCK_SIZE. - num_gen = num_seqs - num_contexts - if update_generation and num_gen > 0: - gen_tok_start = 0 - metadata_token_offset = getattr(metadata, "num_ctx_tokens", 0) - if phase == "all": - # Scalar offset: number of context tokens packed before gen tokens. - gen_tok_start = int(metadata.prompt_lens_cpu_runtime[:num_contexts].sum().item()) - metadata_token_offset = gen_tok_start - num_gen_tokens = latent_cache.shape[0] - gen_tok_start - if num_gen_tokens < 0: - raise RuntimeError( - "FP4 MLA HP generation update received fewer latent tokens than " - f"the context prefix: latent_tokens={latent_cache.shape[0]}, " - f"context_tokens={gen_tok_start}." - ) - if num_gen_tokens == 0: - return - - gen_token_lens = _host_int_list_during_forward( - getattr(metadata, "prompt_lens_cpu_runtime", None), num_contexts, num_seqs - ) - if gen_token_lens is not None: - max_gen_len = max(gen_token_lens) - elif num_gen_tokens % num_gen == 0: - max_gen_len = num_gen_tokens // num_gen - else: - max_gen_len = num_gen_tokens - _snapshot_hp_kv_for_mtp_generation( - metadata, - pool, - local_layer, - num_gen=num_gen, - num_gen_tokens=num_gen_tokens, - max_gen_len=max_gen_len, - metadata_token_offset=metadata_token_offset, - head_dim=head_dim, - pool_head_dim=pool_head_dim, - ) - - _fp4_mla_debug( - "hp generation launch: " - f"grid=({num_gen_tokens},) gen_tok_start={gen_tok_start} " - f"metadata_token_offset={metadata_token_offset}" - ) - _hp_kv_store_gen_kernel[(num_gen_tokens,)]( - pool, - latent_cache, - metadata.seq_slots, - metadata.batch_indices, - metadata.positions, - gen_tok_start, - metadata_token_offset, - num_gen_tokens, - metadata.batch_indices.shape[0], - pool.shape[0], - pool.shape[1], - local_layer, - pool_s0, - pool_s1, - lc_stride, - D=head_dim, - POOL_HEAD_D=pool_head_dim, - BLOCK_D=block_d, - HP_BLOCK=HP_BLOCK_SIZE, + q_fp4=q_fp4, + q_sf=q_sf.contiguous().view(-1), + kv_cache=kv_cache, + sf_cache=sf_cache, + v_sf=v_sf, + global_scale=global_scale, + src_page_ids=src_page_ids, + kv_lens=kv_lens, + p_fp4=p_fp4, + p_sf=p_sf, + max_scores=max_scores, + denom=denom, + output=output, + num_queries=num_queries, + num_heads=num_heads, + head_dim=head_dim, + kv_lora_rank=kv_lora_rank, + q_residual_dim=q_residual_dim, + query_len_per_seq=query_len_per_seq, + max_pages=max_pages, + sm_scale=float(sm_scale), + q_global_scale=q_global_scale, ) - _debug_sync("hp_generation") + return -def _stage_pool_layer_view( +def _hp_pool_layer_view( pool: torch.Tensor, local_layer: int, pool_head_dim: int, - slots: int, ) -> torch.Tensor: - return pool[:, local_layer, 0, :].view(pool.shape[0], slots, pool_head_dim) + return pool[:, local_layer, 0, :].view(pool.shape[0], HP_BLOCK_SIZE, pool_head_dim) -def _snapshot_page_stage_for_mtp_generation( +def _snapshot_hp_kv_for_mtp_generation( metadata: Any, pool: torch.Tensor, local_layer: int, *, - slots: int, num_gen: int, num_gen_tokens: int, max_gen_len: int, @@ -2773,21 +2644,14 @@ def _snapshot_page_stage_for_mtp_generation( head_dim: int, pool_head_dim: int, ) -> None: - """Snapshot staging-pool slots before linear-MTP generation writes. - - Mirror of ``_snapshot_hp_kv_for_mtp_generation`` for the per-page staging - buffer, keyed under ``_FP4_MLA_MTP_STAGE_SNAPSHOTS`` so it never aliases the - HP-pool snapshots. ``slots`` is the page size (vs. the HP pool's - ``HP_BLOCK_SIZE``), so the per-sequence MTP draft length must not exceed it. - """ if getattr(metadata, "is_warmup", False): return if num_gen_tokens <= num_gen: return - if max_gen_len > slots: + if max_gen_len > HP_BLOCK_SIZE: raise NotImplementedError( - "FP4 MLA staging rollback for linear MTP supports at most " - f"{slots} generation tokens per sequence, got {max_gen_len}." + "FP4 MLA HP-pool rollback for linear MTP supports at most " + f"{HP_BLOCK_SIZE} generation tokens per sequence, got {max_gen_len}." ) end_token_offset = metadata_token_offset + num_gen_tokens @@ -2796,18 +2660,18 @@ def _snapshot_page_stage_for_mtp_generation( or end_token_offset > metadata.positions.shape[0] ): raise RuntimeError( - "FP4 MLA staging snapshot would read past generation metadata: " + "FP4 MLA HP-pool snapshot would read past generation metadata: " f"token_offset={metadata_token_offset}, num_gen_tokens={num_gen_tokens}, " f"batch_indices={metadata.batch_indices.shape[0]}, " f"positions={metadata.positions.shape[0]}." ) - snapshots = getattr(metadata, _FP4_MLA_MTP_STAGE_SNAPSHOTS, None) + snapshots = getattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, None) if snapshots is None: snapshots = {} - setattr(metadata, _FP4_MLA_MTP_STAGE_SNAPSHOTS, snapshots) + setattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, snapshots) - snapshot_pool = getattr(metadata, "fp4_mla_page_stage_snapshot_pool", None) + snapshot_pool = getattr(metadata, "fp4_mla_hp_snapshot_pool", None) if snapshot_pool is not None: snapshot_pool[:, local_layer, :, :].copy_(pool[:, local_layer, :, :]) snapshots[int(local_layer)] = { @@ -2816,7 +2680,6 @@ def _snapshot_page_stage_for_mtp_generation( "num_gen_tokens": num_gen_tokens, "head_dim": head_dim, "pool_head_dim": pool_head_dim, - "slots": slots, } return @@ -2830,78 +2693,191 @@ def _snapshot_page_stage_for_mtp_generation( batch_indices = metadata.batch_indices[token_indices].to(torch.long) positions = metadata.positions[token_indices].to(torch.long) seq_slots = metadata.seq_slots[batch_indices].to(torch.long) - stage_slots = torch.remainder(positions, slots).to(torch.long) + hp_slots = torch.remainder(positions, HP_BLOCK_SIZE).to(torch.long) first_new_positions = metadata.kv_lens_cuda_runtime[batch_indices].to( torch.long ) - metadata.prompt_lens_cuda_runtime[batch_indices].to(torch.long) - pool_view = _stage_pool_layer_view(pool, local_layer, pool_head_dim, slots) - values = pool_view[seq_slots, stage_slots, :head_dim].clone() + pool_view = _hp_pool_layer_view(pool, local_layer, pool_head_dim) + values = pool_view[seq_slots, hp_slots, :head_dim].clone() snapshots[int(local_layer)] = { "mode": "values", "batch_indices": batch_indices, "seq_slots": seq_slots, - "hp_slots": stage_slots, + "hp_slots": hp_slots, "positions": positions, "first_new_positions": first_new_positions, "values": values, "head_dim": head_dim, "pool_head_dim": pool_head_dim, - "slots": slots, } -def update_page_stage_for_fp4_mla( +def repair_fp4_mla_hp_kv_for_mtp_rejection( + metadata: Any, + num_accepted_tokens: torch.Tensor, +) -> None: + """Restore HP-pool slots that belonged to rejected linear-MTP tokens. + + Packed FP4 pages past the accepted logical KV length are harmless because + later attention ignores them. The BF16 HP pool is a circular tail mirror, so + rejected speculative writes must be rolled back before the next tile rewrite. + """ + snapshots = getattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, None) + if not snapshots: + return + + keep_snapshots = False + try: + pool = getattr(metadata, "high_precision_kv_pool", None) + if pool is None: + return + accepted_tokens = num_accepted_tokens.to(device=pool.device) + for local_layer, snapshot in snapshots.items(): + head_dim = snapshot["head_dim"] + pool_head_dim = snapshot["pool_head_dim"] + block_d = triton.next_power_of_2(head_dim) + if snapshot.get("mode") == "pool": + keep_snapshots = True + _hp_kv_restore_rejected_from_pool_kernel[(snapshot["num_gen_tokens"],)]( + pool, + metadata.fp4_mla_hp_snapshot_pool, + metadata.batch_indices, + metadata.positions, + metadata.seq_slots, + metadata.kv_lens_cuda_runtime, + metadata.prompt_lens_cuda_runtime, + accepted_tokens, + snapshot["metadata_token_offset"], + snapshot["num_gen_tokens"], + metadata.batch_indices.shape[0], + metadata.num_seqs, + accepted_tokens.shape[0], + pool.shape[0], + pool.shape[1], + int(local_layer), + pool.stride(0), + pool.stride(1), + metadata.fp4_mla_hp_snapshot_pool.stride(0), + metadata.fp4_mla_hp_snapshot_pool.stride(1), + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + else: + positions = snapshot["positions"] + batch_indices = snapshot["batch_indices"] + seq_slots = snapshot["seq_slots"] + hp_slots = snapshot["hp_slots"] + first_new_positions = snapshot["first_new_positions"] + if positions.shape[0] == 0: + continue + _hp_kv_restore_rejected_from_values_kernel[(positions.shape[0],)]( + pool, + snapshot["values"], + batch_indices, + positions, + seq_slots, + hp_slots, + first_new_positions, + accepted_tokens, + positions.shape[0], + accepted_tokens.shape[0], + pool.shape[0], + pool.shape[1], + int(local_layer), + pool.stride(0), + pool.stride(1), + snapshot["values"].stride(0), + snapshot["values"].stride(1), + D=head_dim, + POOL_HEAD_D=pool_head_dim, + BLOCK_D=block_d, + HP_BLOCK=HP_BLOCK_SIZE, + ) + finally: + if not keep_snapshots: + setattr(metadata, _FP4_MLA_MTP_HP_SNAPSHOTS, None) + + +def update_hp_kv_for_fp4_mla( metadata: Any, latent_cache: Optional[torch.Tensor], local_layer: int, *, phase: _HPUpdatePhase = "all", ) -> None: - """Store recent KV tokens at BF16 into the per-page staging buffer. - - Separate twin of ``update_hp_kv_for_fp4_mla`` for the per-page - dynamic-scale path. Identical circular-buffer semantics, but the buffer - holds a whole page (``metadata.page_size`` slots) instead of the HP pool's - ``HP_BLOCK_SIZE`` slots, so the active page can be re-quantized to FP4 with - its exact per-page amax. Reuses the HP store kernels with ``HP_BLOCK`` set - to the page size. Ownership tracking is intentionally omitted -- the HP - update already validates it on this path. + """Store recent KV tokens at BF16 into the high-precision pool. + + Called on every layer before the attention kernel. The pool acts as a + circular buffer of HP_BLOCK_SIZE slots per sequence: + + Context phase stores the last ``kv_len % HP_BLOCK_SIZE`` new tokens of + each request into buffer positions [0, remainder). These are the + tail tokens that do not fill a complete FP4 block of 16. + + Generation phase stores every new token for each request into position + ``position % HP_BLOCK_SIZE``, overwriting the oldest entries in the + circular buffer. This supports linear MTP where a request contributes + more than one generation token in a forward pass. + + The Triton kernels use the GPU ``seq_slots`` tensor for scatter indexing + and are CUDA-graph-compatible for the generation phase. + + Args: + metadata: Attention metadata exposing ``num_contexts``, ``num_seqs``, + ``seq_slots`` / ``seq_slots_cpu``, ``request_ids``, + ``is_cuda_graph``, ``is_warmup``, + ``high_precision_kv_pool``, ``prompt_lens_cpu_runtime``, + ``prompt_lens_cuda_runtime``, ``kv_lens_cuda_runtime``. + latent_cache: MLA latent cache for the current tokens, shape + [num_tokens, head_dim]. When ``None``, only ownership tracking + runs (no data is written to the pool). + local_layer: Layer index within the local pipeline-parallel slice. + phase: Which portion of ``latent_cache`` is present. ``"all"`` means + context tokens followed by generation tokens, ``"context"`` means + only context tokens, and ``"generation"`` means only generation + tokens. """ if phase not in ("all", "context", "generation"): - raise ValueError(f"Unexpected FP4 MLA staging update phase: {phase}") - pool = getattr(metadata, "fp4_mla_page_stage_pool", None) - if pool is None or latent_cache is None: + raise ValueError(f"Unexpected FP4 MLA HP update phase: {phase}") + if metadata.high_precision_kv_pool is None: return - num_contexts = metadata.num_contexts num_seqs = metadata.num_seqs update_context = phase in ("all", "context") update_generation = phase in ("all", "generation") - slots = metadata.page_size + if latent_cache is None: + return + + # ------------------------------------------------------------------ + # Triton kernel dispatch - runs on every layer, CUDA-graph-safe. + # ------------------------------------------------------------------ + pool = metadata.high_precision_kv_pool head_dim = latent_cache.shape[-1] - pool_head_dim = pool.shape[-1] // slots + pool_head_dim = pool.shape[-1] // HP_BLOCK_SIZE if pool_head_dim < head_dim: raise RuntimeError( - f"FP4 MLA staging pool head dimension is too small: got " + f"FP4 MLA HP pool head dimension is too small: got " f"{pool_head_dim}, need at least {head_dim}." ) block_d = triton.next_power_of_2(head_dim) - pool_s0 = pool.stride(0) - pool_s1 = pool.stride(1) + pool_s0 = pool.stride(0) # stride across sequence slots + pool_s1 = pool.stride(1) # stride across layers lc_stride = latent_cache.stride(0) - # Context phase: store last (kv_len % page_size) new tokens. + # Context phase: store last (kv_len % HP_BLOCK_SIZE) new tokens. if update_context and num_contexts > 0: prompt_lens_cpu = metadata.prompt_lens_cpu_runtime[:num_contexts] + # Exclusive prefix sum: token offset in latent_cache for each ctx seq. token_offsets_cpu = torch.zeros(num_contexts, dtype=torch.int32, device="cpu") if num_contexts > 1: token_offsets_cpu[1:].copy_(torch.cumsum(prompt_lens_cpu[:-1].to(torch.int32), dim=0)) token_offsets_gpu = token_offsets_cpu.to(pool.device, non_blocking=False) prompt_lens_gpu = metadata.prompt_lens_cuda_runtime[:num_contexts] - - _hp_kv_store_context_kernel[(num_contexts, slots)]( + _hp_kv_store_context_kernel[(num_contexts, HP_BLOCK_SIZE)]( pool, latent_cache, metadata.seq_slots, @@ -2917,23 +2893,23 @@ def update_page_stage_for_fp4_mla( D=head_dim, POOL_HEAD_D=pool_head_dim, BLOCK_D=block_d, - HP_BLOCK=slots, + HP_BLOCK=HP_BLOCK_SIZE, ) - _debug_sync("stage_context") - # Generation phase: store current tokens at position % page_size. + # Generation phase: store current tokens at position % HP_BLOCK_SIZE. num_gen = num_seqs - num_contexts if update_generation and num_gen > 0: gen_tok_start = 0 metadata_token_offset = getattr(metadata, "num_ctx_tokens", 0) if phase == "all": + # Scalar offset: number of context tokens packed before gen tokens. gen_tok_start = int(metadata.prompt_lens_cpu_runtime[:num_contexts].sum().item()) metadata_token_offset = gen_tok_start num_gen_tokens = latent_cache.shape[0] - gen_tok_start if num_gen_tokens < 0: raise RuntimeError( - "FP4 MLA staging generation update received fewer latent tokens " - f"than the context prefix: latent_tokens={latent_cache.shape[0]}, " + "FP4 MLA HP generation update received fewer latent tokens than " + f"the context prefix: latent_tokens={latent_cache.shape[0]}, " f"context_tokens={gen_tok_start}." ) if num_gen_tokens == 0: @@ -2948,11 +2924,10 @@ def update_page_stage_for_fp4_mla( max_gen_len = num_gen_tokens // num_gen else: max_gen_len = num_gen_tokens - _snapshot_page_stage_for_mtp_generation( + _snapshot_hp_kv_for_mtp_generation( metadata, pool, local_layer, - slots=slots, num_gen=num_gen, num_gen_tokens=num_gen_tokens, max_gen_len=max_gen_len, @@ -2960,7 +2935,6 @@ def update_page_stage_for_fp4_mla( head_dim=head_dim, pool_head_dim=pool_head_dim, ) - _hp_kv_store_gen_kernel[(num_gen_tokens,)]( pool, latent_cache, @@ -2980,95 +2954,5 @@ def update_page_stage_for_fp4_mla( D=head_dim, POOL_HEAD_D=pool_head_dim, BLOCK_D=block_d, - HP_BLOCK=slots, + HP_BLOCK=HP_BLOCK_SIZE, ) - _debug_sync("stage_generation") - - -def repair_fp4_mla_page_stage_for_mtp_rejection( - metadata: Any, - num_accepted_tokens: torch.Tensor, -) -> None: - """Restore staging-pool slots that belonged to rejected linear-MTP tokens. - - Separate twin of ``repair_fp4_mla_hp_kv_for_mtp_rejection`` for the per-page - staging buffer. Reuses the HP restore kernels with ``HP_BLOCK`` set to the - snapshot's page size. - """ - snapshots = getattr(metadata, _FP4_MLA_MTP_STAGE_SNAPSHOTS, None) - if not snapshots: - return - - keep_snapshots = False - try: - pool = getattr(metadata, "fp4_mla_page_stage_pool", None) - if pool is None: - return - accepted_tokens = num_accepted_tokens.to(device=pool.device) - for local_layer, snapshot in snapshots.items(): - head_dim = snapshot["head_dim"] - pool_head_dim = snapshot["pool_head_dim"] - slots = snapshot["slots"] - block_d = triton.next_power_of_2(head_dim) - if snapshot.get("mode") == "pool": - keep_snapshots = True - _hp_kv_restore_rejected_from_pool_kernel[(snapshot["num_gen_tokens"],)]( - pool, - metadata.fp4_mla_page_stage_snapshot_pool, - metadata.batch_indices, - metadata.positions, - metadata.seq_slots, - metadata.kv_lens_cuda_runtime, - metadata.prompt_lens_cuda_runtime, - accepted_tokens, - snapshot["metadata_token_offset"], - snapshot["num_gen_tokens"], - metadata.batch_indices.shape[0], - metadata.num_seqs, - accepted_tokens.shape[0], - pool.shape[0], - pool.shape[1], - int(local_layer), - pool.stride(0), - pool.stride(1), - metadata.fp4_mla_page_stage_snapshot_pool.stride(0), - metadata.fp4_mla_page_stage_snapshot_pool.stride(1), - D=head_dim, - POOL_HEAD_D=pool_head_dim, - BLOCK_D=block_d, - HP_BLOCK=slots, - ) - else: - positions = snapshot["positions"] - batch_indices = snapshot["batch_indices"] - seq_slots = snapshot["seq_slots"] - stage_slots = snapshot["hp_slots"] - first_new_positions = snapshot["first_new_positions"] - if positions.shape[0] == 0: - continue - _hp_kv_restore_rejected_from_values_kernel[(positions.shape[0],)]( - pool, - snapshot["values"], - batch_indices, - positions, - seq_slots, - stage_slots, - first_new_positions, - accepted_tokens, - positions.shape[0], - accepted_tokens.shape[0], - pool.shape[0], - pool.shape[1], - int(local_layer), - pool.stride(0), - pool.stride(1), - snapshot["values"].stride(0), - snapshot["values"].stride(1), - D=head_dim, - POOL_HEAD_D=pool_head_dim, - BLOCK_D=block_d, - HP_BLOCK=slots, - ) - finally: - if not keep_snapshots: - setattr(metadata, _FP4_MLA_MTP_STAGE_SNAPSHOTS, None) diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py.bak b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py.bak deleted file mode 100644 index 3cd4f7bdb94c..000000000000 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py.bak +++ /dev/null @@ -1,3239 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# SPDX-License-Identifier: Apache-2.0 - -"""FP4 MLA paged decode attention using Triton. - -The kernels are adapted from TensorRT-LLM's FP4 MLA decode path. This module -exposes the attention path for already-packed FP4 Q/K/V tensors and swizzled -FP8 block-scale tensors; quantization and KV-cache update helpers remain outside -this internal op. -""" - -import os -from typing import Optional - -import torch -import triton -import triton.language as tl - -FP4_BLOCK_SIZE = 16 -FP4_MLA_P_GLOBAL_SCALE = 448.0 * 6.0 - - -def _ceil_div(lhs: int, rhs: int) -> int: - return (lhs + rhs - 1) // rhs - - -def _env_int(name: str) -> Optional[int]: - value = os.environ.get(name) - if value is None or value == "": - return None - return int(value) - - -def _swizzled_scale_size(rows: int, logical_cols: int) -> int: - scale_cols = _ceil_div(logical_cols, FP4_BLOCK_SIZE) - padded_cols = _ceil_div(scale_cols, 4) * 4 - return _ceil_div(rows, 128) * 128 * padded_cols - - -def _get_kv_cache_strides(kv_cache: torch.Tensor) -> tuple[int, int, int, int, int, int]: - if kv_cache.dim() == 3: - num_pages, page_size, packed_dim = kv_cache.shape - return ( - num_pages, - page_size, - packed_dim, - kv_cache.stride(0), - kv_cache.stride(1), - kv_cache.stride(2), - ) - if kv_cache.dim() >= 5: - num_pages = kv_cache.shape[0] - page_size = kv_cache.shape[2] - packed_dim = kv_cache.shape[4] - return ( - num_pages, - page_size, - packed_dim, - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - ) - raise ValueError( - "kv_cache must be shaped (num_pages, page_size, packed_dim) or (num_pages, ..., page_size, ..., packed_dim)." - ) - - -def _workspace_tensor( - workspace: Optional[torch.Tensor], - shape: tuple[int, ...], - *, - dtype: torch.dtype, - device: torch.device, - name: str, -) -> torch.Tensor: - if workspace is None: - if torch.cuda.is_current_stream_capturing(): - raise ValueError( - f"Cannot allocate {name} while capturing a CUDA graph. " - "Pass a preallocated workspace tensor." - ) - return torch.empty(shape, dtype=dtype, device=device) - - invalid = ( - workspace.dtype != dtype - or workspace.device != device - or len(workspace.shape) != len(shape) - or any(workspace.shape[idx] < dim for idx, dim in enumerate(shape)) - ) - if invalid: - raise ValueError( - f"{name} workspace must have shape at least {shape}, dtype={dtype}, " - f"and device={device}; got shape={tuple(workspace.shape)}, " - f"dtype={workspace.dtype}, device={workspace.device}." - ) - - slices = tuple(slice(0, dim) for dim in shape) - return workspace[slices] - - -@triton.jit -def _fp4_mla_swizzled_sf_offset(row_idx, col_idx, SF_PER_TOKEN: tl.constexpr): - padded_cols = ((SF_PER_TOKEN + 3) // 4) * 4 - col_in_group = col_idx % 4 - col_group = col_idx // 4 - row_in_group0 = row_idx % 32 - row_in_group1 = (row_idx % 128) // 32 - row_group = row_idx // 128 - return ( - col_in_group - + col_group * (4 * 128) - + row_in_group0 * 16 - + row_in_group1 * 4 - + row_group * (128 * padded_cols) - ) - - -@triton.jit -def _fp4_mla_swizzled_sf_offset_row_block( - row_group, row_offsets, col_idx, SF_PER_TOKEN: tl.constexpr -): - padded_cols = ((SF_PER_TOKEN + 3) // 4) * 4 - col_part = (col_idx % 4) + (col_idx // 4) * (4 * 128) - row_part = (row_offsets % 32) * 16 + ((row_offsets % 128) // 32) * 4 - return col_part + row_part + row_group * (128 * padded_cols) - - -@triton.jit -def _fp4_e2m1_quantize(x): - abs_x = tl.abs(x) - magnitude = tl.where( - abs_x < 0.25, - 0, - tl.where( - abs_x < 0.75, - 1, - tl.where( - abs_x < 1.25, - 2, - tl.where( - abs_x < 1.75, - 3, - tl.where(abs_x < 2.5, 4, tl.where(abs_x < 3.5, 5, tl.where(abs_x < 5.0, 6, 7))), - ), - ), - ), - ) - sign = tl.where(x < 0.0, 8, 0) - return (magnitude | sign).to(tl.uint8) - - -@triton.jit -def _fp4_e2m1_quantize_packed(even, odd): - return tl.inline_asm_elementwise( - """ - { - .reg .b8 r; - cvt.rn.satfinite.e2m1x2.f32 r, $1, $2; - mov.b32 $0, {r, r, r, r}; - } - """, - constraints="=r,f,f", - args=[odd.to(tl.float32), even.to(tl.float32)], - dtype=tl.uint8, - is_pure=True, - pack=1, - ) - - -@triton.jit -def _fp4_pack_low_nibbles(even_packed, odd_packed): - return tl.inline_asm_elementwise( - """ - { - .reg .b32 lo; - .reg .b32 hi; - and.b32 lo, $1, 15; - and.b32 hi, $2, 15; - shl.b32 hi, hi, 4; - or.b32 $0, lo, hi; - } - """, - constraints="=r,r,r", - args=[even_packed, odd_packed], - dtype=tl.uint8, - is_pure=True, - pack=1, - ) - - -@triton.jit -def _fp4_pack_high_nibbles(even_packed, odd_packed): - return tl.inline_asm_elementwise( - """ - { - .reg .b32 lo; - .reg .b32 hi; - shr.u32 lo, $1, 4; - and.b32 lo, lo, 15; - and.b32 hi, $2, 240; - or.b32 $0, lo, hi; - } - """, - constraints="=r,r,r", - args=[even_packed, odd_packed], - dtype=tl.uint8, - is_pure=True, - pack=1, - ) - - -@triton.jit -def _fp4_pack_nibbles(even_packed, odd_packed): - return tl.inline_asm_elementwise( - """ - { - .reg .b32 lo; - .reg .b32 hi; - and.b32 lo, $2, 15; - and.b32 hi, $3, 15; - shl.b32 hi, hi, 4; - or.b32 $0, lo, hi; - - shr.u32 lo, $2, 4; - and.b32 lo, lo, 15; - and.b32 hi, $3, 240; - or.b32 $1, lo, hi; - } - """, - constraints="=r,=r,r,r", - args=[even_packed, odd_packed], - dtype=(tl.uint8, tl.uint8), - is_pure=True, - pack=1, - ) - - -@triton.jit -def _fp4_mla_attention_v_repack_kernel( - v_packed_ptr, - kv_cache_ptr, - num_pages, - kv_s0, - kv_s2, - kv_s4, - V_HEAD_D: tl.constexpr, - PAGE_SIZE: tl.constexpr, - BLOCK_V: tl.constexpr, - occupancy: tl.constexpr = 1, -): - page_idx = tl.program_id(0) - dim_block = tl.program_id(1) - - tl.assume(kv_s0 % 8 == 0) - tl.assume(kv_s2 % 8 == 0) - tl.assume(kv_s4 == 1) - v_view = tl.ext.make_view( - base=kv_cache_ptr, - shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], - tile_dim_map=[0, 1, 2], - ) - v_tile = tl.ext.load_view_tko( - v_view, - [ - page_idx.to(tl.int32), - 0, - (dim_block * (BLOCK_V // 2)).to(tl.int32), - ], - ) - v_tile = v_tile.to(tl.uint8, bitcast=True) - v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) - v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) - even_packed, odd_packed = tl.split(v_pairs) - low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) - v_vals = tl.reshape( - tl.join(low_vals, high_vals).permute(0, 2, 1), - (BLOCK_V, PAGE_SIZE // 2), - ) - - out_desc = tl.make_tensor_descriptor( - v_packed_ptr, - shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], - strides=[PAGE_SIZE // 2, 1], - block_shape=[BLOCK_V, PAGE_SIZE // 2], - ) - row_base = (page_idx * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V - out_desc.store([row_base.to(tl.int32), 0], v_vals) - - -@triton.jit -def _fp4_mla_attention_v_repack_pages_kernel( - v_packed_ptr, - kv_cache_ptr, - page_ids_ptr, - num_pages, - kv_s0, - kv_s2, - kv_s4, - V_HEAD_D: tl.constexpr, - PAGE_SIZE: tl.constexpr, - BLOCK_V: tl.constexpr, - occupancy: tl.constexpr = 1, -): - page_list_idx = tl.program_id(0) - dim_block = tl.program_id(1) - page_idx = tl.load(page_ids_ptr + page_list_idx).to(tl.int64) - - tl.assume(kv_s0 % 8 == 0) - tl.assume(kv_s2 % 8 == 0) - tl.assume(kv_s4 == 1) - v_view = tl.ext.make_view( - base=kv_cache_ptr, - shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], - tile_dim_map=[0, 1, 2], - ) - v_tile = tl.ext.load_view_tko( - v_view, - [ - page_idx.to(tl.int32), - 0, - (dim_block * (BLOCK_V // 2)).to(tl.int32), - ], - ) - v_tile = v_tile.to(tl.uint8, bitcast=True) - v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) - v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) - even_packed, odd_packed = tl.split(v_pairs) - low_vals, high_vals = _fp4_pack_nibbles(even_packed, odd_packed) - v_vals = tl.reshape( - tl.join(low_vals, high_vals).permute(0, 2, 1), - (BLOCK_V, PAGE_SIZE // 2), - ) - - out_desc = tl.make_tensor_descriptor( - v_packed_ptr, - shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], - strides=[PAGE_SIZE // 2, 1], - block_shape=[BLOCK_V, PAGE_SIZE // 2], - ) - row_base = (page_idx * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V - out_desc.store([row_base.to(tl.int32), 0], v_vals) - - -def fp4_mla_repack_v_cache( - v_packed: torch.Tensor, - kv_cache: torch.Tensor, - page_ids: Optional[torch.Tensor] = None, - *, - v_head_dim: int, - page_size: int, - block_v: int = 128, - kernel_occupancy: int = 2, - kernel_num_stages: int = 1, -) -> None: - """Populate the V-packed auxiliary cache consumed by the prepacked PV kernel.""" - if v_head_dim % block_v != 0: - raise ValueError(f"v_head_dim={v_head_dim} must be divisible by block_v={block_v}.") - if kv_cache.ndim < 5: - raise ValueError( - f"kv_cache must expose the paged FP4 layout, got shape={tuple(kv_cache.shape)}." - ) - num_pages = kv_cache.shape[0] - num_dim_blocks = triton.cdiv(v_head_dim, block_v) - launch_meta = { - "occupancy": int(kernel_occupancy), - "num_stages": int(kernel_num_stages), - } - if page_ids is None: - if num_pages == 0: - return - _fp4_mla_attention_v_repack_kernel[(num_pages, num_dim_blocks)]( - v_packed, - kv_cache, - num_pages, - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - V_HEAD_D=v_head_dim, - PAGE_SIZE=page_size, - BLOCK_V=block_v, - **launch_meta, - ) - return - - if page_ids.numel() == 0: - return - _fp4_mla_attention_v_repack_pages_kernel[(page_ids.numel(), num_dim_blocks)]( - v_packed, - kv_cache, - page_ids, - num_pages, - kv_cache.stride(0), - kv_cache.stride(2), - kv_cache.stride(4), - V_HEAD_D=v_head_dim, - PAGE_SIZE=page_size, - BLOCK_V=block_v, - **launch_meta, - ) - - -@triton.jit -def _fp4_mla_qk_scores_tile( - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - src_page_ids_ptr, - compact_page, - q_row_base, - head_start, - head_offsets, - token_offsets, - q_num_rows, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - page_ids_len, - num_pages, - Q_HEAD_D: tl.constexpr, - K_HEAD_D: tl.constexpr, - Q_RESIDUAL_D: tl.constexpr, - FP4_BLOCK: tl.constexpr, - Q_SF_PER_TOKEN: tl.constexpr, - K_SF_PER_TOKEN: tl.constexpr, - BLOCK_H: tl.constexpr, - BLOCK_T: tl.constexpr, - BLOCK_K: tl.constexpr, - FULL_BLOCK_END: tl.constexpr, - TAIL_BLOCK_K: tl.constexpr, - NUM_HEADS: tl.constexpr, - USE_TMA_DATA_LOAD: tl.constexpr, - ASSUME_FULL_HEADS: tl.constexpr, - ASSUME_VALID_PAGES: tl.constexpr, -): - if ASSUME_VALID_PAGES: - safe_compact_page = compact_page - physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) - safe_physical_page = physical_page - else: - valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) - safe_compact_page = tl.where(valid_compact_page, compact_page, 0) - physical_page = tl.load( - src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 - ).to(tl.int64) - valid_physical_page = ( - valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) - ) - safe_physical_page = tl.where(valid_physical_page, physical_page, 0) - q_rows = q_row_base + head_offsets - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_q_rows = q_rows - else: - mask_h = head_offsets < NUM_HEADS - safe_q_rows = tl.where(mask_h, q_rows, q_row_base) - scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) - - if ( - USE_TMA_DATA_LOAD - and ASSUME_FULL_HEADS - and ASSUME_VALID_PAGES - and Q_HEAD_D == 640 - and K_HEAD_D == 576 - and Q_RESIDUAL_D == 64 - and BLOCK_H == 128 - and BLOCK_T == 128 - and BLOCK_K == 512 - and FULL_BLOCK_END == 512 - and TAIL_BLOCK_K == 128 - ): - tl.assume(q_fp4_s0 % 8 == 0) - tl.assume(q_fp4_s1 == 1) - tl.assume(kv_s0 % 8 == 0) - tl.assume(kv_s2 % 8 == 0) - tl.assume(kv_s4 == 1) - q_desc = tl.make_tensor_descriptor( - q_fp4_ptr, - shape=[q_num_rows, Q_HEAD_D // 2], - strides=[q_fp4_s0, q_fp4_s1], - block_shape=[BLOCK_H, 256], - ) - k_desc = tl.make_tensor_descriptor( - kv_cache_ptr, - shape=[num_pages, BLOCK_T, K_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - block_shape=[1, BLOCK_T, 256], - ) - q_tail_desc = tl.make_tensor_descriptor( - q_fp4_ptr, - shape=[q_num_rows, Q_HEAD_D // 2], - strides=[q_fp4_s0, q_fp4_s1], - block_shape=[BLOCK_H, 64], - ) - k_tail_desc = tl.make_tensor_descriptor( - kv_cache_ptr, - shape=[num_pages, BLOCK_T, K_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - block_shape=[1, BLOCK_T, 32], - ) - q_sf_full_view = tl.ext.make_view( - base=q_sf_ptr, - shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], - strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, 8, 2, 256], - tile_dim_map=[0, 1, 2, 3], - ) - k_sf_full_view = tl.ext.make_view( - base=sf_cache_ptr, - shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], - strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, 1, 8, 2, 256], - tile_dim_map=[0, 1, 2, 3, 4], - ) - k_sf_tail_view = tl.ext.make_view( - base=sf_cache_ptr, - shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], - strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, 1, 1, 2, 256], - tile_dim_map=[0, 1, 2, 3, 4], - ) - - full_q_vals = q_desc.load([(q_row_base + head_start).to(tl.int32), 0]) - full_k_vals = k_desc.load([safe_physical_page.to(tl.int32), 0, 0]) - full_k_vals = tl.reshape(full_k_vals, (BLOCK_T, 256)) - q_row_group = q_row_base // 128 - full_q_scales = tl.ext.load_view_tko(q_sf_full_view, [q_row_group.to(tl.int32), 0, 0, 0]) - full_q_scales = full_q_scales.reshape([1, 8, 32, 4, 4]).trans(0, 3, 2, 1, 4) - full_q_scales = full_q_scales.reshape([BLOCK_H, 32]) - full_k_scales = tl.ext.load_view_tko( - k_sf_full_view, [safe_physical_page.to(tl.int32), 0, 0, 0, 0] - ) - full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) - full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) - scores = tl.dot_scaled( - full_q_vals, - full_q_scales, - "e2m1", - full_k_vals.T, - full_k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) - - tail_k_vals = k_tail_desc.load([safe_physical_page.to(tl.int32), 0, 256]) - tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) - tail_k_scales = tl.ext.load_view_tko( - k_sf_tail_view, [safe_physical_page.to(tl.int32), 0, 8, 0, 0] - ) - tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) - tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) - q_tail_vals = q_tail_desc.load([(q_row_base + head_start).to(tl.int32), 256]) - # Map Q tail groups [0, 1, ..., 7] onto K tail groups [0, 0, 1, 1, ..., 3, 3]. - q_tail_vals = q_tail_vals.reshape([BLOCK_H, 4, 2, 8]).trans(0, 1, 3, 2) - q_even_vals, q_odd_vals = tl.split(q_tail_vals) - q_even_vals = q_even_vals.reshape([BLOCK_H, 32]) - q_odd_vals = q_odd_vals.reshape([BLOCK_H, 32]) - - q_tail_sf_cols = 32 + tl.arange(0, 8) - q_tail_sf_offsets = _fp4_mla_swizzled_sf_offset( - q_rows[:, None], q_tail_sf_cols[None, :], Q_SF_PER_TOKEN - ) - q_tail_scales = tl.load(q_sf_ptr + q_tail_sf_offsets) - q_tail_scales = q_tail_scales.reshape([BLOCK_H, 4, 2]) - q_even_scales, q_odd_scales = tl.split(q_tail_scales) - scores = tl.dot_scaled( - q_even_vals, - q_even_scales, - "e2m1", - tail_k_vals.T, - tail_k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) - scores = tl.dot_scaled( - q_odd_vals, - q_odd_scales, - "e2m1", - tail_k_vals.T, - tail_k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) - return scores - - packed_k_offsets = tl.arange(0, BLOCK_K // 2) - scale_offsets = tl.arange(0, BLOCK_K // FP4_BLOCK) - residual_groups = Q_RESIDUAL_D // FP4_BLOCK - non_residual_groups = K_HEAD_D // FP4_BLOCK - residual_groups - if USE_TMA_DATA_LOAD and FULL_BLOCK_END > 0: - tl.assume(q_fp4_s0 % 8 == 0) - tl.assume(q_fp4_s1 == 1) - tl.assume(kv_s0 % 8 == 0) - tl.assume(kv_s2 % 8 == 0) - tl.assume(kv_s4 == 1) - q_desc = tl.make_tensor_descriptor( - q_fp4_ptr, - shape=[q_num_rows, Q_HEAD_D // 2], - strides=[q_fp4_s0, q_fp4_s1], - block_shape=[BLOCK_H, BLOCK_K // 2], - ) - k_desc = tl.make_tensor_descriptor( - kv_cache_ptr, - shape=[num_pages, BLOCK_T, K_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - block_shape=[1, BLOCK_T, BLOCK_K // 2], - ) - if USE_TMA_DATA_LOAD and Q_RESIDUAL_D == 64 and TAIL_BLOCK_K == 128: - q_tail_desc = tl.make_tensor_descriptor( - q_fp4_ptr, - shape=[q_num_rows, Q_HEAD_D // 2], - strides=[q_fp4_s0, q_fp4_s1], - block_shape=[BLOCK_H, 64], - ) - k_tail_desc = tl.make_tensor_descriptor( - kv_cache_ptr, - shape=[num_pages, BLOCK_T, K_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - block_shape=[1, BLOCK_T, 32], - ) - for q_start in tl.range(0, FULL_BLOCK_END, BLOCK_K): - q_elem_offsets = q_start + packed_k_offsets * 2 - q_group_offsets = q_elem_offsets // FP4_BLOCK - k_group_offsets = tl.where( - q_group_offsets < non_residual_groups, - q_group_offsets, - non_residual_groups + (q_group_offsets - non_residual_groups) // 2, - ) - byte_offsets_in_group = (q_elem_offsets % FP4_BLOCK) // 2 - packed_q_cols = q_start // 2 + packed_k_offsets - packed_k_cols = k_group_offsets * (FP4_BLOCK // 2) + byte_offsets_in_group - mask_k = q_elem_offsets < Q_HEAD_D - safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) - safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) - if ( - USE_TMA_DATA_LOAD - and FULL_BLOCK_END > 0 - and q_start + BLOCK_K <= non_residual_groups * FP4_BLOCK - ): - q_vals = q_desc.load([(q_row_base + head_start).to(tl.int32), q_start // 2]) - k_vals = k_desc.load([safe_physical_page.to(tl.int32), 0, q_start // 2]) - k_vals = tl.reshape(k_vals, (BLOCK_T, BLOCK_K // 2)) - if not ASSUME_VALID_PAGES: - k_vals = tl.where(valid_physical_page, k_vals, 0) - else: - q_vals = tl.load( - q_fp4_ptr - + safe_q_rows[:, None] * q_fp4_s0 - + safe_packed_q_cols[None, :] * q_fp4_s1, - mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], - other=0, - ) - k_vals = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + token_offsets[:, None].to(tl.int64) * kv_s2 - + safe_packed_k_cols[None, :] * kv_s4, - mask=mask_k[None, :] - if ASSUME_VALID_PAGES - else valid_physical_page & mask_k[None, :], - other=0, - ) - - q_sf_cols = q_start // FP4_BLOCK + scale_offsets - k_sf_cols = tl.where( - q_sf_cols < non_residual_groups, - q_sf_cols, - non_residual_groups + (q_sf_cols - non_residual_groups) // 2, - ) - mask_sf = q_sf_cols < Q_SF_PER_TOKEN - safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) - safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) - q_sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN - ) - k_sf_offsets = _fp4_mla_swizzled_sf_offset( - token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN - ) - q_scales = tl.load(q_sf_ptr + q_sf_offsets) - k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) - scores = tl.dot_scaled( - q_vals, - q_scales, - "e2m1", - k_vals.T, - k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) - - if FULL_BLOCK_END < Q_HEAD_D: - q_start = FULL_BLOCK_END - if Q_RESIDUAL_D == 64 and TAIL_BLOCK_K == 128: - residual_packed_offsets = tl.arange(0, 32) - residual_scale_offsets = tl.arange(0, 4) - packed_k_cols = non_residual_groups * (FP4_BLOCK // 2) + residual_packed_offsets - if USE_TMA_DATA_LOAD: - k_vals = k_tail_desc.load( - [ - safe_physical_page.to(tl.int32), - 0, - (non_residual_groups * (FP4_BLOCK // 2)).to(tl.int32), - ] - ) - k_vals = tl.reshape(k_vals, (BLOCK_T, 32)) - if not ASSUME_VALID_PAGES: - k_vals = tl.where(valid_physical_page, k_vals, 0) - elif ASSUME_VALID_PAGES: - k_vals = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + token_offsets[:, None].to(tl.int64) * kv_s2 - + packed_k_cols[None, :] * kv_s4, - ) - else: - k_vals = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + token_offsets[:, None].to(tl.int64) * kv_s2 - + packed_k_cols[None, :] * kv_s4, - mask=valid_physical_page, - other=0, - ) - k_sf_cols = non_residual_groups + residual_scale_offsets - k_sf_offsets = _fp4_mla_swizzled_sf_offset( - token_offsets[:, None], k_sf_cols[None, :], K_SF_PER_TOKEN - ) - k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) - - q_tail_cols = q_start // 2 + tl.arange(0, 64) - if USE_TMA_DATA_LOAD and ASSUME_FULL_HEADS: - q_tail_vals = q_tail_desc.load( - [(q_row_base + head_start).to(tl.int32), q_start // 2] - ) - elif ASSUME_FULL_HEADS: - q_tail_vals = tl.load( - q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1 - ) - else: - q_tail_vals = tl.load( - q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1, - mask=mask_h[:, None], - other=0, - ) - # Map Q tail groups [0, 1, ..., 7] onto K tail groups [0, 0, 1, 1, ..., 3, 3]. - q_tail_vals = q_tail_vals.reshape([BLOCK_H, 4, 2, 8]).trans(0, 1, 3, 2) - q_even_vals, q_odd_vals = tl.split(q_tail_vals) - q_even_vals = q_even_vals.reshape([BLOCK_H, 32]) - q_odd_vals = q_odd_vals.reshape([BLOCK_H, 32]) - - q_tail_sf_cols = q_start // FP4_BLOCK + tl.arange(0, 8) - q_tail_sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_q_rows[:, None], q_tail_sf_cols[None, :], Q_SF_PER_TOKEN - ) - q_tail_scales = tl.load(q_sf_ptr + q_tail_sf_offsets) - q_tail_scales = q_tail_scales.reshape([BLOCK_H, 4, 2]) - q_even_scales, q_odd_scales = tl.split(q_tail_scales) - scores = tl.dot_scaled( - q_even_vals, - q_even_scales, - "e2m1", - k_vals.T, - k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) - scores = tl.dot_scaled( - q_odd_vals, - q_odd_scales, - "e2m1", - k_vals.T, - k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) - else: - tail_packed_offsets = tl.arange(0, TAIL_BLOCK_K // 2) - tail_scale_offsets = tl.arange(0, TAIL_BLOCK_K // FP4_BLOCK) - q_elem_offsets = q_start + tail_packed_offsets * 2 - q_group_offsets = q_elem_offsets // FP4_BLOCK - k_group_offsets = tl.where( - q_group_offsets < non_residual_groups, - q_group_offsets, - non_residual_groups + (q_group_offsets - non_residual_groups) // 2, - ) - byte_offsets_in_group = (q_elem_offsets % FP4_BLOCK) // 2 - packed_q_cols = q_start // 2 + tail_packed_offsets - packed_k_cols = k_group_offsets * (FP4_BLOCK // 2) + byte_offsets_in_group - mask_k = q_elem_offsets < Q_HEAD_D - safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) - safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) - q_vals = tl.load( - q_fp4_ptr - + safe_q_rows[:, None] * q_fp4_s0 - + safe_packed_q_cols[None, :] * q_fp4_s1, - mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], - other=0, - ) - k_vals = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + token_offsets[:, None].to(tl.int64) * kv_s2 - + safe_packed_k_cols[None, :] * kv_s4, - mask=mask_k[None, :] - if ASSUME_VALID_PAGES - else valid_physical_page & mask_k[None, :], - other=0, - ) - - q_sf_cols = q_start // FP4_BLOCK + tail_scale_offsets - k_sf_cols = tl.where( - q_sf_cols < non_residual_groups, - q_sf_cols, - non_residual_groups + (q_sf_cols - non_residual_groups) // 2, - ) - mask_sf = q_sf_cols < Q_SF_PER_TOKEN - safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) - safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) - q_sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN - ) - k_sf_offsets = _fp4_mla_swizzled_sf_offset( - token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN - ) - q_scales = tl.load(q_sf_ptr + q_sf_offsets) - k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) - scores = tl.dot_scaled( - q_vals, - q_scales, - "e2m1", - k_vals.T, - k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) - - return scores - - -@triton.jit -def _fp4_mla_attention_stats_kernel( - max_ptr, - denom_ptr, - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - global_scale_ptr, - src_page_ids_ptr, - paged_kv_indptr_decode_ptr, - kv_lens_ptr, - page_ids_len, - num_pages, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - stats_s0, - q_num_rows, - sm_scale, - NUM_HEADS: tl.constexpr, - Q_HEAD_D: tl.constexpr, - K_HEAD_D: tl.constexpr, - Q_RESIDUAL_D: tl.constexpr, - PAGE_SIZE: tl.constexpr, - FP4_BLOCK: tl.constexpr, - Q_SF_PER_TOKEN: tl.constexpr, - K_SF_PER_TOKEN: tl.constexpr, - QUERY_LEN_PER_SEQ: tl.constexpr, - MAX_PAGES: tl.constexpr, - BLOCK_H: tl.constexpr, - BLOCK_T: tl.constexpr, - BLOCK_K: tl.constexpr, - FULL_BLOCK_END: tl.constexpr, - TAIL_BLOCK_K: tl.constexpr, - USE_TMA_DATA_LOAD: tl.constexpr, - ASSUME_FULL_HEADS: tl.constexpr, - ASSUME_FULL_PAGES: tl.constexpr, - ASSUME_VALID_PAGES: tl.constexpr, - occupancy: tl.constexpr = 1, -): - query_idx = tl.program_id(0) - head_block = tl.program_id(1) - seq_idx = query_idx // QUERY_LEN_PER_SEQ - query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ - - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_offs_h = offs_h - else: - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - offs_t = tl.arange(0, BLOCK_T) - q_row_base = query_idx * NUM_HEADS - kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) - kv_len = tl.maximum(kv_len, 0) - page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) - - max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) - denom = tl.zeros((BLOCK_H,), dtype=tl.float32) - global_scale = tl.load(global_scale_ptr) - qk_scale = sm_scale / (global_scale * global_scale) - for page_rel in tl.range(0, MAX_PAGES): - page_start = page_rel * PAGE_SIZE - if ASSUME_FULL_PAGES or page_start < kv_len: - compact_page = page_table_start + page_rel - scores = _fp4_mla_qk_scores_tile( - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - src_page_ids_ptr, - compact_page, - q_row_base, - head_block * BLOCK_H, - offs_h, - offs_t, - q_num_rows, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - page_ids_len, - num_pages, - Q_HEAD_D, - K_HEAD_D, - Q_RESIDUAL_D, - FP4_BLOCK, - Q_SF_PER_TOKEN, - K_SF_PER_TOKEN, - BLOCK_H, - BLOCK_T, - BLOCK_K, - FULL_BLOCK_END, - TAIL_BLOCK_K, - NUM_HEADS, - USE_TMA_DATA_LOAD, - ASSUME_FULL_HEADS, - ASSUME_VALID_PAGES, - ) - if ASSUME_FULL_PAGES: - scores = tl.where(mask_h[:, None], scores * qk_scale, -float("inf")) - else: - valid_t = page_start + offs_t < kv_len - scores = tl.where( - mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf") - ) - page_max = tl.max(scores, axis=1) - new_max = tl.maximum(max_score, page_max) - denom = denom * tl.math.exp2((max_score - new_max) * 1.4426950408889634) + tl.sum( - tl.math.exp2((scores - new_max[:, None]) * 1.4426950408889634), axis=1 - ) - max_score = new_max - - tl.store(max_ptr + query_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) - tl.store(denom_ptr + query_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) - - -@triton.jit -def _fp4_mla_attention_page_stats_kernel( - page_max_ptr, - page_sum_ptr, - p_fp4_ptr, - p_sf_ptr, - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - global_scale_ptr, - src_page_ids_ptr, - paged_kv_indptr_decode_ptr, - kv_lens_ptr, - page_ids_len, - num_pages, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - page_stats_s0, - page_stats_s1, - p_s0, - p_s1, - p_num_rows, - q_num_rows, - sm_scale, - NUM_HEADS: tl.constexpr, - Q_HEAD_D: tl.constexpr, - K_HEAD_D: tl.constexpr, - Q_RESIDUAL_D: tl.constexpr, - PAGE_SIZE: tl.constexpr, - FP4_BLOCK: tl.constexpr, - Q_SF_PER_TOKEN: tl.constexpr, - K_SF_PER_TOKEN: tl.constexpr, - SF_PER_PAGE: tl.constexpr, - P_GLOBAL_SCALE: tl.constexpr, - QUERY_LEN_PER_SEQ: tl.constexpr, - MAX_PAGES: tl.constexpr, - BLOCK_H: tl.constexpr, - BLOCK_T: tl.constexpr, - BLOCK_K: tl.constexpr, - FULL_BLOCK_END: tl.constexpr, - TAIL_BLOCK_K: tl.constexpr, - USE_TMA_DATA_LOAD: tl.constexpr, - PACK_PROBS: tl.constexpr, - ASSUME_FULL_HEADS: tl.constexpr, - ASSUME_FULL_PAGES: tl.constexpr, - ASSUME_VALID_PAGES: tl.constexpr, - occupancy: tl.constexpr = 1, -): - query_idx = tl.program_id(0) - head_block = tl.program_id(1) - page_rel = tl.program_id(2) - seq_idx = query_idx // QUERY_LEN_PER_SEQ - query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ - - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_offs_h = offs_h - else: - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - offs_t = tl.arange(0, BLOCK_T) - q_row_base = query_idx * NUM_HEADS - if ASSUME_FULL_PAGES: - kv_len = 0 - else: - kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) - kv_len = tl.maximum(kv_len, 0) - page_start = page_rel * PAGE_SIZE - page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) - out_offsets = query_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h - - page_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) - page_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) - if USE_TMA_DATA_LOAD and PACK_PROBS and ASSUME_FULL_HEADS and ASSUME_VALID_PAGES: - tl.assume(p_s0 % 8 == 0) - tl.assume(p_s1 == 1) - p_desc = tl.make_tensor_descriptor( - p_fp4_ptr, - shape=[p_num_rows, PAGE_SIZE // 2], - strides=[p_s0, p_s1], - block_shape=[BLOCK_H, PAGE_SIZE // 2], - ) - if ASSUME_FULL_PAGES or page_start < kv_len: - scores = _fp4_mla_qk_scores_tile( - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - src_page_ids_ptr, - page_table_start + page_rel, - q_row_base, - head_block * BLOCK_H, - offs_h, - offs_t, - q_num_rows, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - page_ids_len, - num_pages, - Q_HEAD_D, - K_HEAD_D, - Q_RESIDUAL_D, - FP4_BLOCK, - Q_SF_PER_TOKEN, - K_SF_PER_TOKEN, - BLOCK_H, - BLOCK_T, - BLOCK_K, - FULL_BLOCK_END, - TAIL_BLOCK_K, - NUM_HEADS, - USE_TMA_DATA_LOAD, - ASSUME_FULL_HEADS, - ASSUME_VALID_PAGES, - ) - if ASSUME_FULL_PAGES: - valid_t = tl.full([BLOCK_T], True, dtype=tl.int1) - else: - valid_t = page_start + offs_t < kv_len - global_scale = tl.load(global_scale_ptr) - qk_scale = sm_scale / (global_scale * global_scale) - if ASSUME_FULL_HEADS and ASSUME_FULL_PAGES: - scores = scores * qk_scale - page_max = tl.max(scores, axis=1) - exp_scores = tl.math.exp2((scores - page_max[:, None]) * 1.4426950408889634) - page_sum = tl.sum(exp_scores, axis=1) - else: - scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) - page_max = tl.max(scores, axis=1) - safe_page_max = tl.where(mask_h, page_max, 0.0) - exp_scores = tl.math.exp2((scores - safe_page_max[:, None]) * 1.4426950408889634) - exp_scores = tl.where(mask_h[:, None] & valid_t[None, :], exp_scores, 0.0) - page_sum = tl.sum(exp_scores, axis=1) - - if PACK_PROBS: - grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) - amax = tl.max(grouped_probs, axis=2) - inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) - stored_scale = tl.where( - amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0 - ) - scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) - pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) - even_probs, odd_probs = tl.split(pairs) - packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) - - if not ASSUME_VALID_PAGES: - valid_compact_page = (page_table_start + page_rel >= 0) & ( - page_table_start + page_rel < page_ids_len - ) - p_page = query_idx * MAX_PAGES + page_rel - p_rows = p_page * NUM_HEADS + offs_h - safe_p_rows = ( - p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) - ) - scale_cols = tl.arange(0, SF_PER_PAGE) - if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: - sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( - p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE - ) - else: - sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE - ) - if ASSUME_FULL_HEADS: - if ASSUME_VALID_PAGES: - tl.store(p_sf_ptr + sf_offsets, stored_scale) - else: - tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_compact_page) - else: - tl.store( - p_sf_ptr + sf_offsets, - stored_scale, - mask=mask_h[:, None] - if ASSUME_VALID_PAGES - else valid_compact_page & mask_h[:, None], - ) - - byte_offsets = tl.arange(0, FP4_BLOCK // 2) - byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] - if USE_TMA_DATA_LOAD and ASSUME_FULL_HEADS and ASSUME_VALID_PAGES: - p_desc.store( - [(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], - tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), - ) - elif ASSUME_FULL_HEADS: - if ASSUME_VALID_PAGES: - tl.store( - p_fp4_ptr - + safe_p_rows[:, None, None] * p_s0 - + byte_cols[None, :, :] * p_s1, - packed, - ) - else: - tl.store( - p_fp4_ptr - + safe_p_rows[:, None, None] * p_s0 - + byte_cols[None, :, :] * p_s1, - packed, - mask=valid_compact_page, - ) - else: - tl.store( - p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, - packed, - mask=mask_h[:, None, None] - if ASSUME_VALID_PAGES - else valid_compact_page & mask_h[:, None, None], - ) - - if ASSUME_FULL_HEADS: - tl.store(page_max_ptr + out_offsets, page_max) - tl.store(page_sum_ptr + out_offsets, page_sum) - else: - tl.store(page_max_ptr + out_offsets, page_max, mask=mask_h) - tl.store(page_sum_ptr + out_offsets, page_sum, mask=mask_h) - - -@triton.jit -def _fp4_mla_attention_page_stats_grouped_kernel( - page_max_ptr, - page_sum_ptr, - p_fp4_ptr, - p_sf_ptr, - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - global_scale_ptr, - src_page_ids_ptr, - paged_kv_indptr_decode_ptr, - kv_lens_ptr, - page_ids_len: tl.constexpr, - num_pages: tl.constexpr, - q_fp4_s0: tl.constexpr, - q_fp4_s1: tl.constexpr, - kv_s0: tl.constexpr, - kv_s2: tl.constexpr, - kv_s4: tl.constexpr, - sf_s0: tl.constexpr, - page_stats_s0: tl.constexpr, - page_stats_s1: tl.constexpr, - p_s0: tl.constexpr, - p_s1: tl.constexpr, - p_num_rows: tl.constexpr, - q_num_rows: tl.constexpr, - sm_scale: tl.constexpr, - NUM_HEADS: tl.constexpr, - Q_HEAD_D: tl.constexpr, - K_HEAD_D: tl.constexpr, - Q_RESIDUAL_D: tl.constexpr, - PAGE_SIZE: tl.constexpr, - FP4_BLOCK: tl.constexpr, - Q_SF_PER_TOKEN: tl.constexpr, - K_SF_PER_TOKEN: tl.constexpr, - SF_PER_PAGE: tl.constexpr, - P_GLOBAL_SCALE: tl.constexpr, - QUERY_LEN_PER_SEQ: tl.constexpr, - MAX_PAGES: tl.constexpr, - BLOCK_H: tl.constexpr, - BLOCK_T: tl.constexpr, - BLOCK_K: tl.constexpr, - FULL_BLOCK_END: tl.constexpr, - TAIL_BLOCK_K: tl.constexpr, - USE_TMA_DATA_LOAD: tl.constexpr, - PACK_PROBS: tl.constexpr, - GROUP_REDUCE_STATS: tl.constexpr, - ASSUME_FULL_HEADS: tl.constexpr, - ASSUME_FULL_PAGES: tl.constexpr, - ASSUME_VALID_PAGES: tl.constexpr, - GROUP_PAGES: tl.constexpr = 2, - occupancy: tl.constexpr = 1, -): - query_idx = tl.program_id(0) - head_block = tl.program_id(1) - page_group = tl.program_id(2) - seq_idx = query_idx // QUERY_LEN_PER_SEQ - - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - scale_cols = tl.arange(0, SF_PER_PAGE) - q_row_base = query_idx * NUM_HEADS - page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) - global_scale = tl.load(global_scale_ptr) - qk_scale = sm_scale / (global_scale * global_scale) - - tl.assume(q_fp4_s0 % 8 == 0) - tl.assume(q_fp4_s1 == 1) - tl.assume(kv_s0 % 8 == 0) - tl.assume(kv_s2 % 8 == 0) - tl.assume(kv_s4 == 1) - tl.assume(p_s0 % 8 == 0) - tl.assume(p_s1 == 1) - - q_desc = tl.make_tensor_descriptor( - q_fp4_ptr, - shape=[q_num_rows, Q_HEAD_D // 2], - strides=[q_fp4_s0, q_fp4_s1], - block_shape=[BLOCK_H, 256], - ) - k_desc = tl.make_tensor_descriptor( - kv_cache_ptr, - shape=[num_pages, BLOCK_T, K_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - block_shape=[1, BLOCK_T, 256], - ) - q_tail_desc = tl.make_tensor_descriptor( - q_fp4_ptr, - shape=[q_num_rows, Q_HEAD_D // 2], - strides=[q_fp4_s0, q_fp4_s1], - block_shape=[BLOCK_H, 32], - ) - k_tail_desc = tl.make_tensor_descriptor( - kv_cache_ptr, - shape=[num_pages, BLOCK_T, K_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - block_shape=[1, BLOCK_T, 32], - ) - p_desc = tl.make_tensor_descriptor( - p_fp4_ptr, - shape=[p_num_rows, PAGE_SIZE // 2], - strides=[p_s0, p_s1], - block_shape=[BLOCK_H, PAGE_SIZE // 2], - ) - q_sf_full_view = tl.ext.make_view( - base=q_sf_ptr, - shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], - strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, 8, 2, 256], - tile_dim_map=[0, 1, 2, 3], - ) - q_sf_tail_view = tl.ext.make_view( - base=q_sf_ptr, - shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], - strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, 1, 2, 256], - tile_dim_map=[0, 1, 2, 3], - ) - k_sf_full_view = tl.ext.make_view( - base=sf_cache_ptr, - shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], - strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, 1, 8, 2, 256], - tile_dim_map=[0, 1, 2, 3, 4], - ) - k_sf_tail_view = tl.ext.make_view( - base=sf_cache_ptr, - shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], - strides=[sf_s0, 128 * (((K_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, 1, 1, 2, 256], - tile_dim_map=[0, 1, 2, 3, 4], - ) - - q_row_start = (q_row_base + head_block * BLOCK_H).to(tl.int32) - q_row_group = q_row_base // 128 - full_q_vals = q_desc.load([q_row_start, 0]) - full_q_scales = tl.ext.load_view_tko(q_sf_full_view, [q_row_group.to(tl.int32), 0, 0, 0]) - full_q_scales = full_q_scales.reshape([1, 8, 32, 4, 4]).trans(0, 3, 2, 1, 4) - full_q_scales = full_q_scales.reshape([BLOCK_H, 32]) - q0_vals = q_tail_desc.load([q_row_start, 256]) - q1_vals = q_tail_desc.load([q_row_start, 288]) - q0_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group.to(tl.int32), 8, 0, 0]) - q0_scales = q0_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) - q0_scales = q0_scales.reshape([BLOCK_H, 4]) - q1_scales = tl.ext.load_view_tko(q_sf_tail_view, [q_row_group.to(tl.int32), 9, 0, 0]) - q1_scales = q1_scales.reshape([1, 1, 32, 4, 4]).trans(0, 3, 2, 1, 4) - q1_scales = q1_scales.reshape([BLOCK_H, 4]) - - group_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) - group_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) - for page_group_off in tl.range(0, GROUP_PAGES): - page_rel = page_group * GROUP_PAGES + page_group_off - compact_page = page_table_start + page_rel - physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) - - scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) - full_k_vals = k_desc.load([physical_page.to(tl.int32), 0, 0]) - full_k_vals = tl.reshape(full_k_vals, (BLOCK_T, 256)) - full_k_scales = tl.ext.load_view_tko( - k_sf_full_view, [physical_page.to(tl.int32), 0, 0, 0, 0] - ) - full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) - full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) - scores = tl.dot_scaled( - full_q_vals, - full_q_scales, - "e2m1", - full_k_vals.T, - full_k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) - - tail_k_vals = k_tail_desc.load([physical_page.to(tl.int32), 0, 256]) - tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) - tail_k_scales = tl.ext.load_view_tko( - k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0] - ) - tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) - tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) - scores = tl.dot_scaled( - q0_vals, - q0_scales, - "e2m1", - tail_k_vals.T, - tail_k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) - scores = tl.dot_scaled( - q1_vals, - q1_scales, - "e2m1", - tail_k_vals.T, - tail_k_scales, - "e2m1", - acc=scores, - fast_math=True, - rhs_k_pack=True, - ) - - scores = scores * qk_scale - page_max = tl.max(scores, axis=1) - exp_scores = tl.math.exp2((scores - page_max[:, None]) * 1.4426950408889634) - page_sum = tl.sum(exp_scores, axis=1) - if GROUP_REDUCE_STATS: - next_group_max = tl.maximum(group_max, page_max) - group_sum = group_sum * tl.math.exp2( - (group_max - next_group_max) * 1.4426950408889634 - ) + page_sum * tl.math.exp2((page_max - next_group_max) * 1.4426950408889634) - group_max = next_group_max - - grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) - amax = tl.max(grouped_probs, axis=2) - inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) - stored_scale = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) - scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) - pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) - even_probs, odd_probs = tl.split(pairs) - packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) - - p_page = query_idx * MAX_PAGES + page_rel - out_offsets = query_idx * page_stats_s0 + page_rel * page_stats_s1 + offs_h - tl.store(page_max_ptr + out_offsets, page_max) - if not GROUP_REDUCE_STATS: - tl.store(page_sum_ptr + out_offsets, page_sum) - - sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( - p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE - ) - tl.store(p_sf_ptr + sf_offsets, stored_scale) - p_desc.store( - [(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], - tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), - ) - if GROUP_REDUCE_STATS: - group_max_offsets = query_idx * page_stats_s0 + (page_group * 2) * page_stats_s1 + offs_h - group_sum_offsets = group_max_offsets + page_stats_s1 - tl.store(page_sum_ptr + group_max_offsets, group_max) - tl.store(page_sum_ptr + group_sum_offsets, group_sum) - - -@triton.jit -def _fp4_mla_attention_reduce_stats_kernel( - max_ptr, - denom_ptr, - page_max_ptr, - page_sum_ptr, - stats_s0, - page_stats_s0, - page_stats_s1, - NUM_HEADS: tl.constexpr, - MAX_PAGES: tl.constexpr, - BLOCK_H: tl.constexpr, - GROUP_REDUCE_STATS: tl.constexpr = False, - GROUP_PAGES: tl.constexpr = 1, - occupancy: tl.constexpr = 1, -): - gen_idx = tl.program_id(0) - head_block = tl.program_id(1) - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - - max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) - if GROUP_REDUCE_STATS: - for group_rel in tl.range(0, MAX_PAGES // GROUP_PAGES): - group_max = tl.load( - page_sum_ptr - + gen_idx * page_stats_s0 - + (group_rel * 2) * page_stats_s1 - + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - max_score = tl.maximum(max_score, group_max) - else: - for page_rel in tl.range(0, MAX_PAGES): - page_max = tl.load( - page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - max_score = tl.maximum(max_score, page_max) - - denom = tl.zeros((BLOCK_H,), dtype=tl.float32) - if GROUP_REDUCE_STATS: - for group_rel in tl.range(0, MAX_PAGES // GROUP_PAGES): - group_max = tl.load( - page_sum_ptr - + gen_idx * page_stats_s0 - + (group_rel * 2) * page_stats_s1 - + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - group_sum = tl.load( - page_sum_ptr - + gen_idx * page_stats_s0 - + (group_rel * 2 + 1) * page_stats_s1 - + safe_offs_h, - mask=mask_h, - other=0.0, - ) - denom += tl.where( - group_sum > 0.0, - group_sum * tl.math.exp2((group_max - max_score) * 1.4426950408889634), - 0.0, - ) - else: - for page_rel in tl.range(0, MAX_PAGES): - page_max = tl.load( - page_max_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - page_sum = tl.load( - page_sum_ptr + gen_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=0.0, - ) - denom += tl.where( - page_sum > 0.0, - page_sum * tl.math.exp2((page_max - max_score) * 1.4426950408889634), - 0.0, - ) - - tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) - tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) - - -@triton.jit -def _fp4_mla_attention_prob_scale_kernel( - p_sf_ptr, - max_ptr, - denom_ptr, - page_max_ptr, - paged_kv_indptr_decode_ptr, - kv_lens_ptr, - page_ids_len, - stats_s0, - page_stats_s0, - page_stats_s1, - NUM_HEADS: tl.constexpr, - PAGE_SIZE: tl.constexpr, - SF_PER_PAGE: tl.constexpr, - QUERY_LEN_PER_SEQ: tl.constexpr, - MAX_PAGES: tl.constexpr, - BLOCK_H: tl.constexpr, - ASSUME_FULL_HEADS: tl.constexpr, - ASSUME_FULL_PAGES: tl.constexpr, - ASSUME_VALID_PAGES: tl.constexpr, - occupancy: tl.constexpr = 1, -): - query_idx = tl.program_id(0) - head_block = tl.program_id(1) - page_rel = tl.program_id(2) - seq_idx = query_idx // QUERY_LEN_PER_SEQ - query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ - - page_start = page_rel * PAGE_SIZE - if ASSUME_FULL_PAGES: - kv_len = 0 - else: - kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) - kv_len = tl.maximum(kv_len, 0) - if page_start >= kv_len: - return - - page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) - compact_page = page_table_start + page_rel - if not ASSUME_VALID_PAGES: - if (compact_page < 0) | (compact_page >= page_ids_len): - return - - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_offs_h = offs_h - else: - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - page_max = tl.load( - page_max_ptr + query_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=-float("inf"), - ) - max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - factor = tl.where( - denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0 - ) - - p_page = query_idx * MAX_PAGES + page_rel - p_rows = p_page * NUM_HEADS + offs_h - safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) - scale_cols = tl.arange(0, SF_PER_PAGE) - if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: - sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( - p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE - ) - else: - sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE - ) - scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) - tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) - - -@triton.jit -def _fp4_mla_attention_prob_store_page_kernel( - probs_ptr, - max_ptr, - denom_ptr, - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - global_scale_ptr, - src_page_ids_ptr, - paged_kv_indptr_decode_ptr, - kv_lens_ptr, - page_rel, - page_ids_len, - num_pages, - probs_s0, - probs_s1, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - stats_s0, - q_num_rows, - sm_scale, - NUM_HEADS: tl.constexpr, - Q_HEAD_D: tl.constexpr, - K_HEAD_D: tl.constexpr, - Q_RESIDUAL_D: tl.constexpr, - PAGE_SIZE: tl.constexpr, - FP4_BLOCK: tl.constexpr, - Q_SF_PER_TOKEN: tl.constexpr, - K_SF_PER_TOKEN: tl.constexpr, - QUERY_LEN_PER_SEQ: tl.constexpr, - BLOCK_H: tl.constexpr, - BLOCK_K: tl.constexpr, - FULL_BLOCK_END: tl.constexpr, - TAIL_BLOCK_K: tl.constexpr, - USE_TMA_DATA_LOAD: tl.constexpr, - ASSUME_FULL_HEADS: tl.constexpr, - ASSUME_FULL_PAGES: tl.constexpr, - ASSUME_VALID_PAGES: tl.constexpr, - occupancy: tl.constexpr = 1, -): - query_idx = tl.program_id(0) - head_block = tl.program_id(1) - seq_idx = query_idx // QUERY_LEN_PER_SEQ - query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ - - kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) - kv_len = tl.maximum(kv_len, 0) - page_start = page_rel * PAGE_SIZE - if (not ASSUME_FULL_PAGES) and page_start >= kv_len: - return - - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_offs_h = offs_h - else: - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - offs_t = tl.arange(0, PAGE_SIZE) - if ASSUME_FULL_PAGES: - valid_t = tl.full([PAGE_SIZE], True, dtype=tl.int1) - else: - valid_t = page_start + offs_t < kv_len - page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) - compact_page = page_table_start + page_rel - if (compact_page < 0) | (compact_page >= page_ids_len): - return - q_row_base = query_idx * NUM_HEADS - - scores = _fp4_mla_qk_scores_tile( - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - src_page_ids_ptr, - compact_page, - q_row_base, - head_block * BLOCK_H, - offs_h, - offs_t, - q_num_rows, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - page_ids_len, - num_pages, - Q_HEAD_D, - K_HEAD_D, - Q_RESIDUAL_D, - FP4_BLOCK, - Q_SF_PER_TOKEN, - K_SF_PER_TOKEN, - BLOCK_H, - PAGE_SIZE, - BLOCK_K, - FULL_BLOCK_END, - TAIL_BLOCK_K, - NUM_HEADS, - USE_TMA_DATA_LOAD, - ASSUME_FULL_HEADS, - ASSUME_VALID_PAGES, - ) - max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) - global_scale = tl.load(global_scale_ptr) - qk_scale = sm_scale / (global_scale * global_scale) - denom_valid = denom > 0.0 - safe_denom = tl.where(denom_valid, denom, 1.0) - safe_max = tl.where(denom_valid, max_score, 0.0) - scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) - probs = tl.math.exp2((scores - safe_max[:, None]) * 1.4426950408889634) / safe_denom[:, None] - probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) - - prob_rows = query_idx * NUM_HEADS + offs_h - safe_prob_rows = tl.where(mask_h, prob_rows, query_idx * NUM_HEADS) - tl.store( - probs_ptr + safe_prob_rows[:, None] * probs_s0 + offs_t[None, :] * probs_s1, - probs, - mask=mask_h[:, None], - ) - - -@triton.jit -def _fp4_mla_attention_prob_pack_page_kernel( - p_fp4_ptr, - p_sf_ptr, - probs_ptr, - paged_kv_indptr_decode_ptr, - kv_lens_ptr, - page_rel, - page_ids_len, - p_s0, - p_s1, - probs_s0, - probs_s1, - NUM_HEADS: tl.constexpr, - PAGE_SIZE: tl.constexpr, - FP4_BLOCK: tl.constexpr, - SF_PER_PAGE: tl.constexpr, - P_GLOBAL_SCALE: tl.constexpr, - QUERY_LEN_PER_SEQ: tl.constexpr, - MAX_PAGES: tl.constexpr, - BLOCK_H: tl.constexpr, - ASSUME_FULL_HEADS: tl.constexpr, - ASSUME_FULL_PAGES: tl.constexpr, - occupancy: tl.constexpr = 1, -): - query_idx = tl.program_id(0) - token_group = tl.program_id(1) - head_block = tl.program_id(2) - seq_idx = query_idx // QUERY_LEN_PER_SEQ - query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ - - kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) - kv_len = tl.maximum(kv_len, 0) - page_start = page_rel * PAGE_SIZE - if (not ASSUME_FULL_PAGES) and page_start >= kv_len: - return - - page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) - compact_page = page_table_start + page_rel - if (compact_page < 0) | (compact_page >= page_ids_len): - return - - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - else: - mask_h = offs_h < NUM_HEADS - byte_offsets = tl.arange(0, FP4_BLOCK // 2) - token_base = token_group * FP4_BLOCK - even_t = token_base + byte_offsets * 2 - odd_t = even_t + 1 - valid_even = page_start + even_t < kv_len - valid_odd = page_start + odd_t < kv_len - - prob_rows = query_idx * NUM_HEADS + offs_h - safe_prob_rows = tl.where(mask_h, prob_rows, query_idx * NUM_HEADS) - even_probs = tl.load( - probs_ptr + safe_prob_rows[:, None] * probs_s0 + even_t[None, :] * probs_s1, - mask=mask_h[:, None] & valid_even[None, :], - other=0.0, - ) - odd_probs = tl.load( - probs_ptr + safe_prob_rows[:, None] * probs_s0 + odd_t[None, :] * probs_s1, - mask=mask_h[:, None] & valid_odd[None, :], - other=0.0, - ) - amax = tl.maximum(tl.max(tl.abs(even_probs), axis=1), tl.max(tl.abs(odd_probs), axis=1)) - local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) - stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) - - p_page = query_idx * MAX_PAGES + page_rel - p_rows = p_page * NUM_HEADS + offs_h - safe_p_rows = tl.where(mask_h, p_rows, p_page * NUM_HEADS) - sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows, token_group, SF_PER_PAGE) - tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h) - - even_quant = _fp4_e2m1_quantize(even_probs / local_scale[:, None]) - odd_quant = _fp4_e2m1_quantize(odd_probs / local_scale[:, None]) - packed = even_quant | (odd_quant << 4) - byte_cols = token_group * (FP4_BLOCK // 2) + byte_offsets - tl.store( - p_fp4_ptr + safe_p_rows[:, None] * p_s0 + byte_cols[None, :] * p_s1, - packed, - mask=mask_h[:, None], - ) - - -@triton.jit -def _fp4_mla_attention_prob_pack_page_fused_kernel( - p_fp4_ptr, - p_sf_ptr, - max_ptr, - denom_ptr, - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - global_scale_ptr, - src_page_ids_ptr, - paged_kv_indptr_decode_ptr, - kv_lens_ptr, - page_rel, - page_ids_len, - num_pages, - p_s0, - p_s1, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - stats_s0, - q_num_rows, - sm_scale, - NUM_HEADS: tl.constexpr, - Q_HEAD_D: tl.constexpr, - K_HEAD_D: tl.constexpr, - Q_RESIDUAL_D: tl.constexpr, - PAGE_SIZE: tl.constexpr, - FP4_BLOCK: tl.constexpr, - Q_SF_PER_TOKEN: tl.constexpr, - K_SF_PER_TOKEN: tl.constexpr, - SF_PER_PAGE: tl.constexpr, - P_GLOBAL_SCALE: tl.constexpr, - QUERY_LEN_PER_SEQ: tl.constexpr, - MAX_PAGES: tl.constexpr, - BLOCK_H: tl.constexpr, - BLOCK_K: tl.constexpr, - FULL_BLOCK_END: tl.constexpr, - TAIL_BLOCK_K: tl.constexpr, - USE_TMA_DATA_LOAD: tl.constexpr, - PAGE_REL_FROM_GRID: tl.constexpr = False, - ASSUME_FULL_HEADS: tl.constexpr = False, - ASSUME_FULL_PAGES: tl.constexpr = False, - ASSUME_VALID_PAGES: tl.constexpr = False, - occupancy: tl.constexpr = 1, -): - query_idx = tl.program_id(0) - head_block = tl.program_id(1) - if PAGE_REL_FROM_GRID: - page_rel = tl.program_id(2) - seq_idx = query_idx // QUERY_LEN_PER_SEQ - query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ - - kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) - kv_len = tl.maximum(kv_len, 0) - page_start = page_rel * PAGE_SIZE - if (not ASSUME_FULL_PAGES) and page_start >= kv_len: - return - - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_offs_h = offs_h - else: - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - offs_t = tl.arange(0, PAGE_SIZE) - if ASSUME_FULL_PAGES: - valid_t = tl.full([PAGE_SIZE], True, dtype=tl.int1) - else: - valid_t = page_start + offs_t < kv_len - page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) - compact_page = page_table_start + page_rel - if (compact_page < 0) | (compact_page >= page_ids_len): - return - q_row_base = query_idx * NUM_HEADS - - scores = _fp4_mla_qk_scores_tile( - q_fp4_ptr, - q_sf_ptr, - kv_cache_ptr, - sf_cache_ptr, - src_page_ids_ptr, - compact_page, - q_row_base, - head_block * BLOCK_H, - offs_h, - offs_t, - q_num_rows, - q_fp4_s0, - q_fp4_s1, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - page_ids_len, - num_pages, - Q_HEAD_D, - K_HEAD_D, - Q_RESIDUAL_D, - FP4_BLOCK, - Q_SF_PER_TOKEN, - K_SF_PER_TOKEN, - BLOCK_H, - PAGE_SIZE, - BLOCK_K, - FULL_BLOCK_END, - TAIL_BLOCK_K, - NUM_HEADS, - USE_TMA_DATA_LOAD, - ASSUME_FULL_HEADS, - ASSUME_VALID_PAGES, - ) - max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=1.0) - global_scale = tl.load(global_scale_ptr) - qk_scale = sm_scale / (global_scale * global_scale) - denom_valid = denom > 0.0 - safe_denom = tl.where(denom_valid, denom, 1.0) - safe_max = tl.where(denom_valid, max_score, 0.0) - scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) - probs = tl.math.exp2((scores - safe_max[:, None]) * 1.4426950408889634) / safe_denom[:, None] - probs = tl.where(mask_h[:, None] & valid_t[None, :] & denom_valid[:, None], probs, 0.0) - - grouped_probs = tl.reshape(probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) - amax = tl.max(tl.abs(grouped_probs), axis=2) - local_scale = tl.where(amax > 0.0, amax / 6.0, 1.0) - stored_scale = tl.where(amax > 0.0, tl.minimum(local_scale * P_GLOBAL_SCALE, 448.0), 1.0) - scaled_probs = grouped_probs / tl.reshape(local_scale, (BLOCK_H, SF_PER_PAGE, 1)) - pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) - even_probs, odd_probs = tl.split(pairs) - packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) - - p_page = query_idx * MAX_PAGES + page_rel - p_rows = p_page * NUM_HEADS + offs_h - safe_p_rows = tl.where(mask_h, p_rows, p_page * NUM_HEADS) - scale_cols = tl.arange(0, SF_PER_PAGE) - sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) - tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=mask_h[:, None]) - - byte_offsets = tl.arange(0, FP4_BLOCK // 2) - byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] - tl.store( - p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, - packed, - mask=mask_h[:, None, None], - ) - - -@triton.jit -def _fp4_mla_attention_pv_kernel( - out_ptr, - p_fp4_ptr, - p_sf_ptr, - kv_cache_ptr, - v_sf_ptr, - v_packed_ptr, - global_scale_ptr, - src_page_ids_ptr, - paged_kv_indptr_decode_ptr, - kv_lens_ptr, - page_ids_len, - num_pages, - out_s0, - out_s1, - out_s2, - out_num_rows, - p_s0, - p_s1, - p_num_rows, - kv_s0, - kv_s2, - kv_s4, - vsf_s0, - NUM_HEADS: tl.constexpr, - V_HEAD_D: tl.constexpr, - PAGE_SIZE: tl.constexpr, - FP4_BLOCK: tl.constexpr, - SF_PER_PAGE: tl.constexpr, - QUERY_LEN_PER_SEQ: tl.constexpr, - MAX_PAGES: tl.constexpr, - P_GLOBAL_SCALE: tl.constexpr, - BLOCK_H: tl.constexpr, - BLOCK_V: tl.constexpr, - USE_TMA_P_LOAD: tl.constexpr, - USE_TMA_V_LOAD: tl.constexpr, - USE_PREPACKED_V: tl.constexpr, - PV_LOOP_STAGES: tl.constexpr, - ASSUME_FULL_HEADS: tl.constexpr, - ASSUME_FULL_PAGES: tl.constexpr, - ASSUME_FULL_V: tl.constexpr, - ASSUME_VALID_PAGES: tl.constexpr, - occupancy: tl.constexpr = 1, -): - query_idx = tl.program_id(0) - head_block = tl.program_id(1) - dim_block = tl.program_id(2) - seq_idx = query_idx // QUERY_LEN_PER_SEQ - query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ - - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_offs_h = offs_h - else: - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - if ASSUME_FULL_V: - mask_v = tl.full([BLOCK_V], True, dtype=tl.int1) - safe_offs_v = offs_v - else: - mask_v = offs_v < V_HEAD_D - safe_offs_v = tl.where(mask_v, offs_v, 0) - packed_t = tl.arange(0, PAGE_SIZE // 2) - scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) - even_t = packed_t * 2 - odd_t = even_t + 1 - v_packed_offsets = safe_offs_v // 2 - v_use_high_nibble = (safe_offs_v & 1) != 0 - if ASSUME_FULL_V and BLOCK_V == 128: - v_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( - dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE - ) - else: - v_sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE - ) - if USE_TMA_P_LOAD: - tl.assume(p_s0 % 8 == 0) - tl.assume(p_s1 == 1) - p_desc = tl.make_tensor_descriptor( - p_fp4_ptr, - shape=[p_num_rows, PAGE_SIZE // 2], - strides=[p_s0, p_s1], - block_shape=[BLOCK_H, PAGE_SIZE // 2], - ) - if USE_TMA_V_LOAD and ASSUME_FULL_HEADS and ASSUME_FULL_V: - tl.assume(out_s1 % 8 == 0) - tl.assume(out_s2 == 1) - out_desc = tl.make_tensor_descriptor( - out_ptr, - shape=[out_num_rows, V_HEAD_D], - strides=[out_s1, out_s2], - block_shape=[BLOCK_H, BLOCK_V], - ) - if USE_TMA_V_LOAD: - tl.assume(kv_s0 % 8 == 0) - tl.assume(kv_s2 % 8 == 0) - tl.assume(kv_s4 == 1) - v_desc = tl.make_tensor_descriptor( - kv_cache_ptr, - shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - block_shape=[1, PAGE_SIZE, BLOCK_V // 2], - ) - if USE_PREPACKED_V: - v_packed_desc = tl.make_tensor_descriptor( - v_packed_ptr, - shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], - strides=[PAGE_SIZE // 2, 1], - block_shape=[BLOCK_V, PAGE_SIZE // 2], - ) - - if ( - USE_TMA_P_LOAD - and USE_TMA_V_LOAD - and ASSUME_FULL_HEADS - and ASSUME_FULL_PAGES - and ASSUME_FULL_V - and ASSUME_VALID_PAGES - and NUM_HEADS == 128 - and V_HEAD_D == 512 - and PAGE_SIZE == 128 - and BLOCK_H == 128 - and BLOCK_V == 128 - and SF_PER_PAGE == 8 - ): - p_view = tl.ext.make_view( - base=p_fp4_ptr, - shapes=[p_num_rows, PAGE_SIZE // 2], - strides=[p_s0, p_s1], - tile_shape=[BLOCK_H, PAGE_SIZE // 2], - tile_dim_map=[0, 1], - ) - p_sf_view = tl.ext.make_view( - base=p_sf_ptr, - shapes=[p_num_rows // 128, SF_PER_PAGE // 4, 2, 256], - strides=[128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, SF_PER_PAGE // 4, 2, 256], - tile_dim_map=[0, 1, 2, 3], - ) - v_sf_view = tl.ext.make_view( - base=v_sf_ptr, - shapes=[num_pages, V_HEAD_D // 128, SF_PER_PAGE // 4, 2, 256], - strides=[vsf_s0, 128 * (((SF_PER_PAGE + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, 1, SF_PER_PAGE // 4, 2, 256], - tile_dim_map=[0, 1, 2, 3, 4], - ) - if not USE_PREPACKED_V: - v_view = tl.ext.make_view( - base=kv_cache_ptr, - shapes=[num_pages, PAGE_SIZE, V_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - tile_shape=[1, PAGE_SIZE, BLOCK_V // 2], - tile_dim_map=[0, 1, 2], - ) - page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) - global_scale = tl.load(global_scale_ptr) - out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) - acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) - for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): - compact_page = page_table_start + page_rel - physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) - p_page = query_idx * MAX_PAGES + page_rel - - p_vals = tl.ext.load_view_tko( - p_view, - [(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0], - ) - p_vals = p_vals.to(tl.uint8, bitcast=True) - p_scales = tl.ext.load_view_tko(p_sf_view, [p_page.to(tl.int32), 0, 0, 0]) - p_scales = p_scales.reshape([1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 3, 2, 1, 4) - p_scales = p_scales.reshape([BLOCK_H, SF_PER_PAGE]) - - if USE_PREPACKED_V: - v_row = (physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V - v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) - else: - v_tile = tl.ext.load_view_tko( - v_view, - [ - physical_page.to(tl.int32), - 0, - (dim_block * (BLOCK_V // 2)).to(tl.int32), - ], - ) - v_tile = v_tile.to(tl.uint8, bitcast=True) - v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) - v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) - even_packed, odd_packed = tl.split(v_pairs) - low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) - high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) - v_vals = tl.reshape( - tl.join(low_vals, high_vals).permute(0, 2, 1), - (BLOCK_V, PAGE_SIZE // 2), - ) - v_scales = tl.ext.load_view_tko( - v_sf_view, - [ - physical_page.to(tl.int32), - dim_block, - 0, - 0, - 0, - ], - ) - v_scales = v_scales.reshape([1, 1, SF_PER_PAGE // 4, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) - v_scales = v_scales.reshape([BLOCK_V, SF_PER_PAGE]) - acc = tl.dot_scaled( - p_vals, - p_scales, - "e2m1", - v_vals.T, - v_scales, - "e2m1", - acc=acc, - fast_math=True, - rhs_k_pack=True, - ) - - out_vals = acc * out_scale - if out_ptr.dtype.element_ty == tl.bfloat16: - out_vals = out_vals.to(tl.bfloat16) - elif out_ptr.dtype.element_ty == tl.float16: - out_vals = out_vals.to(tl.float16) - out_desc.store( - [ - (query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), - (dim_block * BLOCK_V).to(tl.int32), - ], - out_vals, - ) - return - - if ASSUME_FULL_PAGES: - kv_len = 0 - else: - kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) - kv_len = tl.maximum(kv_len, 0) - page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) - global_scale = tl.load(global_scale_ptr) - out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) - acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) - for page_rel in tl.range(0, MAX_PAGES, num_stages=PV_LOOP_STAGES): - page_start = page_rel * PAGE_SIZE - if ASSUME_FULL_PAGES or page_start < kv_len: - compact_page = page_table_start + page_rel - if ASSUME_VALID_PAGES: - safe_compact_page = compact_page - physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) - safe_physical_page = physical_page - else: - valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) - safe_compact_page = tl.where(valid_compact_page, compact_page, 0) - physical_page = tl.load( - src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 - ).to(tl.int64) - valid_physical_page = ( - valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) - ) - safe_physical_page = tl.where(valid_physical_page, physical_page, 0) - - p_page = query_idx * MAX_PAGES + page_rel - p_rows = p_page * NUM_HEADS + offs_h - safe_p_rows = ( - p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) - ) - if USE_TMA_P_LOAD: - p_vals = p_desc.load([(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) - else: - p_vals = tl.load( - p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, - mask=mask_h[:, None] - if ASSUME_VALID_PAGES - else valid_compact_page & mask_h[:, None], - other=0, - ) - p_sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE - ) - p_scales = tl.load(p_sf_ptr + p_sf_offsets) - - if ASSUME_FULL_PAGES: - valid_even_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) - valid_odd_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) - else: - valid_even_t = page_start + even_t < kv_len - valid_odd_t = page_start + odd_t < kv_len - if USE_PREPACKED_V: - v_row = (safe_physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V - v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) - elif USE_TMA_V_LOAD: - v_tile = v_desc.load( - [ - safe_physical_page.to(tl.int32), - 0, - (dim_block * (BLOCK_V // 2)).to(tl.int32), - ] - ) - v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) - if not ASSUME_VALID_PAGES: - v_tile = tl.where(valid_physical_page, v_tile, 0) - v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) - even_packed, odd_packed = tl.split(v_pairs) - if not ASSUME_FULL_PAGES: - even_packed = tl.where(valid_even_t[None, :], even_packed, 0) - odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) - low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) - high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) - v_vals = tl.reshape( - tl.join(low_vals, high_vals).permute(0, 2, 1), - (BLOCK_V, PAGE_SIZE // 2), - ) - else: - even_packed = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + even_t[None, :].to(tl.int64) * kv_s2 - + v_packed_offsets[:, None] * kv_s4, - mask=(mask_v[:, None] & valid_even_t[None, :]) - if ASSUME_VALID_PAGES - else (valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), - other=0, - ) - odd_packed = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + odd_t[None, :].to(tl.int64) * kv_s2 - + v_packed_offsets[:, None] * kv_s4, - mask=(mask_v[:, None] & valid_odd_t[None, :]) - if ASSUME_VALID_PAGES - else (valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), - other=0, - ) - even_low = even_packed & 0x0F - even_high = (even_packed >> 4) & 0x0F - odd_low = odd_packed & 0x0F - even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) - odd_nibble = tl.where(v_use_high_nibble[:, None], odd_packed >> 4, odd_low) - v_vals = even_nibble | (odd_nibble << 4) - v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) - acc = tl.dot_scaled( - p_vals, - p_scales, - "e2m1", - v_vals.T, - v_scales, - "e2m1", - acc=acc, - fast_math=True, - rhs_k_pack=True, - ) - - if ASSUME_FULL_HEADS and ASSUME_FULL_V: - out_vals = acc * out_scale - if USE_TMA_V_LOAD: - if out_ptr.dtype.element_ty == tl.bfloat16: - out_vals = out_vals.to(tl.bfloat16) - elif out_ptr.dtype.element_ty == tl.float16: - out_vals = out_vals.to(tl.float16) - out_desc.store( - [ - (query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), - (dim_block * BLOCK_V).to(tl.int32), - ], - out_vals, - ) - else: - tl.store( - out_ptr + query_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, - out_vals, - ) - else: - tl.store( - out_ptr - + query_idx * out_s0 - + safe_offs_h[:, None] * out_s1 - + safe_offs_v[None, :] * out_s2, - acc * out_scale, - mask=mask_h[:, None] & mask_v[None, :], - ) - - -def fp4_mla_paged_attention_internal( - q_fp4: torch.Tensor, - q_sf: torch.Tensor, - kv_cache: torch.Tensor, - sf_cache: torch.Tensor, - v_sf: torch.Tensor, - global_scale: torch.Tensor, - src_page_ids: torch.Tensor, - paged_kv_indptr_decode: torch.Tensor, - kv_lens: torch.Tensor, - output: Optional[torch.Tensor] = None, - *, - sm_scale: float, - num_heads: Optional[int] = None, - v_head_dim: Optional[int] = None, - page_size: Optional[int] = None, - q_residual_dim: int = 0, - p_global_scale: float = FP4_MLA_P_GLOBAL_SCALE, - block_h: int = 128, - block_k: Optional[int] = None, - block_v: int = 128, - output_dtype: torch.dtype = torch.bfloat16, - max_pages: Optional[int] = None, - page_pipeline_streams: Optional[int] = None, - kernel_occupancy: Optional[int] = None, - kernel_num_ctas: Optional[int] = None, - kernel_num_stages: Optional[int] = None, - kernel_num_warps: Optional[int] = None, - pv_loop_stages: int = 1, - parallel_page_stats: Optional[bool] = None, - fused_prob_pack: Optional[bool] = None, - use_tma_data_load: Optional[bool] = None, - fused_prob_pack_single_launch: Optional[bool] = None, - pack_prob_in_page_stats: Optional[bool] = None, - page_stats_group_size: Optional[int] = None, - assume_full_pages: Optional[bool] = None, - assume_valid_pages: Optional[bool] = None, - query_len_per_seq: int = 1, - prepack_v_for_pv: bool = False, - use_prepacked_v_for_pv: bool = False, - p_fp4_workspace: Optional[torch.Tensor] = None, - p_sf_workspace: Optional[torch.Tensor] = None, - v_packed_workspace: Optional[torch.Tensor] = None, - p_probs_workspace: Optional[torch.Tensor] = None, - max_scores_workspace: Optional[torch.Tensor] = None, - denom_workspace: Optional[torch.Tensor] = None, - page_max_workspace: Optional[torch.Tensor] = None, - page_sum_workspace: Optional[torch.Tensor] = None, - **kwargs, -) -> torch.Tensor: - del kwargs - if not hasattr(tl, "dot_scaled"): - raise NotImplementedError( - "fp4_mla_paged_attention requires a Triton build with tl.dot_scaled." - ) - if not q_fp4.is_cuda: - raise ValueError("q_fp4 must be a CUDA tensor.") - if q_fp4.dtype != torch.uint8 or kv_cache.dtype != torch.uint8: - raise TypeError("q_fp4 and kv_cache must be packed FP4 tensors with dtype torch.uint8.") - if global_scale.numel() < 1: - raise ValueError("global_scale must contain at least one element.") - if q_fp4.dim() == 3: - inferred_num_queries, inferred_num_heads, packed_q_dim = q_fp4.shape - if num_heads is not None and num_heads != inferred_num_heads: - raise ValueError( - f"num_heads={num_heads} does not match q_fp4.shape[1]={inferred_num_heads}." - ) - num_queries = inferred_num_queries - num_heads = inferred_num_heads - q_fp4_2d = q_fp4.reshape(num_queries * num_heads, packed_q_dim) - elif q_fp4.dim() == 2: - if num_heads is None: - raise ValueError("num_heads is required when q_fp4 is 2D.") - if q_fp4.shape[0] % num_heads != 0: - raise ValueError("q_fp4.shape[0] must be divisible by num_heads.") - num_queries = q_fp4.shape[0] // num_heads - packed_q_dim = q_fp4.shape[1] - q_fp4_2d = q_fp4 - else: - raise ValueError("q_fp4 must be 2D or 3D.") - if query_len_per_seq <= 0: - raise ValueError(f"query_len_per_seq must be positive, got {query_len_per_seq}.") - if num_queries % query_len_per_seq != 0: - raise ValueError( - "q_fp4 query rows must be divisible by query_len_per_seq, got " - f"{num_queries} rows and query_len_per_seq={query_len_per_seq}." - ) - num_gen_seqs = num_queries // query_len_per_seq - - num_pages, inferred_page_size, packed_k_dim, kv_s0, kv_s2, kv_s4 = _get_kv_cache_strides( - kv_cache - ) - if page_size is None: - page_size = inferred_page_size - if page_size != inferred_page_size: - raise ValueError( - f"page_size={page_size} does not match kv_cache page dimension {inferred_page_size}." - ) - if page_size % FP4_BLOCK_SIZE != 0: - raise ValueError(f"page_size must be divisible by {FP4_BLOCK_SIZE}.") - - q_head_dim = packed_q_dim * 2 - k_head_dim = packed_k_dim * 2 - if q_residual_dim < 0 or q_residual_dim % FP4_BLOCK_SIZE != 0: - raise ValueError(f"q_residual_dim must be a non-negative multiple of {FP4_BLOCK_SIZE}.") - if q_head_dim - q_residual_dim != k_head_dim: - raise ValueError( - f"q_head_dim - q_residual_dim must match K head dim: {q_head_dim} - {q_residual_dim} != {k_head_dim}." - ) - if q_head_dim % FP4_BLOCK_SIZE != 0 or k_head_dim % FP4_BLOCK_SIZE != 0: - raise ValueError(f"Q/K head dims must be divisible by {FP4_BLOCK_SIZE}.") - if v_head_dim is None: - v_head_dim = k_head_dim - if v_head_dim <= 0 or v_head_dim > k_head_dim: - raise ValueError(f"v_head_dim must be in (0, {k_head_dim}], got {v_head_dim}.") - q_sf_flat = q_sf.contiguous().view(-1) - if sf_cache.shape[0] < num_pages or v_sf.shape[0] < num_pages: - raise ValueError("sf_cache and v_sf must have a leading physical-page dimension.") - if q_sf_flat.numel() < _swizzled_scale_size(num_queries * num_heads, q_head_dim): - raise ValueError("q_sf is too small for the swizzled Q scale layout.") - if sf_cache.numel() < sf_cache.shape[0] * _swizzled_scale_size(page_size, k_head_dim): - raise ValueError("sf_cache is too small for the swizzled K scale layout.") - if v_sf.numel() < v_sf.shape[0] * _swizzled_scale_size(v_head_dim, page_size): - raise ValueError("v_sf is too small for the swizzled V scale layout.") - - if output is None: - output = torch.empty( - (num_queries, num_heads, v_head_dim), dtype=output_dtype, device=q_fp4.device - ) - elif output.shape != (num_queries, num_heads, v_head_dim): - raise ValueError( - f"output must have shape {(num_queries, num_heads, v_head_dim)}, got {tuple(output.shape)}." - ) - - if num_queries == 0: - return output - triton_backend = "nvt" - env_block_h = _env_int("TRTLLM_FP4_MLA_BLOCK_H") - env_block_k = _env_int("TRTLLM_FP4_MLA_BLOCK_K") - env_block_v = _env_int("TRTLLM_FP4_MLA_BLOCK_V") - env_pv_loop_stages = _env_int("TRTLLM_FP4_MLA_PV_LOOP_STAGES") - env_occupancy = _env_int("TRTLLM_FP4_MLA_OCCUPANCY") - env_num_warps = _env_int("TRTLLM_FP4_MLA_NUM_WARPS") - env_num_stages = _env_int("TRTLLM_FP4_MLA_NUM_STAGES") - env_group_pages = _env_int("TRTLLM_FP4_MLA_GROUP_PAGES") - env_group_reduce_stats = _env_int("TRTLLM_FP4_MLA_GROUP_REDUCE_STATS") - env_page_pipeline_streams = _env_int("TRTLLM_FP4_MLA_PAGE_PIPELINE_STREAMS") - if env_block_h is not None: - block_h = env_block_h - if block_k is None: - block_k = env_block_k or (512 if triton_backend == "nvt" else 256) - elif env_block_k is not None: - block_k = env_block_k - if env_block_v is not None: - block_v = env_block_v - if env_pv_loop_stages is not None: - pv_loop_stages = env_pv_loop_stages - full_block_end = (q_head_dim // block_k) * block_k - tail_k = q_head_dim - full_block_end - tail_block_k = 1 << (tail_k - 1).bit_length() if tail_k > 0 else block_k - if max_pages is None: - if paged_kv_indptr_decode.numel() >= num_gen_seqs + 1: - page_counts = ( - paged_kv_indptr_decode[1 : num_gen_seqs + 1] - paged_kv_indptr_decode[:num_gen_seqs] - ) - max_pages = int(page_counts.max().item()) if page_counts.numel() > 0 else 0 - else: - max_pages = _ceil_div(int(kv_lens[:num_gen_seqs].max().item()), page_size) - if max_pages <= 0: - output.zero_() - return output - - q_sf_per_token = q_head_dim // FP4_BLOCK_SIZE - k_sf_per_token = k_head_dim // FP4_BLOCK_SIZE - sf_per_page = page_size // FP4_BLOCK_SIZE - block_t = page_size - num_head_blocks = triton.cdiv(num_heads, block_h) - assume_full_heads = num_heads % block_h == 0 - assume_full_v = v_head_dim % block_v == 0 - if assume_full_pages is None: - assume_full_pages = False - assume_full_pages = bool(assume_full_pages) and query_len_per_seq == 1 - if assume_valid_pages is None: - assume_valid_pages = False - assume_valid_pages = bool(assume_valid_pages) - if ( - not assume_valid_pages - and assume_full_pages - and src_page_ids.numel() == num_gen_seqs * max_pages - ): - # Exactly-sized full-page decode tables can skip compact-page and - # physical-page validity masks. Keeping this inside the kernel wrapper - # lets call sites stay conservative. - assume_valid_pages = True - total_p_rows = max(num_queries * max_pages * num_heads, 1) - if page_pipeline_streams is None and env_page_pipeline_streams is not None: - page_pipeline_streams = env_page_pipeline_streams - if page_pipeline_streams is None: - if triton_backend == "nvt" and max_pages >= 8 and num_queries >= 128: - page_pipeline_streams = 2 - else: - page_pipeline_streams = 1 - page_pipeline_streams = max(1, min(int(page_pipeline_streams), max_pages)) - launch_meta = {} - if kernel_occupancy is None: - if env_occupancy is not None: - kernel_occupancy = env_occupancy - elif triton_backend == "nvt": - kernel_occupancy = 2 - if kernel_occupancy is not None: - launch_meta["occupancy"] = int(kernel_occupancy) - if kernel_num_ctas is not None: - launch_meta["num_ctas"] = int(kernel_num_ctas) - if kernel_num_stages is None and env_num_stages is not None: - kernel_num_stages = env_num_stages - if kernel_num_stages is not None: - launch_meta["num_stages"] = int(kernel_num_stages) - if kernel_num_warps is None and env_num_warps is not None: - kernel_num_warps = env_num_warps - if kernel_num_warps is not None: - launch_meta["num_warps"] = int(kernel_num_warps) - if fused_prob_pack is None: - fused_prob_pack = triton_backend == "nvt" - if fused_prob_pack_single_launch is None: - fused_prob_pack_single_launch = triton_backend == "nvt" and max_pages >= 8 - if use_tma_data_load is None: - use_tma_data_load = triton_backend == "nvt" - use_tma_data_load = bool(use_tma_data_load and hasattr(tl, "make_tensor_descriptor")) - if use_tma_data_load: - # Device-side descriptors may need Triton's allocator for descriptor scratch storage. - def alloc_fn(size: int, alignment: int, stream: Optional[int]): - return torch.empty(size, device=q_fp4.device, dtype=torch.int8) - - triton.set_allocator(alloc_fn) - - p_fp4 = _workspace_tensor( - p_fp4_workspace, - (total_p_rows, page_size // 2), - dtype=torch.uint8, - device=q_fp4.device, - name="p_fp4", - ) - p_sf = _workspace_tensor( - p_sf_workspace, - (_swizzled_scale_size(total_p_rows, page_size),), - dtype=q_sf.dtype, - device=q_fp4.device, - name="p_sf", - ) - num_dim_blocks = triton.cdiv(v_head_dim, block_v) - v_repack_block_v = block_v - num_repack_dim_blocks = num_dim_blocks - auto_prepack_v_for_pv = ( - triton_backend == "nvt" - and use_tma_data_load - and assume_full_heads - and assume_full_pages - and assume_valid_pages - and v_head_dim == 512 - and page_size == 128 - and block_h in (64, 128) - and block_v == 128 - and sf_per_page == 8 - ) - if not prepack_v_for_pv and not use_prepacked_v_for_pv: - env_prepack_v = os.environ.get("TRTLLM_FP4_MLA_PREPACK_V") - prepack_v_for_pv = auto_prepack_v_for_pv if env_prepack_v is None else env_prepack_v == "1" - wants_prepacked_v_for_pv = bool(prepack_v_for_pv) or bool(use_prepacked_v_for_pv) - if use_prepacked_v_for_pv and v_packed_workspace is None: - raise ValueError("use_prepacked_v_for_pv requires v_packed_workspace to be provided.") - can_use_prepacked_v_for_pv = ( - wants_prepacked_v_for_pv - and use_tma_data_load - and assume_full_heads - and assume_full_pages - and assume_valid_pages - and v_head_dim == 512 - and page_size == 128 - and block_h in (64, 128) - and block_v == 128 - and sf_per_page == 8 - ) - if wants_prepacked_v_for_pv and not can_use_prepacked_v_for_pv: - raise ValueError( - "prepacked V PV path requires TMA, full heads/pages, valid pages, " - "v_head_dim=512, page_size=128, block_h in (64, 128), " - "block_v=128, and sf_per_page=8." - ) - if can_use_prepacked_v_for_pv: - v_packed = _workspace_tensor( - v_packed_workspace, - (num_pages * num_dim_blocks * block_v, page_size // 2), - dtype=torch.uint8, - device=q_fp4.device, - name="v_packed", - ) - else: - v_packed = kv_cache - if fused_prob_pack: - p_probs = None - else: - p_probs_shape = (max(num_queries * num_heads, 1), page_size) - if page_pipeline_streams > 1: - p_probs = _workspace_tensor( - p_probs_workspace, - (page_pipeline_streams, *p_probs_shape), - dtype=torch.float32, - device=q_fp4.device, - name="p_probs", - ) - else: - p_probs = _workspace_tensor( - p_probs_workspace, - p_probs_shape, - dtype=torch.float32, - device=q_fp4.device, - name="p_probs", - ) - max_scores = _workspace_tensor( - max_scores_workspace, - (num_queries, num_heads), - dtype=torch.float32, - device=q_fp4.device, - name="max_scores", - ) - denom = _workspace_tensor( - denom_workspace, - (num_queries, num_heads), - dtype=torch.float32, - device=q_fp4.device, - name="denom", - ) - - v_repack_stream = None - if bool(prepack_v_for_pv) and can_use_prepacked_v_for_pv: - current_stream = torch.cuda.current_stream(q_fp4.device) - v_repack_stream = torch.cuda.Stream(device=q_fp4.device) - v_repack_stream.wait_stream(current_stream) - with torch.cuda.stream(v_repack_stream): - _fp4_mla_attention_v_repack_kernel[(num_pages, num_repack_dim_blocks)]( - v_packed, - kv_cache, - num_pages, - kv_s0, - kv_s2, - kv_s4, - V_HEAD_D=v_head_dim, - PAGE_SIZE=page_size, - BLOCK_V=v_repack_block_v, - **launch_meta, - ) - - if parallel_page_stats is None: - parallel_page_stats = triton_backend == "nvt" and max_pages >= 8 - if pack_prob_in_page_stats is None: - pack_prob_in_page_stats = parallel_page_stats and fused_prob_pack - pack_prob_in_page_stats = bool( - pack_prob_in_page_stats and parallel_page_stats and fused_prob_pack - ) - page_stats_group_sizes = (2, 4, 8, 16, 32, 64, 128, 256, 512, 1024) - if page_stats_group_size is None and env_group_pages is not None: - page_stats_group_size = env_group_pages - if page_stats_group_size is None: - page_stats_group_size = 1 - else: - page_stats_group_size = int(page_stats_group_size) - can_group_page_stats = ( - page_stats_group_size in page_stats_group_sizes - and triton_backend == "nvt" - and parallel_page_stats - and pack_prob_in_page_stats - and use_tma_data_load - and assume_full_heads - and assume_full_pages - and assume_valid_pages - and query_len_per_seq == 1 - and num_heads == 128 - and q_head_dim == 640 - and k_head_dim == 576 - and q_residual_dim == 64 - and page_size == 128 - and block_h == 128 - and block_k == 512 - and full_block_end == 512 - and tail_block_k == 128 - and sf_per_page == 8 - and max_pages % page_stats_group_size == 0 - ) - page_stats_group_size = page_stats_group_size if can_group_page_stats else 1 - group_reduce_stats = ( - env_group_reduce_stats != 0 - if env_group_reduce_stats is not None - else triton_backend == "nvt" - ) and page_stats_group_size > 1 - if parallel_page_stats: - page_stats_shape = (num_queries, max_pages, num_heads) - page_max = _workspace_tensor( - page_max_workspace, - page_stats_shape, - dtype=torch.float32, - device=q_fp4.device, - name="page_max", - ) - page_sum = _workspace_tensor( - page_sum_workspace, - page_stats_shape, - dtype=torch.float32, - device=q_fp4.device, - name="page_sum", - ) - if page_stats_group_size > 1: - _fp4_mla_attention_page_stats_grouped_kernel[ - (num_queries, num_head_blocks, max_pages // page_stats_group_size) - ]( - page_max, - page_sum, - p_fp4, - p_sf, - q_fp4_2d, - q_sf_flat, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], - num_pages, - q_fp4_2d.stride(0), - q_fp4_2d.stride(1), - kv_s0, - kv_s2, - kv_s4, - sf_cache.stride(0), - page_max.stride(0), - page_max.stride(1), - p_fp4.stride(0), - p_fp4.stride(1), - p_fp4.shape[0], - q_fp4_2d.shape[0], - sm_scale=sm_scale, - NUM_HEADS=num_heads, - Q_HEAD_D=q_head_dim, - K_HEAD_D=k_head_dim, - Q_RESIDUAL_D=q_residual_dim, - PAGE_SIZE=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - Q_SF_PER_TOKEN=q_sf_per_token, - K_SF_PER_TOKEN=k_sf_per_token, - SF_PER_PAGE=sf_per_page, - P_GLOBAL_SCALE=p_global_scale, - QUERY_LEN_PER_SEQ=query_len_per_seq, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - BLOCK_T=block_t, - BLOCK_K=block_k, - FULL_BLOCK_END=full_block_end, - TAIL_BLOCK_K=tail_block_k, - USE_TMA_DATA_LOAD=use_tma_data_load, - PACK_PROBS=pack_prob_in_page_stats, - GROUP_REDUCE_STATS=group_reduce_stats, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - GROUP_PAGES=page_stats_group_size, - **launch_meta, - ) - else: - _fp4_mla_attention_page_stats_kernel[(num_queries, num_head_blocks, max_pages)]( - page_max, - page_sum, - p_fp4, - p_sf, - q_fp4_2d, - q_sf_flat, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], - num_pages, - q_fp4_2d.stride(0), - q_fp4_2d.stride(1), - kv_s0, - kv_s2, - kv_s4, - sf_cache.stride(0), - page_max.stride(0), - page_max.stride(1), - p_fp4.stride(0), - p_fp4.stride(1), - p_fp4.shape[0], - q_fp4_2d.shape[0], - sm_scale=sm_scale, - NUM_HEADS=num_heads, - Q_HEAD_D=q_head_dim, - K_HEAD_D=k_head_dim, - Q_RESIDUAL_D=q_residual_dim, - PAGE_SIZE=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - Q_SF_PER_TOKEN=q_sf_per_token, - K_SF_PER_TOKEN=k_sf_per_token, - SF_PER_PAGE=sf_per_page, - P_GLOBAL_SCALE=p_global_scale, - QUERY_LEN_PER_SEQ=query_len_per_seq, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - BLOCK_T=block_t, - BLOCK_K=block_k, - FULL_BLOCK_END=full_block_end, - TAIL_BLOCK_K=tail_block_k, - USE_TMA_DATA_LOAD=use_tma_data_load, - PACK_PROBS=pack_prob_in_page_stats, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, - ) - _fp4_mla_attention_reduce_stats_kernel[(num_queries, num_head_blocks)]( - max_scores, - denom, - page_max, - page_sum, - max_scores.stride(0), - page_max.stride(0), - page_max.stride(1), - NUM_HEADS=num_heads, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - GROUP_REDUCE_STATS=group_reduce_stats, - GROUP_PAGES=page_stats_group_size, - **launch_meta, - ) - if pack_prob_in_page_stats: - _fp4_mla_attention_prob_scale_kernel[(num_queries, num_head_blocks, max_pages)]( - p_sf, - max_scores, - denom, - page_max, - paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], - max_scores.stride(0), - page_max.stride(0), - page_max.stride(1), - NUM_HEADS=num_heads, - PAGE_SIZE=page_size, - SF_PER_PAGE=sf_per_page, - QUERY_LEN_PER_SEQ=query_len_per_seq, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, - ) - else: - _fp4_mla_attention_stats_kernel[(num_queries, num_head_blocks)]( - max_scores, - denom, - q_fp4_2d, - q_sf_flat, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], - num_pages, - q_fp4_2d.stride(0), - q_fp4_2d.stride(1), - kv_s0, - kv_s2, - kv_s4, - sf_cache.stride(0), - max_scores.stride(0), - q_fp4_2d.shape[0], - sm_scale=sm_scale, - NUM_HEADS=num_heads, - Q_HEAD_D=q_head_dim, - K_HEAD_D=k_head_dim, - Q_RESIDUAL_D=q_residual_dim, - PAGE_SIZE=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - Q_SF_PER_TOKEN=q_sf_per_token, - K_SF_PER_TOKEN=k_sf_per_token, - QUERY_LEN_PER_SEQ=query_len_per_seq, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - BLOCK_T=block_t, - BLOCK_K=block_k, - FULL_BLOCK_END=full_block_end, - TAIL_BLOCK_K=tail_block_k, - USE_TMA_DATA_LOAD=use_tma_data_load, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, - ) - - def _launch_prob_page(page_rel: int, p_probs_slot: Optional[torch.Tensor] = None): - if fused_prob_pack: - _fp4_mla_attention_prob_pack_page_fused_kernel[(num_queries, num_head_blocks)]( - p_fp4, - p_sf, - max_scores, - denom, - q_fp4_2d, - q_sf_flat, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - page_rel, - src_page_ids.shape[0], - num_pages, - p_fp4.stride(0), - p_fp4.stride(1), - q_fp4_2d.stride(0), - q_fp4_2d.stride(1), - kv_s0, - kv_s2, - kv_s4, - sf_cache.stride(0), - max_scores.stride(0), - q_fp4_2d.shape[0], - sm_scale=sm_scale, - NUM_HEADS=num_heads, - Q_HEAD_D=q_head_dim, - K_HEAD_D=k_head_dim, - Q_RESIDUAL_D=q_residual_dim, - PAGE_SIZE=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - Q_SF_PER_TOKEN=q_sf_per_token, - K_SF_PER_TOKEN=k_sf_per_token, - SF_PER_PAGE=sf_per_page, - P_GLOBAL_SCALE=p_global_scale, - QUERY_LEN_PER_SEQ=query_len_per_seq, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - BLOCK_K=block_k, - FULL_BLOCK_END=full_block_end, - TAIL_BLOCK_K=tail_block_k, - USE_TMA_DATA_LOAD=use_tma_data_load, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, - ) - return - assert p_probs_slot is not None - _fp4_mla_attention_prob_store_page_kernel[(num_queries, num_head_blocks)]( - p_probs_slot, - max_scores, - denom, - q_fp4_2d, - q_sf_flat, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - page_rel, - src_page_ids.shape[0], - num_pages, - p_probs_slot.stride(0), - p_probs_slot.stride(1), - q_fp4_2d.stride(0), - q_fp4_2d.stride(1), - kv_s0, - kv_s2, - kv_s4, - sf_cache.stride(0), - max_scores.stride(0), - q_fp4_2d.shape[0], - sm_scale=sm_scale, - NUM_HEADS=num_heads, - Q_HEAD_D=q_head_dim, - K_HEAD_D=k_head_dim, - Q_RESIDUAL_D=q_residual_dim, - PAGE_SIZE=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - Q_SF_PER_TOKEN=q_sf_per_token, - K_SF_PER_TOKEN=k_sf_per_token, - QUERY_LEN_PER_SEQ=query_len_per_seq, - BLOCK_H=block_h, - BLOCK_K=block_k, - FULL_BLOCK_END=full_block_end, - TAIL_BLOCK_K=tail_block_k, - USE_TMA_DATA_LOAD=use_tma_data_load, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, - ) - _fp4_mla_attention_prob_pack_page_kernel[(num_queries, sf_per_page, num_head_blocks)]( - p_fp4, - p_sf, - p_probs_slot, - paged_kv_indptr_decode, - kv_lens, - page_rel, - src_page_ids.shape[0], - p_fp4.stride(0), - p_fp4.stride(1), - p_probs_slot.stride(0), - p_probs_slot.stride(1), - NUM_HEADS=num_heads, - PAGE_SIZE=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - SF_PER_PAGE=sf_per_page, - P_GLOBAL_SCALE=p_global_scale, - QUERY_LEN_PER_SEQ=query_len_per_seq, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - **launch_meta, - ) - - if pack_prob_in_page_stats: - pass - elif fused_prob_pack and fused_prob_pack_single_launch: - _fp4_mla_attention_prob_pack_page_fused_kernel[(num_queries, num_head_blocks, max_pages)]( - p_fp4, - p_sf, - max_scores, - denom, - q_fp4_2d, - q_sf_flat, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - 0, - src_page_ids.shape[0], - num_pages, - p_fp4.stride(0), - p_fp4.stride(1), - q_fp4_2d.stride(0), - q_fp4_2d.stride(1), - kv_s0, - kv_s2, - kv_s4, - sf_cache.stride(0), - max_scores.stride(0), - q_fp4_2d.shape[0], - sm_scale=sm_scale, - NUM_HEADS=num_heads, - Q_HEAD_D=q_head_dim, - K_HEAD_D=k_head_dim, - Q_RESIDUAL_D=q_residual_dim, - PAGE_SIZE=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - Q_SF_PER_TOKEN=q_sf_per_token, - K_SF_PER_TOKEN=k_sf_per_token, - SF_PER_PAGE=sf_per_page, - P_GLOBAL_SCALE=p_global_scale, - QUERY_LEN_PER_SEQ=query_len_per_seq, - MAX_PAGES=max_pages, - BLOCK_H=block_h, - BLOCK_K=block_k, - FULL_BLOCK_END=full_block_end, - TAIL_BLOCK_K=tail_block_k, - USE_TMA_DATA_LOAD=use_tma_data_load, - PAGE_REL_FROM_GRID=True, - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, - ) - elif page_pipeline_streams == 1: - for page_rel in range(max_pages): - _launch_prob_page(page_rel, p_probs) - else: - current_stream = torch.cuda.current_stream(q_fp4.device) - streams = [torch.cuda.Stream(device=q_fp4.device) for _ in range(page_pipeline_streams)] - for stream in streams: - stream.wait_stream(current_stream) - for page_rel in range(max_pages): - stream_idx = page_rel % page_pipeline_streams - with torch.cuda.stream(streams[stream_idx]): - if fused_prob_pack: - _launch_prob_page(page_rel) - else: - assert p_probs is not None - _launch_prob_page(page_rel, p_probs[stream_idx]) - for stream in streams: - current_stream.wait_stream(stream) - - if v_repack_stream is not None: - torch.cuda.current_stream(q_fp4.device).wait_stream(v_repack_stream) - - _fp4_mla_attention_pv_kernel[ - ( - num_queries, - num_head_blocks, - num_dim_blocks, - ) - ]( - output, - p_fp4, - p_sf, - kv_cache, - v_sf, - v_packed, - global_scale, - src_page_ids, - paged_kv_indptr_decode, - kv_lens, - src_page_ids.shape[0], - num_pages, - output.stride(0), - output.stride(1), - output.stride(2), - output.shape[0] * output.shape[1], - p_fp4.stride(0), - p_fp4.stride(1), - p_fp4.shape[0], - kv_s0, - kv_s2, - kv_s4, - v_sf.stride(0), - NUM_HEADS=num_heads, - V_HEAD_D=v_head_dim, - PAGE_SIZE=page_size, - FP4_BLOCK=FP4_BLOCK_SIZE, - SF_PER_PAGE=sf_per_page, - QUERY_LEN_PER_SEQ=query_len_per_seq, - MAX_PAGES=max_pages, - P_GLOBAL_SCALE=p_global_scale, - BLOCK_H=block_h, - BLOCK_V=block_v, - USE_TMA_P_LOAD=use_tma_data_load and assume_full_heads and assume_valid_pages, - USE_TMA_V_LOAD=use_tma_data_load and v_head_dim % block_v == 0, - USE_PREPACKED_V=can_use_prepacked_v_for_pv, - PV_LOOP_STAGES=int(pv_loop_stages), - ASSUME_FULL_HEADS=assume_full_heads, - ASSUME_FULL_PAGES=assume_full_pages, - ASSUME_FULL_V=assume_full_v, - ASSUME_VALID_PAGES=assume_valid_pages, - **launch_meta, - ) - return output - - -fp4_mla_paged_attention = fp4_mla_paged_attention_internal diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py index c1f853553d38..63ab10f36398 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py @@ -566,7 +566,7 @@ def _fp4_mla_v_scale_store_context_tokens_kernel( tl.max(tl.abs(odd_values), axis=1), ) tile_amax = tl.max(amax_per_token, axis=0) - global_scale = tl.load(global_scale_ptr) + kv_global_scale = tl.load(global_scale_ptr) # K consumes scales as [token, dim-block], while V consumes scales as # [dim, token-block]. Only the compressed-KV prefix has both views, so # tail K-only dims keep K's per-token scale. @@ -574,8 +574,11 @@ def _fp4_mla_v_scale_store_context_tokens_kernel( tile_scale = tl.where(tile_amax > 0.0, tile_amax / 6.0, 1.0) token_scale = tl.where(amax_per_token > 0.0, amax_per_token / 6.0, 1.0) local_scale = tl.where(shared_tile, tile_scale, token_scale) - stored_scale = local_scale * global_scale - v_stored_scale = tile_scale * global_scale + # A capped page amax (or a block above the static reference amax) can push an + # outlier block's e4m3 scale above the e4m3 ceiling (448); clamp so it clips + # gracefully instead of overflowing the scale's e4m3 representation. + stored_scale = tl.minimum(local_scale * kv_global_scale, 448.0) + v_stored_scale = tl.minimum(tile_scale * kv_global_scale, 448.0) low = _fp4_e2m1_quantize(even_values / local_scale[:, None]) high = _fp4_e2m1_quantize(odd_values / local_scale[:, None]) diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py index 6fd1920e19e4..07d09ab007f4 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py @@ -22,9 +22,14 @@ perfect (NUM_HEADS=128, BLOCK_H=128) case used by page-stats and PV. * ``tl.assume()`` stride hints in front of every ``make_tensor_descriptor`` call -- helps Triton vectorize TMA loads. +* Optional prepacked-V PV path using ``tl.make_tensor_descriptor`` only. It + stores V as ``[page, dim-block, v, packed-token-pair]`` so PV can skip the + per-query V transpose/repack inside the page loop. * Pipelined PV loop via ``tl.range(..., num_stages=PV_LOOP_STAGES)``. """ +from typing import Any, Optional + import triton import triton.language as tl @@ -155,6 +160,170 @@ def _fp4_e2m1_quantize_packed(even, odd): ) +@triton.jit +def _fp4_mla_attention_v_repack_kernel( + v_packed_ptr, + kv_cache_ptr, + num_pages, + kv_s0, + kv_s2, + kv_s4, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BLOCK_V: tl.constexpr, + occupancy: tl.constexpr = 1, +): + page_idx = tl.program_id(0) + dim_block = tl.program_id(1) + + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, PAGE_SIZE, BLOCK_V // 2], + ) + v_tile = v_desc.load( + [ + page_idx.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ] + ) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) + high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + + out_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + row_base = (page_idx * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + out_desc.store([row_base.to(tl.int32), 0], v_vals) + + +@triton.jit +def _fp4_mla_attention_v_repack_pages_kernel( + v_packed_ptr, + kv_cache_ptr, + page_ids_ptr, + num_pages, + kv_s0, + kv_s2, + kv_s4, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BLOCK_V: tl.constexpr, + occupancy: tl.constexpr = 1, +): + page_list_idx = tl.program_id(0) + dim_block = tl.program_id(1) + page_idx = tl.load(page_ids_ptr + page_list_idx).to(tl.int64) + + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, PAGE_SIZE, BLOCK_V // 2], + ) + v_tile = v_desc.load( + [ + page_idx.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ] + ) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) + high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + + out_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + row_base = (page_idx * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + out_desc.store([row_base.to(tl.int32), 0], v_vals) + + +def fp4_mla_repack_v_cache( + v_packed: Any, + kv_cache: Any, + page_ids: Optional[Any] = None, + *, + v_head_dim: int, + page_size: int, + block_v: int = 128, + kernel_occupancy: int = 8, + kernel_num_stages: int = 1, +) -> None: + """Populate the public-Triton V-packed auxiliary cache.""" + if v_head_dim % block_v != 0: + raise ValueError(f"v_head_dim={v_head_dim} must be divisible by block_v={block_v}.") + if kv_cache.ndim < 5: + raise ValueError( + f"kv_cache must expose the paged FP4 layout, got shape={tuple(kv_cache.shape)}." + ) + num_pages = kv_cache.shape[0] + num_dim_blocks = triton.cdiv(v_head_dim, block_v) + launch_meta = { + "occupancy": int(kernel_occupancy), + "num_stages": int(kernel_num_stages), + } + if page_ids is None: + if num_pages == 0: + return + _fp4_mla_attention_v_repack_kernel[(num_pages, num_dim_blocks)]( + v_packed, + kv_cache, + num_pages, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + BLOCK_V=block_v, + **launch_meta, + ) + return + + if page_ids.numel() == 0: + return + _fp4_mla_attention_v_repack_pages_kernel[(page_ids.numel(), num_dim_blocks)]( + v_packed, + kv_cache, + page_ids, + num_pages, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + V_HEAD_D=v_head_dim, + PAGE_SIZE=page_size, + BLOCK_V=block_v, + **launch_meta, + ) + + @triton.jit def _fp4_mla_qk_scores_tile( q_fp4_ptr, @@ -500,7 +669,6 @@ def _fp4_mla_attention_page_stats_kernel( sf_cache_ptr, global_scale_ptr, q_global_scale_ptr, - page_scale_ptr, src_page_ids_ptr, paged_kv_indptr_decode_ptr, kv_lens_ptr, @@ -519,8 +687,6 @@ def _fp4_mla_attention_page_stats_kernel( p_num_rows, q_num_rows, sm_scale, - local_layer, - pscale_s0, NUM_HEADS: tl.constexpr, Q_HEAD_D: tl.constexpr, K_HEAD_D: tl.constexpr, @@ -543,7 +709,6 @@ def _fp4_mla_attention_page_stats_kernel( ASSUME_FULL_HEADS: tl.constexpr, ASSUME_FULL_PAGES: tl.constexpr, ASSUME_VALID_PAGES: tl.constexpr, - USE_PER_PAGE_SCALE: tl.constexpr = False, occupancy: tl.constexpr = 1, ): """Page-stats fused QK + softmax-stats + (optional) FP4 P pack. @@ -630,29 +795,12 @@ def _fp4_mla_attention_page_stats_kernel( else: valid_t = page_start + offs_t < kv_len global_scale = tl.load(global_scale_ptr) - if USE_PER_PAGE_SCALE: - # Independent dynamic Q scale and per-page (K/V shared) KV scale. - # page_gscale also folds into the stored P scale below so the per-page - # V scaling cancels inside the fused PV dot. - compact_page = page_table_start + page_rel - if ASSUME_VALID_PAGES: - phys_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) - else: - valid_cp = (compact_page >= 0) & (compact_page < page_ids_len) - phys_page = tl.load( - src_page_ids_ptr + tl.where(valid_cp, compact_page, 0), - mask=valid_cp, - other=0, - ).to(tl.int64) - phys_page = tl.where((phys_page >= 0) & (phys_page < num_pages), phys_page, 0) - page_gscale = tl.load(page_scale_ptr + local_layer * pscale_s0 + phys_page) - q_gscale = tl.load(q_global_scale_ptr) - qk_scale = sm_scale / (q_gscale * page_gscale) - else: - # Static scale: global_scale == 1.0, so page_gscale == 1.0 makes the - # stored-P fold below a no-op and qk_scale == sm_scale. - page_gscale = global_scale - qk_scale = sm_scale / (global_scale * global_scale) + # Static per-layer scales: QK encodes q * q_gscale and k * kv_gscale, + # so divide the scores by (q_gscale * kv_gscale). When Q and KV share + # one global scale, this reduces to global_scale^2. + # PV applies the final KV global_scale divisor separately. + q_gscale = tl.load(q_global_scale_ptr) + qk_scale = sm_scale / (q_gscale * global_scale) if ASSUME_FULL_HEADS and ASSUME_FULL_PAGES: scores = scores * qk_scale page_max = tl.max(scores, axis=1) @@ -670,13 +818,11 @@ def _fp4_mla_attention_page_stats_kernel( grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) amax = tl.max(grouped_probs, axis=2) inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) - # Fold 1/page_gscale into the stored P block scale (page_gscale == 1.0 - # in the static path, so this is a no-op there). Combined with V's - # baked page_gscale, the page scale cancels in the PV dot and the end - # out_scale = 1/(global_scale*P_GLOBAL_SCALE) = 1/P_GLOBAL_SCALE stays. + # Keep P scales static-like. PV cancels each page's V scale before + # accumulating pages together. stored_scale = tl.where( amax > 0.0, - tl.minimum(amax * (P_GLOBAL_SCALE / 6.0) / page_gscale, 448.0), + tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0, ) scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) @@ -757,160 +903,1051 @@ def _fp4_mla_attention_page_stats_kernel( @triton.jit -def _fp4_mla_attention_reduce_stats_kernel( - max_ptr, - denom_ptr, +def _fp4_mla_attention_page_stats_grouped_kernel( page_max_ptr, page_sum_ptr, - num_pages, - stats_s0, - page_stats_s0, - page_stats_s1, - NUM_HEADS: tl.constexpr, - MAX_PAGES: tl.constexpr, - BLOCK_H: tl.constexpr, - occupancy: tl.constexpr = 1, -): - """Combine per-page max/sum into a global max + denom per (query, head). - - Keep the page dimension in a loop instead of a 2D [pages, heads] vector. - Large decode batches can push max pages above 256, where the bulk vector - form becomes too large for a single Triton program. - """ - gen_idx = tl.program_id(0) - head_block = tl.program_id(1) - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - - max_score = tl.full((BLOCK_H,), -float("inf"), tl.float32) - for page_idx in tl.range(0, MAX_PAGES): - page_valid = page_idx < num_pages - page_offsets = gen_idx * page_stats_s0 + page_idx * page_stats_s1 + safe_offs_h - page_max = tl.load( - page_max_ptr + page_offsets, mask=mask_h & page_valid, other=-float("inf") - ) - max_score = tl.maximum(max_score, page_max) - - safe_max = tl.where(max_score > -float("inf"), max_score, 0.0) - denom = tl.zeros((BLOCK_H,), tl.float32) - for page_idx in tl.range(0, MAX_PAGES): - page_valid = page_idx < num_pages - page_offsets = gen_idx * page_stats_s0 + page_idx * page_stats_s1 + safe_offs_h - page_max = tl.load( - page_max_ptr + page_offsets, mask=mask_h & page_valid, other=-float("inf") - ) - page_sum = tl.load(page_sum_ptr + page_offsets, mask=mask_h & page_valid, other=0.0) - weights = tl.math.exp2((page_max - safe_max) * _LOG2_E) - denom += tl.where(page_sum > 0.0, page_sum * weights, 0.0) - - tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) - tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) - - -@triton.jit -def _fp4_mla_attention_prob_scale_kernel( + p_fp4_ptr, p_sf_ptr, - max_ptr, - denom_ptr, - page_max_ptr, + q_fp4_ptr, + q_sf_ptr, + kv_cache_ptr, + sf_cache_ptr, + global_scale_ptr, + q_global_scale_ptr, + src_page_ids_ptr, paged_kv_indptr_decode_ptr, kv_lens_ptr, page_ids_len, - stats_s0, + num_pages, + q_fp4_s0, + q_fp4_s1, + kv_s0, + kv_s2, + kv_s4, + sf_s0, page_stats_s0, page_stats_s1, + p_s0, + p_s1, + p_num_rows, + q_num_rows, + sm_scale, NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, QUERY_LEN_PER_SEQ: tl.constexpr, MAX_PAGES: tl.constexpr, BLOCK_H: tl.constexpr, - ASSUME_FULL_HEADS: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + GROUP_PAGES: tl.constexpr, ASSUME_FULL_PAGES: tl.constexpr, ASSUME_VALID_PAGES: tl.constexpr, + PAGE_LOOP_STAGES: tl.constexpr, occupancy: tl.constexpr = 1, ): - """Apply per-page softmax correction by scaling p_sf in place.""" + """Grouped page-stats QK + softmax-stats + FP4 P pack. + + Functionally identical to ``_fp4_mla_attention_page_stats_kernel`` for the + perfect decode shape (NUM_HEADS == BLOCK_H, TMA + PACK_PROBS, the standard + 640/576/64 residual-Q layout) but each program owns one + ``(query, head_block, page_group)`` and walks ``GROUP_PAGES`` pages in a + pipelined loop. Q (and its scales) and the TMA descriptors are loaded once + and reused across the group, eliminating the per-page Q reload and the tiny + per-CTA prologue that made the one-page-per-CTA kernel work-bound at long + context. Per-page outputs (page_max/page_sum and packed P) are written + exactly as the one-page kernel writes them, so every downstream stage is + unchanged. + """ query_idx = tl.program_id(0) head_block = tl.program_id(1) - page_rel = tl.program_id(2) + page_group = tl.program_id(2) seq_idx = query_idx // QUERY_LEN_PER_SEQ query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ - page_start = page_rel * PAGE_SIZE + head_start = head_block * BLOCK_H + offs_h = head_start + tl.arange(0, BLOCK_H) + offs_t = tl.arange(0, BLOCK_T) + q_row_base = query_idx * NUM_HEADS + if ASSUME_FULL_PAGES: kv_len = 0 else: kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) kv_len = tl.maximum(kv_len, 0) - if page_start >= kv_len: - return - page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) - compact_page = page_table_start + page_rel - if not ASSUME_VALID_PAGES: - if (compact_page < 0) | (compact_page >= page_ids_len): - return - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - if ASSUME_FULL_HEADS: - mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) - safe_offs_h = offs_h - else: - mask_h = offs_h < NUM_HEADS - safe_offs_h = tl.where(mask_h, offs_h, 0) - page_max = tl.load( - page_max_ptr + query_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, - mask=mask_h, - other=-float("inf"), + global_scale = tl.load(global_scale_ptr) + q_gscale = tl.load(q_global_scale_ptr) + qk_scale = sm_scale / (q_gscale * global_scale) + + residual_groups = Q_RESIDUAL_D // FP4_BLOCK + non_residual_groups = K_HEAD_D // FP4_BLOCK - residual_groups + + # ---- Hoisted, page-independent index tensors and Q tiles ---- + # Main window (q_start == 0): the first FULL_BLOCK_END elements sit entirely + # in the non-residual region, so Q and K map 1:1 and load contiguously. + # Q rows are global rows (query_idx * NUM_HEADS + head); the swizzle's + # row_group term selects this query's scale/tail block, so q_row_base must + # be folded in (the main q_vals descriptor load already does this via its + # row coordinate). Keep int32 to match the one-page kernel -- int64 swizzle + # math is emulated and was measured ~2x slower; the max global row index + # (num_queries * NUM_HEADS * stride) stays well within int32. + q_rows = q_row_base + offs_h + scale_offsets = tl.arange(0, BLOCK_K // FP4_BLOCK) + q_sf_cols = scale_offsets + q_sf_offsets = _fp4_mla_swizzled_sf_offset(q_rows[:, None], q_sf_cols[None, :], Q_SF_PER_TOKEN) + k_sf_offsets_main = _fp4_mla_swizzled_sf_offset( + offs_t[:, None], q_sf_cols[None, :], K_SF_PER_TOKEN + ) + q_scales = tl.load(q_sf_ptr + q_sf_offsets) + + # Tail window (q_start == FULL_BLOCK_END): the residual-Q groups, each of + # which maps onto a duplicated K residual group. + tail_packed_offsets = tl.arange(0, TAIL_BLOCK_K // 2) + tail_scale_offsets = tl.arange(0, TAIL_BLOCK_K // FP4_BLOCK) + qt_elem = FULL_BLOCK_END + tail_packed_offsets * 2 + qt_group = qt_elem // FP4_BLOCK + kt_group = tl.where( + qt_group < non_residual_groups, + qt_group, + non_residual_groups + (qt_group - non_residual_groups) // 2, + ) + byte_t = (qt_elem % FP4_BLOCK) // 2 + packed_qt_cols = FULL_BLOCK_END // 2 + tail_packed_offsets + packed_kt_cols = kt_group * (FP4_BLOCK // 2) + byte_t + qt_sf_cols = FULL_BLOCK_END // FP4_BLOCK + tail_scale_offsets + kt_sf_cols = tl.where( + qt_sf_cols < non_residual_groups, + qt_sf_cols, + non_residual_groups + (qt_sf_cols - non_residual_groups) // 2, + ) + qt_sf_offsets = _fp4_mla_swizzled_sf_offset( + q_rows[:, None], qt_sf_cols[None, :], Q_SF_PER_TOKEN + ) + kt_sf_offsets = _fp4_mla_swizzled_sf_offset( + offs_t[:, None], kt_sf_cols[None, :], K_SF_PER_TOKEN + ) + q_tail_scales = tl.load(q_sf_ptr + qt_sf_offsets) + q_tail_vals = tl.load( + q_fp4_ptr + q_rows[:, None] * q_fp4_s0 + packed_qt_cols[None, :] * q_fp4_s1 ) - max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - factor = tl.where(denom > 0.0, tl.math.exp2((page_max - max_score) * _LOG2_E) / denom, 0.0) - p_page = query_idx * MAX_PAGES + page_rel - p_rows = p_page * NUM_HEADS + offs_h - safe_p_rows = ( - p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) - ).to(tl.int64) + tl.assume(q_fp4_s0 % 8 == 0) + tl.assume(q_fp4_s1 == 1) + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + q_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, BLOCK_K // 2], + ) + k_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, BLOCK_K // 2], + ) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + q_vals = q_desc.load([(q_row_base + head_start).to(tl.int32), 0]) + scale_cols = tl.arange(0, SF_PER_PAGE) - if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: - sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( - p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE - ) - else: - sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE - ) - scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) - tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + + page_lo = page_group * GROUP_PAGES + page_hi = page_lo + GROUP_PAGES + for page_rel in tl.range(page_lo, page_hi, num_stages=PAGE_LOOP_STAGES): + if page_rel < MAX_PAGES: + page_start = page_rel * PAGE_SIZE + page_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + page_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) + if ASSUME_FULL_PAGES or page_start < kv_len: + compact_page = page_table_start + page_rel + if ASSUME_VALID_PAGES: + physical_page = tl.load(src_page_ids_ptr + compact_page).to(tl.int64) + safe_physical_page = physical_page + else: + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + ) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + + k_vals = k_desc.load([safe_physical_page.to(tl.int32), 0, 0]) + k_vals = tl.reshape(k_vals, (BLOCK_T, BLOCK_K // 2)) + if not ASSUME_VALID_PAGES: + k_vals = tl.where(valid_physical_page, k_vals, 0) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets_main) + scores = tl.dot_scaled( + q_vals, + q_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + fast_math=True, + rhs_k_pack=True, + ) + + kt_ptrs = ( + kv_cache_ptr + + safe_physical_page * kv_s0 + + offs_t[:, None].to(tl.int64) * kv_s2 + + packed_kt_cols[None, :] * kv_s4 + ) + if ASSUME_VALID_PAGES: + kt_vals = tl.load(kt_ptrs) + else: + kt_vals = tl.load(kt_ptrs, mask=valid_physical_page, other=0) + kt_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + kt_sf_offsets) + scores = tl.dot_scaled( + q_tail_vals, + q_tail_scales, + "e2m1", + kt_vals.T, + kt_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + if ASSUME_FULL_PAGES: + scores = scores * qk_scale + page_max = tl.max(scores, axis=1) + exp_scores = tl.math.exp2((scores - page_max[:, None]) * _LOG2_E) + page_sum = tl.sum(exp_scores, axis=1) + else: + valid_t = page_start + offs_t < kv_len + scores = tl.where(valid_t[None, :], scores * qk_scale, -float("inf")) + page_max = tl.max(scores, axis=1) + exp_scores = tl.math.exp2((scores - page_max[:, None]) * _LOG2_E) + exp_scores = tl.where(valid_t[None, :], exp_scores, 0.0) + page_sum = tl.sum(exp_scores, axis=1) + + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + stored_scale = tl.where( + amax > 0.0, + tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), + 1.0, + ) + scaled_probs = grouped_probs * tl.reshape( + inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1) + ) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + p_page = query_idx * MAX_PAGES + page_rel + if ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + p_rows = (p_page * NUM_HEADS + offs_h).to(tl.int64) + sf_offsets = _fp4_mla_swizzled_sf_offset( + p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + if ASSUME_VALID_PAGES: + tl.store(p_sf_ptr + sf_offsets, stored_scale) + p_desc.store( + [(p_page * NUM_HEADS + head_start).to(tl.int32), 0], + tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), + ) + else: + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_compact_page) + p_rows = (p_page * NUM_HEADS + offs_h).to(tl.int64) + tl.store( + p_fp4_ptr + p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=valid_compact_page, + ) + + out_offsets = query_idx * page_stats_s0 + page_rel * page_stats_s1 + offs_h + tl.store(page_max_ptr + out_offsets, page_max) + tl.store(page_sum_ptr + out_offsets, page_sum) @triton.jit -def _fp4_mla_attention_pv_kernel( - out_ptr, +def _fp4_mla_attention_page_stats_mtp_kernel( + page_max_ptr, + page_sum_ptr, p_fp4_ptr, p_sf_ptr, + q_fp4_ptr, + q_sf_ptr, kv_cache_ptr, - v_sf_ptr, + sf_cache_ptr, global_scale_ptr, + q_global_scale_ptr, src_page_ids_ptr, paged_kv_indptr_decode_ptr, kv_lens_ptr, page_ids_len, num_pages, - out_s0, - out_s1, - out_s2, - out_num_rows, - p_s0, - p_s1, - p_num_rows, + q_fp4_s0, + q_fp4_s1, kv_s0, kv_s2, kv_s4, - vsf_s0, + sf_s0, + page_stats_s0, + page_stats_s1, + p_s0, + p_s1, + p_num_rows, + q_num_rows, + sm_scale, + NUM_HEADS: tl.constexpr, + Q_HEAD_D: tl.constexpr, + K_HEAD_D: tl.constexpr, + Q_RESIDUAL_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + Q_SF_PER_TOKEN: tl.constexpr, + K_SF_PER_TOKEN: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, + FULL_BLOCK_END: tl.constexpr, + TAIL_BLOCK_K: tl.constexpr, + occupancy: tl.constexpr = 1, +): + """MTP-fused page-stats: one CTA owns (seq, head_block, page) and processes + all QUERY_LEN_PER_SEQ linear-MTP query rows of the sequence, loading the + page's K (and K scales) once and reusing it across the q_len QK matmuls. + + The per-query-row kernel reloads K once per query row (q_len times per + page); at decode the QK is load-latency bound (one K load feeds one MMA), + so amortizing the K load over q_len rows lifts the load:MMA ratio. Per-row + outputs are written identically to the one-page kernel's masked path + (ASSUME_FULL_PAGES/VALID_PAGES are always False for q_len>1), so all + downstream stages are unchanged. Restricted to the perfect decode shape. + """ + seq_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_rel = tl.program_id(2) + head_start = head_block * BLOCK_H + offs_h = head_start + tl.arange(0, BLOCK_H) + offs_t = tl.arange(0, BLOCK_T) + page_start = page_rel * PAGE_SIZE + + kv_len_base = tl.load(kv_lens_ptr + seq_idx) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + q_gscale = tl.load(q_global_scale_ptr) + qk_scale = sm_scale / (q_gscale * global_scale) + + residual_groups = Q_RESIDUAL_D // FP4_BLOCK + non_residual_groups = K_HEAD_D // FP4_BLOCK - residual_groups + + # ---- r-independent index tensors (main + residual-tail column maps) ---- + scale_offsets = tl.arange(0, BLOCK_K // FP4_BLOCK) + q_sf_cols = scale_offsets + k_sf_offsets_main = _fp4_mla_swizzled_sf_offset( + offs_t[:, None], q_sf_cols[None, :], K_SF_PER_TOKEN + ) + tail_packed_offsets = tl.arange(0, TAIL_BLOCK_K // 2) + tail_scale_offsets = tl.arange(0, TAIL_BLOCK_K // FP4_BLOCK) + qt_elem = FULL_BLOCK_END + tail_packed_offsets * 2 + qt_group = qt_elem // FP4_BLOCK + kt_group = tl.where( + qt_group < non_residual_groups, + qt_group, + non_residual_groups + (qt_group - non_residual_groups) // 2, + ) + byte_t = (qt_elem % FP4_BLOCK) // 2 + packed_qt_cols = FULL_BLOCK_END // 2 + tail_packed_offsets + packed_kt_cols = kt_group * (FP4_BLOCK // 2) + byte_t + qt_sf_cols = FULL_BLOCK_END // FP4_BLOCK + tail_scale_offsets + kt_sf_cols = tl.where( + qt_sf_cols < non_residual_groups, + qt_sf_cols, + non_residual_groups + (qt_sf_cols - non_residual_groups) // 2, + ) + kt_sf_offsets = _fp4_mla_swizzled_sf_offset( + offs_t[:, None], kt_sf_cols[None, :], K_SF_PER_TOKEN + ) + scale_cols = tl.arange(0, SF_PER_PAGE) + byte_offsets = tl.arange(0, FP4_BLOCK // 2) + byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + + tl.assume(q_fp4_s0 % 8 == 0) + tl.assume(q_fp4_s1 == 1) + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + q_desc = tl.make_tensor_descriptor( + q_fp4_ptr, + shape=[q_num_rows, Q_HEAD_D // 2], + strides=[q_fp4_s0, q_fp4_s1], + block_shape=[BLOCK_H, BLOCK_K // 2], + ) + k_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, BLOCK_T, K_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, BLOCK_T, BLOCK_K // 2], + ) + # p_desc = tl.make_tensor_descriptor( + # p_fp4_ptr, + # shape=[p_num_rows, PAGE_SIZE // 2], + # strides=[p_s0, p_s1], + # block_shape=[BLOCK_H, PAGE_SIZE // 2], + # ) + + # ---- Load this page's K once (shared across all query rows). ---- + compact_page = page_table_start + page_rel + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + + k_vals = k_desc.load([safe_physical_page.to(tl.int32), 0, 0]) + k_vals = tl.reshape(k_vals, (BLOCK_T, BLOCK_K // 2)) + k_vals = tl.where(valid_physical_page, k_vals, 0) + k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets_main) + kt_vals = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + offs_t[:, None].to(tl.int64) * kv_s2 + + packed_kt_cols[None, :] * kv_s4, + mask=valid_physical_page, + other=0, + ) + kt_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + kt_sf_offsets) + + for r in tl.static_range(QUERY_LEN_PER_SEQ): + query_idx_r = seq_idx * QUERY_LEN_PER_SEQ + r + kv_len_r = tl.maximum(kv_len_base - (QUERY_LEN_PER_SEQ - 1 - r), 0) + q_row_base_r = query_idx_r * NUM_HEADS + q_rows_r = q_row_base_r + offs_h + page_max = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) + page_sum = tl.zeros((BLOCK_H,), dtype=tl.float32) + if page_start < kv_len_r: + q_vals = q_desc.load([(q_row_base_r + head_start).to(tl.int32), 0]) + q_sf_offsets = _fp4_mla_swizzled_sf_offset( + q_rows_r[:, None], q_sf_cols[None, :], Q_SF_PER_TOKEN + ) + q_scales = tl.load(q_sf_ptr + q_sf_offsets) + scores = tl.dot_scaled( + q_vals, + q_scales, + "e2m1", + k_vals.T, + k_scales, + "e2m1", + fast_math=True, + rhs_k_pack=True, + ) + qt_sf_offsets = _fp4_mla_swizzled_sf_offset( + q_rows_r[:, None], qt_sf_cols[None, :], Q_SF_PER_TOKEN + ) + q_tail_scales = tl.load(q_sf_ptr + qt_sf_offsets) + q_tail_vals = tl.load( + q_fp4_ptr + q_rows_r[:, None] * q_fp4_s0 + packed_qt_cols[None, :] * q_fp4_s1 + ) + scores = tl.dot_scaled( + q_tail_vals, + q_tail_scales, + "e2m1", + kt_vals.T, + kt_scales, + "e2m1", + acc=scores, + fast_math=True, + rhs_k_pack=True, + ) + + valid_t = page_start + offs_t < kv_len_r + scores = tl.where(valid_t[None, :], scores * qk_scale, -float("inf")) + page_max = tl.max(scores, axis=1) + exp_scores = tl.math.exp2((scores - page_max[:, None]) * _LOG2_E) + exp_scores = tl.where(valid_t[None, :], exp_scores, 0.0) + page_sum = tl.sum(exp_scores, axis=1) + + grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) + amax = tl.max(grouped_probs, axis=2) + inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) + stored_scale = tl.where( + amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0 + ) + scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) + pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) + even_probs, odd_probs = tl.split(pairs) + packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) + + p_page = query_idx_r * MAX_PAGES + page_rel + p_rows = (p_page * NUM_HEADS + offs_h).to(tl.int64) + sf_offsets = _fp4_mla_swizzled_sf_offset( + p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_compact_page) + tl.store( + p_fp4_ptr + p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + packed, + mask=valid_compact_page, + ) + + out_offsets = query_idx_r * page_stats_s0 + page_rel * page_stats_s1 + offs_h + tl.store(page_max_ptr + out_offsets, page_max) + tl.store(page_sum_ptr + out_offsets, page_sum) + + +@triton.jit +def _fp4_mla_attention_reduce_stats_kernel( + max_ptr, + denom_ptr, + page_max_ptr, + page_sum_ptr, + num_pages, + stats_s0, + page_stats_s0, + page_stats_s1, + NUM_HEADS: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + occupancy: tl.constexpr = 1, +): + """Combine per-page max/sum into a global max + denom per (query, head). + + Keep the page dimension in a loop instead of a 2D [pages, heads] vector. + Large decode batches can push max pages above 256, where the bulk vector + form becomes too large for a single Triton program. + """ + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + + max_score = tl.full((BLOCK_H,), -float("inf"), tl.float32) + for page_idx in tl.range(0, MAX_PAGES): + page_valid = page_idx < num_pages + page_offsets = gen_idx * page_stats_s0 + page_idx * page_stats_s1 + safe_offs_h + page_max = tl.load( + page_max_ptr + page_offsets, mask=mask_h & page_valid, other=-float("inf") + ) + max_score = tl.maximum(max_score, page_max) + + safe_max = tl.where(max_score > -float("inf"), max_score, 0.0) + denom = tl.zeros((BLOCK_H,), tl.float32) + for page_idx in tl.range(0, MAX_PAGES): + page_valid = page_idx < num_pages + page_offsets = gen_idx * page_stats_s0 + page_idx * page_stats_s1 + safe_offs_h + page_max = tl.load( + page_max_ptr + page_offsets, mask=mask_h & page_valid, other=-float("inf") + ) + page_sum = tl.load(page_sum_ptr + page_offsets, mask=mask_h & page_valid, other=0.0) + weights = tl.math.exp2((page_max - safe_max) * _LOG2_E) + denom += tl.where(page_sum > 0.0, page_sum * weights, 0.0) + + tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) + tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) + + +@triton.jit +def _fp4_mla_attention_group_reduce_stats_kernel( + group_max_ptr, + group_sum_ptr, + page_max_ptr, + page_sum_ptr, + num_pages, + group_stats_s0, + group_stats_s1, + page_stats_s0, + page_stats_s1, + NUM_HEADS: tl.constexpr, + GROUP_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + PIPELINE_STAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + """First level of a two-level softmax-stats reduction. + + Each program owns one ``(query, head_block, page_group)`` and combines the + ``GROUP_PAGES`` per-page ``(max, sum)`` pairs in its group into a single + online-softmax partial ``(group_max, group_denom)``. ``group_denom`` is the + page sums rescaled to ``group_max``. A small follow-up combine + (``_fp4_mla_attention_reduce_stats_kernel`` over the compact group buffer) + folds the groups into the global ``(max, denom)``. + + This parallelizes the page reduction across the grid's third axis so the + decode-stats reduction is no longer a handful of CTAs each serially walking + every page (the bottleneck the single-level reduce hit at small batch). + """ + gen_idx = tl.program_id(0) + head_block = tl.program_id(1) + group_idx = tl.program_id(2) + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + + group_max = tl.full((BLOCK_H,), -float("inf"), tl.float32) + group_sum = tl.zeros((BLOCK_H,), tl.float32) + page_lo = group_idx * GROUP_PAGES + page_hi = page_lo + GROUP_PAGES + for page_idx in tl.range(page_lo, page_hi, num_stages=PIPELINE_STAGES): + page_valid = page_idx < num_pages + page_offsets = gen_idx * page_stats_s0 + page_idx * page_stats_s1 + safe_offs_h + page_max = tl.load( + page_max_ptr + page_offsets, mask=mask_h & page_valid, other=-float("inf") + ) + page_sum = tl.load(page_sum_ptr + page_offsets, mask=mask_h & page_valid, other=0.0) + next_group_max = tl.maximum(group_max, page_max) + # Guard the rescale deltas so an empty accumulator (group_sum == 0, + # group_max == -inf) or an empty page (page_sum == 0) never feeds + # (-inf) - (-inf) = NaN into exp2. + old_delta = tl.where(group_sum > 0.0, group_max - next_group_max, 0.0) + new_delta = tl.where(page_sum > 0.0, page_max - next_group_max, 0.0) + group_sum = group_sum * tl.math.exp2(old_delta * _LOG2_E) + page_sum * tl.math.exp2( + new_delta * _LOG2_E + ) + group_max = next_group_max + + out_offsets = gen_idx * group_stats_s0 + group_idx * group_stats_s1 + safe_offs_h + tl.store(group_max_ptr + out_offsets, group_max, mask=mask_h) + tl.store(group_sum_ptr + out_offsets, group_sum, mask=mask_h) + + +@triton.jit +def _fp4_mla_attention_prob_scale_kernel( + p_sf_ptr, + max_ptr, + denom_ptr, + page_max_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + stats_s0, + page_stats_s0, + page_stats_s1, + NUM_HEADS: tl.constexpr, + PAGE_SIZE: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + """Apply per-page softmax correction by scaling p_sf in place.""" + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + page_rel = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + if page_start >= kv_len: + return + + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + compact_page = page_table_start + page_rel + if not ASSUME_VALID_PAGES: + if (compact_page < 0) | (compact_page >= page_ids_len): + return + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + page_max = tl.load( + page_max_ptr + query_idx * page_stats_s0 + page_rel * page_stats_s1 + safe_offs_h, + mask=mask_h, + other=-float("inf"), + ) + max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) + factor = tl.where(denom > 0.0, tl.math.exp2((page_max - max_score) * _LOG2_E) / denom, 0.0) + + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + ).to(tl.int64) + scale_cols = tl.arange(0, SF_PER_PAGE) + if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: + sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) + tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) + + +@triton.jit +def _fp4_mla_attention_pv_kernel( + out_ptr, + p_fp4_ptr, + p_sf_ptr, + kv_cache_ptr, + v_packed_ptr, + v_sf_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + out_s0, + out_s1, + out_s2, + out_num_rows, + p_s0, + p_s1, + p_num_rows, + kv_s0, + kv_s2, + kv_s4, + vsf_s0, + NUM_HEADS: tl.constexpr, + V_HEAD_D: tl.constexpr, + PAGE_SIZE: tl.constexpr, + FP4_BLOCK: tl.constexpr, + SF_PER_PAGE: tl.constexpr, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, + USE_TMA_P_LOAD: tl.constexpr, + USE_TMA_V_LOAD: tl.constexpr, + USE_PREPACKED_V: tl.constexpr, + PV_LOOP_STAGES: tl.constexpr, + ASSUME_FULL_HEADS: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_FULL_V: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + PAGE_SPLIT: tl.constexpr = 1, + PAGES_PER_SPLIT: tl.constexpr = 0, + PARTIAL_OUT: tl.constexpr = False, + partial_out_ptr=None, + partial_s0: tl.constexpr = 0, + partial_s1: tl.constexpr = 0, + partial_s2: tl.constexpr = 0, + partial_s3: tl.constexpr = 0, + occupancy: tl.constexpr = 1, +): + query_idx = tl.program_id(0) + head_block = tl.program_id(1) + # When PAGE_SPLIT > 1 we encode (dim_block, split_idx) into program_id(2). + # The outer loop over pages is partitioned across split_idx programs so the + # grid grows by PAGE_SPLIT× — this lifts the bs<=32 PV grid out of the + # 0.5-wave-per-SM regime that the ncu report flagged. + prog2 = tl.program_id(2) + if PAGE_SPLIT > 1: + dim_block = prog2 // PAGE_SPLIT + split_idx = prog2 - dim_block * PAGE_SPLIT + else: + dim_block = prog2 + split_idx = 0 + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) + if ASSUME_FULL_HEADS: + mask_h = tl.full([BLOCK_H], True, dtype=tl.int1) + safe_offs_h = offs_h + else: + mask_h = offs_h < NUM_HEADS + safe_offs_h = tl.where(mask_h, offs_h, 0) + if ASSUME_FULL_V: + mask_v = tl.full([BLOCK_V], True, dtype=tl.int1) + safe_offs_v = offs_v + else: + mask_v = offs_v < V_HEAD_D + safe_offs_v = tl.where(mask_v, offs_v, 0) + packed_t = tl.arange(0, PAGE_SIZE // 2) + scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) + even_t = packed_t * 2 + odd_t = even_t + 1 + v_packed_offsets = safe_offs_v // 2 + v_use_high_nibble = (safe_offs_v & 1) != 0 + if ASSUME_FULL_V and BLOCK_V == 128: + v_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + else: + v_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) + if USE_TMA_P_LOAD: + tl.assume(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + p_desc = tl.make_tensor_descriptor( + p_fp4_ptr, + shape=[p_num_rows, PAGE_SIZE // 2], + strides=[p_s0, p_s1], + block_shape=[BLOCK_H, PAGE_SIZE // 2], + ) + if not PARTIAL_OUT and USE_TMA_V_LOAD and ASSUME_FULL_HEADS and ASSUME_FULL_V: + tl.assume(out_s1 % 8 == 0) + tl.assume(out_s2 == 1) + out_desc = tl.make_tensor_descriptor( + out_ptr, + shape=[out_num_rows, V_HEAD_D], + strides=[out_s1, out_s2], + block_shape=[BLOCK_H, BLOCK_V], + ) + if USE_PREPACKED_V: + v_packed_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) + elif USE_TMA_V_LOAD: + tl.assume(kv_s0 % 8 == 0) + tl.assume(kv_s2 % 8 == 0) + tl.assume(kv_s4 == 1) + v_desc = tl.make_tensor_descriptor( + kv_cache_ptr, + shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], + strides=[kv_s0, kv_s2, kv_s4], + block_shape=[1, PAGE_SIZE, BLOCK_V // 2], + ) + + if ASSUME_FULL_PAGES: + kv_len = 0 + else: + kv_len = tl.load(kv_lens_ptr + seq_idx) - (QUERY_LEN_PER_SEQ - 1 - query_offset) + kv_len = tl.maximum(kv_len, 0) + page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) + acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + if PAGE_SPLIT > 1: + page_lo = split_idx * PAGES_PER_SPLIT + page_hi = tl.minimum(page_lo + PAGES_PER_SPLIT, MAX_PAGES) + else: + page_lo = 0 + page_hi = MAX_PAGES + for page_rel in tl.range(page_lo, page_hi, num_stages=PV_LOOP_STAGES): + page_start = page_rel * PAGE_SIZE + if ASSUME_FULL_PAGES or page_start < kv_len: + compact_page = page_table_start + page_rel + if ASSUME_VALID_PAGES: + safe_compact_page = compact_page + physical_page = tl.load(src_page_ids_ptr + safe_compact_page).to(tl.int64) + safe_physical_page = physical_page + else: + valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) + safe_compact_page = tl.where(valid_compact_page, compact_page, 0) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + ) + safe_physical_page = tl.where(valid_physical_page, physical_page, 0) + + p_page = query_idx * MAX_PAGES + page_rel + p_rows = p_page * NUM_HEADS + offs_h + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + ).to(tl.int64) + if USE_TMA_P_LOAD: + p_vals = p_desc.load([(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) + else: + p_vals = tl.load( + p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, + mask=mask_h[:, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None], + other=0, + ) + p_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + p_scales = tl.load(p_sf_ptr + p_sf_offsets) + + if ASSUME_FULL_PAGES: + valid_even_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + valid_odd_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) + else: + valid_even_t = page_start + even_t < kv_len + valid_odd_t = page_start + odd_t < kv_len + if USE_PREPACKED_V: + v_row = (safe_physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + if not ASSUME_VALID_PAGES: + v_vals = tl.where(valid_physical_page, v_vals, 0) + elif USE_TMA_V_LOAD: + v_tile = v_desc.load( + [ + safe_physical_page.to(tl.int32), + 0, + (dim_block * (BLOCK_V // 2)).to(tl.int32), + ] + ) + v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) + if not ASSUME_VALID_PAGES: + v_tile = tl.where(valid_physical_page, v_tile, 0) + v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) + even_packed, odd_packed = tl.split(v_pairs) + if not ASSUME_FULL_PAGES: + even_packed = tl.where(valid_even_t[None, :], even_packed, 0) + odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) + low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) + high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) + v_vals = tl.reshape( + tl.join(low_vals, high_vals).permute(0, 2, 1), + (BLOCK_V, PAGE_SIZE // 2), + ) + else: + even_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + even_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_even_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), + other=0, + ) + odd_packed = tl.load( + kv_cache_ptr + + safe_physical_page * kv_s0 + + odd_t[None, :].to(tl.int64) * kv_s2 + + v_packed_offsets[:, None] * kv_s4, + mask=(mask_v[:, None] & valid_odd_t[None, :]) + if ASSUME_VALID_PAGES + else (valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), + other=0, + ) + even_low = even_packed & 0x0F + even_high = (even_packed >> 4) & 0x0F + odd_low = odd_packed & 0x0F + even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) + odd_nibble = tl.where(v_use_high_nibble[:, None], odd_packed >> 4, odd_low) + v_vals = even_nibble | (odd_nibble << 4) + v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) + acc = tl.dot_scaled( + p_vals, + p_scales, + "e2m1", + v_vals.T, + v_scales, + "e2m1", + acc=acc, + fast_math=True, + rhs_k_pack=True, + ) + + if PARTIAL_OUT: + # Write the unscaled partial accumulator to a float32 workspace; the + # reduce-PV kernel sums splits and applies out_scale + dtype cast. + # Layout: partial_out[query_idx, split_idx, head_offset, v_offset]. + base = ( + query_idx * partial_s0 + + split_idx * partial_s1 + + safe_offs_h[:, None] * partial_s2 + + safe_offs_v[None, :] * partial_s3 + ) + if ASSUME_FULL_HEADS and ASSUME_FULL_V: + tl.store(partial_out_ptr + base, acc) + else: + tl.store(partial_out_ptr + base, acc, mask=mask_h[:, None] & mask_v[None, :]) + elif ASSUME_FULL_HEADS and ASSUME_FULL_V: + out_vals = acc * out_scale + if USE_TMA_V_LOAD: + if out_ptr.dtype.element_ty == tl.bfloat16: + out_vals = out_vals.to(tl.bfloat16) + elif out_ptr.dtype.element_ty == tl.float16: + out_vals = out_vals.to(tl.float16) + out_desc.store( + [ + (query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ], + out_vals, + ) + else: + tl.store( + out_ptr + query_idx * out_s0 + offs_h[:, None] * out_s1 + offs_v[None, :] * out_s2, + out_vals, + ) + else: + tl.store( + out_ptr + + query_idx * out_s0 + + safe_offs_h[:, None] * out_s1 + + safe_offs_v[None, :] * out_s2, + acc * out_scale, + mask=mask_h[:, None] & mask_v[None, :], + ) + + +@triton.jit +def _fp4_mla_attention_pv_prepacked_v_kernel( + out_ptr, + p_fp4_ptr, + p_sf_ptr, + v_packed_ptr, + v_sf_ptr, + global_scale_ptr, + src_page_ids_ptr, + paged_kv_indptr_decode_ptr, + kv_lens_ptr, + page_ids_len, + num_pages, + out_s0, + out_s1, + out_s2, + out_num_rows, + p_s0, + p_s1, + p_num_rows, + vsf_s0, NUM_HEADS: tl.constexpr, V_HEAD_D: tl.constexpr, PAGE_SIZE: tl.constexpr, @@ -922,7 +1959,7 @@ def _fp4_mla_attention_pv_kernel( BLOCK_H: tl.constexpr, BLOCK_V: tl.constexpr, USE_TMA_P_LOAD: tl.constexpr, - USE_TMA_V_LOAD: tl.constexpr, + USE_TMA_OUT_STORE: tl.constexpr, PV_LOOP_STAGES: tl.constexpr, ASSUME_FULL_HEADS: tl.constexpr, ASSUME_FULL_PAGES: tl.constexpr, @@ -940,10 +1977,6 @@ def _fp4_mla_attention_pv_kernel( ): query_idx = tl.program_id(0) head_block = tl.program_id(1) - # When PAGE_SPLIT > 1 we encode (dim_block, split_idx) into program_id(2). - # The outer loop over pages is partitioned across split_idx programs so the - # grid grows by PAGE_SPLIT× — this lifts the bs<=32 PV grid out of the - # 0.5-wave-per-SM regime that the ncu report flagged. prog2 = tl.program_id(2) if PAGE_SPLIT > 1: dim_block = prog2 // PAGE_SPLIT @@ -970,10 +2003,6 @@ def _fp4_mla_attention_pv_kernel( safe_offs_v = tl.where(mask_v, offs_v, 0) packed_t = tl.arange(0, PAGE_SIZE // 2) scale_cols = tl.arange(0, PAGE_SIZE // FP4_BLOCK) - even_t = packed_t * 2 - odd_t = even_t + 1 - v_packed_offsets = safe_offs_v // 2 - v_use_high_nibble = (safe_offs_v & 1) != 0 if ASSUME_FULL_V and BLOCK_V == 128: v_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE @@ -991,7 +2020,7 @@ def _fp4_mla_attention_pv_kernel( strides=[p_s0, p_s1], block_shape=[BLOCK_H, PAGE_SIZE // 2], ) - if USE_TMA_V_LOAD and ASSUME_FULL_HEADS and ASSUME_FULL_V: + if not PARTIAL_OUT and USE_TMA_OUT_STORE and ASSUME_FULL_HEADS and ASSUME_FULL_V: tl.assume(out_s1 % 8 == 0) tl.assume(out_s2 == 1) out_desc = tl.make_tensor_descriptor( @@ -1000,17 +2029,12 @@ def _fp4_mla_attention_pv_kernel( strides=[out_s1, out_s2], block_shape=[BLOCK_H, BLOCK_V], ) - if USE_TMA_V_LOAD: - tl.assume(kv_s0 % 8 == 0) - tl.assume(kv_s2 % 8 == 0) - tl.assume(kv_s4 == 1) - v_desc = tl.make_tensor_descriptor( - kv_cache_ptr, - shape=[num_pages, PAGE_SIZE, V_HEAD_D // 2], - strides=[kv_s0, kv_s2, kv_s4], - block_shape=[1, PAGE_SIZE, BLOCK_V // 2], - ) - + v_packed_desc = tl.make_tensor_descriptor( + v_packed_ptr, + shape=[num_pages * (V_HEAD_D // BLOCK_V) * BLOCK_V, PAGE_SIZE // 2], + strides=[PAGE_SIZE // 2, 1], + block_shape=[BLOCK_V, PAGE_SIZE // 2], + ) if ASSUME_FULL_PAGES: kv_len = 0 else: @@ -1019,7 +2043,7 @@ def _fp4_mla_attention_pv_kernel( page_table_start = tl.load(paged_kv_indptr_decode_ptr + seq_idx).to(tl.int64) global_scale = tl.load(global_scale_ptr) out_scale = 1.0 / (global_scale * P_GLOBAL_SCALE) - acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) + acc = tl.zeros((BLOCK_V, BLOCK_H), dtype=tl.float32) if PAGE_SPLIT > 1: page_lo = split_idx * PAGES_PER_SPLIT page_hi = tl.minimum(page_lo + PAGES_PER_SPLIT, MAX_PAGES) @@ -1060,96 +2084,44 @@ def _fp4_mla_attention_pv_kernel( else valid_compact_page & mask_h[:, None], other=0, ) - p_sf_offsets = _fp4_mla_swizzled_sf_offset( - safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE - ) - p_scales = tl.load(p_sf_ptr + p_sf_offsets) - - if ASSUME_FULL_PAGES: - valid_even_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) - valid_odd_t = tl.full([PAGE_SIZE // 2], True, dtype=tl.int1) - else: - valid_even_t = page_start + even_t < kv_len - valid_odd_t = page_start + odd_t < kv_len - if USE_TMA_V_LOAD: - v_tile = v_desc.load( - [ - safe_physical_page.to(tl.int32), - 0, - (dim_block * (BLOCK_V // 2)).to(tl.int32), - ] - ) - v_tile = tl.reshape(v_tile, (PAGE_SIZE, BLOCK_V // 2)) - if not ASSUME_VALID_PAGES: - v_tile = tl.where(valid_physical_page, v_tile, 0) - v_pairs = tl.reshape(v_tile.T, (BLOCK_V // 2, PAGE_SIZE // 2, 2)) - even_packed, odd_packed = tl.split(v_pairs) - if not ASSUME_FULL_PAGES: - even_packed = tl.where(valid_even_t[None, :], even_packed, 0) - odd_packed = tl.where(valid_odd_t[None, :], odd_packed, 0) - low_vals = _fp4_pack_low_nibbles(even_packed, odd_packed) - high_vals = _fp4_pack_high_nibbles(even_packed, odd_packed) - v_vals = tl.reshape( - tl.join(low_vals, high_vals).permute(0, 2, 1), - (BLOCK_V, PAGE_SIZE // 2), + if ASSUME_FULL_HEADS and NUM_HEADS == 128 and BLOCK_H == 128: + p_sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( + p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) else: - even_packed = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + even_t[None, :].to(tl.int64) * kv_s2 - + v_packed_offsets[:, None] * kv_s4, - mask=(mask_v[:, None] & valid_even_t[None, :]) - if ASSUME_VALID_PAGES - else (valid_physical_page & mask_v[:, None] & valid_even_t[None, :]), - other=0, - ) - odd_packed = tl.load( - kv_cache_ptr - + safe_physical_page * kv_s0 - + odd_t[None, :].to(tl.int64) * kv_s2 - + v_packed_offsets[:, None] * kv_s4, - mask=(mask_v[:, None] & valid_odd_t[None, :]) - if ASSUME_VALID_PAGES - else (valid_physical_page & mask_v[:, None] & valid_odd_t[None, :]), - other=0, + p_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE ) - even_low = even_packed & 0x0F - even_high = (even_packed >> 4) & 0x0F - odd_low = odd_packed & 0x0F - even_nibble = tl.where(v_use_high_nibble[:, None], even_high, even_low) - odd_nibble = tl.where(v_use_high_nibble[:, None], odd_packed >> 4, odd_low) - v_vals = even_nibble | (odd_nibble << 4) + p_scales = tl.load(p_sf_ptr + p_sf_offsets) + v_row = (safe_physical_page * (V_HEAD_D // BLOCK_V) + dim_block) * BLOCK_V + v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) + if not ASSUME_VALID_PAGES: + v_vals = tl.where(valid_physical_page, v_vals, 0) v_scales = tl.load(v_sf_ptr + safe_physical_page * vsf_s0 + v_sf_offsets) acc = tl.dot_scaled( - p_vals, - p_scales, - "e2m1", - v_vals.T, + v_vals, v_scales, "e2m1", + p_vals.T, + p_scales, + "e2m1", acc=acc, fast_math=True, rhs_k_pack=True, ) + out_vals = acc.T if PARTIAL_OUT: - # Write the unscaled partial accumulator to a float32 workspace; the - # reduce-PV kernel sums splits and applies out_scale + dtype cast. - # Layout: partial_out[query_idx, split_idx, head_offset, v_offset]. base = ( query_idx * partial_s0 + split_idx * partial_s1 + safe_offs_h[:, None] * partial_s2 + safe_offs_v[None, :] * partial_s3 ) - if ASSUME_FULL_HEADS and ASSUME_FULL_V: - tl.store(partial_out_ptr + base, acc) - else: - tl.store(partial_out_ptr + base, acc, mask=mask_h[:, None] & mask_v[None, :]) + tl.store(partial_out_ptr + base, out_vals, mask=mask_h[:, None] & mask_v[None, :]) elif ASSUME_FULL_HEADS and ASSUME_FULL_V: - out_vals = acc * out_scale - if USE_TMA_V_LOAD: + out_vals = out_vals * out_scale + if USE_TMA_OUT_STORE: if out_ptr.dtype.element_ty == tl.bfloat16: out_vals = out_vals.to(tl.bfloat16) elif out_ptr.dtype.element_ty == tl.float16: @@ -1172,7 +2144,7 @@ def _fp4_mla_attention_pv_kernel( + query_idx * out_s0 + safe_offs_h[:, None] * out_s1 + safe_offs_v[None, :] * out_s2, - acc * out_scale, + out_vals * out_scale, mask=mask_h[:, None] & mask_v[None, :], ) @@ -1240,297 +2212,3 @@ def _fp4_mla_attention_pv_reduce_kernel( out_vals, mask=mask_h[:, None] & mask_v[None, :], ) - - -# --------------------------------------------------------------------------- -# Per-page dynamic-scale store path (triton only) -# -# Two passes per decode step over the *active* page (the page that holds the -# current step's new tokens): -# Pass A (_fp4_mla_page_scale_gen_kernel): compute the page amax over the -# shared K/V latent (old tokens read from the BF16 staging pool, new tokens -# from latent_cache) and write page_gscale = P_GLOBAL_SCALE / page_amax into -# the [num_layers, num_pages] fp32 page-scale pool. -# Pass B (_fp4_mla_page_requant_gen_kernel): re-quantize *every* FP4 tile of -# the active page from the same (staging + latent) source, baking page_gscale -# into the stored K and V block scales. Completed pages are never revisited, -# so their scale is frozen at the value computed on the step that filled them. -# -# Both passes assume the step's new tokens land in a single page (always true -# for 1-token decode; the Python dispatch guards the MTP boundary-cross case). -# --------------------------------------------------------------------------- - - -@triton.jit -def _fp4_mla_page_scale_gen_kernel( - page_scale_ptr, - stage_pool_ptr, - latent_cache_ptr, - seq_slots_ptr, - kv_lens_ptr, - prompt_lens_ptr, - page_ids_ptr, - paged_kv_indptr_ptr, - page_ids_len, - indptr_len, - num_seq_slots, - num_pages, - num_layers, - local_layer, - page_size, - pscale_s0, - pool_s0, - pool_s1, - lc_s0, - lc_s1, - HEAD_D: tl.constexpr, - POOL_HEAD_D: tl.constexpr, - FP4_BLOCK: tl.constexpr, - PAGE_SLOTS: tl.constexpr, - BLOCK_D: tl.constexpr, - P_GLOBAL_SCALE: tl.constexpr, -): - seq_idx = tl.program_id(0) - if (local_layer < 0) | (local_layer >= num_layers): - return - if seq_idx + 1 >= indptr_len: - return - - kv_len = tl.load(kv_lens_ptr + seq_idx) - gen_len = tl.load(prompt_lens_ptr + seq_idx) - if (kv_len <= 0) | (gen_len <= 0): - return - first_new_pos = kv_len - gen_len - active_page = (kv_len - 1) // page_size - # New tokens must all land in the active page (boundary-cross guarded by - # the Python dispatch; bail defensively here too). - if first_new_pos // page_size != active_page: - return - page_pos_start = active_page * page_size - fill = kv_len - page_pos_start - - page_start = tl.load(paged_kv_indptr_ptr + seq_idx).to(tl.int64) - page_end = tl.load(paged_kv_indptr_ptr + seq_idx + 1).to(tl.int64) - physical_page_offset = page_start + active_page - if ( - (physical_page_offset < page_start) - | (physical_page_offset >= page_end) - | (physical_page_offset < 0) - | (physical_page_offset >= page_ids_len) - ): - return - physical_page = tl.load(page_ids_ptr + physical_page_offset).to(tl.int64) - if (physical_page < 0) | (physical_page >= num_pages): - return - seq_slot = tl.load(seq_slots_ptr + seq_idx).to(tl.int64) - if (seq_slot < 0) | (seq_slot >= num_seq_slots): - return - - offs_d = tl.arange(0, BLOCK_D) - mask_d = offs_d < HEAD_D - safe_d = tl.where(mask_d, offs_d, 0) - amax = 0.0 - for tile_idx in tl.range(0, PAGE_SLOTS // FP4_BLOCK): - token_offsets = tile_idx * FP4_BLOCK + tl.arange(0, FP4_BLOCK) - abs_pos = page_pos_start + token_offsets - valid = token_offsets < fill - from_latent = abs_pos >= first_new_pos - slot = abs_pos % PAGE_SLOTS - stage_vals = tl.load( - stage_pool_ptr - + seq_slot * pool_s0 - + local_layer * pool_s1 - + slot[:, None] * POOL_HEAD_D - + safe_d[None, :], - mask=valid[:, None] & (~from_latent)[:, None] & mask_d[None, :], - other=0.0, - ).to(tl.float32) - latent_tok = seq_idx * gen_len + (abs_pos - first_new_pos) - safe_latent = tl.where(valid & from_latent, latent_tok, 0).to(tl.int64) - latent_vals = tl.load( - latent_cache_ptr + safe_latent[:, None] * lc_s0 + safe_d[None, :] * lc_s1, - mask=valid[:, None] & from_latent[:, None] & mask_d[None, :], - other=0.0, - ).to(tl.float32) - vals = stage_vals + latent_vals - amax = tl.maximum(amax, tl.max(tl.abs(vals))) - - gscale = tl.where(amax > 0.0, P_GLOBAL_SCALE / amax, 1.0) - tl.store(page_scale_ptr + local_layer * pscale_s0 + physical_page, gscale) - - -@triton.jit -def _fp4_mla_page_requant_gen_kernel( - kv_cache_ptr, - sf_cache_ptr, - v_sf_ptr, - stage_pool_ptr, - latent_cache_ptr, - page_scale_ptr, - seq_slots_ptr, - kv_lens_ptr, - prompt_lens_ptr, - page_ids_ptr, - paged_kv_indptr_ptr, - page_ids_len, - indptr_len, - num_seq_slots, - num_pages, - num_layers, - local_layer, - page_size, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - pool_s0, - pool_s1, - lc_s0, - lc_s1, - vsf_s0, - vsf_s1, - pscale_s0, - HEAD_D: tl.constexpr, - POOL_HEAD_D: tl.constexpr, - V_HEAD_D: tl.constexpr, - PAGE_SLOTS: tl.constexpr, - FP4_BLOCK: tl.constexpr, - SF_PER_TOKEN: tl.constexpr, - SF_PER_PAGE: tl.constexpr, -): - seq_idx = tl.program_id(0) - tile_idx = tl.program_id(1) - dim_block = tl.program_id(2) - if (local_layer < 0) | (local_layer >= num_layers): - return - if seq_idx + 1 >= indptr_len: - return - - kv_len = tl.load(kv_lens_ptr + seq_idx) - gen_len = tl.load(prompt_lens_ptr + seq_idx) - if gen_len <= 0: - return - first_new_pos = kv_len - gen_len - active_page = (kv_len - 1) // page_size - if first_new_pos // page_size != active_page: - return - # Re-quantize every tile of the active page (page_gscale changed), not just - # the tiles that the new tokens touch. - block_base_pos = active_page * page_size + tile_idx * FP4_BLOCK - if block_base_pos >= kv_len: - return - - page_pos = block_base_pos - active_page * page_size - page_start = tl.load(paged_kv_indptr_ptr + seq_idx).to(tl.int64) - page_end = tl.load(paged_kv_indptr_ptr + seq_idx + 1).to(tl.int64) - physical_page_offset = page_start + active_page - if ( - (page_pos < 0) - | (page_pos >= page_size) - | (physical_page_offset < page_start) - | (physical_page_offset >= page_end) - | (physical_page_offset < 0) - | (physical_page_offset >= page_ids_len) - ): - return - physical_page = tl.load(page_ids_ptr + physical_page_offset).to(tl.int64) - if (physical_page < 0) | (physical_page >= num_pages): - return - seq_slot = tl.load(seq_slots_ptr + seq_idx).to(tl.int64) - if (seq_slot < 0) | (seq_slot >= num_seq_slots): - return - - page_gscale = tl.load(page_scale_ptr + local_layer * pscale_s0 + physical_page) - - byte_offsets = tl.arange(0, FP4_BLOCK // 2) - token_offsets = tl.arange(0, FP4_BLOCK) - even_d = dim_block * FP4_BLOCK + byte_offsets * 2 - odd_d = even_d + 1 - all_d = dim_block * FP4_BLOCK + tl.arange(0, FP4_BLOCK) - mask_even_d = even_d < HEAD_D - mask_odd_d = odd_d < HEAD_D - mask_all_d = all_d < HEAD_D - safe_even_d = tl.where(mask_even_d, even_d, 0) - safe_odd_d = tl.where(mask_odd_d, odd_d, 0) - safe_all_d = tl.where(mask_all_d, all_d, 0) - - abs_positions = block_base_pos + token_offsets - valid_tokens = abs_positions < kv_len - from_latent = abs_positions >= first_new_pos - slot = abs_positions % PAGE_SLOTS - new_token_offsets = abs_positions - first_new_pos - latent_tokens = seq_idx * gen_len + new_token_offsets - safe_latent_tokens = tl.where(valid_tokens & from_latent, latent_tokens, 0).to(tl.int64) - - stage_even = tl.load( - stage_pool_ptr - + seq_slot * pool_s0 - + local_layer * pool_s1 - + slot[:, None] * POOL_HEAD_D - + safe_even_d[None, :], - mask=valid_tokens[:, None] & (~from_latent)[:, None] & mask_even_d[None, :], - other=0.0, - ).to(tl.float32) - stage_odd = tl.load( - stage_pool_ptr - + seq_slot * pool_s0 - + local_layer * pool_s1 - + slot[:, None] * POOL_HEAD_D - + safe_odd_d[None, :], - mask=valid_tokens[:, None] & (~from_latent)[:, None] & mask_odd_d[None, :], - other=0.0, - ).to(tl.float32) - latent_even = tl.load( - latent_cache_ptr + safe_latent_tokens[:, None] * lc_s0 + safe_even_d[None, :] * lc_s1, - mask=valid_tokens[:, None] & from_latent[:, None] & mask_even_d[None, :], - other=0.0, - ).to(tl.float32) - latent_odd = tl.load( - latent_cache_ptr + safe_latent_tokens[:, None] * lc_s0 + safe_odd_d[None, :] * lc_s1, - mask=valid_tokens[:, None] & from_latent[:, None] & mask_odd_d[None, :], - other=0.0, - ).to(tl.float32) - even_values = stage_even + latent_even - odd_values = stage_odd + latent_odd - - amax_per_token = tl.maximum( - tl.max(tl.abs(even_values), axis=1), - tl.max(tl.abs(odd_values), axis=1), - ) - tile_amax = tl.max(amax_per_token, axis=0) - # K consumes scales as [token, dim-block], V as [dim, token-block]. Only the - # compressed-KV prefix has both views; tail K-only dims keep K's per-token - # scale. page_gscale is the shared per-page global scale. - shared_tile = dim_block * FP4_BLOCK < V_HEAD_D - tile_scale = tl.where(tile_amax > 0.0, tile_amax / 6.0, 1.0) - token_scale = tl.where(amax_per_token > 0.0, amax_per_token / 6.0, 1.0) - local_scale = tl.where(shared_tile, tile_scale, token_scale) - stored_scale = local_scale * page_gscale - v_stored_scale = tile_scale * page_gscale - - low = _fp4_e2m1_quantize(even_values / local_scale[:, None]) - high = _fp4_e2m1_quantize(odd_values / local_scale[:, None]) - packed = low | (high << 4) - - packed_cols = dim_block * (FP4_BLOCK // 2) + byte_offsets - page_positions = page_pos + token_offsets - kv_base = physical_page * kv_s0 - tl.store( - kv_cache_ptr + kv_base + page_positions[:, None] * kv_s2 + packed_cols[None, :] * kv_s4, - packed, - mask=valid_tokens[:, None] & mask_even_d[None, :], - ) - - k_sf_offsets = _fp4_mla_swizzled_sf_offset(page_positions, dim_block, SF_PER_TOKEN) - tl.store(sf_cache_ptr + physical_page * sf_s0 + k_sf_offsets, stored_scale, mask=valid_tokens) - - token_scale_col = page_pos // FP4_BLOCK - sf_offsets = _fp4_mla_swizzled_sf_offset(safe_all_d, token_scale_col, SF_PER_PAGE) - v_sf_base = tl.cast(local_layer, tl.int64) * tl.cast( - vsf_s0, tl.int64 - ) + physical_page * tl.cast(vsf_s1, tl.int64) - tl.store( - v_sf_ptr + v_sf_base + sf_offsets.to(tl.int64), - v_stored_scale, - mask=mask_all_d & (all_d < V_HEAD_D), - ) diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index 1b46fa5c09bc..07d80d80b0f4 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -48,6 +48,7 @@ from .sparse.skip_softmax import SkipSoftmaxParams + @functools.cache def generate_spec_decoding_position_offsets(max_num_requests: int, draft_len: int) -> torch.Tensor: @@ -164,10 +165,6 @@ class TrtllmAttentionMetadata(AttentionMetadata): # Standalone tensor, not part of the block-based paged KV cache. high_precision_kv_pool: Optional[torch.Tensor] = None fp4_mla_hp_snapshot_pool: Optional[torch.Tensor] = None - # Ownership tracking: maps seq_slot to request_id that last wrote it. - # Plain Python dict, updated during context phase, checked during decode. - # Debug only; runs outside CUDA graph. - hp_pool_owners: Optional[dict] = None # Pre-computed FlashMLA tile-scheduler metadata and num_splits. # Computed once per forward pass in TrtllmAttention.forward() and reused across layers. @@ -462,7 +459,6 @@ def _post_init_with_buffers(self, buffers) -> None: device='cpu', pin_memory=prefer_pinned(), ) - self.hp_pool_owners = {} num_local_layers = self.kv_cache_manager.num_local_layers head_dim = self.kv_cache_manager.head_dim kv_factor = self.kv_cache_manager.kv_factor diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 4b95dd4095f1..83689a37b7f9 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -2124,16 +2124,10 @@ def _preprocess_inputs(self, inputs: Dict[str, Any]): ) else: md.kv_lens_cuda_runtime[num_ctx_requests:num_seqs] += ( - self.previous_kv_lens_offsets_cuda[:num_gen_requests] - ) + self. + previous_kv_lens_offsets_cuda[:num_gen_requests]) md.on_update_kv_lens() md._populate_fp4_mla_batch_indices_positions() - # Opt-in: also rebuild the read-side decode paging - # (num_blocks / paged_kv_indptr_decode / paged_kv_indices / - # last_page_len) from the corrected kv_lens, since prepare() - # derived those from the all-draft-accepted over-estimate. - if hasattr(md, "repage_fp4_mla_decode_from_kv_lens"): - md.repage_fp4_mla_decode_from_kv_lens() if self.guided_decoder is not None: self.guided_decoder.token_event.record() @@ -2200,8 +2194,8 @@ def _postprocess_inputs(self, inputs: Dict[str, Any]): ) else: md.kv_lens_cuda_runtime[num_ctx_requests:num_seqs] -= ( - self.previous_kv_lens_offsets_cuda[:num_gen_requests] - ) + self. + previous_kv_lens_offsets_cuda[:num_gen_requests]) def _get_all_rank_num_tokens(self, attn_metadata: AttentionMetadata): if self.enable_attention_dp: diff --git a/tensorrt_llm/_torch/speculative/mtp.py b/tensorrt_llm/_torch/speculative/mtp.py index 6ec658280562..4a2fdfbe6679 100644 --- a/tensorrt_llm/_torch/speculative/mtp.py +++ b/tensorrt_llm/_torch/speculative/mtp.py @@ -50,13 +50,10 @@ def _repair_fp4_mla_hp_kv_after_mtp_acceptance( "_fp4_mla_mtp_hp_snapshots", None): return - from ..attention_backend.fp4_mla import ( - repair_fp4_mla_hp_kv_for_mtp_rejection, - repair_fp4_mla_page_stage_for_mtp_rejection) + from ..attention_backend.fp4_mla import \ + repair_fp4_mla_hp_kv_for_mtp_rejection repair_fp4_mla_hp_kv_for_mtp_rejection(attn_metadata, num_accepted_tokens) - repair_fp4_mla_page_stage_for_mtp_rejection(attn_metadata, - num_accepted_tokens) class MTPHiddenStatesManager(BaseResourceManager): From 233176ab0370328ede1d56c612c6243bf9bc6881 Mon Sep 17 00:00:00 2001 From: Tracin <10434017+Tracin@users.noreply.github.com> Date: Wed, 10 Jun 2026 00:55:58 -0700 Subject: [PATCH 09/11] Fix MTP + CUDA graph. Signed-off-by: Tracin <10434017+Tracin@users.noreply.github.com> --- .../_torch/attention_backend/flashinfer.py | 38 ++++- .../_torch/attention_backend/fp4_mla.py | 160 +++++++++++------- 2 files changed, 137 insertions(+), 61 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index 891880e9c717..5f71414796b7 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -133,6 +133,21 @@ class FlashInferAttentionMetadata(AttentionMetadata): # so set kv_layout as "HND" here kv_layout: Literal["NHD", "HND"] = "HND" + # Speculative-decoding placeholders used by shared one-model drafting + # code. FlashInfer does not consume TRTLLM XQA-style packed masks, so + # update_spec_dec_param intentionally leaves these tensors unset. + is_spec_decoding_enabled: bool = False + use_spec_decoding: bool = False + is_spec_dec_tree: bool = False + is_spec_dec_dynamic_tree: bool = False + spec_decoding_position_offsets: Optional[torch.Tensor] = None + spec_decoding_position_offsets_cpp: Optional[torch.Tensor] = None + spec_decoding_packed_mask: Optional[torch.Tensor] = None + spec_decoding_generation_lengths: Optional[torch.Tensor] = None + spec_decoding_bl_tree_mask_offset: Optional[torch.Tensor] = None + spec_decoding_bl_tree_mask: Optional[torch.Tensor] = None + spec_bl_tree_first_sparse_mask_offset_kv: Optional[torch.Tensor] = None + paged_kv_indptr_decode: torch.Tensor = field(init=False) paged_kv_indptr_prefill: torch.Tensor = field(init=False) _paged_kv_indices: torch.Tensor = field(init=False, repr=False) @@ -855,7 +870,16 @@ def _populate_fp4_mla_batch_indices_positions(self) -> None: seq_lens, output_size=self.num_tokens) if self.kv_lens_cuda_runtime is not None: - cached_token_lens = self.kv_lens_cuda_runtime[:num_seqs] - seq_lens + # Subtract the prompt_lens alias (the per-step append count that + # kv_lens was built from), not the live seq_lens. Under CUDA graph / + # one-engine MTP both kv_lens and prompt_lens aliases can lag at the + # decode anchor while seq_lens is the real 1 + draft_len; only the + # mutually-consistent aliases recover the true cached length. They are + # equal on the non-stale path, so this is a no-op there. + append_lens = (self.prompt_lens_cuda_runtime[:num_seqs] + if self.prompt_lens_cuda_runtime is not None else + seq_lens) + cached_token_lens = self.kv_lens_cuda_runtime[:num_seqs] - append_lens else: cached_token_lens = self.cached_token_lens[:num_seqs].to( torch.int32) @@ -882,9 +906,15 @@ def update_for_spec_dec(self) -> None: self._prompt_lens_cuda_buf[:num_seqs].copy_(prompt_lens, non_blocking=True) self.prompt_lens_cuda_runtime = self._prompt_lens_cuda_buf[:num_seqs] - self._prompt_lens_cpu_buf[:num_seqs].copy_(prompt_lens.cpu(), - non_blocking=False) - self.prompt_lens_cpu_runtime = self._prompt_lens_cpu_buf[:num_seqs] + # Refreshing the host mirror needs a D2H copy, which is illegal while a + # CUDA graph is capturing (e.g. the captured spec-dec draft loop). The + # captured kernels read only the device aliases, and host-side consumers + # short-circuit during capture, so update the mirror only when not + # capturing; it keeps its prior value during capture/replay. + if not torch.cuda.is_current_stream_capturing(): + self._prompt_lens_cpu_buf[:num_seqs].copy_(prompt_lens.cpu(), + non_blocking=False) + self.prompt_lens_cpu_runtime = self._prompt_lens_cpu_buf[:num_seqs] if self.num_tokens > 0: self._populate_fp4_mla_batch_indices_positions() diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla.py b/tensorrt_llm/_torch/attention_backend/fp4_mla.py index fd4453d16959..ef97ee2b3a78 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla.py @@ -315,6 +315,34 @@ def _scatter_fp4_mla_kv_cache_2d_context( ) +def _fp4_mla_uniform_generation_lengths( + metadata: Any, num_gen_tokens: int, num_gen: int +) -> tuple[torch.Tensor, torch.Tensor]: + """Return per-sequence ``(kv_len, gen_len)`` tensors for the generation segment. + + The no-dequant FP4 MLA generation kernels need the true number of tokens each + sequence appends this step (``1 + draft_len`` for linear MTP). Under CUDA-graph + / one-engine MTP the ``prompt_lens``/``kv_lens`` runtime aliases can lag at the + decode anchor (``seq_lens == 1``) while the real per-step query length is + ``1 + draft_len`` (the extra tokens are carried in ``num_tokens``). The two + aliases are populated together from the same ``seq_lens`` + (``kv_lens == cached + seq_lens``), so ``cached = kv_len - prompt_len`` is + representation independent and the corrected total is ``cached + per_seq``. + For uniform linear MTP ``per_seq == num_gen_tokens // num_gen``. When the + aliases already match (e.g. the chunked-context path) the returned tensors + equal the metadata slices, i.e. this is a no-op. + """ + num_contexts = metadata.num_contexts + num_seqs = metadata.num_seqs + kv_lens_gen = metadata.kv_lens_cuda_runtime[num_contexts:num_seqs] + prompt_lens_gen = metadata.prompt_lens_cuda_runtime[num_contexts:num_seqs] + if num_gen <= 0 or num_gen_tokens % num_gen != 0: + return kv_lens_gen, prompt_lens_gen + per_seq = num_gen_tokens // num_gen + cached = kv_lens_gen - prompt_lens_gen + return cached + per_seq, torch.full_like(prompt_lens_gen, per_seq) + + def _scatter_fp4_mla_kv_cache_2d_generation( metadata: Any, latent_cache: torch.Tensor, @@ -336,32 +364,22 @@ def _scatter_fp4_mla_kv_cache_2d_generation( num_gen = num_seqs - num_contexts if num_gen <= 0: return - gen_token_lens = _host_int_list_during_forward( - getattr(metadata, "prompt_lens_cpu_runtime", None), num_contexts, num_seqs - ) - if gen_token_lens is not None: - expected_num_tokens = sum(gen_token_lens) - if num_tokens != expected_num_tokens: - raise RuntimeError( - "FP4 MLA 2D generation scatter token count mismatch: " - f"expected {expected_num_tokens} generation tokens from " - f"per-sequence lengths {gen_token_lens}, got {num_tokens}." - ) - if min(gen_token_lens) != max(gen_token_lens): - raise NotImplementedError( - "FP4 MLA no-dequant generation scatter currently supports " - f"uniform linear MTP lengths only, got {gen_token_lens}." - ) - elif num_tokens < num_gen: + if num_tokens < num_gen: raise RuntimeError( f"FP4 MLA 2D generation scatter needs at least {num_gen} generation " f"tokens, got {num_tokens}." ) - elif num_tokens % num_gen != 0: + if num_tokens % num_gen != 0: raise NotImplementedError( - "FP4 MLA no-dequant generation scatter requires a uniform " + "FP4 MLA no-dequant generation scatter requires a uniform linear MTP " f"generation length, got {num_tokens} tokens for {num_gen} sequences." ) + # The prompt_lens/kv_lens runtime aliases can lag at the decode anchor + # (seq_lens == 1) under CUDA graph / one-engine MTP while each generation + # sequence really appends num_tokens // num_gen tokens this step. Recover the + # true per-sequence lengths for the no-dequant kernel below (a no-op when the + # aliases already match). + kv_lens_gen, gen_lens_gen = _fp4_mla_uniform_generation_lengths(metadata, num_tokens, num_gen) pool = getattr(metadata, "high_precision_kv_pool", None) if pool is None: @@ -373,7 +391,7 @@ def _scatter_fp4_mla_kv_cache_2d_generation( f"{hp_head_dim}." ) - max_gen_len = max(gen_token_lens) if gen_token_lens is not None else num_tokens // num_gen + max_gen_len = num_tokens // num_gen max_gen_tiles = _ceil_div(max_gen_len + HP_BLOCK_SIZE - 1, HP_BLOCK_SIZE) page_ids = metadata.paged_kv_indices[metadata.num_context_blocks :] _fp4_mla_v_scale_store_generation_tiles_kernel[ @@ -390,8 +408,8 @@ def _scatter_fp4_mla_kv_cache_2d_generation( latent_cache, global_scale, metadata.seq_slots[num_contexts:num_seqs], - metadata.kv_lens_cuda_runtime[num_contexts:num_seqs], - metadata.prompt_lens_cuda_runtime[num_contexts:num_seqs], + kv_lens_gen, + gen_lens_gen, page_ids, metadata.paged_kv_indptr_decode, page_ids.shape[0], @@ -1456,10 +1474,22 @@ def _get_linear_mtp_query_len_per_seq( num_queries: int, num_gen_seqs: int, ) -> int: - """Return the uniform generation query length required by linear MTP.""" + """Return the uniform generation query length required by linear MTP. + + Derives the length from the real query-token count (``num_queries``, taken + from the q shape) and the generation sequence count, which are reliable in + every representation. The host ``prompt_lens``/``seq_lens`` mirror can lag at + the decode anchor (== 1) under CUDA graph / one-engine MTP, so it is only + consulted to produce a precise diagnostic when the counts do not divide + evenly (a genuinely non-uniform batch, which the no-dequant path does not + support). + """ if num_gen_seqs <= 0: return 1 + if num_queries % num_gen_seqs == 0: + return num_queries // num_gen_seqs + start = metadata.num_contexts end = metadata.num_seqs query_lens = _host_int_list_during_forward( @@ -1467,33 +1497,11 @@ def _get_linear_mtp_query_len_per_seq( ) if query_lens is None: query_lens = _host_int_list_during_forward(getattr(metadata, "seq_lens", None), start, end) - - if query_lens is None: - if num_queries % num_gen_seqs != 0: - raise NotImplementedError( - "FP4 MLA linear MTP requires a uniform generation query length; " - f"got {num_queries} query tokens for {num_gen_seqs} sequences." - ) - return num_queries // num_gen_seqs - - if sum(query_lens) != num_queries and num_queries == num_gen_seqs: - return 1 - if sum(query_lens) != num_queries: - raise RuntimeError( - "FP4 MLA generation query metadata does not match q shape: " - f"query_lens={query_lens}, total={sum(query_lens)}, " - f"q_tokens={num_queries}." - ) - if not query_lens: - return 1 - if min(query_lens) <= 0: - raise RuntimeError(f"FP4 MLA generation query lengths must be positive, got {query_lens}.") - if min(query_lens) != max(query_lens): - raise NotImplementedError( - "FP4 MLA no-dequant attention currently supports linear MTP with " - f"a uniform generation length per sequence, got {query_lens}." - ) - return query_lens[0] + raise NotImplementedError( + "FP4 MLA no-dequant attention requires a uniform linear MTP generation " + f"query length; got {num_queries} query tokens for {num_gen_seqs} " + f"sequences (per-sequence lengths {query_lens})." + ) def _run_triton_attention_decode( @@ -1917,6 +1925,27 @@ def _tma_alloc(size: int, alignment: int, stream): target_l1_ctas = 3 * sm_count num_reduce_groups = _ceil_div(target_l1_ctas, max(seqhead_ctas, 1)) num_reduce_groups = max(1, min(num_reduce_groups, max_pages, 64)) + # The grouped (two-level) reduce needs an auxiliary workspace, and + # _ensure_workspace_tensor can only (re)allocate it outside CUDA graph + # capture. If a warmup forward did not already size that workspace (e.g. the + # warmup batch took the single-level path), fall back to the single-level + # reduce during capture so we never allocate mid-capture. The single-level + # reduce is numerically identical (it just launches fewer CTAs). + if num_reduce_groups > 1 and torch.cuda.is_current_stream_capturing(): + gmax = getattr(metadata, "_fp4_mla_attention_group_max_buf", None) + gsum = getattr(metadata, "_fp4_mla_attention_group_sum_buf", None) + groups_ready = ( + gmax is not None + and gsum is not None + and gmax.shape[0] >= num_queries + and gmax.shape[1] >= num_reduce_groups + and gmax.shape[2] >= num_heads + and gsum.shape[0] >= num_queries + and gsum.shape[1] >= num_reduce_groups + and gsum.shape[2] >= num_heads + ) + if not groups_ready: + num_reduce_groups = 1 if num_reduce_groups <= 1: _attn_reduce_stats_kernel[(num_queries, num_head_blocks)]( max_scores, @@ -2049,6 +2078,21 @@ def _tma_alloc(size: int, alignment: int, stream): if max_pages % p == 0 and max_pages // p >= 16 and base_grid * p <= 148 * 4: page_split = p break + # The page-split PV path needs a partial-output workspace, which + # _ensure_workspace_tensor can only (re)allocate outside CUDA graph capture. + # Fall back to the unsplit PV (numerically identical) during capture unless a + # warmup forward already sized that workspace, so capture never allocates. + if page_split > 1 and torch.cuda.is_current_stream_capturing(): + pbuf = getattr(metadata, "_fp4_mla_attention_pv_partial_buf", None) + partial_ready = ( + pbuf is not None + and pbuf.shape[0] >= num_queries + and pbuf.shape[1] >= page_split + and pbuf.shape[2] >= num_heads + and pbuf.shape[3] >= kv_lora_rank + ) + if not partial_ready: + page_split = 1 if page_split > 1: pages_per_split = max_pages // page_split partial_out = _ensure_workspace_tensor( @@ -2372,7 +2416,11 @@ def run_fp4_mla_attention_decode( src_page_ids = metadata.paged_kv_indices[ metadata.num_context_blocks : metadata.num_context_blocks + num_gen_blocks ] - kv_lens = metadata.kv_lens_cuda_runtime[metadata.num_contexts : metadata.num_seqs] + # The kv_lens runtime alias can lag at the decode anchor (seq_lens == 1) under + # CUDA graph / one-engine MTP; recover the true total per sequence so the + # per-query causal masking sees the full 1 + draft_len window (no-op when the + # alias already matches). + kv_lens, _ = _fp4_mla_uniform_generation_lengths(metadata, num_queries, num_gen_seqs) max_pages = _max_generation_pages(metadata) if max_pages == 0: return @@ -2915,12 +2963,10 @@ def update_hp_kv_for_fp4_mla( if num_gen_tokens == 0: return - gen_token_lens = _host_int_list_during_forward( - getattr(metadata, "prompt_lens_cpu_runtime", None), num_contexts, num_seqs - ) - if gen_token_lens is not None: - max_gen_len = max(gen_token_lens) - elif num_gen_tokens % num_gen == 0: + # Linear MTP is uniform, so derive the per-sequence generation length from + # the real token count rather than the prompt_lens host mirror, which can + # lag at the decode anchor (== 1) under CUDA graph / one-engine MTP. + if num_gen_tokens % num_gen == 0: max_gen_len = num_gen_tokens // num_gen else: max_gen_len = num_gen_tokens From a78bc11935101aa2522c46aba6658b27c08c72dc Mon Sep 17 00:00:00 2001 From: Yuxian Qiu <142763828+yuxianq@users.noreply.github.com> Date: Mon, 15 Jun 2026 09:41:34 +0000 Subject: [PATCH 10/11] [TRTLLM-12807][feat] Add TRTLLM FP4 MLA FMHA lib Signed-off-by: Yuxian Qiu <142763828+yuxianq@users.noreply.github.com> --- .../kernels/cutlass_kernels/CMakeLists.txt | 12 +- .../_torch/attention_backend/flashinfer.py | 543 +------- .../_torch/attention_backend/fmha/__init__.py | 2 + .../_torch/attention_backend/fmha/fallback.py | 13 + .../_torch/attention_backend/fmha/fp4_mla.py | 433 ++++++ .../_torch/attention_backend/fmha/phased.py | 6 + .../_torch/attention_backend/fmha/registry.py | 2 + .../_torch/attention_backend/fp4_mla.py | 508 +++---- .../attention_backend/fp4_mla_kernels.py | 274 ---- .../attention_backend/fp4_mla_triton.py | 7 +- .../_torch/attention_backend/trtllm.py | 286 +++- .../_torch/pyexecutor/model_engine.py | 57 +- .../_torch/pyexecutor/py_executor_creator.py | 32 +- .../_torch/pyexecutor/resource_manager.py | 56 +- .../attention/test_flashinfer_attention.py | 7 +- .../unittest/_torch/attention/test_fp4_mla.py | 1182 +++++++---------- .../executor/test_mla_tokens_per_block.py | 20 +- 17 files changed, 1507 insertions(+), 1933 deletions(-) create mode 100644 tensorrt_llm/_torch/attention_backend/fmha/fp4_mla.py diff --git a/cpp/tensorrt_llm/kernels/cutlass_kernels/CMakeLists.txt b/cpp/tensorrt_llm/kernels/cutlass_kernels/CMakeLists.txt index 7c1d0791c316..b5bdbfb4faf0 100644 --- a/cpp/tensorrt_llm/kernels/cutlass_kernels/CMakeLists.txt +++ b/cpp/tensorrt_llm/kernels/cutlass_kernels/CMakeLists.txt @@ -1,5 +1,5 @@ # -# SPDX-FileCopyrightText: Copyright (c) 1993-2024 NVIDIA CORPORATION & +# SPDX-FileCopyrightText: Copyright (c) 1993-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 @@ -244,11 +244,13 @@ if(USING_OSS_CUTLASS_MOE_GEMM) list(FILTER MOE_GEMM_SRC_CU EXCLUDE REGEX ".*moe_gemm_kernels_(bf16|fp16)_fp4.*") set(MOE_GEMM_SRC_CU_FP4 ${MOE_GEMM_SRC_CU}) - list(FILTER MOE_GEMM_SRC_CU_FP4 INCLUDE REGEX ".*fp4.*") - list(FILTER MOE_GEMM_SRC_CU EXCLUDE REGEX ".*fp4.*") + list(FILTER MOE_GEMM_SRC_CU_FP4 INCLUDE REGEX + ".*/moe_gemm/[^/]*fp4[^/]*\\.cu$") + list(FILTER MOE_GEMM_SRC_CU EXCLUDE REGEX ".*/moe_gemm/[^/]*fp4[^/]*\\.cu$") set(MOE_GEMM_SRC_CU_FP8 ${MOE_GEMM_SRC_CU}) - list(FILTER MOE_GEMM_SRC_CU_FP8 INCLUDE REGEX ".*fp8.*") - list(FILTER MOE_GEMM_SRC_CU EXCLUDE REGEX ".*fp8.*") + list(FILTER MOE_GEMM_SRC_CU_FP8 INCLUDE REGEX + ".*/moe_gemm/[^/]*fp8[^/]*\\.cu$") + list(FILTER MOE_GEMM_SRC_CU EXCLUDE REGEX ".*/moe_gemm/[^/]*fp8[^/]*\\.cu$") add_library(moe_gemm_src STATIC ${MOE_GEMM_SRC_CU} ${GROUPED_SRC_CPP}) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index 5f71414796b7..65515b6bf346 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -10,7 +10,6 @@ from flashinfer.jit.core import check_cuda_arch from typing_extensions import Self -import tensorrt_llm.bindings from tensorrt_llm._torch.pyexecutor.sampling_utils import torch_multi_arange from tensorrt_llm._utils import prefer_pinned from tensorrt_llm.functional import AttentionMaskType @@ -19,19 +18,11 @@ from ..metadata import KVCacheParams from ..utils import get_global_attrs, get_model_extra_attrs -from .fp4_mla import (FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, - FP4_MLA_KV_GLOBAL_SCALE, HP_BLOCK_SIZE, - get_fp4_mla_decode_cache, - is_flashinfer_fp4_mla_attention_enabled, - run_fp4_mla_attention_decode, scatter_fp4_mla_kv_cache, - update_hp_kv_for_fp4_mla) from .interface import (AttentionBackend, AttentionForwardArgs, AttentionInputType, AttentionMetadata, CustomAttentionMask, MLAParams, PredefinedAttentionMask, merge_attention_forward_args) -_DataType = tensorrt_llm.bindings.DataType - try: check_cuda_arch() except RuntimeError: @@ -179,62 +170,9 @@ class FlashInferAttentionMetadata(AttentionMetadata): _mla_qo_indptr_buf: Optional[torch.Tensor] = field(init=False, default=None) _mla_kv_len_arr_buf: Optional[torch.Tensor] = field(init=False, default=None) - _fp4_mla_decode_kv_indices_buf: Optional[torch.Tensor] = field(init=False, - default=None) - _fp4_mla_decode_cache_buf: Optional[torch.Tensor] = field(init=False, - default=None) - _fp4_mla_global_scale: Optional[torch.Tensor] = field(init=False, - default=None) - - # --- MLA FP4 KV cache machinery (mirrors TrtllmAttentionMetadata). --- - # Per-request stable seq_slot from SeqSlotManager, for indexing into the - # high-precision BF16 KV pool. Populated by PyTorchModelEngine via - # hasattr(metadata, 'seq_slots'). - seq_slots: Optional[torch.Tensor] = field(init=False, default=None) - seq_slots_cpu: Optional[torch.Tensor] = field(init=False, default=None) - # BF16 circular buffer, shape - # [max_num_sequences, num_local_layers, kv_factor=1, HP_BLOCK_SIZE * head_dim]. - # Holds up to HP_BLOCK_SIZE most-recent latent vectors per seq. The - # dequant fallback overlays BF16 tail tokens from it; the no-dequant path - # uses it to requantize the active 16-token FP4 KV tile. - high_precision_kv_pool: Optional[torch.Tensor] = field(init=False, - default=None) - fp4_mla_hp_snapshot_pool: Optional[torch.Tensor] = field(init=False, - default=None) - # Auxiliary FP4 MLA V-scale pool for the no-dequant PV path. The - # physical storage is flat per [local_layer, physical_page]; callers view - # it with get_fp4_mla_v_scale_pool_view(..., v_head_dim=kv_lora_rank). - fp4_mla_v_scale_pool: Optional[torch.Tensor] = field(init=False, - default=None) - _fp4_mla_attention_q_buf: Optional[torch.Tensor] = field(init=False, - default=None) - _fp4_mla_attention_p_buf: Optional[torch.Tensor] = field(init=False, - default=None) - _fp4_mla_attention_p_sf_buf: Optional[torch.Tensor] = field(init=False, - default=None) - _fp4_mla_attention_max_buf: Optional[torch.Tensor] = field(init=False, - default=None) - _fp4_mla_attention_denom_buf: Optional[torch.Tensor] = field(init=False, - default=None) # True during warmup forward passes (dummy requests, no real data). is_warmup: bool = field(init=False, default=False) - # Runtime aliases consumed by the shared HP-pool update helper. - # Same naming as TrtllmAttentionMetadata so the helper can duck-type. - kv_lens_cuda_runtime: Optional[torch.Tensor] = field(init=False, - default=None) - prompt_lens_cuda_runtime: Optional[torch.Tensor] = field(init=False, - default=None) - prompt_lens_cpu_runtime: Optional[torch.Tensor] = field(init=False, - default=None) - # Stable backing buffers for the runtime slices above (allocated only when - # NVFP4 KV + MLA is active, to avoid memory cost on the default path). - _kv_lens_cuda_buf: Optional[torch.Tensor] = field(init=False, default=None) - _prompt_lens_cuda_buf: Optional[torch.Tensor] = field(init=False, - default=None) - _prompt_lens_cpu_buf: Optional[torch.Tensor] = field(init=False, - default=None) - def needs_plan(self, plan_params: PlanParams) -> bool: if plan_params not in self._plan_params_to_wrappers: return True @@ -330,15 +268,12 @@ def plan_mla_decode( by prepare() so it runs outside of CUDA graph capture. """ if self._mla_decode_wrapper is None: - kv_indices_buf = (self._fp4_mla_decode_kv_indices_buf - if self.high_precision_kv_pool is not None else - self._paged_kv_indices) self._mla_decode_wrapper = flashinfer.mla.BatchMLAPagedAttentionWrapper( self.workspace_buffer, use_cuda_graph=self.is_cuda_graph, qo_indptr=self._mla_qo_indptr_buf, kv_indptr=self.paged_kv_indptr_decode, - kv_indices=kv_indices_buf, + kv_indices=self._paged_kv_indices, kv_len_arr=self._mla_kv_len_arr_buf, backend="auto", ) @@ -445,19 +380,10 @@ def _do_plan_mla_decode(self, plan_params: MLAPlanParams) -> None: """ num_gen = self.num_generations kv_indptr = self.paged_kv_indptr_decode[:num_gen + 1] - if self.high_precision_kv_pool is not None: - kv_indices_buf = self._fp4_mla_decode_kv_indices_buf[:self. - num_generation_blocks] - compact_indices = torch.arange(self.num_generation_blocks, - dtype=torch.int32, - device=kv_indices_buf.device) - kv_indices_buf.copy_(compact_indices, non_blocking=True) - kv_indices = kv_indices_buf.clone() - else: - kv_indices_start = self.num_context_blocks - kv_indices_end = kv_indices_start + self.num_generation_blocks - kv_indices = self._paged_kv_indices[ - kv_indices_start:kv_indices_end].clone() + kv_indices_start = self.num_context_blocks + kv_indices_end = kv_indices_start + self.num_generation_blocks + kv_indices = self._paged_kv_indices[ + kv_indices_start:kv_indices_end].clone() kv_last_page = self._paged_kv_last_page_len[self.num_contexts:self. num_contexts + num_gen] @@ -687,238 +613,6 @@ def _post_init_with_buffers(self, buffers) -> None: self._mla_context_planned = False self._mla_decode_planned = False - # --- MLA FP4 KV cache buffers (only allocated when MLA + NVFP4). --- - if (self.kv_cache_manager is not None - and self.kv_cache_manager.kv_factor == 1 - and self.kv_cache_manager.dtype == _DataType.NVFP4): - self._allocate_fp4_mla_buffers(buffers, capture_graph) - - def _allocate_fp4_mla_buffers(self, buffers, capture_graph: bool) -> None: - """Allocate the HP BF16 KV pool, seq_slots, and runtime-alias backing - buffers used by ``fp4_mla.update_hp_kv_for_fp4_mla``.""" - max_num_sequences = (self.max_num_sequences if self.max_num_sequences - is not None else self.max_num_requests) - - self.seq_slots = self.get_empty( - buffers, - (max_num_sequences, ), - cache_name="seq_slots", - dtype=torch.int32, - capture_graph=capture_graph, - ) - self.seq_slots_cpu = torch.empty( - max_num_sequences, - dtype=torch.int32, - device='cpu', - pin_memory=prefer_pinned(), - ) - - num_local_layers = self.kv_cache_manager.num_local_layers - head_dim = self.kv_cache_manager.head_dim - kv_factor = self.kv_cache_manager.kv_factor - hp_pool_shape = [ - max_num_sequences, num_local_layers, kv_factor, - HP_BLOCK_SIZE * head_dim - ] - existing_hp_pool = self.high_precision_kv_pool - if (capture_graph and existing_hp_pool is not None - and existing_hp_pool.dtype == torch.bfloat16 - and existing_hp_pool.device.type == "cuda" - and len(existing_hp_pool.shape) == len(hp_pool_shape) - and all(existing_hp_pool.shape[idx] >= dim - for idx, dim in enumerate(hp_pool_shape))): - # HP KV is persistent seq-slot state. CUDA graph metadata must - # share it instead of reserving one full pool per captured graph. - self.high_precision_kv_pool = existing_hp_pool - else: - self.high_precision_kv_pool = self.get_empty( - buffers, - hp_pool_shape, - cache_name="high_precision_kv_pool", - dtype=torch.bfloat16, - capture_graph=capture_graph, - ) - if capture_graph: - self.fp4_mla_hp_snapshot_pool = self.get_empty( - buffers, - hp_pool_shape, - cache_name="fp4_mla_hp_snapshot_pool", - dtype=torch.bfloat16, - capture_graph=capture_graph, - ) - else: - self.fp4_mla_hp_snapshot_pool = None - - max_num_pages = self.kv_cache_manager.blocks_in_primary_pool - self._fp4_mla_decode_kv_indices_buf = self.get_empty( - buffers, - (max_num_pages, ), - cache_name="_fp4_mla_decode_kv_indices_buf", - dtype=torch.int32, - capture_graph=capture_graph, - ) - self._fp4_mla_global_scale = self.get_empty( - buffers, - (1, ), - cache_name="_fp4_mla_global_scale", - dtype=torch.float32, - capture_graph=capture_graph, - ) - self._fp4_mla_global_scale.fill_(FP4_MLA_KV_GLOBAL_SCALE) - - if is_flashinfer_fp4_mla_attention_enabled(): - self.fp4_mla_v_scale_pool = self.kv_cache_manager.get_mla_v_scale_pool( - ) - if self.fp4_mla_v_scale_pool is None: - raise RuntimeError( - "FP4 MLA attention requires the C++ KV cache manager to " - "allocate the V-scale pool.") - - # Runtime-alias backing buffers: GPU for kv/prompt lens, CPU pinned for - # the helper's prompt_lens_cpu read. - self._kv_lens_cuda_buf = self.get_empty( - buffers, - (max_num_sequences, ), - cache_name="fp4_mla_kv_lens_cuda", - dtype=torch.int32, - capture_graph=capture_graph, - ) - self._prompt_lens_cuda_buf = self.get_empty( - buffers, - (max_num_sequences, ), - cache_name="fp4_mla_prompt_lens_cuda", - dtype=torch.int32, - capture_graph=capture_graph, - ) - self._prompt_lens_cpu_buf = torch.empty( - max_num_sequences, - dtype=torch.int32, - device='cpu', - pin_memory=prefer_pinned(), - ) - - logger.info( - f"FlashInfer MLA + NVFP4 KV: HP pool shape=" - f"{list(self.high_precision_kv_pool.shape)}, " - f"size={self.high_precision_kv_pool.nbytes / (1 << 20):.1f} MB") - if self.fp4_mla_v_scale_pool is not None: - logger.info( - f"FlashInfer MLA + NVFP4 KV: V scale pool shape=" - f"{list(self.fp4_mla_v_scale_pool.shape)}, " - f"size={self.fp4_mla_v_scale_pool.nbytes / (1 << 20):.1f} MB") - - def _populate_fp4_mla_runtime_aliases(self, kv_lens: torch.Tensor) -> None: - """Populate kv_lens / prompt_lens runtime-alias slices for use by the - shared HP-pool update helper. - - ``kv_lens`` is the CPU int tensor holding total KV length per sequence - after the current forward pass (see ``prepare`` where it's computed as - ``cached_token_lens + seq_lens_kv_cuda``). - """ - num_seqs = self.num_contexts + self.num_generations - kv_lens_int32 = kv_lens[:num_seqs].to(torch.int32) - self._kv_lens_cuda_buf[:num_seqs].copy_(kv_lens_int32, - non_blocking=True) - self.kv_lens_cuda_runtime = self._kv_lens_cuda_buf[:num_seqs] - - # prompt_lens = number of new tokens per sequence this forward pass. - # For self-attention this equals seq_lens_kv (= seq_lens). Use the - # int32 view stored on self.seq_lens_kv_cuda to keep the dtype stable. - prompt_lens = self.seq_lens_kv_cuda[:num_seqs].to(torch.int32) - self._prompt_lens_cuda_buf[:num_seqs].copy_(prompt_lens, - non_blocking=True) - self.prompt_lens_cuda_runtime = self._prompt_lens_cuda_buf[:num_seqs] - - # CPU pinned mirror (used by the helper to compute token offsets). - self._prompt_lens_cpu_buf[:num_seqs].copy_(prompt_lens.cpu(), - non_blocking=False) - self.prompt_lens_cpu_runtime = self._prompt_lens_cpu_buf[:num_seqs] - - def _populate_fp4_mla_batch_indices_positions(self) -> None: - """Populate append/scatter token metadata without FlashInfer helpers. - - FlashInfer's helper kernels are not reliable for the 128-token pages - required by the no-dequant FP4 MLA path. FP4 MLA only needs the generic - ragged append metadata: - batch_indices[token] = sequence index in this scheduled batch - positions[token] = absolute KV position written by that token - """ - num_seqs = self.num_contexts + self.num_generations - if num_seqs == 0 or self.num_tokens == 0: - return - - device = self._batch_indices.device - seq_lens = self.seq_lens_kv_cuda[:num_seqs].to(torch.int32) - seq_ids = torch.arange(num_seqs, dtype=torch.int32, device=device) - batch_indices = torch.repeat_interleave(seq_ids, - seq_lens, - output_size=self.num_tokens) - - # Per-sequence offsets of NEW KV tokens in the ragged batch. - # Compute this from seq_lens_kv_cuda instead of qo_indptr so the - # append positions stay correct if query lengths and KV lengths diverge. - kv_token_starts = torch.empty((num_seqs, ), - dtype=torch.int32, - device=device) - kv_token_starts[0].zero_() - if num_seqs > 1: - torch.cumsum(seq_lens[:-1], - dim=0, - dtype=torch.int32, - out=kv_token_starts[1:]) - kv_start = torch.repeat_interleave(kv_token_starts, - seq_lens, - output_size=self.num_tokens) - if self.kv_lens_cuda_runtime is not None: - # Subtract the prompt_lens alias (the per-step append count that - # kv_lens was built from), not the live seq_lens. Under CUDA graph / - # one-engine MTP both kv_lens and prompt_lens aliases can lag at the - # decode anchor while seq_lens is the real 1 + draft_len; only the - # mutually-consistent aliases recover the true cached length. They are - # equal on the non-stale path, so this is a no-op there. - append_lens = (self.prompt_lens_cuda_runtime[:num_seqs] - if self.prompt_lens_cuda_runtime is not None else - seq_lens) - cached_token_lens = self.kv_lens_cuda_runtime[:num_seqs] - append_lens - else: - cached_token_lens = self.cached_token_lens[:num_seqs].to( - torch.int32) - cached_start = torch.repeat_interleave(cached_token_lens, - seq_lens, - output_size=self.num_tokens) - token_offsets = torch.arange(self.num_tokens, - dtype=torch.int32, - device=device) - positions = token_offsets - kv_start + cached_start - - self._batch_indices[:self.num_tokens].copy_(batch_indices, - non_blocking=True) - self._positions[:self.num_tokens].copy_(positions, non_blocking=True) - - def update_for_spec_dec(self) -> None: - if self.high_precision_kv_pool is None: - return - if self.kv_lens_cuda_runtime is None: - return - - num_seqs = self.num_seqs - prompt_lens = self.seq_lens_kv_cuda[:num_seqs].to(torch.int32) - self._prompt_lens_cuda_buf[:num_seqs].copy_(prompt_lens, - non_blocking=True) - self.prompt_lens_cuda_runtime = self._prompt_lens_cuda_buf[:num_seqs] - # Refreshing the host mirror needs a D2H copy, which is illegal while a - # CUDA graph is capturing (e.g. the captured spec-dec draft loop). The - # captured kernels read only the device aliases, and host-side consumers - # short-circuit during capture, so update the mirror only when not - # capturing; it keeps its prior value during capture/replay. - if not torch.cuda.is_current_stream_capturing(): - self._prompt_lens_cpu_buf[:num_seqs].copy_(prompt_lens.cpu(), - non_blocking=False) - self.prompt_lens_cpu_runtime = self._prompt_lens_cpu_buf[:num_seqs] - - if self.num_tokens > 0: - self._populate_fp4_mla_batch_indices_positions() - def create_cuda_graph_metadata(self, max_batch_size: int, sub_cross_metadata: bool = False, @@ -1104,11 +798,6 @@ def prepare(self) -> None: # number of tokens needed in the kv cache for each sequence after the next pass kv_lens = self.cached_token_lens + self.seq_lens_kv_cuda - # Populate runtime aliases consumed by the shared HP-pool update helper - # (fp4_mla.update_hp_kv_for_fp4_mla). Only active when MLA + NVFP4. - if self.high_precision_kv_pool is not None: - self._populate_fp4_mla_runtime_aliases(kv_lens) - # start and end indices of each sequence in the ragged key and value # for self attention it's the same as qo_indptr so avoid computing twice. if self.is_cross: @@ -1216,20 +905,17 @@ def prepare(self) -> None: # For cross attention, num_tokens is 0 during decode, and we don't need to update kv cache. if self.num_tokens > 0: - if self.high_precision_kv_pool is not None: - self._populate_fp4_mla_batch_indices_positions() - else: - batch_indices, positions = flashinfer.get_batch_indices_positions( - self.kv_indptr, - flashinfer.get_seq_lens(self.paged_kv_indptr, - self.paged_kv_last_page_len, - self.page_size), - self.num_tokens, - ) - self._batch_indices[:batch_indices.size(0)].copy_( - batch_indices, non_blocking=True) - self._positions[:positions.size(0)].copy_(positions, - non_blocking=True) + batch_indices, positions = flashinfer.get_batch_indices_positions( + self.kv_indptr, + flashinfer.get_seq_lens(self.paged_kv_indptr, + self.paged_kv_last_page_len, + self.page_size), + self.num_tokens, + ) + self._batch_indices[:batch_indices.size(0)].copy_(batch_indices, + non_blocking=True) + self._positions[:positions.size(0)].copy_(positions, + non_blocking=True) # Multi-wrapper case (Gemma4 hybrid: different head_dim per layer) # shares one workspace_buffer; eager plan() would overwrite earlier @@ -1571,16 +1257,11 @@ def __init__( self.qk_nope_head_dim = mla_params.qk_nope_head_dim self.v_head_dim = mla_params.v_head_dim - # Phase 1 restriction: NVFP4 KV cache on the FlashInfer backend is - # only supported for MLA (DeepSeek-style) models. Non-MLA dense/GQA - # FlashInfer paths would need a separate scatter/dequant path that has - # not been implemented. - if getattr(self, "has_fp4_kv_cache", False) and not self.is_mla_enable: + if getattr(self, "has_fp4_kv_cache", False): raise NotImplementedError( - "NVFP4 KV cache on the FlashInfer attention backend is only " - "supported for MLA models. Set attn_backend='TRTLLM' for " - "non-MLA FP4 KV cache, or use BF16/FP8 KV cache with " - "FlashInfer.") + "NVFP4 KV cache is not supported on the FlashInfer attention " + "backend. Set attn_backend='TRTLLM' for FP4 KV cache, or use " + "BF16/FP8 KV cache with FlashInfer.") def update_quant_config(self, new_quant_config: Optional[QuantConfig]): self.quant_config = new_quant_config @@ -1749,13 +1430,7 @@ def _mla_forward_context( output: torch.Tensor, latent_cache: torch.Tensor, ) -> None: - """MLA context phase: append latent to MLA caches, run ragged prefill. - - With NVFP4 KV cache, attention still runs in BF16 over the caller's - BF16 q/k/v (no history to read during context). The only side effect - of the cache write is a quantize-and-scatter of the new ckv/kpe into - the FP4 paged pool plus a BF16 mirror of the tail into the HP pool. - """ + """MLA context phase: append latent to MLA caches, run ragged prefill.""" # 1. Append latent_cache to separate ckv/kpe paged caches. # latent_cache shape: [num_ctx_tokens, kv_lora_rank + qk_rope_head_dim] num_ctx_tokens = metadata.num_ctx_tokens @@ -1768,47 +1443,22 @@ def _mla_forward_context( append_ckv = append_ckv.to(kv_dtype) append_kpe = append_kpe.to(kv_dtype) - if self.has_fp4_kv_cache: - # MLA.forward_impl slices latent_cache to [:num_ctx_tokens] before - # dispatching to the context path. The FP4 scatter's grid is - # latent_cache.shape[0], so a regression that passes a full-batch - # latent here would over-write gen tokens. Fail loudly instead of - # silently corrupting the cache. - assert latent_cache.shape[0] == num_ctx_tokens, ( - f"FP4 MLA context scatter expected latent_cache of shape " - f"[num_ctx_tokens={num_ctx_tokens}, ...] but got " - f"{list(latent_cache.shape)}. Did MLA.forward_impl stop " - f"pre-slicing latent_cache per phase?") - scatter_fp4_mla_kv_cache( - metadata, - latent_cache, - self.layer_idx, - token_offset=0, - phase="context", - local_layer=self._local_layer_idx(metadata), - v_head_dim=self.kv_lora_rank, - ) - update_hp_kv_for_fp4_mla(metadata, - latent_cache, - self._local_layer_idx(metadata), - phase="context") - else: - ckv_cache, kpe_cache = self._get_mla_caches(metadata) + ckv_cache, kpe_cache = self._get_mla_caches(metadata) - ctx_batch_indices = metadata.batch_indices[:num_ctx_tokens] - ctx_positions = metadata.positions[:num_ctx_tokens] + ctx_batch_indices = metadata.batch_indices[:num_ctx_tokens] + ctx_positions = metadata.positions[:num_ctx_tokens] - flashinfer.page.append_paged_mla_kv_cache( - append_ckv, - append_kpe, - ctx_batch_indices, - ctx_positions, - ckv_cache, - kpe_cache, - metadata.paged_kv_indices, - metadata.paged_kv_indptr, - metadata.paged_kv_last_page_len, - ) + flashinfer.page.append_paged_mla_kv_cache( + append_ckv, + append_kpe, + ctx_batch_indices, + ctx_positions, + ckv_cache, + kpe_cache, + metadata.paged_kv_indices, + metadata.paged_kv_indptr, + metadata.paged_kv_last_page_len, + ) # 2. Run ragged prefill with expanded q, k, v num_contexts = metadata.num_contexts @@ -1853,88 +1503,31 @@ def _mla_forward_generation( if self.has_fp8_kv_cache: kv_dtype = torch.float8_e4m3fn - if self.has_fp4_kv_cache: - use_fp4_attention = is_flashinfer_fp4_mla_attention_enabled() - if latent_cache is not None: - # MLA.forward_impl slices latent_cache to [num_ctx_tokens:] - # before dispatching to the generation path. The FP4 scatter's - # grid is latent_cache.shape[0], so a regression that passes a - # full-batch latent here would OOB-read batch_indices. Fail - # loudly instead of silently corrupting the cache. - num_gen_tokens = metadata.num_tokens - metadata.num_ctx_tokens - assert latent_cache.shape[0] == num_gen_tokens, ( - f"FP4 MLA generation scatter expected latent_cache of " - f"shape [num_gen_tokens={num_gen_tokens}, ...] but got " - f"{list(latent_cache.shape)}. Did MLA.forward_impl stop " - f"pre-slicing latent_cache per phase?") - if (use_fp4_attention - and num_gen_tokens != metadata.num_generations - and os.getenv( - FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, - "triton").lower() not in ("triton", "cutile")): - use_fp4_attention = False - if use_fp4_attention: - scatter_fp4_mla_kv_cache( - metadata, - latent_cache, - self.layer_idx, - token_offset=metadata.num_ctx_tokens, - phase="generation", - local_layer=self._local_layer_idx(metadata), - v_head_dim=self.kv_lora_rank, - ) - update_hp_kv_for_fp4_mla(metadata, - latent_cache, - self._local_layer_idx(metadata), - phase="generation") - else: - scatter_fp4_mla_kv_cache( - metadata, - latent_cache, - self.layer_idx, - token_offset=metadata.num_ctx_tokens, - ) - update_hp_kv_for_fp4_mla(metadata, - latent_cache, - self._local_layer_idx(metadata), - phase="generation") - if not use_fp4_attention: - combined_cache = get_fp4_mla_decode_cache( - metadata, - self.layer_idx, - self._local_layer_idx(metadata), - head_dim=self.kv_lora_rank + self.qk_rope_head_dim, - dtype=kv_dtype, - ) - ckv_cache = combined_cache[..., :self.kv_lora_rank] - kpe_cache = combined_cache[..., self.kv_lora_rank:] - else: - use_fp4_attention = False - ckv_cache, kpe_cache = self._get_mla_caches(metadata) - - # If latent_cache is provided, append it to the paged MLA KV cache first. - # latent_cache shape: [num_tokens, kv_lora_rank + qk_rope_head_dim] - # RoPE must already be applied to the k_pe portion before calling this. - if latent_cache is not None: - append_ckv = latent_cache[:, :self.kv_lora_rank] - append_kpe = latent_cache[:, self.kv_lora_rank:] - if self.has_fp8_kv_cache: - append_ckv = append_ckv.to(kv_dtype) - append_kpe = append_kpe.to(kv_dtype) - num_ctx_tokens = metadata.num_ctx_tokens - gen_batch_indices = metadata.batch_indices[num_ctx_tokens:] - gen_positions = metadata.positions[num_ctx_tokens:] - flashinfer.page.append_paged_mla_kv_cache( - append_ckv, - append_kpe, - gen_batch_indices, - gen_positions, - ckv_cache, - kpe_cache, - metadata.paged_kv_indices, - metadata.paged_kv_indptr, - metadata.paged_kv_last_page_len, - ) + ckv_cache, kpe_cache = self._get_mla_caches(metadata) + + # If latent_cache is provided, append it to the paged MLA KV cache first. + # latent_cache shape: [num_tokens, kv_lora_rank + qk_rope_head_dim] + # RoPE must already be applied to the k_pe portion before calling this. + if latent_cache is not None: + append_ckv = latent_cache[:, :self.kv_lora_rank] + append_kpe = latent_cache[:, self.kv_lora_rank:] + if self.has_fp8_kv_cache: + append_ckv = append_ckv.to(kv_dtype) + append_kpe = append_kpe.to(kv_dtype) + num_ctx_tokens = metadata.num_ctx_tokens + gen_batch_indices = metadata.batch_indices[num_ctx_tokens:] + gen_positions = metadata.positions[num_ctx_tokens:] + flashinfer.page.append_paged_mla_kv_cache( + append_ckv, + append_kpe, + gen_batch_indices, + gen_positions, + ckv_cache, + kpe_cache, + metadata.paged_kv_indices, + metadata.paged_kv_indptr, + metadata.paged_kv_last_page_len, + ) # fused_q layout: [num_tokens, num_heads * (kv_lora_rank + qk_rope_head_dim)] # Split into q_nope (absorbed) and q_pe (rope) @@ -1950,22 +1543,6 @@ def _mla_forward_generation( else: sm_scale = 1.0 / math.sqrt(qk_head_dim) - if use_fp4_attention: - out_view = output[:num_tokens].view(-1, self.num_heads, - self.kv_lora_rank) - run_fp4_mla_attention_decode( - metadata, - self.layer_idx, - self._local_layer_idx(metadata), - q_nope, - q_pe, - out_view, - sm_scale=sm_scale, - kv_lora_rank=self.kv_lora_rank, - qk_rope_head_dim=self.qk_rope_head_dim, - ) - return - plan_params = MLAPlanParams( num_heads=self.num_heads, kv_lora_rank=self.kv_lora_rank, diff --git a/tensorrt_llm/_torch/attention_backend/fmha/__init__.py b/tensorrt_llm/_torch/attention_backend/fmha/__init__.py index 1c3981abcf91..6a0a1307f913 100644 --- a/tensorrt_llm/_torch/attention_backend/fmha/__init__.py +++ b/tensorrt_llm/_torch/attention_backend/fmha/__init__.py @@ -15,6 +15,7 @@ from .fallback import FallbackFmha from .flashinfer_trtllm_gen import FlashInferTrtllmGenFmha +from .fp4_mla import Fp4MlaFmha from .interface import Fmha from .phased import FmhaParams, PhasedFmha from .registry import DEFAULT_FMHA_LIBS, FMHA_LIBS, FmhaCls, get_enabled_fmha_lib_classes @@ -24,6 +25,7 @@ "FMHA_LIBS", "FallbackFmha", "FlashInferTrtllmGenFmha", + "Fp4MlaFmha", "Fmha", "FmhaCls", "FmhaParams", diff --git a/tensorrt_llm/_torch/attention_backend/fmha/fallback.py b/tensorrt_llm/_torch/attention_backend/fmha/fallback.py index 73c5778379d8..117bfd755307 100644 --- a/tensorrt_llm/_torch/attention_backend/fmha/fallback.py +++ b/tensorrt_llm/_torch/attention_backend/fmha/fallback.py @@ -51,6 +51,19 @@ class FallbackFmha(Fmha): """Fallback FMHA implementation using the fused TRT-LLM thop attention op.""" + def is_supported( + self, + q: torch.Tensor, + k: Optional[torch.Tensor], + v: Optional[torch.Tensor], + metadata: "TrtllmAttentionMetadata", + forward_args: AttentionForwardArgs, + ) -> bool: + attn = self.attn + if attn.is_mla_enable and attn.has_fp4_kv_cache: + return False + return True + def forward( self, q: torch.Tensor, diff --git a/tensorrt_llm/_torch/attention_backend/fmha/fp4_mla.py b/tensorrt_llm/_torch/attention_backend/fmha/fp4_mla.py new file mode 100644 index 000000000000..ffd42d69863d --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fmha/fp4_mla.py @@ -0,0 +1,433 @@ +# 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. + +from typing import TYPE_CHECKING, Optional, Tuple + +import torch +import torch.nn.functional as F + +from tensorrt_llm._torch.attention_backend.fp4_mla import ( + FP4_MLA_Q_RESIDUAL_DIM, + FP4_MLA_TOKENS_PER_BLOCK, + HP_BLOCK_SIZE, + apply_fp4_mla_rope, + run_fp4_mla_attention_decode, + scatter_fp4_mla_kv_cache, + update_hp_kv_for_fp4_mla, +) +from tensorrt_llm._torch.attention_backend.interface import ( + AttentionForwardArgs, + AttentionInputType, + PredefinedAttentionMask, +) +from tensorrt_llm._utils import get_sm_version, is_sm_100f +from tensorrt_llm.bindings import DataType +from tensorrt_llm.logger import logger +from tensorrt_llm.quantization.mode import QuantMode + +from .phased import FmhaParams, PhasedFmha + +if TYPE_CHECKING: + from tensorrt_llm._torch.attention_backend.trtllm import ( + TrtllmAttention, + TrtllmAttentionMetadata, + ) + + +class Fp4MlaFmha(PhasedFmha): + """TRTLLM FMHA library for the no-dequant NVFP4 MLA decode kernel.""" + + SUPPORTED_Q_DTYPES = {torch.bfloat16, torch.float8_e4m3fn} + SUPPORTED_CONTEXT_DTYPES = {torch.float16, torch.bfloat16} + SUPPORTED_OUTPUT_DTYPES = {torch.float16, torch.bfloat16} + + @classmethod + def is_available(cls, attn: "TrtllmAttention") -> bool: + if not attn.is_mla_enable: + logger.debug("FP4 MLA FMHA is unavailable: requires MLA.") + return False + if not QuantMode(attn.quant_mode).has_fp4_kv_cache(): + logger.debug("FP4 MLA FMHA is unavailable: requires NVFP4 KV cache quantization.") + return False + if attn.attention_chunk_size not in (None, 0): + logger.debug("FP4 MLA FMHA is unavailable: chunked attention is not supported.") + return False + if attn.predicted_tokens_per_seq > HP_BLOCK_SIZE: + logger.debug( + "FP4 MLA FMHA is unavailable: linear MTP length exceeds " + "FP4 MLA HP rollback support." + ) + return False + if attn.kv_lora_rank is None or attn.qk_rope_head_dim is None: + logger.debug("FP4 MLA FMHA is unavailable: missing MLA dimensions.") + return False + if attn.qk_nope_head_dim is None or attn.v_head_dim is None: + logger.debug("FP4 MLA FMHA is unavailable: missing MLA context dimensions.") + return False + if attn.qk_rope_head_dim != FP4_MLA_Q_RESIDUAL_DIM: + logger.debug( + f"FP4 MLA FMHA is unavailable: requires qk_rope_head_dim=" + f"{FP4_MLA_Q_RESIDUAL_DIM}, got {attn.qk_rope_head_dim}." + ) + return False + context_head_dim = attn.qk_nope_head_dim + attn.qk_rope_head_dim + fused_head_dim = attn.kv_lora_rank + attn.qk_rope_head_dim + if attn.head_dim not in (context_head_dim, fused_head_dim): + logger.debug( + "FP4 MLA FMHA is unavailable: head_dim must equal either " + f"qk_nope_head_dim + qk_rope_head_dim ({context_head_dim}) or " + f"kv_lora_rank + qk_rope_head_dim ({fused_head_dim})." + ) + return False + sm = get_sm_version() + if not is_sm_100f(sm): + logger.debug(f"FP4 MLA FMHA is unavailable: requires SM100 or SM103, got SM{sm}.") + return False + if not hasattr(torch.ops, "trtllm") or not hasattr( + torch.ops.trtllm, "fp4_quantize_with_residual" + ): + logger.debug("FP4 MLA FMHA is unavailable: missing trtllm FP4 quantization op.") + return False + return True + + def is_supported( + self, + q: torch.Tensor, + k: Optional[torch.Tensor], + v: Optional[torch.Tensor], + metadata: "TrtllmAttentionMetadata", + forward_args: AttentionForwardArgs, + ) -> bool: + supported, reason = self._is_supported_with_reason( + q, k, v, self.attn, metadata, forward_args + ) + if not supported: + logger.debug(f"FP4 MLA FMHA does not support request: {reason}") + return supported + + @classmethod + def _is_supported_with_reason( + cls, + q: torch.Tensor, + k: Optional[torch.Tensor], + v: Optional[torch.Tensor], + attn: "TrtllmAttention", + meta: "TrtllmAttentionMetadata", + fwd: AttentionForwardArgs, + ) -> Tuple[bool, str]: + if fwd.attention_input_type == AttentionInputType.context_only: + return cls._is_context_supported_with_reason(q, k, v, attn, meta, fwd) + + if fwd.attention_input_type != AttentionInputType.generation_only: + return False, "supports generation-only attention." + if meta.num_generations <= 0: + return False, "requires generation requests." + if k is not None or v is not None: + return False, "expects fused MLA query input." + if fwd.output is None: + return False, "requires output." + if fwd.latent_cache is None: + return False, "requires latent_cache." + if q.dtype not in cls.SUPPORTED_Q_DTYPES: + return False, f"unsupported query dtype {q.dtype}." + if fwd.output.dtype not in cls.SUPPORTED_OUTPUT_DTYPES: + return False, f"unsupported output dtype {fwd.output.dtype}." + if fwd.output_sf is not None: + return False, "does not support quantized attention output." + if fwd.attention_mask != PredefinedAttentionMask.CAUSAL: + return False, "requires causal mask." + if fwd.attention_mask_data is not None: + return False, "does not support custom attention masks." + if fwd.attention_sinks is not None: + return False, "does not support attention sinks." + if fwd.sage_attn_num_elts_per_blk_q > 0 or fwd.sage_attn_num_elts_per_blk_k > 0: + return False, "does not support sage attention." + if fwd.sage_attn_num_elts_per_blk_v > 0: + return False, "does not support sage attention." + sparse = fwd.sparse_prediction + if ( + (sparse.sparse_kv_indices is not None and sparse.sparse_kv_indices.numel() > 0) + or (sparse.sparse_attn_indices is not None and sparse.sparse_attn_indices.numel() > 0) + or meta.num_sparse_topk > 0 + ): + return False, "does not support sparse attention." + if meta.helix_position_offsets is not None: + return False, "does not support helix parallelism." + if meta.use_spec_decoding and meta.is_spec_dec_tree: + return False, "does not support speculative decoding trees." + if meta.kv_cache_manager is None: + return False, "requires a KV cache manager." + if meta.kv_cache_manager.dtype != DataType.NVFP4: + return False, f"requires NVFP4 KV cache storage, got {meta.kv_cache_manager.dtype}." + if meta.kv_cache_manager.kv_factor != 1: + return False, "requires MLA SELF-K-only KV cache." + if meta.kv_cache_block_offsets is None: + return False, "requires paged KV cache block offsets." + if meta.high_precision_kv_pool is None: + return False, "requires high-precision KV pool." + if meta.fp4_mla_v_scale_pool is None: + return False, "requires FP4 MLA V-scale pool." + if meta.batch_indices is None or meta.positions is None: + return False, "requires FP4 MLA append metadata." + if ( + meta._paged_kv_indptr is None + or meta.paged_kv_indptr_decode is None + or meta._paged_kv_indices is None + ): + return False, "requires FP4 MLA page metadata." + if meta.tokens_per_block != FP4_MLA_TOKENS_PER_BLOCK: + return ( + False, + f"requires tokens_per_block={FP4_MLA_TOKENS_PER_BLOCK}, " + f"got {meta.tokens_per_block}.", + ) + if fwd.attention_window_size is not None and fwd.attention_window_size < meta.max_seq_len: + return False, "does not support sliding-window attention." + if meta.beam_width != 1: + return False, f"does not support beam search, got beam_width={meta.beam_width}." + fused_head_dim = attn.kv_lora_rank + attn.qk_rope_head_dim + if q.shape[-1] != attn.num_heads * fused_head_dim: + return False, f"unexpected fused query hidden size {q.shape[-1]}." + if fwd.latent_cache.shape[-1] != fused_head_dim: + return False, f"unexpected latent_cache hidden size {fwd.latent_cache.shape[-1]}." + if q.shape[0] != fwd.latent_cache.shape[0]: + return False, "query and latent_cache token counts do not match." + if q.shape[0] < meta.num_generations: + return False, "not enough query tokens for generation batch." + if q.shape[0] % meta.num_generations != 0: + return False, "requires uniform linear MTP generation length." + + return True, "" + + @classmethod + def _is_context_supported_with_reason( + cls, + q: torch.Tensor, + k: Optional[torch.Tensor], + v: Optional[torch.Tensor], + attn: "TrtllmAttention", + meta: "TrtllmAttentionMetadata", + fwd: AttentionForwardArgs, + ) -> Tuple[bool, str]: + if meta.num_contexts <= 0: + return False, "requires context requests." + if getattr(meta, "num_ctx_cached_tokens", 0) != 0: + return False, "does not support cached-context FP4 MLA prefill." + if k is None or v is None: + return False, "requires expanded context K and V tensors." + if fwd.output is None: + return False, "requires output." + if fwd.latent_cache is None: + return False, "requires latent_cache." + if q.dtype not in cls.SUPPORTED_CONTEXT_DTYPES: + return False, f"unsupported context query dtype {q.dtype}." + if k.dtype != q.dtype or v.dtype != q.dtype: + return False, "requires matching context q/k/v dtypes." + if fwd.output.dtype not in cls.SUPPORTED_OUTPUT_DTYPES: + return False, f"unsupported output dtype {fwd.output.dtype}." + if fwd.output_sf is not None: + return False, "does not support quantized context output." + if fwd.attention_mask != PredefinedAttentionMask.CAUSAL: + return False, "requires causal mask." + if fwd.attention_mask_data is not None: + return False, "does not support custom attention masks." + if fwd.attention_sinks is not None: + return False, "does not support attention sinks." + if fwd.sage_attn_num_elts_per_blk_q > 0 or fwd.sage_attn_num_elts_per_blk_k > 0: + return False, "does not support sage attention." + if fwd.sage_attn_num_elts_per_blk_v > 0: + return False, "does not support sage attention." + sparse = fwd.sparse_prediction + if ( + (sparse.sparse_kv_indices is not None and sparse.sparse_kv_indices.numel() > 0) + or (sparse.sparse_attn_indices is not None and sparse.sparse_attn_indices.numel() > 0) + or meta.num_sparse_topk > 0 + ): + return False, "does not support sparse attention." + if meta.helix_position_offsets is not None: + return False, "does not support helix parallelism." + if meta.kv_cache_manager is None: + return False, "requires a KV cache manager." + if meta.kv_cache_manager.dtype != DataType.NVFP4: + return False, f"requires NVFP4 KV cache storage, got {meta.kv_cache_manager.dtype}." + if meta.kv_cache_manager.kv_factor != 1: + return False, "requires MLA SELF-K-only KV cache." + if meta.kv_cache_block_offsets is None: + return False, "requires paged KV cache block offsets." + if meta.high_precision_kv_pool is None: + return False, "requires high-precision KV pool." + if meta.fp4_mla_v_scale_pool is None: + return False, "requires FP4 MLA V-scale pool." + if meta.batch_indices is None or meta.positions is None: + return False, "requires FP4 MLA append metadata." + if meta.paged_kv_indptr_decode is None or meta._paged_kv_indices is None: + return False, "requires FP4 MLA page metadata." + if meta.tokens_per_block != FP4_MLA_TOKENS_PER_BLOCK: + return ( + False, + f"requires tokens_per_block={FP4_MLA_TOKENS_PER_BLOCK}, " + f"got {meta.tokens_per_block}.", + ) + if fwd.attention_window_size is not None and fwd.attention_window_size < meta.max_seq_len: + return False, "does not support sliding-window attention." + if meta.beam_width != 1: + return False, f"does not support beam search, got beam_width={meta.beam_width}." + qk_head_dim = attn.qk_nope_head_dim + attn.qk_rope_head_dim + if q.shape[-1] != attn.num_heads * qk_head_dim: + return False, f"unexpected context query hidden size {q.shape[-1]}." + if k.shape[-1] != attn.num_heads * qk_head_dim: + return False, f"unexpected context key hidden size {k.shape[-1]}." + if v.shape[-1] != attn.num_heads * attn.v_head_dim: + return False, f"unexpected context value hidden size {v.shape[-1]}." + if fwd.latent_cache.shape[-1] != attn.kv_lora_rank + attn.qk_rope_head_dim: + return False, f"unexpected latent_cache hidden size {fwd.latent_cache.shape[-1]}." + if q.shape[0] != meta.num_ctx_tokens: + return False, "query token count must match num_ctx_tokens." + if k.shape[0] != q.shape[0] or v.shape[0] != q.shape[0]: + return False, "context q/k/v token counts do not match." + if fwd.latent_cache.shape[0] < q.shape[0]: + return False, "latent_cache does not contain all context tokens." + + return True, "" + + def run_mla_context(self, params: FmhaParams) -> None: + attn = params.attn + meta = params.meta + fwd = params.fwd + if params.qkv_input is None: + raise RuntimeError("FP4 MLA context requires q input.") + if params.k_input is None or params.v_input is None: + raise RuntimeError("FP4 MLA context requires expanded k/v inputs.") + if params.context_buf is None: + raise RuntimeError("FP4 MLA context requires context_buf.") + if fwd.latent_cache is None: + raise RuntimeError("FP4 MLA context requires latent_cache.") + if meta.positions is None: + raise RuntimeError("FP4 MLA context requires positions.") + + local_layer = attn.get_local_layer_idx(meta) + kv_lora_rank = attn.kv_lora_rank or 0 + qk_nope_head_dim = attn.qk_nope_head_dim or 0 + qk_rope_head_dim = attn.qk_rope_head_dim or 0 + v_head_dim = attn.v_head_dim or 0 + qk_head_dim = qk_nope_head_dim + qk_rope_head_dim + num_tokens = params.num_tokens + + positions = meta.positions[:num_tokens] + q_ctx = params.qkv_input.view(num_tokens, attn.num_heads, qk_head_dim) + k_ctx = params.k_input.view(num_tokens, attn.num_heads, qk_head_dim) + v_ctx = params.v_input.view(num_tokens, attn.num_heads, v_head_dim) + + q_nope = q_ctx[..., :qk_nope_head_dim] + q_pe = q_ctx[..., qk_nope_head_dim:] + k_nope = k_ctx[..., :qk_nope_head_dim] + k_pe = fwd.latent_cache[:num_tokens, kv_lora_rank:].unsqueeze(1) + + q_pe = apply_fp4_mla_rope( + q_pe, + positions, + attn.rotary_cos_sin, + attn.rope_params.max_positions, + qk_rope_head_dim, + ) + k_pe = apply_fp4_mla_rope( + k_pe, + positions, + attn.rotary_cos_sin, + attn.rope_params.max_positions, + qk_rope_head_dim, + ).squeeze(1) + + latent_cache = torch.empty_like(fwd.latent_cache[:num_tokens]) + latent_cache[..., :kv_lora_rank].copy_(fwd.latent_cache[:num_tokens, :kv_lora_rank]) + latent_cache[..., kv_lora_rank:].copy_(k_pe) + scatter_fp4_mla_kv_cache( + meta, + latent_cache, + attn.layer_idx, + token_offset=0, + phase="context", + local_layer=local_layer, + v_head_dim=kv_lora_rank, + ) + update_hp_kv_for_fp4_mla(meta, latent_cache, local_layer, phase="context") + + q_ctx = torch.cat((q_nope, q_pe), dim=-1) + k_ctx = torch.cat((k_nope, k_pe.unsqueeze(1).expand(-1, attn.num_heads, -1)), dim=-1) + output = params.context_buf.view(num_tokens, attn.num_heads, v_head_dim) + sm_scale = 1.0 / (attn.q_scaling * qk_head_dim**0.5) + + host_context_lengths = meta.prompt_lens_cpu_runtime[: meta.num_contexts].tolist() + token_offset = 0 + for context_len in host_context_lengths: + if context_len == 0: + continue + next_offset = token_offset + int(context_len) + q_seq = q_ctx[token_offset:next_offset].transpose(0, 1).unsqueeze(0) + k_seq = k_ctx[token_offset:next_offset].transpose(0, 1).unsqueeze(0) + v_seq = v_ctx[token_offset:next_offset].transpose(0, 1).unsqueeze(0) + out_seq = F.scaled_dot_product_attention( + q_seq, + k_seq, + v_seq, + is_causal=True, + scale=sm_scale, + ) + output[token_offset:next_offset].copy_(out_seq.squeeze(0).transpose(0, 1)) + token_offset = next_offset + + def run_mla_generation(self, params: FmhaParams) -> None: + attn = params.attn + meta = params.meta + fwd = params.fwd + if params.qkv_input is None: + raise RuntimeError("FP4 MLA generation requires qkv_input.") + if params.context_buf is None: + raise RuntimeError("FP4 MLA generation requires context_buf.") + if fwd.latent_cache is None: + raise RuntimeError("FP4 MLA generation requires latent_cache.") + + local_layer = attn.get_local_layer_idx(meta) + kv_lora_rank = attn.kv_lora_rank or 0 + qk_rope_head_dim = attn.qk_rope_head_dim or 0 + fused_head_dim = kv_lora_rank + qk_rope_head_dim + + scatter_fp4_mla_kv_cache( + meta, + fwd.latent_cache, + attn.layer_idx, + token_offset=getattr(meta, "num_ctx_tokens", 0), + phase="generation", + local_layer=local_layer, + v_head_dim=kv_lora_rank, + ) + update_hp_kv_for_fp4_mla(meta, fwd.latent_cache, local_layer, phase="generation") + + query = params.qkv_input.view(params.num_tokens, attn.num_heads, fused_head_dim) + q_nope = query[..., :kv_lora_rank] + q_pe = query[..., kv_lora_rank:] + output = params.context_buf.view(params.num_tokens, attn.num_heads, kv_lora_rank) + sm_scale = 1.0 / (attn.q_scaling * (attn.qk_nope_head_dim + qk_rope_head_dim) ** 0.5) + run_fp4_mla_attention_decode( + meta, + attn.layer_idx, + local_layer, + q_nope, + q_pe, + output, + sm_scale=sm_scale, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + ) diff --git a/tensorrt_llm/_torch/attention_backend/fmha/phased.py b/tensorrt_llm/_torch/attention_backend/fmha/phased.py index aced12909cc9..ad596c0d7b9d 100644 --- a/tensorrt_llm/_torch/attention_backend/fmha/phased.py +++ b/tensorrt_llm/_torch/attention_backend/fmha/phased.py @@ -37,6 +37,8 @@ class FmhaParams: workspace: torch.Tensor attention_input: Optional[torch.Tensor] = None qkv_input: Optional[torch.Tensor] = None + k_input: Optional[torch.Tensor] = None + v_input: Optional[torch.Tensor] = None context_buf: Optional[torch.Tensor] = None sequence_lengths: Optional[torch.Tensor] = None context_lengths: Optional[torch.Tensor] = None @@ -203,6 +205,8 @@ def forward( params.attention_input = q[token_offset : token_offset + num_ctx_tokens] params.qkv_input = params.attention_input + params.k_input = None if k is None else k[token_offset : token_offset + num_ctx_tokens] + params.v_input = None if v is None else v[token_offset : token_offset + num_ctx_tokens] params.context_buf = out_tensor[token_offset : token_offset + num_ctx_tokens] params.sequence_lengths = sequence_length[seq_offset:] params.context_lengths = context_lengths[seq_offset:] @@ -240,6 +244,8 @@ def forward( params.attention_input = q[token_offset : token_offset + num_gen_tokens] params.qkv_input = params.attention_input + params.k_input = None if k is None else k[token_offset : token_offset + num_gen_tokens] + params.v_input = None if v is None else v[token_offset : token_offset + num_gen_tokens] params.context_buf = out_tensor[token_offset : token_offset + num_gen_tokens] params.sequence_lengths = sequence_length[seq_offset:] params.max_past_kv_length = max_past_kv_len diff --git a/tensorrt_llm/_torch/attention_backend/fmha/registry.py b/tensorrt_llm/_torch/attention_backend/fmha/registry.py index 97467657e9b5..e0a5217b23bd 100644 --- a/tensorrt_llm/_torch/attention_backend/fmha/registry.py +++ b/tensorrt_llm/_torch/attention_backend/fmha/registry.py @@ -18,11 +18,13 @@ from .fallback import FallbackFmha from .flashinfer_trtllm_gen import FlashInferTrtllmGenFmha +from .fp4_mla import Fp4MlaFmha from .interface import Fmha FmhaCls: TypeAlias = type[Fmha] FMHA_LIBS: dict[str, FmhaCls] = { + "fp4_mla": Fp4MlaFmha, "flashinfer_trtllm_gen": FlashInferTrtllmGenFmha, "fallback": FallbackFmha, } diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla.py b/tensorrt_llm/_torch/attention_backend/fp4_mla.py index ef97ee2b3a78..3bff8b7d9540 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla.py @@ -9,10 +9,7 @@ do not yet fill a complete FP4 quant block of 16 elements along the sequence dimension. -Used by both ``TrtllmAttention`` (via an internal C++ attention op that reads -both pools) and ``FlashInferAttention`` (via either explicit Python-side -dequant into a BF16 workspace before calling FlashInfer MLA wrappers, or an -env-gated Triton attention path that reads packed FP4 Q, K, and V directly). +Used by the TRTLLM attention backend FP4 MLA FMHA path. """ import os @@ -23,9 +20,6 @@ import triton.language as tl from .fp4_mla_kernels import ( - _fp4_mla_dequant_kernel, - _fp4_mla_overlay_hp_tail_kernel, - _fp4_mla_scatter_kernel, _fp4_mla_v_scale_store_context_tokens_kernel, _fp4_mla_v_scale_store_generation_tiles_kernel, _hp_kv_restore_rejected_from_pool_kernel, @@ -44,8 +38,7 @@ # Max finite e4m3 magnitude for FP4 MLA block-scale clamping. FP4_MLA_E4M3_MAX: float = 448.0 FP4_MLA_Q_RESIDUAL_DIM: int = 64 -FLASHINFER_FP4_MLA_ATTENTION_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION" -FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION_BACKEND" +FP4_MLA_ATTENTION_BACKEND_ENV = "TRTLLM_FP4_MLA_ATTENTION_BACKEND" _HPUpdatePhase = Literal["all", "context", "generation"] _FP4_MLA_MTP_HP_SNAPSHOTS = "_fp4_mla_mtp_hp_snapshots" @@ -53,15 +46,6 @@ # Environment helpers -def _env_enabled(name: str) -> bool: - return os.getenv(name, "0").lower() in ( - "1", - "true", - "yes", - "on", - ) - - def _env_enabled_default(name: str, default: bool) -> bool: value = os.getenv(name) if value is None or value == "": @@ -81,13 +65,12 @@ def _env_int(name: str) -> Optional[int]: return int(value) -def is_flashinfer_fp4_mla_attention_enabled() -> bool: - """Return whether FlashInfer MLA should allocate no-dequant FP4 attention buffers.""" - return _env_enabled(FLASHINFER_FP4_MLA_ATTENTION_ENV) +def _fp4_mla_attention_backend() -> str: + return os.getenv(FP4_MLA_ATTENTION_BACKEND_ENV, "triton").lower() -def _fp4_mla_attention_backend() -> str: - return os.getenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, "triton").lower() +def _cutile_backend_available() -> bool: + return hasattr(tl, "ext") def _ceil_div(lhs: int, rhs: int) -> int: @@ -107,6 +90,43 @@ def _get_sm_count(device: torch.device) -> int: return count +def apply_fp4_mla_rope( + target: torch.Tensor, + positions: torch.Tensor, + rotary_cos_sin: torch.Tensor, + max_positions: int, + rope_dim: int, +) -> torch.Tensor: + """Apply TRTLLM MLA RoPE using adjacent GPT-J-style element pairs.""" + if rope_dim % 2 != 0: + raise RuntimeError(f"MLA RoPE dimension must be even, got {rope_dim}.") + if target.shape[-1] != rope_dim: + raise RuntimeError( + f"MLA RoPE target last dimension must be {rope_dim}, got {target.shape[-1]}." + ) + + positions = positions[: target.shape[0]].to(torch.long) + cos_sin = rotary_cos_sin.view(max_positions, rope_dim, 2).index_select(0, positions) + pair_count = rope_dim // 2 + + math_dtype = target.dtype + if math_dtype == torch.float8_e4m3fn: + math_dtype = torch.bfloat16 + target_math = target.to(math_dtype) + cos = cos_sin[:, :pair_count, 0].to(math_dtype) + sin = cos_sin[:, :pair_count, 1].to(math_dtype) + while cos.ndim < target_math.ndim: + cos = cos.unsqueeze(1) + sin = sin.unsqueeze(1) + + target_even = target_math[..., 0::2] + target_odd = target_math[..., 1::2] + target_roped = torch.empty_like(target_math) + target_roped[..., 0::2] = target_even * cos - target_odd * sin + target_roped[..., 1::2] = target_odd * cos + target_even * sin + return target_roped.to(target.dtype) + + def _host_int_list_during_forward(value: Any, start: int, end: int) -> Optional[list[int]]: if torch.cuda.is_current_stream_capturing(): return None @@ -138,10 +158,6 @@ def _get_fp4_mla_swizzled_scale_size(rows: int, cols: int) -> int: return row_groups * col_groups * 32 * 16 -def _use_fp4_mla_swizzled_sf() -> bool: - return is_flashinfer_fp4_mla_attention_enabled() - - def _get_fp4_mla_context_start_positions(metadata: Any, num_contexts: int) -> torch.Tensor: kv_cache_params = getattr(metadata, "kv_cache_params", None) cached_token_lens = getattr(kv_cache_params, "num_cached_tokens_per_seq", None) @@ -438,65 +454,6 @@ def _scatter_fp4_mla_kv_cache_2d_generation( ) -def _scatter_fp4_mla_kv_cache_1d( - metadata: Any, - latent_cache: torch.Tensor, - kv_cache: torch.Tensor, - sf_cache: torch.Tensor, - global_scale: torch.Tensor, - *, - layer_idx: int, - token_offset: int, - num_tokens: int, - head_dim: int, - sf_per_token: int, - use_swizzled_sf: bool, -) -> None: - q_fp4, q_sf = torch.ops.trtllm.fp4_quantize( - latent_cache, global_scale, FP4_BLOCK_SIZE, False, False - ) - q_sf = q_sf.view(num_tokens, head_dim // FP4_BLOCK_SIZE) - - packed_dim = head_dim // 2 - block_packed_dim = triton.next_power_of_2(packed_dim) - block_sf = triton.next_power_of_2(sf_per_token) - - _fp4_mla_scatter_kernel[(num_tokens,)]( - kv_cache, - sf_cache, - q_fp4, - q_sf, - metadata.batch_indices, - metadata.positions, - metadata.paged_kv_indices, - metadata.paged_kv_indptr, - metadata.paged_kv_indices.shape[0], - metadata.paged_kv_indptr.shape[0], - kv_cache.shape[0], - token_offset, - metadata.page_size, - kv_cache.stride(0), - kv_cache.stride(1), - kv_cache.stride(2), - kv_cache.stride(3), - kv_cache.stride(4), - sf_cache.stride(0), - sf_cache.stride(1), - sf_cache.stride(2), - sf_cache.stride(3), - sf_cache.stride(4), - q_fp4.stride(0), - q_fp4.stride(1), - q_sf.stride(0), - q_sf.stride(1), - PACKED_D=packed_dim, - SF_PER_TOKEN=sf_per_token, - BLOCK_PACKED_D=block_packed_dim, - BLOCK_SF=block_sf, - USE_SWIZZLED_SF=use_swizzled_sf, - ) - - # Public cache update and decode entry points @@ -521,12 +478,11 @@ def scatter_fp4_mla_kv_cache( slices ``latent_cache[:num_ctx_tokens]`` for context and ``latent_cache[num_ctx_tokens:]`` for generation before dispatching. - When the no-dequant FP4 MLA attention path is enabled, callers should pass - ``phase``, ``local_layer``, and ``v_head_dim``. Context scatter then writes - the final FP4 tile representation directly: dimensions below + Callers must pass ``phase``, ``local_layer``, and ``v_head_dim``. Context + scatter writes the final FP4 tile representation directly: dimensions below ``v_head_dim`` use one shared 16-token by 16-dim scale written into both - K's token-major scale layout and V's dim-major scale layout. Tail K-only - dimensions use K's per-token 1D scales. Generation scatter rewrites each + K's token-major scale layout and V's dim-major scale layout. Tail K-only + dimensions use K's per-token 1D scales. Generation scatter rewrites each touched 16-token tile by reading old tokens from the HP pool and new tokens from ``latent_cache``; callers then update the HP pool after scatter. """ @@ -552,175 +508,94 @@ def scatter_fp4_mla_kv_cache( "range (see MLA.forward_impl)." ) - use_swizzled_sf = _use_fp4_mla_swizzled_sf() - if use_swizzled_sf: - _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) + _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) global_scale = _get_fp4_mla_global_scale(metadata, latent_cache.device) kv_cache, sf_cache = _get_fp4_mla_kv_cache_tensors(metadata, layer_idx) sf_per_token = head_dim // FP4_BLOCK_SIZE - use_2d_scatter = ( - use_swizzled_sf - and phase in ("context", "generation") - and getattr(metadata, "fp4_mla_v_scale_pool", None) is not None - ) - if use_2d_scatter: - assert phase is not None - if local_layer is None or v_head_dim is None: - raise ValueError("Real FP4 MLA scatter requires local_layer and v_head_dim.") - if metadata.page_size % HP_BLOCK_SIZE != 0: - raise ValueError( - f"FP4 MLA scatter requires page_size divisible by " - f"{HP_BLOCK_SIZE}, got {metadata.page_size}." - ) - if v_head_dim > head_dim: - raise ValueError(f"FP4 MLA v_head_dim={v_head_dim} cannot exceed head_dim={head_dim}.") - if v_head_dim % FP4_BLOCK_SIZE != 0: - raise ValueError( - f"FP4 MLA v_head_dim must be divisible by {FP4_BLOCK_SIZE}, got {v_head_dim}." - ) + if phase not in ("context", "generation"): + raise ValueError("FP4 MLA scatter requires phase='context' or 'generation'.") + if getattr(metadata, "fp4_mla_v_scale_pool", None) is None: + raise RuntimeError("FP4 MLA scatter requires the auxiliary V scale pool.") + if local_layer is None or v_head_dim is None: + raise ValueError("FP4 MLA scatter requires local_layer and v_head_dim.") + if metadata.page_size % HP_BLOCK_SIZE != 0: + raise ValueError( + f"FP4 MLA scatter requires page_size divisible by " + f"{HP_BLOCK_SIZE}, got {metadata.page_size}." + ) + if v_head_dim > head_dim: + raise ValueError(f"FP4 MLA v_head_dim={v_head_dim} cannot exceed head_dim={head_dim}.") + if v_head_dim % FP4_BLOCK_SIZE != 0: + raise ValueError( + f"FP4 MLA v_head_dim must be divisible by {FP4_BLOCK_SIZE}, got {v_head_dim}." + ) - sf_cache = sf_cache.view(torch.float8_e4m3fn) - v_sf = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=v_head_dim) - num_dim_blocks = triton.cdiv(head_dim, FP4_BLOCK_SIZE) - sf_per_page = metadata.page_size // HP_BLOCK_SIZE + sf_cache = sf_cache.view(torch.float8_e4m3fn) + v_sf = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=v_head_dim) + num_dim_blocks = triton.cdiv(head_dim, FP4_BLOCK_SIZE) + sf_per_page = metadata.page_size // HP_BLOCK_SIZE - if phase == "context": - _scatter_fp4_mla_kv_cache_2d_context( - metadata, - latent_cache, - kv_cache, - sf_cache, - v_sf, - global_scale, - token_offset=token_offset, - local_layer=local_layer, - v_head_dim=v_head_dim, - head_dim=head_dim, - num_tokens=num_tokens, - num_dim_blocks=num_dim_blocks, - sf_per_token=sf_per_token, - sf_per_page=sf_per_page, - ) - else: - _scatter_fp4_mla_kv_cache_2d_generation( - metadata, - latent_cache, - kv_cache, - sf_cache, - v_sf, - global_scale, - local_layer=local_layer, - v_head_dim=v_head_dim, - head_dim=head_dim, - num_tokens=num_tokens, - num_dim_blocks=num_dim_blocks, - sf_per_token=sf_per_token, - sf_per_page=sf_per_page, - ) - if phase == "context": - v_pack_page_ids = metadata.paged_kv_indices - else: - num_gen_blocks = metadata.num_generation_blocks - v_pack_page_ids = metadata.paged_kv_indices[ - metadata.num_context_blocks : metadata.num_context_blocks + num_gen_blocks - ] - _maybe_update_cutile_v_packed_cache( + if phase == "context": + _scatter_fp4_mla_kv_cache_2d_context( metadata, - layer_idx, + latent_cache, kv_cache, - v_pack_page_ids, - v_head_dim=v_head_dim, - page_size=metadata.page_size, + sf_cache, + v_sf, + global_scale, + token_offset=token_offset, local_layer=local_layer, - v_sf=v_sf[local_layer], + v_head_dim=v_head_dim, + head_dim=head_dim, + num_tokens=num_tokens, + num_dim_blocks=num_dim_blocks, + sf_per_token=sf_per_token, + sf_per_page=sf_per_page, ) - _maybe_update_triton_v_packed_cache( + v_pack_page_ids = metadata.paged_kv_indices + else: + _scatter_fp4_mla_kv_cache_2d_generation( metadata, - layer_idx, + latent_cache, kv_cache, - v_pack_page_ids, - num_queries=num_tokens, - v_head_dim=v_head_dim, - page_size=metadata.page_size, + sf_cache, + v_sf, + global_scale, local_layer=local_layer, - v_sf=v_sf[local_layer], + v_head_dim=v_head_dim, + head_dim=head_dim, + num_tokens=num_tokens, + num_dim_blocks=num_dim_blocks, + sf_per_token=sf_per_token, + sf_per_page=sf_per_page, ) - return - - _scatter_fp4_mla_kv_cache_1d( + num_gen_blocks = metadata.num_generation_blocks + v_pack_page_ids = metadata.paged_kv_indices[ + metadata.num_context_blocks : metadata.num_context_blocks + num_gen_blocks + ] + _maybe_update_cutile_v_packed_cache( metadata, - latent_cache, + layer_idx, kv_cache, - sf_cache, - global_scale, - layer_idx=layer_idx, - token_offset=token_offset, - num_tokens=num_tokens, - head_dim=head_dim, - sf_per_token=sf_per_token, - use_swizzled_sf=use_swizzled_sf, - ) - - -def _ensure_decode_workspace( - metadata: Any, - head_dim: int, - dtype: torch.dtype, -) -> torch.Tensor: - num_blocks = _get_decode_workspace_num_blocks(metadata) - workspace = getattr(metadata, "_fp4_mla_decode_cache_buf", None) - needs_alloc = ( - workspace is None - or workspace.shape[0] < max(num_blocks, 1) - or workspace.shape[1] != metadata.page_size - or workspace.shape[2] != head_dim - or workspace.dtype != dtype + v_pack_page_ids, + v_head_dim=v_head_dim, + page_size=metadata.page_size, + local_layer=local_layer, + v_sf=v_sf[local_layer], ) - if needs_alloc: - if torch.cuda.is_current_stream_capturing(): - raise ValueError( - "Cannot allocate FlashInfer FP4 MLA decode workspace while " - "capturing a CUDA graph. Run a warmup prepare/forward first." - ) - workspace = torch.empty( - (max(num_blocks, 1), metadata.page_size, head_dim), - dtype=dtype, - device=metadata.paged_kv_indices.device, - ) - metadata._fp4_mla_decode_cache_buf = workspace - return workspace[:num_blocks] - - -def _get_decode_workspace_num_blocks(metadata: Any) -> int: - if metadata.is_cuda_graph: - max_blocks_per_seq = ( - metadata.kv_cache_manager.max_seq_len + metadata.page_size - 1 - ) // metadata.page_size - max_graph_blocks = metadata.max_num_requests * max_blocks_per_seq - return min( - metadata.kv_cache_manager.blocks_in_primary_pool, - max_graph_blocks, - ) - return metadata.num_generation_blocks - - -def _get_decode_src_page_ids(metadata: Any, num_blocks: int) -> torch.Tensor: - page_ids = ( - metadata._paged_kv_indices - if metadata.is_cuda_graph and hasattr(metadata, "_paged_kv_indices") - else metadata.paged_kv_indices + _maybe_update_triton_v_packed_cache( + metadata, + layer_idx, + kv_cache, + v_pack_page_ids, + num_queries=num_tokens, + v_head_dim=v_head_dim, + page_size=metadata.page_size, + local_layer=local_layer, + v_sf=v_sf[local_layer], ) - src_page_ids = page_ids[metadata.num_context_blocks : metadata.num_context_blocks + num_blocks] - if src_page_ids.numel() != num_blocks: - raise RuntimeError( - f"FP4 MLA dequant needs {num_blocks} decode page ids from " - f"paged_kv_indices[{metadata.num_context_blocks}:" - f"{metadata.num_context_blocks + num_blocks}], got " - f"{src_page_ids.numel()}." - ) - return src_page_ids def _validate_fp4_mla_cache_shape(page_size: int, head_dim: int) -> None: @@ -759,85 +634,6 @@ def _validate_fp4_mla_attention_q_shape(head_dim: int, q_residual_dim: int) -> N ) -def get_fp4_mla_decode_cache( - metadata: Any, - layer_idx: int, - local_layer: int, - *, - head_dim: int, - dtype: torch.dtype, -) -> torch.Tensor: - """Build a compact dequantized MLA cache for FlashInfer decode.""" - # Must match scatter_fp4_mla_kv_cache: the env var picks the SF layout, - # not the page size. When the dequant fallback is the read path - # (env disabled), scatter wrote linear scales and we must read linear. - use_swizzled_sf = _use_fp4_mla_swizzled_sf() - if use_swizzled_sf: - _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) - combined = _ensure_decode_workspace(metadata, head_dim, dtype) - num_blocks = combined.shape[0] - if num_blocks == 0: - return combined - - kv_cache, sf_cache = _get_fp4_mla_kv_cache_tensors(metadata, layer_idx) - sf_cache = sf_cache.view(torch.float8_e4m3fn) - global_scale = _get_fp4_mla_global_scale(metadata, combined.device) - src_page_ids = _get_decode_src_page_ids(metadata, num_blocks) - block_d = triton.next_power_of_2(head_dim) - - _fp4_mla_dequant_kernel[(num_blocks, metadata.page_size)]( - combined, - kv_cache, - sf_cache, - global_scale, - src_page_ids, - src_page_ids.shape[0], - kv_cache.shape[0], - kv_cache.stride(0), - kv_cache.stride(1), - kv_cache.stride(2), - kv_cache.stride(3), - kv_cache.stride(4), - sf_cache.stride(0), - sf_cache.stride(1), - sf_cache.stride(2), - sf_cache.stride(3), - sf_cache.stride(4), - combined.stride(0), - combined.stride(1), - combined.stride(2), - D=head_dim, - FP4_BLOCK=FP4_BLOCK_SIZE, - BLOCK_D=block_d, - USE_SWIZZLED_SF=use_swizzled_sf, - ) - - num_gen = metadata.num_seqs - metadata.num_contexts - if num_gen > 0 and metadata.high_precision_kv_pool is not None: - pool = metadata.high_precision_kv_pool - _fp4_mla_overlay_hp_tail_kernel[(num_gen, HP_BLOCK_SIZE)]( - combined, - pool, - metadata.seq_slots[metadata.num_contexts : metadata.num_seqs], - metadata.kv_lens_cuda_runtime[metadata.num_contexts : metadata.num_seqs], - metadata.paged_kv_indptr_decode, - pool.shape[0], - pool.shape[1], - combined.shape[0], - local_layer, - metadata.page_size, - combined.stride(0), - combined.stride(1), - combined.stride(2), - pool.stride(0), - pool.stride(1), - D=head_dim, - BLOCK_D=block_d, - HP_BLOCK=HP_BLOCK_SIZE, - ) - return combined - - def _ensure_workspace_tensor( metadata: Any, attr_name: str, @@ -870,6 +666,8 @@ def _ensure_workspace_tensor( def _cutile_persistent_v_pack_enabled() -> bool: if _fp4_mla_attention_backend() != "cutile": return False + if not _cutile_backend_available(): + return False return os.getenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1").lower() not in ( "0", "false", @@ -2334,11 +2132,6 @@ def run_fp4_mla_attention_decode( V-view scale pool. No BF16 dequantized KV workspace is materialized on this path. """ - if not is_flashinfer_fp4_mla_attention_enabled(): - raise RuntimeError( - f"FP4 MLA attention decode requires {FLASHINFER_FP4_MLA_ATTENTION_ENV}=1." - ) - head_dim = kv_lora_rank + qk_rope_head_dim _validate_fp4_mla_cache_shape(metadata.page_size, head_dim) if metadata.page_size != FP4_MLA_TOKENS_PER_BLOCK: @@ -2426,6 +2219,12 @@ def run_fp4_mla_attention_decode( return if backend == "cutile": + if not _cutile_backend_available(): + raise RuntimeError( + "FP4 MLA cutile attention backend requires Triton tl.ext APIs, " + "which are unavailable in this runtime. Set " + f"{FP4_MLA_ATTENTION_BACKEND_ENV}=triton." + ) from .fp4_mla_cutile import fp4_mla_paged_attention total_p_rows = num_queries * max_pages * num_heads @@ -2635,41 +2434,38 @@ def run_fp4_mla_attention_decode( if backend != "triton": raise ValueError( f"Unsupported FP4 MLA attention backend '{backend}'. " - f"Set {FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV} to " - "'triton' or 'cutile'." + f"Set {FP4_MLA_ATTENTION_BACKEND_ENV} to 'triton' or 'cutile'." ) # Self-contained public-Triton path: TMA-loaded QK + fused page-stats pack, # reduce-stats, prob-scale, and PV with an optional prepacked V cache. - if backend == "triton": - _run_triton_attention_decode( - metadata=metadata, - layer_idx=layer_idx, - local_layer=local_layer, - q_fp4=q_fp4, - q_sf=q_sf.contiguous().view(-1), - kv_cache=kv_cache, - sf_cache=sf_cache, - v_sf=v_sf, - global_scale=global_scale, - src_page_ids=src_page_ids, - kv_lens=kv_lens, - p_fp4=p_fp4, - p_sf=p_sf, - max_scores=max_scores, - denom=denom, - output=output, - num_queries=num_queries, - num_heads=num_heads, - head_dim=head_dim, - kv_lora_rank=kv_lora_rank, - q_residual_dim=q_residual_dim, - query_len_per_seq=query_len_per_seq, - max_pages=max_pages, - sm_scale=float(sm_scale), - q_global_scale=q_global_scale, - ) - return + _run_triton_attention_decode( + metadata=metadata, + layer_idx=layer_idx, + local_layer=local_layer, + q_fp4=q_fp4, + q_sf=q_sf.contiguous().view(-1), + kv_cache=kv_cache, + sf_cache=sf_cache, + v_sf=v_sf, + global_scale=global_scale, + src_page_ids=src_page_ids, + kv_lens=kv_lens, + p_fp4=p_fp4, + p_sf=p_sf, + max_scores=max_scores, + denom=denom, + output=output, + num_queries=num_queries, + num_heads=num_heads, + head_dim=head_dim, + kv_lora_rank=kv_lora_rank, + q_residual_dim=q_residual_dim, + query_len_per_seq=query_len_per_seq, + max_pages=max_pages, + sm_scale=float(sm_scale), + q_global_scale=q_global_scale, + ) def _hp_pool_layer_view( diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py index 63ab10f36398..3ffcc79437b7 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py @@ -297,104 +297,6 @@ def _fp4_mla_swizzled_sf_offset( ) -@triton.jit -def _fp4_mla_scatter_kernel( - kv_cache_ptr, - sf_cache_ptr, - q_fp4_ptr, - q_sf_ptr, - batch_indices_ptr, - positions_ptr, - paged_kv_indices_ptr, - paged_kv_indptr_ptr, - page_ids_len, - indptr_len, - num_pages, - token_offset, - page_size, - kv_s0, - kv_s1, - kv_s2, - kv_s3, - kv_s4, - sf_s0, - sf_s1, - sf_s2, - sf_s3, - sf_s4, - q_fp4_s0, - q_fp4_s1, - q_sf_s0, - q_sf_s1, - PACKED_D: tl.constexpr, - SF_PER_TOKEN: tl.constexpr, - BLOCK_PACKED_D: tl.constexpr, - BLOCK_SF: tl.constexpr, - USE_SWIZZLED_SF: tl.constexpr, -): - token_idx = tl.program_id(0) - metadata_token_idx = token_offset + token_idx - - # Keep page address math in int64; 128-token FP4 pages can overflow - # int32 offsets in large KV pools. - batch_idx = tl.load(batch_indices_ptr + metadata_token_idx).to(tl.int64) - position = tl.load(positions_ptr + metadata_token_idx).to(tl.int64) - if (batch_idx < 0) | (batch_idx + 1 >= indptr_len) | (position < 0): - return - - page_size_i64 = tl.cast(page_size, tl.int64) - page_idx = position // page_size_i64 - page_pos = position - page_idx * page_size_i64 - page_start = tl.load(paged_kv_indptr_ptr + batch_idx).to(tl.int64) - page_end = tl.load(paged_kv_indptr_ptr + batch_idx + 1).to(tl.int64) - page_table_offset = page_start + page_idx - if ( - (page_pos < 0) - | (page_pos >= page_size_i64) - | (page_table_offset < page_start) - | (page_table_offset >= page_end) - | (page_table_offset < 0) - | (page_table_offset >= page_ids_len) - ): - return - physical_page = tl.load(paged_kv_indices_ptr + page_table_offset).to(tl.int64) - if (physical_page < 0) | (physical_page >= num_pages): - return - - offs_packed = tl.arange(0, BLOCK_PACKED_D) - mask_packed = offs_packed < PACKED_D - safe_offs_packed = tl.where(mask_packed, offs_packed, 0) - q_vals = tl.load( - q_fp4_ptr + token_idx * q_fp4_s0 + safe_offs_packed * q_fp4_s1, - mask=mask_packed, - other=0, - ) - kv_dst = physical_page * kv_s0 + page_pos * kv_s2 - tl.store(kv_cache_ptr + kv_dst + safe_offs_packed * kv_s4, q_vals, mask=mask_packed) - - offs_sf = tl.arange(0, BLOCK_SF) - mask_sf = offs_sf < SF_PER_TOKEN - # Masked lanes are predicated off, but the address arithmetic still runs. - # In the swizzled layout, out-of-range cols land past the per-page stride; - # in the linear layout, they spill ~(BLOCK_SF - SF_PER_TOKEN) bytes past - # each page row. For the last physical page either case can fall outside - # the sf_cache allocation. Pin masked lanes to col 0 so all computed - # addresses stay in-bounds regardless of allocator slack. - safe_offs_sf = tl.where(mask_sf, offs_sf, 0) - sf_vals = tl.load( - q_sf_ptr + token_idx * q_sf_s0 + safe_offs_sf * q_sf_s1, - mask=mask_sf, - other=0, - ) - if USE_SWIZZLED_SF: - sf_offsets = _fp4_mla_swizzled_sf_offset(page_pos, safe_offs_sf, SF_PER_TOKEN) - sf_dst = physical_page * sf_s0 - tl.store(sf_cache_ptr + sf_dst + sf_offsets, sf_vals, mask=mask_sf) - else: - sf_dst = physical_page * sf_s0 + page_pos * sf_s2 - tl.store(sf_cache_ptr + sf_dst + safe_offs_sf * sf_s4, sf_vals, mask=mask_sf) - - # FP4 conversion and cache kernels @@ -924,127 +826,6 @@ def _fp4_mla_v_scale_store_generation_tiles_kernel( ) -@triton.jit -def _fp4_mla_load_values( - kv_cache_ptr, - sf_cache_ptr, - physical_page, - page_pos, - offs_d, - mask_d, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - sf_s2, - sf_s4, - D: tl.constexpr, - FP4_BLOCK: tl.constexpr, - SF_PER_TOKEN: tl.constexpr, - USE_SWIZZLED_SF: tl.constexpr, -): - packed_offsets = offs_d // 2 - packed = tl.load( - kv_cache_ptr + physical_page * kv_s0 + page_pos * kv_s2 + packed_offsets * kv_s4, - mask=mask_d, - other=0, - ) - low = packed & 0x0F - high = (packed >> 4) & 0x0F - nibble = tl.where((offs_d & 1) == 0, low, high) - - scale_offsets = offs_d // FP4_BLOCK - # See _fp4_mla_scatter_kernel for the rationale: pin masked lanes to col 0 - # so the address arithmetic stays inside the per-page stride for the last - # physical page, regardless of which SF layout is in use. - safe_scale_offsets = tl.where(mask_d, scale_offsets, 0) - if USE_SWIZZLED_SF: - swizzled_sf_offsets = _fp4_mla_swizzled_sf_offset( - page_pos, safe_scale_offsets, SF_PER_TOKEN - ) - scale = tl.load( - sf_cache_ptr + physical_page * sf_s0 + swizzled_sf_offsets, - mask=mask_d, - other=0.0, - ).to(tl.float32) - else: - scale = tl.load( - sf_cache_ptr + physical_page * sf_s0 + page_pos * sf_s2 + safe_scale_offsets * sf_s4, - mask=mask_d, - other=0.0, - ).to(tl.float32) - return _fp4_e2m1_to_f32(nibble) * scale - - -@triton.jit -def _fp4_mla_dequant_kernel( - out_ptr, - kv_cache_ptr, - sf_cache_ptr, - global_scale_ptr, - src_page_ids_ptr, - page_ids_len, - num_pages, - kv_s0, - kv_s1, - kv_s2, - kv_s3, - kv_s4, - sf_s0, - sf_s1, - sf_s2, - sf_s3, - sf_s4, - out_s0, - out_s1, - out_s2, - D: tl.constexpr, - FP4_BLOCK: tl.constexpr, - BLOCK_D: tl.constexpr, - USE_SWIZZLED_SF: tl.constexpr, -): - compact_page = tl.program_id(0).to(tl.int64) - page_pos = tl.program_id(1).to(tl.int64) - valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) - safe_compact_page = tl.where(valid_compact_page, compact_page, 0) - physical_page = tl.load( - src_page_ids_ptr + safe_compact_page, - mask=valid_compact_page, - other=-1, - ).to(tl.int64) - valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) - safe_physical_page = tl.where(valid_physical_page, physical_page, 0) - global_scale = tl.load(global_scale_ptr) - - offs_d = tl.arange(0, BLOCK_D) - mask_d = offs_d < D - safe_offs_d = tl.where(mask_d, offs_d, 0) - value = _fp4_mla_load_values( - kv_cache_ptr, - sf_cache_ptr, - safe_physical_page, - page_pos, - safe_offs_d, - mask_d & valid_physical_page, - kv_s0, - kv_s2, - kv_s4, - sf_s0, - sf_s2, - sf_s4, - D, - FP4_BLOCK, - D // FP4_BLOCK, - USE_SWIZZLED_SF, - ) - - tl.store( - out_ptr + compact_page * out_s0 + page_pos * out_s1 + safe_offs_d * out_s2, - value / global_scale, - mask=mask_d, - ) - - @triton.jit def _fp4_mla_qk_scores_tile( q_fp4_ptr, @@ -1576,58 +1357,3 @@ def _fp4_mla_attention_pv_kernel( acc / (global_scale * P_GLOBAL_SCALE), mask=mask_h[:, None] & mask_v[None, :], ) - - -@triton.jit -def _fp4_mla_overlay_hp_tail_kernel( - out_ptr, - pool_ptr, - seq_slots_ptr, - kv_lens_ptr, - paged_kv_indptr_decode_ptr, - num_seq_slots, - num_layers, - num_pages, - layer_idx, - page_size, - out_s0, - out_s1, - out_s2, - pool_stride_seq, - pool_stride_layer, - D: tl.constexpr, - BLOCK_D: tl.constexpr, - HP_BLOCK: tl.constexpr, -): - gen_idx = tl.program_id(0) - tail_idx = tl.program_id(1) - if (layer_idx < 0) | (layer_idx >= num_layers): - return - - kv_len = tl.load(kv_lens_ptr + gen_idx) - tail_count = kv_len % HP_BLOCK - if tail_idx >= tail_count: - return - - abs_pos = (kv_len - tail_count + tail_idx).to(tl.int64) - page_size_i64 = tl.cast(page_size, tl.int64) - rel_page = abs_pos // page_size_i64 - page_pos = abs_pos - rel_page * page_size_i64 - compact_page = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) + rel_page - if (compact_page < 0) | (compact_page >= num_pages): - return - hp_slot = abs_pos % HP_BLOCK - seq_slot = tl.load(seq_slots_ptr + gen_idx).to(tl.int64) - if (seq_slot < 0) | (seq_slot >= num_seq_slots): - return - - offs_d = tl.arange(0, BLOCK_D) - mask_d = offs_d < D - safe_offs_d = tl.where(mask_d, offs_d, 0) - src_base = seq_slot * pool_stride_seq + layer_idx * pool_stride_layer + hp_slot * D - value = tl.load(pool_ptr + src_base + safe_offs_d, mask=mask_d, other=0.0) - tl.store( - out_ptr + compact_page * out_s0 + page_pos * out_s1 + safe_offs_d * out_s2, - value, - mask=mask_d, - ) diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py index 07d09ab007f4..c29a00913c8b 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py @@ -2,10 +2,9 @@ # SPDX-License-Identifier: Apache-2.0 """Triton FP4 MLA decode-path kernels. -This file owns the no-dequant ``triton`` attention backend selected via -``TRTLLM_FLASHINFER_FP4_MLA_ATTENTION_BACKEND=triton``. It is separate from -``fp4_mla_kernels.py``, which holds shared KV-cache scatter/dequant and HP-pool -helper kernels. +This file owns the no-dequant Triton attention kernels used by the TRTLLM +FP4 MLA FMHA path. It is separate from ``fp4_mla_kernels.py``, which holds +shared KV-cache scatter and HP-pool helper kernels. Optimizations -- self-contained, public-Triton compatible (no ``tl.ext.*`` or any private-Triton extension): diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index 07d80d80b0f4..95ec6f65ec49 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -38,7 +38,7 @@ from ..utils import (compute_swizzled_sf_shape, get_global_attrs, get_model_extra_attrs) -from .fp4_mla import HP_BLOCK_SIZE, update_hp_kv_for_fp4_mla +from .fp4_mla import FP4_MLA_KV_GLOBAL_SCALE, HP_BLOCK_SIZE, apply_fp4_mla_rope from .interface import (AttentionBackend, AttentionForwardArgs, AttentionInputType, AttentionMask, AttentionMetadata, KVCacheParams, MLAParams, PositionalEmbeddingParams, @@ -48,7 +48,6 @@ from .sparse.skip_softmax import SkipSoftmaxParams - @functools.cache def generate_spec_decoding_position_offsets(max_num_requests: int, draft_len: int) -> torch.Tensor: @@ -165,6 +164,16 @@ class TrtllmAttentionMetadata(AttentionMetadata): # Standalone tensor, not part of the block-based paged KV cache. high_precision_kv_pool: Optional[torch.Tensor] = None fp4_mla_hp_snapshot_pool: Optional[torch.Tensor] = None + fp4_mla_v_scale_pool: Optional[torch.Tensor] = None + _fp4_mla_global_scale: Optional[torch.Tensor] = None + batch_indices: Optional[torch.Tensor] = None + positions: Optional[torch.Tensor] = None + _paged_kv_indptr: Optional[torch.Tensor] = None + paged_kv_indptr_decode: Optional[torch.Tensor] = None + _paged_kv_indices: Optional[torch.Tensor] = None + num_blocks: Optional[List[int]] = None + num_context_blocks: int = 0 + num_generation_blocks: int = 0 # Pre-computed FlashMLA tile-scheduler metadata and num_splits. # Computed once per forward pass in TrtllmAttention.forward() and reused across layers. @@ -242,6 +251,33 @@ def tokens_per_block(self) -> Optional[int]: """ return self.kv_cache_manager.tokens_per_block if self.kv_cache_manager is not None else None + @property + def page_size(self) -> int: + """ + Number of tokens per cache page. + """ + assert self.kv_cache_manager is not None, "page_size requires a KV cache manager" + return self.kv_cache_manager.tokens_per_block + + @property + def paged_kv_indices(self) -> torch.Tensor: + """ + Compact flattened page table used by FP4 MLA helper kernels. + """ + if self._paged_kv_indices is None: + raise RuntimeError("paged_kv_indices is not allocated.") + total_blocks = self.num_context_blocks + self.num_generation_blocks + return self._paged_kv_indices[:total_blocks] + + @property + def paged_kv_indptr(self) -> torch.Tensor: + """ + Compact page-table indptr used by FP4 MLA helper kernels. + """ + if self._paged_kv_indptr is None: + raise RuntimeError("paged_kv_indptr is not allocated.") + return self._paged_kv_indptr[:self.num_seqs + 1] + @property def host_kv_cache_pool_pointers(self) -> Optional[torch.Tensor]: """ @@ -485,6 +521,51 @@ def _post_init_with_buffers(self, buffers) -> None: ) else: self.fp4_mla_hp_snapshot_pool = None + self.batch_indices = self.get_empty( + buffers, + (self.max_num_tokens, ), + cache_name="fp4_mla_batch_indices", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self.positions = self.get_empty( + buffers, + (self.max_num_tokens, ), + cache_name="fp4_mla_positions", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self._paged_kv_indices = self.get_empty( + buffers, + (self.kv_cache_manager.blocks_in_primary_pool, ), + cache_name="fp4_mla_paged_kv_indices", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self._paged_kv_indptr = self.get_empty( + buffers, + (self.max_num_sequences + 1, ), + cache_name="fp4_mla_paged_kv_indptr", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self.paged_kv_indptr_decode = self.get_empty( + buffers, + (self.max_num_sequences + 1, ), + cache_name="fp4_mla_paged_kv_indptr_decode", + dtype=torch.int32, + capture_graph=capture_graph, + ) + self._fp4_mla_global_scale = self.get_empty( + buffers, + (1, ), + cache_name="fp4_mla_global_scale", + dtype=torch.float32, + capture_graph=capture_graph, + ) + self._fp4_mla_global_scale.fill_(FP4_MLA_KV_GLOBAL_SCALE) + self.fp4_mla_v_scale_pool = self.kv_cache_manager.get_mla_v_scale_pool( + ) logger.info( f"Allocated high-precision BF16 KV pool: shape=" f"{list(self.high_precision_kv_pool.shape)}, " @@ -529,6 +610,15 @@ def update_for_spec_dec(self) -> None: # so that forward() recomputes it for the next sub-step. if self.enable_flash_mla: self._flash_mla_metadata_valid = False + if self.high_precision_kv_pool is None: + return + + num_seqs = self.num_seqs + self.prompt_lens_cuda_runtime = self.seq_lens_kv_cuda[:num_seqs] + if not torch.cuda.is_current_stream_capturing(): + self.prompt_lens_cpu_runtime = self.seq_lens_kv[:num_seqs] + if self.num_tokens > 0: + self._populate_fp4_mla_batch_indices_positions() def update_helix_param( self, @@ -656,6 +746,106 @@ def prepare(self) -> None: self.host_request_types_runtime = self.host_request_types[:self. num_seqs] + if self.high_precision_kv_pool is not None: + self._populate_fp4_mla_runtime_metadata(kv_lens) + self._populate_fp4_mla_page_metadata(kv_lens) + if self.num_tokens > 0: + self._populate_fp4_mla_batch_indices_positions() + + def _populate_fp4_mla_runtime_metadata(self, kv_lens: torch.Tensor) -> None: + """Reuse runtime length fields with FP4 MLA append-length semantics.""" + num_seqs = self.num_contexts + self.num_generations + if num_seqs == 0: + return + self.kv_lens_cuda_runtime = self.kv_lens_cuda[:num_seqs] + self.kv_lens_runtime = kv_lens[:num_seqs] + self.prompt_lens_cuda_runtime = self.seq_lens_kv_cuda[:num_seqs] + self.prompt_lens_cpu_runtime = self.seq_lens_kv[:num_seqs] + + def _populate_fp4_mla_page_metadata(self, kv_lens: torch.Tensor) -> None: + """Build the compact page table consumed by FP4 MLA helper kernels.""" + if self.kv_cache_manager is None or self.request_ids is None: + return + assert self._paged_kv_indices is not None + assert self._paged_kv_indptr is not None + assert self.paged_kv_indptr_decode is not None + + num_blocks_tensor = ((kv_lens[:self.num_seqs] + self.page_size - 1) // + self.page_size) + self.num_blocks = [int(item) for item in num_blocks_tensor.tolist()] + self.num_context_blocks = sum(self.num_blocks[:self.num_contexts]) + self.num_generation_blocks = sum(self.num_blocks[self.num_contexts:]) + + block_ids_per_seq = self.kv_cache_manager.get_batch_cache_indices( + self.request_ids) + paged_kv_indices_list = [] + for seq_idx, block_ids in enumerate(block_ids_per_seq): + paged_kv_indices_list.extend(block_ids[:self.num_blocks[seq_idx]]) + paged_kv_indices = torch.tensor(paged_kv_indices_list, + dtype=torch.int32) + if paged_kv_indices.numel() > 0: + self._paged_kv_indices[:paged_kv_indices.numel()].copy_( + paged_kv_indices, non_blocking=True) + + paged_kv_indptr = torch.cumsum(torch.tensor([0] + self.num_blocks, + dtype=torch.int32), + dtype=torch.int32, + dim=0) + self._paged_kv_indptr[:paged_kv_indptr.numel()].copy_(paged_kv_indptr, + non_blocking=True) + + paged_kv_indptr_decode = torch.cumsum( + torch.tensor([0] + self.num_blocks[self.num_contexts:], + dtype=torch.int32), + dtype=torch.int32, + dim=0, + ) + self.paged_kv_indptr_decode[:paged_kv_indptr_decode.numel()].copy_( + paged_kv_indptr_decode, non_blocking=True) + + def _populate_fp4_mla_batch_indices_positions(self) -> None: + """Populate per-token append metadata for FP4 MLA scatter/HP kernels.""" + num_seqs = self.num_contexts + self.num_generations + if num_seqs == 0 or self.num_tokens == 0: + return + assert self.batch_indices is not None + assert self.positions is not None + assert self.kv_lens_cuda_runtime is not None + assert self.prompt_lens_cuda_runtime is not None + + device = self.batch_indices.device + seq_lens = self.seq_lens_kv_cuda[:num_seqs].to(torch.int32) + seq_ids = torch.arange(num_seqs, dtype=torch.int32, device=device) + batch_indices = torch.repeat_interleave(seq_ids, + seq_lens, + output_size=self.num_tokens) + + kv_token_starts = torch.empty((num_seqs, ), + dtype=torch.int32, + device=device) + kv_token_starts[0].zero_() + if num_seqs > 1: + torch.cumsum(seq_lens[:-1], + dim=0, + dtype=torch.int32, + out=kv_token_starts[1:]) + kv_start = torch.repeat_interleave(kv_token_starts, + seq_lens, + output_size=self.num_tokens) + append_lens = self.prompt_lens_cuda_runtime[:num_seqs] + cached_token_lens = self.kv_lens_cuda_runtime[:num_seqs] - append_lens + cached_start = torch.repeat_interleave(cached_token_lens, + seq_lens, + output_size=self.num_tokens) + token_offsets = torch.arange(self.num_tokens, + dtype=torch.int32, + device=device) + positions = token_offsets - kv_start + cached_start + + self.batch_indices[:self.num_tokens].copy_(batch_indices, + non_blocking=True) + self.positions[:self.num_tokens].copy_(positions, non_blocking=True) + def prepare_encoder_only(self) -> None: """Fast path for encoder-only forward (eager + CUDA graph capture).""" extra_attrs = get_model_extra_attrs() @@ -1423,6 +1613,13 @@ def create_output(self, q, *, is_quantize_output: bool, **kwargs) -> List[torch.Tensor]: use_nvfp4_output = False out_dtype = None + if self.is_mla_enable and self.has_fp4_kv_cache: + # The first FP4 MLA FMHA library writes an unquantized decode + # result. Keep attention output in BF16/FP16 and let the + # following projection quantize its input when needed. + is_quantize_output = False + if q.dtype == torch.float8_e4m3fn: + out_dtype = torch.bfloat16 if is_quantize_output: use_nvfp4_output = self.use_nvfp4_output(metadata, attention_mask) out_dtype = self.get_quantize_output_dtype(use_nvfp4_output) @@ -1492,25 +1689,6 @@ def create_fmha_libs(self) -> None: if fmha_cls.is_available(self): self.fmha_libs.append(fmha_cls(self)) - def _update_high_precision_kv_for_fp4_mla( - self, - metadata: TrtllmAttentionMetadata, - latent_cache: Optional[torch.Tensor], - attention_input_type: AttentionInputType = AttentionInputType.mixed, - ) -> None: - """Thin wrapper over the shared HP-pool update helper (see - ``fp4_mla.update_hp_kv_for_fp4_mla`` for the full contract).""" - if attention_input_type == AttentionInputType.context_only: - phase = "context" - elif attention_input_type == AttentionInputType.generation_only: - phase = "generation" - else: - phase = "all" - update_hp_kv_for_fp4_mla(metadata, - latent_cache, - self.get_local_layer_idx(metadata), - phase=phase) - def forward( self, q: torch.Tensor, @@ -1709,6 +1887,9 @@ def forward( # call site where ``output_sf`` is always ``None``. if forward_args.output_sf is not None and forward_args.out_scale_sf is not None: forward_args.out_scale = forward_args.out_scale_sf + elif self.is_mla_enable and self.has_fp4_kv_cache: + forward_args.out_scale = None + forward_args.out_scale_sf = None # Default ``forward_args.kv_scale_*`` to the layer-level mirrors when # the caller didn't populate them. ``modules/attention.py`` only sets @@ -1728,11 +1909,6 @@ def forward( assert metadata.kv_cache_manager is None assert metadata.num_contexts == metadata.num_seqs - if metadata.high_precision_kv_pool is not None and self.is_mla_enable: - self._update_high_precision_kv_for_fp4_mla( - metadata, forward_args.latent_cache, - forward_args.attention_input_type) - if not self.fmha_libs: self.create_fmha_libs() @@ -1997,6 +2173,10 @@ def mla_rope_generation( # kernel reads it. self._ensure_rope_table_size(metadata.max_seq_len) + if self.has_fp4_kv_cache: + self._fp4_mla_rope_generation(fused_q, q_pe, latent_cache, metadata) + return + helix_tensor_params = [ metadata.helix_position_offsets, metadata.helix_is_inactive_rank ] @@ -2041,3 +2221,57 @@ def mla_rope_generation( self.v_head_dim, self.rope_append, ) + + def _fp4_mla_rope_generation( + self, + fused_q: torch.Tensor, + q_pe: torch.Tensor, + latent_cache: torch.Tensor, + metadata: TrtllmAttentionMetadata, + ) -> None: + """Apply the MLA generation RoPE step without appending the FP4 KV cache.""" + assert self.kv_lora_rank is not None + assert self.qk_rope_head_dim is not None + + if metadata.positions is None: + raise RuntimeError( + "FP4 MLA generation requires per-token positions.") + if q_pe.shape[-1] != self.qk_rope_head_dim: + raise RuntimeError( + f"FP4 MLA q_pe last dimension must be {self.qk_rope_head_dim}, " + f"got {q_pe.shape[-1]}.") + if latent_cache.shape[-1] != self.kv_lora_rank + self.qk_rope_head_dim: + raise RuntimeError( + "FP4 MLA latent_cache last dimension must be " + f"{self.kv_lora_rank + self.qk_rope_head_dim}, got " + f"{latent_cache.shape[-1]}.") + + num_tokens = q_pe.shape[0] + token_offset = getattr(metadata, "num_ctx_tokens", 0) + positions = metadata.positions[token_offset:token_offset + + num_tokens].to(torch.long) + if num_tokens > 0 and torch.cuda.is_current_stream_capturing(): + max_position = metadata.max_seq_len + elif num_tokens > 0: + max_position = int(positions.max().item()) + 1 + else: + max_position = 0 + self._ensure_rope_table_size(max(max_position, metadata.max_seq_len)) + + q_roped = apply_fp4_mla_rope( + q_pe, + positions, + self.rotary_cos_sin, + self.rope_params.max_positions, + self.qk_rope_head_dim, + ) + k_pe = latent_cache[..., self.kv_lora_rank:] + k_roped = apply_fp4_mla_rope( + k_pe.unsqueeze(1), + positions, + self.rotary_cos_sin, + self.rope_params.max_positions, + self.qk_rope_head_dim, + ).squeeze(1) + fused_q[..., self.kv_lora_rank:].copy_(q_roped.to(fused_q.dtype)) + k_pe.copy_(k_roped.to(k_pe.dtype)) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 83689a37b7f9..ad06c7e17e5b 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -8,7 +8,8 @@ import weakref from abc import ABC, abstractmethod from contextlib import contextmanager -from typing import Any, Callable, Dict, List, Optional, Tuple, Union +from typing import (TYPE_CHECKING, Any, Callable, Dict, List, Optional, Tuple, + Union) import torch import torch._dynamo.config @@ -81,6 +82,9 @@ from .sampler import SampleStateTensors from .scheduler import ScheduledRequests +if TYPE_CHECKING: + from tensorrt_llm.llmapi.llm_args import DecodingBaseConfig + class ModelEngine(ABC): @@ -2099,35 +2103,6 @@ def _preprocess_inputs(self, inputs: Dict[str, Any]): previous_kv_lens_offsets_cuda[:num_gen_requests] ) inputs['attn_metadata'].on_update_kv_lens() - elif getattr(inputs['attn_metadata'], 'kv_lens_cuda_runtime', - None) is not None: - # FlashInfer NVFP4 MLA does not expose kv_lens_cuda; it keeps - # the per-sequence KV length in kv_lens_cuda_runtime and - # precomputes the FP4 KV-cache write positions / batch indices - # from it during prepare(). The overlap scheduler builds the - # generation metadata from the all-draft-accepted estimate - # (num_cached_tokens_per_seq = past_seen + runtime_draft_len + - # 1), so without the same previous_kv_lens_offsets correction - # the FP4 KV scatter and the BF16 HP-pool overlay would write - # at over-estimated positions whenever the previous MTP step - # rejected a draft token, corrupting the KV cache. Apply the - # correction, then rebuild positions/batch indices from the - # corrected lengths so the captured graph replays them. - md = inputs['attn_metadata'] - if num_ctx_requests >= num_chunked_ctx_requests and num_chunked_ctx_requests > 0: - md.kv_lens_cuda_runtime[ - num_ctx_requests - - num_chunked_ctx_requests:num_ctx_requests] += ( - self. - previous_kv_lens_offsets_cuda[: - num_chunked_ctx_requests] - ) - else: - md.kv_lens_cuda_runtime[num_ctx_requests:num_seqs] += ( - self. - previous_kv_lens_offsets_cuda[:num_gen_requests]) - md.on_update_kv_lens() - md._populate_fp4_mla_batch_indices_positions() if self.guided_decoder is not None: self.guided_decoder.token_event.record() @@ -2175,27 +2150,6 @@ def _postprocess_inputs(self, inputs: Dict[str, Any]): self. previous_kv_lens_offsets_cuda[:num_gen_requests] ) - elif getattr(inputs['attn_metadata'], 'kv_lens_cuda_runtime', - None) is not None: - # Undo the FlashInfer NVFP4 MLA kv_lens correction applied in - # _preprocess_inputs so the captured graph re-applies it from - # the original (over-estimated) lengths on the post-capture - # replay. positions/batch indices are rebuilt by the captured - # _populate_fp4_mla_batch_indices_positions on every replay, so - # they do not need to be restored here. - md = inputs['attn_metadata'] - if num_ctx_requests >= num_chunked_ctx_requests and num_chunked_ctx_requests > 0: - md.kv_lens_cuda_runtime[ - num_ctx_requests - - num_chunked_ctx_requests:num_ctx_requests] -= ( - self. - previous_kv_lens_offsets_cuda[: - num_chunked_ctx_requests] - ) - else: - md.kv_lens_cuda_runtime[num_ctx_requests:num_seqs] -= ( - self. - previous_kv_lens_offsets_cuda[:num_gen_requests]) def _get_all_rank_num_tokens(self, attn_metadata: AttentionMetadata): if self.enable_attention_dp: @@ -4291,7 +4245,6 @@ def _prepare_star_attention_inputs( anchor_len) all_cache_indices = all_cache_indices[ num_kvblocks_per_ctx_block:] - cache_indices = all_cache_indices[:num_kv_blocks] last_query_pos_id = request.ctx_position_blocks[-1][-1] position_ids.append(last_query_pos_id + request.gen_iters + 1) block_ids_per_seq.extend([all_cache_indices]) diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index 7d3fc729cec4..d0497a47be96 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -46,7 +46,6 @@ FLASH_MLA_TOKENS_PER_BLOCK = 64 FP4_MLA_TOKENS_PER_BLOCK = 128 -FLASHINFER_FP4_MLA_ATTENTION_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION" class _ExecutorMemoryMonitor: @@ -201,20 +200,8 @@ def _has_fp4_kv_cache(model_config, kv_cache_config) -> bool: or kv_cache_quant_algo in fp4_quant_values) -def _enable_fp4_mla_attention() -> bool: - return os.getenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "0").lower() in ( - "1", - "true", - "yes", - "on", - ) - - -def _is_fp4_mla_flashinfer_attention_requested(llm_args: TorchLlmArgs, - kv_cache_config) -> bool: - return (llm_args.attn_backend == "FLASHINFER" - and _has_fp4_kv_cache(llm_args, kv_cache_config) - and _enable_fp4_mla_attention()) +def _enable_fp4_mla_attention(config=None) -> bool: + return getattr(config, "attn_backend", None) == "TRTLLM" def _select_mla_tokens_per_block(config, model_config, kv_cache_config, @@ -223,7 +210,7 @@ def _select_mla_tokens_per_block(config, model_config, kv_cache_config, return tokens_per_block if (_has_fp4_kv_cache(model_config, kv_cache_config) - and _enable_fp4_mla_attention()): + and _enable_fp4_mla_attention(model_config)): tokens_per_block = FP4_MLA_TOKENS_PER_BLOCK logger.info( f"Change tokens_per_block to: {tokens_per_block} for using FP4 MLA attention" @@ -479,23 +466,14 @@ def create_py_executor( "token-match acceptance. This can change the sampled output " "distribution; set use_rejection_sampling=True to use exact " "one-model speculative sampling.") - # Regular FlashInfer decode expects one query token per sequence. The - # FP4 MLA no-dequant path has its own linear-MTP handling, so allow that - # explicit configuration through. - fp4_mla_flashinfer = _is_fp4_mla_flashinfer_attention_requested( - llm_args, kv_cache_config) - if llm_args.attn_backend == "FLASHINFER" and not fp4_mla_flashinfer: + # Regular FlashInfer decode expects one query token per sequence. + if llm_args.attn_backend == "FLASHINFER": raise ValueError( f"FLASHINFER attention backend is not supported with one-engine speculative " f"decoding mode '{spec_config.spec_dec_mode.name}'. The FLASHINFER backend's " f"decode path expects exactly 1 token per sequence, but one-engine speculative " f"decoding requires multiple tokens per sequence. Please use 'TRTLLM' attention " f"backend instead by setting attn_backend='TRTLLM'.") - if fp4_mla_flashinfer: - # FlashInfer metadata prepares page tables against one KV manager. - # Keep one-model MTP draft layers in the main manager so global - # draft layer ids are present in layer_offsets. - spec_config._allow_separate_draft_kv_cache = False if mm_encoder_only: llm_args.mm_encoder_only = True diff --git a/tensorrt_llm/_torch/pyexecutor/resource_manager.py b/tensorrt_llm/_torch/pyexecutor/resource_manager.py index 9a46a9164956..794811caa3f6 100644 --- a/tensorrt_llm/_torch/pyexecutor/resource_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/resource_manager.py @@ -8,8 +8,8 @@ from abc import ABC, abstractmethod from collections import OrderedDict, defaultdict, deque from dataclasses import dataclass -from typing import (TYPE_CHECKING, Dict, Iterable, List, Optional, Set, Tuple, - Union) +from typing import (TYPE_CHECKING, Dict, Iterable, List, Optional, Sequence, + Set, Tuple, Union) import torch from mpi4py import MPI @@ -17,15 +17,39 @@ import tensorrt_llm import tensorrt_llm.bindings from tensorrt_llm._torch.distributed.communicator import Distributed, ReduceOp -from tensorrt_llm._utils import (get_size_in_bytes, mpi_comm, mpi_disabled, +from tensorrt_llm._utils import (TensorWrapper, convert_to_torch_tensor, + get_size_in_bytes, mpi_comm, mpi_disabled, prefer_pinned, torch_comm) from tensorrt_llm.bindings.internal.batch_manager import ( - LinearAttentionMetadata, LinearCacheType) + KvCacheStats, LinearAttentionMetadata, LinearCacheType) +from tensorrt_llm.bindings.internal.batch_manager.kv_cache_manager_v2_utils import ( + IndexMapper, copy_batch_block_offsets_to_device) from tensorrt_llm.bindings.internal.runtime import TaskLayerModuleConfig from tensorrt_llm.llmapi.llm_args import KvCacheConfig, PeftCacheConfig from tensorrt_llm.lora_helper import LoraConfig from tensorrt_llm.lora_manager import LoraManager, LoraModelConfig from tensorrt_llm.runtime import ModelConfig as ModelConfigPython +# yapf: disable +from tensorrt_llm.runtime.kv_cache_manager_v2 import (DEFAULT_BEAM_INDEX, + AttentionLayerConfig, + BufferConfig, + CacheTierConfig, + DiskCacheTierConfig, + GpuCacheTierConfig, + HostCacheTierConfig) +from tensorrt_llm.runtime.kv_cache_manager_v2 import \ + KVCacheManager as KVCacheManagerPy +from tensorrt_llm.runtime.kv_cache_manager_v2 import \ + KVCacheManagerConfig as KVCacheManagerConfigPy +from tensorrt_llm.runtime.kv_cache_manager_v2 import (LayerId, ReuseScope, + TokenIdExt, _KVCache) +from tensorrt_llm.runtime.kv_cache_manager_v2._common import (BAD_PAGE_INDEX, + GPU_LEVEL) +from tensorrt_llm.runtime.kv_cache_manager_v2._config import DataRole +from tensorrt_llm.runtime.kv_cache_manager_v2._utils import (exact_div, + typed_range) + +# yapf: enable # isort: off # isort: on @@ -58,16 +82,15 @@ BlocksPerWindow = Dict[int, Tuple[ int, int]] # window_size -> (blocks_in_primary_pool, blocks_in_secondary_pool) -FLASHINFER_FP4_MLA_ATTENTION_ENV = "TRTLLM_FLASHINFER_FP4_MLA_ATTENTION" -def _flashinfer_fp4_mla_attention_enabled() -> bool: - return os.getenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "0").lower() in ( - "1", - "true", - "yes", - "on", - ) +class Role: + KEY = DataRole("key") + VALUE = DataRole("value") + KEY_BLOCK_SCALE = DataRole("key_block_scale") + VALUE_BLOCK_SCALE = DataRole("value_block_scale") + INDEX_KEY = DataRole("index_key") + ALL = DataRole("all") @dataclass @@ -1161,8 +1184,7 @@ def get_cache_bytes_per_token(self): def _enable_mla_v_scale_pool(self) -> bool: return (self.dtype == DataType.NVFP4 - and self.kv_cache_type == CacheTypeCpp.SELFKONLY - and _flashinfer_fp4_mla_attention_enabled()) + and self.kv_cache_type == CacheTypeCpp.SELFKONLY) def _get_mla_v_scale_bytes_per_token(self, num_layers: int) -> int: if not self._enable_mla_v_scale_pool(): @@ -2238,6 +2260,7 @@ def reset_reuse_state(self): """Reset the reuse state of the KV cache manager.""" self.impl.reset_reuse_state() + class KVCacheManagerV2(BaseResourceManager): def __init__( @@ -3155,6 +3178,11 @@ def _augment_tokens_for_block_reuse( or req.multimodal_lengths is None): return tokens[chunk_start:chunk_end] if is_sliced else tokens + from .kv_cache_manager_v2 import ( + _augment_tokens_with_contiguous_mm_metadata, + _augment_tokens_with_mm_run_metadata, + _resolve_multimodal_run_metadata) + result: list[TokenIdExt] = list(tokens[chunk_start:chunk_end]) run_metadata = _resolve_multimodal_run_metadata(req) if run_metadata is not None: diff --git a/tests/unittest/_torch/attention/test_flashinfer_attention.py b/tests/unittest/_torch/attention/test_flashinfer_attention.py index d91a45cd3bb0..08fe40dc190c 100644 --- a/tests/unittest/_torch/attention/test_flashinfer_attention.py +++ b/tests/unittest/_torch/attention/test_flashinfer_attention.py @@ -672,12 +672,12 @@ class TestFlashInferFp4KvGuards(unittest.TestCase): """Guards that FlashInfer + NVFP4 KV cache rejects unsupported configs at init time with a clear NotImplementedError, rather than failing mid-forward. - Phase 1 of NVFP4 KV cache support on the FlashInfer backend only covers - MLA; non-MLA FP4 should error out early pointing users to attn_backend + NVFP4 KV cache is supported only by the TRTLLM attention backend, so + FlashInfer should error out early and point users to attn_backend ``TRTLLM`` or a BF16/FP8 KV cache. """ - def test_non_mla_fp4_kv_raises_not_implemented(self): + def test_fp4_kv_raises_not_implemented(self): from tensorrt_llm.models.modeling_utils import QuantConfig from tensorrt_llm.quantization.mode import QuantAlgo @@ -694,5 +694,4 @@ def test_non_mla_fp4_kv_raises_not_implemented(self): # self-serve without reading the backend source. msg = str(ctx.exception) self.assertIn("NVFP4 KV cache", msg) - self.assertIn("MLA", msg) self.assertIn("TRTLLM", msg) diff --git a/tests/unittest/_torch/attention/test_fp4_mla.py b/tests/unittest/_torch/attention/test_fp4_mla.py index 913f305f0f2a..e6be4e4dcaee 100644 --- a/tests/unittest/_torch/attention/test_fp4_mla.py +++ b/tests/unittest/_torch/attention/test_fp4_mla.py @@ -1,12 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Roundtrip tests for the FP4 MLA KV-cache kernels. - -Exercises ``scatter_fp4_mla_kv_cache`` and ``get_fp4_mla_decode_cache`` -(plus ``update_hp_kv_for_fp4_mla``) as a pair on a tiny V1 ``KVCacheManager``. -The goal is to catch stride / page-id / SF-layout bugs without standing up a -real model or FlashInfer wrapper. -""" +"""TRTLLM FP4 MLA helper tests.""" import os from types import SimpleNamespace @@ -15,36 +9,65 @@ import torch import tensorrt_llm +import tensorrt_llm._torch.attention_backend.fmha.fp4_mla as fp4_mla_fmha +from tensorrt_llm._torch.attention_backend.fmha.fp4_mla import Fp4MlaFmha from tensorrt_llm._torch.attention_backend.fp4_mla import ( - FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, - FLASHINFER_FP4_MLA_ATTENTION_ENV, FP4_BLOCK_SIZE, + FP4_MLA_ATTENTION_BACKEND_ENV, FP4_MLA_KV_GLOBAL_SCALE, FP4_MLA_P_GLOBAL_SCALE, FP4_MLA_Q_RESIDUAL_DIM, FP4_MLA_TOKENS_PER_BLOCK, HP_BLOCK_SIZE, - get_fp4_mla_decode_cache, + _cutile_backend_available, + _get_cutile_v_packed_cache, + _maybe_update_cutile_v_packed_cache, get_fp4_mla_v_scale_pool_shape, get_fp4_mla_v_scale_pool_size, get_fp4_mla_v_scale_pool_view, - is_flashinfer_fp4_mla_attention_enabled, repair_fp4_mla_hp_kv_for_mtp_rejection, run_fp4_mla_attention_decode, scatter_fp4_mla_kv_cache, update_hp_kv_for_fp4_mla, - _get_cutile_v_packed_cache, - _maybe_update_cutile_v_packed_cache, ) from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings.executor import KvCacheConfig from tensorrt_llm.mapping import Mapping +from tensorrt_llm.quantization.mode import QuantMode _DataType = tensorrt_llm.bindings.DataType _CacheType = tensorrt_llm.bindings.internal.batch_manager.CacheType _TEST_GLOBAL_SCALE = FP4_MLA_KV_GLOBAL_SCALE +def _fp4_mla_attn_for_availability(head_dim: int): + return SimpleNamespace( + is_mla_enable=True, + quant_mode=int(QuantMode(0).set_fp4_kv_cache()), + attention_chunk_size=None, + predicted_tokens_per_seq=1, + kv_lora_rank=512, + qk_nope_head_dim=128, + qk_rope_head_dim=FP4_MLA_Q_RESIDUAL_DIM, + v_head_dim=128, + head_dim=head_dim, + ) + + +def test_fp4_mla_fmha_available_for_context_and_generation_head_dims(monkeypatch): + monkeypatch.setattr(fp4_mla_fmha, "get_sm_version", lambda: 100) + monkeypatch.setattr(fp4_mla_fmha, "is_sm_100f", lambda sm: True) + monkeypatch.setattr( + fp4_mla_fmha.torch, + "ops", + SimpleNamespace(trtllm=SimpleNamespace(fp4_quantize_with_residual=object())), + ) + + assert Fp4MlaFmha.is_available(_fp4_mla_attn_for_availability(128 + 64)) + assert Fp4MlaFmha.is_available(_fp4_mla_attn_for_availability(512 + 64)) + assert not Fp4MlaFmha.is_available(_fp4_mla_attn_for_availability(128)) + + def _swizzled_sf_offset(row_idx: int, col_idx: int, sf_per_token: int) -> int: padded_cols = ((sf_per_token + 3) // 4) * 4 return ( @@ -60,6 +83,17 @@ def _is_pre_blackwell() -> bool: return not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] < 10 +def _is_cutile_unavailable() -> bool: + return _is_pre_blackwell() or not _cutile_backend_available() + + +def _reset_triton_allocator() -> None: + import triton + import triton.runtime._allocation as triton_allocation + + triton.set_allocator(triton_allocation._NULL_ALLOCATOR) + + def _dequant_fp4_swizzled( fp4_tensor: torch.Tensor, sf_tensor: torch.Tensor, @@ -83,11 +117,11 @@ def _dequant_fp4_swizzled( for row_idx in range(fp4_bytes.shape[0]): for sf_col in range(sf_per_token): - start = sf_col * 16 + start = sf_col * FP4_BLOCK_SIZE packed = fp4_bytes[row_idx, start // 2 : start // 2 + 8] low = packed & 0x0F high = (packed >> 4) & 0x0F - vals = torch.empty(16, dtype=torch.float32, device=fp4_tensor.device) + vals = torch.empty(FP4_BLOCK_SIZE, dtype=torch.float32, device=fp4_tensor.device) low_sign = torch.where( (low & 0x08) != 0, -torch.ones_like(low, dtype=torch.float32), @@ -101,14 +135,18 @@ def _dequant_fp4_swizzled( vals[0::2] = fp4_values[(low & 0x07).long()] * low_sign vals[1::2] = fp4_values[(high & 0x07).long()] * high_sign sf_offset = _swizzled_sf_offset(row_idx, sf_col, sf_per_token) - out[row_idx, start : start + 16] = vals * sf_flat[sf_offset].float() / global_scale + out[row_idx, start : start + FP4_BLOCK_SIZE] = ( + vals * sf_flat[sf_offset].float() / global_scale + ) return out def _duplicate_tail_groups(tensor: torch.Tensor, residual_dim: int) -> torch.Tensor: prefix = tensor[..., :-residual_dim] - tail = tensor[..., -residual_dim:].reshape(*tensor.shape[:-1], residual_dim // 16, 16) + tail = tensor[..., -residual_dim:].reshape( + *tensor.shape[:-1], residual_dim // FP4_BLOCK_SIZE, FP4_BLOCK_SIZE + ) duplicated_tail = tail.repeat_interleave(2, dim=-2).reshape( *tensor.shape[:-1], residual_dim * 2, @@ -116,44 +154,6 @@ def _duplicate_tail_groups(tensor: torch.Tensor, residual_dim: int) -> torch.Ten return torch.cat((prefix, duplicated_tail), dim=-1) -def test_flashinfer_fp4_mla_attention_env(monkeypatch): - monkeypatch.delenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, raising=False) - assert not is_flashinfer_fp4_mla_attention_enabled() - - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "on") - assert is_flashinfer_fp4_mla_attention_enabled() - - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "0") - assert not is_flashinfer_fp4_mla_attention_enabled() - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_fp4_mla_v_scale_pool_view_shape(): - device = torch.device("cuda") - num_layers = 2 - num_pages = 3 - kv_lora_rank = 512 - qk_rope_head_dim = 64 - head_dim = kv_lora_rank + qk_rope_head_dim - page_size = FP4_MLA_TOKENS_PER_BLOCK - - allocated_page_elems = get_fp4_mla_v_scale_pool_size(head_dim, page_size) - pool = torch.empty( - (num_layers, num_pages, allocated_page_elems), - dtype=torch.float8_e4m3fn, - device=device, - ) - metadata = SimpleNamespace(page_size=page_size, fp4_mla_v_scale_pool=pool) - - view = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=kv_lora_rank) - - assert tuple(view.shape) == get_fp4_mla_v_scale_pool_shape( - num_layers, num_pages, kv_lora_rank, page_size - ) - assert view.data_ptr() == pool.data_ptr() - assert view.numel() == num_layers * num_pages * kv_lora_rank * (page_size // 16) - - def _build_metadata(kv_cache_manager, *, num_tokens, page_size, num_layers): """Build a minimal metadata namespace that satisfies the kernels' field expectations. Single sequence, single layer slice, no draft tokens.""" @@ -179,7 +179,6 @@ def _build_metadata(kv_cache_manager, *, num_tokens, page_size, num_layers): ) seq_slots = torch.zeros(1, dtype=torch.int32, device=device) - seq_slots_cpu = torch.zeros(1, dtype=torch.int32, device="cpu") kv_lens = torch.tensor([num_tokens], dtype=torch.int32, device=device) prompt_lens_cuda = torch.tensor([num_tokens], dtype=torch.int32, device=device) prompt_lens_cpu = torch.tensor([num_tokens], dtype=torch.int32) @@ -201,7 +200,6 @@ def _build_metadata(kv_cache_manager, *, num_tokens, page_size, num_layers): fp4_mla_v_scale_pool=kv_cache_manager.get_mla_v_scale_pool(), hp_pool_owners={}, seq_slots=seq_slots, - seq_slots_cpu=seq_slots_cpu, kv_lens_cuda_runtime=kv_lens, prompt_lens_cuda_runtime=prompt_lens_cuda, prompt_lens_cpu_runtime=prompt_lens_cpu, @@ -249,8 +247,8 @@ def _build_multi_seq_metadata(kv_cache_manager, *, seq_lens, page_size, num_laye dtype=torch.bfloat16, device=device, ) + seq_slots = torch.arange(num_seqs, dtype=torch.int32, device=device) - seq_slots_cpu = torch.arange(num_seqs, dtype=torch.int32, device="cpu") kv_lens = torch.tensor(seq_lens, dtype=torch.int32, device=device) prompt_lens_cuda = torch.tensor(seq_lens, dtype=torch.int32, device=device) prompt_lens_cpu = torch.tensor(seq_lens, dtype=torch.int32) @@ -273,7 +271,6 @@ def _build_multi_seq_metadata(kv_cache_manager, *, seq_lens, page_size, num_laye fp4_mla_v_scale_pool=kv_cache_manager.get_mla_v_scale_pool(), hp_pool_owners={}, seq_slots=seq_slots, - seq_slots_cpu=seq_slots_cpu, kv_lens_cuda_runtime=kv_lens, prompt_lens_cuda_runtime=prompt_lens_cuda, prompt_lens_cpu_runtime=prompt_lens_cpu, @@ -284,6 +281,34 @@ def _build_multi_seq_metadata(kv_cache_manager, *, seq_lens, page_size, num_laye ) +def _materialize_reference_cache(metadata, layer_idx: int, head_dim: int) -> torch.Tensor: + kv_cache = metadata.kv_cache_manager.get_buffers(layer_idx).view(torch.uint8) + sf_cache = metadata.kv_cache_manager.get_block_scale_buffers(layer_idx).view( + torch.float8_e4m3fn + ) + pages = [] + src_page_ids = metadata.paged_kv_indices[ + metadata.num_context_blocks : metadata.num_context_blocks + metadata.num_generation_blocks + ] + for page_id in src_page_ids.tolist(): + fp4_page = kv_cache[page_id, 0, :, 0, :] + sf_page = sf_cache[page_id] + pages.append( + _dequant_fp4_swizzled( + fp4_page, + sf_page, + logical_dim=head_dim, + sf_per_token=head_dim // FP4_BLOCK_SIZE, + global_scale=_TEST_GLOBAL_SCALE, + ) + ) + if not pages: + return torch.empty( + (0, metadata.page_size, head_dim), dtype=torch.float32, device=kv_cache.device + ) + return torch.stack(pages, dim=0) + + def _build_fp4_mla_attention_decode_case(*, seq_lens, num_heads, seed, query_len_per_seq=1): torch.manual_seed(seed) device = torch.device("cuda") @@ -372,19 +397,7 @@ def _fp4_mla_attention_decode_reference( qk_rope_head_dim, ): head_dim = kv_lora_rank + qk_rope_head_dim - high_precision_kv_pool = metadata.high_precision_kv_pool - metadata.high_precision_kv_pool = None - try: - dequant_cache = get_fp4_mla_decode_cache( - metadata, - layer_idx=0, - local_layer=0, - head_dim=head_dim, - dtype=torch.bfloat16, - ) - finally: - metadata.high_precision_kv_pool = high_precision_kv_pool - + dequant_cache = _materialize_reference_cache(metadata, 0, head_dim) num_heads = q_nope.shape[1] global_scale = metadata._fp4_mla_global_scale q_full = torch.cat((q_nope, q_pe), dim=-1).reshape(-1, head_dim) @@ -399,7 +412,7 @@ def _fp4_mla_attention_decode_reference( q_fp4, q_sf.view(torch.float8_e4m3fn), logical_dim=q_logical_dim, - sf_per_token=q_logical_dim // 16, + sf_per_token=q_logical_dim // FP4_BLOCK_SIZE, global_scale=_TEST_GLOBAL_SCALE, ) @@ -409,7 +422,7 @@ def _fp4_mla_attention_decode_reference( metadata._fp4_mla_attention_p_buf, metadata._fp4_mla_attention_p_sf_buf, logical_dim=metadata.page_size, - sf_per_token=metadata.page_size // 16, + sf_per_token=metadata.page_size // FP4_BLOCK_SIZE, global_scale=FP4_MLA_P_GLOBAL_SCALE, ) @@ -449,25 +462,10 @@ def _fp4_mla_attention_decode_reference( exact_probs.append(probs) quantized_probs.append(p) - outputs.append(torch.matmul(probs, cache[:, :kv_lora_rank].float())) + outputs.append(torch.matmul(p, cache[:, :kv_lora_rank].float())) return torch.stack(outputs, dim=0), exact_probs, quantized_probs -def _cuda_event_benchmark(fn, *, warmup_iters=10, iters=100): - for _ in range(warmup_iters): - fn() - torch.cuda.synchronize() - - start = torch.cuda.Event(enable_timing=True) - end = torch.cuda.Event(enable_timing=True) - start.record() - for _ in range(iters): - fn() - end.record() - torch.cuda.synchronize() - return start.elapsed_time(end) / iters - - def _assert_fp4_mla_attention_decode_accuracy( monkeypatch, *, @@ -475,12 +473,11 @@ def _assert_fp4_mla_attention_decode_accuracy( num_heads: int, seq_lens: list[int], seed: int, - check_probs: bool, + check_probs: bool = False, query_len_per_seq: int = 1, ) -> None: - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, backend) - + _reset_triton_allocator() + monkeypatch.setenv(FP4_MLA_ATTENTION_BACKEND_ENV, backend) ( kv_cache_manager, metadata, @@ -532,301 +529,222 @@ def _assert_fp4_mla_attention_decode_accuracy( torch.testing.assert_close( output.float(), ref_output, - atol=1e-1, - rtol=1e-1, + atol=1.5e-1, + rtol=1.5e-1, msg=f"{backend} FP4 MLA attention decode output diverged from reference", ) finally: + torch.cuda.synchronize() + _reset_triton_allocator() kv_cache_manager.shutdown() + torch.cuda.synchronize() + torch.cuda.empty_cache() -def _ceil_div(lhs: int, rhs: int) -> int: - return (lhs + rhs - 1) // rhs - +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_fp4_mla_v_scale_pool_view_shape(): + device = torch.device("cuda") + num_layers = 2 + num_pages = 3 + kv_lora_rank = 512 + page_size = FP4_MLA_TOKENS_PER_BLOCK -def _estimate_fp4_mla_attention_decode_mbu_bytes( - *, - batch_size: int, - seq_len: int, - num_heads: int, - kv_lora_rank: int, - qk_rope_head_dim: int, - page_size: int, -) -> int: - """Estimate logical global-memory traffic for the FP4 MLA decode benchmark.""" - fp4_block_size = 16 - head_block_size = 128 - bf16_bytes = 2 - fp32_bytes = 4 + allocated_page_elems = get_fp4_mla_v_scale_pool_size(kv_lora_rank, page_size) + pool = torch.empty( + (num_layers, num_pages, allocated_page_elems), + dtype=torch.float8_e4m3fn, + device=device, + ) + metadata = SimpleNamespace(page_size=page_size, fp4_mla_v_scale_pool=pool) - q_input_dim = kv_lora_rank + qk_rope_head_dim - qk_dim = q_input_dim + FP4_MLA_Q_RESIDUAL_DIM - pages_per_seq = _ceil_div(seq_len, page_size) - padded_seq_len = pages_per_seq * page_size - head_blocks = _ceil_div(num_heads, head_block_size) - q_rows = batch_size * num_heads + view = get_fp4_mla_v_scale_pool_view(metadata, v_head_dim=kv_lora_rank) - qk_fp4_bytes_per_token = qk_dim // 2 + _ceil_div(qk_dim, fp4_block_size) - v_fp4_bytes_per_token = kv_lora_rank // 2 + _ceil_div(kv_lora_rank, fp4_block_size) - p_bytes_per_seq = padded_seq_len // 2 + _ceil_div(padded_seq_len, fp4_block_size) + assert tuple(view.shape) == get_fp4_mla_v_scale_pool_shape( + num_layers, num_pages, kv_lora_rank, page_size + ) + assert view.data_ptr() == pool.data_ptr() + assert view.numel() == num_layers * num_pages * kv_lora_rank * (page_size // FP4_BLOCK_SIZE) - q_setup_bytes = q_rows * q_input_dim * bf16_bytes * 3 + q_rows * qk_fp4_bytes_per_token - qk_cache_bytes = 2 * batch_size * head_blocks * padded_seq_len * qk_fp4_bytes_per_token - qk_q_bytes = 2 * batch_size * num_heads * pages_per_seq * qk_fp4_bytes_per_token - stats_bytes = batch_size * num_heads * 2 * fp32_bytes * (1 + pages_per_seq) - p_prob_bytes = batch_size * num_heads * padded_seq_len * fp32_bytes * 2 - p_quant_bytes = 2 * batch_size * num_heads * p_bytes_per_seq - pv_cache_bytes = batch_size * head_blocks * padded_seq_len * v_fp4_bytes_per_token - output_bytes = q_rows * kv_lora_rank * bf16_bytes - return ( - q_setup_bytes - + qk_cache_bytes - + qk_q_bytes - + stats_bytes - + p_prob_bytes - + p_quant_bytes - + pv_cache_bytes - + output_bytes +def _materialize_single_seq_cache_with_hp_tail( + metadata, *, layer_idx: int, local_layer: int, head_dim: int, num_tokens: int +) -> torch.Tensor: + flat = ( + _materialize_reference_cache(metadata, layer_idx, head_dim) + .reshape(-1, head_dim)[:num_tokens] + .to(torch.bfloat16) ) + tail = num_tokens % HP_BLOCK_SIZE + if tail == 0: + return flat + + seq_slot = int(metadata.seq_slots[0].item()) + pool_head_dim = metadata.high_precision_kv_pool.shape[-1] // HP_BLOCK_SIZE + hp_view = metadata.high_precision_kv_pool[seq_slot, local_layer, 0, :].view( + HP_BLOCK_SIZE, pool_head_dim + ) + flat[-tail:] = hp_view[:tail, :head_dim] + return flat -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") @pytest.mark.parametrize( - ("num_tokens", "page_size"), - [(20, 16), (32, 16), (17, 16), (129, 128), (144, 128)], - ids=[ - "tail4_page16", - "aligned32_page16", - "tail1_page16", - "tail1_page128", - "aligned144_page128", - ], + "num_tokens", + [32, 129, 144], + ids=["aligned32", "tail1", "aligned144"], ) -def test_fp4_mla_scatter_gather_roundtrip(num_tokens: int, page_size: int, monkeypatch): - """Write BF16 latent through scatter + HP update, then read via the - dequant-gather + HP-overlay path and verify the two halves of the output: - - * Positions that fall in the FP4 region (before the last ``kv_len % 16`` - tokens) must match the input up to NVFP4 quant error. - * Positions covered by the HP overlay (the last ``kv_len % 16`` tokens) - must match the input exactly (BF16 roundtrip). - """ - monkeypatch.delenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, raising=False) +def test_fp4_mla_scatter_gather_roundtrip(num_tokens: int): torch.manual_seed(0) device = torch.device("cuda") - - # MLA shapes (DeepSeek-V3-Lite style, scaled down). kv_lora_rank = 512 qk_rope_head_dim = 64 - head_dim = kv_lora_rank + qk_rope_head_dim # 576, divisible by 16. + head_dim = kv_lora_rank + qk_rope_head_dim + page_size = FP4_MLA_TOKENS_PER_BLOCK num_layers = 1 - max_seq_len = max(64, ((num_tokens + page_size - 1) // page_size) * page_size) - max_batch_size = 1 + max_seq_len = ((num_tokens + page_size - 1) // page_size) * page_size mapping = Mapping(world_size=1, tp_size=1, rank=0) - # max_tokens must cover at least ceil(max_seq_len / page_size) pages. - kv_cache_config = KvCacheConfig(max_tokens=max_seq_len, enable_block_reuse=False) kv_cache_manager = KVCacheManager( - kv_cache_config, + KvCacheConfig(max_tokens=max_seq_len, enable_block_reuse=False), _CacheType.SELFKONLY, num_layers=num_layers, num_kv_heads=1, head_dim=head_dim, tokens_per_block=page_size, max_seq_len=max_seq_len, - max_batch_size=max_batch_size, + max_batch_size=1, mapping=mapping, dtype=_DataType.NVFP4, ) try: kv_cache_manager.add_dummy_requests([0], [num_tokens]) - - # Zero the underlying data + scale pools so stale bytes can't mask bugs. - data_buf = kv_cache_manager.get_buffers(0).view(torch.uint8) - data_buf.zero_() - sf_buf = kv_cache_manager.get_block_scale_buffers(0) - assert sf_buf is not None, "V1 NVFP4 manager must expose block scales" - sf_buf.zero_() - + kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() + kv_cache_manager.get_block_scale_buffers(0).zero_() metadata = _build_metadata( kv_cache_manager, num_tokens=num_tokens, page_size=page_size, num_layers=num_layers ) - - # Stay inside the configured global-scale range so FP4 quantization - # does not saturate; narrower latents keep the FP4 tolerance reasonable. latent = ( - torch.randn(num_tokens, head_dim, dtype=torch.bfloat16, device=device) * 1.5 - ).clamp_(-5.0, 5.0) + torch.randn(num_tokens, head_dim, dtype=torch.bfloat16, device=device) * 0.25 + ).clamp_(-1.0, 1.0) - # --- Write path ------------------------------------------------- - scatter_fp4_mla_kv_cache(metadata, latent, layer_idx=0, token_offset=0) + scatter_fp4_mla_kv_cache( + metadata, + latent, + layer_idx=0, + token_offset=0, + phase="context", + local_layer=0, + v_head_dim=kv_lora_rank, + ) update_hp_kv_for_fp4_mla(metadata, latent, local_layer=0, phase="context") + torch.cuda.synchronize() - # Switch the metadata into "decode" shape for the read path. The - # decode kernels gather the entire 0..num_tokens range (as if every - # token were in the KV history for an upcoming decode step). - metadata.num_contexts = 0 - metadata.num_seqs = 1 - - # --- Read path -------------------------------------------------- - combined = get_fp4_mla_decode_cache( + recovered = _materialize_single_seq_cache_with_hp_tail( metadata, layer_idx=0, local_layer=0, head_dim=head_dim, - dtype=torch.bfloat16, + num_tokens=num_tokens, ) - # combined shape: [num_blocks, page_size, head_dim]. - flat = combined.reshape(-1, head_dim)[:num_tokens] - - # --- Assertions ------------------------------------------------- tail = num_tokens % HP_BLOCK_SIZE - fp4_end = num_tokens - tail # exclusive - - # FP4-dequantized region: allow NVFP4 quant error. With unit global - # scale, worst-case absolute error is ~0.5 of the largest FP4 step - # within the value's block; 1.0 is a safe bound for the clamped - # latent range [-5, 5]. + fp4_end = num_tokens - tail if fp4_end > 0: torch.testing.assert_close( - flat[:fp4_end].float(), - latent[:fp4_end].float(), - atol=1.0, - rtol=0.5, - msg=f"FP4 region mismatch for num_tokens={num_tokens}", + recovered[:fp4_end].float(), latent[:fp4_end].float(), atol=1.0, rtol=0.5 ) - - # HP overlay region: must be exact BF16 roundtrip. if tail > 0: torch.testing.assert_close( - flat[fp4_end:].float(), - latent[fp4_end:].float(), - atol=0.0, - rtol=0.0, - msg=f"HP overlay mismatch for num_tokens={num_tokens}", + recovered[fp4_end:].float(), latent[fp4_end:].float(), atol=0, rtol=0 ) finally: kv_cache_manager.shutdown() -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -@pytest.mark.parametrize( - ("page_size", "ctx_tokens"), - [(16, 32), (128, 128)], - ids=["page16", "page128"], -) -def test_fp4_mla_hp_overlay_generation_phase(page_size: int, ctx_tokens: int, monkeypatch): - """After an aligned context, perform one decode step and verify that the - decode token surfaces through the HP overlay at the first tail slot.""" - monkeypatch.delenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, raising=False) +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +def test_fp4_mla_hp_overlay_generation_phase(): torch.manual_seed(1) device = torch.device("cuda") - kv_lora_rank = 512 qk_rope_head_dim = 64 head_dim = kv_lora_rank + qk_rope_head_dim - num_layers = 1 - max_seq_len = max(64, page_size * 2) - # Aligned context -> no tail until the first decode token. + page_size = FP4_MLA_TOKENS_PER_BLOCK + ctx_tokens = page_size total_tokens = ctx_tokens + 1 + num_layers = 1 mapping = Mapping(world_size=1, tp_size=1, rank=0) kv_cache_manager = KVCacheManager( - KvCacheConfig(max_tokens=max_seq_len, enable_block_reuse=False), + KvCacheConfig(max_tokens=page_size * 2, enable_block_reuse=False), _CacheType.SELFKONLY, num_layers=num_layers, num_kv_heads=1, head_dim=head_dim, tokens_per_block=page_size, - max_seq_len=max_seq_len, + max_seq_len=page_size * 2, max_batch_size=1, mapping=mapping, dtype=_DataType.NVFP4, ) try: kv_cache_manager.add_dummy_requests([0], [total_tokens]) - kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() kv_cache_manager.get_block_scale_buffers(0).zero_() - - # Write the full context (32 tokens) in one scatter, then the single - # decode token separately, matching the production flow. metadata = _build_metadata( kv_cache_manager, num_tokens=total_tokens, page_size=page_size, num_layers=num_layers ) ctx_latent = ( - torch.randn(ctx_tokens, head_dim, dtype=torch.bfloat16, device=device) * 1.5 - ).clamp_(-5.0, 5.0) - gen_latent = (torch.randn(1, head_dim, dtype=torch.bfloat16, device=device) * 1.5).clamp_( - -5.0, 5.0 + torch.randn(ctx_tokens, head_dim, dtype=torch.bfloat16, device=device) * 0.25 + ).clamp_(-1.0, 1.0) + gen_latent = (torch.randn(1, head_dim, dtype=torch.bfloat16, device=device) * 0.25).clamp_( + -1.0, 1.0 ) - # Context scatter + HP update: kv_len temporarily = 32. metadata.kv_lens_cuda_runtime = torch.tensor([ctx_tokens], dtype=torch.int32, device=device) - metadata.prompt_lens_cuda_runtime = torch.tensor( - [ctx_tokens], dtype=torch.int32, device=device - ) + metadata.prompt_lens_cuda_runtime = metadata.kv_lens_cuda_runtime metadata.prompt_lens_cpu_runtime = torch.tensor([ctx_tokens], dtype=torch.int32) - metadata.num_contexts = 1 - metadata.num_seqs = 1 - # Only ctx_tokens are visible to the scatter kernel this call. metadata.positions = torch.arange(ctx_tokens, dtype=torch.int32, device=device) metadata.batch_indices = torch.zeros(ctx_tokens, dtype=torch.int32, device=device) - scatter_fp4_mla_kv_cache(metadata, ctx_latent, layer_idx=0, token_offset=0) + scatter_fp4_mla_kv_cache( + metadata, + ctx_latent, + layer_idx=0, + token_offset=0, + phase="context", + local_layer=0, + v_head_dim=kv_lora_rank, + ) update_hp_kv_for_fp4_mla(metadata, ctx_latent, local_layer=0, phase="context") - # Decode scatter + HP update: append the single gen token at position 32. metadata.kv_lens_cuda_runtime = torch.tensor( [total_tokens], dtype=torch.int32, device=device ) + metadata.prompt_lens_cuda_runtime = torch.tensor([1], dtype=torch.int32, device=device) + metadata.prompt_lens_cpu_runtime = torch.tensor([1], dtype=torch.int32) metadata.positions = torch.tensor([ctx_tokens], dtype=torch.int32, device=device) metadata.batch_indices = torch.zeros(1, dtype=torch.int32, device=device) metadata.num_contexts = 0 - metadata.prompt_lens_cuda_runtime = torch.tensor([1], dtype=torch.int32, device=device) - metadata.prompt_lens_cpu_runtime = torch.tensor([1], dtype=torch.int32) - scatter_fp4_mla_kv_cache(metadata, gen_latent, layer_idx=0, token_offset=0) - update_hp_kv_for_fp4_mla(metadata, gen_latent, local_layer=0, phase="generation") - - # Now read the full 33-token history back. - num_blocks = (total_tokens + page_size - 1) // page_size - metadata.num_generation_blocks = num_blocks - metadata.paged_kv_indices = torch.tensor( - kv_cache_manager.get_batch_cache_indices([0])[0][:num_blocks], - dtype=torch.int32, - device=device, - ) - metadata.paged_kv_indptr_decode = torch.tensor( - [0, num_blocks], dtype=torch.int32, device=device + scatter_fp4_mla_kv_cache( + metadata, + gen_latent, + layer_idx=0, + token_offset=0, + phase="generation", + local_layer=0, + v_head_dim=kv_lora_rank, ) + update_hp_kv_for_fp4_mla(metadata, gen_latent, local_layer=0, phase="generation") + torch.cuda.synchronize() - combined = get_fp4_mla_decode_cache( + recovered = _materialize_single_seq_cache_with_hp_tail( metadata, layer_idx=0, local_layer=0, head_dim=head_dim, - dtype=torch.bfloat16, - ) - flat = combined.reshape(-1, head_dim)[:total_tokens] - - # Position 32 is the lone tail token; HP overlay must return exactly - # gen_latent[0]. - torch.testing.assert_close( - flat[ctx_tokens].float(), - gen_latent[0].float(), - atol=0.0, - rtol=0.0, - msg="HP overlay did not restore the decode token", + num_tokens=total_tokens, ) - - # Positions 0..31 come from FP4 dequant; accept NVFP4 roundtrip noise. torch.testing.assert_close( - flat[:ctx_tokens].float(), - ctx_latent.float(), - atol=1.0, - rtol=0.5, - msg="FP4 region mismatch on context tokens", + recovered[ctx_tokens].float(), gen_latent[0].float(), atol=0, rtol=0 ) finally: kv_cache_manager.shutdown() @@ -840,37 +758,34 @@ def test_fp4_mla_hp_pool_restores_rejected_linear_mtp_tokens(preallocated_snapsh old_len = 30 gen_len = 4 accepted_len = 2 - initial_hp = ( torch.arange(HP_BLOCK_SIZE * head_dim, dtype=torch.float32, device=device) .reshape(HP_BLOCK_SIZE, head_dim) .to(torch.bfloat16) ) hp_pool = initial_hp.reshape(1, 1, 1, HP_BLOCK_SIZE * head_dim).clone() - gen_latent = ( - torch.arange(gen_len * head_dim, dtype=torch.float32, device=device) - .reshape(gen_len, head_dim) - .add_(1000.0) - .to(torch.bfloat16) - ) metadata = SimpleNamespace( high_precision_kv_pool=hp_pool, - hp_pool_owners={0: 0}, - seq_slots=torch.zeros(1, dtype=torch.int32, device=device), - seq_slots_cpu=torch.zeros(1, dtype=torch.int32, device="cpu"), + fp4_mla_hp_snapshot_pool=None, + seq_slots=torch.tensor([0], dtype=torch.int32, device=device), + batch_indices=torch.zeros(gen_len, dtype=torch.int32, device=device), + positions=torch.arange(old_len, old_len + gen_len, dtype=torch.int32, device=device), kv_lens_cuda_runtime=torch.tensor([old_len + gen_len], dtype=torch.int32, device=device), prompt_lens_cuda_runtime=torch.tensor([gen_len], dtype=torch.int32, device=device), prompt_lens_cpu_runtime=torch.tensor([gen_len], dtype=torch.int32), - batch_indices=torch.zeros(gen_len, dtype=torch.int32, device=device), - positions=torch.arange(old_len, old_len + gen_len, dtype=torch.int32, device=device), num_contexts=0, num_seqs=1, - request_ids=[0], - is_cuda_graph=False, + num_tokens=gen_len, is_warmup=False, ) if preallocated_snapshot: metadata.fp4_mla_hp_snapshot_pool = torch.empty_like(hp_pool) + gen_latent = ( + torch.arange(gen_len * head_dim, dtype=torch.float32, device=device) + .reshape(gen_len, head_dim) + .add_(1000) + .to(torch.bfloat16) + ) update_hp_kv_for_fp4_mla(metadata, gen_latent, local_layer=0, phase="generation") repair_fp4_mla_hp_kv_for_mtp_rejection( @@ -881,251 +796,40 @@ def test_fp4_mla_hp_pool_restores_rejected_linear_mtp_tokens(preallocated_snapsh hp_view = hp_pool.view(HP_BLOCK_SIZE, head_dim) accepted_slots = [(old_len + idx) % HP_BLOCK_SIZE for idx in range(accepted_len)] rejected_slots = [(old_len + idx) % HP_BLOCK_SIZE for idx in range(accepted_len, gen_len)] - torch.testing.assert_close( - hp_view[accepted_slots].float(), - gen_latent[:accepted_len].float(), - atol=0.0, - rtol=0.0, - ) - torch.testing.assert_close( - hp_view[rejected_slots].float(), - initial_hp[rejected_slots].float(), - atol=0.0, - rtol=0.0, - ) + for idx, slot in enumerate(accepted_slots): + torch.testing.assert_close(hp_view[slot].float(), gen_latent[idx].float()) + for idx, slot in enumerate(rejected_slots, start=accepted_len): + torch.testing.assert_close( + hp_view[slot].float(), + initial_hp[slot].float(), + msg=f"rejected token {idx} was not restored", + ) if preallocated_snapshot: assert getattr(metadata, "_fp4_mla_mtp_hp_snapshots") else: assert getattr(metadata, "_fp4_mla_mtp_hp_snapshots") is None -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -@pytest.mark.parametrize("attention_env", ["0", "1"], ids=["linear_sf", "swizzled_sf"]) -def test_fp4_mla_scatter_last_page_no_oob(attention_env: str, monkeypatch): - """Scatter must not write past a page's scale region in either SF layout. - - With ``tokens_per_block=128`` the scatter's ``BLOCK_SF=64`` Triton block - has ``SF_PER_TOKEN=36`` valid lanes plus 28 masked-out lanes. For those - masked lanes the unconstrained offset (linear or swizzled) can exceed - the per-page stride. If masked-lane addresses are not pinned in-bounds, - the last physical page's masked stores fall past the sf_cache allocation, - which crashes with "illegal memory access" on Blackwell. - - The SF layout is gated by ``FLASHINFER_FP4_MLA_ATTENTION_ENV``; cover both - settings so a regression in either path is caught. The test writes tokens - that land exclusively on the LAST physical page and asserts: - 1. No bytes outside the target page are modified. - 2. The data round-trips correctly through the dequant path (so the - valid lanes still wrote the right values). - """ - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, attention_env) - torch.manual_seed(3) +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +def test_fp4_mla_real_scatter_writes_shared_2d_scales(): + torch.manual_seed(4) device = torch.device("cuda") - kv_lora_rank = 512 qk_rope_head_dim = 64 head_dim = kv_lora_rank + qk_rope_head_dim page_size = FP4_MLA_TOKENS_PER_BLOCK - num_pages = 4 - max_seq_len = num_pages * page_size + num_tokens = 32 num_layers = 1 mapping = Mapping(world_size=1, tp_size=1, rank=0) kv_cache_manager = KVCacheManager( - KvCacheConfig(max_tokens=max_seq_len, enable_block_reuse=False), + KvCacheConfig(max_tokens=page_size, enable_block_reuse=False), _CacheType.SELFKONLY, num_layers=num_layers, num_kv_heads=1, head_dim=head_dim, tokens_per_block=page_size, - max_seq_len=max_seq_len, - max_batch_size=1, - mapping=mapping, - dtype=_DataType.NVFP4, - ) - try: - kv_cache_manager.add_dummy_requests([0], [max_seq_len]) - kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() - sf_buf = kv_cache_manager.get_block_scale_buffers(0) - assert sf_buf is not None - - block_ids = kv_cache_manager.get_batch_cache_indices([0])[0][:num_pages] - last_physical_page = block_ids[-1] - # Sentinel non-target pages so a stray cross-page write is observable. - sf_buf.view(torch.uint8).fill_(0xA5) - sf_buf[last_physical_page].zero_() - kv_cache_manager.get_buffers(0).view(torch.uint8).fill_(0xA5) - kv_cache_manager.get_buffers(0).view(torch.uint8)[last_physical_page].zero_() - snapshot_sf = sf_buf.view(torch.uint8).clone() - snapshot_kv = kv_cache_manager.get_buffers(0).view(torch.uint8).clone() - - # Write tokens that land on the last physical page only. - last_start = (num_pages - 1) * page_size - num_tokens = page_size - paged_kv_indices = torch.tensor(block_ids, dtype=torch.int32, device=device) - paged_kv_indptr = torch.tensor([0, num_pages], dtype=torch.int32, device=device) - batch_indices = torch.zeros(num_tokens, dtype=torch.int32, device=device) - positions = torch.arange( - last_start, last_start + num_tokens, dtype=torch.int32, device=device - ) - - # Single-sequence read-back metadata (one big seq covering all pages). - metadata = _build_metadata( - kv_cache_manager, - num_tokens=max_seq_len, - page_size=page_size, - num_layers=num_layers, - ) - # Override scatter-only fields to write just the last-page slice. - metadata.batch_indices = batch_indices - metadata.positions = positions - metadata.paged_kv_indices = paged_kv_indices - metadata.paged_kv_indptr = paged_kv_indptr - - latent = ( - torch.randn(num_tokens, head_dim, dtype=torch.bfloat16, device=device) * 1.5 - ).clamp_(-5.0, 5.0) - scatter_fp4_mla_kv_cache(metadata, latent, layer_idx=0, token_offset=0) - torch.cuda.synchronize() - - # Bytes for any non-target page must be unchanged. - sf_after = sf_buf.view(torch.uint8) - kv_after = kv_cache_manager.get_buffers(0).view(torch.uint8) - for pid in range(sf_after.shape[0]): - if pid == last_physical_page: - continue - torch.testing.assert_close( - sf_after[pid], - snapshot_sf[pid], - atol=0, - rtol=0, - msg=f"scatter wrote to sf of non-target page {pid}", - ) - torch.testing.assert_close( - kv_after[pid], - snapshot_kv[pid], - atol=0, - rtol=0, - msg=f"scatter wrote to kv of non-target page {pid}", - ) - - # Read back via dequant and verify round-trip correctness. - metadata.num_contexts = 0 - metadata.num_seqs = 1 - combined = get_fp4_mla_decode_cache( - metadata, - layer_idx=0, - local_layer=0, - head_dim=head_dim, - dtype=torch.bfloat16, - ).reshape(-1, head_dim) - recovered = combined[last_start : last_start + num_tokens] - torch.testing.assert_close( - recovered.float(), - latent.float(), - atol=1.0, - rtol=0.5, - msg="last-page FP4 round-trip mismatch", - ) - finally: - kv_cache_manager.shutdown() - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_fp4_mla_dequant_invalid_page_ids_no_oob(monkeypatch): - """Dequant should reject a short page-id slice and guard invalid physical - pages in-kernel instead of doing unchecked page-stride arithmetic.""" - monkeypatch.delenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, raising=False) - device = torch.device("cuda") - - kv_lora_rank = 512 - qk_rope_head_dim = 64 - head_dim = kv_lora_rank + qk_rope_head_dim - page_size = FP4_MLA_TOKENS_PER_BLOCK - num_layers = 1 - - mapping = Mapping(world_size=1, tp_size=1, rank=0) - kv_cache_manager = KVCacheManager( - KvCacheConfig(max_tokens=page_size, enable_block_reuse=False), - _CacheType.SELFKONLY, - num_layers=num_layers, - num_kv_heads=1, - head_dim=head_dim, - tokens_per_block=page_size, - max_seq_len=page_size, - max_batch_size=1, - mapping=mapping, - dtype=_DataType.NVFP4, - ) - try: - kv_cache_manager.add_dummy_requests([0], [page_size]) - kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() - kv_cache_manager.get_block_scale_buffers(0).zero_() - - metadata = _build_metadata( - kv_cache_manager, - num_tokens=page_size, - page_size=page_size, - num_layers=num_layers, - ) - metadata.num_contexts = 0 - metadata.num_seqs = 0 - metadata.num_context_blocks = 0 - metadata.num_generation_blocks = 1 - metadata.paged_kv_indices = torch.empty(0, dtype=torch.int32, device=device) - - with pytest.raises(RuntimeError, match="needs 1 decode page ids"): - get_fp4_mla_decode_cache( - metadata, - layer_idx=0, - local_layer=0, - head_dim=head_dim, - dtype=torch.bfloat16, - ) - - invalid_page = kv_cache_manager.get_buffers(0).shape[0] - metadata.paged_kv_indices = torch.tensor([invalid_page], dtype=torch.int32, device=device) - combined = get_fp4_mla_decode_cache( - metadata, - layer_idx=0, - local_layer=0, - head_dim=head_dim, - dtype=torch.bfloat16, - ) - torch.cuda.synchronize() - torch.testing.assert_close( - combined, - torch.zeros_like(combined), - atol=0, - rtol=0, - ) - finally: - kv_cache_manager.shutdown() - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_fp4_mla_real_scatter_writes_shared_2d_scales(monkeypatch): - """Real FP4 scatter shares 16x16 scales only where K and V overlap.""" - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") - torch.manual_seed(4) - device = torch.device("cuda") - - kv_lora_rank = 512 - qk_rope_head_dim = 64 - head_dim = kv_lora_rank + qk_rope_head_dim - page_size = FP4_MLA_TOKENS_PER_BLOCK - num_tokens = 32 - num_layers = 1 - - mapping = Mapping(world_size=1, tp_size=1, rank=0) - kv_cache_manager = KVCacheManager( - KvCacheConfig(max_tokens=page_size, enable_block_reuse=False), - _CacheType.SELFKONLY, - num_layers=num_layers, - num_kv_heads=1, - head_dim=head_dim, - tokens_per_block=page_size, - max_seq_len=page_size, + max_seq_len=page_size, max_batch_size=1, mapping=mapping, dtype=_DataType.NVFP4, @@ -1136,13 +840,8 @@ def test_fp4_mla_real_scatter_writes_shared_2d_scales(monkeypatch): kv_cache_manager.get_block_scale_buffers(0).zero_() metadata = _build_metadata( - kv_cache_manager, - num_tokens=num_tokens, - page_size=page_size, - num_layers=num_layers, + kv_cache_manager, num_tokens=num_tokens, page_size=page_size, num_layers=num_layers ) - assert metadata.fp4_mla_v_scale_pool is not None - row = ( torch.arange(num_tokens, dtype=torch.float32, device=device) % HP_BLOCK_SIZE + 1.0 ).view(num_tokens, 1) @@ -1164,8 +863,8 @@ def test_fp4_mla_real_scatter_writes_shared_2d_scales(monkeypatch): torch.cuda.synchronize() physical_page = kv_cache_manager.get_batch_cache_indices([0])[0][0] - sf_per_token = head_dim // 16 - sf_per_page = page_size // 16 + sf_per_token = head_dim // FP4_BLOCK_SIZE + sf_per_page = page_size // FP4_BLOCK_SIZE k_page = ( kv_cache_manager.get_block_scale_buffers(0) .view(torch.float8_e4m3fn)[physical_page] @@ -1191,30 +890,20 @@ def test_fp4_mla_real_scatter_writes_shared_2d_scales(monkeypatch): ) k_bytes = k_page[k_offsets] assert bool((k_bytes[0] != 0).item()) - torch.testing.assert_close( - k_bytes, - k_bytes[0].expand_as(k_bytes), - atol=0, - rtol=0, - msg=f"K scales are not shared for dim block {dim_block}", - ) + torch.testing.assert_close(k_bytes, k_bytes[0].expand_as(k_bytes), atol=0, rtol=0) v_offsets = torch.tensor( [ - _swizzled_sf_offset(dim_block * 16 + row_idx, token_block, sf_per_page) - for row_idx in range(16) + _swizzled_sf_offset( + dim_block * FP4_BLOCK_SIZE + row_idx, token_block, sf_per_page + ) + for row_idx in range(FP4_BLOCK_SIZE) ], dtype=torch.long, device=device, ) v_bytes = v_page[v_offsets] - torch.testing.assert_close( - v_bytes, - k_bytes[0].expand_as(v_bytes), - atol=0, - rtol=0, - msg=f"K/V scales disagree for dim block {dim_block}", - ) + torch.testing.assert_close(v_bytes, k_bytes[0].expand_as(v_bytes), atol=0, rtol=0) tail_dim_block = kv_lora_rank // FP4_BLOCK_SIZE tail_offsets = torch.tensor( @@ -1227,18 +916,105 @@ def test_fp4_mla_real_scatter_writes_shared_2d_scales(monkeypatch): ) tail_bytes = k_page[tail_offsets] assert bool((tail_bytes[0] != 0).item()) - assert int(torch.unique(tail_bytes).numel()) > 1, "K-only tail scales must be per-token" + assert int(torch.unique(tail_bytes).numel()) > 1 finally: kv_cache_manager.shutdown() -@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 tensor cores") -def test_fp4_mla_attention_decode_residual_qk_duplicates_k_tail(monkeypatch): - """Residual-Q QK must use the same cached K tail for main and residual groups.""" - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") - torch.manual_seed(6) +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +def test_fp4_mla_scatter_last_page_no_oob(): + torch.manual_seed(3) device = torch.device("cuda") + kv_lora_rank = 512 + qk_rope_head_dim = 64 + head_dim = kv_lora_rank + qk_rope_head_dim + page_size = FP4_MLA_TOKENS_PER_BLOCK + num_pages = 4 + max_seq_len = num_pages * page_size + num_layers = 1 + mapping = Mapping(world_size=1, tp_size=1, rank=0) + kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=max_seq_len, enable_block_reuse=False), + _CacheType.SELFKONLY, + num_layers=num_layers, + num_kv_heads=1, + head_dim=head_dim, + tokens_per_block=page_size, + max_seq_len=max_seq_len, + max_batch_size=1, + mapping=mapping, + dtype=_DataType.NVFP4, + ) + try: + kv_cache_manager.add_dummy_requests([0], [max_seq_len]) + kv_cache = kv_cache_manager.get_buffers(0).view(torch.uint8) + sf_buf = kv_cache_manager.get_block_scale_buffers(0) + assert sf_buf is not None + + block_ids = kv_cache_manager.get_batch_cache_indices([0])[0][:num_pages] + last_physical_page = block_ids[-1] + sf_buf.view(torch.uint8).fill_(0xA5) + sf_buf[last_physical_page].zero_() + kv_cache.fill_(0xA5) + kv_cache[last_physical_page].zero_() + snapshot_sf = sf_buf.view(torch.uint8).clone() + snapshot_kv = kv_cache.clone() + + last_start = (num_pages - 1) * page_size + paged_kv_indices = torch.tensor(block_ids, dtype=torch.int32, device=device) + paged_kv_indptr = torch.tensor([0, num_pages], dtype=torch.int32, device=device) + metadata = _build_metadata( + kv_cache_manager, num_tokens=max_seq_len, page_size=page_size, num_layers=num_layers + ) + metadata.batch_indices = torch.zeros(page_size, dtype=torch.int32, device=device) + metadata.positions = torch.arange( + last_start, last_start + page_size, dtype=torch.int32, device=device + ) + metadata.paged_kv_indices = paged_kv_indices + metadata.paged_kv_indptr = paged_kv_indptr + metadata.kv_lens_cuda_runtime = torch.tensor( + [max_seq_len], dtype=torch.int32, device=device + ) + metadata.prompt_lens_cuda_runtime = torch.tensor( + [page_size], dtype=torch.int32, device=device + ) + metadata.prompt_lens_cpu_runtime = torch.tensor([page_size], dtype=torch.int32) + + latent = ( + torch.randn(page_size, head_dim, dtype=torch.bfloat16, device=device) * 0.25 + ).clamp_(-1.0, 1.0) + scatter_fp4_mla_kv_cache( + metadata, + latent, + layer_idx=0, + token_offset=0, + phase="context", + local_layer=0, + v_head_dim=kv_lora_rank, + ) + torch.cuda.synchronize() + + sf_after = sf_buf.view(torch.uint8) + kv_after = kv_cache + for pid in range(sf_after.shape[0]): + if pid == last_physical_page: + continue + torch.testing.assert_close(sf_after[pid], snapshot_sf[pid], atol=0, rtol=0) + torch.testing.assert_close(kv_after[pid], snapshot_kv[pid], atol=0, rtol=0) + + recovered = _materialize_reference_cache(metadata, 0, head_dim).reshape(-1, head_dim)[ + last_start : last_start + page_size + ] + torch.testing.assert_close(recovered.float(), latent.float(), atol=1.0, rtol=0.5) + finally: + kv_cache_manager.shutdown() + + +def _fp4_mla_attention_decode_residual_qk_duplicates_k_tail_impl(monkeypatch): + monkeypatch.setenv("TRTLLM_FP4_MLA_TRITON_PREPACK_V", "0") + torch.manual_seed(6) + device = torch.device("cuda") kv_lora_rank = 512 qk_rope_head_dim = 64 head_dim = kv_lora_rank + qk_rope_head_dim @@ -1264,28 +1040,15 @@ def test_fp4_mla_attention_decode_residual_qk_duplicates_k_tail(monkeypatch): kv_cache_manager.add_dummy_requests([0], [num_tokens]) kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() kv_cache_manager.get_block_scale_buffers(0).zero_() - metadata = _build_metadata( - kv_cache_manager, - num_tokens=num_tokens, - page_size=page_size, - num_layers=num_layers, + kv_cache_manager, num_tokens=num_tokens, page_size=page_size, num_layers=num_layers ) - assert metadata.fp4_mla_v_scale_pool is not None token_pattern = torch.linspace( - -0.4, - 0.4, - num_tokens, - dtype=torch.float32, - device=device, + -0.4, 0.4, num_tokens, dtype=torch.float32, device=device ).view(num_tokens, 1) dim_pattern = torch.linspace( - -0.7, - 0.7, - FP4_MLA_Q_RESIDUAL_DIM, - dtype=torch.float32, - device=device, + -0.7, 0.7, FP4_MLA_Q_RESIDUAL_DIM, dtype=torch.float32, device=device ).view(1, FP4_MLA_Q_RESIDUAL_DIM) latent = torch.zeros(num_tokens, head_dim, dtype=torch.bfloat16, device=device) latent[:, -FP4_MLA_Q_RESIDUAL_DIM:] = (token_pattern + dim_pattern).to(torch.bfloat16) @@ -1299,20 +1062,12 @@ def test_fp4_mla_attention_decode_residual_qk_duplicates_k_tail(monkeypatch): local_layer=0, v_head_dim=kv_lora_rank, ) - torch.cuda.synchronize() - metadata.num_contexts = 0 - metadata.num_seqs = 1 - + metadata.prompt_lens_cuda_runtime = torch.ones(1, dtype=torch.int32, device=device) + metadata.prompt_lens_cpu_runtime = torch.ones(1, dtype=torch.int32) q_nope = torch.zeros(1, num_heads, kv_lora_rank, dtype=torch.bfloat16, device=device) q_pe = ( - torch.linspace( - -0.9, - 0.9, - qk_rope_head_dim, - dtype=torch.float32, - device=device, - ) + torch.linspace(-0.9, 0.9, qk_rope_head_dim, dtype=torch.float32, device=device) .view(1, num_heads, qk_rope_head_dim) .to(torch.bfloat16) ) @@ -1344,35 +1099,40 @@ def test_fp4_mla_attention_decode_residual_qk_duplicates_k_tail(monkeypatch): q_fp4, q_sf, logical_dim=q_logical_dim, - sf_per_token=q_logical_dim // 16, + sf_per_token=q_logical_dim // FP4_BLOCK_SIZE, global_scale=_TEST_GLOBAL_SCALE, ) - dequant_cache = get_fp4_mla_decode_cache( - metadata, - layer_idx=0, - local_layer=0, - head_dim=head_dim, - dtype=torch.bfloat16, - ).reshape(-1, head_dim)[:num_tokens] + dequant_cache = _materialize_reference_cache(metadata, 0, head_dim).reshape(-1, head_dim)[ + :num_tokens + ] logical_k = _duplicate_tail_groups(dequant_cache.float(), FP4_MLA_Q_RESIDUAL_DIM) ref_scores = torch.matmul(q_dequant, logical_k.transpose(0, 1)) * sm_scale ref_probs = torch.softmax(ref_scores, dim=-1) - probs = metadata._fp4_mla_attention_p_prob_buf[:num_heads, :num_tokens] - + p_dequant = _dequant_fp4_swizzled( + metadata._fp4_mla_attention_p_buf, + metadata._fp4_mla_attention_p_sf_buf, + logical_dim=metadata.page_size, + sf_per_token=metadata.page_size // FP4_BLOCK_SIZE, + global_scale=FP4_MLA_P_GLOBAL_SCALE, + ) + probs = p_dequant[:num_heads, :num_tokens] torch.testing.assert_close( probs, ref_probs, - atol=2e-2, - rtol=2e-2, + atol=8e-2, + rtol=8e-2, msg="FP4 MLA residual-Q probabilities did not match duplicated K-tail reference", ) finally: + torch.cuda.synchronize() + _reset_triton_allocator() kv_cache_manager.shutdown() + torch.cuda.synchronize() + torch.cuda.empty_cache() @pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 tensor cores") -def test_fp4_mla_attention_decode_multi_seq_matches_reference(monkeypatch): - """Multiple decode sequences and heads must match a QK-softmax-PV reference.""" +def test_fp4_mla_attention_decode_matches_reference(monkeypatch): _assert_fp4_mla_attention_decode_accuracy( monkeypatch, backend="triton", @@ -1385,7 +1145,6 @@ def test_fp4_mla_attention_decode_multi_seq_matches_reference(monkeypatch): @pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 tensor cores") def test_fp4_mla_attention_decode_linear_mtp_matches_reference(monkeypatch): - """Linear MTP query rows must use per-query causal KV lengths.""" _assert_fp4_mla_attention_decode_accuracy( monkeypatch, backend="triton", @@ -1397,9 +1156,16 @@ def test_fp4_mla_attention_decode_linear_mtp_matches_reference(monkeypatch): ) -@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 tensor cores") +def test_fp4_mla_attention_decode_residual_qk_duplicates_k_tail(monkeypatch): + _fp4_mla_attention_decode_residual_qk_duplicates_k_tail_impl(monkeypatch) + + +@pytest.mark.skipif( + _is_cutile_unavailable(), + reason="requires Blackwell FP4 tensor cores and Triton tl.ext", +) def test_fp4_mla_attention_decode_cutile_matches_reference(monkeypatch): - """CuTile decode backend must preserve the FP4 MLA residual-tail contract.""" _assert_fp4_mla_attention_decode_accuracy( monkeypatch, backend="cutile", @@ -1410,9 +1176,11 @@ def test_fp4_mla_attention_decode_cutile_matches_reference(monkeypatch): ) -@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +@pytest.mark.skipif( + _is_cutile_unavailable(), + reason="requires Blackwell FP4 tensor cores and Triton tl.ext", +) def test_fp4_mla_attention_decode_cutile_shared_v_pack_matches_reference(monkeypatch): - """Shared V-packed storage must preserve the prepacked PV fast path.""" monkeypatch.setenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1") monkeypatch.setenv("TRTLLM_FP4_MLA_SHARE_V_PACK_STORAGE", "1") _assert_fp4_mla_attention_decode_accuracy( @@ -1425,9 +1193,11 @@ def test_fp4_mla_attention_decode_cutile_shared_v_pack_matches_reference(monkeyp ) -@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +@pytest.mark.skipif( + _is_cutile_unavailable(), + reason="requires Blackwell FP4 tensor cores and Triton tl.ext", +) def test_fp4_mla_attention_decode_cutile_grouped_tail_matches_reference(monkeypatch): - """Grouped page-stats must handle a partial final page group.""" monkeypatch.setenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1") monkeypatch.setenv("TRTLLM_FP4_MLA_SHARE_V_PACK_STORAGE", "1") monkeypatch.setenv("TRTLLM_FP4_MLA_GROUP_PAGES", "8") @@ -1441,10 +1211,47 @@ def test_fp4_mla_attention_decode_cutile_grouped_tail_matches_reference(monkeypa ) -@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +@pytest.mark.skipif( + _is_cutile_unavailable(), + reason="requires Blackwell FP4 tensor cores and Triton tl.ext", +) +def test_fp4_mla_attention_decode_cutile_linear_mtp_matches_reference(monkeypatch): + _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + backend="cutile", + num_heads=128, + seq_lens=[32, 128], + seed=13, + check_probs=False, + query_len_per_seq=3, + ) + + +@pytest.mark.skipif( + _is_cutile_unavailable(), + reason="requires Blackwell FP4 tensor cores and Triton tl.ext", +) +def test_fp4_mla_attention_decode_cutile_grouped_mtp_matches_reference(monkeypatch): + monkeypatch.setenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1") + monkeypatch.setenv("TRTLLM_FP4_MLA_SHARE_V_PACK_STORAGE", "1") + monkeypatch.setenv("TRTLLM_FP4_MLA_GROUP_PAGES", "8") + _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + backend="cutile", + num_heads=128, + seq_lens=[9 * FP4_MLA_TOKENS_PER_BLOCK, 9 * FP4_MLA_TOKENS_PER_BLOCK], + seed=23, + check_probs=False, + query_len_per_seq=4, + ) + + +@pytest.mark.skipif( + _is_cutile_unavailable(), + reason="requires Blackwell FP4 support and Triton tl.ext", +) def test_fp4_mla_cutile_shared_v_pack_storage_is_layer_tagged(monkeypatch): - """Layer ownership is metadata state; storage is reused across layers.""" - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, "cutile") + monkeypatch.setenv(FP4_MLA_ATTENTION_BACKEND_ENV, "cutile") monkeypatch.setenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1") monkeypatch.setenv("TRTLLM_FP4_MLA_SHARE_V_PACK_STORAGE", "1") @@ -1482,16 +1289,19 @@ def test_fp4_mla_cutile_shared_v_pack_storage_is_layer_tagged(monkeypatch): shared = metadata._fp4_mla_attention_v_packed_buf shared_ptr = shared.data_ptr() assert not hasattr(metadata, "_fp4_mla_attention_v_packed_buf_l0") - assert _get_cutile_v_packed_cache( - metadata, - 0, - kv_cache, - v_head_dim=v_head_dim, - page_size=page_size, - local_layer=0, - v_sf=v_sf[0], - page_ids=page_ids, - ) is not None + assert ( + _get_cutile_v_packed_cache( + metadata, + 0, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + local_layer=0, + v_sf=v_sf[0], + page_ids=page_ids, + ) + is not None + ) assert ( _get_cutile_v_packed_cache( metadata, @@ -1545,69 +1355,96 @@ def test_fp4_mla_cutile_shared_v_pack_storage_is_layer_tagged(monkeypatch): ) is None ) - assert _get_cutile_v_packed_cache( - metadata, - 1, - kv_cache, - v_head_dim=v_head_dim, - page_size=page_size, - local_layer=1, - v_sf=v_sf[1], - page_ids=page_ids, - ) is not None + assert ( + _get_cutile_v_packed_cache( + metadata, + 1, + kv_cache, + v_head_dim=v_head_dim, + page_size=page_size, + local_layer=1, + v_sf=v_sf[1], + page_ids=page_ids, + ) + is not None + ) -@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") -def test_fp4_mla_attention_decode_cutile_linear_mtp_matches_reference(monkeypatch): - """CuTile linear MTP rows must use per-query causal KV lengths.""" - _assert_fp4_mla_attention_decode_accuracy( - monkeypatch, - backend="cutile", - num_heads=128, - seq_lens=[32, 128], - seed=13, - check_probs=False, - query_len_per_seq=3, - ) +def _ceil_div(lhs: int, rhs: int) -> int: + return (lhs + rhs - 1) // rhs -@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 support") -def test_fp4_mla_attention_decode_cutile_grouped_mtp_matches_reference(monkeypatch): - """CuTile grouped page-stats must mask MTP future tokens on the final page.""" - monkeypatch.setenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1") - monkeypatch.setenv("TRTLLM_FP4_MLA_SHARE_V_PACK_STORAGE", "1") - monkeypatch.setenv("TRTLLM_FP4_MLA_GROUP_PAGES", "8") - _assert_fp4_mla_attention_decode_accuracy( - monkeypatch, - backend="cutile", - num_heads=128, - seq_lens=[9 * FP4_MLA_TOKENS_PER_BLOCK, 9 * FP4_MLA_TOKENS_PER_BLOCK], - seed=23, - check_probs=False, - query_len_per_seq=4, +def _cuda_event_benchmark(fn, *, warmup_iters=10, iters=100): + for _ in range(warmup_iters): + fn() + torch.cuda.synchronize() + + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(iters): + fn() + end.record() + torch.cuda.synchronize() + return start.elapsed_time(end) / iters + + +def _estimate_fp4_mla_attention_decode_mbu_bytes( + *, + batch_size: int, + seq_len: int, + num_heads: int, + kv_lora_rank: int, + qk_rope_head_dim: int, + page_size: int, +) -> int: + fp4_block_size = 16 + head_block_size = 128 + bf16_bytes = 2 + fp32_bytes = 4 + + q_input_dim = kv_lora_rank + qk_rope_head_dim + qk_dim = q_input_dim + FP4_MLA_Q_RESIDUAL_DIM + pages_per_seq = _ceil_div(seq_len, page_size) + padded_seq_len = pages_per_seq * page_size + head_blocks = _ceil_div(num_heads, head_block_size) + q_rows = batch_size * num_heads + + qk_fp4_bytes_per_token = qk_dim // 2 + _ceil_div(qk_dim, fp4_block_size) + v_fp4_bytes_per_token = kv_lora_rank // 2 + _ceil_div(kv_lora_rank, fp4_block_size) + p_bytes_per_seq = padded_seq_len // 2 + _ceil_div(padded_seq_len, fp4_block_size) + + q_setup_bytes = q_rows * q_input_dim * bf16_bytes * 3 + q_rows * qk_fp4_bytes_per_token + qk_cache_bytes = 2 * batch_size * head_blocks * padded_seq_len * qk_fp4_bytes_per_token + qk_q_bytes = 2 * batch_size * num_heads * pages_per_seq * qk_fp4_bytes_per_token + stats_bytes = batch_size * num_heads * 2 * fp32_bytes * (1 + pages_per_seq) + p_prob_bytes = batch_size * num_heads * padded_seq_len * fp32_bytes * 2 + p_quant_bytes = 2 * batch_size * num_heads * p_bytes_per_seq + pv_cache_bytes = batch_size * head_blocks * padded_seq_len * v_fp4_bytes_per_token + output_bytes = q_rows * kv_lora_rank * bf16_bytes + + return ( + q_setup_bytes + + qk_cache_bytes + + qk_q_bytes + + stats_bytes + + p_prob_bytes + + p_quant_bytes + + pv_cache_bytes + + output_bytes ) @pytest.mark.skipif( os.environ.get("TRTLLM_RUN_FP4_MLA_ATTENTION_BENCHMARK") != "1", - reason=("Manual perf benchmark; set TRTLLM_RUN_FP4_MLA_ATTENTION_BENCHMARK=1 to run"), + reason="manual perf benchmark", ) @pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 tensor cores") @pytest.mark.parametrize("batch_size", [16, 32, 64, 128, 256], ids=lambda x: f"bs{x}") @pytest.mark.parametrize("seq_len", [8192], ids=lambda x: f"seq{x}") -def test_fp4_mla_attention_decode_perf_benchmark( - batch_size, - seq_len, - monkeypatch, -): - """Opt-in microbenchmark for ``run_fp4_mla_attention_decode``. - - Run manually with: - ``TRTLLM_RUN_FP4_MLA_ATTENTION_BENCHMARK=1 pytest -s -k fp4_mla_attention_decode_perf``. - """ - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") - backend = os.environ.get(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, "triton") - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_BACKEND_ENV, backend) +def test_fp4_mla_attention_decode_perf_benchmark(batch_size, seq_len, monkeypatch): + backend = os.environ.get(FP4_MLA_ATTENTION_BACKEND_ENV, "triton") + monkeypatch.setenv(FP4_MLA_ATTENTION_BACKEND_ENV, backend) num_heads = 128 ( kv_cache_manager, @@ -1667,12 +1504,9 @@ def run_decode(): @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_fp4_mla_shared_tile_rejects_unaligned_context_start(monkeypatch): - """The no-dequant FP4 path requires context chunks to start on a 16-token boundary.""" - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") +def test_fp4_mla_shared_tile_rejects_unaligned_context_start(): torch.manual_seed(5) device = torch.device("cuda") - kv_lora_rank = 512 qk_rope_head_dim = 64 head_dim = kv_lora_rank + qk_rope_head_dim @@ -1699,15 +1533,12 @@ def test_fp4_mla_shared_tile_rejects_unaligned_context_start(monkeypatch): kv_cache_manager.add_dummy_requests([0], [total_tokens]) kv_cache_manager.get_buffers(0).view(torch.uint8).zero_() kv_cache_manager.get_block_scale_buffers(0).zero_() - metadata = _build_metadata( - kv_cache_manager, - num_tokens=cached_tokens, - page_size=page_size, - num_layers=num_layers, + kv_cache_manager, num_tokens=new_tokens, page_size=page_size, num_layers=num_layers + ) + metadata.positions = torch.arange( + cached_tokens, total_tokens, dtype=torch.int32, device=device ) - assert metadata.fp4_mla_v_scale_pool is not None - new_latent = ( torch.randn(new_tokens, head_dim, dtype=torch.bfloat16, device=device) * 1.5 ).clamp_(-5.0, 5.0) @@ -1719,10 +1550,7 @@ def test_fp4_mla_shared_tile_rejects_unaligned_context_start(monkeypatch): ) metadata.prompt_lens_cpu_runtime = torch.tensor([new_tokens], dtype=torch.int32) - with pytest.raises( - ValueError, - match="start position.*16-token aligned", - ): + with pytest.raises(ValueError, match="start position.*16-token aligned"): scatter_fp4_mla_kv_cache( metadata, new_latent, diff --git a/tests/unittest/_torch/executor/test_mla_tokens_per_block.py b/tests/unittest/_torch/executor/test_mla_tokens_per_block.py index 5a42fe67cd01..6a586ddba205 100644 --- a/tests/unittest/_torch/executor/test_mla_tokens_per_block.py +++ b/tests/unittest/_torch/executor/test_mla_tokens_per_block.py @@ -5,7 +5,6 @@ from tensorrt_llm._torch.pyexecutor.py_executor_creator import ( FLASH_MLA_TOKENS_PER_BLOCK, - FLASHINFER_FP4_MLA_ATTENTION_ENV, FP4_MLA_TOKENS_PER_BLOCK, _select_mla_tokens_per_block, ) @@ -20,9 +19,11 @@ def _non_mla_config(): return SimpleNamespace() -def _model_config(kv_cache_quant_algo=None, enable_flash_mla=False): +def _model_config(kv_cache_quant_algo=None, enable_flash_mla=False, attn_backend=None): quant_config = SimpleNamespace(kv_cache_quant_algo=kv_cache_quant_algo) - return SimpleNamespace(quant_config=quant_config, enable_flash_mla=enable_flash_mla) + return SimpleNamespace( + quant_config=quant_config, enable_flash_mla=enable_flash_mla, attn_backend=attn_backend + ) def _kv_cache_config(dtype="auto", tokens_per_block=32): @@ -57,8 +58,7 @@ def test_flash_mla_non_fp4_uses_flash_mla_tokens_per_block(): assert kv_cache_config.tokens_per_block == FLASH_MLA_TOKENS_PER_BLOCK -def test_fp4_mla_dequant_flow_uses_flash_mla_tokens_per_block(monkeypatch): - monkeypatch.delenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, raising=False) +def test_fp4_mla_dequant_flow_uses_flash_mla_tokens_per_block(): kv_cache_config = _kv_cache_config(tokens_per_block=32) tokens_per_block = _select_mla_tokens_per_block( @@ -72,13 +72,12 @@ def test_fp4_mla_dequant_flow_uses_flash_mla_tokens_per_block(monkeypatch): assert kv_cache_config.tokens_per_block == FLASH_MLA_TOKENS_PER_BLOCK -def test_fp4_mla_attention_uses_128_tokens_per_block_from_quant_config(monkeypatch): - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") +def test_fp4_mla_attention_uses_128_tokens_per_block_from_quant_config(): kv_cache_config = _kv_cache_config(tokens_per_block=32) tokens_per_block = _select_mla_tokens_per_block( _mla_config(), - _model_config(kv_cache_quant_algo=QuantAlgo.NVFP4, enable_flash_mla=True), + _model_config(kv_cache_quant_algo=QuantAlgo.NVFP4, attn_backend="TRTLLM"), kv_cache_config, kv_cache_config.tokens_per_block, ) @@ -87,13 +86,12 @@ def test_fp4_mla_attention_uses_128_tokens_per_block_from_quant_config(monkeypat assert kv_cache_config.tokens_per_block == FP4_MLA_TOKENS_PER_BLOCK -def test_fp4_mla_attention_uses_128_tokens_per_block_from_kv_cache_dtype(monkeypatch): - monkeypatch.setenv(FLASHINFER_FP4_MLA_ATTENTION_ENV, "1") +def test_fp4_mla_attention_uses_128_tokens_per_block_from_kv_cache_dtype(): kv_cache_config = _kv_cache_config(dtype="nvfp4", tokens_per_block=32) tokens_per_block = _select_mla_tokens_per_block( _mla_config(), - _model_config(enable_flash_mla=True), + _model_config(enable_flash_mla=True, attn_backend="TRTLLM"), kv_cache_config, kv_cache_config.tokens_per_block, ) From 3549726f78a609bb1a8071c76e5fe5ca57ad4557 Mon Sep 17 00:00:00 2001 From: Yuxian Qiu <142763828+yuxianq@users.noreply.github.com> Date: Sat, 20 Jun 2026 00:19:57 +0000 Subject: [PATCH 11/11] [TRTLLM-12807][fix] Fix PR gate checks Signed-off-by: Yuxian Qiu <142763828+yuxianq@users.noreply.github.com> --- .../attention_backend/fp4_mla_cutile.py | 613 +++++++++++++----- .../_torch/attention_backend/trtllm.py | 1 - 2 files changed, 437 insertions(+), 177 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py index f3c097ca3e12..c518c312668b 100644 --- a/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py @@ -17,7 +17,6 @@ import triton import triton.language as tl - FP4_BLOCK_SIZE = 16 FP4_MLA_P_GLOBAL_SCALE = 448.0 * 6.0 @@ -42,12 +41,26 @@ def _swizzled_scale_size(rows: int, logical_cols: int) -> int: def _get_kv_cache_strides(kv_cache: torch.Tensor) -> tuple[int, int, int, int, int, int]: if kv_cache.dim() == 3: num_pages, page_size, packed_dim = kv_cache.shape - return num_pages, page_size, packed_dim, kv_cache.stride(0), kv_cache.stride(1), kv_cache.stride(2) + return ( + num_pages, + page_size, + packed_dim, + kv_cache.stride(0), + kv_cache.stride(1), + kv_cache.stride(2), + ) if kv_cache.dim() >= 5: num_pages = kv_cache.shape[0] page_size = kv_cache.shape[2] packed_dim = kv_cache.shape[4] - return num_pages, page_size, packed_dim, kv_cache.stride(0), kv_cache.stride(2), kv_cache.stride(4) + return ( + num_pages, + page_size, + packed_dim, + kv_cache.stride(0), + kv_cache.stride(2), + kv_cache.stride(4), + ) raise ValueError( "kv_cache must be shaped (num_pages, page_size, packed_dim) or (num_pages, ..., page_size, ..., packed_dim)." ) @@ -94,12 +107,18 @@ def _fp4_mla_swizzled_sf_offset(row_idx, col_idx, SF_PER_TOKEN: tl.constexpr): row_in_group1 = (row_idx % 128) // 32 row_group = row_idx // 128 return ( - col_in_group + col_group * (4 * 128) + row_in_group0 * 16 + row_in_group1 * 4 + row_group * (128 * padded_cols) + col_in_group + + col_group * (4 * 128) + + row_in_group0 * 16 + + row_in_group1 * 4 + + row_group * (128 * padded_cols) ) @triton.jit -def _fp4_mla_swizzled_sf_offset_row_block(row_group, row_offsets, col_idx, SF_PER_TOKEN: tl.constexpr): +def _fp4_mla_swizzled_sf_offset_row_block( + row_group, row_offsets, col_idx, SF_PER_TOKEN: tl.constexpr +): padded_cols = ((SF_PER_TOKEN + 3) // 4) * 4 col_part = (col_idx % 4) + (col_idx // 4) * (4 * 128) row_part = (row_offsets % 32) * 16 + ((row_offsets % 128) // 32) * 4 @@ -369,7 +388,9 @@ def fp4_mla_repack_v_cache( if v_head_dim % block_v != 0: raise ValueError(f"v_head_dim={v_head_dim} must be divisible by block_v={block_v}.") if kv_cache.ndim < 5: - raise ValueError(f"kv_cache must expose the paged FP4 layout, got shape={tuple(kv_cache.shape)}.") + raise ValueError( + f"kv_cache must expose the paged FP4 layout, got shape={tuple(kv_cache.shape)}." + ) num_pages = kv_cache.shape[0] num_dim_blocks = triton.cdiv(v_head_dim, block_v) launch_meta = { @@ -454,8 +475,12 @@ def _fp4_mla_qk_scores_tile( else: valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) safe_compact_page = tl.where(valid_compact_page, compact_page, 0) - physical_page = tl.load(src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1).to(tl.int64) - valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) + ) safe_physical_page = tl.where(valid_physical_page, physical_page, 0) q_rows = q_row_base + head_offsets if ASSUME_FULL_HEADS: @@ -515,13 +540,6 @@ def _fp4_mla_qk_scores_tile( tile_shape=[1, 8, 2, 256], tile_dim_map=[0, 1, 2, 3], ) - q_sf_tail_view = tl.ext.make_view( - base=q_sf_ptr, - shapes=[q_num_rows // 128, ((Q_SF_PER_TOKEN + 3) // 4), 2, 256], - strides=[128 * (((Q_SF_PER_TOKEN + 3) // 4) * 4), 512, 256, 1], - tile_shape=[1, 1, 2, 256], - tile_dim_map=[0, 1, 2, 3], - ) k_sf_full_view = tl.ext.make_view( base=sf_cache_ptr, shapes=[num_pages, 1, ((K_SF_PER_TOKEN + 3) // 4), 2, 256], @@ -544,7 +562,9 @@ def _fp4_mla_qk_scores_tile( full_q_scales = tl.ext.load_view_tko(q_sf_full_view, [q_row_group.to(tl.int32), 0, 0, 0]) full_q_scales = full_q_scales.reshape([1, 8, 32, 4, 4]).trans(0, 3, 2, 1, 4) full_q_scales = full_q_scales.reshape([BLOCK_H, 32]) - full_k_scales = tl.ext.load_view_tko(k_sf_full_view, [safe_physical_page.to(tl.int32), 0, 0, 0, 0]) + full_k_scales = tl.ext.load_view_tko( + k_sf_full_view, [safe_physical_page.to(tl.int32), 0, 0, 0, 0] + ) full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) scores = tl.dot_scaled( @@ -561,7 +581,9 @@ def _fp4_mla_qk_scores_tile( tail_k_vals = k_tail_desc.load([safe_physical_page.to(tl.int32), 0, 256]) tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) - tail_k_scales = tl.ext.load_view_tko(k_sf_tail_view, [safe_physical_page.to(tl.int32), 0, 8, 0, 0]) + tail_k_scales = tl.ext.load_view_tko( + k_sf_tail_view, [safe_physical_page.to(tl.int32), 0, 8, 0, 0] + ) tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) q_tail_vals = q_tail_desc.load([(q_row_base + head_start).to(tl.int32), 256]) @@ -572,7 +594,9 @@ def _fp4_mla_qk_scores_tile( q_odd_vals = q_odd_vals.reshape([BLOCK_H, 32]) q_tail_sf_cols = 32 + tl.arange(0, 8) - q_tail_sf_offsets = _fp4_mla_swizzled_sf_offset(q_rows[:, None], q_tail_sf_cols[None, :], Q_SF_PER_TOKEN) + q_tail_sf_offsets = _fp4_mla_swizzled_sf_offset( + q_rows[:, None], q_tail_sf_cols[None, :], Q_SF_PER_TOKEN + ) q_tail_scales = tl.load(q_sf_ptr + q_tail_sf_offsets) q_tail_scales = q_tail_scales.reshape([BLOCK_H, 4, 2]) q_even_scales, q_odd_scales = tl.split(q_tail_scales) @@ -650,15 +674,23 @@ def _fp4_mla_qk_scores_tile( mask_k = q_elem_offsets < Q_HEAD_D safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) - if USE_TMA_DATA_LOAD and FULL_BLOCK_END > 0 and q_start + BLOCK_K <= non_residual_groups * FP4_BLOCK: + if ( + USE_TMA_DATA_LOAD + and FULL_BLOCK_END > 0 + and q_start + BLOCK_K <= non_residual_groups * FP4_BLOCK + ): q_vals = q_desc.load([(q_row_base + head_start).to(tl.int32), q_start // 2]) - k_vals = k_desc.load([safe_physical_page.to(tl.int32), token_start.to(tl.int32), q_start // 2]) + k_vals = k_desc.load( + [safe_physical_page.to(tl.int32), token_start.to(tl.int32), q_start // 2] + ) k_vals = tl.reshape(k_vals, (BLOCK_T, BLOCK_K // 2)) if not ASSUME_VALID_PAGES: k_vals = tl.where(valid_physical_page, k_vals, 0) else: q_vals = tl.load( - q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + safe_packed_q_cols[None, :] * q_fp4_s1, + q_fp4_ptr + + safe_q_rows[:, None] * q_fp4_s0 + + safe_packed_q_cols[None, :] * q_fp4_s1, mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], other=0, ) @@ -667,7 +699,9 @@ def _fp4_mla_qk_scores_tile( + safe_physical_page * kv_s0 + token_offsets[:, None].to(tl.int64) * kv_s2 + safe_packed_k_cols[None, :] * kv_s4, - mask=mask_k[None, :] if ASSUME_VALID_PAGES else valid_physical_page & mask_k[None, :], + mask=mask_k[None, :] + if ASSUME_VALID_PAGES + else valid_physical_page & mask_k[None, :], other=0, ) @@ -680,8 +714,12 @@ def _fp4_mla_qk_scores_tile( mask_sf = q_sf_cols < Q_SF_PER_TOKEN safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) - q_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN) - k_sf_offsets = _fp4_mla_swizzled_sf_offset(token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN) + q_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN + ) + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN + ) q_scales = tl.load(q_sf_ptr + q_sf_offsets) k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) scores = tl.dot_scaled( @@ -730,14 +768,20 @@ def _fp4_mla_qk_scores_tile( other=0, ) k_sf_cols = non_residual_groups + residual_scale_offsets - k_sf_offsets = _fp4_mla_swizzled_sf_offset(token_offsets[:, None], k_sf_cols[None, :], K_SF_PER_TOKEN) + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], k_sf_cols[None, :], K_SF_PER_TOKEN + ) k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) q_tail_cols = q_start // 2 + tl.arange(0, 64) if USE_TMA_DATA_LOAD and ASSUME_FULL_HEADS: - q_tail_vals = q_tail_desc.load([(q_row_base + head_start).to(tl.int32), q_start // 2]) + q_tail_vals = q_tail_desc.load( + [(q_row_base + head_start).to(tl.int32), q_start // 2] + ) elif ASSUME_FULL_HEADS: - q_tail_vals = tl.load(q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1) + q_tail_vals = tl.load( + q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1 + ) else: q_tail_vals = tl.load( q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + q_tail_cols[None, :] * q_fp4_s1, @@ -796,7 +840,9 @@ def _fp4_mla_qk_scores_tile( safe_packed_q_cols = tl.where(mask_k, packed_q_cols, 0) safe_packed_k_cols = tl.where(mask_k, packed_k_cols, 0) q_vals = tl.load( - q_fp4_ptr + safe_q_rows[:, None] * q_fp4_s0 + safe_packed_q_cols[None, :] * q_fp4_s1, + q_fp4_ptr + + safe_q_rows[:, None] * q_fp4_s0 + + safe_packed_q_cols[None, :] * q_fp4_s1, mask=mask_k[None, :] if ASSUME_FULL_HEADS else mask_h[:, None] & mask_k[None, :], other=0, ) @@ -805,7 +851,9 @@ def _fp4_mla_qk_scores_tile( + safe_physical_page * kv_s0 + token_offsets[:, None].to(tl.int64) * kv_s2 + safe_packed_k_cols[None, :] * kv_s4, - mask=mask_k[None, :] if ASSUME_VALID_PAGES else valid_physical_page & mask_k[None, :], + mask=mask_k[None, :] + if ASSUME_VALID_PAGES + else valid_physical_page & mask_k[None, :], other=0, ) @@ -818,8 +866,12 @@ def _fp4_mla_qk_scores_tile( mask_sf = q_sf_cols < Q_SF_PER_TOKEN safe_q_sf_cols = tl.where(mask_sf, q_sf_cols, 0) safe_k_sf_cols = tl.where(mask_sf, k_sf_cols, 0) - q_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN) - k_sf_offsets = _fp4_mla_swizzled_sf_offset(token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN) + q_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_q_rows[:, None], safe_q_sf_cols[None, :], Q_SF_PER_TOKEN + ) + k_sf_offsets = _fp4_mla_swizzled_sf_offset( + token_offsets[:, None], safe_k_sf_cols[None, :], K_SF_PER_TOKEN + ) q_scales = tl.load(q_sf_ptr + q_sf_offsets) k_scales = tl.load(sf_cache_ptr + safe_physical_page * sf_s0 + k_sf_offsets) scores = tl.dot_scaled( @@ -947,7 +999,9 @@ def _fp4_mla_attention_stats_kernel( scores = tl.where(mask_h[:, None], scores * qk_scale, -float("inf")) else: valid_t = page_start + offs_t < kv_len - scores = tl.where(mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf")) + scores = tl.where( + mask_h[:, None] & valid_t[None, :], scores * qk_scale, -float("inf") + ) page_max = tl.max(scores, axis=1) new_max = tl.maximum(max_score, page_max) denom = denom * tl.math.exp2((max_score - new_max) * 1.4426950408889634) + tl.sum( @@ -1108,7 +1162,9 @@ def _fp4_mla_attention_page_stats_kernel( grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) amax = tl.max(grouped_probs, axis=2) inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) - stored_scale = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) + stored_scale = tl.where( + amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0 + ) scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) even_probs, odd_probs = tl.split(pairs) @@ -1117,21 +1173,27 @@ def _fp4_mla_attention_page_stats_kernel( if ASSUME_VALID_PAGES: safe_compact_page = page_table_start + page_rel else: - valid_compact_page = (page_table_start + page_rel >= 0) & (page_table_start + page_rel < page_ids_len) + valid_compact_page = (page_table_start + page_rel >= 0) & ( + page_table_start + page_rel < page_ids_len + ) safe_compact_page = tl.where(valid_compact_page, page_table_start + page_rel, 0) if P_BY_QUERY: p_page = query_idx * MAX_PAGES + page_rel else: p_page = safe_compact_page p_rows = p_page * NUM_HEADS + offs_h - safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + ) scale_cols = tl.arange(0, SF_PER_PAGE) if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) else: - sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) if ASSUME_FULL_HEADS: if ASSUME_VALID_PAGES: tl.store(p_sf_ptr + sf_offsets, stored_scale) @@ -1141,7 +1203,9 @@ def _fp4_mla_attention_page_stats_kernel( tl.store( p_sf_ptr + sf_offsets, stored_scale, - mask=mask_h[:, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None], + mask=mask_h[:, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None], ) byte_offsets = tl.arange(0, FP4_BLOCK // 2) @@ -1154,12 +1218,16 @@ def _fp4_mla_attention_page_stats_kernel( elif ASSUME_FULL_HEADS: if ASSUME_VALID_PAGES: tl.store( - p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + p_fp4_ptr + + safe_p_rows[:, None, None] * p_s0 + + byte_cols[None, :, :] * p_s1, packed, ) else: tl.store( - p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, + p_fp4_ptr + + safe_p_rows[:, None, None] * p_s0 + + byte_cols[None, :, :] * p_s1, packed, mask=valid_compact_page, ) @@ -1167,7 +1235,9 @@ def _fp4_mla_attention_page_stats_kernel( tl.store( p_fp4_ptr + safe_p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, packed, - mask=mask_h[:, None, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None, None], + mask=mask_h[:, None, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None, None], ) if ASSUME_FULL_HEADS: @@ -1361,7 +1431,9 @@ def _fp4_mla_attention_page_stats_grouped_kernel( scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) full_k_vals = k_desc.load([physical_page.to(tl.int32), 0, 0]) full_k_vals = tl.reshape(full_k_vals, (BLOCK_T, 256)) - full_k_scales = tl.ext.load_view_tko(k_sf_full_view, [physical_page.to(tl.int32), 0, 0, 0, 0]) + full_k_scales = tl.ext.load_view_tko( + k_sf_full_view, [physical_page.to(tl.int32), 0, 0, 0, 0] + ) full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) scores = tl.dot_scaled( @@ -1378,7 +1450,9 @@ def _fp4_mla_attention_page_stats_grouped_kernel( tail_k_vals = k_tail_desc.load([physical_page.to(tl.int32), 0, 256]) tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) - tail_k_scales = tl.ext.load_view_tko(k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0]) + tail_k_scales = tl.ext.load_view_tko( + k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0] + ) tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) if DUPLICATE_TAIL_K: @@ -1427,10 +1501,11 @@ def _fp4_mla_attention_page_stats_grouped_kernel( ) scores = scores * qk_scale - full_softmax_page = ASSUME_FULL_PAGES or (MASK_MTP_FINAL_PAGE_ONLY and page_rel < MAX_PAGES - 1) + full_softmax_page = ASSUME_FULL_PAGES or ( + MASK_MTP_FINAL_PAGE_ONLY and page_rel < MAX_PAGES - 1 + ) if ASSUME_FULL_HEADS and ( - (ASSUME_FULL_PAGES and not ALLOW_PARTIAL_GROUPS) - or full_softmax_page + (ASSUME_FULL_PAGES and not ALLOW_PARTIAL_GROUPS) or full_softmax_page ): page_max = tl.max(scores, axis=1) exp_scores = tl.math.exp2((scores - page_max[:, None]) * 1.4426950408889634) @@ -1449,9 +1524,9 @@ def _fp4_mla_attention_page_stats_grouped_kernel( next_group_max = tl.maximum(group_max, page_max) old_delta = tl.where(group_sum > 0.0, group_max - next_group_max, 0.0) new_delta = tl.where(page_sum > 0.0, page_max - next_group_max, 0.0) - group_sum = group_sum * tl.math.exp2(old_delta * 1.4426950408889634) + page_sum * tl.math.exp2( - new_delta * 1.4426950408889634 - ) + group_sum = group_sum * tl.math.exp2( + old_delta * 1.4426950408889634 + ) + page_sum * tl.math.exp2(new_delta * 1.4426950408889634) group_max = next_group_max grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) @@ -1464,9 +1539,15 @@ def _fp4_mla_attention_page_stats_grouped_kernel( packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) out_offsets = query_idx * page_stats_s0 + page_rel * page_stats_s1 + offs_h - tl.store(page_max_ptr + out_offsets, page_max, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) + tl.store( + page_max_ptr + out_offsets, page_max, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS) + ) if not GROUP_REDUCE_STATS: - tl.store(page_sum_ptr + out_offsets, page_sum, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) + tl.store( + page_sum_ptr + out_offsets, + page_sum, + mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS), + ) if P_BY_QUERY: p_page = query_idx * MAX_PAGES + page_rel @@ -1475,7 +1556,9 @@ def _fp4_mla_attention_page_stats_grouped_kernel( sf_offsets = _fp4_mla_swizzled_sf_offset_row_block( p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) - tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) + tl.store( + p_sf_ptr + sf_offsets, stored_scale, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS) + ) if ALLOW_PARTIAL_GROUPS: byte_offsets = tl.arange(0, FP4_BLOCK // 2) byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] @@ -1491,7 +1574,9 @@ def _fp4_mla_attention_page_stats_grouped_kernel( tl.reshape(packed, (BLOCK_H, PAGE_SIZE // 2)), ) if GROUP_REDUCE_STATS: - group_max_offsets = query_idx * page_stats_s0 + (logical_page_group * 2) * page_stats_s1 + offs_h + group_max_offsets = ( + query_idx * page_stats_s0 + (logical_page_group * 2) * page_stats_s1 + offs_h + ) group_sum_offsets = group_max_offsets + page_stats_s1 tl.store(page_sum_ptr + group_max_offsets, group_max) tl.store(page_sum_ptr + group_sum_offsets, group_sum) @@ -1697,12 +1782,16 @@ def _fp4_mla_attention_page_stats_grouped_mtp_pair_kernel( full_k_vals = k_desc.load([physical_page.to(tl.int32), 0, 0]) full_k_vals = tl.reshape(full_k_vals, (BLOCK_T, 256)) - full_k_scales = tl.ext.load_view_tko(k_sf_full_view, [physical_page.to(tl.int32), 0, 0, 0, 0]) + full_k_scales = tl.ext.load_view_tko( + k_sf_full_view, [physical_page.to(tl.int32), 0, 0, 0, 0] + ) full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) tail_k_vals = k_tail_desc.load([physical_page.to(tl.int32), 0, 256]) tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) - tail_k_scales = tl.ext.load_view_tko(k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0]) + tail_k_scales = tl.ext.load_view_tko( + k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0] + ) tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) @@ -1762,9 +1851,9 @@ def _fp4_mla_attention_page_stats_grouped_mtp_pair_kernel( next_group_max0 = tl.maximum(group_max0, page_max0) old_delta0 = tl.where(group_sum0 > 0.0, group_max0 - next_group_max0, 0.0) new_delta0 = tl.where(page_sum0 > 0.0, page_max0 - next_group_max0, 0.0) - group_sum0 = group_sum0 * tl.math.exp2(old_delta0 * 1.4426950408889634) + page_sum0 * tl.math.exp2( - new_delta0 * 1.4426950408889634 - ) + group_sum0 = group_sum0 * tl.math.exp2( + old_delta0 * 1.4426950408889634 + ) + page_sum0 * tl.math.exp2(new_delta0 * 1.4426950408889634) group_max0 = next_group_max0 grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) amax = tl.max(grouped_probs, axis=2) @@ -1775,14 +1864,26 @@ def _fp4_mla_attention_page_stats_grouped_mtp_pair_kernel( even_probs, odd_probs = tl.split(pairs) packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) out_offsets0 = query_idx0 * page_stats_s0 + page_rel * page_stats_s1 + offs_h - tl.store(page_max_ptr + out_offsets0, page_max0, mask=valid_group_page0 | (not ALLOW_PARTIAL_GROUPS)) + tl.store( + page_max_ptr + out_offsets0, + page_max0, + mask=valid_group_page0 | (not ALLOW_PARTIAL_GROUPS), + ) if not GROUP_REDUCE_STATS: - tl.store(page_sum_ptr + out_offsets0, page_sum0, mask=valid_group_page0 | (not ALLOW_PARTIAL_GROUPS)) + tl.store( + page_sum_ptr + out_offsets0, + page_sum0, + mask=valid_group_page0 | (not ALLOW_PARTIAL_GROUPS), + ) p_page0 = query_idx0 * MAX_PAGES + page_rel sf_offsets0 = _fp4_mla_swizzled_sf_offset_row_block( p_page0, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) - tl.store(p_sf_ptr + sf_offsets0, stored_scale, mask=valid_group_page0 | (not ALLOW_PARTIAL_GROUPS)) + tl.store( + p_sf_ptr + sf_offsets0, + stored_scale, + mask=valid_group_page0 | (not ALLOW_PARTIAL_GROUPS), + ) if ALLOW_PARTIAL_GROUPS: byte_offsets = tl.arange(0, FP4_BLOCK // 2) byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] @@ -1854,9 +1955,9 @@ def _fp4_mla_attention_page_stats_grouped_mtp_pair_kernel( next_group_max1 = tl.maximum(group_max1, page_max1) old_delta1 = tl.where(group_sum1 > 0.0, group_max1 - next_group_max1, 0.0) new_delta1 = tl.where(page_sum1 > 0.0, page_max1 - next_group_max1, 0.0) - group_sum1 = group_sum1 * tl.math.exp2(old_delta1 * 1.4426950408889634) + page_sum1 * tl.math.exp2( - new_delta1 * 1.4426950408889634 - ) + group_sum1 = group_sum1 * tl.math.exp2( + old_delta1 * 1.4426950408889634 + ) + page_sum1 * tl.math.exp2(new_delta1 * 1.4426950408889634) group_max1 = next_group_max1 grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) amax = tl.max(grouped_probs, axis=2) @@ -1867,14 +1968,26 @@ def _fp4_mla_attention_page_stats_grouped_mtp_pair_kernel( even_probs, odd_probs = tl.split(pairs) packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) out_offsets1 = query_idx1 * page_stats_s0 + page_rel * page_stats_s1 + offs_h - tl.store(page_max_ptr + out_offsets1, page_max1, mask=valid_group_page1 | (not ALLOW_PARTIAL_GROUPS)) + tl.store( + page_max_ptr + out_offsets1, + page_max1, + mask=valid_group_page1 | (not ALLOW_PARTIAL_GROUPS), + ) if not GROUP_REDUCE_STATS: - tl.store(page_sum_ptr + out_offsets1, page_sum1, mask=valid_group_page1 | (not ALLOW_PARTIAL_GROUPS)) + tl.store( + page_sum_ptr + out_offsets1, + page_sum1, + mask=valid_group_page1 | (not ALLOW_PARTIAL_GROUPS), + ) p_page1 = query_idx1 * MAX_PAGES + page_rel sf_offsets1 = _fp4_mla_swizzled_sf_offset_row_block( p_page1, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) - tl.store(p_sf_ptr + sf_offsets1, stored_scale, mask=valid_group_page1 | (not ALLOW_PARTIAL_GROUPS)) + tl.store( + p_sf_ptr + sf_offsets1, + stored_scale, + mask=valid_group_page1 | (not ALLOW_PARTIAL_GROUPS), + ) if ALLOW_PARTIAL_GROUPS: byte_offsets = tl.arange(0, FP4_BLOCK // 2) byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] @@ -2039,9 +2152,9 @@ def _fp4_mla_attention_page_stats_grouped_generic_kernel( next_group_max = tl.maximum(group_max, page_max) old_delta = tl.where(group_sum > 0.0, group_max - next_group_max, 0.0) new_delta = tl.where(page_sum > 0.0, page_max - next_group_max, 0.0) - group_sum = group_sum * tl.math.exp2(old_delta * 1.4426950408889634) + page_sum * tl.math.exp2( - new_delta * 1.4426950408889634 - ) + group_sum = group_sum * tl.math.exp2( + old_delta * 1.4426950408889634 + ) + page_sum * tl.math.exp2(new_delta * 1.4426950408889634) group_max = next_group_max grouped_probs = tl.reshape(exp_scores, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK)) @@ -2054,13 +2167,23 @@ def _fp4_mla_attention_page_stats_grouped_generic_kernel( packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) out_offsets = gen_idx * page_stats_s0 + page_rel * page_stats_s1 + offs_h - tl.store(page_max_ptr + out_offsets, page_max, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) + tl.store( + page_max_ptr + out_offsets, page_max, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS) + ) if not GROUP_REDUCE_STATS: - tl.store(page_sum_ptr + out_offsets, page_sum, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) + tl.store( + page_sum_ptr + out_offsets, + page_sum, + mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS), + ) safe_p_rows = safe_compact_page * NUM_HEADS + offs_h - sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) - tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS)) + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) + tl.store( + p_sf_ptr + sf_offsets, stored_scale, mask=valid_group_page | (not ALLOW_PARTIAL_GROUPS) + ) if ALLOW_PARTIAL_GROUPS: byte_offsets = tl.arange(0, FP4_BLOCK // 2) byte_cols = scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] @@ -2076,7 +2199,9 @@ def _fp4_mla_attention_page_stats_grouped_generic_kernel( ) if GROUP_REDUCE_STATS: - group_max_offsets = gen_idx * page_stats_s0 + (logical_page_group * 2) * page_stats_s1 + offs_h + group_max_offsets = ( + gen_idx * page_stats_s0 + (logical_page_group * 2) * page_stats_s1 + offs_h + ) group_sum_offsets = group_max_offsets + page_stats_s1 tl.store(page_sum_ptr + group_max_offsets, group_max) tl.store(page_sum_ptr + group_sum_offsets, group_sum) @@ -2109,7 +2234,10 @@ def _fp4_mla_attention_reduce_stats_kernel( if GROUP_REDUCE_STATS: for group_rel in tl.range(0, NUM_PAGE_GROUPS): group_max = tl.load( - page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, + page_sum_ptr + + gen_idx * page_stats_s0 + + (group_rel * 2) * page_stats_s1 + + safe_offs_h, mask=mask_h, other=-float("inf"), ) @@ -2127,16 +2255,26 @@ def _fp4_mla_attention_reduce_stats_kernel( if GROUP_REDUCE_STATS: for group_rel in tl.range(0, NUM_PAGE_GROUPS): group_max = tl.load( - page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, + page_sum_ptr + + gen_idx * page_stats_s0 + + (group_rel * 2) * page_stats_s1 + + safe_offs_h, mask=mask_h, other=-float("inf"), ) group_sum = tl.load( - page_sum_ptr + gen_idx * page_stats_s0 + (group_rel * 2 + 1) * page_stats_s1 + safe_offs_h, + page_sum_ptr + + gen_idx * page_stats_s0 + + (group_rel * 2 + 1) * page_stats_s1 + + safe_offs_h, mask=mask_h, other=0.0, ) - denom += tl.where(group_sum > 0.0, group_sum * tl.math.exp2((group_max - max_score) * 1.4426950408889634), 0.0) + denom += tl.where( + group_sum > 0.0, + group_sum * tl.math.exp2((group_max - max_score) * 1.4426950408889634), + 0.0, + ) else: for page_rel in tl.range(0, MAX_PAGES): page_max = tl.load( @@ -2149,7 +2287,11 @@ def _fp4_mla_attention_reduce_stats_kernel( mask=mask_h, other=0.0, ) - denom += tl.where(page_sum > 0.0, page_sum * tl.math.exp2((page_max - max_score) * 1.4426950408889634), 0.0) + denom += tl.where( + page_sum > 0.0, + page_sum * tl.math.exp2((page_max - max_score) * 1.4426950408889634), + 0.0, + ) tl.store(max_ptr + gen_idx * stats_s0 + safe_offs_h, max_score, mask=mask_h) tl.store(denom_ptr + gen_idx * stats_s0 + safe_offs_h, denom, mask=mask_h) @@ -2282,16 +2424,18 @@ def _fp4_mla_attention_page_stats_half_grouped_kernel( next_group_max = tl.maximum(group_max, tile_max) old_delta = tl.where(group_sum > 0.0, group_max - next_group_max, 0.0) new_delta = tl.where(tile_sum > 0.0, tile_max - next_group_max, 0.0) - group_sum = group_sum * tl.math.exp2(old_delta * 1.4426950408889634) + tile_sum * tl.math.exp2( - new_delta * 1.4426950408889634 - ) + group_sum = group_sum * tl.math.exp2( + old_delta * 1.4426950408889634 + ) + tile_sum * tl.math.exp2(new_delta * 1.4426950408889634) group_max = next_group_max grouped_probs = tl.reshape(exp_scores, (BLOCK_H, BLOCK_T // FP4_BLOCK, FP4_BLOCK)) amax = tl.max(grouped_probs, axis=2) inv_local_scale = tl.where(amax > 0.0, 6.0 / amax, 1.0) stored_scale = tl.where(amax > 0.0, tl.minimum(amax * (P_GLOBAL_SCALE / 6.0), 448.0), 1.0) - scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, BLOCK_T // FP4_BLOCK, 1)) + scaled_probs = grouped_probs * tl.reshape( + inv_local_scale, (BLOCK_H, BLOCK_T // FP4_BLOCK, 1) + ) pairs = tl.reshape(scaled_probs, (BLOCK_H, BLOCK_T // FP4_BLOCK, FP4_BLOCK // 2, 2)) even_probs, odd_probs = tl.split(pairs) packed = _fp4_e2m1_quantize_packed(even_probs, odd_probs) @@ -2308,7 +2452,9 @@ def _fp4_mla_attention_page_stats_half_grouped_kernel( tl.store(p_sf_ptr + sf_offsets, stored_scale, mask=valid_tile) byte_offsets = tl.arange(0, FP4_BLOCK // 2) - byte_cols = token_base // 2 + local_scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + byte_cols = ( + token_base // 2 + local_scale_cols[:, None] * (FP4_BLOCK // 2) + byte_offsets[None, :] + ) p_rows = safe_compact_page * NUM_HEADS + offs_h tl.store( p_fp4_ptr + p_rows[:, None, None] * p_s0 + byte_cols[None, :, :] * p_s1, @@ -2382,7 +2528,9 @@ def _fp4_mla_attention_prob_scale_kernel( ) max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - factor = tl.where(denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0) + factor = tl.where( + denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0 + ) if P_BY_QUERY: p_page = query_idx * MAX_PAGES + page_rel @@ -2396,7 +2544,9 @@ def _fp4_mla_attention_prob_scale_kernel( p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) else: - sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) @@ -2464,7 +2614,9 @@ def _fp4_mla_attention_prob_scale_half_kernel( ) max_score = tl.load(max_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) denom = tl.load(denom_ptr + query_idx * stats_s0 + safe_offs_h, mask=mask_h, other=0.0) - factor = tl.where(denom > 0.0, tl.math.exp2((tile_max - max_score) * 1.4426950408889634) / denom, 0.0) + factor = tl.where( + denom > 0.0, tl.math.exp2((tile_max - max_score) * 1.4426950408889634) / denom, 0.0 + ) scale_cols = tile_in_page * (BLOCK_T // 16) + tl.arange(0, BLOCK_T // 16) if ASSUME_FULL_HEADS and ASSUME_VALID_PAGES and NUM_HEADS == 128 and BLOCK_H == 128: @@ -2473,8 +2625,12 @@ def _fp4_mla_attention_prob_scale_half_kernel( ) else: p_rows = compact_page * NUM_HEADS + offs_h - safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, compact_page * NUM_HEADS) - sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, compact_page * NUM_HEADS) + ) + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) @@ -2540,7 +2696,10 @@ def _fp4_mla_attention_prob_scale_from_group_stats_kernel( max_score = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) for group_rel in tl.range(0, NUM_PAGE_GROUPS): group_max = tl.load( - page_sum_ptr + query_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, + page_sum_ptr + + query_idx * page_stats_s0 + + (group_rel * 2) * page_stats_s1 + + safe_offs_h, mask=mask_h, other=-float("inf"), ) @@ -2549,12 +2708,18 @@ def _fp4_mla_attention_prob_scale_from_group_stats_kernel( denom = tl.zeros((BLOCK_H,), dtype=tl.float32) for group_rel in tl.range(0, NUM_PAGE_GROUPS): group_max = tl.load( - page_sum_ptr + query_idx * page_stats_s0 + (group_rel * 2) * page_stats_s1 + safe_offs_h, + page_sum_ptr + + query_idx * page_stats_s0 + + (group_rel * 2) * page_stats_s1 + + safe_offs_h, mask=mask_h, other=-float("inf"), ) group_sum = tl.load( - page_sum_ptr + query_idx * page_stats_s0 + (group_rel * 2 + 1) * page_stats_s1 + safe_offs_h, + page_sum_ptr + + query_idx * page_stats_s0 + + (group_rel * 2 + 1) * page_stats_s1 + + safe_offs_h, mask=mask_h, other=0.0, ) @@ -2564,7 +2729,9 @@ def _fp4_mla_attention_prob_scale_from_group_stats_kernel( 0.0, ) - factor = tl.where(denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0) + factor = tl.where( + denom > 0.0, tl.math.exp2((page_max - max_score) * 1.4426950408889634) / denom, 0.0 + ) if P_BY_QUERY: p_page = query_idx * MAX_PAGES + page_rel @@ -2578,7 +2745,9 @@ def _fp4_mla_attention_prob_scale_from_group_stats_kernel( p_page, offs_h[:, None], scale_cols[None, :], SF_PER_PAGE ) else: - sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) scales = tl.load(p_sf_ptr + sf_offsets, mask=mask_h[:, None], other=1.0).to(tl.float32) tl.store(p_sf_ptr + sf_offsets, scales * factor[:, None], mask=mask_h[:, None]) @@ -3039,7 +3208,9 @@ def _fp4_mla_attention_pv_kernel( dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE ) else: - v_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE) + v_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) if USE_TMA_P_LOAD: tl.assume(p_s0 % 8 == 0) tl.assume(p_s1 == 1) @@ -3181,7 +3352,10 @@ def _fp4_mla_attention_pv_kernel( out_vals = out_vals.to(tl.float16) if USE_TMA_OUT_STORE: out_desc.store( - [(query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], + [ + (query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ], out_vals, ) else: @@ -3211,10 +3385,12 @@ def _fp4_mla_attention_pv_kernel( else: valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) safe_compact_page = tl.where(valid_compact_page, compact_page, 0) - physical_page = tl.load(src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1).to( - tl.int64 + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) ) - valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) safe_physical_page = tl.where(valid_physical_page, physical_page, 0) if P_BY_QUERY: @@ -3222,16 +3398,22 @@ def _fp4_mla_attention_pv_kernel( else: p_page = safe_compact_page p_rows = p_page * NUM_HEADS + offs_h - safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + ) if USE_TMA_P_LOAD: p_vals = p_desc.load([(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) else: p_vals = tl.load( p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, - mask=mask_h[:, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None], + mask=mask_h[:, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None], other=0, ) - p_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + p_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) p_scales = tl.load(p_sf_ptr + p_sf_offsets) if ASSUME_FULL_PAGES: @@ -3309,7 +3491,10 @@ def _fp4_mla_attention_pv_kernel( elif out_ptr.dtype.element_ty == tl.float16: out_vals = out_vals.to(tl.float16) out_desc.store( - [(query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], + [ + (query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ], out_vals, ) else: @@ -3319,7 +3504,10 @@ def _fp4_mla_attention_pv_kernel( ) else: tl.store( - out_ptr + query_idx * out_s0 + safe_offs_h[:, None] * out_s1 + safe_offs_v[None, :] * out_s2, + out_ptr + + query_idx * out_s0 + + safe_offs_h[:, None] * out_s1 + + safe_offs_v[None, :] * out_s2, acc * out_scale, mask=mask_h[:, None] & mask_v[None, :], ) @@ -3412,7 +3600,9 @@ def _fp4_mla_attention_pv_prepacked_v_kernel( dim_block, offs_v[:, None], scale_cols[None, :], SF_PER_PAGE ) else: - v_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE) + v_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) if USE_TMA_P_LOAD: tl.assume(p_s0 % 8 == 0) tl.assume(p_s1 == 1) @@ -3580,7 +3770,9 @@ def _fp4_mla_attention_pv_prepacked_v_kernel( 0.0, ) if PV_SCALE_IN_SF: - p_scales_with_factor = (p_scales.to(tl.float32) * factor[:, None]).to(tl.float8e4nv) + p_scales_with_factor = (p_scales.to(tl.float32) * factor[:, None]).to( + tl.float8e4nv + ) acc = tl.ext.dot_scaled( v_vals, v_scales, @@ -3625,7 +3817,10 @@ def _fp4_mla_attention_pv_prepacked_v_kernel( out_vals = out_vals.to(tl.float16) if USE_TMA_OUT_STORE: out_desc.store( - [(query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], + [ + (query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ], out_vals, ) else: @@ -3655,10 +3850,12 @@ def _fp4_mla_attention_pv_prepacked_v_kernel( else: valid_compact_page = (compact_page >= 0) & (compact_page < page_ids_len) safe_compact_page = tl.where(valid_compact_page, compact_page, 0) - physical_page = tl.load(src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1).to( - tl.int64 + physical_page = tl.load( + src_page_ids_ptr + safe_compact_page, mask=valid_compact_page, other=-1 + ).to(tl.int64) + valid_physical_page = ( + valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) ) - valid_physical_page = valid_compact_page & (physical_page >= 0) & (physical_page < num_pages) safe_physical_page = tl.where(valid_physical_page, physical_page, 0) if P_BY_QUERY: @@ -3666,16 +3863,22 @@ def _fp4_mla_attention_pv_prepacked_v_kernel( else: p_page = safe_compact_page p_rows = p_page * NUM_HEADS + offs_h - safe_p_rows = p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + safe_p_rows = ( + p_rows if ASSUME_FULL_HEADS else tl.where(mask_h, p_rows, p_page * NUM_HEADS) + ) if USE_TMA_P_LOAD: p_vals = p_desc.load([(p_page * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), 0]) else: p_vals = tl.load( p_fp4_ptr + safe_p_rows[:, None] * p_s0 + packed_t[None, :] * p_s1, - mask=mask_h[:, None] if ASSUME_VALID_PAGES else valid_compact_page & mask_h[:, None], + mask=mask_h[:, None] + if ASSUME_VALID_PAGES + else valid_compact_page & mask_h[:, None], other=0, ) - p_sf_offsets = _fp4_mla_swizzled_sf_offset(safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE) + p_sf_offsets = _fp4_mla_swizzled_sf_offset( + safe_p_rows[:, None], scale_cols[None, :], SF_PER_PAGE + ) p_scales = tl.load(p_sf_ptr + p_sf_offsets) if ASSUME_FULL_PAGES: @@ -3758,7 +3961,10 @@ def _fp4_mla_attention_pv_prepacked_v_kernel( elif out_ptr.dtype.element_ty == tl.float16: out_vals = out_vals.to(tl.float16) out_desc.store( - [(query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32)], + [ + (query_idx * NUM_HEADS + head_block * BLOCK_H).to(tl.int32), + (dim_block * BLOCK_V).to(tl.int32), + ], out_vals, ) else: @@ -3768,7 +3974,10 @@ def _fp4_mla_attention_pv_prepacked_v_kernel( ) else: tl.store( - out_ptr + query_idx * out_s0 + safe_offs_h[:, None] * out_s1 + safe_offs_v[None, :] * out_s2, + out_ptr + + query_idx * out_s0 + + safe_offs_h[:, None] * out_s1 + + safe_offs_v[None, :] * out_s2, acc * out_scale, mask=mask_h[:, None] & mask_v[None, :], ) @@ -3996,7 +4205,6 @@ def _fp4_mla_attention_online_qkpv_group_kernel( dim_block = combo - page_group * NUM_DIM_BLOCKS offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - offs_t = tl.arange(0, BLOCK_T) offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) packed_t = tl.arange(0, PAGE_SIZE // 2) scale_cols = tl.arange(0, SF_PER_PAGE) @@ -4100,7 +4308,9 @@ def _fp4_mla_attention_online_qkpv_group_kernel( scores = tl.zeros((BLOCK_H, BLOCK_T), dtype=tl.float32) full_k_vals = k_desc.load([physical_page.to(tl.int32), 0, 0]) full_k_vals = tl.reshape(full_k_vals, (BLOCK_T, 256)) - full_k_scales = tl.ext.load_view_tko(k_sf_full_view, [physical_page.to(tl.int32), 0, 0, 0, 0]) + full_k_scales = tl.ext.load_view_tko( + k_sf_full_view, [physical_page.to(tl.int32), 0, 0, 0, 0] + ) full_k_scales = full_k_scales.reshape([1, 1, 8, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) full_k_scales = full_k_scales.reshape([BLOCK_T, 32]) scores = tl.dot_scaled( @@ -4117,7 +4327,9 @@ def _fp4_mla_attention_online_qkpv_group_kernel( tail_k_vals = k_tail_desc.load([physical_page.to(tl.int32), 0, 256]) tail_k_vals = tl.reshape(tail_k_vals, (BLOCK_T, 32)) - tail_k_scales = tl.ext.load_view_tko(k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0]) + tail_k_scales = tl.ext.load_view_tko( + k_sf_tail_view, [physical_page.to(tl.int32), 0, 8, 0, 0] + ) tail_k_scales = tail_k_scales.reshape([1, 1, 1, 32, 4, 4]).trans(0, 1, 4, 3, 2, 5) tail_k_scales = tail_k_scales.reshape([BLOCK_T, 4]) scores = tl.dot_scaled( @@ -4165,8 +4377,12 @@ def _fp4_mla_attention_online_qkpv_group_kernel( scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) even_probs, odd_probs = tl.split(pairs) - p_vals = tl.reshape(_fp4_e2m1_quantize_packed(even_probs, odd_probs), (BLOCK_H, PAGE_SIZE // 2)) - v_scale_offsets = _fp4_mla_swizzled_sf_offset(offs_v[:, None], scale_cols[None, :], SF_PER_PAGE) + p_vals = tl.reshape( + _fp4_e2m1_quantize_packed(even_probs, odd_probs), (BLOCK_H, PAGE_SIZE // 2) + ) + v_scale_offsets = _fp4_mla_swizzled_sf_offset( + offs_v[:, None], scale_cols[None, :], SF_PER_PAGE + ) v_scales = tl.load(v_sf_ptr + physical_page * vsf_s0 + v_scale_offsets) page_o = tl.ext.dot_scaled( v_vals, @@ -4181,7 +4397,9 @@ def _fp4_mla_attention_online_qkpv_group_kernel( ) else: v_scale_cols = packed_t // (FP4_BLOCK // 2) - v_scale_offsets = _fp4_mla_swizzled_sf_offset(offs_v[:, None], v_scale_cols[None, :], SF_PER_PAGE) + v_scale_offsets = _fp4_mla_swizzled_sf_offset( + offs_v[:, None], v_scale_cols[None, :], SF_PER_PAGE + ) v_scales = tl.load(v_sf_ptr + physical_page * vsf_s0 + v_scale_offsets).to(tl.float32) v_low = _fp4_e2m1_to_f32(v_vals & 0x0F) * v_scales v_high = _fp4_e2m1_to_f32((v_vals >> 4) & 0x0F) * v_scales @@ -4201,10 +4419,7 @@ def _fp4_mla_attention_online_qkpv_group_kernel( group_m = next_m partial_o_offsets = ( - page_group * po_s0 - + gen_idx * po_s1 - + offs_h[:, None] * po_s2 - + offs_v[None, :] * po_s3 + page_group * po_s0 + gen_idx * po_s1 + offs_h[:, None] * po_s2 + offs_v[None, :] * po_s3 ) tl.store(partial_o_ptr + partial_o_offsets, group_o.T) if dim_block == 0: @@ -4415,7 +4630,9 @@ def _fp4_mla_attention_gen_qkpv_group_kernel( scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) even_probs, odd_probs = tl.split(pairs) - p_vals = tl.reshape(_fp4_e2m1_quantize_packed(even_probs, odd_probs), (BLOCK_H, PAGE_SIZE // 2)) + p_vals = tl.reshape( + _fp4_e2m1_quantize_packed(even_probs, odd_probs), (BLOCK_H, PAGE_SIZE // 2) + ) page_o0 = _fp4_mla_pv_page_o_prepacked_raw( v_sf_ptr, @@ -4661,7 +4878,9 @@ def _fp4_mla_attention_mtp_fused_qkpv_group_kernel( scaled_probs = grouped_probs * tl.reshape(inv_local_scale, (BLOCK_H, SF_PER_PAGE, 1)) pairs = tl.reshape(scaled_probs, (BLOCK_H, SF_PER_PAGE, FP4_BLOCK // 2, 2)) even_probs, odd_probs = tl.split(pairs) - p_vals = tl.reshape(_fp4_e2m1_quantize_packed(even_probs, odd_probs), (BLOCK_H, PAGE_SIZE // 2)) + p_vals = tl.reshape( + _fp4_e2m1_quantize_packed(even_probs, odd_probs), (BLOCK_H, PAGE_SIZE // 2) + ) v_row = (physical_page * NUM_DIM_BLOCKS + dim_block) * BLOCK_V v_vals = v_packed_desc.load([v_row.to(tl.int32), 0]) @@ -4733,10 +4952,7 @@ def _fp4_mla_attention_mtp_fused_qkpv_group_kernel( group_m = next_m partial_o_offsets = ( - page_group * po_s0 - + query_idx * po_s1 - + offs_h[:, None] * po_s2 - + offs_v[None, :] * po_s3 + page_group * po_s0 + query_idx * po_s1 + offs_h[:, None] * po_s2 + offs_v[None, :] * po_s3 ) tl.store(partial_o_ptr + partial_o_offsets, group_o.T) partial_ml_offsets = page_group * pm_s0 + query_idx * pm_s1 + offs_h @@ -4785,7 +5001,9 @@ def _fp4_mla_attention_online_qkpv_reduce_kernel( for group_idx in tl.range(0, NUM_PAGE_GROUPS): group_m = tl.load(partial_m_ptr + group_idx * pm_s0 + gen_idx * pm_s1 + offs_h) group_l = tl.load(partial_l_ptr + group_idx * pm_s0 + gen_idx * pm_s1 + offs_h) - scale = tl.where(group_l > 0.0, tl.math.exp2((group_m - global_m) * 1.4426950408889634), 0.0) + scale = tl.where( + group_l > 0.0, tl.math.exp2((group_m - global_m) * 1.4426950408889634), 0.0 + ) partial_o = tl.load( partial_o_ptr + group_idx * po_s0 @@ -4850,8 +5068,6 @@ def _fp4_mla_attention_pv_group_partial_prepacked_v_kernel( page_group = combo // NUM_DIM_BLOCKS dim_block = combo - page_group * NUM_DIM_BLOCKS - offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) - offs_v = dim_block * BLOCK_V + tl.arange(0, BLOCK_V) page_table_start = tl.load(paged_kv_indptr_decode_ptr + gen_idx).to(tl.int64) tl.assume(p_s0 % 8 == 0) @@ -4957,11 +5173,9 @@ def _fp4_mla_attention_pv_group_partial_prepacked_v_kernel( partial_o_desc.store( [ - ( - page_group * (po_s0 // po_s2) - + gen_idx * (po_s1 // po_s2) - + head_block * BLOCK_H - ).to(tl.int32), + (page_group * (po_s0 // po_s2) + gen_idx * (po_s1 // po_s2) + head_block * BLOCK_H).to( + tl.int32 + ), (dim_block * BLOCK_V).to(tl.int32), ], acc.T, @@ -5008,21 +5222,27 @@ def _fp4_mla_attention_pv_group_partial_reduce_kernel( global_m = tl.full((BLOCK_H,), -float("inf"), dtype=tl.float32) for group_idx in tl.range(0, NUM_PAGE_GROUPS): - group_m = tl.load(page_sum_ptr + gen_idx * page_stats_s0 + (group_idx * 2) * page_stats_s1 + offs_h) + group_m = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + (group_idx * 2) * page_stats_s1 + offs_h + ) global_m = tl.maximum(global_m, group_m) global_l = tl.zeros((BLOCK_H,), dtype=tl.float32) acc = tl.zeros((BLOCK_H, BLOCK_V), dtype=tl.float32) for group_idx in tl.range(0, NUM_PAGE_GROUPS): - group_m = tl.load(page_sum_ptr + gen_idx * page_stats_s0 + (group_idx * 2) * page_stats_s1 + offs_h) - group_l = tl.load(page_sum_ptr + gen_idx * page_stats_s0 + (group_idx * 2 + 1) * page_stats_s1 + offs_h) - scale = tl.where(group_l > 0.0, tl.math.exp2((group_m - global_m) * 1.4426950408889634), 0.0) + group_m = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + (group_idx * 2) * page_stats_s1 + offs_h + ) + group_l = tl.load( + page_sum_ptr + gen_idx * page_stats_s0 + (group_idx * 2 + 1) * page_stats_s1 + offs_h + ) + scale = tl.where( + group_l > 0.0, tl.math.exp2((group_m - global_m) * 1.4426950408889634), 0.0 + ) partial_o = partial_o_desc.load( [ ( - group_idx * (po_s0 // po_s2) - + gen_idx * (po_s1 // po_s2) - + head_block * BLOCK_H + group_idx * (po_s0 // po_s2) + gen_idx * (po_s1 // po_s2) + head_block * BLOCK_H ).to(tl.int32), (dim_block * BLOCK_V).to(tl.int32), ] @@ -5284,7 +5504,9 @@ def fp4_mla_paged_attention_internal( ) -> torch.Tensor: del kwargs if not hasattr(tl, "dot_scaled"): - raise NotImplementedError("fp4_mla_paged_attention requires a Triton build with tl.dot_scaled.") + raise NotImplementedError( + "fp4_mla_paged_attention requires a Triton build with tl.dot_scaled." + ) if not q_fp4.is_cuda: raise ValueError("q_fp4 must be a CUDA tensor.") if q_fp4.dtype != torch.uint8 or kv_cache.dtype != torch.uint8: @@ -5294,7 +5516,9 @@ def fp4_mla_paged_attention_internal( if q_fp4.dim() == 3: inferred_num_gen, inferred_num_heads, packed_q_dim = q_fp4.shape if num_heads is not None and num_heads != inferred_num_heads: - raise ValueError(f"num_heads={num_heads} does not match q_fp4.shape[1]={inferred_num_heads}.") + raise ValueError( + f"num_heads={num_heads} does not match q_fp4.shape[1]={inferred_num_heads}." + ) num_gen = inferred_num_gen num_heads = inferred_num_heads q_fp4_2d = q_fp4.reshape(num_gen * num_heads, packed_q_dim) @@ -5318,11 +5542,15 @@ def fp4_mla_paged_attention_internal( ) num_gen_seqs = num_gen // query_len_per_seq - num_pages, inferred_page_size, packed_k_dim, kv_s0, kv_s2, kv_s4 = _get_kv_cache_strides(kv_cache) + num_pages, inferred_page_size, packed_k_dim, kv_s0, kv_s2, kv_s4 = _get_kv_cache_strides( + kv_cache + ) if page_size is None: page_size = inferred_page_size if page_size != inferred_page_size: - raise ValueError(f"page_size={page_size} does not match kv_cache page dimension {inferred_page_size}.") + raise ValueError( + f"page_size={page_size} does not match kv_cache page dimension {inferred_page_size}." + ) if page_size % FP4_BLOCK_SIZE != 0: raise ValueError(f"page_size must be divisible by {FP4_BLOCK_SIZE}.") @@ -5351,9 +5579,13 @@ def fp4_mla_paged_attention_internal( raise ValueError("v_sf is too small for the swizzled V scale layout.") if output is None: - output = torch.empty((num_gen, num_heads, v_head_dim), dtype=output_dtype, device=q_fp4.device) + output = torch.empty( + (num_gen, num_heads, v_head_dim), dtype=output_dtype, device=q_fp4.device + ) elif output.shape != (num_gen, num_heads, v_head_dim): - raise ValueError(f"output must have shape {(num_gen, num_heads, v_head_dim)}, got {tuple(output.shape)}.") + raise ValueError( + f"output must have shape {(num_gen, num_heads, v_head_dim)}, got {tuple(output.shape)}." + ) if num_gen == 0: return output @@ -5404,9 +5636,7 @@ def fp4_mla_paged_attention_internal( env_duplicate_tail_k = _env_int("TRTLLM_FP4_MLA_DUPLICATE_TAIL_K") env_debug_page_stats_pack = _env_int("TRTLLM_FP4_MLA_DEBUG_PAGE_STATS_PACK") env_debug_stop_after_page_stats = _env_int("TRTLLM_FP4_MLA_DEBUG_STOP_AFTER_PAGE_STATS") - env_mtp_page_stats_pair_group_pages = _env_int( - "TRTLLM_FP4_MLA_MTP_PAGE_STATS_PAIR_GROUP_PAGES" - ) + env_mtp_page_stats_pair_group_pages = _env_int("TRTLLM_FP4_MLA_MTP_PAGE_STATS_PAIR_GROUP_PAGES") env_pv_tma = _env_int("TRTLLM_FP4_MLA_PV_TMA") env_pv_p_tma = _env_int("TRTLLM_FP4_MLA_PV_P_TMA") env_pv_v_tma = _env_int("TRTLLM_FP4_MLA_PV_V_TMA") @@ -5427,8 +5657,7 @@ def fp4_mla_paged_attention_internal( if max_pages is None: if paged_kv_indptr_decode.numel() >= num_gen_seqs + 1: page_counts = ( - paged_kv_indptr_decode[1 : num_gen_seqs + 1] - - paged_kv_indptr_decode[:num_gen_seqs] + paged_kv_indptr_decode[1 : num_gen_seqs + 1] - paged_kv_indptr_decode[:num_gen_seqs] ) max_pages = int(page_counts.max().item()) if page_counts.numel() > 0 else 0 else: @@ -5446,7 +5675,9 @@ def fp4_mla_paged_attention_internal( if pv_block_h <= 0: raise ValueError(f"TRTLLM_FP4_MLA_PV_BLOCK_H must be positive, got {pv_block_h}.") if pv_block_h not in (64, 128): - raise ValueError(f"TRTLLM_FP4_MLA_PV_BLOCK_H currently supports 64 or 128, got {pv_block_h}.") + raise ValueError( + f"TRTLLM_FP4_MLA_PV_BLOCK_H currently supports 64 or 128, got {pv_block_h}." + ) num_pv_head_blocks = triton.cdiv(num_heads, pv_block_h) assume_full_pv_heads = num_heads % pv_block_h == 0 assume_full_v = v_head_dim % block_v == 0 @@ -5455,7 +5686,9 @@ def fp4_mla_paged_attention_internal( assume_full_pages = bool(assume_full_pages) and query_len_per_seq == 1 mask_mtp_final_page_only = ( bool(assume_full_pages_except_mtp_tail) - and (env_mtp_final_page_fast_path != 0 if env_mtp_final_page_fast_path is not None else True) + and ( + env_mtp_final_page_fast_path != 0 if env_mtp_final_page_fast_path is not None else True + ) and query_len_per_seq > 1 and query_len_per_seq <= page_size ) @@ -5698,7 +5931,9 @@ def _debug_report() -> None: "TRTLLM_FP4_MLA_MTP_FUSED_QKPV_BLOCK_H currently supports 16, 32, 64, or 128, " f"got {mtp_fused_block_h}." ) - mtp_fused_block_v = env_mtp_fused_qkpv_block_v if env_mtp_fused_qkpv_block_v is not None else v_head_dim + mtp_fused_block_v = ( + env_mtp_fused_qkpv_block_v if env_mtp_fused_qkpv_block_v is not None else v_head_dim + ) if env_mtp_fused_qkpv == 1 and mtp_fused_block_v not in (128, 256, 512): raise ValueError( "TRTLLM_FP4_MLA_MTP_FUSED_QKPV_BLOCK_V currently supports 128, 256, or 512, " @@ -5709,7 +5944,9 @@ def _debug_report() -> None: f"TRTLLM_FP4_MLA_MTP_FUSED_QKPV_BLOCK_V={mtp_fused_block_v} must divide v_head_dim={v_head_dim}." ) num_mtp_fused_dim_blocks = triton.cdiv(v_head_dim, mtp_fused_block_v) - mtp_fused_group_pages = env_mtp_fused_qkpv_group_pages if env_mtp_fused_qkpv_group_pages is not None else 128 + mtp_fused_group_pages = ( + env_mtp_fused_qkpv_group_pages if env_mtp_fused_qkpv_group_pages is not None else 128 + ) mtp_fused_group_pages = max(1, min(int(mtp_fused_group_pages), max_pages)) mtp_fused_launch_meta = dict(page_stats_launch_meta) mtp_fused_reduce_meta = dict(pv_launch_meta) @@ -5767,7 +6004,11 @@ def _debug_report() -> None: name="mtp_fused_partial_l", ) _fp4_mla_attention_mtp_fused_qkpv_group_kernel[ - (num_gen, num_mtp_fused_head_blocks, num_mtp_fused_page_groups * num_mtp_fused_dim_blocks) + ( + num_gen, + num_mtp_fused_head_blocks, + num_mtp_fused_page_groups * num_mtp_fused_dim_blocks, + ) ]( mtp_fused_partial_o, mtp_fused_partial_m, @@ -6001,7 +6242,9 @@ def _debug_report() -> None: ) return output - online_qkpv_max_batch = env_online_qkpv_max_batch if env_online_qkpv_max_batch is not None else 32 + online_qkpv_max_batch = ( + env_online_qkpv_max_batch if env_online_qkpv_max_batch is not None else 32 + ) can_use_online_qkpv = ( env_online_qkpv == 1 and can_use_prepacked_v_for_pv @@ -6028,7 +6271,9 @@ def _debug_report() -> None: if v_repack_stream is not None: torch.cuda.current_stream(q_fp4.device).wait_stream(v_repack_stream) v_repack_stream = None - online_group_pages = env_online_qkpv_group_pages if env_online_qkpv_group_pages is not None else 128 + online_group_pages = ( + env_online_qkpv_group_pages if env_online_qkpv_group_pages is not None else 128 + ) online_group_pages = max(1, min(int(online_group_pages), max_pages)) num_online_page_groups = _ceil_div(max_pages, online_group_pages) online_partial_o = _workspace_tensor( @@ -6137,7 +6382,9 @@ def _debug_report() -> None: parallel_page_stats = triton_backend == "nvt" and max_pages >= 8 if pack_prob_in_page_stats is None: pack_prob_in_page_stats = parallel_page_stats and fused_prob_pack - pack_prob_in_page_stats = bool(pack_prob_in_page_stats and parallel_page_stats and fused_prob_pack) + pack_prob_in_page_stats = bool( + pack_prob_in_page_stats and parallel_page_stats and fused_prob_pack + ) debug_pack_prob_in_page_stats = ( pack_prob_in_page_stats if env_debug_page_stats_pack is None @@ -6175,7 +6422,9 @@ def _debug_report() -> None: # large TileIR kernels while still sharing K loads across query pairs. page_stats_group_size = min(page_stats_group_size, 2) grouped_storage_full_pages = assume_full_pages or (p_by_query and can_address_all_compact_pages) - grouped_assume_valid_pages = assume_valid_pages or (p_by_query and can_address_all_compact_pages) + grouped_assume_valid_pages = assume_valid_pages or ( + p_by_query and can_address_all_compact_pages + ) can_group_page_stats = ( page_stats_group_size in page_stats_group_sizes and triton_backend == "nvt" @@ -6197,12 +6446,15 @@ def _debug_report() -> None: and sf_per_page == 8 ) page_stats_group_size = page_stats_group_size if can_group_page_stats else 1 - num_page_groups = _ceil_div(max_pages, page_stats_group_size) if page_stats_group_size > 1 else max_pages + num_page_groups = ( + _ceil_div(max_pages, page_stats_group_size) if page_stats_group_size > 1 else max_pages + ) allow_partial_page_groups = page_stats_group_size > 1 and max_pages % page_stats_group_size != 0 group_reduce_stats = ( - (env_group_reduce_stats != 0 if env_group_reduce_stats is not None else triton_backend == "nvt") - and page_stats_group_size > 1 - ) + env_group_reduce_stats != 0 + if env_group_reduce_stats is not None + else triton_backend == "nvt" + ) and page_stats_group_size > 1 page_stats_max_entries = max_pages page_stats_num_groups = num_page_groups pv_apply_prob_scale = False @@ -6352,7 +6604,12 @@ def _debug_report() -> None: page_group_launches = ( (num_page_groups, assume_full_pages, allow_partial_page_groups, 0), ) - for grid_page_groups, launch_assume_full_pages, launch_allow_partial_groups, page_group_offset in page_group_launches: + for ( + grid_page_groups, + launch_assume_full_pages, + launch_allow_partial_groups, + page_group_offset, + ) in page_group_launches: if grid_page_groups <= 0: continue page_stats_grouped_kernel[(num_gen, num_head_blocks, grid_page_groups)]( @@ -6497,7 +6754,10 @@ def _debug_report() -> None: torch.cuda.current_stream(q_fp4.device).wait_stream(v_repack_stream) v_repack_stream = None group_pv_partial_dtype = torch.float32 - if os.environ.get("TRTLLM_FP4_MLA_GROUP_PV_PARTIAL_DTYPE", "").lower() in ("bf16", "bfloat16"): + if os.environ.get("TRTLLM_FP4_MLA_GROUP_PV_PARTIAL_DTYPE", "").lower() in ( + "bf16", + "bfloat16", + ): group_pv_partial_dtype = torch.bfloat16 group_partial_o = _workspace_tensor( None, @@ -6566,7 +6826,10 @@ def _debug_report() -> None: ) return output scale_from_group_stats = ( - (env_scale_from_group_stats == 1 or (env_scale_from_group_stats is None and num_gen >= 64)) + ( + env_scale_from_group_stats == 1 + or (env_scale_from_group_stats is None and num_gen >= 64) + ) and pack_prob_in_page_stats and group_reduce_stats and page_stats_group_size > 1 @@ -6737,9 +7000,7 @@ def _launch_prob_scale_from_group_stats( PV_LOOP_STAGES=int(pv_loop_stages), **pv_launch_meta, ) - _fp4_mla_attention_cast_acc_kernel[ - (num_gen, num_pv_head_blocks, num_dim_blocks) - ]( + _fp4_mla_attention_cast_acc_kernel[(num_gen, num_pv_head_blocks, num_dim_blocks)]( output, out_acc, output.stride(0), diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index 95ec6f65ec49..5190829261f4 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -32,7 +32,6 @@ from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal import thop from tensorrt_llm.functional import AttentionMaskType -from tensorrt_llm.llmapi import SkipSoftmaxAttentionConfig from tensorrt_llm.logger import logger from tensorrt_llm.models.modeling_utils import QuantConfig