diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/minimax_sparse.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/minimax_sparse.py index a4ee9fcf2..046d7ff7c 100644 --- a/python/sglang/srt/layers/attention/minimax_sparse_ops/minimax_sparse.py +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/minimax_sparse.py @@ -1,5 +1,6 @@ # Copyright 2025 XunhaoLai. All rights reserved. +import logging from typing import Callable, List, Optional, Tuple 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.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( q: torch.Tensor, # [total_extend_tokens, num_q_heads, qk_head_dim] @@ -96,21 +111,42 @@ def minimax_sparse_prefill( # replaces this step; the indexer above is unchanged. MSA has no attn-sink # input, so keep the Triton path when sink is present. if use_msa and sink is None: - from .msa import msa_sparse_prefill_main + from .msa import MSAUnavailableError, msa_sparse_prefill_main - o = msa_sparse_prefill_main( - q=q, - k_cache=k_cache, - v_cache=v_cache, - topk_idx=topk_idx, - req_to_token=req_to_token, - slot_ids=slot_ids, - cu_seqlens=cu_seqlens, - seq_lens=seq_lens, - prefix_lens=prefix_lens, - block_size_k=block_size_k, - sm_scale=sm_scale, - ) + try: + o = msa_sparse_prefill_main( + q=q, + k_cache=k_cache, + v_cache=v_cache, + topk_idx=topk_idx, + req_to_token=req_to_token, + slot_ids=slot_ids, + cu_seqlens=cu_seqlens, + seq_lens=seq_lens, + prefix_lens=prefix_lens, + block_size_k=block_size_k, + 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: o = flash_prefill_with_gqa_share_sparse( q=q, @@ -203,21 +239,36 @@ def minimax_sparse_decode( # Step 3: Sparse attention using topk index (main head). The MSA path # only replaces this step; keep the Triton path when sink is present. if use_msa and sink is None: - from .msa import msa_sparse_decode_main + from .msa import MSAUnavailableError, msa_sparse_decode_main - o = msa_sparse_decode_main( - q=q, - k_cache=k_cache, - v_cache=v_cache, - topk_idx=topk_idx, - req_to_token=req_to_token, - slot_ids=slot_ids, - seq_lens=seq_lens, - block_size_k=block_size_k, - sm_scale=sm_scale, - kv_indices=msa_kv_indices, - plan=msa_plan, - ) + try: + o = msa_sparse_decode_main( + q=q, + k_cache=k_cache, + v_cache=v_cache, + topk_idx=topk_idx, + req_to_token=req_to_token, + slot_ids=slot_ids, + seq_lens=seq_lens, + block_size_k=block_size_k, + sm_scale=sm_scale, + kv_indices=msa_kv_indices, + 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: o = flash_decode_with_gqa_share_sparse( q=q, diff --git a/python/sglang/srt/layers/attention/minimax_sparse_ops/msa.py b/python/sglang/srt/layers/attention/minimax_sparse_ops/msa.py index b79fc3227..a9086c582 100644 --- a/python/sglang/srt/layers/attention/minimax_sparse_ops/msa.py +++ b/python/sglang/srt/layers/attention/minimax_sparse_ops/msa.py @@ -12,9 +12,34 @@ from typing import Optional 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) 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: cap = torch.cuda.get_device_capability() except Exception: @@ -24,10 +49,9 @@ def msa_available() -> bool: if cap[0] != 10 or cap[1] not in (0, 3): return False try: - import fmha_sm100 # noqa: F401 - + _load_fmha_sm100() return True - except Exception: + except MSAUnavailableError: return False @@ -82,7 +106,7 @@ def msa_sparse_prefill_main( 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 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]. 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, seq_lens.to(torch.int32), 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 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 P = block_size_k if max_slots % P != 0: @@ -156,7 +178,7 @@ def build_msa_decode_meta( B = slot_ids.shape[0] kv_indices = _build_page_table(req_to_token, slot_ids, seq_lens, P) seq_lens_i32 = seq_lens.to(torch.int32) - plan = fmha_sm100_plan( + plan = _run_fmha_sm100_plan( torch.ones(B, dtype=torch.int32), seq_lens_i32, 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 undocumented fmha_sm100 internals.""" 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; " "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." ) missing = [k for k in _MSA_CG_LEN_KEYS if not torch.is_tensor(plan[3].get(k))] if missing: - raise RuntimeError( + raise MSAUnavailableError( f"fmha_sm100 plan is missing length tensors {missing}; 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." @@ -228,13 +250,11 @@ def build_msa_decode_cg_plan( 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``. """ - from fmha_sm100 import fmha_sm100_plan - P = block_size_k ref_len = topk * P # length at which all topk blocks exist -> full worklist qo = torch.ones(batch_size, dtype=torch.int32) kv = torch.full((batch_size,), ref_len, dtype=torch.int32) - plan = fmha_sm100_plan( + plan = _run_fmha_sm100_plan( qo, kv, 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 reads the same plan-tensor and ``kv_indices_buf`` addresses on replay. """ - from fmha_sm100 import fmha_sm100_plan - P = block_size_k B = seq_lens.shape[0] 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. # 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). - fresh = fmha_sm100_plan( + fresh = _run_fmha_sm100_plan( torch.ones(B, dtype=torch.int32), seq_lens_i32, 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 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 H = q.shape[1]