Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -134,7 +134,12 @@ def run_msa_paged_gqa(
head_dim = attn.head_dim
kv_cache_manager = metadata.kv_cache_manager
num_tokens = int(q.shape[0])
if k is not None and v is not None:
# The fused per-layer scatter (msa_write_layer_caches) may have written
# this layer's K/V already; consume the marker so it never goes stale.
prewritten = getattr(metadata, "_msa_prewritten_layer", None) == layer_idx
if prewritten:
metadata._msa_prewritten_layer = None
if k is not None and v is not None and not prewritten:
write_msa_main_kv(
kv_cache_manager, layer_idx, metadata.msa_out_cache_loc[:num_tokens], k, v
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -254,6 +254,10 @@ class MiniMaxM3MsaSparseAttentionMetadata(TrtllmAttentionMetadata):
msa_kv_indices: Optional[torch.Tensor] = None
msa_max_score: Optional[torch.Tensor] = None
msa_n_valid_blocks: Optional[torch.Tensor] = None
# Layer whose K/V/index-K caches were already written this step by the
# fused scatter (msa_write_layer_caches); run_msa_paged_gqa consumes and
# clears it so the legacy per-cache writes are skipped exactly once.
_msa_prewritten_layer: Optional[int] = None

# _msa_buffers_ready gates the once-only device buffers;
# _msa_fields_ready marks that the current step's buffers are populated.
Expand Down Expand Up @@ -839,6 +843,9 @@ def _build_msa_fields(self) -> None:
buffers. The transient builder tensors are discarded.
"""
self._msa_fields_ready = False
# Drop any prewritten marker a failed prior step left unconsumed, so
# it can never suppress a later step's cache write.
self._msa_prewritten_layer = None
if not self._msa_buffers_ready:
return
request_ids = self.request_ids
Expand Down Expand Up @@ -920,6 +927,53 @@ def msa_write_idx_k(self, layer_idx: int, idx_k: torch.Tensor) -> None:
layout="HND",
)

def msa_write_layer_caches(
self,
layer_idx: int,
k: torch.Tensor,
v: torch.Tensor,
idx_k: Optional[torch.Tensor] = None,
) -> None:
"""Write a layer's new-token K, V, and (sparse layers) index-K.

One fused kernel launch when the source/cache layouts allow it, else
the legacy per-cache writes. Runs before the indexer's proxy pass
reads the index-K cache; the layer is recorded in
_msa_prewritten_layer so run_msa_paged_gqa skips its own K/V write.
Requires prepared metadata (msa_out_cache_loc filled), the same
contract as the writes it replaces.
"""
from .msa_scatter import fused_write_layer_caches

buffers = self.kv_cache_manager.get_buffers(layer_idx, kv_layout="HND")
k_view, v_view = buffers[:, 0], buffers[:, 1]
idx_cache = self.msa_idx_k_cache(layer_idx) if idx_k is not None else None
num_tokens = int(k.shape[0])
out_cache_loc = self.msa_out_cache_loc[:num_tokens]
if not fused_write_layer_caches(k_view, v_view, idx_cache, out_cache_loc, k, v, idx_k):
num_kv_heads = int(k_view.shape[1])
head_dim = int(k_view.shape[3])
write_kv_slots(
k_view,
out_cache_loc,
k.reshape(num_tokens, num_kv_heads, head_dim),
layout="HND",
)
write_kv_slots(
v_view,
out_cache_loc,
v.reshape(num_tokens, num_kv_heads, head_dim),
layout="HND",
)
if idx_k is not None:
write_kv_slots(
idx_cache,
out_cache_loc,
idx_k.reshape(num_tokens, 1, int(idx_cache.shape[-1])),
layout="HND",
)
self._msa_prewritten_layer = layer_idx

def msa_proxy_max_score_view(
self, num_index_heads: int, plan_max_k_tiles: int, num_tokens: int
) -> torch.Tensor:
Expand Down Expand Up @@ -1014,13 +1068,16 @@ def run_indexer(
metadata,
*,
idx_sm_scale: Optional[float] = None,
idx_k_prewritten: bool = False,
) -> torch.Tensor:
"""Write the index-K cache and return the selected block indices.

The model layer runs this before forward and threads the result through
forward_args.topk_indices. Returns [total_q, num_kv_heads, topk].
Decode uses the prebuilt graph-safe proxy plan; prefill and mixed
batches use the prebuilt eager proxy plan.
batches use the prebuilt eager proxy plan. `idx_k_prewritten` marks
that the fused per-layer cache write (msa_write_layer_caches) already
stored this layer's index-K.
"""
config = self.m3_config
idx_sm_scale = idx_sm_scale if idx_sm_scale is not None else config.sparse_index_dim**-0.5
Expand All @@ -1030,10 +1087,17 @@ def run_indexer(
# scatter below both honor the source strides.
idx_q_view = idx_q.reshape(num_tokens, config.num_index_heads, config.sparse_index_dim)
idx_k_cache = metadata.msa_idx_k_cache(self.layer_idx)
if idx_k is not None:
# Index-K may already be in the cache by two routes: the fused per-layer
# write (msa_write_layer_caches, idx_k_prewritten=True) stored a live
# bf16 idx_k, or the FP8 producer inserted FP8 index-K and passed
# idx_k=None. Write here only when neither owns it — i.e. a live idx_k
# that was not pre-written.
if idx_k is not None and not idx_k_prewritten:
idx_k_view = idx_k.reshape(num_tokens, 1, config.sparse_index_dim)
metadata.msa_write_idx_k(self.layer_idx, idx_k_view)
elif idx_k_cache.dtype != torch.float8_e4m3fn or idx_q_view.dtype != torch.float8_e4m3fn:
elif idx_k is None and (
idx_k_cache.dtype != torch.float8_e4m3fn or idx_q_view.dtype != torch.float8_e4m3fn
):
raise ValueError(
"A missing live index-K is valid only when the fused MiniMax-M3 "
"producer already emitted FP8 index-Q and inserted FP8 index-K."
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,163 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
"""Fused paged-cache scatter for the MiniMax-M3 MSA backend.

One Triton launch writes a layer's new-token main K, main V, and (sparse
layers) index-K into their paged HND caches at the step's write slots.
The legacy path costs three aten advanced-indexing writes per layer plus
their index preprocessing; at 60 layers per forward step, all captured
into decode CUDA graphs, the launch count dominates the cost. The kernel
derives each token's (page, within-page) split from ``out_cache_loc``
in-register, so it needs no precomputed index tensors at all.

Sources may be strided row views (slices of the fused QKV projection);
only the innermost [num_heads * head_dim] extent must be contiguous.
Stores cast to the cache dtype, which folds the FP8 KV-cache cast in.
"""

from __future__ import annotations

from typing import Optional

import torch
import triton
import triton.language as tl


@triton.jit
def _fused_paged_scatter_kernel(
k_src,
v_src,
idx_src,
k_cache,
v_cache,
idx_cache,
out_cache_loc,
k_src_row_stride,
v_src_row_stride,
idx_src_row_stride,
kc_stride_page,
kc_stride_head,
kc_stride_tok,
vc_stride_page,
vc_stride_head,
vc_stride_tok,
ic_stride_page,
ic_stride_tok,
tokens_per_block,
H: tl.constexpr,
D: tl.constexpr,
HAS_IDX: tl.constexpr,
):
# int64 throughout: t * row_stride can exceed 2^31 elements on large
# eager prefill steps (num_tokens up to max_num_tokens times the fused
# QKV row stride), and the slot * page-stride products likewise.
t = tl.program_id(0).to(tl.int64)
slot = tl.load(out_cache_loc + t).to(tl.int64)
page = slot // tokens_per_block
within = slot % tokens_per_block
d = tl.arange(0, D)
for h in tl.static_range(H):
k_vals = tl.load(k_src + t * k_src_row_stride + h * D + d)
v_vals = tl.load(v_src + t * v_src_row_stride + h * D + d)
k_dst = k_cache + page * kc_stride_page + h * kc_stride_head + within * kc_stride_tok + d
v_dst = v_cache + page * vc_stride_page + h * vc_stride_head + within * vc_stride_tok + d
tl.store(k_dst, k_vals.to(k_cache.dtype.element_ty))
tl.store(v_dst, v_vals.to(v_cache.dtype.element_ty))
if HAS_IDX:
i_vals = tl.load(idx_src + t * idx_src_row_stride + d)
i_dst = idx_cache + page * ic_stride_page + within * ic_stride_tok + d
tl.store(i_dst, i_vals.to(idx_cache.dtype.element_ty))


def _row_stride_if_fusable(src: torch.Tensor, inner: int) -> Optional[int]:
"""Row stride (elements) if `src` is a [T, inner] row view with contiguous
rows (e.g. a column slice of the fused QKV projection); None otherwise."""
if src.dim() != 2 or src.shape[1] != inner or src.stride(1) != 1:
return None
return src.stride(0)


def fused_write_layer_caches(
k_cache: torch.Tensor,
v_cache: torch.Tensor,
idx_cache: Optional[torch.Tensor],
out_cache_loc: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
idx_k: Optional[torch.Tensor],
) -> bool:
"""Fused single-launch write of new-token K/V (+index-K) into paged HND
caches. Returns False when a layout precondition fails, so the caller can
keep the legacy per-cache writes.

`k_cache`/`v_cache` are [num_pages, num_kv_heads, tokens_per_block,
head_dim] HND views; `idx_cache` is the MQA index-K view with head dim 1.
`k`/`v` are the layer's new-token values, [T, H*D] or [T, H, D] row views;
`idx_k` is [T, D] or [T, 1, D].
"""
if not (k.is_cuda and k_cache.is_cuda):
return False
if k_cache.dim() != 4 or v_cache.dim() != 4:
return False
if k_cache.stride(-1) != 1 or v_cache.stride(-1) != 1:
return False
num_pages, num_heads, tokens_per_block, head_dim = k_cache.shape
if (head_dim & (head_dim - 1)) != 0:
return False
inner = num_heads * head_dim
k_stride = _row_stride_if_fusable(k, inner)
v_stride = _row_stride_if_fusable(v, inner)
if k_stride is None or v_stride is None:
return False

has_idx = idx_k is not None
idx_stride = 0
ic_stride_page = 0
ic_stride_tok = 0
if has_idx:
if idx_cache is None or idx_cache.dim() != 4 or idx_cache.stride(-1) != 1:
return False
if int(idx_cache.shape[1]) != 1 or int(idx_cache.shape[3]) != head_dim:
return False
if int(idx_cache.shape[2]) != tokens_per_block:
return False
idx_stride = _row_stride_if_fusable(idx_k, head_dim)
if idx_stride is None:
return False
ic_stride_page = idx_cache.stride(0)
ic_stride_tok = idx_cache.stride(2)

num_tokens = int(out_cache_loc.shape[0])
if num_tokens == 0:
return True

_fused_paged_scatter_kernel[(num_tokens,)](
k,
v,
idx_k if has_idx else k, # unused when HAS_IDX=False
k_cache,
v_cache,
idx_cache if has_idx else k_cache, # unused when HAS_IDX=False
out_cache_loc,
k_stride,
v_stride,
idx_stride,
k_cache.stride(0),
k_cache.stride(1),
k_cache.stride(2),
v_cache.stride(0),
v_cache.stride(1),
v_cache.stride(2),
ic_stride_page,
ic_stride_tok,
tokens_per_block,
H=num_heads,
D=head_dim,
HAS_IDX=has_idx,
num_warps=2,
)
return True


__all__ = ["fused_write_layer_caches"]
15 changes: 14 additions & 1 deletion tensorrt_llm/_torch/models/modeling_minimaxm3.py
Original file line number Diff line number Diff line change
Expand Up @@ -1489,11 +1489,24 @@ def _msa_attention_core(
"""
if self.is_sparse_attention_layer:
assert idx_q is not None
# One launch writes this layer's K/V and, on the bf16 path, index-K,
# ahead of the proxy pass that reads the index-K cache. On the FP8
# indexer path idx_k is None because the fused producer already
# inserted FP8 index-K into the side cache; the fused write then
# stores K/V only.
attn_metadata.msa_write_layer_caches(self.attn.layer_idx, k, v, idx_k)
# Publish the selected blocks so the FMHA runs the sparse path.
kv_block_indexes = self.attn.run_indexer(idx_q, idx_k, attn_metadata)
# idx_k_prewritten marks that index-K is already in the cache (via
# the fused write above on bf16, or the FP8 producer when idx_k is
# None), so run_indexer must not write it again.
kv_block_indexes = self.attn.run_indexer(
idx_q, idx_k, attn_metadata, idx_k_prewritten=True
)
forward_args = AttentionForwardArgs(output=output, topk_indices=kv_block_indexes)
else:
assert idx_q is None and idx_k is None
# Dense layers get the same fused K/V write.
attn_metadata.msa_write_layer_caches(self.attn.layer_idx, k, v)
# No top-k selection means the FMHA attends the full page table.
forward_args = AttentionForwardArgs(output=output)
self.attn.forward(q, k, v, attn_metadata, forward_args=forward_args)
Expand Down
Loading
Loading