From 71e89e9003f51715869a60a9387cce71acc73727 Mon Sep 17 00:00:00 2001 From: ori <39351881+froststeam@users.noreply.github.com> Date: Fri, 1 May 2026 02:26:47 +0800 Subject: [PATCH] [MUSA][19/N] Support qwen series models (#23654) Co-authored-by: zhiguo.qin --- 3rdparty/amd/wheel/sglang/pyproject.toml | 2 +- python/pyproject_other.toml | 2 +- .../musa/attention/flashattention_backend.py | 42 ++- .../srt/hardware_backend/musa/kernels/topk.py | 300 ++++++++++++++++++ python/sglang/srt/layers/attention/vision.py | 14 +- python/sglang/srt/layers/moe/topk.py | 28 +- .../srt/layers/quantization/fp8_kernel.py | 5 - .../srt/layers/quantization/fp8_utils.py | 20 +- .../srt/layers/rotary_embedding/base.py | 3 +- .../sglang/srt/layers/utils/multi_platform.py | 5 +- python/sglang/srt/server_args.py | 4 +- python/sglang/srt/speculative/dflash_utils.py | 4 +- python/sglang/srt/speculative/draft_utils.py | 25 +- python/sglang/srt/speculative/eagle_info.py | 4 +- .../sglang/srt/speculative/eagle_info_v2.py | 7 +- python/sglang/srt/speculative/eagle_worker.py | 1 + .../sglang/srt/speculative/eagle_worker_v2.py | 6 +- python/sglang/srt/speculative/ngram_info.py | 4 +- python/sglang/srt/speculative/spec_utils.py | 7 +- sgl-kernel/pyproject_musa.toml | 2 +- 20 files changed, 426 insertions(+), 59 deletions(-) create mode 100644 python/sglang/srt/hardware_backend/musa/kernels/topk.py diff --git a/3rdparty/amd/wheel/sglang/pyproject.toml b/3rdparty/amd/wheel/sglang/pyproject.toml index 323985c89..df1f27f58 100644 --- a/3rdparty/amd/wheel/sglang/pyproject.toml +++ b/3rdparty/amd/wheel/sglang/pyproject.toml @@ -123,7 +123,7 @@ srt_musa = [ "sglang[runtime_common]", "torch", "torch_musa", - "torchada>=0.1.50", + "torchada>=0.1.53", "mthreads-ml-py", "mate>=0.2.0", "deep-gemm>=0.1.3", diff --git a/python/pyproject_other.toml b/python/pyproject_other.toml index ee27fc1ae..46cd08aa0 100755 --- a/python/pyproject_other.toml +++ b/python/pyproject_other.toml @@ -115,7 +115,7 @@ srt_musa = [ "sglang[runtime_common]", "torch", "torch_musa", - "torchada>=0.1.50", + "torchada>=0.1.53", "mthreads-ml-py", "mate>=0.2.0", "deep-gemm>=0.1.3", 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 04dd7d6a0..17fb35ae6 100644 --- a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py +++ b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py @@ -15,6 +15,7 @@ from sglang.srt.hardware_backend.musa.layers.utils.cp_utils import ( ) from sglang.srt.layers.attention.flashattention_backend import ( FlashAttentionBackend, + FlashAttentionMultiStepBackend, merge_state_v2_wrapper, ) from sglang.srt.layers.radix_attention import AttentionType @@ -235,17 +236,6 @@ class MusaFlashAttentionBackend(FlashAttentionBackend): self._current_max_seqlen_k = max_seqlen_k self._current_can_run_tbo = can_run_tbo - def init_forward_metadata(self, forward_batch: ForwardBatch): - super().init_forward_metadata(forward_batch) - metadata = self.forward_metadata - if not hasattr(metadata, "extend_with_prefix"): - metadata.extend_with_prefix = False - - if forward_batch.forward_mode.is_extend_or_draft_extend_or_mixed( - include_draft_extend_v2=True - ): - metadata.extend_with_prefix = any(forward_batch.extend_prefix_lens_cpu) - def forward_extend( self, q: torch.Tensor, @@ -420,7 +410,10 @@ class MusaFlashAttentionBackend(FlashAttentionBackend): _fa_cp_attn, ) elif ( - metadata.extend_with_prefix + ( + forward_batch.extend_prefix_lens_cpu is not None + and any(forward_batch.extend_prefix_lens_cpu) + ) or forward_batch.forward_mode.is_target_verify() or forward_batch.forward_mode.is_draft_extend() ): @@ -918,3 +911,28 @@ class MusaFlashAttentionBackend(FlashAttentionBackend): o = result return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) + + +class MusaFlashAttentionMultiStepBackend(FlashAttentionMultiStepBackend): + + def __init__( + self, + model_runner: ModelRunner, + topk: int, + speculative_num_steps: int, + fa_impl_ver: int = 3, + ): + self.model_runner = model_runner + self.topk = topk + self.speculative_num_steps = speculative_num_steps + self.attn_backends = [] + for i in range(self.speculative_num_steps - 1): + self.attn_backends.append( + MusaFlashAttentionBackend( + model_runner, + speculative_step_id=i, + topk=self.topk, + speculative_num_steps=self.speculative_num_steps, + fa_impl_ver=fa_impl_ver, + ) + ) diff --git a/python/sglang/srt/hardware_backend/musa/kernels/topk.py b/python/sglang/srt/hardware_backend/musa/kernels/topk.py new file mode 100644 index 000000000..9785a10d7 --- /dev/null +++ b/python/sglang/srt/hardware_backend/musa/kernels/topk.py @@ -0,0 +1,300 @@ +from typing import ( + Optional, +) + +import torch +import triton +import triton.language as tl + + +@triton.jit +def tanh(x): + # Tanh is just a scaled sigmoid + return 2 * tl.sigmoid(2 * x) - 1 + + +@triton.autotune( + configs=[ + triton.Config({}, num_warps=1, num_stages=1), + triton.Config({}, num_warps=1, num_stages=2), + triton.Config({}, num_warps=2, num_stages=1), + triton.Config({}, num_warps=2, num_stages=2), + triton.Config({}, num_warps=4, num_stages=1), + triton.Config({}, num_warps=4, num_stages=2), + triton.Config({}, num_warps=4, num_stages=3), + triton.Config({}, num_warps=8, num_stages=1), + triton.Config({}, num_warps=8, num_stages=2), + triton.Config({}, num_warps=8, num_stages=3), + triton.Config({}, num_warps=16, num_stages=1), + triton.Config({}, num_warps=16, num_stages=2), + triton.Config({}, num_warps=16, num_stages=3), + triton.Config({}, num_warps=32, num_stages=1), + triton.Config({}, num_warps=32, num_stages=2), + ], + key=["num_tokens", "num_experts", "has_correction_bias"], +) +@triton.jit +def topk_softmax_triton_kernel( + gating_output_ptr, + selected_expert_ptr, + moe_weights_ptr, + renormalize_flag, + num_experts, + num_tokens, # for autotune key + moe_softcapping, + correction_bias_ptr, + has_correction_bias: tl.constexpr, + K: tl.constexpr, + BLOCK_K: tl.constexpr, + BLOCK_WIDTH_SIZE_UP: tl.constexpr, +): + curr_row_idx = tl.program_id(0) + + FLOAT_MINIMUM = -10000.0 + LOG2E = 1.4426950408889634 + + weights_local_final = tl.zeros((BLOCK_K,), dtype=tl.float32) + selected_local_final = tl.zeros((BLOCK_K,), dtype=tl.int32) + + offset = tl.arange(0, BLOCK_WIDTH_SIZE_UP) + k_offset = tl.arange(0, BLOCK_K) + mask_expert = offset < num_experts + mask_topk = k_offset < K + + row_offset = curr_row_idx * num_experts + + logits = tl.load( + gating_output_ptr + row_offset + offset, mask=mask_expert, other=FLOAT_MINIMUM + ) + logits = tl.cast(logits, tl.float32) + + if has_correction_bias: + bias = tl.load(correction_bias_ptr + offset, mask=mask_expert, other=0.0) + logits = logits + bias + + if moe_softcapping > 0.0: + logits = moe_softcapping * tanh(logits / moe_softcapping) + + row_max = tl.max(logits, axis=0) + probs = tl.exp2((logits - row_max) * LOG2E) + row_sum = tl.sum(probs, axis=0) + inv_row_sum = 1.0 / row_sum + probs = probs * inv_row_sum + probs = tl.where(mask_expert, probs, FLOAT_MINIMUM) + + weights_selected_sum = 0.0 + for k_idx in range(K): + top_k_index = tl.argmax(probs, axis=0) + mask = offset == top_k_index + top_k_value = tl.sum(tl.where(mask, probs, 0.0)) + + weights_local_final = tl.where( + k_offset == k_idx, top_k_value, weights_local_final + ) + selected_local_final = tl.where( + k_offset == k_idx, top_k_index, selected_local_final + ) + weights_selected_sum += top_k_value + + probs = tl.where(offset == top_k_index, FLOAT_MINIMUM, probs) + + if renormalize_flag: + weights_local_final = weights_local_final / weights_selected_sum + + tl.store( + moe_weights_ptr + curr_row_idx * K + k_offset, + weights_local_final, + mask=mask_topk, + ) + tl.store( + selected_expert_ptr + curr_row_idx * K + k_offset, + selected_local_final, + mask=mask_topk, + ) + + +def topk_softmax( + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + gating_output: torch.Tensor, + renormalize: bool = False, + moe_softcapping: float = 0.0, + correction_bias: Optional[torch.Tensor] = None, +) -> None: + """ + Compute top-k softmax for MoE routing. + + Args: + topk_weights: Output tensor for top-k weights [num_tokens, topk] + topk_ids: Output tensor for top-k expert indices [num_tokens, topk] + gating_output: Gating logits [num_tokens, num_experts] + renormalize: Whether to renormalize the top-k weights + moe_softcapping: Tanh softcapping value (0.0 to disable) + correction_bias: Per-expert bias correction [num_experts], must be float32 if provided + """ + + num_tokens, num_experts = gating_output.shape + topk = topk_weights.shape[-1] + has_correction_bias = correction_bias is not None + + block_width_up = triton.next_power_of_2(num_experts) + grid = (num_tokens,) + + topk_softmax_triton_kernel[grid]( + gating_output, + topk_ids, + topk_weights, + renormalize, + num_experts, + num_tokens, + moe_softcapping, + correction_bias, + has_correction_bias, + K=topk, + BLOCK_K=triton.next_power_of_2(topk), + BLOCK_WIDTH_SIZE_UP=block_width_up, + ) + + +@triton.autotune( + configs=[ + triton.Config({}, num_warps=1, num_stages=1), + triton.Config({}, num_warps=1, num_stages=2), + triton.Config({}, num_warps=2, num_stages=1), + triton.Config({}, num_warps=2, num_stages=2), + triton.Config({}, num_warps=4, num_stages=1), + triton.Config({}, num_warps=4, num_stages=2), + triton.Config({}, num_warps=4, num_stages=3), + triton.Config({}, num_warps=8, num_stages=1), + triton.Config({}, num_warps=8, num_stages=2), + triton.Config({}, num_warps=8, num_stages=3), + triton.Config({}, num_warps=16, num_stages=1), + triton.Config({}, num_warps=16, num_stages=2), + triton.Config({}, num_warps=16, num_stages=3), + triton.Config({}, num_warps=32, num_stages=1), + triton.Config({}, num_warps=32, num_stages=2), + ], + key=["num_tokens", "num_experts"], +) +@triton.jit +def topk_sigmoid_triton_kernel( + gating_output_ptr, + selected_expert_ptr, + moe_weights_ptr, + renormalize_flag, + correction_bias_ptr, + has_correction_bias: tl.constexpr, + num_experts, + num_tokens, # for autotune key + K: tl.constexpr, + BLOCK_K: tl.constexpr, + BLOCK_WIDTH_SIZE_UP: tl.constexpr, +): + curr_row_idx = tl.program_id(0) + + FLOAT_MINIMUM = -10000.0 + LOG2E = 1.4426950408889634 + + weights_local_final = tl.zeros((BLOCK_K,), dtype=tl.float32) + selected_local_final = tl.zeros((BLOCK_K,), dtype=tl.int32) + + offset = tl.arange(0, BLOCK_WIDTH_SIZE_UP) + k_offset = tl.arange(0, BLOCK_K) + mask_expert = offset < num_experts + mask_topk = k_offset < K + + row_offset = curr_row_idx * num_experts + + x = tl.load( + gating_output_ptr + row_offset + offset, mask=mask_expert, other=FLOAT_MINIMUM + ) + x = tl.cast(x, tl.float32) + + # Compute sigmoid(x) + is_positive = x >= 0 + neg_x = tl.where(is_positive, -x, x) + exp_neg_x = tl.exp2(neg_x * LOG2E) + probs = tl.where( + is_positive, + 1.0 / (1.0 + exp_neg_x), + exp_neg_x / (1.0 + exp_neg_x), + ) + + if has_correction_bias: + bias = tl.load(correction_bias_ptr + offset, mask=mask_expert, other=0.0) + probs_for_choice = probs + bias + else: + probs_for_choice = probs + + probs_for_choice = tl.where(mask_expert, probs_for_choice, FLOAT_MINIMUM) + + weights_selected_sum = 0.0 + for k_idx in range(K): + top_k_index = tl.argmax(probs_for_choice, axis=0) + mask = offset == top_k_index + top_k_value = tl.sum(tl.where(mask, probs, 0.0)) + + weights_local_final = tl.where( + k_offset == k_idx, top_k_value, weights_local_final + ) + selected_local_final = tl.where( + k_offset == k_idx, top_k_index, selected_local_final + ) + weights_selected_sum += top_k_value + + probs_for_choice = tl.where( + offset == top_k_index, FLOAT_MINIMUM, probs_for_choice + ) + + if renormalize_flag: + weights_local_final = weights_local_final / weights_selected_sum + + tl.store( + moe_weights_ptr + curr_row_idx * K + k_offset, + weights_local_final, + mask=mask_topk, + ) + tl.store( + selected_expert_ptr + curr_row_idx * K + k_offset, + selected_local_final, + mask=mask_topk, + ) + + +def topk_sigmoid( + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + gating_output: torch.Tensor, + renormalize: bool = False, + correction_bias: Optional[torch.Tensor] = None, +) -> None: + """ + Compute top-k sigmoid for MoE routing. + + Args: + topk_weights: Output tensor for top-k weights [num_tokens, topk] + topk_ids: Output tensor for top-k expert indices [num_tokens, topk] + gating_output: Gating logits [num_tokens, num_experts] + renormalize: Whether to renormalize the top-k weights + correction_bias: Per-expert bias correction [num_experts], must be float32 if provided + """ + num_tokens, num_experts = gating_output.shape + topk = topk_weights.shape[-1] + has_correction_bias = correction_bias is not None + + block_width_up = triton.next_power_of_2(num_experts) + grid = (num_tokens,) + + topk_sigmoid_triton_kernel[grid]( + gating_output, + topk_ids, + topk_weights, + renormalize, + correction_bias, + has_correction_bias, + num_experts, + num_tokens, + K=topk, + BLOCK_K=triton.next_power_of_2(topk), + BLOCK_WIDTH_SIZE_UP=block_width_up, + ) diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py index 29ccd7606..6791d0c79 100644 --- a/python/sglang/srt/layers/attention/vision.py +++ b/python/sglang/srt/layers/attention/vision.py @@ -22,6 +22,7 @@ from sglang.srt.utils import ( is_blackwell_supported, is_cuda, is_hip, + is_musa, is_npu, is_xpu, print_info_once, @@ -32,6 +33,7 @@ from sglang.srt.utils.multi_stream_utils import ( ) _is_cuda = is_cuda() +_is_musa = is_musa() _is_npu = is_npu() _is_hip = is_hip() _is_xpu = is_xpu() @@ -43,6 +45,9 @@ if _is_cuda: flash_attn_varlen_func, ) +if _is_musa: + from flash_attn_interface import flash_attn_varlen_func + if _is_npu: import torch_npu @@ -382,8 +387,8 @@ class VisionFlash3Attention(nn.Module): self, **kwargs, ): - if not _is_cuda: - raise Exception("VisionFlash3Attention is only available for cuda") + if not (_is_cuda or _is_musa): + raise Exception("VisionFlash3Attention is only available for cuda or musa") super().__init__() use_data_parallel = ( kwargs["use_data_parallel"] if "use_data_parallel" in kwargs else False @@ -950,6 +955,11 @@ class VisionAttention(nn.Module): backend = "fa4" else: backend = "triton_attn" + elif _is_musa: + if get_device_capability() >= (3, 1): + backend = "fa3" + else: + backend = "triton_attn" elif _is_hip: if get_device_capability() >= (9, 4) and _use_aiter: backend = "aiter_attn" diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index 02fac83a0..2eba4f094 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -83,7 +83,7 @@ _is_xpu = is_xpu() _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip _is_musa = is_musa() -if _is_cuda or _is_musa: +if _is_cuda: from sgl_kernel import moe_fused_gate try: @@ -126,7 +126,7 @@ if _is_cuda or _is_musa: except ImportError as e: pass -if _is_cuda or _is_hip or _is_xpu or _is_musa: +if _is_cuda or _is_hip or _is_xpu: from sgl_kernel import topk_softmax try: @@ -139,6 +139,13 @@ if _use_aiter: from aiter.fused_moe import fused_topk as aiter_fused_topk except ImportError: raise ImportError("aiter is required when SGLANG_USE_AITER is set to True") +if _is_musa: + try: + from mate import moe_fused_gate + except ImportError as e: + raise ImportError("mate is required for the biased grouped topk.") + + from sglang.srt.hardware_backend.musa.kernels.topk import topk_sigmoid, topk_softmax # -------------------------------- TopKConfig --------------------------------------- @@ -864,7 +871,7 @@ def biased_grouped_topk_gpu( return topk_weights, topk_ids elif ( - (_is_cuda or _is_musa) + _is_cuda # moe_fused_gate kernel ensures that num_experts/num_expert_group does not exceed MAX_VPT=32 now. And when kernel can handle MAX_VPT > 32, we can remove this assertion. and experts_per_group <= 32 and is_power_of_two(num_experts) @@ -902,6 +909,21 @@ def biased_grouped_topk_gpu( routed_scaling_factor if routed_scaling_factor is not None else 1.0, ) return topk_weights, topk_ids + elif _is_musa and ( + gating_output.shape[1] // num_expert_group <= 32 + or (num_expert_group == 1 and gating_output.shape[1] in {160, 256, 384}) + ): + topk_weights, topk_ids = moe_fused_gate( + gating_output.to(dtype=torch.float32), + correction_bias, + num_expert_group, + topk_group, + topk, + num_fused_shared_experts, + routed_scaling_factor if routed_scaling_factor is not None else 1.0, + True, + apply_routed_scaling_factor_on_output, + ) else: # Use optimized path for Kimi K2 (384 experts with num_expert_group=1) num_experts = gating_output.shape[1] diff --git a/python/sglang/srt/layers/quantization/fp8_kernel.py b/python/sglang/srt/layers/quantization/fp8_kernel.py index bc83b166a..ac3498657 100644 --- a/python/sglang/srt/layers/quantization/fp8_kernel.py +++ b/python/sglang/srt/layers/quantization/fp8_kernel.py @@ -1113,11 +1113,6 @@ def w8a8_block_fp8_matmul_deepgemm( # Deepgemm only supports output tensor type as bfloat16 assert C.dtype == torch.bfloat16 and deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM - if _is_musa: - # XXX (MUSA): `deep_gemm_fp8_fp8_bf16_nt` on MUSA requires contiguous tensors - As = As.contiguous() - Bs = Bs.contiguous() - deep_gemm_fp8_fp8_bf16_nt(A, As, B, Bs, C) return C diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index 0f178877f..d4073c29b 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -675,13 +675,19 @@ def deepgemm_w8a8_block_fp8_linear_with_fallback( input_2d = input.view(-1, input.shape[-1]) output_shape = [*input.shape[:-1], weight.shape[0]] - q_input, x_scale = sglang_per_token_group_quant_fp8( - input_2d, - block_size[1], - column_major_scales=True, - scale_tma_aligned=True, - scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0, - ) + if not _is_musa: + q_input, x_scale = sglang_per_token_group_quant_fp8( + input_2d, + block_size[1], + column_major_scales=True, + scale_tma_aligned=True, + scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0, + ) + else: + q_input, x_scale = sglang_per_token_group_quant_fp8( + input_2d, + block_size[1], + ) output = w8a8_block_fp8_matmul_deepgemm( q_input, weight, x_scale, weight_scale, block_size, output_dtype=output_dtype diff --git a/python/sglang/srt/layers/rotary_embedding/base.py b/python/sglang/srt/layers/rotary_embedding/base.py index 2ccfdddfc..2b13c1594 100644 --- a/python/sglang/srt/layers/rotary_embedding/base.py +++ b/python/sglang/srt/layers/rotary_embedding/base.py @@ -99,7 +99,8 @@ class RotaryEmbedding(MultiPlatformOp): self._apply_rotary_emb_wrapped = apply_rotary_emb - if get_global_server_args().rl_on_policy_target is not None: + # XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend + if get_global_server_args().rl_on_policy_target is not None or _is_musa: self._forward_method = self.forward_native self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)( apply_rotary_emb diff --git a/python/sglang/srt/layers/utils/multi_platform.py b/python/sglang/srt/layers/utils/multi_platform.py index d482b4530..893248a21 100644 --- a/python/sglang/srt/layers/utils/multi_platform.py +++ b/python/sglang/srt/layers/utils/multi_platform.py @@ -98,10 +98,7 @@ class MultiPlatformOp(nn.Module): return self.forward_native(*args, **kwargs) def forward_musa(self, *args, **kwargs): - # XXX (MUSA): MUSA kernels follow the CUDA path by default. - # At this stage, sgl-kernel support for MUSA is still under active - # development, so we fall back to the PyTorch-native implementation. - return self.forward_native(*args, **kwargs) + return self.forward_cuda(*args, **kwargs) def forward_hpu(self, *args, **kwargs): return self.forward_native(*args, **kwargs) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index dc4fd7dd2..aec174ee1 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2391,8 +2391,8 @@ class ServerArgs: ) assert ( - is_cuda() - ), "Mamba extra_buffer is only supported on CUDA devices with FLA backend" + is_cuda() or is_musa() + ), "Mamba extra_buffer is only supported on CUDA and MUSA devices with FLA backend" if self.speculative_num_draft_tokens is not None: assert ( self.mamba_track_interval >= self.speculative_num_draft_tokens diff --git a/python/sglang/srt/speculative/dflash_utils.py b/python/sglang/srt/speculative/dflash_utils.py index 4f0d10733..2d7963532 100644 --- a/python/sglang/srt/speculative/dflash_utils.py +++ b/python/sglang/srt/speculative/dflash_utils.py @@ -8,7 +8,7 @@ import torch import torch.nn.functional as F from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod -from sglang.srt.utils import is_cuda +from sglang.srt.utils import is_cuda, is_musa DEFAULT_DFLASH_MASK_TOKEN = "<|MASK|>" @@ -25,7 +25,7 @@ _DFLASH_VERIFY_SKIP_CUSTOM_MASK_BACKENDS = frozenset( ) -if is_cuda(): +if is_cuda() or is_musa(): try: from sgl_kernel import ( top_k_renorm_prob, diff --git a/python/sglang/srt/speculative/draft_utils.py b/python/sglang/srt/speculative/draft_utils.py index 271dc00d0..94c85ac90 100644 --- a/python/sglang/srt/speculative/draft_utils.py +++ b/python/sglang/srt/speculative/draft_utils.py @@ -1,7 +1,7 @@ import logging from sglang.srt.server_args import ServerArgs, get_global_server_args -from sglang.srt.utils.common import is_blackwell +from sglang.srt.utils.common import is_blackwell, is_musa logger = logging.getLogger(__name__) @@ -142,9 +142,14 @@ class DraftBackendFactory: ) def _create_fa_decode_backend(self, fa_impl_ver: int = 3): - from sglang.srt.layers.attention.flashattention_backend import ( - FlashAttentionMultiStepBackend, - ) + if not is_musa(): + from sglang.srt.layers.attention.flashattention_backend import ( + FlashAttentionMultiStepBackend, + ) + else: + from sglang.srt.hardware_backend.musa.attention.flashattention_backend import ( + MusaFlashAttentionMultiStepBackend as FlashAttentionMultiStepBackend, + ) return FlashAttentionMultiStepBackend( self.draft_model_runner, @@ -225,10 +230,14 @@ class DraftBackendFactory: return AiterAttnBackend(self.draft_model_runner, skip_prefill=False) def _create_fa_prefill_backend(self, fa_impl_ver: int = 3): - from sglang.srt.layers.attention.flashattention_backend import ( - FlashAttentionBackend, - ) - + if not is_musa(): + from sglang.srt.layers.attention.flashattention_backend import ( + FlashAttentionBackend, + ) + else: + from sglang.srt.hardware_backend.musa.attention.flashattention_backend import ( + MusaFlashAttentionBackend as FlashAttentionBackend, + ) return FlashAttentionBackend( self.draft_model_runner, skip_prefill=False, fa_impl_ver=fa_impl_ver ) diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index 26cb028a1..91932a92f 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -44,9 +44,9 @@ from sglang.srt.speculative.spec_utils import ( get_src_tgt_cache_loc, get_target_cache_loc, ) -from sglang.srt.utils import is_cuda, next_power_of_2 +from sglang.srt.utils import is_cuda, is_musa, next_power_of_2 -if is_cuda(): +if is_cuda() or is_musa(): from sgl_kernel import ( top_k_renorm_prob, top_p_renorm_prob, diff --git a/python/sglang/srt/speculative/eagle_info_v2.py b/python/sglang/srt/speculative/eagle_info_v2.py index 4c7535ced..bc9f35e67 100644 --- a/python/sglang/srt/speculative/eagle_info_v2.py +++ b/python/sglang/srt/speculative/eagle_info_v2.py @@ -35,11 +35,12 @@ from sglang.srt.speculative.spec_utils import ( SIMULATE_ACC_LEN, generate_simulated_accept_index, ) -from sglang.srt.utils.common import is_cuda, is_hip, is_npu, next_power_of_2 +from sglang.srt.utils.common import is_cuda, is_hip, is_musa, is_npu, next_power_of_2 _is_cuda = is_cuda() _is_hip = is_hip() _is_npu = is_npu() +_is_musa = is_musa() if TYPE_CHECKING: from sglang.srt.managers.tp_worker import TpModelWorker @@ -48,7 +49,7 @@ if TYPE_CHECKING: ) from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput -if is_cuda(): +if is_cuda() or is_musa(): from sgl_kernel import ( top_k_renorm_prob, top_p_renorm_prob, @@ -539,7 +540,7 @@ def assign_extend_cache_locs_func( draft_token_num: int, device, ) -> torch.Tensor: - if _is_cuda or _is_hip: + if _is_cuda or _is_hip or _is_musa: out_cache_loc = torch.empty( (batch_size * draft_token_num,), dtype=torch.int64, diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index 79ea70ed0..4eb122556 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -272,6 +272,7 @@ class EAGLEWorker(TpModelWorker): Device2DraftCudaGraphRunner = { "npu": EAGLEDraftNpuGraphRunner, "cuda": EAGLEDraftCudaGraphRunner, + "musa": EAGLEDraftCudaGraphRunner, } # Capture draft if self.speculative_num_steps > 1: diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 6f477f4cf..3fa43aa37 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -63,6 +63,7 @@ from sglang.srt.utils.common import ( get_available_gpu_memory, is_cuda, is_hip, + is_musa, is_npu, next_power_of_2, ) @@ -70,6 +71,7 @@ from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions _is_npu = is_npu() _is_cuda = is_cuda() +_is_musa = is_musa() _is_hip = is_hip() logger = logging.getLogger(__name__) @@ -265,6 +267,7 @@ class EagleDraftWorker(BaseDraftWorker): Device2DraftCudaGraphRunner = { "npu": EAGLEDraftNpuGraphRunner, "cuda": EAGLEDraftCudaGraphRunner, + "musa": EAGLEDraftCudaGraphRunner, } # Capture draft if self.speculative_num_steps > 1: @@ -284,6 +287,7 @@ class EagleDraftWorker(BaseDraftWorker): Device2ExtendCudaGraphRunner = { "npu": EAGLEDraftExtendNpuGraphRunner, "cuda": EAGLEDraftExtendCudaGraphRunner, + "musa": EAGLEDraftCudaGraphRunner, } supports_hip_aiter_draft_extend_graph = False if _is_hip: @@ -296,7 +300,7 @@ class EagleDraftWorker(BaseDraftWorker): self.draft_attn_backend, AiterMultiStepDraftBackend ) - supports_cuda_draft_extend_graph = _is_cuda and ( + supports_cuda_draft_extend_graph = (_is_cuda or _is_musa) and ( isinstance(self.draft_extend_attn_backend, TritonAttnBackend) or isinstance(self.draft_extend_attn_backend, TRTLLMMLABackend) ) diff --git a/python/sglang/srt/speculative/ngram_info.py b/python/sglang/srt/speculative/ngram_info.py index efb389e5f..4a68b1731 100644 --- a/python/sglang/srt/speculative/ngram_info.py +++ b/python/sglang/srt/speculative/ngram_info.py @@ -34,9 +34,9 @@ from sglang.srt.speculative.spec_utils import ( get_src_tgt_cache_loc, get_target_cache_loc, ) -from sglang.srt.utils import is_cuda, is_hip, next_power_of_2 +from sglang.srt.utils import is_cuda, is_hip, is_musa, next_power_of_2 -if is_cuda(): +if is_cuda() or is_musa(): from sgl_kernel import ( top_k_renorm_prob, top_p_renorm_prob, diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 2833f1970..46f6c00b1 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -20,11 +20,12 @@ from sglang.srt.environ import envs from sglang.srt.managers.schedule_batch import Req from sglang.srt.mem_cache.common import get_last_loc from sglang.srt.server_args import ServerArgs, get_global_server_args -from sglang.srt.utils import is_cuda, is_hip, is_npu, next_power_of_2 +from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu, next_power_of_2 _is_cuda = is_cuda() _is_hip = is_hip() _is_npu = is_npu() +_is_musa = is_musa() if TYPE_CHECKING: from sglang.srt.speculative.eagle_info import EagleVerifyInput @@ -46,7 +47,9 @@ SIMULATE_ACC_LEN = envs.SGLANG_SIMULATE_ACC_LEN.get() # turn off if < 0 SIMULATE_ACC_METHOD = envs.SGLANG_SIMULATE_ACC_METHOD.get() TREE_TRAVERSE_TIME_THRESHOLD = 1 # TODO: set this properly -TREE_SPEC_KERNEL_AVAILABLE = _is_cuda # This kernel is only available for CUDA now +TREE_SPEC_KERNEL_AVAILABLE = ( + _is_cuda or _is_musa +) # This kernel is only available for CUDA and MUSA now def spec_need_hidden_states(server_args: Optional[ServerArgs] = None) -> bool: diff --git a/sgl-kernel/pyproject_musa.toml b/sgl-kernel/pyproject_musa.toml index a1a1230c1..57ddc0cc4 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.50", + "torchada>=0.1.53", "wheel", ] build-backend = "setuptools.build_meta"