Fix stale GLM MoE routing after runtime weight updates (#35883)
Co-authored-by: Jiajun Li <jiajun.li@radixark.ai>
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user