diff --git a/docs_new/docs/sglang-diffusion/environment_variables.mdx b/docs_new/docs/sglang-diffusion/environment_variables.mdx index b864521e9..7558b1725 100644 --- a/docs_new/docs/sglang-diffusion/environment_variables.mdx +++ b/docs_new/docs/sglang-diffusion/environment_variables.mdx @@ -152,6 +152,11 @@ description: "Configure SGLang diffusion behavior with environment variables." not set Optional FlashInfer FP4 GEMM backend override for diffusion NVFP4. When unset, SGLang defaults to flashinfer_trtllm. + + SGLANG_DIFFUSION_ENABLE_W8A8_FP8_GEMM + false + Experimental opt-in for fused W8A8 FP8 GEMM in diffusion weight-only FP8 linears. When disabled, FP8 weights are dequantized to the compute dtype before matmul. Enabling this dynamically quantizes activations to FP8 and may change output quality. + diff --git a/python/sglang/multimodal_gen/envs.py b/python/sglang/multimodal_gen/envs.py index 122acdbe3..5e5182521 100644 --- a/python/sglang/multimodal_gen/envs.py +++ b/python/sglang/multimodal_gen/envs.py @@ -58,6 +58,7 @@ if TYPE_CHECKING: # model loading SGLANG_USE_RUNAI_MODEL_STREAMER: bool = True SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND: str | None = None + SGLANG_DIFFUSION_ENABLE_W8A8_FP8_GEMM: bool = False SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D: str = "auto" SGLANG_USE_CUDA_HUNYUANVIDEO_GROUP_NORM_SILU: bool = False SGLANG_USE_ROCM_VAE: bool = False @@ -305,6 +306,11 @@ environment_variables: dict[str, Callable[[], Any]] = { "SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND": _lazy_str( "SGLANG_DIFFUSION_FLASHINFER_FP4_GEMM_BACKEND" ), + # Experimental opt-in for W8A8 FP8 GEMM in diffusion weight-only FP8 linears. + # When disabled, FP8 weights are dequantized to compute dtype before matmul. + "SGLANG_DIFFUSION_ENABLE_W8A8_FP8_GEMM": _lazy_bool( + "SGLANG_DIFFUSION_ENABLE_W8A8_FP8_GEMM" + ), # ROCm: use AITer GroupNorm in VAE for improved performance "SGLANG_USE_ROCM_VAE": _lazy_bool("SGLANG_USE_ROCM_VAE"), # ROCm: enable cudnn.benchmark (MIOpen auto-tuning) for VAE conv layers diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/weight_only_fp8.py b/python/sglang/multimodal_gen/runtime/layers/quantization/weight_only_fp8.py index 5c84a5754..b7bd5491e 100644 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/weight_only_fp8.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/weight_only_fp8.py @@ -1,9 +1,12 @@ # SPDX-License-Identifier: Apache-2.0 +import logging + import torch import torch.nn as nn import torch.nn.functional as F +import sglang.multimodal_gen.envs as envs from sglang.multimodal_gen.runtime.distributed import ( divide, get_tp_group, @@ -13,6 +16,27 @@ from sglang.multimodal_gen.runtime.layers.utils import get_group_rank, get_group from sglang.multimodal_gen.runtime.models.utils import set_weight_attrs FP8_WEIGHT_DTYPE = torch.float8_e4m3fn +W8A8_FP8_GEMM_ENV = "SGLANG_DIFFUSION_ENABLE_W8A8_FP8_GEMM" + +logger = logging.getLogger(__name__) +_w8a8_fp8_gemm_warning_logged = False + + +def _can_apply_fused_w8a8_fp8_linear( + x: torch.Tensor, + weight: torch.Tensor, + weight_scale: torch.Tensor, + compute_dtype: torch.dtype, +) -> bool: + return ( + x.device.type == "cuda" + and weight.device.type == "cuda" + and weight_scale.device.type == "cuda" + and not x.is_meta + and not weight.is_meta + and not weight_scale.is_meta + and compute_dtype in (torch.float16, torch.bfloat16) + ) def dequantize_rowwise_fp8_weight( @@ -30,8 +54,53 @@ def dequantize_rowwise_fp8_weight( return weight.to(dtype) * weight_scale.to(dtype).unsqueeze(1) +def _apply_srt_w8a8_fp8_linear(*args, **kwargs) -> torch.Tensor: + from sglang.srt.layers.quantization.fp8_utils import apply_fp8_linear + + return apply_fp8_linear(*args, **kwargs) + + +def _is_cutlass_fp8_supported() -> bool: + from sglang.srt.layers.quantization.fp8_utils import cutlass_fp8_supported + + return cutlass_fp8_supported() + + +def _apply_weight_only_fp8_linear( + x: torch.Tensor, + weight: torch.Tensor, + weight_scale: torch.Tensor, + bias: torch.Tensor | None, + compute_dtype: torch.dtype, + enable_fused_w8a8: bool, +) -> torch.Tensor: + x = x.to(compute_dtype) + bias = bias.to(compute_dtype) if bias is not None else None + if enable_fused_w8a8 and _can_apply_fused_w8a8_fp8_linear( + x, weight, weight_scale, compute_dtype + ): + try: + # The fused kernel uses W8A8 compute; fallback keeps BF16/FP16 + # activations after dequantizing the FP8 weights. + output = _apply_srt_w8a8_fp8_linear( + input=x, + weight=weight.t(), + weight_scale=weight_scale, + input_scale=None, + bias=bias, + cutlass_fp8_supported=_is_cutlass_fp8_supported(), + ) + _log_w8a8_fp8_gemm_warning_once() + return output + except (ImportError, NotImplementedError): + pass + + dequant_weight = dequantize_rowwise_fp8_weight(weight, weight_scale, compute_dtype) + return F.linear(x, dequant_weight, bias) + + class WeightOnlyFP8Linear(nn.Module): - """Storage-only e4m3 FP8 linear with row-wise dequantization before matmul.""" + """Storage-only e4m3 FP8 linear with row-wise weight scales.""" def __init__( self, @@ -39,11 +108,13 @@ class WeightOnlyFP8Linear(nn.Module): out_features: int, bias: bool = True, compute_dtype: torch.dtype | None = None, + enable_fused_w8a8: bool | None = None, ) -> None: super().__init__() self.in_features = in_features self.out_features = out_features self.compute_dtype = compute_dtype + self.enable_fused_w8a8 = _resolve_enable_fused_w8a8(enable_fused_w8a8) self.weight = nn.Parameter( torch.empty(out_features, in_features, dtype=FP8_WEIGHT_DTYPE), requires_grad=False, @@ -65,15 +136,18 @@ class WeightOnlyFP8Linear(nn.Module): def forward(self, x: torch.Tensor) -> torch.Tensor: compute_dtype = self.compute_dtype or x.dtype - weight = dequantize_rowwise_fp8_weight( - self.weight, self.weight_scale, compute_dtype + return _apply_weight_only_fp8_linear( + x, + self.weight, + self.weight_scale, + self.bias, + compute_dtype, + self.enable_fused_w8a8, ) - bias = self.bias.to(compute_dtype) if self.bias is not None else None - return F.linear(x.to(compute_dtype), weight, bias) class WeightOnlyFP8ColumnParallelLinear(nn.Module): - """Column-parallel e4m3 FP8 linear with row-wise dequantization.""" + """Column-parallel storage-only e4m3 FP8 linear.""" def __init__( self, @@ -83,12 +157,14 @@ class WeightOnlyFP8ColumnParallelLinear(nn.Module): compute_dtype: torch.dtype | None = None, gather_output: bool = True, tp_group=None, + enable_fused_w8a8: bool | None = None, ) -> None: super().__init__() self.in_features = in_features self.out_features = out_features self.compute_dtype = compute_dtype self.gather_output = gather_output + self.enable_fused_w8a8 = _resolve_enable_fused_w8a8(enable_fused_w8a8) self.tp_group = tp_group or get_tp_group() self.tp_size = get_group_size(self.tp_group) self.tp_rank = get_group_rank(self.tp_group) @@ -154,11 +230,14 @@ class WeightOnlyFP8ColumnParallelLinear(nn.Module): def forward(self, x: torch.Tensor) -> torch.Tensor: compute_dtype = self.compute_dtype or x.dtype - weight = dequantize_rowwise_fp8_weight( - self.weight, self.weight_scale, compute_dtype + output_parallel = _apply_weight_only_fp8_linear( + x, + self.weight, + self.weight_scale, + self.bias, + compute_dtype, + self.enable_fused_w8a8, ) - bias = self.bias.to(compute_dtype) if self.bias is not None else None - output_parallel = F.linear(x.to(compute_dtype), weight, bias) if self.gather_output: return tensor_model_parallel_all_gather( output_parallel, tp_group=self.tp_group @@ -166,6 +245,25 @@ class WeightOnlyFP8ColumnParallelLinear(nn.Module): return output_parallel +def _resolve_enable_fused_w8a8(value: bool | None) -> bool: + if value is not None: + return value + return envs.SGLANG_DIFFUSION_ENABLE_W8A8_FP8_GEMM + + +def _log_w8a8_fp8_gemm_warning_once() -> None: + global _w8a8_fp8_gemm_warning_logged + if _w8a8_fp8_gemm_warning_logged: + return + logger.warning( + "%s=1 enables W8A8 FP8 GEMM for weight-only FP8 linears; activations " + "are dynamically quantized to FP8 and outputs may differ from the " + "official weight-only FP8 path.", + W8A8_FP8_GEMM_ENV, + ) + _w8a8_fp8_gemm_warning_logged = True + + def swap_linears_to_weight_only_fp8(module: nn.Module) -> None: """Recursively replace nn.Linear with WeightOnlyFP8Linear. diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py index 0c7775aa0..1e68c5ba7 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3vl.py @@ -98,8 +98,11 @@ def _make_text_linear( out_features, bias=bias, gather_output=gather_output, + enable_fused_w8a8=False, ) - return WeightOnlyFP8Linear(in_features, out_features, bias=bias) + return WeightOnlyFP8Linear( + in_features, out_features, bias=bias, enable_fused_w8a8=False + ) if quant_config is not None: if use_column_parallel: return Qwen3VLColumnParallelLinear( diff --git a/python/sglang/multimodal_gen/test/unit/test_ideogram4.py b/python/sglang/multimodal_gen/test/unit/test_ideogram4.py index 9e02d1570..9e86cf0e5 100644 --- a/python/sglang/multimodal_gen/test/unit/test_ideogram4.py +++ b/python/sglang/multimodal_gen/test/unit/test_ideogram4.py @@ -34,6 +34,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import ( ) from sglang.multimodal_gen.runtime.layers.quantization.weight_only_fp8 import ( FP8_WEIGHT_DTYPE, + W8A8_FP8_GEMM_ENV, WeightOnlyFP8ColumnParallelLinear, WeightOnlyFP8Linear, dequantize_rowwise_fp8_weight, @@ -909,6 +910,24 @@ class TestIdeogram4(unittest.TestCase): self.assertEqual(model.weight.dtype, FP8_WEIGHT_DTYPE) self.assertEqual(model.weight_scale.dtype, torch.float32) + def test_weight_only_fp8_w8a8_gemm_defaults_to_off(self): + with patch.dict(os.environ, {W8A8_FP8_GEMM_ENV: "0"}): + model = WeightOnlyFP8Linear(3, 2, bias=False) + + self.assertFalse(model.enable_fused_w8a8) + + def test_weight_only_fp8_w8a8_gemm_env_opt_in(self): + with patch.dict(os.environ, {W8A8_FP8_GEMM_ENV: "1"}): + model = WeightOnlyFP8Linear(3, 2, bias=False) + + self.assertTrue(model.enable_fused_w8a8) + + def test_weight_only_fp8_w8a8_gemm_explicit_flag_overrides_env(self): + with patch.dict(os.environ, {W8A8_FP8_GEMM_ENV: "1"}): + model = WeightOnlyFP8Linear(3, 2, bias=False, enable_fused_w8a8=False) + + self.assertFalse(model.enable_fused_w8a8) + def test_ideogram_text_encoder_post_config_hook_preserves_local_arch(self): config = Ideogram4TextEncoderConfig() config.arch_config.architectures = ["RemoteQwen3VLTextModel"] @@ -961,13 +980,20 @@ class TestIdeogram4(unittest.TestCase): set_global_server_args( SimpleNamespace(attention_backend="torch_sdpa", comfyui_mode=False) ) - with torch.device("meta"): + with ( + patch.dict(os.environ, {W8A8_FP8_GEMM_ENV: "1"}), + torch.device("meta"), + ): encoder = IdeogramQwen3VLTextEncoder(config) finally: set_global_server_args(prev_args) - self.assertTrue( - any(isinstance(module, WeightOnlyFP8Linear) for module in encoder.modules()) - ) + fp8_linears = [ + module + for module in encoder.modules() + if isinstance(module, WeightOnlyFP8Linear) + ] + self.assertTrue(fp8_linears) + self.assertTrue(all(not module.enable_fused_w8a8 for module in fp8_linears)) self.assertFalse( any(isinstance(module, torch.nn.Linear) for module in encoder.modules()) )