Fix stale GLM MoE routing after runtime weight updates (#35883)

Co-authored-by: Jiajun Li <jiajun.li@radixark.ai>
This commit is contained in:
Yuzhen Zhou
2026-08-30 14:13:34 -07:00
committed by GitHub
co-authored by Jiajun Li
parent 5ab97c4f44
commit 8a87079dbb
6 changed files with 68 additions and 32 deletions
+4 -7
View File
@@ -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
+4 -7
View File
@@ -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
@@ -55,7 +55,6 @@ _NON_PERSISTENT_BUFFER_PATTERNS = (
"cos_sin_cache",
"inv_freq",
"freqs_cis",
"_weight_fp32",
"expert_mask_gpu",
)
@@ -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."""
@@ -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()
@@ -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()