From 887d380acedb3ffee6f48a072ba2cce08e43303a Mon Sep 17 00:00:00 2001 From: ori <39351881+froststeam@users.noreply.github.com> Date: Thu, 23 Apr 2026 11:22:20 +0800 Subject: [PATCH] [MUSA] Resolve output garbage in Context Parallel on MusaFlashAttentionBackend (#23270) Co-authored-by: zhiguo.qin --- 3rdparty/amd/wheel/sglang/pyproject.toml | 5 +- python/pyproject_other.toml | 18 +-- .../musa/attention/flashattention_backend.py | 107 ++++++++++-------- .../musa/layers/utils/__init__.py | 0 .../musa/layers/utils/cp_utils.py | 57 ++++++++++ sgl-kernel/pyproject_musa.toml | 2 +- 6 files changed, 128 insertions(+), 61 deletions(-) create mode 100644 python/sglang/srt/hardware_backend/musa/layers/utils/__init__.py create mode 100644 python/sglang/srt/hardware_backend/musa/layers/utils/cp_utils.py diff --git a/3rdparty/amd/wheel/sglang/pyproject.toml b/3rdparty/amd/wheel/sglang/pyproject.toml index 99a99c0b5..323985c89 100644 --- a/3rdparty/amd/wheel/sglang/pyproject.toml +++ b/3rdparty/amd/wheel/sglang/pyproject.toml @@ -123,8 +123,11 @@ srt_musa = [ "sglang[runtime_common]", "torch", "torch_musa", - "torchada>=0.1.45", + "torchada>=0.1.50", "mthreads-ml-py", + "mate>=0.2.0", + "deep-gemm>=0.1.3", + "flash_attn_3>=0.1.4", "numpy<2.0", ] diff --git a/python/pyproject_other.toml b/python/pyproject_other.toml index faa42ef84..93192b644 100755 --- a/python/pyproject_other.toml +++ b/python/pyproject_other.toml @@ -112,15 +112,15 @@ srt_hpu = ["sglang[runtime_common]"] # https://docs.sglang.io/platforms/mthreads_gpu.md srt_musa = [ - "sglang[runtime_common]", - "torch", - "torch_musa", - "torchada>=0.1.48", - "mthreads-ml-py", - "mate", - "mate-deep_gemm", - "mate-flash-attention", - "numpy<2.0", + "sglang[runtime_common]", + "torch", + "torch_musa", + "torchada>=0.1.50", + "mthreads-ml-py", + "mate>=0.2.0", + "deep-gemm>=0.1.3", + "flash_attn_3>=0.1.4", + "numpy<2.0", ] diffusion_musa = [ diff --git a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py index 4583c1d82..04dd7d6a0 100644 --- a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py +++ b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py @@ -4,12 +4,15 @@ import threading from typing import TYPE_CHECKING, Optional, Tuple, Union import torch -from flash_attn import flash_attn_varlen_func -from flash_attn import flash_attn_with_kvcache as mate_flash_attn_with_kvcache -from flash_attn import get_scheduler_metadata +from flash_attn_interface import flash_attn_varlen_func +from flash_attn_interface import flash_attn_with_kvcache as mate_flash_attn_with_kvcache +from flash_attn_interface import get_scheduler_metadata from sglang.srt.distributed import get_pp_group, get_pp_indices from sglang.srt.environ import envs +from sglang.srt.hardware_backend.musa.layers.utils.cp_utils import ( + musa_cp_attn_forward_extend as cp_attn_forward_extend, +) from sglang.srt.layers.attention.flashattention_backend import ( FlashAttentionBackend, merge_state_v2_wrapper, @@ -17,7 +20,6 @@ from sglang.srt.layers.attention.flashattention_backend import ( from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.layers.utils.cp_utils import ( cp_allgather_and_save_kv_cache, - cp_attn_forward_extend, ) from sglang.srt.server_args import get_global_server_args @@ -143,10 +145,15 @@ def flash_attn_with_kvcache( pack_gqa=None, sm_margin: int = 0, return_softmax_lse: bool = False, - ver: int = 3, - **kwargs, + sinks=None, + score_mod=None, + aux_tensors=None, + ver=3, ): """MUSA flash_attn_with_kvcache wrapper that auto-injects scheduler_metadata.""" + if ver != 3: + raise ValueError("Only ver=3 is supported for MUSA FA3.") + if scheduler_metadata is None and _CURRENT_BACKEND is not None: backend = _CURRENT_BACKEND # Ensure backend has been properly set up for this call @@ -195,6 +202,7 @@ def flash_attn_with_kvcache( pack_gqa=pack_gqa, sm_margin=sm_margin, return_softmax_lse=return_softmax_lse, + sinks=sinks, ) @@ -211,6 +219,9 @@ class MusaFlashAttentionBackend(FlashAttentionBackend): self._current_max_seqlen_k: int = 0 self._current_can_run_tbo: bool = False + # Disable default scheduler metadata for fa3 + self._get_scheduler_metadata = None + # Register this backend as the global current instance for the wrapper global _CURRENT_BACKEND _CURRENT_BACKEND = self @@ -402,6 +413,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend): ) result = cp_attn_forward_extend( + self, forward_batch, q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim), self.device, @@ -431,50 +443,6 @@ class MusaFlashAttentionBackend(FlashAttentionBackend): num_splits=self.num_splits, **kwargs, ) - - if use_cascade_attn: - # Update state for the second call - self._current_prefix = "forward_extend_use_cascade_attn" - self._current_max_seqlen_k = ( - self.forward_metadata_spec_decode_expand.max_seq_len_k - ) - - o, softmax_lse, *rest = result - o_expand, softmax_lse_expand, *rest_expand = ( - flash_attn_with_kvcache( - q=q.contiguous().view( - -1, layer.tp_q_head_num, layer.head_dim - ), - k_cache=key_cache.view( - -1, 1, layer.tp_k_head_num, layer.head_dim - ), - v_cache=value_cache.view( - -1, 1, layer.tp_v_head_num, layer.head_dim - ), - page_table=self.forward_metadata_spec_decode_expand.page_table, - cache_seqlens=self.forward_metadata_spec_decode_expand.cache_seqlens_int32, - cu_seqlens_q=self.forward_metadata_spec_decode_expand.cu_seqlens_q, - cu_seqlens_k_new=self.forward_metadata_spec_decode_expand.cu_seqlens_k, - max_seqlen_q=self.forward_metadata_spec_decode_expand.max_seq_len_q, - softmax_scale=layer.scaling, - causal=False, - window_size=window_size, - softcap=layer.logit_cap, - k_descale=k_descale, - v_descale=v_descale, - return_softmax_lse=True, - num_splits=self.num_splits, - **kwargs, - ) - ) - o, _ = merge_state_v2_wrapper( - o, - softmax_lse.T.contiguous(), - o_expand, - softmax_lse_expand.T.contiguous(), - ) - else: - o = result else: output = flash_attn_varlen_func( q=q.view(-1, layer.tp_q_head_num, layer.head_dim), @@ -487,6 +455,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend): softmax_scale=layer.scaling, causal=True, return_softmax_lse=forward_batch.mha_return_lse, + **kwargs, ) if forward_batch.mha_return_lse: output, lse, *rest = output @@ -496,6 +465,44 @@ class MusaFlashAttentionBackend(FlashAttentionBackend): lse, ) return output.view(-1, layer.tp_q_head_num * layer.v_head_dim) + + if use_cascade_attn: + # Update state for the second call + self._current_prefix = "forward_extend_use_cascade_attn" + self._current_max_seqlen_k = ( + self.forward_metadata_spec_decode_expand.max_seq_len_k + ) + + o, softmax_lse, *rest = result + o_expand, softmax_lse_expand, *rest_expand = flash_attn_with_kvcache( + q=q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim), + k_cache=key_cache.view(-1, 1, layer.tp_k_head_num, layer.head_dim), + v_cache=value_cache.view( + -1, 1, layer.tp_v_head_num, layer.head_dim + ), + page_table=self.forward_metadata_spec_decode_expand.page_table, + cache_seqlens=self.forward_metadata_spec_decode_expand.cache_seqlens_int32, + cu_seqlens_q=self.forward_metadata_spec_decode_expand.cu_seqlens_q, + cu_seqlens_k_new=self.forward_metadata_spec_decode_expand.cu_seqlens_k, + max_seqlen_q=self.forward_metadata_spec_decode_expand.max_seq_len_q, + softmax_scale=layer.scaling, + causal=False, + window_size=window_size, + softcap=layer.logit_cap, + k_descale=k_descale, + v_descale=v_descale, + return_softmax_lse=True, + num_splits=self.num_splits, + **kwargs, + ) + o, _ = merge_state_v2_wrapper( + o, + softmax_lse.T.contiguous(), + o_expand, + softmax_lse_expand.T.contiguous(), + ) + else: + o = result else: if ( forward_batch.attn_attend_prefix_cache is not None diff --git a/python/sglang/srt/hardware_backend/musa/layers/utils/__init__.py b/python/sglang/srt/hardware_backend/musa/layers/utils/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/srt/hardware_backend/musa/layers/utils/cp_utils.py b/python/sglang/srt/hardware_backend/musa/layers/utils/cp_utils.py new file mode 100644 index 000000000..d1eadc659 --- /dev/null +++ b/python/sglang/srt/hardware_backend/musa/layers/utils/cp_utils.py @@ -0,0 +1,57 @@ +from typing import TYPE_CHECKING, Callable + +import torch + +if TYPE_CHECKING: + from sglang.srt.hardware_backend.musa.attention.flashattention_backend import ( + MusaFlashAttentionBackend, + ) + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + + +def musa_cp_attn_forward_extend( + musa_fa_backend: "MusaFlashAttentionBackend", + forward_batch: "ForwardBatch", + q: torch.Tensor, + device: torch.device, + attn_fn: Callable[[torch.Tensor, torch.Tensor, torch.Tensor, int], torch.Tensor], +) -> torch.Tensor: + """ + Split q into prev/next zigzag halves based on CP metadata, call the + backend-specific attention function twice with appropriate per-half + metadata, and concatenate the results. + + attn_fn signature: + attn_fn(q, cu_seqlens_q, cache_seqlens, max_seqlen_q) -> result + where only these four CP-varying parameters differ between halves. + All other backend-specific args should be captured in the closure. + """ + cp_meta = forward_batch.attn_cp_metadata + + q_prev, q_next = torch.chunk(q, 2, dim=0) + + cu_seqlens_q_prev = torch.tensor( + [0, cp_meta.actual_seq_q_prev], device=device, dtype=torch.int32 + ) + if hasattr(musa_fa_backend, "_current_prefix"): + musa_fa_backend._current_prefix = "forward_extend_cp_prev" + result_prev = attn_fn( + q_prev, + cu_seqlens_q_prev, + cp_meta.kv_len_prev_tensor, + cp_meta.actual_seq_q_prev, + ) + + cu_seqlens_q_next = torch.tensor( + [0, cp_meta.actual_seq_q_next], device=device, dtype=torch.int32 + ) + if hasattr(musa_fa_backend, "_current_prefix"): + musa_fa_backend._current_prefix = "forward_extend_cp_next" + result_next = attn_fn( + q_next, + cu_seqlens_q_next, + cp_meta.kv_len_next_tensor, + cp_meta.actual_seq_q_next, + ) + + return torch.concat([result_prev, result_next], dim=0) diff --git a/sgl-kernel/pyproject_musa.toml b/sgl-kernel/pyproject_musa.toml index 5692b48db..99584d8d3 100644 --- a/sgl-kernel/pyproject_musa.toml +++ b/sgl-kernel/pyproject_musa.toml @@ -3,7 +3,7 @@ requires = [ "setuptools>=75.0", "scikit-build-core>=0.10", "torch", - "torchada>=0.1.45", + "torchada>=0.1.50", "wheel", ] build-backend = "setuptools.build_meta"