diff --git a/docs/references/environment_variables.md b/docs/references/environment_variables.md index d8419c48e..f378305c9 100644 --- a/docs/references/environment_variables.md +++ b/docs/references/environment_variables.md @@ -122,7 +122,6 @@ SGLang supports various environment variables that can be used to configure its | Environment Variable | Description | Default Value | | --- | --- | --- | | `SGLANG_INT4_WEIGHT` | Enable INT4 weight quantization | `false` | -| `SGLANG_PER_TOKEN_GROUP_QUANT_8BIT_V2` | Apply per token group quantization kernel with fused silu and mul and masked m | `false` | | `SGLANG_FORCE_FP8_MARLIN` | Force using FP8 MARLIN kernels even if other FP8 kernels are available | `false` | | `SGLANG_NVFP4_CKPT_FP8_GEMM_IN_ATTN` | Quantize q_b_proj from BF16 to FP8 when launching DeepSeek NVFP4 checkpoint | `false` | | `SGLANG_MOE_NVFP4_DISPATCH` | Use nvfp4 for moe dispatch (on flashinfer_cutlass or flashinfer_cutedsl moe runner backend) | `"false"` | diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 78babcb7a..d5afb072d 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -3,7 +3,7 @@ import subprocess import warnings from contextlib import ExitStack, contextmanager from enum import IntEnum -from typing import Any +from typing import Any, Optional @contextmanager @@ -341,7 +341,6 @@ class Envs: SGLANG_FORCE_FP8_MARLIN = EnvBool(False) SGLANG_MOE_NVFP4_DISPATCH = EnvBool(False) SGLANG_NVFP4_CKPT_FP8_GEMM_IN_ATTN = EnvBool(False) - SGLANG_PER_TOKEN_GROUP_QUANT_8BIT_V2 = EnvBool(False) SGLANG_NVFP4_CKPT_FP8_NEXTN_MOE = EnvBool(False) SGLANG_QUANT_ALLOW_DOWNCASTING = EnvBool(False) SGLANG_FP8_IGNORED_LAYERS = EnvStr("") @@ -552,12 +551,15 @@ envs = Envs() EnvField._allow_set_name = False -def _print_deprecated_env(new_name: str, old_name: str): +def _print_deprecated_env(old_name: str, new_name: Optional[str] = None): if old_name in os.environ: - warnings.warn( - f"Environment variable {old_name} will be deprecated, please use {new_name} instead" - ) - os.environ[new_name] = os.environ[old_name] + if new_name is None: + warnings.warn(f"Environment variable {old_name} has been deprecated.") + else: + warnings.warn( + f"Environment variable {old_name} will be deprecated, please use {new_name} instead" + ) + os.environ[new_name] = os.environ[old_name] def _warn_deprecated_env_to_cli_flag(env_name: str, suggestion: str): @@ -570,14 +572,15 @@ def _warn_deprecated_env_to_cli_flag(env_name: str, suggestion: str): def _convert_SGL_to_SGLANG(): - _print_deprecated_env("SGLANG_LOG_GC", "SGLANG_GC_LOG") + _print_deprecated_env("SGLANG_GC_LOG", "SGLANG_LOG_GC") _print_deprecated_env( - "SGLANG_MOE_NVFP4_DISPATCH", "SGLANG_CUTEDSL_MOE_NVFP4_DISPATCH" + "SGLANG_CUTEDSL_MOE_NVFP4_DISPATCH", "SGLANG_MOE_NVFP4_DISPATCH" ) _print_deprecated_env( - "SGLANG_ENABLE_TP_MEMORY_INBALANCE_CHECK", "SGL_DISABLE_TP_MEMORY_INBALANCE_CHECK", + "SGLANG_ENABLE_TP_MEMORY_INBALANCE_CHECK", ) + _print_deprecated_env("SGLANG_PER_TOKEN_GROUP_QUANT_8BIT_V2") _deprecated_ms_to_s = { "SGLANG_QUEUED_TIMEOUT_MS": "SGLANG_REQ_WAITING_TIMEOUT", "SGLANG_FORWARD_TIMEOUT_MS": "SGLANG_REQ_RUNNING_TIMEOUT", diff --git a/python/sglang/srt/layers/quantization/fp8_kernel.py b/python/sglang/srt/layers/quantization/fp8_kernel.py index 016c836fe..bf7b6ba17 100644 --- a/python/sglang/srt/layers/quantization/fp8_kernel.py +++ b/python/sglang/srt/layers/quantization/fp8_kernel.py @@ -324,10 +324,6 @@ def _per_token_group_quant_8bit_raw( return x_q, x_s -# backward compatibility -per_token_group_quant_fp8 = _per_token_group_quant_8bit_raw - - def _per_token_group_quant_8bit_fuse_silu_and_mul( x: torch.Tensor, group_size: int, @@ -507,6 +503,11 @@ def sglang_per_token_group_quant_fp8( scale_ue8m0=scale_ue8m0, ) + # Enable v2 kernel by default on supported group sizes + _V2_KERNEL_SUPPORTED_GROUP_SIZES = [16, 32, 64, 128] + if enable_v2 is None: + enable_v2 = group_size in _V2_KERNEL_SUPPORTED_GROUP_SIZES + if x.shape[0] > 0: # Temporary if enable_sgl_per_token_group_quant_8bit: @@ -606,6 +607,12 @@ def sglang_per_token_quant_fp8( return x_q, x_s +if _is_cuda: + per_token_group_quant_fp8 = sglang_per_token_group_quant_fp8 +else: + per_token_group_quant_fp8 = _per_token_group_quant_8bit_raw + + @triton.jit def _static_quant_fp8( # Pointers to inputs and output diff --git a/sgl-kernel/python/sgl_kernel/gemm.py b/sgl-kernel/python/sgl_kernel/gemm.py index 6d320aa9d..a6e65cd6b 100644 --- a/sgl-kernel/python/sgl_kernel/gemm.py +++ b/sgl-kernel/python/sgl_kernel/gemm.py @@ -109,10 +109,9 @@ def sgl_per_token_group_quant_8bit( masked_m: Optional[torch.Tensor] = None, enable_v2: Optional[bool] = None, ) -> None: + _V2_KERNEL_SUPPORTED_GROUP_SIZES = [16, 32, 64, 128] if enable_v2 is None: - from sglang.srt.utils import get_bool_env_var - - enable_v2 = get_bool_env_var("SGLANG_PER_TOKEN_GROUP_QUANT_8BIT_V2") + enable_v2 = group_size in _V2_KERNEL_SUPPORTED_GROUP_SIZES if enable_v2: return torch.ops.sgl_kernel.sgl_per_token_group_quant_8bit_v2.default( diff --git a/test/registered/quant/test_fp8_kernel.py b/test/registered/quant/test_fp8_kernel.py index 5ae5c0485..dcd5ce057 100644 --- a/test/registered/quant/test_fp8_kernel.py +++ b/test/registered/quant/test_fp8_kernel.py @@ -96,7 +96,9 @@ class TestPerTokenGroupQuantFP8(TestFP8Base): A, A_quant_gt, scale_gt = self._make_A( M=self.M, K=self.K, group_size=self.group_size, out_dtype=self.quant_type ) - A_quant, scale = per_token_group_quant_fp8(x=A, group_size=self.group_size) + A_quant, scale = per_token_group_quant_fp8( + x=A.to(torch.bfloat16), group_size=self.group_size + ) torch.testing.assert_close(scale, scale_gt) diff = (A_quant.to(torch.float16) - A_quant_gt.to(torch.float16)).abs() diff_count = (diff > 1e-5).count_nonzero()