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__()
|
super().__init__()
|
||||||
self.weight = nn.Parameter(
|
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(
|
self.e_score_correction_bias = nn.Parameter(
|
||||||
torch.empty((config.n_routed_experts), dtype=torch.float32)
|
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):
|
def forward(self, hidden_states):
|
||||||
if self._weight_fp32 is None:
|
logits = F.linear(hidden_states.to(torch.float32), self.weight, None)
|
||||||
self._weight_fp32 = self.weight.data.to(torch.float32)
|
|
||||||
logits = F.linear(hidden_states.to(torch.float32), self._weight_fp32, None)
|
|
||||||
return logits
|
return logits
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -156,19 +156,16 @@ class Glm4MoeLiteGate(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.is_nextn = is_nextn
|
self.is_nextn = is_nextn
|
||||||
self.weight = nn.Parameter(
|
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(
|
self.e_score_correction_bias = nn.Parameter(
|
||||||
torch.empty((config.n_routed_experts), dtype=torch.float32)
|
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):
|
def forward(self, hidden_states):
|
||||||
if self._weight_fp32 is None:
|
logits = F.linear(hidden_states.to(torch.float32), self.weight, None)
|
||||||
self._weight_fp32 = self.weight.data.to(torch.float32)
|
|
||||||
logits = F.linear(hidden_states.to(torch.float32), self._weight_fp32, None)
|
|
||||||
return logits
|
return logits
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -55,7 +55,6 @@ _NON_PERSISTENT_BUFFER_PATTERNS = (
|
|||||||
"cos_sin_cache",
|
"cos_sin_cache",
|
||||||
"inv_freq",
|
"inv_freq",
|
||||||
"freqs_cis",
|
"freqs_cis",
|
||||||
"_weight_fp32",
|
|
||||||
"expert_mask_gpu",
|
"expert_mask_gpu",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -211,7 +211,6 @@ class TestWeightCheckerE2E(CustomTestCase):
|
|||||||
self.assertNotIn("cos_sin_cache", name)
|
self.assertNotIn("cos_sin_cache", name)
|
||||||
self.assertNotIn("inv_freq", name)
|
self.assertNotIn("inv_freq", name)
|
||||||
self.assertNotIn("freqs_cis", name)
|
self.assertNotIn("freqs_cis", name)
|
||||||
self.assertNotIn("_weight_fp32", name)
|
|
||||||
|
|
||||||
def test_z_snapshot_reset_compare_detects_diff(self):
|
def test_z_snapshot_reset_compare_detects_diff(self):
|
||||||
"""Destructive: leaves weights randomized. Named test_z_* so it runs last."""
|
"""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.w = nn.Parameter(torch.randn(4, 4), requires_grad=False)
|
||||||
self.b = nn.Parameter(torch.zeros(4), requires_grad=False)
|
self.b = nn.Parameter(torch.zeros(4), requires_grad=False)
|
||||||
self.register_buffer("running_mean", torch.zeros(4))
|
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_cos_sin_cache", torch.full((8,), 3.14))
|
||||||
self.register_buffer("rotary_emb_freqs_cis", torch.full((8,), 2.71))
|
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))
|
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))],
|
[("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):
|
def test_substring_match_not_endswith(self):
|
||||||
# Pattern can appear anywhere in the name, not just at the end.
|
# Pattern can appear anywhere in the name, not just at the end.
|
||||||
t = torch.randn(4)
|
t = torch.randn(4)
|
||||||
@@ -545,10 +538,12 @@ class TestResetTensors(_WeightCheckerTestBase):
|
|||||||
self.checker._reset_tensors()
|
self.checker._reset_tensors()
|
||||||
torch.testing.assert_close(self.model.rotary_emb_freqs_cis, before)
|
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 = self.model.gate_proj_weight_fp32_cache.clone()
|
||||||
|
before_ptr = self.model.gate_proj_weight_fp32_cache.data_ptr()
|
||||||
self.checker._reset_tensors()
|
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):
|
class TestCompare(_WeightCheckerTestBase):
|
||||||
@@ -645,11 +640,6 @@ class TestIsNonPersistentBufferName(CustomTestCase):
|
|||||||
def test_matches_freqs_cis_substring(self):
|
def test_matches_freqs_cis_substring(self):
|
||||||
self.assertTrue(_is_non_persistent_buffer_name("model.rotary_emb.freqs_cis"))
|
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):
|
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.layers.0.mlp.weight"))
|
||||||
self.assertFalse(_is_non_persistent_buffer_name("model.embed_tokens.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("w", names)
|
||||||
self.assertIn("b", names)
|
self.assertIn("b", names)
|
||||||
self.assertIn("running_mean", names)
|
self.assertIn("running_mean", names)
|
||||||
|
self.assertIn("gate_proj_weight_fp32_cache", names)
|
||||||
# Non-persistent buffer patterns are filtered out.
|
# Non-persistent buffer patterns are filtered out.
|
||||||
self.assertNotIn("rotary_emb_cos_sin_cache", names)
|
self.assertNotIn("rotary_emb_cos_sin_cache", names)
|
||||||
self.assertNotIn("rotary_emb_freqs_cis", names)
|
self.assertNotIn("rotary_emb_freqs_cis", names)
|
||||||
self.assertNotIn("gate_proj_weight_fp32_cache", names)
|
|
||||||
|
|
||||||
def test_hashes_are_hex_strings(self):
|
def test_hashes_are_hex_strings(self):
|
||||||
out = self.checker._compute_checksum()
|
out = self.checker._compute_checksum()
|
||||||
|
|||||||
Reference in New Issue
Block a user