[Ascend] qwen optimization (#12078)
Co-authored-by: Even Zhou <even.y.zhou@outlook.com>
This commit is contained in:
+9
-10
@@ -10,11 +10,10 @@ ARG PIP_INDEX_URL="https://pypi.org/simple/"
|
|||||||
ARG APTMIRROR=""
|
ARG APTMIRROR=""
|
||||||
ARG PYTORCH_VERSION="2.8.0"
|
ARG PYTORCH_VERSION="2.8.0"
|
||||||
ARG TORCHVISION_VERSION="0.23.0"
|
ARG TORCHVISION_VERSION="0.23.0"
|
||||||
ARG PTA_VERSION="v7.2.0-pytorch${PYTORCH_VERSION}"
|
ARG PTA_URL="https://sglang-ascend.obs.cn-east-3.myhuaweicloud.com/sglang/torch_npu/torch_npu-2.8.0.post2.dev20251113-cp311-cp311-manylinux_2_28_aarch64.whl"
|
||||||
ARG PTA_NAME="torch_npu-${PYTORCH_VERSION}-cp311-cp311-manylinux_2_28_aarch64.whl"
|
ARG TRITON_ASCEND_URL="https://sglang-ascend.obs.cn-east-3.myhuaweicloud.com/sglang/triton_ascend/triton_ascend-3.2.0.dev2025112116-cp311-cp311-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl"
|
||||||
ARG PTA_URL="https://gitcode.com/Ascend/pytorch/releases/download/${PTA_VERSION}/${PTA_NAME}"
|
ARG BISHENG_NAME="Ascend-BiSheng-toolkit_aarch64_20251121.run"
|
||||||
ARG TRITON_ASCEND_URL="https://sglang-ascend.obs.cn-east-3.myhuaweicloud.com/sglang/triton_ascend-3.2.0%2Bgitb0ea0850-cp311-cp311-linux_aarch64.whl"
|
ARG BISHENG_URL="https://sglang-ascend.obs.cn-east-3.myhuaweicloud.com/sglang/triton_ascend/${BISHENG_NAME}"
|
||||||
ARG BISHENG_URL="https://sglang-ascend.obs.cn-east-3.myhuaweicloud.com/sglang/Ascend-BiSheng-toolkit_aarch64.run"
|
|
||||||
ARG SGLANG_TAG=main
|
ARG SGLANG_TAG=main
|
||||||
ARG ASCEND_CANN_PATH=/usr/local/Ascend/ascend-toolkit
|
ARG ASCEND_CANN_PATH=/usr/local/Ascend/ascend-toolkit
|
||||||
ARG SGLANG_KERNEL_NPU_TAG=main
|
ARG SGLANG_KERNEL_NPU_TAG=main
|
||||||
@@ -64,13 +63,13 @@ RUN ${PIP_INSTALL} sglang-router
|
|||||||
|
|
||||||
|
|
||||||
### Install PyTorch and PTA
|
### Install PyTorch and PTA
|
||||||
RUN (${PIP_INSTALL} torch==${PYTORCH_VERSION} torchvision==${TORCHVISION_VERSION} --index-url https://download.pytorch.org/whl/cpu) && \
|
RUN (${PIP_INSTALL} torch==${PYTORCH_VERSION} torchvision==${TORCHVISION_VERSION} --index-url https://download.pytorch.org/whl/cpu) \
|
||||||
(wget -O "${PTA_NAME}" "${PTA_URL}" && ${PIP_INSTALL} "./${PTA_NAME}" && rm "./${PTA_NAME}")
|
&& (${PIP_INSTALL} ${PTA_URL})
|
||||||
|
|
||||||
|
|
||||||
# TODO: install from pypi released triton-ascend
|
# TODO: install from pypi released triton-ascend
|
||||||
RUN ${PIP_INSTALL} attrs==24.2.0 numpy==1.26.4 scipy==1.13.1 decorator==5.1.1 psutil==6.0.0 pytest==8.3.2 pytest-xdist==3.6.1 pyyaml pybind11 && \
|
RUN (${PIP_INSTALL} pybind11) \
|
||||||
${PIP_INSTALL} ${TRITON_ASCEND_URL}
|
&& (${PIP_INSTALL} ${TRITON_ASCEND_URL})
|
||||||
|
|
||||||
# Install SGLang
|
# Install SGLang
|
||||||
RUN git clone https://github.com/sgl-project/sglang --branch $SGLANG_TAG && \
|
RUN git clone https://github.com/sgl-project/sglang --branch $SGLANG_TAG && \
|
||||||
@@ -96,6 +95,6 @@ RUN wget https://sglang-ascend.obs.cn-east-3.myhuaweicloud.com/ops/CANN-custom_o
|
|||||||
${PIP_INSTALL} ./custom_ops-1.0.$DEVICE_TYPE-cp311-cp311-linux_aarch64.whl
|
${PIP_INSTALL} ./custom_ops-1.0.$DEVICE_TYPE-cp311-cp311-linux_aarch64.whl
|
||||||
|
|
||||||
# Install Bisheng
|
# Install Bisheng
|
||||||
RUN wget ${BISHENG_URL} && chmod a+x Ascend-BiSheng-toolkit_aarch64.run && ./Ascend-BiSheng-toolkit_aarch64.run --install && rm Ascend-BiSheng-toolkit_aarch64.run
|
RUN wget -O "${BISHENG_NAME}" "${BISHENG_URL}" && chmod a+x "${BISHENG_NAME}" && "./${BISHENG_NAME}" --install && rm "${BISHENG_NAME}"
|
||||||
|
|
||||||
CMD ["/bin/bash"]
|
CMD ["/bin/bash"]
|
||||||
|
|||||||
@@ -625,53 +625,93 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if not self.use_mla:
|
if not self.use_mla:
|
||||||
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(
|
num_tokens = q.shape[0]
|
||||||
layer.layer_id
|
"""PA will support bs<tp in the later version of CANN"""
|
||||||
).view(-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim)
|
if num_tokens < get_attention_tp_size():
|
||||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(
|
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||||
layer.layer_id
|
layer.layer_id
|
||||||
).view(-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim)
|
).view(-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim)
|
||||||
query = q.reshape(-1, 1, layer.tp_q_head_num * layer.qk_head_dim)
|
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(
|
||||||
if self.forward_metadata.seq_lens_cpu_int is None:
|
layer.layer_id
|
||||||
actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list
|
).view(-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim)
|
||||||
|
query = q.reshape(-1, 1, layer.tp_q_head_num * layer.qk_head_dim)
|
||||||
|
if self.forward_metadata.seq_lens_cpu_int is None:
|
||||||
|
actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list
|
||||||
|
else:
|
||||||
|
actual_seq_len_kv = (
|
||||||
|
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
|
||||||
|
)
|
||||||
|
num_tokens = query.shape[0]
|
||||||
|
workspace = (
|
||||||
|
torch_npu._npu_fused_infer_attention_score_get_max_workspace(
|
||||||
|
query,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
block_table=self.forward_metadata.block_tables,
|
||||||
|
block_size=self.page_size,
|
||||||
|
num_heads=layer.tp_q_head_num,
|
||||||
|
num_key_value_heads=layer.tp_k_head_num,
|
||||||
|
input_layout="BSH",
|
||||||
|
scale=layer.scaling,
|
||||||
|
actual_seq_lengths_kv=actual_seq_len_kv,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
output = torch.empty(
|
||||||
|
(num_tokens, 1, layer.tp_q_head_num * layer.v_head_dim),
|
||||||
|
dtype=q.dtype,
|
||||||
|
device=q.device,
|
||||||
|
)
|
||||||
|
softmax_lse = torch.empty(1, dtype=q.dtype, device=q.device)
|
||||||
|
torch_npu.npu_fused_infer_attention_score.out(
|
||||||
|
query,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
block_table=self.forward_metadata.block_tables,
|
||||||
|
block_size=self.page_size,
|
||||||
|
num_heads=layer.tp_q_head_num,
|
||||||
|
num_key_value_heads=layer.tp_k_head_num,
|
||||||
|
input_layout="BSH",
|
||||||
|
scale=layer.scaling,
|
||||||
|
actual_seq_lengths_kv=actual_seq_len_kv,
|
||||||
|
workspace=workspace,
|
||||||
|
out=[output, softmax_lse],
|
||||||
|
)
|
||||||
|
return output.view(num_tokens, layer.tp_q_head_num * layer.v_head_dim)
|
||||||
else:
|
else:
|
||||||
actual_seq_len_kv = (
|
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||||
self.forward_metadata.seq_lens_cpu_int.cpu().int().tolist()
|
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(
|
||||||
|
layer.layer_id
|
||||||
|
)
|
||||||
|
query = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
||||||
|
num_tokens = query.shape[0]
|
||||||
|
attn_output = torch.empty(
|
||||||
|
(num_tokens, layer.tp_q_head_num, layer.v_head_dim),
|
||||||
|
dtype=query.dtype,
|
||||||
|
device=query.device,
|
||||||
|
)
|
||||||
|
if self.forward_metadata.seq_lens_cpu_int is None:
|
||||||
|
actual_seq_len_kv = torch.from_numpy(
|
||||||
|
np.array(self.forward_metadata.seq_lens_cpu_list).astype(
|
||||||
|
np.int32
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_int
|
||||||
|
|
||||||
|
torch_npu._npu_paged_attention(
|
||||||
|
query=query,
|
||||||
|
key_cache=k_cache,
|
||||||
|
value_cache=v_cache,
|
||||||
|
num_heads=layer.tp_q_head_num,
|
||||||
|
num_kv_heads=layer.tp_k_head_num,
|
||||||
|
scale_value=layer.scaling,
|
||||||
|
block_table=self.forward_metadata.block_tables,
|
||||||
|
context_lens=actual_seq_len_kv,
|
||||||
|
out=attn_output,
|
||||||
|
)
|
||||||
|
return attn_output.view(
|
||||||
|
num_tokens, layer.tp_q_head_num * layer.v_head_dim
|
||||||
)
|
)
|
||||||
num_tokens = query.shape[0]
|
|
||||||
workspace = torch_npu._npu_fused_infer_attention_score_get_max_workspace(
|
|
||||||
query,
|
|
||||||
k_cache,
|
|
||||||
v_cache,
|
|
||||||
block_table=self.forward_metadata.block_tables,
|
|
||||||
block_size=self.page_size,
|
|
||||||
num_heads=layer.tp_q_head_num,
|
|
||||||
num_key_value_heads=layer.tp_k_head_num,
|
|
||||||
input_layout="BSH",
|
|
||||||
scale=layer.scaling,
|
|
||||||
actual_seq_lengths_kv=actual_seq_len_kv,
|
|
||||||
)
|
|
||||||
output = torch.empty(
|
|
||||||
(num_tokens, 1, layer.tp_q_head_num * layer.v_head_dim),
|
|
||||||
dtype=q.dtype,
|
|
||||||
device=q.device,
|
|
||||||
)
|
|
||||||
softmax_lse = torch.empty(1, dtype=q.dtype, device=q.device)
|
|
||||||
torch_npu.npu_fused_infer_attention_score.out(
|
|
||||||
query,
|
|
||||||
k_cache,
|
|
||||||
v_cache,
|
|
||||||
block_table=self.forward_metadata.block_tables,
|
|
||||||
block_size=self.page_size,
|
|
||||||
num_heads=layer.tp_q_head_num,
|
|
||||||
num_key_value_heads=layer.tp_k_head_num,
|
|
||||||
input_layout="BSH",
|
|
||||||
scale=layer.scaling,
|
|
||||||
actual_seq_lengths_kv=actual_seq_len_kv,
|
|
||||||
workspace=workspace,
|
|
||||||
out=[output, softmax_lse],
|
|
||||||
)
|
|
||||||
return output.view(num_tokens, layer.tp_q_head_num * layer.v_head_dim)
|
|
||||||
else:
|
else:
|
||||||
c_kv, k_rope = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
c_kv, k_rope = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id)
|
||||||
k_rope_cache = k_rope.view(
|
k_rope_cache = k_rope.view(
|
||||||
|
|||||||
@@ -45,6 +45,10 @@ elif _is_npu:
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
if _is_npu:
|
||||||
|
import torch_npu
|
||||||
|
|
||||||
|
|
||||||
class DeepEPMoE(FusedMoE):
|
class DeepEPMoE(FusedMoE):
|
||||||
"""
|
"""
|
||||||
MoE Expert Parallel Impl based on DeepEP (https://github.com/deepseek-ai/DeepEP/tree/main)
|
MoE Expert Parallel Impl based on DeepEP (https://github.com/deepseek-ai/DeepEP/tree/main)
|
||||||
@@ -411,9 +415,142 @@ def npu_fused_moe_without_routing_weights_bf16(
|
|||||||
return hidden_states
|
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,
|
||||||
|
):
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
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 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 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 _process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||||
|
w13 = self.release_weight_cache(layer.w13_weight)
|
||||||
|
torch_npu.npu_format_cast_(w13, 2)
|
||||||
|
cpu_w13 = w13.cpu()
|
||||||
|
w13 = self.reshape_w13_weight(cpu_w13, -1).npu()
|
||||||
|
torch_npu.npu_format_cast_(w13, 29)
|
||||||
|
layer.w13_weight = torch.nn.Parameter(w13, requires_grad=False)
|
||||||
|
|
||||||
|
w2 = torch_npu.npu_format_cast(layer.w2_weight.data, 29)
|
||||||
|
layer.w2_weight = torch.nn.Parameter(w2, requires_grad=False)
|
||||||
|
|
||||||
|
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
|
||||||
|
)
|
||||||
|
|
||||||
|
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]):
|
||||||
if get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mooncake():
|
if get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_mooncake():
|
||||||
return DeepEPMoE
|
return DeepEPMoE
|
||||||
|
if get_moe_a2a_backend().is_ascend_fuseep():
|
||||||
|
return NpuFuseEPMoE
|
||||||
|
|
||||||
# NEW: Direct FP4 detection (bypasses EP requirements)
|
# NEW: Direct FP4 detection (bypasses EP requirements)
|
||||||
# Check for FP4 quantization with TRTLLM flag, regardless of EP
|
# Check for FP4 quantization with TRTLLM flag, regardless of EP
|
||||||
|
|||||||
@@ -93,6 +93,18 @@ 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,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError(f"Unsupported a2a backend: {a2a_backend}")
|
raise NotImplementedError(f"Unsupported a2a backend: {a2a_backend}")
|
||||||
|
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ from sglang.srt.layers.moe.token_dispatcher.deepep import (
|
|||||||
DeepEPNormalCombineInput,
|
DeepEPNormalCombineInput,
|
||||||
DeepEPNormalDispatchOutput,
|
DeepEPNormalDispatchOutput,
|
||||||
)
|
)
|
||||||
|
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,
|
||||||
@@ -48,4 +49,5 @@ __all__ = [
|
|||||||
"DeepEPLLDispatchOutput",
|
"DeepEPLLDispatchOutput",
|
||||||
"DeepEPLLCombineInput",
|
"DeepEPLLCombineInput",
|
||||||
"DeepEPNormalCombineInput",
|
"DeepEPNormalCombineInput",
|
||||||
|
"NpuFuseEPDispatcher",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -0,0 +1,97 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import NamedTuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
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
|
||||||
|
from sglang.srt.utils import get_int_env_var
|
||||||
|
|
||||||
|
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 = get_int_env_var(
|
||||||
|
"SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK", 128
|
||||||
|
)
|
||||||
|
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
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,
|
||||||
|
)
|
||||||
@@ -106,6 +106,7 @@ if _use_aiter:
|
|||||||
raise ImportError("aiter is required when SGLANG_USE_AITER is set to True")
|
raise ImportError("aiter is required when SGLANG_USE_AITER is set to True")
|
||||||
if _is_npu:
|
if _is_npu:
|
||||||
import torch_npu
|
import torch_npu
|
||||||
|
from sgl_kernel_npu.norm.l1_norm import l1_norm
|
||||||
|
|
||||||
# -------------------------------- TopKConfig ---------------------------------------
|
# -------------------------------- TopKConfig ---------------------------------------
|
||||||
|
|
||||||
@@ -363,15 +364,14 @@ class TopK(CustomOp):
|
|||||||
router_logits,
|
router_logits,
|
||||||
k=self.topk_config.top_k,
|
k=self.topk_config.top_k,
|
||||||
)
|
)
|
||||||
topk_weights = topk_weights.to(torch.float32)
|
|
||||||
|
|
||||||
if renormalize:
|
if renormalize:
|
||||||
topk_weights_sum = (
|
topk_weights = l1_norm(
|
||||||
topk_weights.sum(dim=-1, keepdim=True)
|
topk_weights
|
||||||
if self.topk_config.num_fused_shared_experts == 0
|
if self.topk_config.num_fused_shared_experts == 0
|
||||||
else topk_weights[:, :-1].sum(dim=-1, keepdim=True)
|
else topk_weights[:, :-1]
|
||||||
)
|
)
|
||||||
topk_weights = topk_weights / topk_weights_sum
|
topk_weights = topk_weights.to(torch.float32)
|
||||||
|
|
||||||
if expert_location_dispatch_info is not None:
|
if expert_location_dispatch_info is not None:
|
||||||
topk_ids = topk_ids_logical_to_physical(
|
topk_ids = topk_ids_logical_to_physical(
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ class MoeA2ABackend(Enum):
|
|||||||
NONE = "none"
|
NONE = "none"
|
||||||
DEEPEP = "deepep"
|
DEEPEP = "deepep"
|
||||||
MOONCAKE = "mooncake"
|
MOONCAKE = "mooncake"
|
||||||
|
ASCEND_FUSEEP = "ascend_fuseep"
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _missing_(cls, value):
|
def _missing_(cls, value):
|
||||||
@@ -43,6 +44,9 @@ class MoeA2ABackend(Enum):
|
|||||||
def is_mooncake(self):
|
def is_mooncake(self):
|
||||||
return self == MoeA2ABackend.MOONCAKE
|
return self == MoeA2ABackend.MOONCAKE
|
||||||
|
|
||||||
|
def is_ascend_fuseep(self):
|
||||||
|
return self == MoeA2ABackend.ASCEND_FUSEEP
|
||||||
|
|
||||||
|
|
||||||
class MoeRunnerBackend(Enum):
|
class MoeRunnerBackend(Enum):
|
||||||
|
|
||||||
|
|||||||
@@ -32,6 +32,7 @@ from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
|
|||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
apply_module_patch,
|
apply_module_patch,
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
|
get_bool_env_var,
|
||||||
is_cpu,
|
is_cpu,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
is_npu,
|
is_npu,
|
||||||
@@ -63,6 +64,8 @@ elif _is_npu:
|
|||||||
else:
|
else:
|
||||||
useMindIETurbo = True
|
useMindIETurbo = True
|
||||||
|
|
||||||
|
from sgl_kernel_npu.norm.add_rmsnorm_bias import add_rmsnorm_bias
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@@ -87,10 +90,13 @@ def npu_wrapper_rmsnorm_forward(func):
|
|||||||
if not x.is_contiguous():
|
if not x.is_contiguous():
|
||||||
x = x.contiguous()
|
x = x.contiguous()
|
||||||
if residual is not None:
|
if residual is not None:
|
||||||
out, _, residual_out = torch_npu.npu_add_rms_norm(
|
out, residual_out = add_rmsnorm_bias(
|
||||||
residual, x, self.weight.data, self.variance_epsilon
|
x,
|
||||||
|
residual,
|
||||||
|
self.weight.data,
|
||||||
|
self.bias,
|
||||||
|
self.variance_epsilon,
|
||||||
)
|
)
|
||||||
out = out + self.bias
|
|
||||||
return out.to(x.dtype), residual_out
|
return out.to(x.dtype), residual_out
|
||||||
|
|
||||||
out = torch_npu.npu_rms_norm(x, self.weight.data, self.variance_epsilon)[0]
|
out = torch_npu.npu_rms_norm(x, self.weight.data, self.variance_epsilon)[0]
|
||||||
@@ -1072,15 +1078,23 @@ class NPU_W8A8MoEMethod(FusedMoEMethodBase):
|
|||||||
layer.register_parameter("w2_weight_offset", w2_weight_offset)
|
layer.register_parameter("w2_weight_offset", w2_weight_offset)
|
||||||
set_weight_attrs(w2_weight_offset, extra_weight_attrs)
|
set_weight_attrs(w2_weight_offset, extra_weight_attrs)
|
||||||
|
|
||||||
|
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 process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||||
layer.w13_weight = Parameter(
|
weight_data = self.release_weight_cache(layer.w13_weight.data)
|
||||||
layer.w13_weight.data.transpose(1, 2).contiguous(), requires_grad=False
|
layer.w13_weight = Parameter(weight_data, requires_grad=False)
|
||||||
)
|
|
||||||
layer.w2_weight = Parameter(
|
weight_data = self.release_weight_cache(layer.w2_weight.data)
|
||||||
layer.w2_weight.data.transpose(1, 2).contiguous(), requires_grad=False
|
layer.w2_weight = Parameter(weight_data, requires_grad=False)
|
||||||
)
|
|
||||||
layer.w13_weight_scale = Parameter(
|
layer.w13_weight_scale = Parameter(
|
||||||
layer.w13_weight_scale.data.squeeze(-1).contiguous(), requires_grad=False
|
layer.w13_weight_scale.data.squeeze(-1).contiguous().to(torch.float32),
|
||||||
|
requires_grad=False,
|
||||||
)
|
)
|
||||||
layer.w2_weight_scale = Parameter(
|
layer.w2_weight_scale = Parameter(
|
||||||
layer.w2_weight_scale.data.squeeze(-1).contiguous(), requires_grad=False
|
layer.w2_weight_scale.data.squeeze(-1).contiguous(), requires_grad=False
|
||||||
@@ -1092,6 +1106,10 @@ class NPU_W8A8MoEMethod(FusedMoEMethodBase):
|
|||||||
layer.w2_weight_offset.data.squeeze(-1).contiguous(), requires_grad=False
|
layer.w2_weight_offset.data.squeeze(-1).contiguous(), requires_grad=False
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if get_bool_env_var("ENABLE_ASCEND_MOE_NZ"):
|
||||||
|
layer.w13_weight.data = torch_npu.npu_format_cast(layer.w13_weight.data, 29)
|
||||||
|
layer.w2_weight.data = torch_npu.npu_format_cast(layer.w2_weight.data, 29)
|
||||||
|
|
||||||
def create_moe_runner(
|
def create_moe_runner(
|
||||||
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
||||||
):
|
):
|
||||||
@@ -1145,7 +1163,7 @@ class NPU_W8A8MoEMethod(FusedMoEMethodBase):
|
|||||||
# act_fn: swiglu
|
# act_fn: swiglu
|
||||||
hidden_states, swiglu_out_scale = torch_npu.npu_dequant_swiglu_quant(
|
hidden_states, swiglu_out_scale = torch_npu.npu_dequant_swiglu_quant(
|
||||||
x=hidden_states,
|
x=hidden_states,
|
||||||
weight_scale=layer.w13_weight_scale.to(torch.float32),
|
weight_scale=layer.w13_weight_scale,
|
||||||
activation_scale=hidden_states_scale,
|
activation_scale=hidden_states_scale,
|
||||||
bias=None,
|
bias=None,
|
||||||
quant_scale=None,
|
quant_scale=None,
|
||||||
|
|||||||
@@ -135,6 +135,7 @@ class RotaryEmbedding(CustomOp):
|
|||||||
self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)(
|
self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)(
|
||||||
self._apply_rotary_emb_wrapped
|
self._apply_rotary_emb_wrapped
|
||||||
)
|
)
|
||||||
|
self.position_cos, self.position_sin = None, None
|
||||||
|
|
||||||
def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor:
|
def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor:
|
||||||
"""Compute the inverse frequency."""
|
"""Compute the inverse frequency."""
|
||||||
@@ -202,6 +203,18 @@ class RotaryEmbedding(CustomOp):
|
|||||||
device=device, dtype=dtype
|
device=device, dtype=dtype
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def get_cos_sin_with_position(self, positions):
|
||||||
|
cos_sin = self.cos_sin_cache.index_select(0, positions.flatten())
|
||||||
|
last_dim = cos_sin.size()[-1]
|
||||||
|
cos, sin = (
|
||||||
|
cos_sin.reshape(-1, 2, last_dim // 2).repeat(1, 1, 2).chunk(2, dim=-2)
|
||||||
|
)
|
||||||
|
# BSNH
|
||||||
|
self.position_cos, self.position_sin = (
|
||||||
|
cos.view(-1, 1, 1, last_dim).contiguous(),
|
||||||
|
sin.view(-1, 1, 1, last_dim).contiguous(),
|
||||||
|
)
|
||||||
|
|
||||||
def forward_native(
|
def forward_native(
|
||||||
self,
|
self,
|
||||||
positions: torch.Tensor,
|
positions: torch.Tensor,
|
||||||
|
|||||||
@@ -19,11 +19,13 @@ import logging
|
|||||||
import os
|
import os
|
||||||
import threading
|
import threading
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Optional, Union
|
from typing import TYPE_CHECKING, Dict, Optional, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import is_deepseek_nsa
|
from sglang.srt.configs.model_config import AttentionArch, is_deepseek_nsa
|
||||||
|
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||||
from sglang.srt.model_executor.cuda_graph_runner import CudaGraphRunner
|
from sglang.srt.model_executor.cuda_graph_runner import CudaGraphRunner
|
||||||
from sglang.srt.utils import is_npu
|
from sglang.srt.utils import is_npu
|
||||||
|
|
||||||
@@ -47,6 +49,20 @@ class NPUGraphRunner(CudaGraphRunner):
|
|||||||
|
|
||||||
def __init__(self, model_runner: ModelRunner):
|
def __init__(self, model_runner: ModelRunner):
|
||||||
super().__init__(model_runner)
|
super().__init__(model_runner)
|
||||||
|
self.update_attr_name = None
|
||||||
|
self.update_attr_type = None
|
||||||
|
self.model_runner = model_runner
|
||||||
|
self._init_arch_map()
|
||||||
|
|
||||||
|
def _init_arch_map(self):
|
||||||
|
self.attr_name: Dict[str, str] = {
|
||||||
|
AttentionArch.MLA: "actual_seq_lengths_kv",
|
||||||
|
AttentionArch.MHA: "context_lens",
|
||||||
|
}
|
||||||
|
self.attr_type: Dict[str, Union[list, torch.Tensor]] = {
|
||||||
|
AttentionArch.MLA: [],
|
||||||
|
AttentionArch.MHA: torch.Tensor(),
|
||||||
|
}
|
||||||
|
|
||||||
def _create_device_graph(self):
|
def _create_device_graph(self):
|
||||||
return torch.npu.NPUGraph()
|
return torch.npu.NPUGraph()
|
||||||
@@ -61,9 +77,22 @@ class NPUGraphRunner(CudaGraphRunner):
|
|||||||
out = run_once_fn()
|
out = run_once_fn()
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
def _get_update_attr_name(self, model_runner):
|
||||||
|
if self.bs < get_attention_tp_size():
|
||||||
|
return self.attr_name[AttentionArch.MLA]
|
||||||
|
return self.attr_name[model_runner.model_config.attention_arch]
|
||||||
|
|
||||||
|
def _get_update_attr_type(self, model_runner):
|
||||||
|
if self.bs < get_attention_tp_size():
|
||||||
|
return self.attr_type[AttentionArch.MLA]
|
||||||
|
return self.attr_type[model_runner.model_config.attention_arch]
|
||||||
|
|
||||||
def _update_inputs(self, seq_lens):
|
def _update_inputs(self, seq_lens):
|
||||||
|
if isinstance(self.update_attr_type, torch.Tensor):
|
||||||
|
seq_lens = torch.from_numpy(np.array(seq_lens).astype(np.int32))
|
||||||
|
|
||||||
self.graphs[self.bs].update(
|
self.graphs[self.bs].update(
|
||||||
cpu_update_input=[{"actual_seq_lengths_kv": seq_lens}]
|
cpu_update_input=[{self.update_attr_name: seq_lens}]
|
||||||
)
|
)
|
||||||
|
|
||||||
def _cache_loc_dtype(self):
|
def _cache_loc_dtype(self):
|
||||||
@@ -110,6 +139,8 @@ class NPUGraphRunner(CudaGraphRunner):
|
|||||||
self.buffers.input_ids[: self.raw_num_token].copy_(forward_batch.input_ids)
|
self.buffers.input_ids[: self.raw_num_token].copy_(forward_batch.input_ids)
|
||||||
self.buffers.positions[: self.raw_num_token].copy_(forward_batch.positions)
|
self.buffers.positions[: self.raw_num_token].copy_(forward_batch.positions)
|
||||||
|
|
||||||
|
self.update_attr_name = self._get_update_attr_name(self.model_runner)
|
||||||
|
self.update_attr_type = self._get_update_attr_type(self.model_runner)
|
||||||
# Replay
|
# Replay
|
||||||
if not is_deepseek_nsa(self.model_runner.model_config.hf_config):
|
if not is_deepseek_nsa(self.model_runner.model_config.hf_config):
|
||||||
if forward_batch.forward_mode.is_target_verify():
|
if forward_batch.forward_mode.is_target_verify():
|
||||||
|
|||||||
@@ -44,6 +44,9 @@ logger = logging.getLogger(__name__)
|
|||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
|
|
||||||
|
if _is_npu:
|
||||||
|
from sgl_kernel_npu.norm.split_qkv_rmsnorm_rope import split_qkv_rmsnorm_rope
|
||||||
|
|
||||||
|
|
||||||
class Qwen3Attention(nn.Module):
|
class Qwen3Attention(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -161,6 +164,33 @@ class Qwen3Attention(nn.Module):
|
|||||||
k = k_by_head.view(k.shape)
|
k = k_by_head.view(k.shape)
|
||||||
return q, k
|
return q, k
|
||||||
|
|
||||||
|
def forward_prepare_native(self, positions, hidden_states):
|
||||||
|
qkv, _ = self.qkv_proj(hidden_states)
|
||||||
|
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||||
|
q, k = self._apply_qk_norm(q, k)
|
||||||
|
q, k = self.rotary_emb(positions, q, k)
|
||||||
|
return q, k, v
|
||||||
|
|
||||||
|
def forward_prepare_npu(self, positions, hidden_states):
|
||||||
|
qkv, _ = self.qkv_proj(hidden_states)
|
||||||
|
|
||||||
|
if self.attn.layer_id == 0:
|
||||||
|
self.rotary_emb.get_cos_sin_with_position(positions)
|
||||||
|
q, k, v = split_qkv_rmsnorm_rope(
|
||||||
|
qkv,
|
||||||
|
self.rotary_emb.position_sin,
|
||||||
|
self.rotary_emb.position_cos,
|
||||||
|
self.q_norm.weight,
|
||||||
|
self.k_norm.weight,
|
||||||
|
self.q_size,
|
||||||
|
self.kv_size,
|
||||||
|
self.head_dim,
|
||||||
|
self.q_norm.variance_epsilon,
|
||||||
|
q_bias=getattr(self.q_norm, "bias", None),
|
||||||
|
k_bias=getattr(self.k_norm, "bias", None),
|
||||||
|
)
|
||||||
|
return q, k, v
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
positions: torch.Tensor,
|
positions: torch.Tensor,
|
||||||
@@ -170,10 +200,16 @@ class Qwen3Attention(nn.Module):
|
|||||||
if get_global_server_args().rl_on_policy_target is not None:
|
if get_global_server_args().rl_on_policy_target is not None:
|
||||||
hidden_states = hidden_states.bfloat16()
|
hidden_states = hidden_states.bfloat16()
|
||||||
|
|
||||||
qkv, _ = self.qkv_proj(hidden_states)
|
if not _is_npu:
|
||||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
q, k, v = self.forward_prepare_native(
|
||||||
q, k = self._apply_qk_norm(q, k)
|
positions=positions,
|
||||||
q, k = self.rotary_emb(positions, q, k)
|
hidden_states=hidden_states,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
q, k, v = self.forward_prepare_npu(
|
||||||
|
positions=positions,
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
)
|
||||||
|
|
||||||
if get_global_server_args().rl_on_policy_target is not None:
|
if get_global_server_args().rl_on_policy_target is not None:
|
||||||
q = q.to(torch.bfloat16)
|
q = q.to(torch.bfloat16)
|
||||||
|
|||||||
@@ -70,6 +70,7 @@ from sglang.srt.utils import (
|
|||||||
is_cuda,
|
is_cuda,
|
||||||
is_flashinfer_available,
|
is_flashinfer_available,
|
||||||
is_non_idle_and_non_empty,
|
is_non_idle_and_non_empty,
|
||||||
|
is_npu,
|
||||||
)
|
)
|
||||||
|
|
||||||
Qwen3MoeConfig = None
|
Qwen3MoeConfig = None
|
||||||
@@ -78,6 +79,10 @@ _is_flashinfer_available = is_flashinfer_available()
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
|
_is_npu = is_npu()
|
||||||
|
|
||||||
|
if _is_npu:
|
||||||
|
from sgl_kernel_npu.norm.split_qkv_rmsnorm_rope import split_qkv_rmsnorm_rope
|
||||||
|
|
||||||
|
|
||||||
class Qwen3MoeSparseMoeBlock(nn.Module):
|
class Qwen3MoeSparseMoeBlock(nn.Module):
|
||||||
@@ -139,7 +144,10 @@ class Qwen3MoeSparseMoeBlock(nn.Module):
|
|||||||
use_reduce_scatter: bool = False,
|
use_reduce_scatter: bool = False,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|
||||||
if not get_moe_a2a_backend().is_deepep():
|
if (
|
||||||
|
not get_moe_a2a_backend().is_deepep()
|
||||||
|
and not get_moe_a2a_backend().is_ascend_fuseep()
|
||||||
|
):
|
||||||
return self.forward_normal(
|
return self.forward_normal(
|
||||||
hidden_states, should_allreduce_fusion, use_reduce_scatter
|
hidden_states, should_allreduce_fusion, use_reduce_scatter
|
||||||
)
|
)
|
||||||
@@ -392,14 +400,37 @@ class Qwen3MoeAttention(nn.Module):
|
|||||||
state.pop("attn_intermediate_state")
|
state.pop("attn_intermediate_state")
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward_prepare(
|
def forward_prepare_npu(
|
||||||
|
self,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
):
|
||||||
|
qkv, _ = self.qkv_proj(hidden_states)
|
||||||
|
if self.attn.layer_id == 0:
|
||||||
|
self.rotary_emb.get_cos_sin_with_position(positions)
|
||||||
|
q, k, v = split_qkv_rmsnorm_rope(
|
||||||
|
qkv,
|
||||||
|
self.rotary_emb.position_sin,
|
||||||
|
self.rotary_emb.position_cos,
|
||||||
|
self.q_norm.weight,
|
||||||
|
self.k_norm.weight,
|
||||||
|
self.q_size,
|
||||||
|
self.kv_size,
|
||||||
|
self.head_dim,
|
||||||
|
self.q_norm.variance_epsilon,
|
||||||
|
q_bias=getattr(self.q_norm, "bias", None),
|
||||||
|
k_bias=getattr(self.k_norm, "bias", None),
|
||||||
|
)
|
||||||
|
inner_state = q, k, v, forward_batch
|
||||||
|
return None, forward_batch, inner_state
|
||||||
|
|
||||||
|
def forward_prepare_native(
|
||||||
self,
|
self,
|
||||||
positions: torch.Tensor,
|
positions: torch.Tensor,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
):
|
):
|
||||||
if hidden_states.shape[0] == 0:
|
|
||||||
return hidden_states, forward_batch, None
|
|
||||||
qkv, _ = self.qkv_proj(hidden_states)
|
qkv, _ = self.qkv_proj(hidden_states)
|
||||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||||
q, k = self._apply_qk_norm(q, k)
|
q, k = self._apply_qk_norm(q, k)
|
||||||
@@ -421,6 +452,27 @@ class Qwen3MoeAttention(nn.Module):
|
|||||||
inner_state = q, k, v, forward_batch
|
inner_state = q, k, v, forward_batch
|
||||||
return None, forward_batch, inner_state
|
return None, forward_batch, inner_state
|
||||||
|
|
||||||
|
def forward_prepare(
|
||||||
|
self,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
):
|
||||||
|
if hidden_states.shape[0] == 0:
|
||||||
|
return hidden_states, forward_batch, None
|
||||||
|
if not _is_npu:
|
||||||
|
return self.forward_prepare_native(
|
||||||
|
positions=positions,
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
return self.forward_prepare_npu(
|
||||||
|
positions=positions,
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
)
|
||||||
|
|
||||||
def forward_core(self, intermediate_state):
|
def forward_core(self, intermediate_state):
|
||||||
hidden_states, forward_batch, inner_state = intermediate_state
|
hidden_states, forward_batch, inner_state = intermediate_state
|
||||||
if inner_state is None:
|
if inner_state is None:
|
||||||
|
|||||||
@@ -404,7 +404,7 @@ class ServerArgs:
|
|||||||
|
|
||||||
# Expert parallelism
|
# Expert parallelism
|
||||||
ep_size: int = 1
|
ep_size: int = 1
|
||||||
moe_a2a_backend: Literal["none", "deepep", "mooncake"] = "none"
|
moe_a2a_backend: Literal["none", "deepep", "mooncake", "ascend_fuseep"] = "none"
|
||||||
moe_runner_backend: str = "auto"
|
moe_runner_backend: str = "auto"
|
||||||
flashinfer_mxfp4_moe_precision: Literal["default", "bf16"] = "default"
|
flashinfer_mxfp4_moe_precision: Literal["default", "bf16"] = "default"
|
||||||
enable_flashinfer_allreduce_fusion: bool = False
|
enable_flashinfer_allreduce_fusion: bool = False
|
||||||
@@ -1516,6 +1516,12 @@ class ServerArgs:
|
|||||||
f"Mooncake MoE is enabled. The expert parallel size is adjusted to be the same as the tensor parallel size[{self.tp_size}]."
|
f"Mooncake MoE is enabled. The expert parallel size is adjusted to be the same as the tensor parallel size[{self.tp_size}]."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if self.moe_a2a_backend == "ascend_fuseep":
|
||||||
|
self.ep_size = self.tp_size
|
||||||
|
logger.warning(
|
||||||
|
f"Ascend fused EP MoE is enabled. The expert parallel size is adjusted to be the same as the tensor parallel size[{self.tp_size}]."
|
||||||
|
)
|
||||||
|
|
||||||
def _handle_eplb_and_dispatch(self):
|
def _handle_eplb_and_dispatch(self):
|
||||||
if self.enable_eplb and (self.expert_distribution_recorder_mode is None):
|
if self.enable_eplb and (self.expert_distribution_recorder_mode is None):
|
||||||
self.expert_distribution_recorder_mode = "stat"
|
self.expert_distribution_recorder_mode = "stat"
|
||||||
@@ -1523,7 +1529,7 @@ class ServerArgs:
|
|||||||
"EPLB is enabled. The expert_distribution_recorder_mode is automatically set."
|
"EPLB is enabled. The expert_distribution_recorder_mode is automatically set."
|
||||||
)
|
)
|
||||||
|
|
||||||
if (self.enable_eplb or (self.init_expert_location is not None)) and (
|
if (self.enable_eplb or (self.init_expert_location != "trivial")) and (
|
||||||
self.ep_dispatch_algorithm is None
|
self.ep_dispatch_algorithm is None
|
||||||
):
|
):
|
||||||
self.ep_dispatch_algorithm = "static"
|
self.ep_dispatch_algorithm = "static"
|
||||||
@@ -2969,7 +2975,7 @@ class ServerArgs:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--moe-a2a-backend",
|
"--moe-a2a-backend",
|
||||||
type=str,
|
type=str,
|
||||||
choices=["none", "deepep", "mooncake"],
|
choices=["none", "deepep", "mooncake", "ascend_fuseep"],
|
||||||
default=ServerArgs.moe_a2a_backend,
|
default=ServerArgs.moe_a2a_backend,
|
||||||
help="Choose the backend for MoE A2A.",
|
help="Choose the backend for MoE A2A.",
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -633,38 +633,48 @@ def get_cmo_stream():
|
|||||||
AIV or communication kernels, aiming to overlap the memory access time.
|
AIV or communication kernels, aiming to overlap the memory access time.
|
||||||
"""
|
"""
|
||||||
global cmo_stream
|
global cmo_stream
|
||||||
if cmo_stream is None:
|
|
||||||
cmo_stream = torch.get_device_module().Stream()
|
|
||||||
return cmo_stream
|
return cmo_stream
|
||||||
|
|
||||||
|
|
||||||
def prepare_weight_cache(handle, cache):
|
def set_cmo_stream(stream):
|
||||||
|
global cmo_stream
|
||||||
|
cmo_stream = stream
|
||||||
|
|
||||||
|
|
||||||
|
def prepare_weight_cache(handle, cache, PREFETCH_MAX_SIZE=1000000000):
|
||||||
|
"""
|
||||||
|
PREFETCH_MAX_SIZE: maximum size (bytes) for each prefetch operation.
|
||||||
|
This affects the time spent in prefetch:
|
||||||
|
time ≈ PREFETCH_MAX_SIZE / system_bandwidth
|
||||||
|
"""
|
||||||
import torch_npu
|
import torch_npu
|
||||||
|
|
||||||
NPU_PREFETCH_MAX_SIZE_BYTES = (
|
|
||||||
1000000000 # 1GB, a large value to prefetch entire weight
|
|
||||||
)
|
|
||||||
stream = get_cmo_stream()
|
stream = get_cmo_stream()
|
||||||
stream.wait_stream(torch.npu.current_stream())
|
if stream is None:
|
||||||
with torch.npu.stream(stream):
|
stream = torch.get_device_module().Stream()
|
||||||
|
set_cmo_stream(stream)
|
||||||
|
stream.wait_stream(torch.get_device_module().current_stream())
|
||||||
|
with torch.get_device_module().stream(stream):
|
||||||
if isinstance(cache, list):
|
if isinstance(cache, list):
|
||||||
for weight in cache:
|
for weight in cache:
|
||||||
torch_npu.npu_prefetch(
|
torch_npu.npu_prefetch(
|
||||||
weight,
|
weight,
|
||||||
handle,
|
handle,
|
||||||
NPU_PREFETCH_MAX_SIZE_BYTES,
|
PREFETCH_MAX_SIZE,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
torch_npu.npu_prefetch(
|
torch_npu.npu_prefetch(
|
||||||
cache,
|
cache,
|
||||||
handle,
|
handle,
|
||||||
NPU_PREFETCH_MAX_SIZE_BYTES,
|
PREFETCH_MAX_SIZE,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def wait_cmo_stream():
|
def wait_cmo_stream():
|
||||||
cur_stream = torch.get_device_module().current_stream()
|
stream = get_cmo_stream()
|
||||||
cur_stream.wait_stream(get_cmo_stream())
|
if stream is not None:
|
||||||
|
cur_stream = torch.get_device_module().current_stream()
|
||||||
|
cur_stream.wait_stream(stream)
|
||||||
|
|
||||||
|
|
||||||
@lru_cache(maxsize=1)
|
@lru_cache(maxsize=1)
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ apt update -y && apt install -y \
|
|||||||
update-ca-certificates
|
update-ca-certificates
|
||||||
${PIP_INSTALL} --upgrade pip
|
${PIP_INSTALL} --upgrade pip
|
||||||
# Pin wheel to 0.45.1, REF: https://github.com/pypa/wheel/issues/662
|
# Pin wheel to 0.45.1, REF: https://github.com/pypa/wheel/issues/662
|
||||||
${PIP_INSTALL} wheel==0.45.1
|
${PIP_INSTALL} wheel==0.45.1 pybind11
|
||||||
|
|
||||||
|
|
||||||
### Install MemFabric
|
### Install MemFabric
|
||||||
@@ -33,22 +33,18 @@ PYTORCH_VERSION="2.8.0"
|
|||||||
TORCHVISION_VERSION="0.23.0"
|
TORCHVISION_VERSION="0.23.0"
|
||||||
${PIP_INSTALL} torch==${PYTORCH_VERSION} torchvision==${TORCHVISION_VERSION} --index-url https://download.pytorch.org/whl/cpu
|
${PIP_INSTALL} torch==${PYTORCH_VERSION} torchvision==${TORCHVISION_VERSION} --index-url https://download.pytorch.org/whl/cpu
|
||||||
|
|
||||||
PTA_VERSION="v7.2.0-pytorch${PYTORCH_VERSION}"
|
PTA_URL="https://sglang-ascend.obs.cn-east-3.myhuaweicloud.com/sglang/torch_npu/torch_npu-2.8.0.post2.dev20251113-cp311-cp311-manylinux_2_28_aarch64.whl"
|
||||||
PTA_NAME="torch_npu-${PYTORCH_VERSION}-cp311-cp311-manylinux_2_28_aarch64.whl"
|
${PIP_INSTALL} ${PTA_URL}
|
||||||
PTA_URL="https://gitcode.com/Ascend/pytorch/releases/download/${PTA_VERSION}/${PTA_NAME}"
|
|
||||||
wget -O "${PTA_NAME}" "${PTA_URL}" && ${PIP_INSTALL} "./${PTA_NAME}"
|
|
||||||
|
|
||||||
|
|
||||||
### Install Triton-Ascend
|
### Install Triton-Ascend
|
||||||
TRITON_ASCEND_NAME="triton_ascend-3.2.0+gitb0ea0850-cp311-cp311-linux_aarch64.whl"
|
TRITON_ASCEND_URL="https://sglang-ascend.obs.cn-east-3.myhuaweicloud.com/sglang/triton_ascend/triton_ascend-3.2.0.dev2025112116-cp311-cp311-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl"
|
||||||
TRITON_ASCEND_URL="https://sglang-ascend.obs.cn-east-3.myhuaweicloud.com/sglang/triton_ascend-3.2.0%2Bgitb0ea0850-cp311-cp311-linux_aarch64.whl"
|
${PIP_INSTALL} ${TRITON_ASCEND_URL}
|
||||||
${PIP_INSTALL} attrs==24.2.0 numpy==1.26.4 scipy==1.13.1 decorator==5.1.1 psutil==6.0.0 pytest==8.3.2 pytest-xdist==3.6.1 pyyaml pybind11
|
|
||||||
wget -O "${TRITON_ASCEND_NAME}" "${TRITON_ASCEND_URL}" && ${PIP_INSTALL} "./${TRITON_ASCEND_NAME}"
|
|
||||||
|
|
||||||
|
|
||||||
### Install BiSheng
|
### Install BiSheng
|
||||||
BISHENG_NAME="Ascend-BiSheng-toolkit_aarch64.run"
|
BISHENG_NAME="Ascend-BiSheng-toolkit_aarch64_20251121.run"
|
||||||
BISHENG_URL="https://sglang-ascend.obs.cn-east-3.myhuaweicloud.com/sglang/${BISHENG_NAME}"
|
BISHENG_URL="https://sglang-ascend.obs.cn-east-3.myhuaweicloud.com/sglang/triton_ascend/${BISHENG_NAME}"
|
||||||
wget -O "${BISHENG_NAME}" "${BISHENG_URL}" && chmod a+x "${BISHENG_NAME}" && "./${BISHENG_NAME}" --install && rm "${BISHENG_NAME}"
|
wget -O "${BISHENG_NAME}" "${BISHENG_URL}" && chmod a+x "${BISHENG_NAME}" && "./${BISHENG_NAME}" --install && rm "${BISHENG_NAME}"
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user