Fix MiniMax MSA fallback when fmha plan is unavailable (#29250)
This commit is contained in:
@@ -1,5 +1,6 @@
|
|||||||
# Copyright 2025 XunhaoLai. All rights reserved.
|
# Copyright 2025 XunhaoLai. All rights reserved.
|
||||||
|
|
||||||
|
import logging
|
||||||
from typing import Callable, List, Optional, Tuple
|
from typing import Callable, List, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -11,6 +12,20 @@ from .decode.topk_sparse import flash_decode_with_gqa_share_sparse
|
|||||||
from .prefill.flash_with_topk_idx import flash_prefill_with_topk_index
|
from .prefill.flash_with_topk_idx import flash_prefill_with_topk_index
|
||||||
from .prefill.topk_sparse import flash_prefill_with_gqa_share_sparse
|
from .prefill.topk_sparse import flash_prefill_with_gqa_share_sparse
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
_msa_fallback_warned = False
|
||||||
|
|
||||||
|
|
||||||
|
def _warn_msa_fallback(err: Exception) -> None:
|
||||||
|
global _msa_fallback_warned
|
||||||
|
if _msa_fallback_warned:
|
||||||
|
return
|
||||||
|
logger.warning(
|
||||||
|
"MiniMax MSA backend is unavailable (%s); falling back to Triton sparse attention.",
|
||||||
|
err,
|
||||||
|
)
|
||||||
|
_msa_fallback_warned = True
|
||||||
|
|
||||||
|
|
||||||
def minimax_sparse_prefill(
|
def minimax_sparse_prefill(
|
||||||
q: torch.Tensor, # [total_extend_tokens, num_q_heads, qk_head_dim]
|
q: torch.Tensor, # [total_extend_tokens, num_q_heads, qk_head_dim]
|
||||||
@@ -96,8 +111,9 @@ def minimax_sparse_prefill(
|
|||||||
# 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.
|
||||||
if use_msa and sink is None:
|
if use_msa and sink is None:
|
||||||
from .msa import msa_sparse_prefill_main
|
from .msa import MSAUnavailableError, msa_sparse_prefill_main
|
||||||
|
|
||||||
|
try:
|
||||||
o = msa_sparse_prefill_main(
|
o = msa_sparse_prefill_main(
|
||||||
q=q,
|
q=q,
|
||||||
k_cache=k_cache,
|
k_cache=k_cache,
|
||||||
@@ -111,6 +127,26 @@ def minimax_sparse_prefill(
|
|||||||
block_size_k=block_size_k,
|
block_size_k=block_size_k,
|
||||||
sm_scale=sm_scale,
|
sm_scale=sm_scale,
|
||||||
)
|
)
|
||||||
|
except MSAUnavailableError as err:
|
||||||
|
_warn_msa_fallback(err)
|
||||||
|
o = flash_prefill_with_gqa_share_sparse(
|
||||||
|
q=q,
|
||||||
|
k_cache=k_cache,
|
||||||
|
v_cache=v_cache,
|
||||||
|
sink=sink,
|
||||||
|
req_to_token=req_to_token,
|
||||||
|
slot_ids=slot_ids,
|
||||||
|
topk_idx=topk_idx,
|
||||||
|
block_size_q=block_size_q,
|
||||||
|
block_size_k=block_size_k,
|
||||||
|
cu_seqlens=cu_seqlens,
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
prefix_lens=prefix_lens,
|
||||||
|
max_seqlen_q=max_seqlen_q,
|
||||||
|
sm_scale=sm_scale,
|
||||||
|
cu_seqblocks_q=cu_seqblocks_q,
|
||||||
|
max_seqblock_q=max_seqblock_q,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
o = flash_prefill_with_gqa_share_sparse(
|
o = flash_prefill_with_gqa_share_sparse(
|
||||||
q=q,
|
q=q,
|
||||||
@@ -203,8 +239,9 @@ def minimax_sparse_decode(
|
|||||||
# 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:
|
||||||
from .msa import msa_sparse_decode_main
|
from .msa import MSAUnavailableError, msa_sparse_decode_main
|
||||||
|
|
||||||
|
try:
|
||||||
o = msa_sparse_decode_main(
|
o = msa_sparse_decode_main(
|
||||||
q=q,
|
q=q,
|
||||||
k_cache=k_cache,
|
k_cache=k_cache,
|
||||||
@@ -218,6 +255,20 @@ def minimax_sparse_decode(
|
|||||||
kv_indices=msa_kv_indices,
|
kv_indices=msa_kv_indices,
|
||||||
plan=msa_plan,
|
plan=msa_plan,
|
||||||
)
|
)
|
||||||
|
except MSAUnavailableError as err:
|
||||||
|
_warn_msa_fallback(err)
|
||||||
|
o = flash_decode_with_gqa_share_sparse(
|
||||||
|
q=q,
|
||||||
|
sink=sink,
|
||||||
|
k_cache=k_cache,
|
||||||
|
v_cache=v_cache,
|
||||||
|
req_to_token=req_to_token,
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
slot_ids=slot_ids,
|
||||||
|
block_size=block_size_k,
|
||||||
|
topk_idx=topk_idx,
|
||||||
|
sm_scale=sm_scale,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
o = flash_decode_with_gqa_share_sparse(
|
o = flash_decode_with_gqa_share_sparse(
|
||||||
q=q,
|
q=q,
|
||||||
|
|||||||
@@ -12,9 +12,34 @@ from typing import Optional
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
class MSAUnavailableError(RuntimeError):
|
||||||
|
"""Raised when fmha_sm100 cannot serve the MiniMax MSA path."""
|
||||||
|
|
||||||
|
|
||||||
|
@functools.lru_cache(maxsize=1)
|
||||||
|
def _load_fmha_sm100():
|
||||||
|
try:
|
||||||
|
from fmha_sm100 import fmha_sm100, fmha_sm100_plan
|
||||||
|
except Exception as err:
|
||||||
|
raise MSAUnavailableError(
|
||||||
|
"fmha_sm100 or fmha_sm100_plan is not importable"
|
||||||
|
) from err
|
||||||
|
if not callable(fmha_sm100) or not callable(fmha_sm100_plan):
|
||||||
|
raise MSAUnavailableError("fmha_sm100 exports must be callable")
|
||||||
|
return fmha_sm100, fmha_sm100_plan
|
||||||
|
|
||||||
|
|
||||||
|
def _run_fmha_sm100_plan(*args, **kwargs):
|
||||||
|
_, fmha_sm100_plan = _load_fmha_sm100()
|
||||||
|
try:
|
||||||
|
return fmha_sm100_plan(*args, **kwargs)
|
||||||
|
except (AttributeError, RuntimeError, TypeError) as err:
|
||||||
|
raise MSAUnavailableError("fmha_sm100_plan failed") from err
|
||||||
|
|
||||||
|
|
||||||
@functools.lru_cache(maxsize=1)
|
@functools.lru_cache(maxsize=1)
|
||||||
def msa_available() -> bool:
|
def msa_available() -> bool:
|
||||||
"""True iff the fmha_sm100 sparse kernels are importable on this device."""
|
"""True iff the fmha_sm100 sparse kernels and plan API are usable here."""
|
||||||
try:
|
try:
|
||||||
cap = torch.cuda.get_device_capability()
|
cap = torch.cuda.get_device_capability()
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -24,10 +49,9 @@ def msa_available() -> bool:
|
|||||||
if cap[0] != 10 or cap[1] not in (0, 3):
|
if cap[0] != 10 or cap[1] not in (0, 3):
|
||||||
return False
|
return False
|
||||||
try:
|
try:
|
||||||
import fmha_sm100 # noqa: F401
|
_load_fmha_sm100()
|
||||||
|
|
||||||
return True
|
return True
|
||||||
except Exception:
|
except MSAUnavailableError:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
@@ -82,7 +106,7 @@ def msa_sparse_prefill_main(
|
|||||||
|
|
||||||
Returns o [total_q, num_q_heads, head_dim].
|
Returns o [total_q, num_q_heads, head_dim].
|
||||||
"""
|
"""
|
||||||
from fmha_sm100 import fmha_sm100, fmha_sm100_plan
|
fmha_sm100, _ = _load_fmha_sm100()
|
||||||
|
|
||||||
max_slots, num_kv_heads, head_dim = k_cache.shape
|
max_slots, num_kv_heads, head_dim = k_cache.shape
|
||||||
num_q_heads = q.shape[1]
|
num_q_heads = q.shape[1]
|
||||||
@@ -105,7 +129,7 @@ def msa_sparse_prefill_main(
|
|||||||
# topk_idx [Hkv, total_q, topk] -> kv_block_indexes [total_q, Hkv, topk].
|
# topk_idx [Hkv, total_q, topk] -> kv_block_indexes [total_q, Hkv, topk].
|
||||||
kv_block_indexes = topk_idx.permute(1, 0, 2).contiguous().to(torch.int32)
|
kv_block_indexes = topk_idx.permute(1, 0, 2).contiguous().to(torch.int32)
|
||||||
|
|
||||||
plan = fmha_sm100_plan(
|
plan = _run_fmha_sm100_plan(
|
||||||
qo_segment_lens,
|
qo_segment_lens,
|
||||||
seq_lens.to(torch.int32),
|
seq_lens.to(torch.int32),
|
||||||
num_q_heads,
|
num_q_heads,
|
||||||
@@ -147,8 +171,6 @@ def build_msa_decode_meta(
|
|||||||
Used only by the standalone parity harnesses; the serving backend builds eager-decode
|
Used only by the standalone parity harnesses; the serving backend builds eager-decode
|
||||||
metadata via ``build_msa_decode_cg_plan`` + ``update_msa_decode_cg_meta`` instead.
|
metadata via ``build_msa_decode_cg_plan`` + ``update_msa_decode_cg_meta`` instead.
|
||||||
"""
|
"""
|
||||||
from fmha_sm100 import fmha_sm100_plan
|
|
||||||
|
|
||||||
max_slots, num_kv_heads, _ = k_cache.shape
|
max_slots, num_kv_heads, _ = k_cache.shape
|
||||||
P = block_size_k
|
P = block_size_k
|
||||||
if max_slots % P != 0:
|
if max_slots % P != 0:
|
||||||
@@ -156,7 +178,7 @@ def build_msa_decode_meta(
|
|||||||
B = slot_ids.shape[0]
|
B = slot_ids.shape[0]
|
||||||
kv_indices = _build_page_table(req_to_token, slot_ids, seq_lens, P)
|
kv_indices = _build_page_table(req_to_token, slot_ids, seq_lens, P)
|
||||||
seq_lens_i32 = seq_lens.to(torch.int32)
|
seq_lens_i32 = seq_lens.to(torch.int32)
|
||||||
plan = fmha_sm100_plan(
|
plan = _run_fmha_sm100_plan(
|
||||||
torch.ones(B, dtype=torch.int32),
|
torch.ones(B, dtype=torch.int32),
|
||||||
seq_lens_i32,
|
seq_lens_i32,
|
||||||
num_q_heads,
|
num_q_heads,
|
||||||
@@ -199,14 +221,14 @@ def _check_cg_plan_layout(plan) -> None:
|
|||||||
assumes (dict at tuple index 3 holding the four length tensors) — these are
|
assumes (dict at tuple index 3 holding the four length tensors) — these are
|
||||||
undocumented fmha_sm100 internals."""
|
undocumented fmha_sm100 internals."""
|
||||||
if not (isinstance(plan, tuple) and len(plan) > 3 and isinstance(plan[3], dict)):
|
if not (isinstance(plan, tuple) and len(plan) > 3 and isinstance(plan[3], dict)):
|
||||||
raise RuntimeError(
|
raise MSAUnavailableError(
|
||||||
"fmha_sm100_plan no longer returns a tuple with a metadata dict at index 3; "
|
"fmha_sm100_plan no longer returns a tuple with a metadata dict at index 3; "
|
||||||
"the MSA CUDA-graph decode path must be revalidated against this fmha_sm100 "
|
"the MSA CUDA-graph decode path must be revalidated against this fmha_sm100 "
|
||||||
"version. Set SGLANG_DISABLE_MSA=1 to use the Triton path meanwhile."
|
"version. Set SGLANG_DISABLE_MSA=1 to use the Triton path meanwhile."
|
||||||
)
|
)
|
||||||
missing = [k for k in _MSA_CG_LEN_KEYS if not torch.is_tensor(plan[3].get(k))]
|
missing = [k for k in _MSA_CG_LEN_KEYS if not torch.is_tensor(plan[3].get(k))]
|
||||||
if missing:
|
if missing:
|
||||||
raise RuntimeError(
|
raise MSAUnavailableError(
|
||||||
f"fmha_sm100 plan is missing length tensors {missing}; the MSA CUDA-graph "
|
f"fmha_sm100 plan is missing length tensors {missing}; the MSA CUDA-graph "
|
||||||
"decode path must be revalidated against this fmha_sm100 version. "
|
"decode path must be revalidated against this fmha_sm100 version. "
|
||||||
"Set SGLANG_DISABLE_MSA=1 to use the Triton path meanwhile."
|
"Set SGLANG_DISABLE_MSA=1 to use the Triton path meanwhile."
|
||||||
@@ -228,13 +250,11 @@ def build_msa_decode_cg_plan(
|
|||||||
every topk block valid; the length-dependent tensors are overwritten each step by
|
every topk block valid; the length-dependent tensors are overwritten each step by
|
||||||
``update_msa_decode_cg_meta``. Returns the plan tuple to pass to ``fmha_sm100``.
|
``update_msa_decode_cg_meta``. Returns the plan tuple to pass to ``fmha_sm100``.
|
||||||
"""
|
"""
|
||||||
from fmha_sm100 import fmha_sm100_plan
|
|
||||||
|
|
||||||
P = block_size_k
|
P = block_size_k
|
||||||
ref_len = topk * P # length at which all topk blocks exist -> full worklist
|
ref_len = topk * P # length at which all topk blocks exist -> full worklist
|
||||||
qo = torch.ones(batch_size, dtype=torch.int32)
|
qo = torch.ones(batch_size, dtype=torch.int32)
|
||||||
kv = torch.full((batch_size,), ref_len, dtype=torch.int32)
|
kv = torch.full((batch_size,), ref_len, dtype=torch.int32)
|
||||||
plan = fmha_sm100_plan(
|
plan = _run_fmha_sm100_plan(
|
||||||
qo,
|
qo,
|
||||||
kv,
|
kv,
|
||||||
num_q_heads,
|
num_q_heads,
|
||||||
@@ -266,8 +286,6 @@ def update_msa_decode_cg_meta(
|
|||||||
capture — i.e. only from ``init_forward_metadata_out_graph``. The captured graph then
|
capture — i.e. only from ``init_forward_metadata_out_graph``. The captured graph then
|
||||||
reads the same plan-tensor and ``kv_indices_buf`` addresses on replay.
|
reads the same plan-tensor and ``kv_indices_buf`` addresses on replay.
|
||||||
"""
|
"""
|
||||||
from fmha_sm100 import fmha_sm100_plan
|
|
||||||
|
|
||||||
P = block_size_k
|
P = block_size_k
|
||||||
B = seq_lens.shape[0]
|
B = seq_lens.shape[0]
|
||||||
seq_lens_i32 = seq_lens.to(torch.int32)
|
seq_lens_i32 = seq_lens.to(torch.int32)
|
||||||
@@ -276,7 +294,7 @@ def update_msa_decode_cg_meta(
|
|||||||
# worklist is identical to the persistent one (topk*P based) and is discarded.
|
# worklist is identical to the persistent one (topk*P based) and is discarded.
|
||||||
# qo_offset is clamped: graph replay pads the batch with seq_len==0 slots
|
# qo_offset is clamped: graph replay pads the batch with seq_len==0 slots
|
||||||
# (masked via kv_segment_lens==0, but seq_len-1 would be -1).
|
# (masked via kv_segment_lens==0, but seq_len-1 would be -1).
|
||||||
fresh = fmha_sm100_plan(
|
fresh = _run_fmha_sm100_plan(
|
||||||
torch.ones(B, dtype=torch.int32),
|
torch.ones(B, dtype=torch.int32),
|
||||||
seq_lens_i32,
|
seq_lens_i32,
|
||||||
num_q_heads,
|
num_q_heads,
|
||||||
@@ -320,7 +338,7 @@ def msa_sparse_decode_main(
|
|||||||
(eager decode only) and passes them in. When omitted (only the standalone parity
|
(eager decode only) and passes them in. When omitted (only the standalone parity
|
||||||
harnesses) they are built here via ``build_msa_decode_meta``.
|
harnesses) they are built here via ``build_msa_decode_meta``.
|
||||||
"""
|
"""
|
||||||
from fmha_sm100 import fmha_sm100
|
fmha_sm100, _ = _load_fmha_sm100()
|
||||||
|
|
||||||
max_slots, num_kv_heads, head_dim = k_cache.shape
|
max_slots, num_kv_heads, head_dim = k_cache.shape
|
||||||
H = q.shape[1]
|
H = q.shape[1]
|
||||||
|
|||||||
Reference in New Issue
Block a user