diff --git a/bench_fp4_mla_decode.py b/bench_fp4_mla_decode.py new file mode 100644 index 000000000000..99c18642bc79 --- /dev/null +++ b/bench_fp4_mla_decode.py @@ -0,0 +1,283 @@ +# 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 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] [--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 +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 import _build_fp4_mla_attention_decode_case # 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_fp8", + "triton", + "cutile", +) + + +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 _seq_lens_for_batch(batch, seq): + return [seq] * 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 = 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 * 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 _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 * 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, 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 + (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. + 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, 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) + # 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, 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, + 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=max_seq, + out=output, + bmm1_scale=0.1, + bmm2_scale=1.0, + backend="trtllm-gen", + ) + + run() + torch.cuda.synchronize() + 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) + tflops = flops / avg_ms / 1e9 + 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={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, + ) + return avg_ms + + +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) + total_seq = sum(seq_lens) + ( + 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=heads, seed=8, query_len_per_seq=q_len + ) + 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, 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) + 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(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_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, + ) + return avg_ms + finally: + kv_cache_manager.shutdown() + + +def main(): + p = argparse.ArgumentParser() + 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( + "--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, + choices=BACKEND_CHOICES, + help=("Backend to benchmark; default runs all fast backends."), + ) + 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] + if args.backend: + backends = [args.backend] + else: + backends = list(BACKEND_CHOICES) + + for b in batches: + for be in backends: + if be == "trtllm_fp8": + 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, args.warmup, args.iters) + + +if __name__ == "__main__": + main() 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/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/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..65515b6bf346 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -11,6 +11,7 @@ from typing_extensions import Self 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 @@ -31,8 +32,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.""" @@ -125,6 +124,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) @@ -156,6 +170,8 @@ 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) + # True during warmup forward passes (dummy requests, no real data). + is_warmup: bool = field(init=False, default=False) def needs_plan(self, plan_params: PlanParams) -> bool: if plan_params not in self._plan_params_to_wrappers: @@ -364,9 +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] - kv_indices = self._paged_kv_indices[self.num_context_blocks:self. - num_context_blocks + - self.num_generation_blocks] + 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] @@ -1240,12 +1257,21 @@ def __init__( self.qk_nope_head_dim = mla_params.qk_nope_head_dim self.v_head_dim = mla_params.v_head_dim + if getattr(self, "has_fp4_kv_cache", False): + raise NotImplementedError( + "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 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 +1384,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", @@ -1469,32 +1502,32 @@ 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. + # 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. - 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 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) 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 new file mode 100644 index 000000000000..3bff8b7d9540 --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla.py @@ -0,0 +1,2800 @@ +# 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 the TRTLLM attention backend FP4 MLA FMHA path. +""" + +import os +from typing import Any, Literal, Optional + +import torch +import triton +import triton.language as tl + +from .fp4_mla_kernels import ( + _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 * 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 +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" + + +# Environment helpers + + +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 == "": + return None + return int(value) + + +def _fp4_mla_attention_backend() -> str: + return os.getenv(FP4_MLA_ATTENTION_BACKEND_ENV, "triton").lower() + + +def _cutile_backend_available() -> bool: + return hasattr(tl, "ext") + + +def _ceil_div(lhs: int, rhs: int) -> int: + return (lhs + rhs - 1) // rhs + + +_SM_COUNT_CACHE: dict[int, int] = {} + + +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 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 + 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 _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, + ) + + +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, + 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 needs at least {num_gen} generation " + f"tokens, got {num_tokens}." + ) + if num_tokens % num_gen != 0: + raise NotImplementedError( + "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: + 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 = 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], + kv_lens_gen, + gen_lens_gen, + 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, + ) + + +# 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. + + 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 + 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)." + ) + + _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 + + 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 + + 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, + ) + v_pack_page_ids = metadata.paged_kv_indices + 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, + ) + 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, + 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], + ) + _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], + ) + + +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 _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 + if not _cutile_backend_available(): + return False + return os.getenv("TRTLLM_FP4_MLA_PERSISTENT_V_PACK", "1").lower() not in ( + "0", + "false", + "no", + "off", + ) + + +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 _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" + 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_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]: + return _v_packed_shape(kv_cache, v_head_dim, page_size, block_v) + + +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, + kv_cache: torch.Tensor, + page_ids: torch.Tensor, + *, + 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 + 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 + + 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, block_v), + 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=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, + ) + + +def _get_cutile_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 _cutile_persistent_v_pack_enabled(): + return None + 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, 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]] + + +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) + + +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", + ) + + +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 + ) + + +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 _triton_v_packed_valid_attr(layer_idx: int) -> str: + return f"_fp4_mla_triton_attention_v_packed_valid_l{layer_idx}" + + +def _triton_shared_v_packed_valid_attr() -> str: + return "_fp4_mla_triton_attention_v_packed_valid_tag" + + +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, + 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 = ( + _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, + ), + ) + + +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, + ) + + +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]] + + +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 + + 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, + attr_name, + _v_packed_shape(kv_cache, v_head_dim, page_size, block_v), + 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=block_v, + ) + _set_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 v_packed + + +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, + page_ids, + v_head_dim=v_head_dim, + page_size=page_size, + block_v=block_v, + local_layer=local_layer, + v_sf=v_sf, + ) + + +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 + + 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. + + 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( + 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) + 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( + *, + 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], + 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), + 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, + GROUP_PAGES=ps_group_pages, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + PAGE_LOOP_STAGES=ps_loop_stages, + **page_stats_launch_meta, + ) + else: + _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, + 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, + 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, + **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)) + # 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, + 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, + ) + 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 + # 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( + 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. + + 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. + """ + 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}." + ) + 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, + 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 + ] + # 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 + + 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 + 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, + ) + 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_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 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 in (128, 256) + and metadata.page_size // FP4_BLOCK_SIZE == 8 + ) + 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, + 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 + ) + 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, + ) + 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_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, + 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, + 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, + ) + 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, + ) + 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 {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. + _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( + 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``, + ``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.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") + + 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) + + # 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, 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, + ) + + # 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 + + # 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 + _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, + ) + _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, + ) 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..c518c312668b --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_cutile.py @@ -0,0 +1,7446 @@ +# 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_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_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 = 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, + 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 + 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) + 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), 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, + 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), + token_start.to(tl.int32), + (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, + P_BY_QUERY: 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 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) + 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( + 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, + P_BY_QUERY: 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, +): + 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 = 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) + + 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 = 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 + 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]) + 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 + 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) + 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 = 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, + 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( + 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 = ( + 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_page_stats_grouped_mtp_pair_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, + 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, +): + seq_idx = tl.program_id(0) + head_block = tl.program_id(1) + 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) + offs_t = tl.arange(0, BLOCK_T) + scale_cols = tl.arange(0, 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: + 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) + + 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, + 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: 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, + 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) + 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) + 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(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], + ) + + 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_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, + GROUP_PAGES: tl.constexpr, + NUM_PAGE_GROUPS: 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) + 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) + + 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, + ) + + 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_half_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, + num_pages, + 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, + 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) + tile_group = tl.program_id(2) + + offs_h = head_block * BLOCK_H + tl.arange(0, BLOCK_H) + local_t = tl.arange(0, BLOCK_T) + local_scale_cols = tl.arange(0, BLOCK_T // FP4_BLOCK) + 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(p_s0 % 8 == 0) + tl.assume(p_s1 == 1) + + 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) + + 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_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, + P_BY_QUERY: 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 + ) + + 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_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], + ) + 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_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, + 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) + + 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, + ) + 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, + ) + 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_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 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, + ) + 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, + ) + else: + 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, + ) + + 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 + ) + acc += partial_o * scale[:, None] + global_l += group_l * scale + + global_scale = tl.load(global_scale_ptr) + 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, + ) + + +@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 + + 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", + 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, + ) + + 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_group_partial_reduce_kernel( + out_ptr, + partial_o_ptr, + page_sum_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, + po_num_rows: tl.constexpr, + page_stats_s0: tl.constexpr, + page_stats_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, + 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) + 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], + ) + + 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 + + global_scale = tl.load(global_scale_ptr) + out_scale = 1.0 / (global_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, + ) + + +@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", + 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, + ) + + 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( + 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_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, + 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_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.") + + 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 + 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" + 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: + 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 + 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) 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_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 + 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: + 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 = {} + 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 + 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) + 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 + 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) + 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: + 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) + 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", + ) + 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", + ) + denom = _workspace_tensor( + denom_workspace, + (num_gen, num_heads), + dtype=torch.float32, + device=q_fp4.device, + name="denom", + ) + 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 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_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 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_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_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, + ) + _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, + ) + 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 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 + 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 + ), + 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 or env_debug_stop_after_page_stats == 1) + 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_h == 128 + and block_k == 512 + and full_block_end == 512 + and tail_block_k == 128 + 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 + ) + 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, page_stats_max_entries, 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", + ) + 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, + 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, + 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, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **page_stats_launch_meta, + ) + _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, + v_sf, + v_packed, + src_page_ids, + paged_kv_indptr_decode, + num_pages, + 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], + 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, + 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_V=block_v, + **pv_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 not scale_from_group_stats: + _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=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, + 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, + 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, + 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, + ASSUME_FULL_HEADS=assume_full_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_VALID_PAGES=assume_valid_pages, + **page_stats_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, + 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, + 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, + **page_stats_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, + 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, + **page_stats_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, + 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, + **page_stats_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, + 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, + 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, + **page_stats_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) + + 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_pv_head_blocks, + num_dim_blocks, + ) + ]( + output, + p_fp4, + p_sf, + kv_cache, + 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, + 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], + 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, + 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_BY_QUERY=p_by_query, + P_GLOBAL_SCALE=p_global_scale, + BLOCK_H=pv_block_h, + BLOCK_V=block_v, + 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), + 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, + **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_pv_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, + 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=pv_block_h, + BLOCK_V=block_v, + 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_pv_heads, + ASSUME_FULL_PAGES=assume_full_pages, + ASSUME_FULL_V=assume_full_v, + ASSUME_VALID_PAGES=assume_valid_pages, + **pv_launch_meta, + ) + _debug_mark("pv") + _debug_report() + 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 new file mode 100644 index 000000000000..3ffcc79437b7 --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_kernels.py @@ -0,0 +1,1359 @@ +# 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_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, + 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 generation tokens into the HP KV pool. + + Grid: (num_generation_tokens,). + Each program stores one token into the circular buffer position + ``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_token_idx = tl.program_id(0) + if (layer_idx < 0) | (layer_idx >= num_layers): + return + if gen_token_idx >= num_tokens: + return + + metadata_token_idx = token_offset + gen_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) | (position < 0): + return + + 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) + 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 _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, + 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) + ) + + +# 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) + 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. + 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) + # 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]) + 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_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_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, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_T: tl.constexpr, + BLOCK_K: tl.constexpr, +): + 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 = 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 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 + 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_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, + QUERY_LEN_PER_SEQ: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_K: tl.constexpr, +): + 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 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 + 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, + 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 + 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.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 = 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, +): + 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 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) + 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_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, + QUERY_LEN_PER_SEQ: tl.constexpr, + MAX_PAGES: tl.constexpr, + P_GLOBAL_SCALE: tl.constexpr, + BLOCK_H: tl.constexpr, + BLOCK_V: tl.constexpr, +): + 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) + 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 + 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): + 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_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], + 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 + + 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_triton.py b/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py new file mode 100644 index 000000000000..c29a00913c8b --- /dev/null +++ b/tensorrt_llm/_torch/attention_backend/fp4_mla_triton.py @@ -0,0 +1,2213 @@ +# 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 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): + +* ``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. +* 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 + +_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_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, + 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, + 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, +): + """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) + # 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) + 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) + # 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), 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_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, + q_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, + GROUP_PAGES: tl.constexpr, + ASSUME_FULL_PAGES: tl.constexpr, + ASSUME_VALID_PAGES: tl.constexpr, + PAGE_LOOP_STAGES: tl.constexpr, + occupancy: tl.constexpr = 1, +): + """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_group = tl.program_id(2) + seq_idx = query_idx // QUERY_LEN_PER_SEQ + query_offset = query_idx - seq_idx * QUERY_LEN_PER_SEQ + + 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) + 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 + + # ---- 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 + ) + + 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) + 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_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, + 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, + 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, + 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, + 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_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, + 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) + 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) + 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_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], + ) + 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: + 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_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) + 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, + ) + 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: + 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) + 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( + 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: + base = ( + query_idx * partial_s0 + + split_idx * partial_s1 + + safe_offs_h[:, None] * partial_s2 + + safe_offs_v[None, :] * partial_s3 + ) + 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 = 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: + 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, + out_vals * 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, :], + ) diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index 93e1a2bfe4c4..5190829261f4 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -29,12 +29,15 @@ 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.logger import logger from tensorrt_llm.models.modeling_utils import QuantConfig from ..utils import (compute_swizzled_sf_shape, get_global_attrs, get_model_extra_attrs) +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, @@ -146,6 +149,31 @@ 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 + 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. flash_mla_tile_scheduler_metadata: Optional[torch.Tensor] = None @@ -222,6 +250,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]: """ @@ -419,6 +474,102 @@ 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(), + ) + 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, + ) + 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 + 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)}, " + 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( @@ -458,6 +609,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, @@ -585,6 +745,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() @@ -1352,6 +1612,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) @@ -1619,6 +1886,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 @@ -1902,6 +2172,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 ] @@ -1946,3 +2220,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/modules/fla/utils.py b/tensorrt_llm/_torch/modules/fla/utils.py index 480051cbcc7d..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,7 +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 = 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 ff1e5f8a5306..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): @@ -2905,6 +2909,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 +2980,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 +3153,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 +3233,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 +3306,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 +3764,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 +3930,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 +3940,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 +4086,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 +4094,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 +4207,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) @@ -4216,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]) @@ -4251,6 +4279,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 diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index 797f2fd48666..d0497a47be96 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -44,6 +44,9 @@ 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 + class _ExecutorMemoryMonitor: """Currently this focuses on tracking memory usage and related errors.""" @@ -181,6 +184,50 @@ 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(config=None) -> bool: + return getattr(config, "attn_backend", None) == "TRTLLM" + + +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(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" + ) + 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(), @@ -406,7 +453,20 @@ def create_py_executor( ) llm_args.disable_overlap_scheduler = True - # Check FLASHINFER compatibility with one-engine speculative decoding + 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. if llm_args.attn_backend == "FLASHINFER": raise ValueError( f"FLASHINFER attention backend is not supported with one-engine speculative " @@ -647,24 +707,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..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 @@ -60,6 +84,15 @@ int]] # window_size -> (blocks_in_primary_pool, blocks_in_secondary_pool) +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 class PoolConfiguration: """Configuration of a single KV pool. @@ -570,6 +603,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 +1178,24 @@ 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) + + 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 +1452,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 +1476,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 +1697,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 @@ -2126,6 +2261,1438 @@ def reset_reuse_state(self): 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 + + 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: + 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: def __init__(self, max_num_requests: int): 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 dda345844c19..4a2fdfbe6679 100644 --- a/tensorrt_llm/_torch/speculative/mtp.py +++ b/tensorrt_llm/_torch/speculative/mtp.py @@ -43,6 +43,19 @@ 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_hp_kv_for_mtp_rejection(attn_metadata, num_accepted_tokens) + + class MTPHiddenStatesManager(BaseResourceManager): def __init__(self, @@ -675,7 +688,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 +702,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) @@ -831,7 +844,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 +862,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 +883,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 +913,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): @@ -1070,7 +1088,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 +1107,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/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/_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) 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_flashinfer_attention.py b/tests/unittest/_torch/attention/test_flashinfer_attention.py index 382e11f591ae..08fe40dc190c 100644 --- a/tests/unittest/_torch/attention/test_flashinfer_attention.py +++ b/tests/unittest/_torch/attention/test_flashinfer_attention.py @@ -666,3 +666,32 @@ 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. + + 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_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("TRTLLM", msg) diff --git a/tests/unittest/_torch/attention/test_fp4_mla.py b/tests/unittest/_torch/attention/test_fp4_mla.py new file mode 100644 index 000000000000..e6be4e4dcaee --- /dev/null +++ b/tests/unittest/_torch/attention/test_fp4_mla.py @@ -0,0 +1,1564 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""TRTLLM FP4 MLA helper tests.""" + +import os +from types import SimpleNamespace + +import pytest +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 ( + 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, + _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, + 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.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 ( + 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 _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, + *, + 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 * FP4_BLOCK_SIZE + packed = fp4_bytes[row_idx, start // 2 : start // 2 + 8] + low = packed & 0x0F + high = (packed >> 4) & 0x0F + 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), + 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 + 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 // FP4_BLOCK_SIZE, FP4_BLOCK_SIZE + ) + duplicated_tail = tail.repeat_interleave(2, dim=-2).reshape( + *tensor.shape[:-1], + residual_dim * 2, + ) + return torch.cat((prefix, duplicated_tail), dim=-1) + + +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) + 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, + 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) + 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, + 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 _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") + + 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 + 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(num_queries, num_heads, kv_lora_rank, dtype=torch.bfloat16, device=device) + * 0.25 + ).clamp_(-1.0, 1.0) + q_pe = ( + torch.randn( + num_queries, + 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 + 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) + 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 // FP4_BLOCK_SIZE, + global_scale=_TEST_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 // FP4_BLOCK_SIZE, + global_scale=FP4_MLA_P_GLOBAL_SCALE, + ) + + 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(num_seqs): + kv_len = kv_lens[seq_idx] + 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(p, cache[:, :kv_lora_rank].float())) + return torch.stack(outputs, dim=0), exact_probs, quantized_probs + + +def _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + *, + backend: str, + num_heads: int, + seq_lens: list[int], + seed: int, + check_probs: bool = False, + query_len_per_seq: int = 1, +) -> None: + _reset_triton_allocator() + monkeypatch.setenv(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, + query_len_per_seq=query_len_per_seq, + ) + 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=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() + + +@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 + + 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) + + 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 // FP4_BLOCK_SIZE) + + +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(_is_pre_blackwell(), reason="requires Blackwell FP4 support") +@pytest.mark.parametrize( + "num_tokens", + [32, 129, 144], + ids=["aligned32", "tail1", "aligned144"], +) +def test_fp4_mla_scatter_gather_roundtrip(num_tokens: int): + torch.manual_seed(0) + 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 + max_seq_len = ((num_tokens + page_size - 1) // page_size) * page_size + + 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], [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 + ) + latent = ( + torch.randn(num_tokens, 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, + ) + update_hp_kv_for_fp4_mla(metadata, latent, local_layer=0, phase="context") + torch.cuda.synchronize() + + recovered = _materialize_single_seq_cache_with_hp_tail( + metadata, + layer_idx=0, + local_layer=0, + head_dim=head_dim, + num_tokens=num_tokens, + ) + tail = num_tokens % HP_BLOCK_SIZE + fp4_end = num_tokens - tail + if fp4_end > 0: + torch.testing.assert_close( + recovered[:fp4_end].float(), latent[:fp4_end].float(), atol=1.0, rtol=0.5 + ) + if tail > 0: + torch.testing.assert_close( + recovered[fp4_end:].float(), latent[fp4_end:].float(), atol=0, rtol=0 + ) + finally: + kv_cache_manager.shutdown() + + +@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 + 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=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=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_() + 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) * 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 + ) + + metadata.kv_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.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, + 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") + + 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 + 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() + + recovered = _materialize_single_seq_cache_with_hp_tail( + metadata, + layer_idx=0, + local_layer=0, + head_dim=head_dim, + num_tokens=total_tokens, + ) + torch.testing.assert_close( + recovered[ctx_tokens].float(), gen_latent[0].float(), atol=0, rtol=0 + ) + finally: + 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() + metadata = SimpleNamespace( + high_precision_kv_pool=hp_pool, + 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), + num_contexts=0, + num_seqs=1, + 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( + 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)] + 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(_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_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 + ) + 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 // 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] + .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) + + v_offsets = torch.tensor( + [ + _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) + + 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 + finally: + kv_cache_manager.shutdown() + + +@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 + 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 + ) + + 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, + ) + metadata.num_contexts = 0 + 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) + .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 // FP4_BLOCK_SIZE, + global_scale=_TEST_GLOBAL_SCALE, + ) + 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) + 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=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_matches_reference(monkeypatch): + _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + backend="triton", + num_heads=5, + seq_lens=[32, 128], + seed=7, + check_probs=True, + ) + + +@pytest.mark.skipif(_is_pre_blackwell(), reason="requires Blackwell FP4 tensor cores") +def test_fp4_mla_attention_decode_linear_mtp_matches_reference(monkeypatch): + _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 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): + _assert_fp4_mla_attention_decode_accuracy( + monkeypatch, + backend="cutile", + num_heads=128, + seq_lens=[32, 128], + seed=7, + check_probs=False, + ) + + +@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): + 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_cutile_unavailable(), + reason="requires Blackwell FP4 tensor cores and Triton tl.ext", +) +def test_fp4_mla_attention_decode_cutile_grouped_tail_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=19, + check_probs=False, + ) + + +@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): + 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") + + 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 + ) + + +def _ceil_div(lhs: int, rhs: int) -> int: + return (lhs + rhs - 1) // rhs + + +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", +) +@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): + backend = os.environ.get(FP4_MLA_ATTENTION_BACKEND_ENV, "triton") + monkeypatch.setenv(FP4_MLA_ATTENTION_BACKEND_ENV, backend) + 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 + 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"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"estimated MBU={mbu_tbps:.2f} TB/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(): + 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=new_tokens, page_size=page_size, num_layers=num_layers + ) + metadata.positions = torch.arange( + cached_tokens, total_tokens, dtype=torch.int32, device=device + ) + 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..6a586ddba205 --- /dev/null +++ b/tests/unittest/_torch/executor/test_mla_tokens_per_block.py @@ -0,0 +1,100 @@ +# 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, + 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, 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, attn_backend=attn_backend + ) + + +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(): + 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(): + 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, attn_backend="TRTLLM"), + 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(): + 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, attn_backend="TRTLLM"), + 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