[MoE Refactor] deprecate forward_npu and NpuFuseEPMoE (#25678)
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.7
parent
a24c374f84
commit
b765faee30
@@ -0,0 +1,171 @@
|
|||||||
|
"""Ascend FuseEP fused dispatch+GEMM+combine forward path.
|
||||||
|
|
||||||
|
Follows the mega_moe shape: a free-function bypass invoked from
|
||||||
|
``FusedMoE.forward`` when ``--moe-a2a-backend ascend_fuseep`` is set, plus a
|
||||||
|
weight-postprocess helper that NPU quant_methods call from their
|
||||||
|
``process_weights_after_loading`` when the same backend is selected.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.distributed import get_tp_group
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.hardware_backend.npu.utils import FusedMoEMode, npu_format_cast
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer
|
||||||
|
from sglang.srt.layers.moe.utils import DeepEPMode
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||||
|
from sglang.srt.layers.moe.topk import TopKOutput
|
||||||
|
|
||||||
|
|
||||||
|
_PARAMS_BYTES = 2 # bf16 — Ascend's Dispatch & Combine does not support fp16
|
||||||
|
|
||||||
|
|
||||||
|
def _get_fuseep_buffer(layer: "FusedMoE"):
|
||||||
|
DeepEPBuffer.set_dispatch_mode_as_low_latency()
|
||||||
|
return DeepEPBuffer.get_deepep_buffer(
|
||||||
|
get_tp_group().device_group,
|
||||||
|
layer.hidden_size,
|
||||||
|
_PARAMS_BYTES,
|
||||||
|
DeepEPMode.LOW_LATENCY,
|
||||||
|
envs.SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get(),
|
||||||
|
layer.num_experts,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def forward_fuseep(
|
||||||
|
layer: "FusedMoE",
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
topk_output: "TopKOutput",
|
||||||
|
) -> torch.Tensor:
|
||||||
|
buf = _get_fuseep_buffer(layer)
|
||||||
|
hidden_states, _ = buf.fused_deep_moe(
|
||||||
|
hidden_states,
|
||||||
|
topk_idx=topk_output.topk_ids,
|
||||||
|
topk_weights=topk_output.topk_weights,
|
||||||
|
gmm1_permuted_weight=layer.w13_weight,
|
||||||
|
gmm1_permuted_weight_scale=layer.w13_weight_scale,
|
||||||
|
gmm2_weight=layer.w2_weight,
|
||||||
|
gmm2_weight_scale=layer.w2_weight_scale,
|
||||||
|
num_max_dispatch_tokens_per_rank=(
|
||||||
|
envs.SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get()
|
||||||
|
),
|
||||||
|
num_experts=layer.num_experts,
|
||||||
|
fuse_mode=envs.SGLANG_NPU_FUSED_MOE_MODE.get(),
|
||||||
|
)
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
def _permute_w13_weight_scale(w: torch.Tensor, tile_n: int) -> torch.Tensor:
|
||||||
|
if tile_n % 2 != 0:
|
||||||
|
raise ValueError(f"tile_n must be even, got {tile_n}")
|
||||||
|
|
||||||
|
*dims, n = w.shape
|
||||||
|
if n % tile_n != 0:
|
||||||
|
raise ValueError(f"Last dimension {n} must be divisible by tile_n {tile_n}")
|
||||||
|
|
||||||
|
w_reshaped = w.reshape(*dims, 2, n // tile_n, tile_n // 2)
|
||||||
|
perm_order = list(range(len(dims))) + [-2, -3, -1]
|
||||||
|
return w_reshaped.permute(perm_order).reshape(*dims, n)
|
||||||
|
|
||||||
|
|
||||||
|
def _reshape_w13_weight(
|
||||||
|
weight: torch.Tensor, dim: int, chunk_size: int = 64
|
||||||
|
) -> torch.Tensor:
|
||||||
|
# Achieving greater computing power through reshape on Ascend.
|
||||||
|
original_shape = weight.shape
|
||||||
|
if dim < 0:
|
||||||
|
dim += len(original_shape)
|
||||||
|
|
||||||
|
if original_shape[dim] % (2 * chunk_size) != 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"Dimension {dim} size {original_shape[dim]} must be divisible by "
|
||||||
|
f"{2 * chunk_size}"
|
||||||
|
)
|
||||||
|
|
||||||
|
new_shape = (
|
||||||
|
*original_shape[:dim],
|
||||||
|
2,
|
||||||
|
original_shape[dim] // (2 * chunk_size),
|
||||||
|
chunk_size,
|
||||||
|
*original_shape[dim + 1 :],
|
||||||
|
)
|
||||||
|
|
||||||
|
weight = weight.view(new_shape)
|
||||||
|
weight = weight.transpose(dim, dim + 1).contiguous()
|
||||||
|
return weight.view(*original_shape[:dim], -1, *original_shape[dim + 1 :])
|
||||||
|
|
||||||
|
|
||||||
|
def _release_weight_cache(weight: torch.Tensor) -> torch.Tensor:
|
||||||
|
# .contiguous() introduces additional memory overhead; release with resize_(0)
|
||||||
|
origin_weight = weight.data.transpose(1, 2)
|
||||||
|
new_weight = origin_weight.contiguous()
|
||||||
|
origin_weight.untyped_storage().resize_(0)
|
||||||
|
return new_weight
|
||||||
|
|
||||||
|
|
||||||
|
def _scale_from_float_to_int64(scale: torch.Tensor) -> torch.nn.Parameter:
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
converted = torch.from_numpy(
|
||||||
|
np.frombuffer(
|
||||||
|
scale.cpu().to(torch.float32).numpy().tobytes(), dtype=np.int32
|
||||||
|
).astype(np.int64)
|
||||||
|
).to(scale.device)
|
||||||
|
return torch.nn.Parameter(converted, requires_grad=False)
|
||||||
|
|
||||||
|
|
||||||
|
def process_fuseep_weights(layer: torch.nn.Module) -> None:
|
||||||
|
"""Apply the Ascend FuseEP-specific weight layout.
|
||||||
|
|
||||||
|
Replaces NPU quant_method weight layouts with the form required by the
|
||||||
|
fused_deep_moe op. Invoked from NPU ``process_weights_after_loading``
|
||||||
|
when ``--moe-a2a-backend ascend_fuseep`` is set.
|
||||||
|
"""
|
||||||
|
if envs.SGLANG_NPU_FUSED_MOE_MODE.get() == FusedMoEMode.DISPATCH_FFN_COMBINE.value:
|
||||||
|
w13_weight = _release_weight_cache(layer.w13_weight)
|
||||||
|
layer.w13_weight.data = npu_format_cast(w13_weight)
|
||||||
|
w2_weight = _release_weight_cache(layer.w2_weight)
|
||||||
|
layer.w2_weight.data = npu_format_cast(w2_weight)
|
||||||
|
|
||||||
|
layer.w13_weight_scale.data = layer.w13_weight_scale.data.view(
|
||||||
|
layer.w13_weight_scale.data.shape[0], -1
|
||||||
|
)
|
||||||
|
w2_scale = layer.w2_weight_scale.data.squeeze(-1).contiguous()
|
||||||
|
layer.w2_weight_scale = torch.nn.Parameter(
|
||||||
|
w2_scale.to(torch.float32), requires_grad=False
|
||||||
|
)
|
||||||
|
|
||||||
|
layer.w13_weight_scale = _scale_from_float_to_int64(layer.w13_weight_scale.data)
|
||||||
|
layer.w2_weight_scale = _scale_from_float_to_int64(layer.w2_weight_scale.data)
|
||||||
|
else:
|
||||||
|
cpu_w13 = layer.w13_weight.data.transpose(1, 2).cpu()
|
||||||
|
layer.w13_weight.data = _reshape_w13_weight(cpu_w13, -1).npu()
|
||||||
|
w13_scale = layer.w13_weight_scale.data.squeeze(-1).contiguous()
|
||||||
|
w13_scale = _permute_w13_weight_scale(w13_scale, 128)
|
||||||
|
layer.w13_weight_scale = torch.nn.Parameter(
|
||||||
|
w13_scale.to(torch.float32), requires_grad=False
|
||||||
|
)
|
||||||
|
layer.w13_weight.data = npu_format_cast(layer.w13_weight.data)
|
||||||
|
layer.w2_weight.data = npu_format_cast(layer.w2_weight.data)
|
||||||
|
|
||||||
|
w2_scale = layer.w2_weight_scale.data.squeeze(-1).contiguous()
|
||||||
|
layer.w2_weight_scale = torch.nn.Parameter(
|
||||||
|
w2_scale.to(torch.float32), requires_grad=False
|
||||||
|
)
|
||||||
|
|
||||||
|
if hasattr(layer, "w13_weight_offset"):
|
||||||
|
layer.w13_weight_offset = torch.nn.Parameter(
|
||||||
|
layer.w13_weight_offset.data.squeeze(-1).contiguous(),
|
||||||
|
requires_grad=False,
|
||||||
|
)
|
||||||
|
if hasattr(layer, "w2_weight_offset"):
|
||||||
|
layer.w2_weight_offset = torch.nn.Parameter(
|
||||||
|
layer.w2_weight_offset.data.squeeze(-1).contiguous(),
|
||||||
|
requires_grad=False,
|
||||||
|
)
|
||||||
@@ -9,7 +9,9 @@ from sglang.srt.layers.quantization.base_config import FusedMoEMethodBase
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.moe.token_dispatcher import (
|
from sglang.srt.layers.moe.token_dispatcher import (
|
||||||
CombineInput,
|
CombineInput,
|
||||||
StandardDispatchOutput,
|
DeepEPLLDispatchOutput,
|
||||||
|
DeepEPNormalDispatchOutput,
|
||||||
|
DispatchOutput,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||||
|
|
||||||
@@ -384,6 +386,93 @@ def fused_moe_npu(
|
|||||||
return final_hidden_states
|
return final_hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
def maybe_apply_deepep_npu(
|
||||||
|
quant_method,
|
||||||
|
layer: torch.nn.Module,
|
||||||
|
dispatch_output: "DispatchOutput",
|
||||||
|
) -> Optional["CombineInput"]:
|
||||||
|
"""Route DeepEP dispatch outputs through the NPU compute path.
|
||||||
|
|
||||||
|
Replaces the deprecated DeepEPMoE.forward_npu wrapper: detects DeepEP
|
||||||
|
normal/LL formats, calls ``quant_method.apply_without_routing_weights``,
|
||||||
|
and wraps the result in the matching CombineInput. Returns None for
|
||||||
|
non-DeepEP formats so the caller falls through to its standard path.
|
||||||
|
"""
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher import (
|
||||||
|
DeepEPLLCombineInput,
|
||||||
|
DeepEPNormalCombineInput,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.base import DispatchOutputChecker
|
||||||
|
|
||||||
|
if not dispatch_output.format.is_deepep():
|
||||||
|
return None
|
||||||
|
|
||||||
|
# NOTE: Ascend's Dispatch & Combine does not support FP16
|
||||||
|
output_dtype = torch.bfloat16
|
||||||
|
group_list_type = 1
|
||||||
|
|
||||||
|
if DispatchOutputChecker.format_is_deepep_normal(dispatch_output):
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
assert isinstance(dispatch_output, DeepEPNormalDispatchOutput)
|
||||||
|
(
|
||||||
|
hidden_states,
|
||||||
|
hidden_states_scale,
|
||||||
|
_,
|
||||||
|
_,
|
||||||
|
num_recv_tokens_per_expert,
|
||||||
|
) = dispatch_output
|
||||||
|
group_list = torch.tensor(
|
||||||
|
num_recv_tokens_per_expert,
|
||||||
|
dtype=torch.int64,
|
||||||
|
device=hidden_states.device,
|
||||||
|
)
|
||||||
|
combine_cls = DeepEPNormalCombineInput
|
||||||
|
else:
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
assert isinstance(dispatch_output, DeepEPLLDispatchOutput)
|
||||||
|
(
|
||||||
|
hidden_states,
|
||||||
|
hidden_states_scale,
|
||||||
|
_,
|
||||||
|
_,
|
||||||
|
group_list,
|
||||||
|
_,
|
||||||
|
) = dispatch_output
|
||||||
|
group_list = group_list.to(torch.int64)
|
||||||
|
combine_cls = DeepEPLLCombineInput
|
||||||
|
|
||||||
|
hidden_states = quant_method.apply_without_routing_weights(
|
||||||
|
layer,
|
||||||
|
hidden_states,
|
||||||
|
hidden_states_scale,
|
||||||
|
group_list_type,
|
||||||
|
group_list,
|
||||||
|
output_dtype,
|
||||||
|
)
|
||||||
|
|
||||||
|
return combine_cls(
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
topk_ids=dispatch_output.topk_ids,
|
||||||
|
topk_weights=dispatch_output.topk_weights,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def maybe_apply_fuseep_weights(layer: torch.nn.Module) -> bool:
|
||||||
|
"""Apply the FuseEP weight layout if --moe-a2a-backend is ascend_fuseep.
|
||||||
|
|
||||||
|
Returns True when the FuseEP layout was applied and the caller should
|
||||||
|
skip its own ``process_weights_after_loading`` body.
|
||||||
|
"""
|
||||||
|
from sglang.srt.layers.moe import get_moe_a2a_backend
|
||||||
|
|
||||||
|
if not get_moe_a2a_backend().is_ascend_fuseep():
|
||||||
|
return False
|
||||||
|
from sglang.srt.hardware_backend.npu.moe.fuseep import process_fuseep_weights
|
||||||
|
|
||||||
|
process_fuseep_weights(layer)
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
class _NPUFusedMoEMethodBase(FusedMoEMethodBase):
|
class _NPUFusedMoEMethodBase(FusedMoEMethodBase):
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -392,6 +481,17 @@ class _NPUFusedMoEMethodBase(FusedMoEMethodBase):
|
|||||||
):
|
):
|
||||||
self.quant_config = quant_config
|
self.quant_config = quant_config
|
||||||
|
|
||||||
|
def _maybe_apply_deepep(
|
||||||
|
self,
|
||||||
|
layer: torch.nn.Module,
|
||||||
|
dispatch_output: "DispatchOutput",
|
||||||
|
) -> Optional["CombineInput"]:
|
||||||
|
return maybe_apply_deepep_npu(self, layer, dispatch_output)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _maybe_apply_fuseep_weights(layer: torch.nn.Module) -> bool:
|
||||||
|
return maybe_apply_fuseep_weights(layer)
|
||||||
|
|
||||||
|
|
||||||
class NPUW4A4Int4DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
class NPUW4A4Int4DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
||||||
|
|
||||||
@@ -444,10 +544,14 @@ class NPUW4A4Int4DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
|||||||
def apply(
|
def apply(
|
||||||
self,
|
self,
|
||||||
layer,
|
layer,
|
||||||
dispatch_output: "StandardDispatchOutput",
|
dispatch_output: "DispatchOutput",
|
||||||
) -> "CombineInput":
|
) -> "CombineInput":
|
||||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
||||||
|
|
||||||
|
combine_input = self._maybe_apply_deepep(layer, dispatch_output)
|
||||||
|
if combine_input is not None:
|
||||||
|
return combine_input
|
||||||
|
|
||||||
x = dispatch_output.hidden_states
|
x = dispatch_output.hidden_states
|
||||||
topk_output = dispatch_output.topk_output
|
topk_output = dispatch_output.topk_output
|
||||||
|
|
||||||
@@ -512,6 +616,8 @@ class NPUW4A4Int4DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
|||||||
class NPUW8A8Int8DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
class NPUW8A8Int8DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
||||||
|
|
||||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||||
|
if self._maybe_apply_fuseep_weights(layer):
|
||||||
|
return
|
||||||
layer.w13_weight.data = npu_format_cast(layer.w13_weight.data.transpose(1, 2))
|
layer.w13_weight.data = npu_format_cast(layer.w13_weight.data.transpose(1, 2))
|
||||||
layer.w2_weight.data = npu_format_cast(layer.w2_weight.data.transpose(1, 2))
|
layer.w2_weight.data = npu_format_cast(layer.w2_weight.data.transpose(1, 2))
|
||||||
layer.w13_weight_scale = torch.nn.Parameter(
|
layer.w13_weight_scale = torch.nn.Parameter(
|
||||||
@@ -544,10 +650,14 @@ class NPUW8A8Int8DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
|||||||
def apply(
|
def apply(
|
||||||
self,
|
self,
|
||||||
layer,
|
layer,
|
||||||
dispatch_output: "StandardDispatchOutput",
|
dispatch_output: "DispatchOutput",
|
||||||
) -> "CombineInput":
|
) -> "CombineInput":
|
||||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
||||||
|
|
||||||
|
combine_input = self._maybe_apply_deepep(layer, dispatch_output)
|
||||||
|
if combine_input is not None:
|
||||||
|
return combine_input
|
||||||
|
|
||||||
# release fp32 scale to save memory
|
# release fp32 scale to save memory
|
||||||
layer.w13_weight_scale = None
|
layer.w13_weight_scale = None
|
||||||
layer.w2_weight_scale = None
|
layer.w2_weight_scale = None
|
||||||
@@ -761,10 +871,14 @@ class NPUW4A8Int8DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
|||||||
def apply(
|
def apply(
|
||||||
self,
|
self,
|
||||||
layer,
|
layer,
|
||||||
dispatch_output: "StandardDispatchOutput",
|
dispatch_output: "DispatchOutput",
|
||||||
) -> "CombineInput":
|
) -> "CombineInput":
|
||||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
||||||
|
|
||||||
|
combine_input = self._maybe_apply_deepep(layer, dispatch_output)
|
||||||
|
if combine_input is not None:
|
||||||
|
return combine_input
|
||||||
|
|
||||||
hidden_states = dispatch_output.hidden_states
|
hidden_states = dispatch_output.hidden_states
|
||||||
topk_output = dispatch_output.topk_output
|
topk_output = dispatch_output.topk_output
|
||||||
|
|
||||||
@@ -1020,10 +1134,14 @@ class NPUW4A16Int4DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
|||||||
def apply(
|
def apply(
|
||||||
self,
|
self,
|
||||||
layer,
|
layer,
|
||||||
dispatch_output: "StandardDispatchOutput",
|
dispatch_output: "DispatchOutput",
|
||||||
) -> "CombineInput":
|
) -> "CombineInput":
|
||||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
||||||
|
|
||||||
|
combine_input = self._maybe_apply_deepep(layer, dispatch_output)
|
||||||
|
if combine_input is not None:
|
||||||
|
return combine_input
|
||||||
|
|
||||||
x = dispatch_output.hidden_states
|
x = dispatch_output.hidden_states
|
||||||
topk_output = dispatch_output.topk_output
|
topk_output = dispatch_output.topk_output
|
||||||
|
|
||||||
|
|||||||
@@ -1,13 +1,12 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import TYPE_CHECKING, Any, Dict, Optional, Union
|
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph
|
from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.hardware_backend.npu.utils import FusedMoEMode, npu_format_cast
|
|
||||||
from sglang.srt.layers import deep_gemm_wrapper
|
from sglang.srt.layers import deep_gemm_wrapper
|
||||||
from sglang.srt.layers.moe import (
|
from sglang.srt.layers.moe import (
|
||||||
get_deepep_mode,
|
get_deepep_mode,
|
||||||
@@ -85,7 +84,7 @@ class DeepEPMoE(FusedMoE):
|
|||||||
if _use_aiter:
|
if _use_aiter:
|
||||||
self.deprecate_flag = True
|
self.deprecate_flag = True
|
||||||
elif _is_npu:
|
elif _is_npu:
|
||||||
self.deprecate_flag = False
|
self.deprecate_flag = True
|
||||||
elif deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM and isinstance(
|
elif deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM and isinstance(
|
||||||
quant_config, Fp8Config
|
quant_config, Fp8Config
|
||||||
):
|
):
|
||||||
@@ -203,10 +202,7 @@ class DeepEPMoE(FusedMoE):
|
|||||||
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher import DispatchOutputChecker
|
from sglang.srt.layers.moe.token_dispatcher import DispatchOutputChecker
|
||||||
|
|
||||||
if _is_npu:
|
if DispatchOutputChecker.format_is_deepep_normal(dispatch_output):
|
||||||
assert DispatchOutputChecker.format_is_deepep(dispatch_output)
|
|
||||||
output = self.forward_npu(dispatch_output)
|
|
||||||
elif DispatchOutputChecker.format_is_deepep_normal(dispatch_output):
|
|
||||||
if self.quant_config is None:
|
if self.quant_config is None:
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"Unquantized DeepEP MoE currently supports low_latency mode only"
|
"Unquantized DeepEP MoE currently supports low_latency mode only"
|
||||||
@@ -269,241 +265,6 @@ class DeepEPMoE(FusedMoE):
|
|||||||
dispatch_output=dispatch_output,
|
dispatch_output=dispatch_output,
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward_npu(
|
|
||||||
self,
|
|
||||||
dispatch_output: Union[DeepEPNormalDispatchOutput, DeepEPLLDispatchOutput],
|
|
||||||
):
|
|
||||||
assert self.quant_method is not None
|
|
||||||
assert self.moe_runner_config.activation == "silu"
|
|
||||||
|
|
||||||
from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import (
|
|
||||||
npu_fused_moe_without_routing_weights_bf16,
|
|
||||||
)
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher import DispatchOutputChecker
|
|
||||||
|
|
||||||
# NOTE: Ascend's Dispatch & Combine does not support FP16
|
|
||||||
output_dtype = torch.bfloat16
|
|
||||||
group_list_type = 1
|
|
||||||
|
|
||||||
if DispatchOutputChecker.format_is_deepep_normal(dispatch_output):
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
assert isinstance(dispatch_output, DeepEPNormalDispatchOutput)
|
|
||||||
hidden_states, hidden_states_scale, _, _, num_recv_tokens_per_expert = (
|
|
||||||
dispatch_output
|
|
||||||
)
|
|
||||||
|
|
||||||
group_list = torch.tensor(
|
|
||||||
num_recv_tokens_per_expert,
|
|
||||||
dtype=torch.int64,
|
|
||||||
device=hidden_states.device,
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.w13_weight.dtype == torch.bfloat16:
|
|
||||||
hidden_states = npu_fused_moe_without_routing_weights_bf16(
|
|
||||||
self, hidden_states, group_list_type, group_list, output_dtype
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
hidden_states = self.quant_method.apply_without_routing_weights(
|
|
||||||
self,
|
|
||||||
hidden_states,
|
|
||||||
hidden_states_scale,
|
|
||||||
group_list_type,
|
|
||||||
group_list,
|
|
||||||
output_dtype,
|
|
||||||
)
|
|
||||||
elif DispatchOutputChecker.format_is_deepep_ll(dispatch_output):
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
assert isinstance(dispatch_output, DeepEPLLDispatchOutput)
|
|
||||||
(
|
|
||||||
hidden_states,
|
|
||||||
hidden_states_scale,
|
|
||||||
topk_ids,
|
|
||||||
topk_weights,
|
|
||||||
group_list,
|
|
||||||
_,
|
|
||||||
) = dispatch_output
|
|
||||||
|
|
||||||
group_list = group_list.to(torch.int64)
|
|
||||||
|
|
||||||
if self.w13_weight.dtype == torch.bfloat16:
|
|
||||||
hidden_states = npu_fused_moe_without_routing_weights_bf16(
|
|
||||||
self, hidden_states, group_list_type, group_list, output_dtype
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
hidden_states = self.quant_method.apply_without_routing_weights(
|
|
||||||
self,
|
|
||||||
hidden_states,
|
|
||||||
hidden_states_scale,
|
|
||||||
group_list_type,
|
|
||||||
group_list,
|
|
||||||
output_dtype,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
raise ValueError(f"Not Supported DeepEP format {dispatch_output.format}")
|
|
||||||
|
|
||||||
return hidden_states
|
|
||||||
|
|
||||||
|
|
||||||
class NpuFuseEPMoE(DeepEPMoE):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
num_experts: int,
|
|
||||||
top_k: int,
|
|
||||||
hidden_size: int,
|
|
||||||
intermediate_size: int,
|
|
||||||
layer_id: int,
|
|
||||||
num_fused_shared_experts: int = 0,
|
|
||||||
params_dtype: Optional[torch.dtype] = None,
|
|
||||||
quant_config: Optional[QuantizationConfig] = None,
|
|
||||||
prefix: str = "",
|
|
||||||
activation: str = "silu",
|
|
||||||
routed_scaling_factor: Optional[float] = None,
|
|
||||||
**kwargs,
|
|
||||||
):
|
|
||||||
super().__init__(
|
|
||||||
num_experts=num_experts,
|
|
||||||
top_k=top_k,
|
|
||||||
hidden_size=hidden_size,
|
|
||||||
intermediate_size=intermediate_size,
|
|
||||||
layer_id=layer_id,
|
|
||||||
num_fused_shared_experts=num_fused_shared_experts,
|
|
||||||
params_dtype=params_dtype,
|
|
||||||
quant_config=quant_config,
|
|
||||||
prefix=prefix,
|
|
||||||
activation=activation,
|
|
||||||
routed_scaling_factor=routed_scaling_factor,
|
|
||||||
**kwargs,
|
|
||||||
)
|
|
||||||
|
|
||||||
self.quant_method.process_weights_after_loading = (
|
|
||||||
self._process_weights_after_loading
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward(
|
|
||||||
self,
|
|
||||||
hidden_states: torch.Tensor,
|
|
||||||
topk_output: TopKOutput,
|
|
||||||
forward_shared_experts=None,
|
|
||||||
alt_stream=None,
|
|
||||||
disable_sbo=False,
|
|
||||||
):
|
|
||||||
return self.dispatcher.dispatch(
|
|
||||||
hidden_states=hidden_states,
|
|
||||||
topk_output=topk_output,
|
|
||||||
gmm1_permuted_weight=self.w13_weight,
|
|
||||||
gmm1_permuted_weight_scale=self.w13_weight_scale,
|
|
||||||
gmm2_weight=self.w2_weight,
|
|
||||||
gmm2_weight_scale=self.w2_weight_scale,
|
|
||||||
).hidden_state
|
|
||||||
|
|
||||||
def permute_w13_weight_scale(self, w: torch.Tensor, tile_n: int):
|
|
||||||
if tile_n % 2 != 0:
|
|
||||||
raise ValueError(f"tile_n must be even, got {tile_n}")
|
|
||||||
|
|
||||||
*dims, n = w.shape
|
|
||||||
if n % tile_n != 0:
|
|
||||||
raise ValueError(f"Last dimension {n} must be divisible by tile_n {tile_n}")
|
|
||||||
|
|
||||||
w_reshaped = w.reshape(*dims, 2, n // tile_n, tile_n // 2)
|
|
||||||
|
|
||||||
# Permute the last two dimensions.
|
|
||||||
perm_order = list(range(len(dims))) + [-2, -3, -1]
|
|
||||||
w_permuted = w_reshaped.permute(perm_order)
|
|
||||||
|
|
||||||
return w_permuted.reshape(*dims, n)
|
|
||||||
|
|
||||||
def reshape_w13_weight(self, weight: torch.Tensor, dim: int, chunk_size: int = 64):
|
|
||||||
# Achieving greater computing power through reshape on Ascend.
|
|
||||||
original_shape = weight.shape
|
|
||||||
if dim < 0:
|
|
||||||
dim += len(original_shape)
|
|
||||||
|
|
||||||
if original_shape[dim] % (2 * chunk_size) != 0:
|
|
||||||
raise ValueError(
|
|
||||||
f"Dimension {dim} size {original_shape[dim]} must be divisible by {2 * chunk_size}"
|
|
||||||
)
|
|
||||||
|
|
||||||
new_shape = (
|
|
||||||
*original_shape[:dim],
|
|
||||||
2,
|
|
||||||
original_shape[dim] // (2 * chunk_size),
|
|
||||||
chunk_size,
|
|
||||||
*original_shape[dim + 1 :],
|
|
||||||
)
|
|
||||||
|
|
||||||
weight = weight.view(new_shape)
|
|
||||||
weight = weight.transpose(dim, dim + 1).contiguous()
|
|
||||||
|
|
||||||
return weight.view(*original_shape[:dim], -1, *original_shape[dim + 1 :])
|
|
||||||
|
|
||||||
def release_weight_cache(self, weight: torch.Tensor):
|
|
||||||
# .contiguous() introduces additional memory overhead and needs to be released using resize_(0)
|
|
||||||
origin_weight = weight.data.transpose(1, 2)
|
|
||||||
new_weight = origin_weight.contiguous()
|
|
||||||
origin_weight.untyped_storage().resize_(0)
|
|
||||||
return new_weight
|
|
||||||
|
|
||||||
def scale_from_float_to_int64(self, scale):
|
|
||||||
import numpy as np
|
|
||||||
|
|
||||||
scale = torch.from_numpy(
|
|
||||||
np.frombuffer(
|
|
||||||
scale.cpu().to(torch.float32).numpy().tobytes(), dtype=np.int32
|
|
||||||
).astype(np.int64)
|
|
||||||
).to(scale.device)
|
|
||||||
return torch.nn.Parameter(scale, requires_grad=False)
|
|
||||||
|
|
||||||
def _process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
|
||||||
if (
|
|
||||||
envs.SGLANG_NPU_FUSED_MOE_MODE.get()
|
|
||||||
== FusedMoEMode.DISPATCH_FFN_COMBINE.value
|
|
||||||
):
|
|
||||||
w13_weight = self.release_weight_cache(layer.w13_weight)
|
|
||||||
layer.w13_weight.data = npu_format_cast(w13_weight)
|
|
||||||
w2_weight = self.release_weight_cache(layer.w2_weight)
|
|
||||||
layer.w2_weight.data = npu_format_cast(w2_weight)
|
|
||||||
|
|
||||||
layer.w13_weight_scale.data = layer.w13_weight_scale.data.view(
|
|
||||||
layer.w13_weight_scale.data.shape[0], -1
|
|
||||||
)
|
|
||||||
w2_scale = layer.w2_weight_scale.data.squeeze(-1).contiguous()
|
|
||||||
layer.w2_weight_scale = torch.nn.Parameter(
|
|
||||||
w2_scale.to(torch.float32), requires_grad=False
|
|
||||||
)
|
|
||||||
|
|
||||||
layer.w13_weight_scale = self.scale_from_float_to_int64(
|
|
||||||
layer.w13_weight_scale.data
|
|
||||||
)
|
|
||||||
layer.w2_weight_scale = self.scale_from_float_to_int64(
|
|
||||||
layer.w2_weight_scale.data
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
cpu_w13 = layer.w13_weight.data.transpose(1, 2).cpu()
|
|
||||||
layer.w13_weight.data = self.reshape_w13_weight(cpu_w13, -1).npu()
|
|
||||||
w13_scale = layer.w13_weight_scale.data.squeeze(-1).contiguous()
|
|
||||||
w13_scale = self.permute_w13_weight_scale(w13_scale, 128)
|
|
||||||
layer.w13_weight_scale = torch.nn.Parameter(
|
|
||||||
w13_scale.to(torch.float32), requires_grad=False
|
|
||||||
)
|
|
||||||
layer.w13_weight.data = npu_format_cast(layer.w13_weight.data)
|
|
||||||
layer.w2_weight.data = npu_format_cast(layer.w2_weight.data)
|
|
||||||
|
|
||||||
w2_scale = layer.w2_weight_scale.data.squeeze(-1).contiguous()
|
|
||||||
layer.w2_weight_scale = torch.nn.Parameter(
|
|
||||||
w2_scale.to(torch.float32), requires_grad=False
|
|
||||||
)
|
|
||||||
|
|
||||||
if hasattr(layer, "w13_weight_offset"):
|
|
||||||
layer.w13_weight_offset = torch.nn.Parameter(
|
|
||||||
layer.w13_weight_offset.data.squeeze(-1).contiguous(),
|
|
||||||
requires_grad=False,
|
|
||||||
)
|
|
||||||
if hasattr(layer, "w2_weight_offset"):
|
|
||||||
layer.w2_weight_offset = torch.nn.Parameter(
|
|
||||||
layer.w2_weight_offset.data.squeeze(-1).contiguous(),
|
|
||||||
requires_grad=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def get_moe_impl_class(quant_config: Optional[QuantizationConfig]):
|
def get_moe_impl_class(quant_config: Optional[QuantizationConfig]):
|
||||||
# [TODO] kk, temporary solution
|
# [TODO] kk, temporary solution
|
||||||
@@ -515,6 +276,8 @@ def get_moe_impl_class(quant_config: Optional[QuantizationConfig]):
|
|||||||
):
|
):
|
||||||
return DeepEPMoE
|
return DeepEPMoE
|
||||||
if get_moe_a2a_backend().is_ascend_fuseep():
|
if get_moe_a2a_backend().is_ascend_fuseep():
|
||||||
return NpuFuseEPMoE
|
# ascend_fuseep bypasses dispatch/combine inside FusedMoE.forward
|
||||||
|
# (see forward_fuseep in hardware_backend/npu/moe/fuseep.py).
|
||||||
|
return FusedMoE
|
||||||
|
|
||||||
return FusedMoE
|
return FusedMoE
|
||||||
|
|||||||
@@ -83,7 +83,14 @@ _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
|||||||
|
|
||||||
def create_moe_dispatcher(moe_runner_config: MoeRunnerConfig) -> BaseDispatcher:
|
def create_moe_dispatcher(moe_runner_config: MoeRunnerConfig) -> BaseDispatcher:
|
||||||
a2a_backend = get_moe_a2a_backend()
|
a2a_backend = get_moe_a2a_backend()
|
||||||
if a2a_backend.is_none() or a2a_backend.is_megamoe():
|
if (
|
||||||
|
a2a_backend.is_none()
|
||||||
|
or a2a_backend.is_megamoe()
|
||||||
|
or a2a_backend.is_ascend_fuseep()
|
||||||
|
):
|
||||||
|
# ascend_fuseep bypasses the dispatcher abstraction (see
|
||||||
|
# forward_fuseep in hardware_backend/npu/moe/fuseep.py); a
|
||||||
|
# StandardDispatcher is created but never invoked.
|
||||||
return StandardDispatcher(moe_runner_config)
|
return StandardDispatcher(moe_runner_config)
|
||||||
elif (
|
elif (
|
||||||
a2a_backend.is_deepep()
|
a2a_backend.is_deepep()
|
||||||
@@ -107,19 +114,6 @@ def create_moe_dispatcher(moe_runner_config: MoeRunnerConfig) -> BaseDispatcher:
|
|||||||
async_finish=True,
|
async_finish=True,
|
||||||
return_recv_hook=True,
|
return_recv_hook=True,
|
||||||
)
|
)
|
||||||
elif a2a_backend.is_ascend_fuseep():
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher import NpuFuseEPDispatcher
|
|
||||||
|
|
||||||
return NpuFuseEPDispatcher(
|
|
||||||
group=get_tp_group().device_group,
|
|
||||||
router_topk=moe_runner_config.top_k,
|
|
||||||
permute_fusion=True,
|
|
||||||
num_experts=moe_runner_config.num_experts,
|
|
||||||
num_local_experts=moe_runner_config.num_local_experts,
|
|
||||||
hidden_size=moe_runner_config.hidden_size,
|
|
||||||
params_dtype=moe_runner_config.params_dtype,
|
|
||||||
)
|
|
||||||
|
|
||||||
elif a2a_backend.is_flashinfer():
|
elif a2a_backend.is_flashinfer():
|
||||||
return FlashinferDispatcher(
|
return FlashinferDispatcher(
|
||||||
group=get_tp_group().device_group,
|
group=get_tp_group().device_group,
|
||||||
@@ -308,6 +302,7 @@ class FusedMoE(torch.nn.Module):
|
|||||||
|
|
||||||
self.quant_method.create_moe_runner(self, self.moe_runner_config)
|
self.quant_method.create_moe_runner(self, self.moe_runner_config)
|
||||||
self.dispatcher = create_moe_dispatcher(self.moe_runner_config)
|
self.dispatcher = create_moe_dispatcher(self.moe_runner_config)
|
||||||
|
self._use_ascend_fuseep = get_moe_a2a_backend().is_ascend_fuseep()
|
||||||
|
|
||||||
if (
|
if (
|
||||||
get_moe_runner_backend().is_flashinfer_trtllm_routed()
|
get_moe_runner_backend().is_flashinfer_trtllm_routed()
|
||||||
@@ -1058,6 +1053,10 @@ class FusedMoE(torch.nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, hidden_states: torch.Tensor, topk_output: TopKOutput):
|
def forward(self, hidden_states: torch.Tensor, topk_output: TopKOutput):
|
||||||
|
if self._use_ascend_fuseep:
|
||||||
|
from sglang.srt.hardware_backend.npu.moe.fuseep import forward_fuseep
|
||||||
|
|
||||||
|
return forward_fuseep(self, hidden_states, topk_output)
|
||||||
if is_in_piecewise_cuda_graph():
|
if is_in_piecewise_cuda_graph():
|
||||||
if TopKOutputChecker.format_is_standard(topk_output):
|
if TopKOutputChecker.format_is_standard(topk_output):
|
||||||
return moe_forward_piecewise_cuda_graph_impl(
|
return moe_forward_piecewise_cuda_graph_impl(
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ from sglang.srt.layers.moe.token_dispatcher.flashinfer import (
|
|||||||
FlashinferDispatcher,
|
FlashinferDispatcher,
|
||||||
FlashinferDispatchOutput,
|
FlashinferDispatchOutput,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.token_dispatcher.fuseep import NpuFuseEPDispatcher
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher.mooncake import (
|
from sglang.srt.layers.moe.token_dispatcher.mooncake import (
|
||||||
MooncakeCombineInput,
|
MooncakeCombineInput,
|
||||||
MooncakeDispatchOutput,
|
MooncakeDispatchOutput,
|
||||||
@@ -75,5 +74,4 @@ __all__ = [
|
|||||||
"DeepEPLLDispatchOutput",
|
"DeepEPLLDispatchOutput",
|
||||||
"DeepEPLLCombineInput",
|
"DeepEPLLCombineInput",
|
||||||
"DeepEPNormalCombineInput",
|
"DeepEPNormalCombineInput",
|
||||||
"NpuFuseEPDispatcher",
|
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -1,98 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import logging
|
|
||||||
from typing import NamedTuple
|
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher.base import (
|
|
||||||
BaseDispatcher,
|
|
||||||
CombineInput,
|
|
||||||
CombineInputFormat,
|
|
||||||
DispatchOutput,
|
|
||||||
DispatchOutputFormat,
|
|
||||||
)
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer
|
|
||||||
from sglang.srt.layers.moe.topk import TopKOutput
|
|
||||||
from sglang.srt.layers.moe.utils import DeepEPMode
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class FuseEPDispatchOutput(NamedTuple):
|
|
||||||
"""DeepEP low latency dispatch output."""
|
|
||||||
|
|
||||||
hidden_state: torch.Tensor
|
|
||||||
|
|
||||||
@property
|
|
||||||
def format(self) -> DispatchOutputFormat:
|
|
||||||
return DispatchOutputFormat.DEEPEP_LL
|
|
||||||
|
|
||||||
|
|
||||||
class FuseEPCombineInput(NamedTuple):
|
|
||||||
"""DeepEP low latency combine input."""
|
|
||||||
|
|
||||||
hidden_state: torch.Tensor
|
|
||||||
|
|
||||||
@property
|
|
||||||
def format(self) -> CombineInputFormat:
|
|
||||||
return CombineInputFormat.DEEPEP_LL
|
|
||||||
|
|
||||||
|
|
||||||
class NpuFuseEPDispatcher(BaseDispatcher):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
group: torch.distributed.ProcessGroup,
|
|
||||||
router_topk: int,
|
|
||||||
permute_fusion: bool = False,
|
|
||||||
num_experts: int = None,
|
|
||||||
num_local_experts: int = None,
|
|
||||||
hidden_size: int = None,
|
|
||||||
params_dtype: torch.dtype = None,
|
|
||||||
deepep_mode: DeepEPMode = DeepEPMode.LOW_LATENCY,
|
|
||||||
):
|
|
||||||
self.group = group
|
|
||||||
self.router_topk = router_topk
|
|
||||||
self.permute_fusion = permute_fusion
|
|
||||||
self.num_experts = num_experts
|
|
||||||
self.num_local_experts = num_local_experts
|
|
||||||
self.hidden_size = hidden_size
|
|
||||||
self.params_dtype = params_dtype
|
|
||||||
self.deepep_mode = deepep_mode
|
|
||||||
|
|
||||||
self.params_bytes = 2
|
|
||||||
self.num_max_dispatch_tokens_per_rank = (
|
|
||||||
envs.SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get()
|
|
||||||
)
|
|
||||||
|
|
||||||
def dispatch(
|
|
||||||
self, hidden_states: torch.Tensor, topk_output: TopKOutput, **kwargs
|
|
||||||
) -> DispatchOutput:
|
|
||||||
hidden_states, _ = self._get_buffer().fused_deep_moe(
|
|
||||||
hidden_states,
|
|
||||||
topk_idx=topk_output.topk_ids,
|
|
||||||
topk_weights=topk_output.topk_weights,
|
|
||||||
gmm1_permuted_weight=kwargs["gmm1_permuted_weight"],
|
|
||||||
gmm1_permuted_weight_scale=kwargs["gmm1_permuted_weight_scale"],
|
|
||||||
gmm2_weight=kwargs["gmm2_weight"],
|
|
||||||
gmm2_weight_scale=kwargs["gmm2_weight_scale"],
|
|
||||||
num_max_dispatch_tokens_per_rank=self.num_max_dispatch_tokens_per_rank,
|
|
||||||
num_experts=self.num_experts,
|
|
||||||
fuse_mode=envs.SGLANG_NPU_FUSED_MOE_MODE.get(),
|
|
||||||
)
|
|
||||||
return FuseEPDispatchOutput(hidden_states)
|
|
||||||
|
|
||||||
def combine(self, combine_input: CombineInput, **kwargs) -> torch.Tensor:
|
|
||||||
pass
|
|
||||||
|
|
||||||
def _get_buffer(self):
|
|
||||||
DeepEPBuffer.set_dispatch_mode_as_low_latency()
|
|
||||||
return DeepEPBuffer.get_deepep_buffer(
|
|
||||||
self.group,
|
|
||||||
self.hidden_size,
|
|
||||||
self.params_bytes,
|
|
||||||
self.deepep_mode,
|
|
||||||
self.num_max_dispatch_tokens_per_rank,
|
|
||||||
self.num_experts,
|
|
||||||
)
|
|
||||||
@@ -1023,7 +1023,6 @@ class CompressedTensorsFusedMoEMethod(FusedMoEMethodBase):
|
|||||||
layer input. See LinearMethodBase for param details
|
layer input. See LinearMethodBase for param details
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
scheme = layer.scheme
|
scheme = layer.scheme
|
||||||
if scheme is None:
|
if scheme is None:
|
||||||
raise ValueError("A scheme must be defined for each layer")
|
raise ValueError("A scheme must be defined for each layer")
|
||||||
|
|||||||
@@ -44,6 +44,7 @@ from sglang.srt.utils import (
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.moe.token_dispatcher import (
|
from sglang.srt.layers.moe.token_dispatcher import (
|
||||||
CombineInput,
|
CombineInput,
|
||||||
|
DispatchOutput,
|
||||||
StandardDispatchOutput,
|
StandardDispatchOutput,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -637,10 +638,14 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
|
|||||||
def forward_npu(
|
def forward_npu(
|
||||||
self,
|
self,
|
||||||
layer: torch.nn.Module,
|
layer: torch.nn.Module,
|
||||||
dispatch_output: StandardDispatchOutput,
|
dispatch_output: "DispatchOutput",
|
||||||
) -> CombineInput:
|
) -> CombineInput:
|
||||||
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.base import DispatchOutputChecker
|
||||||
|
|
||||||
|
if DispatchOutputChecker.format_is_deepep(dispatch_output):
|
||||||
|
return self._forward_npu_deepep(layer, dispatch_output)
|
||||||
|
|
||||||
# x.shape = [B*S, H]
|
# x.shape = [B*S, H]
|
||||||
x = dispatch_output.hidden_states
|
x = dispatch_output.hidden_states
|
||||||
@@ -719,6 +724,46 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
|
|||||||
|
|
||||||
return StandardCombineInput(hidden_states=final_hidden_states)
|
return StandardCombineInput(hidden_states=final_hidden_states)
|
||||||
|
|
||||||
|
def _forward_npu_deepep(
|
||||||
|
self,
|
||||||
|
layer: torch.nn.Module,
|
||||||
|
dispatch_output: "DispatchOutput",
|
||||||
|
) -> CombineInput:
|
||||||
|
from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import (
|
||||||
|
npu_fused_moe_without_routing_weights_bf16,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher import (
|
||||||
|
DeepEPLLCombineInput,
|
||||||
|
DeepEPNormalCombineInput,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.base import DispatchOutputChecker
|
||||||
|
|
||||||
|
# NOTE: Ascend's Dispatch & Combine does not support FP16
|
||||||
|
output_dtype = torch.bfloat16
|
||||||
|
group_list_type = 1
|
||||||
|
|
||||||
|
if DispatchOutputChecker.format_is_deepep_normal(dispatch_output):
|
||||||
|
hidden_states, _, _, _, num_recv_tokens_per_expert = dispatch_output
|
||||||
|
group_list = torch.tensor(
|
||||||
|
num_recv_tokens_per_expert,
|
||||||
|
dtype=torch.int64,
|
||||||
|
device=hidden_states.device,
|
||||||
|
)
|
||||||
|
combine_cls = DeepEPNormalCombineInput
|
||||||
|
else:
|
||||||
|
hidden_states, _, _, _, group_list, _ = dispatch_output
|
||||||
|
group_list = group_list.to(torch.int64)
|
||||||
|
combine_cls = DeepEPLLCombineInput
|
||||||
|
|
||||||
|
hidden_states = npu_fused_moe_without_routing_weights_bf16(
|
||||||
|
layer, hidden_states, group_list_type, group_list, output_dtype
|
||||||
|
)
|
||||||
|
return combine_cls(
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
topk_ids=dispatch_output.topk_ids,
|
||||||
|
topk_weights=dispatch_output.topk_weights,
|
||||||
|
)
|
||||||
|
|
||||||
def forward_tpu(self, *args, **kwargs) -> CombineInput:
|
def forward_tpu(self, *args, **kwargs) -> CombineInput:
|
||||||
raise NotImplementedError("The TPU backend currently does not support MoE.")
|
raise NotImplementedError("The TPU backend currently does not support MoE.")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user