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