diff --git a/benchmark/boolq/bench_sglang.py b/benchmark/boolq/bench_sglang.py index b3ce3c996..b38d7a01e 100644 --- a/benchmark/boolq/bench_sglang.py +++ b/benchmark/boolq/bench_sglang.py @@ -4,7 +4,7 @@ import time import numpy as np -from sglang.api import set_default_backend +from sglang.lang.api import set_default_backend from sglang.test.test_utils import ( add_common_sglang_args_and_parse, select_sglang_backend, diff --git a/docs/advanced_features/server_arguments.md b/docs/advanced_features/server_arguments.md index 7a730cd06..7675b4bbb 100644 --- a/docs/advanced_features/server_arguments.md +++ b/docs/advanced_features/server_arguments.md @@ -119,6 +119,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s | `--modelopt-export-path` | Path to export the quantized model in HuggingFace format after ModelOpt quantization. The exported model can then be used directly with SGLang for inference. If not provided, the model will not be exported. | `None` | Type: str | | `--quantize-and-serve` | Quantize the model with ModelOpt and immediately serve it without exporting. This is useful for development and prototyping. For production, it's recommended to use separate quantization and deployment steps. | `False` | bool flag (set to enable) | | `--rl-quant-profile` | Path to the FlashRL quantization profile. Required when using --load-format flash_rl. | `None` | Type: str | +| `--enable-quant-communications` | Enable INT8 quantization of TP communications (Supported only for NPU for Qwen3 series). | `False` | bool flag (set to enable) | ## Memory and scheduling | Argument | Description | Defaults | Options | diff --git a/python/sglang/srt/distributed/communication_op.py b/python/sglang/srt/distributed/communication_op.py index ea1318b6e..de83c9c81 100644 --- a/python/sglang/srt/distributed/communication_op.py +++ b/python/sglang/srt/distributed/communication_op.py @@ -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_) diff --git a/python/sglang/srt/distributed/device_communicators/npu_communicator.py b/python/sglang/srt/distributed/device_communicators/npu_communicator.py index cb6eb88e3..5518584d4 100644 --- a/python/sglang/srt/distributed/device_communicators/npu_communicator.py +++ b/python/sglang/srt/distributed/device_communicators/npu_communicator.py @@ -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: diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 5b92f9e18..c42760ce5 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -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, diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index d0096f4d0..853cf3ad5 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -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) diff --git a/python/sglang/srt/layers/linear.py b/python/sglang/srt/layers/linear.py index ddd5b6a4f..c69a06436 100644 --- a/python/sglang/srt/layers/linear.py +++ b/python/sglang/srt/layers/linear.py @@ -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 diff --git a/python/sglang/srt/lora/layers.py b/python/sglang/srt/lora/layers.py index b23ba29d2..460591861 100644 --- a/python/sglang/srt/lora/layers.py +++ b/python/sglang/srt/lora/layers.py @@ -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: diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py index 55659586d..39e404884 100644 --- a/python/sglang/srt/models/qwen2.py +++ b/python/sglang/srt/models/qwen2.py @@ -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 diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index 7df12b8ed..30333f999 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -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( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 6b8ac5277..c5bba8351 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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 diff --git a/test/registered/ascend/llm_models/test_npu_llama_2_7b_communications_compression.py b/test/registered/ascend/llm_models/test_npu_llama_2_7b_communications_compression.py new file mode 100644 index 000000000..5a45eb3b0 --- /dev/null +++ b/test/registered/ascend/llm_models/test_npu_llama_2_7b_communications_compression.py @@ -0,0 +1,37 @@ +import unittest + +from sglang.test.ascend.gsm8k_ascend_mixin import GSM8KAscendMixin +from sglang.test.ascend.test_ascend_utils import LLAMA_2_7B_WEIGHTS_PATH +from sglang.test.ci.ci_register import register_npu_ci +from sglang.test.test_utils import CustomTestCase + +register_npu_ci(est_time=400, suite="nightly-2-npu-a3") + + +class TestLlama(GSM8KAscendMixin, CustomTestCase): + """Testcase: Verify that the inference accuracy of the LLM-Research/Llama-2-7B model on the GSM8K dataset with tp communications quantization is no less than 0.18. + + [Test Category] Model + [Test Target] LLM-Research/Llama-2-7B + """ + + model = LLAMA_2_7B_WEIGHTS_PATH + accuracy = 0.18 + other_args = [ + "--trust-remote-code", + "--mem-fraction-static", + 0.8, + "--max-running-requests", + 32, + "--attention-backend", + "ascend", + "--cuda-graph-max-bs", + 32, + "--tp-size", + 2, + "--enable-quant-communications", + ] + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/ascend/llm_models/test_npu_qwen3_8b_communications_quantization.py b/test/registered/ascend/llm_models/test_npu_qwen3_8b_communications_quantization.py new file mode 100644 index 000000000..5c23e336f --- /dev/null +++ b/test/registered/ascend/llm_models/test_npu_qwen3_8b_communications_quantization.py @@ -0,0 +1,37 @@ +import unittest + +from sglang.test.ascend.gsm8k_ascend_mixin import GSM8KAscendMixin +from sglang.test.ascend.test_ascend_utils import QWEN3_8B_WEIGHTS_PATH +from sglang.test.ci.ci_register import register_npu_ci +from sglang.test.test_utils import CustomTestCase + +register_npu_ci(est_time=400, suite="nightly-2-npu-a3") + + +class TestQwen38BCommQuantization(GSM8KAscendMixin, CustomTestCase): + """Testcase: Verify that the inference accuracy of the Qwen/Qwen3-8B model with TP communications quantization on the GSM8K dataset is no less than 0.85. + + [Test Category] Model + [Test Target] Qwen/Qwen3-8B + """ + + model = QWEN3_8B_WEIGHTS_PATH + accuracy = 0.85 + other_args = [ + "--trust-remote-code", + "--mem-fraction-static", + 0.8, + "--max-running-requests", + 32, + "--attention-backend", + "ascend", + "--cuda-graph-max-bs", + 32, + "--tp-size", + 2, + "--enable-quant-communications", + ] + + +if __name__ == "__main__": + unittest.main()