[MUSA] Resolve output garbage in Context Parallel on MusaFlashAttentionBackend (#23270)

Co-authored-by: zhiguo.qin <zhiguo.qin@mthreads.com>
This commit is contained in:
ori
2026-04-22 20:22:20 -07:00
committed by GitHub
co-authored by zhiguo.qin
parent 6490afe36e
commit 887d380ace
6 changed files with 128 additions and 61 deletions
+4 -1
View File
@@ -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",
]
+9 -9
View File
@@ -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 = [
@@ -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
@@ -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)
+1 -1
View File
@@ -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"