[NPU]TP Communications compression For Qwen3 models for NPU (#20520)
Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
co-authored by
ronnie_zheng
parent
ebbaab5597
commit
83bf5d6869
@@ -18,6 +18,11 @@ def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
|
||||
return get_tp_group().all_reduce(input_)
|
||||
|
||||
|
||||
def tensor_model_parallel_quant_all_reduce(input_: torch.Tensor) -> torch.Tensor:
|
||||
"""All-reduce the input tensor across model parallel group."""
|
||||
return get_tp_group().quant_all_reduce(input_)
|
||||
|
||||
|
||||
def tensor_model_parallel_fused_allreduce_rmsnorm(
|
||||
input_: torch.Tensor,
|
||||
residual_inp_: torch.Tensor,
|
||||
@@ -60,6 +65,13 @@ def attention_tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Te
|
||||
return get_attn_tp_group().all_reduce(input_)
|
||||
|
||||
|
||||
def attention_tensor_model_parallel_quant_all_reduce(
|
||||
input_: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""All-reduce the input tensor across attention parallel group."""
|
||||
return get_attn_tp_group().quant_all_reduce(input_)
|
||||
|
||||
|
||||
def moe_tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
|
||||
"""All-reduce the input tensor across moe parallel group."""
|
||||
return get_moe_tp_group().all_reduce(input_)
|
||||
|
||||
@@ -4,11 +4,16 @@ from torch.distributed import ProcessGroup
|
||||
|
||||
from sglang.srt.utils import is_npu
|
||||
|
||||
_is_npu = is_npu()
|
||||
|
||||
if _is_npu:
|
||||
from torch_npu import npu_dynamic_quant
|
||||
|
||||
|
||||
class NpuCommunicator:
|
||||
|
||||
def __init__(self, group: ProcessGroup):
|
||||
if not is_npu():
|
||||
if not _is_npu:
|
||||
self.disabled = True
|
||||
return
|
||||
self.disabled = False
|
||||
@@ -19,6 +24,33 @@ class NpuCommunicator:
|
||||
dist.all_reduce(x, group=self.group)
|
||||
return x
|
||||
|
||||
def quant_all_reduce(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Note:
|
||||
All reduce is split into All gather + reduce.
|
||||
All gather is performed in low precision, but reduce in full precision.
|
||||
"""
|
||||
world_size = self.world_size
|
||||
input_size = x.size()
|
||||
output_size = (input_size[0] * world_size,) + input_size[1:]
|
||||
x_q, scale = npu_dynamic_quant(x, dst_type=torch.int8)
|
||||
# Allocate output tensor.
|
||||
output_tensor = torch.empty(output_size, dtype=x_q.dtype, device=x.device)
|
||||
output_scale = torch.empty(
|
||||
output_size[:1], dtype=scale.dtype, device=scale.device
|
||||
)
|
||||
# All-gather.
|
||||
dist.all_gather_into_tensor(output_tensor, x_q, group=self.group)
|
||||
dist.all_gather_into_tensor(output_scale, scale, group=self.group)
|
||||
|
||||
output_tensor = output_tensor.to(x.dtype) * output_scale.unsqueeze(-1).to(
|
||||
x.dtype
|
||||
)
|
||||
# Reshape
|
||||
output_tensor = output_tensor.reshape((world_size,) + input_size)
|
||||
|
||||
return output_tensor.sum(dim=0)
|
||||
|
||||
def all_gather(self, x: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
world_size = self.world_size
|
||||
if dim < 0:
|
||||
|
||||
@@ -633,6 +633,20 @@ class GroupCoordinator:
|
||||
inplace_all_reduce(input_, group_name=self.unique_name)
|
||||
return input_
|
||||
|
||||
def quant_all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
User-facing quant-all-reduce function similar to all-reduce. (NPU support only)
|
||||
"""
|
||||
# Bypass the function if we are using only 1 GPU.
|
||||
if self.world_size == 1:
|
||||
return input_
|
||||
|
||||
if self.npu_communicator is not None and not self.npu_communicator.disabled:
|
||||
return self.npu_communicator.quant_all_reduce(input_)
|
||||
else:
|
||||
inplace_all_reduce(input_, group_name=self.unique_name)
|
||||
return input_
|
||||
|
||||
def fused_allreduce_rmsnorm(
|
||||
self,
|
||||
input_: torch.Tensor,
|
||||
|
||||
@@ -22,6 +22,7 @@ import torch
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
attention_tensor_model_parallel_all_reduce,
|
||||
attention_tensor_model_parallel_quant_all_reduce,
|
||||
get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size,
|
||||
get_tp_group,
|
||||
@@ -1000,9 +1001,18 @@ class CommunicateWithAllReduceAndLayerNormFn:
|
||||
handled = True
|
||||
|
||||
if not handled:
|
||||
hidden_states = attention_tensor_model_parallel_all_reduce(
|
||||
hidden_states
|
||||
quantize_communications = (
|
||||
not forward_batch.forward_mode.is_decode_or_idle()
|
||||
and get_global_server_args().enable_quant_communications
|
||||
)
|
||||
if quantize_communications:
|
||||
hidden_states = attention_tensor_model_parallel_quant_all_reduce(
|
||||
hidden_states
|
||||
)
|
||||
else:
|
||||
hidden_states = attention_tensor_model_parallel_all_reduce(
|
||||
hidden_states
|
||||
)
|
||||
if _is_npu and context.cache is not None:
|
||||
_ = prepare_weight_cache(hidden_states, context.cache)
|
||||
hidden_states, residual = layernorm(hidden_states, residual)
|
||||
|
||||
@@ -19,6 +19,7 @@ from sglang.srt.distributed import (
|
||||
split_tensor_along_last_dim,
|
||||
tensor_model_parallel_all_gather,
|
||||
tensor_model_parallel_all_reduce,
|
||||
tensor_model_parallel_quant_all_reduce,
|
||||
)
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
use_symmetric_memory,
|
||||
@@ -37,6 +38,7 @@ from sglang.srt.layers.parameter import (
|
||||
_ColumnvLLMParameter,
|
||||
)
|
||||
from sglang.srt.layers.utils import pad_or_narrow_weight
|
||||
from sglang.srt.server_args import get_global_server_args
|
||||
from sglang.srt.utils import get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -1509,7 +1511,7 @@ class RowParallelLinear(LinearBase):
|
||||
# Fallback for parameters that don't accept additional args
|
||||
param.load_row_parallel_weight(loaded_weight)
|
||||
|
||||
def forward(self, input_, skip_all_reduce=False):
|
||||
def forward(self, input_, skip_all_reduce=False, forward_batch=None):
|
||||
if self.input_is_parallel:
|
||||
input_parallel = input_
|
||||
else:
|
||||
@@ -1536,7 +1538,18 @@ class RowParallelLinear(LinearBase):
|
||||
if self.use_dp_attention_reduce:
|
||||
output = get_attention_tp_group().all_reduce(output_parallel)
|
||||
else:
|
||||
output = tensor_model_parallel_all_reduce(output_parallel)
|
||||
quantize_communications = (
|
||||
(
|
||||
not forward_batch.forward_mode.is_decode_or_idle()
|
||||
and get_global_server_args().enable_quant_communications
|
||||
)
|
||||
if forward_batch is not None
|
||||
else False
|
||||
)
|
||||
if quantize_communications:
|
||||
output = tensor_model_parallel_quant_all_reduce(output_parallel)
|
||||
else:
|
||||
output = tensor_model_parallel_all_reduce(output_parallel)
|
||||
else:
|
||||
output = output_parallel
|
||||
|
||||
|
||||
@@ -676,7 +676,7 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
|
||||
)
|
||||
return lora_output
|
||||
|
||||
def forward(self, input_: torch.Tensor, skip_all_reduce=False):
|
||||
def forward(self, input_: torch.Tensor, skip_all_reduce=False, forward_batch=None):
|
||||
if self.base_layer.input_is_parallel:
|
||||
input_parallel = input_
|
||||
else:
|
||||
|
||||
@@ -91,13 +91,17 @@ class Qwen2MLP(nn.Module):
|
||||
)
|
||||
self.act_fn = SiluAndMul()
|
||||
|
||||
def forward(self, x):
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
forward_batch: ForwardBatch = None,
|
||||
) -> torch.Tensor:
|
||||
if get_global_server_args().rl_on_policy_target is not None:
|
||||
x = x.bfloat16()
|
||||
|
||||
gate_up, _ = self.gate_up_proj(x)
|
||||
x = self.act_fn(gate_up)
|
||||
x, _ = self.down_proj(x)
|
||||
x, _ = self.down_proj(x, forward_batch=forward_batch)
|
||||
return x
|
||||
|
||||
|
||||
|
||||
@@ -419,7 +419,7 @@ class Qwen3DecoderLayer(nn.Module):
|
||||
else None
|
||||
),
|
||||
)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
hidden_states = self.mlp(hidden_states, forward_batch=forward_batch)
|
||||
if _is_npu and get_cmo_stream():
|
||||
wait_cmo_stream()
|
||||
hidden_states, residual = self.layer_communicator.postprocess_layer(
|
||||
|
||||
@@ -773,6 +773,9 @@ class ServerArgs:
|
||||
# For forward hooks
|
||||
forward_hooks: Optional[List[dict[str, Any]]] = None
|
||||
|
||||
# For communications compression
|
||||
enable_quant_communications: Optional[bool] = False
|
||||
|
||||
# For msProbe
|
||||
msprobe_dump_config: Optional[str] = None
|
||||
|
||||
@@ -6654,6 +6657,13 @@ class ServerArgs:
|
||||
help="JSON-formatted forward hook specifications to attach to the model.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--enable-quant-communications",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Enable INT8 quantization of TP communications (limited support).",
|
||||
)
|
||||
|
||||
# For msProbe
|
||||
parser.add_argument(
|
||||
"--msprobe-dump-config",
|
||||
@@ -6929,6 +6939,17 @@ class ServerArgs:
|
||||
"When enabling two batch overlap, moe_a2a_backend cannot be 'none'."
|
||||
)
|
||||
|
||||
# Check communications compression
|
||||
if self.enable_quant_communications and self.tp_size == 1:
|
||||
raise ValueError(
|
||||
"Communications quantization is only used with tp_size != 1"
|
||||
)
|
||||
|
||||
if self.enable_quant_communications and self.device != "npu":
|
||||
raise ValueError(
|
||||
"Communications quantization is only supported for NPU device"
|
||||
)
|
||||
|
||||
if (
|
||||
self.enable_grpc
|
||||
and self.grpc_port is not None
|
||||
|
||||
Reference in New Issue
Block a user