From 8a87079dbbf0f5b1543ec25d914dfd988eba42de Mon Sep 17 00:00:00 2001 From: Yuzhen Zhou <82826991+zyzshishui@users.noreply.github.com> Date: Sun, 30 Aug 2026 14:13:34 -0700 Subject: [PATCH] Fix stale GLM MoE routing after runtime weight updates (#35883) Co-authored-by: Jiajun Li --- python/sglang/srt/models/glm4_moe.py | 11 ++-- python/sglang/srt/models/glm4_moe_lite.py | 11 ++-- python/sglang/srt/utils/weight_checker.py | 1 - test/registered/rl/test_weight_checker_e2e.py | 1 - .../unit/models/test_glm_moe_gate_fp32.py | 54 +++++++++++++++++++ .../unit/utils/test_weight_checker.py | 22 +++----- 6 files changed, 68 insertions(+), 32 deletions(-) create mode 100644 test/registered/unit/models/test_glm_moe_gate_fp32.py diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py index a4aa494c7..7f7a6392a 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -368,19 +368,16 @@ class Glm4MoeGate(nn.Module): ): super().__init__() self.weight = nn.Parameter( - torch.empty((config.n_routed_experts, config.hidden_size)) + torch.empty( + (config.n_routed_experts, config.hidden_size), dtype=torch.float32 + ) ) self.e_score_correction_bias = nn.Parameter( torch.empty((config.n_routed_experts), dtype=torch.float32) ) - # GLM requires FP32 gate projection; cache to avoid per-forward cast. - # FIXME: if gate weight is updated at runtime (e.g. expert rebalancing), _weight_fp32 must be invalidated. - self.register_buffer("_weight_fp32", None, persistent=False) def forward(self, hidden_states): - if self._weight_fp32 is None: - self._weight_fp32 = self.weight.data.to(torch.float32) - logits = F.linear(hidden_states.to(torch.float32), self._weight_fp32, None) + logits = F.linear(hidden_states.to(torch.float32), self.weight, None) return logits diff --git a/python/sglang/srt/models/glm4_moe_lite.py b/python/sglang/srt/models/glm4_moe_lite.py index 0388d2180..4d74c41dc 100644 --- a/python/sglang/srt/models/glm4_moe_lite.py +++ b/python/sglang/srt/models/glm4_moe_lite.py @@ -156,19 +156,16 @@ class Glm4MoeLiteGate(nn.Module): super().__init__() self.is_nextn = is_nextn self.weight = nn.Parameter( - torch.empty((config.n_routed_experts, config.hidden_size)) + torch.empty( + (config.n_routed_experts, config.hidden_size), dtype=torch.float32 + ) ) self.e_score_correction_bias = nn.Parameter( torch.empty((config.n_routed_experts), dtype=torch.float32) ) - # GLM requires FP32 gate projection; cache to avoid per-forward cast. - # FIXME: if gate weight is updated at runtime (e.g. expert rebalancing), _weight_fp32 must be invalidated. - self.register_buffer("_weight_fp32", None, persistent=False) def forward(self, hidden_states): - if self._weight_fp32 is None: - self._weight_fp32 = self.weight.data.to(torch.float32) - logits = F.linear(hidden_states.to(torch.float32), self._weight_fp32, None) + logits = F.linear(hidden_states.to(torch.float32), self.weight, None) return logits diff --git a/python/sglang/srt/utils/weight_checker.py b/python/sglang/srt/utils/weight_checker.py index bf0f6bb1e..bb6c3377f 100644 --- a/python/sglang/srt/utils/weight_checker.py +++ b/python/sglang/srt/utils/weight_checker.py @@ -55,7 +55,6 @@ _NON_PERSISTENT_BUFFER_PATTERNS = ( "cos_sin_cache", "inv_freq", "freqs_cis", - "_weight_fp32", "expert_mask_gpu", ) diff --git a/test/registered/rl/test_weight_checker_e2e.py b/test/registered/rl/test_weight_checker_e2e.py index 6a0b40030..bf1de7364 100644 --- a/test/registered/rl/test_weight_checker_e2e.py +++ b/test/registered/rl/test_weight_checker_e2e.py @@ -211,7 +211,6 @@ class TestWeightCheckerE2E(CustomTestCase): self.assertNotIn("cos_sin_cache", name) self.assertNotIn("inv_freq", name) self.assertNotIn("freqs_cis", name) - self.assertNotIn("_weight_fp32", name) def test_z_snapshot_reset_compare_detects_diff(self): """Destructive: leaves weights randomized. Named test_z_* so it runs last.""" diff --git a/test/registered/unit/models/test_glm_moe_gate_fp32.py b/test/registered/unit/models/test_glm_moe_gate_fp32.py new file mode 100644 index 000000000..6489aeca6 --- /dev/null +++ b/test/registered/unit/models/test_glm_moe_gate_fp32.py @@ -0,0 +1,54 @@ +"""Regression tests for GLM MoE gate weights used by FP32 routing.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=4, suite="base-a-test-cpu") + +import unittest +from types import SimpleNamespace + +import torch +import torch.nn.functional as F + +from sglang.srt.model_executor.model_runner_components.weight_updater import ( + _model_load_weights_direct, +) +from sglang.srt.model_loader.utils import set_default_torch_dtype +from sglang.srt.model_loader.weight_utils import default_weight_loader +from sglang.srt.models.glm4_moe import Glm4MoeGate +from sglang.srt.models.glm4_moe_lite import Glm4MoeLiteGate +from sglang.test.test_utils import CustomTestCase + +_CONFIG = SimpleNamespace(n_routed_experts=3, hidden_size=4) + + +class TestGlmMoeGateFp32Weight(CustomTestCase): + def test_bf16_load_updates_fp32_weight_in_place(self): + """BF16 runtime updates must overwrite the canonical FP32 gate weight.""" + hidden_states = torch.arange(8, dtype=torch.bfloat16).reshape(2, 4) + initial = torch.arange(12, dtype=torch.bfloat16).reshape(3, 4) + updated = initial + 1 + + for gate_cls in (Glm4MoeGate, Glm4MoeLiteGate): + with self.subTest(gate=gate_cls.__name__): + with set_default_torch_dtype(torch.bfloat16): + gate = gate_cls(_CONFIG) + + self.assertEqual(gate.weight.dtype, torch.float32) + self.assertFalse(hasattr(gate, "_weight_fp32")) + weight_ptr = gate.weight.data_ptr() + + default_weight_loader(gate.weight, initial) + torch.testing.assert_close(gate.weight, initial.float()) + + _model_load_weights_direct(gate, [("weight", updated)]) + self.assertEqual(gate.weight.data_ptr(), weight_ptr) + torch.testing.assert_close(gate.weight, updated.float()) + torch.testing.assert_close( + gate(hidden_states), + F.linear(hidden_states.float(), updated.float()), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/utils/test_weight_checker.py b/test/registered/unit/utils/test_weight_checker.py index 79fe66aab..e9d893a9d 100644 --- a/test/registered/unit/utils/test_weight_checker.py +++ b/test/registered/unit/utils/test_weight_checker.py @@ -111,7 +111,7 @@ class _TinyModel(nn.Module): self.w = nn.Parameter(torch.randn(4, 4), requires_grad=False) self.b = nn.Parameter(torch.zeros(4), requires_grad=False) self.register_buffer("running_mean", torch.zeros(4)) - # Buffer names that match weight_checker's hard-coded skip patterns. + # Buffer names used to exercise weight checker's hard-coded filters. self.register_buffer("rotary_emb_cos_sin_cache", torch.full((8,), 3.14)) self.register_buffer("rotary_emb_freqs_cis", torch.full((8,), 2.71)) self.register_buffer("gate_proj_weight_fp32_cache", torch.full((8,), 1.41)) @@ -252,13 +252,6 @@ class TestPostprocessTensors(CustomTestCase): [("model.rotary_emb.inv_freq", False, RawComparable(t))], ) - def test_skips_weight_fp32_substring(self): - t = torch.randn(4) - _assert_entries_close( - _build_check_entries({"model.layers.0.mlp.gate._weight_fp32": t}, set()), - [("model.layers.0.mlp.gate._weight_fp32", False, RawComparable(t))], - ) - def test_substring_match_not_endswith(self): # Pattern can appear anywhere in the name, not just at the end. t = torch.randn(4) @@ -545,10 +538,12 @@ class TestResetTensors(_WeightCheckerTestBase): self.checker._reset_tensors() torch.testing.assert_close(self.model.rotary_emb_freqs_cis, before) - def test_skips_weight_fp32(self): + def test_poisons_weight_fp32_cache(self): before = self.model.gate_proj_weight_fp32_cache.clone() + before_ptr = self.model.gate_proj_weight_fp32_cache.data_ptr() self.checker._reset_tensors() - torch.testing.assert_close(self.model.gate_proj_weight_fp32_cache, before) + self.assertEqual(self.model.gate_proj_weight_fp32_cache.data_ptr(), before_ptr) + self.assertFalse(torch.equal(self.model.gate_proj_weight_fp32_cache, before)) class TestCompare(_WeightCheckerTestBase): @@ -645,11 +640,6 @@ class TestIsNonPersistentBufferName(CustomTestCase): def test_matches_freqs_cis_substring(self): self.assertTrue(_is_non_persistent_buffer_name("model.rotary_emb.freqs_cis")) - def test_matches_weight_fp32_substring(self): - self.assertTrue( - _is_non_persistent_buffer_name("model.layers.0.mlp.gate._weight_fp32") - ) - def test_does_not_match_normal_param_names(self): self.assertFalse(_is_non_persistent_buffer_name("model.layers.0.mlp.weight")) self.assertFalse(_is_non_persistent_buffer_name("model.embed_tokens.weight")) @@ -723,10 +713,10 @@ class TestComputeChecksum(_ChecksumTestBase): self.assertIn("w", names) self.assertIn("b", names) self.assertIn("running_mean", names) + self.assertIn("gate_proj_weight_fp32_cache", names) # Non-persistent buffer patterns are filtered out. self.assertNotIn("rotary_emb_cos_sin_cache", names) self.assertNotIn("rotary_emb_freqs_cis", names) - self.assertNotIn("gate_proj_weight_fp32_cache", names) def test_hashes_are_hex_strings(self): out = self.checker._compute_checksum()