diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 46e9c70dc..335879a17 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -983,6 +983,9 @@ class Envs: SGLANG_MOE_NVFP4_DISPATCH = EnvBool(False) SGLANG_NVFP4_CKPT_FP8_GEMM_IN_ATTN = EnvBool(False) SGLANG_NVFP4_CKPT_FP8_NEXTN_MOE = EnvBool(False) + # GLM NextN (MTP): cast the draft layer's bf16 fused MoE to per-channel FP8 + # on load. Unrelated to the NVFP4 block-FP8 NextN path above. + SGLANG_GLM_NEXTN_MOE_PTPC = EnvBool(False) SGLANG_QUANT_ALLOW_DOWNCASTING = EnvBool(False) SGLANG_FP8_IGNORED_LAYERS = EnvStr("") SGLANG_FP4_IGNORED_LAYERS = EnvStr("") diff --git a/python/sglang/srt/layers/quantization/quark/schemes/quark_w8a8_fp8_moe.py b/python/sglang/srt/layers/quantization/quark/schemes/quark_w8a8_fp8_moe.py index 673ab108b..33778e071 100644 --- a/python/sglang/srt/layers/quantization/quark/schemes/quark_w8a8_fp8_moe.py +++ b/python/sglang/srt/layers/quantization/quark/schemes/quark_w8a8_fp8_moe.py @@ -13,7 +13,7 @@ from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo from sglang.srt.layers.quantization.fp8_utils import normalize_e4m3fn_to_e4m3fnuz from sglang.srt.layers.quantization.quark.schemes import QuarkMoEScheme from sglang.srt.layers.quantization.utils import all_close_1d, per_tensor_dequantize -from sglang.srt.utils import get_bool_env_var, is_hip, set_weight_attrs +from sglang.srt.utils import get_bool_env_var, is_hip, print_info_once, set_weight_attrs if TYPE_CHECKING: from sglang.srt.layers.moe.token_dispatcher import ( @@ -31,8 +31,6 @@ _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip if _use_aiter: from aiter.ops.shuffle import shuffle_weight - from sglang.kernels.ops.moe.rocm_moe_utils import rocm_fused_experts_tkw1 - class QuarkW8A8FP8MoE(QuarkMoEScheme): def __init__(self, weight_config: dict[str, Any], input_config: dict[str, Any]): @@ -238,74 +236,84 @@ class QuarkW8A8FP8MoE(QuarkMoEScheme): f"Unsupported weight quantization strategy: {self.weight_qscheme}." ) - if ( - _use_aiter - and self.is_weight_per_channel - and self.moe_runner_config.apply_router_weight_on_input - ): + # Triton reads the canonical layout; only AITER wants the shuffled one, + # which aiter.fused_moe selects on via the is_shuffled tag. + if _use_aiter and self.runner.runner_backend.is_aiter(): with torch.no_grad(): - # Pre-shuffle weights layer.w13_weight = torch.nn.Parameter( shuffle_weight(layer.w13_weight.data, (16, 16)), requires_grad=False, ) + layer.w13_weight.is_shuffled = True torch.cuda.empty_cache() layer.w2_weight = torch.nn.Parameter( shuffle_weight(layer.w2_weight.data, (16, 16)), requires_grad=False, ) + layer.w2_weight.is_shuffled = True torch.cuda.empty_cache() def create_moe_runner( self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig ): + from sglang.srt.layers.moe.utils import ( + get_moe_a2a_backend, + get_moe_runner_backend, + ) + self.moe_runner_config = moe_runner_config - self.runner = MoeRunner(MoeRunnerBackend.TRITON, moe_runner_config) + moe_runner_backend = get_moe_runner_backend() + a2a_supports_aiter = get_moe_a2a_backend().supports_aiter() + # AITER's per_Token fused MoE needs per-channel weight scales; the + # per-tensor scheme has no equivalent there. + use_aiter_runner = ( + _use_aiter + and self.is_weight_per_channel + and a2a_supports_aiter + and (moe_runner_backend.is_auto() or moe_runner_backend.is_aiter()) + ) + self.runner = MoeRunner( + MoeRunnerBackend.AITER if use_aiter_runner else MoeRunnerBackend.TRITON, + moe_runner_config, + ) + print_info_once( + f"QuarkW8A8FP8MoE runner={self.runner.runner_backend.value} " + f"(use_aiter={_use_aiter} per_channel={self.is_weight_per_channel} " + f"a2a_supports_aiter={a2a_supports_aiter} " + f"requested={moe_runner_backend.value})" + ) def apply_weights( self, layer: torch.nn.Module, dispatch_output: StandardDispatchOutput, ) -> CombineInput: - - from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput - - x = dispatch_output.hidden_states - topk_output = dispatch_output.topk_output - - moe_runner_config = self.moe_runner_config - - if ( - _use_aiter - and self.is_weight_per_channel - and moe_runner_config.apply_router_weight_on_input - ): - topk_weights, topk_ids, _ = topk_output - output = rocm_fused_experts_tkw1( - hidden_states=x, - w1=layer.w13_weight, - w2=layer.w2_weight, - topk_weights=topk_weights, - topk_ids=topk_ids, - activation=moe_runner_config.activation, - apply_router_weight_on_input=moe_runner_config.apply_router_weight_on_input, - use_fp8_w8a8=True, - per_channel_quant=self.is_weight_per_channel, - w1_scale=layer.w13_weight_scale, - w2_scale=layer.w2_weight_scale, - a1_scale=layer.w13_input_scale, - a2_scale=layer.w2_input_scale, + if self.runner.runner_backend.is_aiter(): + from sglang.srt.layers.moe.moe_runner.aiter import ( + AiterMoeQuantInfo, + AiterQuantType, ) - return StandardCombineInput(hidden_states=output) - else: - quant_info = TritonMoeQuantInfo( + + quant_info = AiterMoeQuantInfo( w13_weight=layer.w13_weight, w2_weight=layer.w2_weight, - use_fp8_w8a8=True, - per_channel_quant=self.is_weight_per_channel, + quant_type=AiterQuantType.PER_TOKEN, w13_scale=layer.w13_weight_scale, w2_scale=layer.w2_weight_scale, a13_scale=layer.w13_input_scale, a2_scale=layer.w2_input_scale, + expert_mask=layer.dispatcher.expert_mask_gpu, ) return self.runner.run(dispatch_output, quant_info) + + quant_info = TritonMoeQuantInfo( + w13_weight=layer.w13_weight, + w2_weight=layer.w2_weight, + use_fp8_w8a8=True, + per_channel_quant=self.is_weight_per_channel, + w13_scale=layer.w13_weight_scale, + w2_scale=layer.w2_weight_scale, + a13_scale=layer.w13_input_scale, + a2_scale=layer.w2_input_scale, + ) + return self.runner.run(dispatch_output, quant_info) diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py index fbda6fdc0..91584d233 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -14,6 +14,7 @@ """Inference-only GLM-4.5, GLM-4.6 and GLM-4.7 model compatible with HuggingFace weights""" +import copy import logging import re from typing import Any, Dict, Iterable, List, Optional, Tuple, Union @@ -95,6 +96,7 @@ from sglang.srt.utils import ( is_hip, is_non_idle_and_non_empty, is_npu, + log_info_on_rank0, make_layers, ) from sglang.srt.utils.hf_transformers_utils import get_rope_config @@ -110,6 +112,42 @@ _device_sm = get_device_sm() logger = logging.getLogger(__name__) +_GLM_NEXTN_EXPERT_PROJ_RE = re.compile( + r"mlp\.experts\.\d+\.(gate_proj|up_proj|down_proj)\.weight$" +) + + +def enable_glm_nextn_moe_ptpc( + quant_config: Optional[QuantizationConfig], +) -> bool: + return ( + envs.SGLANG_GLM_NEXTN_MOE_PTPC.get() + and quant_config is not None + and quant_config.get_name() == "quark" + ) + + +def glm_nextn_mtp_fused_experts_excluded( + quant_config: Optional[QuantizationConfig], + num_hidden_layers: int, +) -> bool: + exclude_layers = getattr(quant_config, "exclude_layers", None) or [] + layer_prefix = f"model.layers.{num_hidden_layers}." + return any( + name.startswith(layer_prefix) and ".mlp.experts." in name + for name in exclude_layers + ) + + +def should_apply_glm_nextn_moe_ptpc( + quant_config: Optional[QuantizationConfig], + num_hidden_layers: int, +) -> bool: + if not enable_glm_nextn_moe_ptpc(quant_config): + return False + return glm_nextn_mtp_fused_experts_excluded(quant_config, num_hidden_layers) + + if _is_npu: from sgl_kernel_npu.norm.split_qkv_rmsnorm_rope import split_qkv_rmsnorm_rope @@ -1464,10 +1502,58 @@ class GlmMoeDsaForCausalLMNextN(DeepseekV3ForCausalLMNextN): return name.replace(layer_prefix, "model", 1) return name.replace(layer_prefix, "model.decoder", 1) + def _maybe_quant_glm_nextn_moe_to_ptpc(self, weights): + """Cast this GLM-5.2 draft layer's routed experts to per-channel FP8.""" + layer_id = self.config.num_hidden_layers + if not should_apply_glm_nextn_moe_ptpc(self.quant_config, layer_id): + return weights + + layer_prefix = f"model.layers.{layer_id}" + fp8_max = torch.finfo(torch.float8_e4m3fn).max + log_info_on_rank0( + logger, + "GLM NextN MoE PTPC: casting draft expert weights under " + f"{layer_prefix}.mlp to fp8_e4m3 per-channel", + ) + + def _cast() -> Iterable[Tuple[str, torch.Tensor]]: + for name, tensor in weights: + if not ( + name.startswith(layer_prefix + ".") + and _GLM_NEXTN_EXPERT_PROJ_RE.search(name) + ): + yield name, tensor + continue + if tensor.ndim != 2: + raise ValueError( + f"{name}: PTPC cast expects a 2D expert weight, " + f"got {tuple(tensor.shape)}" + ) + weight = tensor.to(torch.float32) + scale = weight.abs().amax(dim=-1, keepdim=True).clamp(min=1e-12) + scale /= fp8_max + yield ( + name, + (weight / scale).clamp(-fp8_max, fp8_max).to(torch.float8_e4m3fn), + ) + yield name[: -len("weight")] + "weight_scale", scale.squeeze(-1) + + return _cast() + + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + weights = self._maybe_quant_glm_nextn_moe_to_ptpc(weights) + return super().load_weights(weights) + def _resolve_nextn_quant_config(self, config, quant_config): if quant_config is None or quant_config.get_name() != "quark": return quant_config + # The caller reuses this QuarkConfig for the target model and lm_head, + # so the draft-only rewrites below need a private copy of the wrapper + # and of the dict its schemes are read from. + quant_config = copy.copy(quant_config) + quant_config.quant_config = copy.deepcopy(quant_config.quant_config) + layer_prefix = f"model.layers.{config.num_hidden_layers}" # Quark's per-module scheme selection (e.g. MTP self_attn in PTPC-FP8 @@ -1500,16 +1586,34 @@ class GlmMoeDsaForCausalLMNextN(DeepseekV3ForCausalLMNextN): names.add(self._map_mtp_ckpt_name(name, layer_prefix)) # Fused routed experts are queried by the coarse module prefix - # "model.decoder.mlp.experts". Expanded per-expert leaf excludes do not - # match that prefix, so add the coarse prefix when any routed expert in - # the MTP layer is excluded. This keeps only that fused MoE module bf16 - # while allowing the remaining draft modules to use their quant config. - if any(".mlp.experts." in name for name in mtp_excluded): + # "model.decoder.mlp.experts", which expanded per-expert leaf excludes + # do not match. So that module needs its own entry: bf16 as in the + # checkpoint, or the scheme matching the on-load PTPC-FP8 cast. + # Same gate as the weight-loader cast (Quark-excluded = bf16 in ckpt). + if should_apply_glm_nextn_moe_ptpc(quant_config, config.num_hidden_layers): + mtp_layer_quant_config = quant_config.quant_config.setdefault( + "layer_quant_config", {} + ) + mtp_layer_quant_config["model.decoder.mlp.experts"] = { + "weight": { + "dtype": "fp8_e4m3", + "is_dynamic": False, + "qscheme": "per_channel", + }, + # Dynamic per_channel is QuarkW8A8FP8MoE's per-token input. + "input_tensors": { + "dtype": "fp8_e4m3", + "is_dynamic": True, + "qscheme": "per_channel", + }, + } + logger.info( + "SGLANG_GLM_NEXTN_MOE_PTPC=1: MTP fused MoE " + "(model.decoder.mlp.experts) runs as PTPC-FP8" + ) + elif any(".mlp.experts." in name for name in mtp_excluded): names.add("model.decoder.mlp.experts") - import copy - - quant_config = copy.copy(quant_config) quant_config.exclude_layers = list(names) return quant_config diff --git a/test/registered/unit/models/test_glm_nextn_moe_ptpc.py b/test/registered/unit/models/test_glm_nextn_moe_ptpc.py new file mode 100644 index 000000000..e047cbf50 --- /dev/null +++ b/test/registered/unit/models/test_glm_nextn_moe_ptpc.py @@ -0,0 +1,160 @@ +"""CI for SGLANG_GLM_NEXTN_MOE_PTPC=1 (GLM-5.2 NextN per-channel FP8 draft MoE). + +The feature is off by default. Without a case that turns the flag on, CI never +touches the Quark scheme rewrite and cannot claim the path works. These tests +exercise that ON wiring on CPU without loading a 70B MXFP4 checkpoint. + +A full serve+generate job still needs the MXFP4 weights in the runner cache; +register that separately as nightly if the checkpoint is present. +""" + +from __future__ import annotations + +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.srt.models.glm4_moe import ( + GlmMoeDsaForCausalLMNextN, + enable_glm_nextn_moe_ptpc, + should_apply_glm_nextn_moe_ptpc, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +LAYER = 78 +PREFIX = f"model.layers.{LAYER}" +EXPERT_LEAF = f"{PREFIX}.mlp.experts.0.w1" +ATTN_LEAF = f"{PREFIX}.self_attn.q_proj" +EXPERT_WEIGHT = f"{PREFIX}.mlp.experts.0.gate_proj.weight" +_PTPC_ENV = "sglang.srt.models.glm4_moe.envs.SGLANG_GLM_NEXTN_MOE_PTPC.get" + + +def _quark_cfg(*, exclude=None, layer_quant=None): + return SimpleNamespace( + get_name=lambda: "quark", + quant_config={"layer_quant_config": dict(layer_quant or {}), "exclude": []}, + exclude_layers=list( + exclude if exclude is not None else [EXPERT_LEAF, ATTN_LEAF] + ), + ) + + +class TestEnableGlmNextnMoePtpc(CustomTestCase): + def test_off_by_default(self): + self.assertFalse(enable_glm_nextn_moe_ptpc(_quark_cfg())) + + def test_on_requires_quark(self): + with patch( + _PTPC_ENV, + return_value=True, + ): + self.assertTrue(enable_glm_nextn_moe_ptpc(_quark_cfg())) + self.assertFalse( + enable_glm_nextn_moe_ptpc(SimpleNamespace(get_name=lambda: "fp8")) + ) + self.assertFalse(enable_glm_nextn_moe_ptpc(None)) + + def test_apply_requires_excluded_mtp_experts(self): + with patch( + _PTPC_ENV, + return_value=True, + ): + self.assertTrue(should_apply_glm_nextn_moe_ptpc(_quark_cfg(), LAYER)) + self.assertFalse( + should_apply_glm_nextn_moe_ptpc(_quark_cfg(exclude=[ATTN_LEAF]), LAYER) + ) + + +class TestResolveNextnQuantConfigPtpcOn(CustomTestCase): + def _resolve(self, cfg, flag: bool): + model = GlmMoeDsaForCausalLMNextN.__new__(GlmMoeDsaForCausalLMNextN) + hf = SimpleNamespace(num_hidden_layers=LAYER) + with patch( + _PTPC_ENV, + return_value=flag, + ): + return model._resolve_nextn_quant_config(hf, cfg) + + def test_flag_off_excludes_fused_experts(self): + src = _quark_cfg() + out = self._resolve(src, flag=False) + self.assertIn("model.decoder.mlp.experts", out.exclude_layers) + self.assertNotIn( + "model.decoder.mlp.experts", + out.quant_config.get("layer_quant_config", {}), + ) + + def test_flag_on_assigns_ptpc_scheme_instead_of_bf16_exclude(self): + src = _quark_cfg() + out = self._resolve(src, flag=True) + self.assertNotIn("model.decoder.mlp.experts", out.exclude_layers) + scheme = out.quant_config["layer_quant_config"]["model.decoder.mlp.experts"] + self.assertEqual(scheme["weight"]["dtype"], "fp8_e4m3") + self.assertEqual(scheme["weight"]["qscheme"], "per_channel") + self.assertFalse(scheme["weight"]["is_dynamic"]) + self.assertEqual(scheme["input_tensors"]["dtype"], "fp8_e4m3") + self.assertEqual(scheme["input_tensors"]["qscheme"], "per_channel") + self.assertTrue(scheme["input_tensors"]["is_dynamic"]) + + def test_flag_on_does_not_mutate_caller_config(self): + src = _quark_cfg() + orig_exclude = list(src.exclude_layers) + orig_layer = dict(src.quant_config.get("layer_quant_config") or {}) + self._resolve(src, flag=True) + self.assertEqual(src.exclude_layers, orig_exclude) + self.assertEqual(src.quant_config.get("layer_quant_config") or {}, orig_layer) + + def test_flag_on_skips_ptpc_when_experts_not_excluded(self): + src = _quark_cfg(exclude=[ATTN_LEAF]) + out = self._resolve(src, flag=True) + self.assertNotIn("model.decoder.mlp.experts", out.exclude_layers) + self.assertNotIn( + "model.decoder.mlp.experts", + out.quant_config.get("layer_quant_config", {}), + ) + + def test_flag_on_skips_ptpc_when_mtp_not_in_exclude(self): + src = _quark_cfg(exclude=[]) + out = self._resolve(src, flag=True) + self.assertNotIn( + "model.decoder.mlp.experts", + out.quant_config.get("layer_quant_config", {}), + ) + + +class TestMaybeQuantGlmNextnMoeToPtpc(CustomTestCase): + def _cast(self, cfg, flag: bool): + loader = GlmMoeDsaForCausalLMNextN.__new__(GlmMoeDsaForCausalLMNextN) + loader.quant_config = cfg + loader.config = SimpleNamespace(num_hidden_layers=LAYER) + weights = [(EXPERT_WEIGHT, torch.ones(4, 8, dtype=torch.bfloat16))] + with patch( + _PTPC_ENV, + return_value=flag, + ): + return list(loader._maybe_quant_glm_nextn_moe_to_ptpc(weights)) + + def test_flag_on_casts_excluded_bf16_experts(self): + out = self._cast(_quark_cfg(), flag=True) + names = [name for name, _ in out] + self.assertIn(EXPERT_WEIGHT, names) + self.assertIn(EXPERT_WEIGHT[: -len("weight")] + "weight_scale", names) + weight = dict(out)[EXPERT_WEIGHT] + self.assertEqual(weight.dtype, torch.float8_e4m3fn) + + def test_flag_on_does_not_cast_when_experts_not_excluded(self): + src = _quark_cfg(exclude=[ATTN_LEAF]) + out = self._cast(src, flag=True) + self.assertEqual(len(out), 1) + name, tensor = out[0] + self.assertEqual(name, EXPERT_WEIGHT) + self.assertEqual(tensor.dtype, torch.bfloat16) + + +if __name__ == "__main__": + unittest.main()