-
Notifications
You must be signed in to change notification settings - Fork 2.7k
[None][perf] Fuse MiniMax-M3 MSA per-layer KV-cache writes into one kernel #18614
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -231,6 +231,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 | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I have some thoughts about code design here, please feedback if you have any concern or other idea, thanks. I think _msa_prewritten_layer shouldn't be in attention metadata, rather it should be in the modeling code. Is there a case where at layer x we use fusion kernel and layer x+1 not? I would prefer to make it a modeling level constant and kept the same during execution, if possible. |
||
|
|
||
| # _msa_buffers_ready gates the once-only device buffers; | ||
| # _msa_fields_ready marks that the current step's buffers are populated. | ||
|
|
@@ -660,6 +664,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 | ||
|
|
@@ -733,6 +740,53 @@ def msa_write_idx_k(self, layer_idx: int, idx_k: torch.Tensor) -> None: | |
| layout="HND", | ||
| ) | ||
|
|
||
| def msa_write_layer_caches( | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Also, I think msa_write_layer_caches also seems shouldn't be a member of metadata, as this is more about modeling and not metadata. It can serve as a individual function. (I do see write_k_indices is in metadata now, maybe it is a common practice to do so.) |
||
| 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: | ||
|
|
@@ -827,13 +881,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.sparse_backend_args. 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 | ||
|
|
@@ -866,13 +923,18 @@ def run_indexer( | |
| "The MiniMax-M3 BF16 indexer requires BF16 index-Q and a live " | ||
| f"BF16 index-K tensor; got Q={idx_q_view.dtype}, K={live_k_dtype}." | ||
| ) | ||
| idx_k_view = idx_k.view(num_tokens, 1, config.sparse_index_dim) | ||
| metadata.msa_write_idx_k(self.layer_idx, idx_k_view) | ||
| # The fused per-layer write (msa_write_layer_caches, signalled by | ||
| # idx_k_prewritten) may already have stored this live bf16 index-K | ||
| # ahead of the proxy pass; write it here only when it did not. | ||
| if not idx_k_prewritten: | ||
| idx_k_view = idx_k.view(num_tokens, 1, config.sparse_index_dim) | ||
| metadata.msa_write_idx_k(self.layer_idx, idx_k_view) | ||
| # The FP8 indexer mirrors vLLM's unscaled E4M3 contract: normalized | ||
| # index Q/K are cast directly and the proxy accumulates their QK scores | ||
| # in FP32. Block ordering is invariant to the omitted positive scale. | ||
| # The fused production path arrives here with E4M3 Q and an already | ||
| # populated cache; the BF16 path writes its live K above. | ||
| # populated cache; the BF16 path writes its live K above unless the | ||
| # fused per-layer write already did. | ||
|
|
||
| # One selection path. Decode passes the graph-safe proxy plan plus the | ||
| # proxy scratch shaped to the live query count. Prefill and mixed batches | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,173 @@ | ||
| # 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 or device 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 one head. | ||
| `k`/`v` are the layer's new-token values as [T, H*D] row views; | ||
| `idx_k` is [T, D]. Their inner dimension must be contiguous. | ||
| """ | ||
| if not k_cache.is_cuda or any( | ||
| tensor.device != k_cache.device for tensor in (k, v, v_cache, out_cache_loc) | ||
| ): | ||
| return False | ||
| if k_cache.dim() != 4 or v_cache.dim() != 4: | ||
| return False | ||
| if v_cache.shape != k_cache.shape: | ||
| 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 idx_k.device != k_cache.device or idx_cache.device != k_cache.device: | ||
| 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 | ||
| if k.shape[0] < num_tokens or v.shape[0] < num_tokens: | ||
| return False | ||
| if has_idx and idx_k.shape[0] < num_tokens: | ||
| return False | ||
|
|
||
| _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"] |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -1402,14 +1402,27 @@ 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 | ||
| ) | ||
|
Comment on lines
+1410
to
+1417
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 📐 Maintainability & Code Quality | 🟡 Minor | ⚡ Quick win Add an MSA model-wiring regression test. The changed path writes K/V and BF16 index-K before Add a CUDA test in As per path instructions, “Leave an INLINE review comment on the smallest relevant changed production-code hunk when a material test coverage gap exists.” 🤖 Prompt for AI AgentsSource: Path instructions |
||
| forward_args = AttentionForwardArgs( | ||
| output=output, | ||
| sparse_backend_args=SparseBackendForwardArgs(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) | ||
|
|
||
Uh oh!
There was an error while loading. Please reload this page.