[Feature] Optimize DeepSeek's DeepEP on Ascend NPU (#8355)
Co-authored-by: ronnie_zheng <zl19940307@163.com> Co-authored-by: Hexq0210 <hexq0809521@gmail.com>
This commit is contained in:
@@ -3,7 +3,18 @@ from __future__ import annotations
|
||||
import importlib
|
||||
import sys
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Mapping, Optional, Tuple, Union, cast
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
List,
|
||||
Mapping,
|
||||
Optional,
|
||||
Tuple,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
import torch
|
||||
from torch.nn.parameter import Parameter
|
||||
@@ -79,22 +90,16 @@ def npu_wrapper_rmsnorm_forward(func):
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
if not x.is_contiguous():
|
||||
x = x.contiguous()
|
||||
original_dtype = x.dtype
|
||||
x = x.to(torch.float32)
|
||||
if residual is not None:
|
||||
x = x + residual.to(torch.float32)
|
||||
residual = x.to(original_dtype)
|
||||
out, _, residual_out = torch_npu.npu_add_rms_norm(
|
||||
residual, x, self.weight.data, self.variance_epsilon
|
||||
)
|
||||
out = out + self.bias
|
||||
return out.to(x.dtype), residual_out
|
||||
|
||||
x = (
|
||||
torch_npu.npu_rms_norm(
|
||||
x, self.weight.to(torch.float32), self.variance_epsilon
|
||||
)[0]
|
||||
+ self.bias
|
||||
)
|
||||
|
||||
if residual is None:
|
||||
return x.to(original_dtype)
|
||||
return x.to(original_dtype), residual
|
||||
out = torch_npu.npu_rms_norm(x, self.weight.data, self.variance_epsilon)[0]
|
||||
out = out + self.bias
|
||||
return out.to(x.dtype)
|
||||
|
||||
return _rmsnorm_forward_oot
|
||||
|
||||
@@ -571,8 +576,10 @@ class NPU_W8A8LinearMethodImpl:
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
tp_rank: Optional[int] = 0,
|
||||
) -> torch.Tensor:
|
||||
# To prevent import loops
|
||||
from sglang.srt.layers.linear import RowParallelLinear
|
||||
|
||||
original_dtype = x.dtype
|
||||
if original_dtype != torch.int8:
|
||||
x = torch_npu.npu_quantize(
|
||||
@@ -583,8 +590,12 @@ class NPU_W8A8LinearMethodImpl:
|
||||
-1,
|
||||
True,
|
||||
)
|
||||
|
||||
quant_bias = layer.quant_bias if tp_rank == 0 else None
|
||||
# Only fuse bias add into GEMM for rank 0 (this ensures that
|
||||
# bias will not get added more than once in Attention TP>1 case)
|
||||
if isinstance(layer, RowParallelLinear) and layer.tp_rank > 0:
|
||||
quant_bias = None
|
||||
else:
|
||||
quant_bias = layer.quant_bias
|
||||
return torch_npu.npu_quant_matmul(
|
||||
x,
|
||||
layer.weight,
|
||||
@@ -651,13 +662,21 @@ class NPU_W8A8LinearMethodMTImpl:
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
tp_rank: Optional[int] = 0,
|
||||
) -> torch.Tensor:
|
||||
# To prevent import loops
|
||||
from sglang.srt.layers.linear import RowParallelLinear
|
||||
|
||||
original_dtype = x.dtype
|
||||
if original_dtype != torch.int8:
|
||||
x = quant_per_tensor(x, layer.input_scale, layer.input_offset)
|
||||
|
||||
quant_bias = layer.quant_bias if tp_rank == 0 else None
|
||||
# Only fuse bias add into GEMM for rank 0 (this ensures that
|
||||
# bias will not get added more than once in Attention TP>1 case)
|
||||
if isinstance(layer, RowParallelLinear) and layer.tp_rank > 0:
|
||||
quant_bias = None
|
||||
else:
|
||||
quant_bias = layer.quant_bias
|
||||
|
||||
return ops.quant_matmul(
|
||||
x=x, weight=layer.weight, deq_scale=layer.deq_scale, deq_bias=quant_bias
|
||||
)
|
||||
@@ -737,11 +756,6 @@ class NPU_W8A8LinearMethod(LinearMethodBase):
|
||||
x: torch.Tensor,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
from sglang.srt.layers.linear import RowParallelLinear
|
||||
|
||||
if isinstance(layer, RowParallelLinear):
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
return self.quant_method.apply(layer, x, bias, tp_rank)
|
||||
return self.quant_method.apply(layer, x, bias)
|
||||
|
||||
|
||||
@@ -780,7 +794,6 @@ class NPU_W8A8DynamicLinearMethodImpl:
|
||||
tp_rank: Optional[int] = 0,
|
||||
) -> torch.Tensor:
|
||||
original_dtype = x.dtype
|
||||
# use ATB quantize
|
||||
quant_out, dynamic_scale = torch_npu.npu_dynamic_quant(x)
|
||||
return torch_npu.npu_quant_matmul(
|
||||
quant_out,
|
||||
@@ -863,11 +876,6 @@ class NPU_W8A8DynamicLinearMethod(LinearMethodBase):
|
||||
x: torch.Tensor,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
from sglang.srt.layers.linear import RowParallelLinear
|
||||
|
||||
if isinstance(layer, RowParallelLinear):
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
return self.quant_method.apply(layer, x, bias, tp_rank)
|
||||
return self.quant_method.apply(layer, x, bias)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user