MiniMax-M3: share the sparse index top-k across layers and reuse the decode top-k buffer (#36527)
Co-authored-by: Kevin Mi <mikevin920@yahoo.com>
This commit is contained in:
@@ -18,22 +18,53 @@ from ..common.utils import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@triton.heuristics(
|
def _prune_decode_configs(configs, named_args, **kwargs):
|
||||||
{
|
"""Drop autotune configs whose token tile is smaller than a sparse block.
|
||||||
"BLOCK_SIZE_H": lambda args: max(
|
|
||||||
16, triton.next_power_of_2(args["gqa_group_size"])
|
BLOCKS_PER_K_BLOCK = BLOCK_SIZE_N // block_size is 0 for those, so they
|
||||||
),
|
cannot compile. Keep the full list if nothing survives (block_size larger
|
||||||
"BLOCK_SIZE_D": lambda args: triton.next_power_of_2(args["head_dim"]),
|
than every tile) and let triton report the real failure.
|
||||||
"BATCH_SIZE_BUCKET": lambda args: triton.next_power_of_2(args["batch_size"]),
|
"""
|
||||||
}
|
block_size = named_args["block_size"]
|
||||||
)
|
kept = [c for c in configs if c.kwargs["BLOCK_SIZE_N"] >= block_size]
|
||||||
@triton.autotune(
|
return kept or list(configs)
|
||||||
configs=[
|
|
||||||
|
|
||||||
|
_ON_HIP = torch.version.hip is not None
|
||||||
|
|
||||||
|
|
||||||
|
def _decode_score_block_n(args):
|
||||||
|
# gfx950 pick: 512 for tiny batches (few CTAs), else 128.
|
||||||
|
return max(512 if args["batch_size"] <= 4 else 128, args["block_size"])
|
||||||
|
|
||||||
|
|
||||||
|
# On ROCm the autotune sweep is replaced by a fixed config plus the
|
||||||
|
# BLOCK_SIZE_N heuristic above: runtime autotune can fire under CUDA-graph
|
||||||
|
# capture and mistune.
|
||||||
|
_DECODE_SCORE_CONFIGS = (
|
||||||
|
[triton.Config({}, num_warps=4, num_stages=1)]
|
||||||
|
if _ON_HIP
|
||||||
|
else [
|
||||||
triton.Config({"BLOCK_SIZE_N": BN}, num_warps=nw, num_stages=ns)
|
triton.Config({"BLOCK_SIZE_N": BN}, num_warps=nw, num_stages=ns)
|
||||||
for BN in [64, 128, 256, 512]
|
for BN in [64, 128, 256, 512]
|
||||||
for nw in [4, 8, 16]
|
for nw in [4, 8, 16]
|
||||||
for ns in [1, 2, 3]
|
for ns in [1, 2, 3]
|
||||||
],
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
_DECODE_SCORE_HEURISTICS = {
|
||||||
|
"BLOCK_SIZE_H": lambda args: max(
|
||||||
|
16, triton.next_power_of_2(args["gqa_group_size"])
|
||||||
|
),
|
||||||
|
"BLOCK_SIZE_D": lambda args: triton.next_power_of_2(args["head_dim"]),
|
||||||
|
"BATCH_SIZE_BUCKET": lambda args: triton.next_power_of_2(args["batch_size"]),
|
||||||
|
**({"BLOCK_SIZE_N": _decode_score_block_n} if _ON_HIP else {}),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@triton.heuristics(_DECODE_SCORE_HEURISTICS)
|
||||||
|
@triton.autotune(
|
||||||
|
configs=_DECODE_SCORE_CONFIGS,
|
||||||
key=[
|
key=[
|
||||||
"BATCH_SIZE_BUCKET",
|
"BATCH_SIZE_BUCKET",
|
||||||
"gqa_group_size",
|
"gqa_group_size",
|
||||||
@@ -41,6 +72,9 @@ from ..common.utils import (
|
|||||||
"block_size",
|
"block_size",
|
||||||
"SCORE_TYPE",
|
"SCORE_TYPE",
|
||||||
],
|
],
|
||||||
|
prune_configs_by=(
|
||||||
|
None if _ON_HIP else {"early_config_prune": _prune_decode_configs}
|
||||||
|
),
|
||||||
)
|
)
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def _decode_score_kernel(
|
def _decode_score_kernel(
|
||||||
@@ -228,6 +262,7 @@ def _decode_score_kernel(
|
|||||||
"HAS_SINK",
|
"HAS_SINK",
|
||||||
"SCORE_TYPE",
|
"SCORE_TYPE",
|
||||||
],
|
],
|
||||||
|
prune_configs_by={"early_config_prune": _prune_decode_configs},
|
||||||
)
|
)
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def _decode_score_attn_kernel(
|
def _decode_score_attn_kernel(
|
||||||
|
|||||||
@@ -1579,6 +1579,14 @@ class Envs:
|
|||||||
# MiniMax-M3 MXFP8 MoE experimental fusion toggles (default off; A/B only).
|
# MiniMax-M3 MXFP8 MoE experimental fusion toggles (default off; A/B only).
|
||||||
SGLANG_MINIMAX_M3_FUSED_SWIGLU_MXFP8 = EnvBool(False)
|
SGLANG_MINIMAX_M3_FUSED_SWIGLU_MXFP8 = EnvBool(False)
|
||||||
SGLANG_MINIMAX_M3_FUSED_MOE_COMBINE = EnvBool(False)
|
SGLANG_MINIMAX_M3_FUSED_MOE_COMBINE = EnvBool(False)
|
||||||
|
|
||||||
|
# MiniMax-M3 sparse-attention toggles for ROCm.
|
||||||
|
# Share one index top-k across every N sparse layers; 1 disables sharing.
|
||||||
|
# Changes which KV blocks the skip layers attend, so it applies on ROCm only
|
||||||
|
# (never under two-batch overlap); elsewhere the backend pins 1.
|
||||||
|
# 2 is the accuracy-safe default: higher values reuse staler selections
|
||||||
|
# in the skip layers.
|
||||||
|
SGLANG_MINIMAX_M3_INDEX_TOPK_FREQ = EnvInt(2)
|
||||||
# MiniMax M3 NPU prefill MAIN-attention: route the sparse main attention through
|
# MiniMax M3 NPU prefill MAIN-attention: route the sparse main attention through
|
||||||
# the native Ascend FA op `torch.ops.npu.npu_fused_infer_attention_score` (FIA)
|
# the native Ascend FA op `torch.ops.npu.npu_fused_infer_attention_score` (FIA)
|
||||||
# with a per-query CUSTOM block_table
|
# with a per-query CUSTOM block_table
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ from sglang.srt.layers.attention.base_attn_backend import (
|
|||||||
AttentionBackend,
|
AttentionBackend,
|
||||||
SharedReadEnds,
|
SharedReadEnds,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.moe.utils import is_tbo_enabled
|
||||||
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
|
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
@@ -26,7 +27,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_spec,
|
get_spec,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import m3_fp8_attn_gemm_enabled
|
from sglang.srt.server_args import m3_fp8_attn_gemm_enabled
|
||||||
from sglang.srt.utils import is_npu
|
from sglang.srt.utils import is_gfx95_supported, is_hip, is_npu
|
||||||
|
|
||||||
if is_npu():
|
if is_npu():
|
||||||
from sglang.kernels.ops.attention.minimax_sparse.common.index import (
|
from sglang.kernels.ops.attention.minimax_sparse.common.index import (
|
||||||
@@ -147,6 +148,12 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
|
|
||||||
# NPU: per-forward cached metadata for the triton paths (rebuilt each forward).
|
# NPU: per-forward cached metadata for the triton paths (rebuilt each forward).
|
||||||
self._prefill_meta: Optional[SimpleNamespace] = None
|
self._prefill_meta: Optional[SimpleNamespace] = None
|
||||||
|
# (owning ForwardBatch, cu_seqlens, seq_lens, prefix_lens, cu_seqblocks_q,
|
||||||
|
# max_seqblock_q, all_seqblock_q). The owner is part of the key because one
|
||||||
|
# metadata init can be followed by more than one ForwardBatch reaching the
|
||||||
|
# layers (two-batch overlap splits into two children with different
|
||||||
|
# extend_seq_lens); a hit requires the SAME object, not just a live cache.
|
||||||
|
self._prefill_seqblock_meta: Optional[tuple] = None
|
||||||
self._extend_meta: Optional[SimpleNamespace] = None
|
self._extend_meta: Optional[SimpleNamespace] = None
|
||||||
self._extend_meta_key: Optional[int] = None
|
self._extend_meta_key: Optional[int] = None
|
||||||
self._decode_seq_lens_i32_cg: dict[int, torch.Tensor] = {}
|
self._decode_seq_lens_i32_cg: dict[int, torch.Tensor] = {}
|
||||||
@@ -276,10 +283,45 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
self.dense_backend: Optional[AttentionBackend] = None
|
self.dense_backend: Optional[AttentionBackend] = None
|
||||||
|
|
||||||
|
self.index_topk_freq = (
|
||||||
|
max(int(envs.SGLANG_MINIMAX_M3_INDEX_TOPK_FREQ.get()), 1)
|
||||||
|
if is_hip() and is_gfx95_supported() and not is_tbo_enabled()
|
||||||
|
else 1
|
||||||
|
)
|
||||||
|
self.index_cache_enabled = self.index_topk_freq > 1
|
||||||
|
# topk_index_reduce widens the last dim to idx_group_size * topk_blocks
|
||||||
|
# (union of the group's selections), so the shared decode buffer must be
|
||||||
|
# that wide. Head split mirrors MiniMaxM3 sparse attention's.
|
||||||
|
self._idx_group_size = 1
|
||||||
|
if self.index_cache_enabled:
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
|
||||||
|
_num_idx_heads = max(
|
||||||
|
sparse_cfg["sparse_num_index_heads"] // get_parallel().attn_tp_size, 1
|
||||||
|
)
|
||||||
|
self._idx_group_size = max(
|
||||||
|
_num_idx_heads // self.kv_pool.main_pool.head_num, 1
|
||||||
|
)
|
||||||
|
# Persistent per-bs device buffer for decode top-k reuse. Allocated eagerly
|
||||||
|
# outside CUDA-graph capture; the captured graph only copies into and reads
|
||||||
|
# from a fixed address.
|
||||||
|
self._decode_topk_buf: dict = {}
|
||||||
|
self._topk_group_of_layer: dict[int, int] = {}
|
||||||
|
self._topk_is_source: dict[int, bool] = {}
|
||||||
|
for ordinal, lid in enumerate(
|
||||||
|
lid for lid in self.sparse_layer_ids if lid in self.disable_value_layer_ids
|
||||||
|
):
|
||||||
|
group = ordinal // self.index_topk_freq
|
||||||
|
self._topk_group_of_layer[lid] = group
|
||||||
|
self._topk_is_source[lid] = (ordinal % self.index_topk_freq) == 0
|
||||||
|
self._topk_cache: dict = {}
|
||||||
|
self._topk_cache_owner: Optional[ForwardBatch] = None
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
f"[MiniMaxSparse] Backend initialized "
|
f"[MiniMaxSparse] Backend initialized "
|
||||||
f"(score_type={self.score_type!r}, "
|
f"(score_type={self.score_type!r}, "
|
||||||
f"main_attn={'MSA' if self.use_msa else 'triton'}, "
|
f"main_attn={'MSA' if self.use_msa else 'triton'}, "
|
||||||
|
f"index_topk_freq={self.index_topk_freq}, "
|
||||||
f"msa_decode={self._use_msa_decode}, "
|
f"msa_decode={self._use_msa_decode}, "
|
||||||
f"msa_owns_decode={self._msa_owns_decode}, "
|
f"msa_owns_decode={self._msa_owns_decode}, "
|
||||||
f"decode_cuda_graph={_decode_cuda_graph}, "
|
f"decode_cuda_graph={_decode_cuda_graph}, "
|
||||||
@@ -330,6 +372,23 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
):
|
):
|
||||||
# getattr covers replay views lacking extend_seq_lens_cpu and TARGET_VERIFY.
|
# getattr covers replay views lacking extend_seq_lens_cpu and TARGET_VERIFY.
|
||||||
self._msa_dec_meta = None
|
self._msa_dec_meta = None
|
||||||
|
# New forward -> drop the per-forward index-cache top-k (prefill only).
|
||||||
|
if self.index_cache_enabled:
|
||||||
|
self._topk_cache = {}
|
||||||
|
self._topk_cache_owner = None
|
||||||
|
# Decode top-k reuse: pre-allocate the per-bs persistent buffer so graph
|
||||||
|
# capture never allocates. num_kv_heads == 1 at TP>=4 for M3.
|
||||||
|
if self.index_cache_enabled and forward_batch.forward_mode.is_decode_or_idle():
|
||||||
|
bs = forward_batch.seq_lens.shape[0]
|
||||||
|
if bs > 0 and bs not in self._decode_topk_buf:
|
||||||
|
_nkv = self.kv_pool.main_pool.head_num
|
||||||
|
self._decode_topk_buf[bs] = torch.empty(
|
||||||
|
(_nkv, bs, self.topk_blocks * self._idx_group_size),
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=forward_batch.seq_lens.device,
|
||||||
|
)
|
||||||
|
# Per-forward cache of the layer-invariant prefill seqblock trio.
|
||||||
|
self._prefill_seqblock_meta = None
|
||||||
if self.is_npu:
|
if self.is_npu:
|
||||||
# Invalidate cached prefill/extend metadata; rebuilt on first sparse layer.
|
# Invalidate cached prefill/extend metadata; rebuilt on first sparse layer.
|
||||||
self._prefill_meta = None
|
self._prefill_meta = None
|
||||||
@@ -1347,7 +1406,46 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
else:
|
else:
|
||||||
idx_k_cache, idx_v_cache = self.kv_pool.get_index_kv_buffer(layer.layer_id)
|
idx_k_cache, idx_v_cache = self.kv_pool.get_index_kv_buffer(layer.layer_id)
|
||||||
|
|
||||||
cu_seqlens, seq_lens, prefix_lens = self._resolve_extend_meta(forward_batch, q)
|
cached = self._prefill_seqblock_meta
|
||||||
|
if cached is None or cached[0] is not forward_batch:
|
||||||
|
cu_seqlens, seq_lens, prefix_lens = self._resolve_extend_meta(
|
||||||
|
forward_batch, q
|
||||||
|
)
|
||||||
|
if self.is_npu:
|
||||||
|
cu_seqblocks_q = max_seqblock_q = all_seqblock_q = None
|
||||||
|
else:
|
||||||
|
from sglang.kernels.ops.attention.minimax_sparse.common.utils import (
|
||||||
|
get_cu_seqblocks,
|
||||||
|
)
|
||||||
|
|
||||||
|
cu_seqblocks_q, max_seqblock_q, all_seqblock_q, _, _, _ = (
|
||||||
|
get_cu_seqblocks(
|
||||||
|
cu_seqlens,
|
||||||
|
self._max_seqlen_q,
|
||||||
|
self.block_size_q,
|
||||||
|
self.block_size_k,
|
||||||
|
forward_batch.extend_seq_lens_cpu,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
cached = (
|
||||||
|
forward_batch,
|
||||||
|
cu_seqlens,
|
||||||
|
seq_lens,
|
||||||
|
prefix_lens,
|
||||||
|
cu_seqblocks_q,
|
||||||
|
max_seqblock_q,
|
||||||
|
all_seqblock_q,
|
||||||
|
)
|
||||||
|
self._prefill_seqblock_meta = cached
|
||||||
|
(
|
||||||
|
_,
|
||||||
|
cu_seqlens,
|
||||||
|
seq_lens,
|
||||||
|
prefix_lens,
|
||||||
|
cu_seqblocks_q,
|
||||||
|
max_seqblock_q,
|
||||||
|
all_seqblock_q,
|
||||||
|
) = cached
|
||||||
|
|
||||||
# DP attention pads q beyond real tokens; trim (CPU list avoids a sync).
|
# DP attention pads q beyond real tokens; trim (CPU list avoids a sync).
|
||||||
if forward_batch.extend_seq_lens_cpu is not None:
|
if forward_batch.extend_seq_lens_cpu is not None:
|
||||||
@@ -1398,7 +1496,24 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
minimax_sparse_prefill,
|
minimax_sparse_prefill,
|
||||||
)
|
)
|
||||||
|
|
||||||
idx_o, o = minimax_sparse_prefill(
|
# Index cache: only for disable_value layers (idx_o is None,
|
||||||
|
# so skipping the indexer has no output side effect). A group's source
|
||||||
|
# layer computes + stores the reduced top-k; the other layers reuse it.
|
||||||
|
use_index_cache = self.index_cache_enabled and disable_value
|
||||||
|
cached_topk_idx = None
|
||||||
|
want_topk = False
|
||||||
|
if use_index_cache:
|
||||||
|
if self._topk_cache_owner is not forward_batch:
|
||||||
|
self._topk_cache = {}
|
||||||
|
self._topk_cache_owner = forward_batch
|
||||||
|
group = self._topk_group_of_layer[layer.layer_id]
|
||||||
|
if self._topk_is_source[layer.layer_id]:
|
||||||
|
want_topk = True # compute and store for this group
|
||||||
|
else:
|
||||||
|
cached_topk_idx = self._topk_cache.get(group)
|
||||||
|
# Miss (e.g. source layer chunked differently) -> recompute safely.
|
||||||
|
|
||||||
|
result = minimax_sparse_prefill(
|
||||||
q,
|
q,
|
||||||
k_cache,
|
k_cache,
|
||||||
v_cache,
|
v_cache,
|
||||||
@@ -1423,13 +1538,23 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
disable_index_value=disable_value,
|
disable_index_value=disable_value,
|
||||||
use_msa=self.use_msa,
|
use_msa=self.use_msa,
|
||||||
seqlens_cpu=forward_batch.extend_seq_lens_cpu,
|
seqlens_cpu=forward_batch.extend_seq_lens_cpu,
|
||||||
|
cu_seqblocks_q=cu_seqblocks_q,
|
||||||
|
max_seqblock_q=max_seqblock_q,
|
||||||
|
all_seqblock_q=all_seqblock_q,
|
||||||
q_scale=layer.q_scale_float,
|
q_scale=layer.q_scale_float,
|
||||||
k_scale=layer.k_scale_float,
|
k_scale=layer.k_scale_float,
|
||||||
v_scale=layer.v_scale_float,
|
v_scale=layer.v_scale_float,
|
||||||
idx_q_scale=layer.idx_q_scale_float,
|
idx_q_scale=layer.idx_q_scale_float,
|
||||||
idx_k_scale=layer.idx_k_scale_float,
|
idx_k_scale=layer.idx_k_scale_float,
|
||||||
idx_v_scale=layer.idx_v_scale_float,
|
idx_v_scale=layer.idx_v_scale_float,
|
||||||
|
cached_topk_idx=cached_topk_idx,
|
||||||
|
return_topk_idx=want_topk,
|
||||||
)
|
)
|
||||||
|
if want_topk:
|
||||||
|
idx_o, o, reduced_topk_idx = result
|
||||||
|
self._topk_cache[group] = reduced_topk_idx
|
||||||
|
else:
|
||||||
|
idx_o, o = result
|
||||||
if actual_num_tokens < original_num_tokens:
|
if actual_num_tokens < original_num_tokens:
|
||||||
pad_len = original_num_tokens - actual_num_tokens
|
pad_len = original_num_tokens - actual_num_tokens
|
||||||
o = torch.cat([o, o.new_zeros(pad_len, *o.shape[1:])], dim=0)
|
o = torch.cat([o, o.new_zeros(pad_len, *o.shape[1:])], dim=0)
|
||||||
@@ -1499,18 +1624,19 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
):
|
):
|
||||||
assert len(kwargs) == 0
|
assert len(kwargs) == 0
|
||||||
disable_value = layer.layer_id in self.disable_value_layer_ids
|
disable_value = layer.layer_id in self.disable_value_layer_ids
|
||||||
self.kv_pool.set_fused_kv_index_buffer(
|
if not self._is_sparse_kv_cached_by_fusion(forward_batch, layer.layer_id):
|
||||||
layer,
|
self.kv_pool.set_fused_kv_index_buffer(
|
||||||
forward_batch.out_cache_loc,
|
layer,
|
||||||
k,
|
forward_batch.out_cache_loc,
|
||||||
v,
|
k,
|
||||||
idx_k,
|
v,
|
||||||
None if disable_value else idx_v,
|
idx_k,
|
||||||
layer.k_scale_float,
|
None if disable_value else idx_v,
|
||||||
layer.v_scale_float,
|
layer.k_scale_float,
|
||||||
layer.idx_k_scale_float,
|
layer.v_scale_float,
|
||||||
layer.idx_v_scale_float,
|
layer.idx_k_scale_float,
|
||||||
)
|
layer.idx_v_scale_float,
|
||||||
|
)
|
||||||
k_cache, v_cache = self.kv_pool.get_kv_buffer(layer.layer_id)
|
k_cache, v_cache = self.kv_pool.get_kv_buffer(layer.layer_id)
|
||||||
if disable_value:
|
if disable_value:
|
||||||
idx_k_cache = self.kv_pool.get_index_k_buffer(layer.layer_id)
|
idx_k_cache = self.kv_pool.get_index_k_buffer(layer.layer_id)
|
||||||
@@ -1565,6 +1691,17 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
minimax_sparse_decode,
|
minimax_sparse_decode,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Decode top-k reuse: group source layer computes+stores; skips reuse.
|
||||||
|
_use_reuse = self.index_cache_enabled and disable_value and attn_fn is None
|
||||||
|
_topk_buf = self._decode_topk_buf.get(q.shape[0]) if _use_reuse else None
|
||||||
|
_cached_topk = None
|
||||||
|
_want_topk = False
|
||||||
|
if _use_reuse and _topk_buf is not None:
|
||||||
|
if self._topk_is_source.get(layer.layer_id, True):
|
||||||
|
_want_topk = True
|
||||||
|
else:
|
||||||
|
_cached_topk = _topk_buf
|
||||||
|
|
||||||
idx_o, o = minimax_sparse_decode(
|
idx_o, o = minimax_sparse_decode(
|
||||||
q,
|
q,
|
||||||
None,
|
None,
|
||||||
@@ -1596,6 +1733,8 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
idx_q_scale=layer.idx_q_scale_float,
|
idx_q_scale=layer.idx_q_scale_float,
|
||||||
idx_k_scale=layer.idx_k_scale_float,
|
idx_k_scale=layer.idx_k_scale_float,
|
||||||
idx_v_scale=layer.idx_v_scale_float,
|
idx_v_scale=layer.idx_v_scale_float,
|
||||||
|
cached_topk_idx=_cached_topk,
|
||||||
|
topk_out=_topk_buf if _want_topk else None,
|
||||||
)
|
)
|
||||||
return (
|
return (
|
||||||
None if idx_o is None else idx_o.reshape(q.shape[0], -1).contiguous(),
|
None if idx_o is None else idx_o.reshape(q.shape[0], -1).contiguous(),
|
||||||
|
|||||||
@@ -73,9 +73,18 @@ def minimax_sparse_prefill(
|
|||||||
idx_q_scale: Optional[float] = None,
|
idx_q_scale: Optional[float] = None,
|
||||||
idx_k_scale: Optional[float] = None,
|
idx_k_scale: Optional[float] = None,
|
||||||
idx_v_scale: Optional[float] = None,
|
idx_v_scale: Optional[float] = None,
|
||||||
|
cached_topk_idx: Optional[torch.Tensor] = None,
|
||||||
|
return_topk_idx: bool = False,
|
||||||
):
|
):
|
||||||
"""Run MiniMax-M3 sparse prefill.
|
"""Run MiniMax-M3 sparse prefill.
|
||||||
|
|
||||||
|
Index cache: when ``cached_topk_idx`` is given, skip Step 1
|
||||||
|
(the flash-index attention + top-k selection) and reuse the provided top-k
|
||||||
|
indices for Step 3's sparse attention. When ``return_topk_idx`` is True, the
|
||||||
|
reduced top-k tensor is returned as a third element so the caller can cache
|
||||||
|
it for later skip layers. Only valid for ``disable_index_value`` layers
|
||||||
|
(idx_o is None there, so skipping the indexer has no output side effect).
|
||||||
|
|
||||||
``cu_seqblocks_q``, ``max_seqblock_q``, and ``all_seqblock_q`` are optional
|
``cu_seqblocks_q``, ``max_seqblock_q``, and ``all_seqblock_q`` are optional
|
||||||
precomputed query-block metadata shared by the index and value sparse
|
precomputed query-block metadata shared by the index and value sparse
|
||||||
kernels. Supplying them avoids recomputing the same block layout twice.
|
kernels. Supplying them avoids recomputing the same block layout twice.
|
||||||
@@ -88,42 +97,52 @@ def minimax_sparse_prefill(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# All seqlen is less than topk, use full attention
|
# All seqlen is less than topk, use full attention
|
||||||
# Step 1: Flash attention with topk index (using index head)
|
if cached_topk_idx is not None:
|
||||||
idx_o, topk_idx = flash_prefill_with_topk_index(
|
# Index cache hit: reuse a prior sparse layer's reduced
|
||||||
q=idx_q,
|
# top-k, skipping Step 1 (flash-index attention + top-k) and Step 2
|
||||||
k_cache=idx_k_cache,
|
# (reduce). idx_o is unused downstream for disable_index_value layers.
|
||||||
v_cache=idx_v_cache,
|
idx_o = None
|
||||||
sink=idx_sink,
|
topk_idx = cached_topk_idx
|
||||||
req_to_token=req_to_token,
|
else:
|
||||||
slot_ids=slot_ids,
|
# Step 1: Flash attention with topk index (using index head)
|
||||||
cu_seqlens=cu_seqlens,
|
idx_o, topk_idx = flash_prefill_with_topk_index(
|
||||||
seq_lens=seq_lens,
|
q=idx_q,
|
||||||
prefix_lens=prefix_lens,
|
k_cache=idx_k_cache,
|
||||||
max_seqlen_q=max_seqlen_q,
|
v_cache=idx_v_cache,
|
||||||
max_seqlen_k=max_seqlen_k,
|
sink=idx_sink,
|
||||||
block_size_q=block_size_q,
|
req_to_token=req_to_token,
|
||||||
block_size_k=block_size_k,
|
slot_ids=slot_ids,
|
||||||
topk=topk,
|
cu_seqlens=cu_seqlens,
|
||||||
init_blocks=init_blocks,
|
seq_lens=seq_lens,
|
||||||
local_blocks=local_blocks,
|
prefix_lens=prefix_lens,
|
||||||
sm_scale=idx_sm_scale,
|
max_seqlen_q=max_seqlen_q,
|
||||||
score_type=score_type,
|
max_seqlen_k=max_seqlen_k,
|
||||||
disable_index_value=disable_index_value,
|
block_size_q=block_size_q,
|
||||||
cu_seqblocks_q=cu_seqblocks_q,
|
block_size_k=block_size_k,
|
||||||
max_seqblock_q=max_seqblock_q,
|
topk=topk,
|
||||||
all_seqblock_q=all_seqblock_q,
|
init_blocks=init_blocks,
|
||||||
q_scale=idx_q_scale,
|
local_blocks=local_blocks,
|
||||||
k_scale=idx_k_scale,
|
sm_scale=idx_sm_scale,
|
||||||
v_scale=idx_v_scale,
|
score_type=score_type,
|
||||||
)
|
disable_index_value=disable_index_value,
|
||||||
# Step 2: Reduce topk idx if num_idx_heads > num_kv_heads
|
cu_seqblocks_q=cu_seqblocks_q,
|
||||||
num_idx_heads = idx_q.shape[1]
|
max_seqblock_q=max_seqblock_q,
|
||||||
num_kv_heads = k_cache.shape[1]
|
all_seqblock_q=all_seqblock_q,
|
||||||
idx_group_size = num_idx_heads // num_kv_heads
|
q_scale=idx_q_scale,
|
||||||
if idx_group_size > 1:
|
k_scale=idx_k_scale,
|
||||||
topk_idx = topk_index_reduce(
|
v_scale=idx_v_scale,
|
||||||
topk_idx.view(num_kv_heads, idx_group_size, -1, topk), dim=1
|
|
||||||
)
|
)
|
||||||
|
# Step 2: Reduce topk idx if num_idx_heads > num_kv_heads
|
||||||
|
num_idx_heads = idx_q.shape[1]
|
||||||
|
num_kv_heads = k_cache.shape[1]
|
||||||
|
idx_group_size = num_idx_heads // num_kv_heads
|
||||||
|
if idx_group_size > 1:
|
||||||
|
topk_idx = topk_index_reduce(
|
||||||
|
topk_idx.view(num_kv_heads, idx_group_size, -1, topk), dim=1
|
||||||
|
)
|
||||||
|
|
||||||
|
# Reduced top-k cached by the caller for subsequent skip layers.
|
||||||
|
reduced_topk_idx = topk_idx
|
||||||
# Step 3: Sparse attention using topk index (main head). The MSA path only
|
# Step 3: Sparse attention using topk index (main head). The MSA path only
|
||||||
# replaces this step; the indexer above is unchanged. MSA has no attn-sink
|
# replaces this step; the indexer above is unchanged. MSA has no attn-sink
|
||||||
# input, so keep the Triton path when sink is present.
|
# input, so keep the Triton path when sink is present.
|
||||||
@@ -192,6 +211,8 @@ def minimax_sparse_prefill(
|
|||||||
k_scale=k_scale,
|
k_scale=k_scale,
|
||||||
v_scale=v_scale,
|
v_scale=v_scale,
|
||||||
)
|
)
|
||||||
|
if return_topk_idx:
|
||||||
|
return idx_o, o, reduced_topk_idx
|
||||||
return idx_o, o
|
return idx_o, o
|
||||||
|
|
||||||
|
|
||||||
@@ -232,45 +253,72 @@ def minimax_sparse_decode(
|
|||||||
idx_q_scale: Optional[float] = None,
|
idx_q_scale: Optional[float] = None,
|
||||||
idx_k_scale: Optional[float] = None,
|
idx_k_scale: Optional[float] = None,
|
||||||
idx_v_scale: Optional[float] = None,
|
idx_v_scale: Optional[float] = None,
|
||||||
|
cached_topk_idx: Optional[torch.Tensor] = None,
|
||||||
|
topk_out: Optional[torch.Tensor] = None,
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
# Step 1: Flash decode with topk index (using index head). When the dense main
|
# Index top-k sharing for DECODE. A group's source layer passes ``topk_out``
|
||||||
# attention is used, the indexer emits the page table directly (fused
|
# (a persistent buffer) and publishes its reduced top-k there; the group's
|
||||||
# transform) instead of block ids, plus the per-query effective KV length.
|
# skip layers pass that buffer back as ``cached_topk_idx`` and skip Step 1
|
||||||
idx_o, topk_idx, real_seq_lens = flash_decode_with_topk_idx(
|
# (flash-index decode + top-k) and Step 2 (reduce) entirely, never reading
|
||||||
q=idx_q,
|
# idx_k_cache. Only valid for disable_index_value layers on the
|
||||||
sink=idx_sink,
|
# non-dense-main path (idx_o is None there). All-device: CUDA-graph safe.
|
||||||
k_cache=idx_k_cache,
|
if cached_topk_idx is not None:
|
||||||
v_cache=idx_v_cache,
|
idx_o = None
|
||||||
req_to_token=req_to_token,
|
real_seq_lens = None
|
||||||
seq_lens=seq_lens,
|
topk_idx = cached_topk_idx
|
||||||
max_seqlen=max_seqlen,
|
_skip_reduce = True
|
||||||
slot_ids=slot_ids,
|
else:
|
||||||
block_size=block_size_k,
|
_skip_reduce = False
|
||||||
topk=topk,
|
# Step 1: Flash decode with topk index (using index head). When the dense main
|
||||||
init_blocks=init_blocks,
|
# attention is used, the indexer emits the page table directly (fused
|
||||||
local_blocks=local_blocks,
|
# transform) instead of block ids, plus the per-query effective KV length.
|
||||||
sm_scale=idx_sm_scale,
|
idx_o, topk_idx, real_seq_lens = flash_decode_with_topk_idx(
|
||||||
score_type=score_type,
|
q=idx_q,
|
||||||
disable_index_value=disable_index_value,
|
sink=idx_sink,
|
||||||
use_dense_main_attn=dense_main_attn_fn is not None,
|
k_cache=idx_k_cache,
|
||||||
page_size=page_size,
|
v_cache=idx_v_cache,
|
||||||
q_scale=idx_q_scale,
|
req_to_token=req_to_token,
|
||||||
k_scale=idx_k_scale,
|
seq_lens=seq_lens,
|
||||||
v_scale=idx_v_scale,
|
max_seqlen=max_seqlen,
|
||||||
)
|
slot_ids=slot_ids,
|
||||||
|
block_size=block_size_k,
|
||||||
|
topk=topk,
|
||||||
|
init_blocks=init_blocks,
|
||||||
|
local_blocks=local_blocks,
|
||||||
|
sm_scale=idx_sm_scale,
|
||||||
|
score_type=score_type,
|
||||||
|
disable_index_value=disable_index_value,
|
||||||
|
use_dense_main_attn=dense_main_attn_fn is not None,
|
||||||
|
page_size=page_size,
|
||||||
|
q_scale=idx_q_scale,
|
||||||
|
k_scale=idx_k_scale,
|
||||||
|
v_scale=idx_v_scale,
|
||||||
|
)
|
||||||
num_idx_heads = idx_q.shape[1]
|
num_idx_heads = idx_q.shape[1]
|
||||||
num_kv_heads = k_cache.shape[1]
|
num_kv_heads = k_cache.shape[1]
|
||||||
idx_group_size = num_idx_heads // num_kv_heads
|
idx_group_size = num_idx_heads // num_kv_heads
|
||||||
|
assert dense_main_attn_fn is None or (
|
||||||
|
cached_topk_idx is None and topk_out is None
|
||||||
|
), "index top-k sharing is not available on the dense main-attention path"
|
||||||
if dense_main_attn_fn is not None:
|
if dense_main_attn_fn is not None:
|
||||||
# topk_idx is the page table; real_seq_lens is the per-query cache_seqlens
|
# topk_idx is the page table; real_seq_lens is the per-query cache_seqlens
|
||||||
assert idx_group_size == 1
|
assert idx_group_size == 1
|
||||||
o = dense_main_attn_fn(q, topk_idx, real_seq_lens)
|
o = dense_main_attn_fn(q, topk_idx, real_seq_lens)
|
||||||
else:
|
else:
|
||||||
# Step 2: Reduce topk idx if num_idx_heads > num_kv_heads
|
# Step 2: Reduce topk idx if num_idx_heads > num_kv_heads
|
||||||
if idx_group_size > 1:
|
if idx_group_size > 1 and not _skip_reduce:
|
||||||
topk_idx = topk_index_reduce(
|
topk_idx = topk_index_reduce(
|
||||||
topk_idx.view(num_kv_heads, idx_group_size, -1, topk), dim=1
|
topk_idx.view(num_kv_heads, idx_group_size, -1, topk), dim=1
|
||||||
)
|
)
|
||||||
|
if topk_out is not None:
|
||||||
|
# Publish the group's selection to the persistent buffer the skip
|
||||||
|
# layers read (fixed address -> CUDA-graph safe).
|
||||||
|
if topk_out.shape != topk_idx.shape:
|
||||||
|
raise ValueError(
|
||||||
|
f"topk_out shape {tuple(topk_out.shape)} does not match "
|
||||||
|
f"reduced top-k shape {tuple(topk_idx.shape)}"
|
||||||
|
)
|
||||||
|
topk_out.copy_(topk_idx)
|
||||||
# Step 3: Sparse attention using topk index (main head). The MSA path
|
# Step 3: Sparse attention using topk index (main head). The MSA path
|
||||||
# only replaces this step; keep the Triton path when sink is present.
|
# only replaces this step; keep the Triton path when sink is present.
|
||||||
if use_msa and sink is None:
|
if use_msa and sink is None:
|
||||||
|
|||||||
Reference in New Issue
Block a user