diff --git a/python/sglang/srt/model_executor/model_runner_components/weight_updater.py b/python/sglang/srt/model_executor/model_runner_components/weight_updater.py index 31c7518df3db..37c38ad5b825 100644 --- a/python/sglang/srt/model_executor/model_runner_components/weight_updater.py +++ b/python/sglang/srt/model_executor/model_runner_components/weight_updater.py @@ -55,6 +55,13 @@ def _unsupported_derived_weight_cache_error( "online weight updates." ) + if model is not None: + # Model-owned caches can publish the same rank-uniform constraint + # without importing individual model implementations in the updater. + for module in model.modules(): + reason = getattr(module, "_derived_weight_cache_error", None) + if reason is not None: + return reason from sglang.kernels.ops.attention.dsv4.gemm import hpc_bf16xfp32_gemm_enabled if hpc_bf16xfp32_gemm_enabled(): diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py index 8c48dd68db1a..39656d722dfc 100644 --- a/python/sglang/srt/models/qwen3_5.py +++ b/python/sglang/srt/models/qwen3_5.py @@ -133,8 +133,12 @@ _is_gfx95 = is_gfx95_supported() _is_hip = is_hip() _QWEN3_5_MOE_TEXT_MODEL_TYPES = ("qwen3_5_moe_text", "qwen4_exp_text") +# qwen4_exp shares these classes, but the ROCm packed path is Qwen3.5-only. +_QWEN3_5_ROCM_PACKED_MODEL_TYPES = ("qwen3_5_text", "qwen3_5_moe_text") _is_xpu = is_xpu() _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip +if _use_aiter: + from aiter.tuned_gemm import tgemm _hip_use_alt_stream = get_bool_env_var("SGLANG_ALT_STREAM") and _is_hip _gdn_use_alt_stream = _is_cuda or ( get_bool_env_var("SGLANG_GDN_QKVZ_BA_ALT_STREAM", "False") and _hip_use_alt_stream @@ -394,6 +398,10 @@ def __init__( self._bind_packed_weight_loaders(self.in_proj_ba) self._fused_in_proj_weight: Optional[torch.Tensor] = None self._fused_in_proj_qkvz_width = 0 + self._fused_in_proj_ba_width = 0 + self._fused_in_proj_scale: Optional[torch.Tensor] = None + self._fused_in_proj_sources = None + self._derived_weight_cache_error = None self._fused_input_proj_cpu_enabled = LazyValue( lambda: ( _is_cpu @@ -653,16 +661,27 @@ def fix_query_key_value_ordering( return query, key, value, z, b, a def finalize_fused_in_proj(self) -> None: - """Stack in_proj_qkvz + in_proj_ba into one GEMM weight; - the module weights become row views of it, - so weight reload and dtype checks still see them.""" - if not _is_cuda or self._fused_in_proj_weight is not None: + """Prepare one BF16 or FP8 GEMM for both input projections. + + BF16 parameters alias the packed rows. FP8 keeps the original layouts + for the separate path and publishes its derived-cache update constraint. + """ + if not (_is_cuda or _use_aiter) or self._fused_in_proj_weight is not None: + return + if ( + _use_aiter + and self.config.model_type not in _QWEN3_5_ROCM_PACKED_MODEL_TYPES + ): return if get_lora().enable_lora or get_lora().lora_paths: # LoRA wraps the individual Linear modules; the fused GEMM would # bypass their adapters. return + if _use_aiter and not get_bool_env_var("SGLANG_QWEN35_PACKED_IN_PROJ", "True"): + return qkvz, ba = self.in_proj_qkvz, self.in_proj_ba + if _use_aiter and self._finalize_fused_fp8_in_proj(): + return if not ( isinstance(qkvz.quant_method, UnquantizedLinearMethod) and isinstance(ba.quant_method, UnquantizedLinearMethod) @@ -678,7 +697,124 @@ def finalize_fused_in_proj(self) -> None: ba.weight.data = fused[self._fused_in_proj_qkvz_width :] self._fused_in_proj_weight = fused + def _finalize_fused_fp8_in_proj(self) -> bool: + """Pack loaded Quark per-channel FP8 projections without requantizing.""" + from aiter.ops.shuffle import shuffle_weight + + from sglang.srt.layers.quantization.fp8_utils import use_aiter_bpreshuffle_gemm + from sglang.srt.layers.quantization.quark.schemes.quark_w8a8_fp8 import ( + QuarkW8A8Fp8, + ) + + projections = (self.in_proj_qkvz, self.in_proj_ba) + if not all( + type(getattr(proj, "scheme", None)) is QuarkW8A8Fp8 + and proj.scheme.weight_qscheme == "per_channel" + and proj.scheme.per_token + and proj.input_scale is None + and proj.bias is None + and proj.weight.dtype == torch.float8_e4m3fn + and proj.weight.is_cuda + and proj.weight.t().is_contiguous() + and proj.weight.shape[1] % 16 == 0 + and proj.weight_scale.numel() == proj.weight.shape[1] + for proj in projections + ): + return False + # Keep the derived-cache update constraint uniform across ranks. With + # PP, an attention-only stage may have no GDN cache to publish it from. + if get_parallel().pp_size != 1: + return False + widths = [proj.weight.shape[1] for proj in projections] + if projections[0].weight.shape[0] != projections[1].weight.shape[0]: + return False + # CK's bpreshuffle GEMM requires N % 64 == 0. TP4 has BA=32, + # so pad only the packed result, preserving the original Linear sizes. + padded_width = (sum(widths) + 63) // 64 * 64 + weight = torch.zeros( + (padded_width, projections[0].weight.shape[0]), + dtype=projections[0].weight.dtype, + device=projections[0].weight.device, + ) + scale = torch.ones((padded_width, 1), dtype=torch.float32, device=weight.device) + offset = 0 + for proj, width in zip(projections, widths): + source = proj.weight.t() + if not use_aiter_bpreshuffle_gemm(width): + source = shuffle_weight(source, (16, 16)) + # AITER's shuffle is local to 16 output rows; concatenating aligned + # shuffled blocks is identical to shuffling their concatenation. + weight[offset : offset + width].copy_(source) + scale[offset : offset + width].copy_(proj.weight_scale.view(-1, 1)) + offset += width + self._fused_in_proj_weight = weight + self._fused_in_proj_scale = scale + self._fused_in_proj_qkvz_width, self._fused_in_proj_ba_width = widths + self._fused_in_proj_sources = tuple( + (value, value.data_ptr(), None if value.is_inference() else value._version) + for proj in projections + for value in (proj.weight, proj.weight_scale) + ) + self._derived_weight_cache_error = ( + "Online weight updates are not supported with packed FP8 GDN input " + "projections: captured graphs retain a derived weight/scale buffer. " + "Restart with SGLANG_QWEN35_PACKED_IN_PROJ=0 before updating weights." + ) + return True + + def _fused_fp8_in_proj_sources_valid(self) -> bool: + current = tuple( + value + for proj in (self.in_proj_qkvz, self.in_proj_ba) + for value in (proj.weight, proj.weight_scale) + ) + return all( + value is source + and value.data_ptr() == pointer + and (version is None or value._version == version) + for value, (source, pointer, version) in zip( + current, self._fused_in_proj_sources + ) + ) + def _forward_input_proj(self, hidden_states: torch.Tensor): + if _use_aiter and self._fused_in_proj_weight is not None: + # Unquantized BF16 projections consume the bf16 side of the fused + # AR+RMSNorm tuple; one aiter GEMM replaces the two separate + # projections. Measured on MI355X the packed GEMM only pays off for + # decode/verify-sized batches, so larger batches keep the separate + # projections below. + x = hidden_states[0] if isinstance(hidden_states, tuple) else hidden_states + if ( + x.dtype == torch.bfloat16 + and 0 < x.shape[0] <= 64 + and not torch.compiler.is_compiling() + ): + if self._fused_in_proj_scale is not None: + if self._fused_fp8_in_proj_sources_valid(): + from sglang.srt.layers.quantization.fp8_utils import ( + apply_fp8_linear, + ) + + fused_out = apply_fp8_linear( + x, + self._fused_in_proj_weight.t(), + self._fused_in_proj_scale, + use_per_token_if_dynamic=True, + ) + split = self._fused_in_proj_qkvz_width + return ( + fused_out[:, :split], + fused_out[:, split : split + self._fused_in_proj_ba_width], + ) + else: + fused_out = tgemm.mm( + x, self._fused_in_proj_weight, None, otype=x.dtype + ) + return ( + fused_out[:, : self._fused_in_proj_qkvz_width], + fused_out[:, self._fused_in_proj_qkvz_width :], + ) # AMD/aiter fused AR+RMSNorm+per-group-quant path ships a # ``(bf16, fp8, scale)`` 3-tuple so the FP8 ``in_proj_qkvz`` can # consume ``(fp8, scale)`` (skipping its internal quant) while the @@ -688,7 +824,8 @@ def _forward_input_proj(self, hidden_states: torch.Tensor): return self._forward_input_proj_fused_quant_amd(hidden_states) if ( - self._fused_in_proj_weight is not None + not _use_aiter + and self._fused_in_proj_weight is not None and hidden_states.dtype == torch.bfloat16 # Measured on cuBLAS above ~1k rows: # the merged (m, 4120) GEMM is ~10% slower than the two separate GEMMs. @@ -1724,6 +1861,15 @@ def get_input_embeddings(self): return self.embed_tokens def prepare_before_cuda_graph_capture(self, model_runner) -> None: + if _use_aiter and self.config.model_type in _QWEN3_5_ROCM_PACKED_MODEL_TYPES: + packed = 0 + for module in self.modules(): + if isinstance(module, Qwen3_5GatedDeltaNet): + module.finalize_fused_in_proj() + packed += int(module._fused_in_proj_weight is not None) + logger.info( + "Packed BF16/FP8 GDN input projection enabled for %d layers", packed + ) if self.flashinfer_mnnvl_cutedsl_fusion is None: return from sglang.srt.layers.moe.qwen35_flashinfer_fusion import ( diff --git a/test/registered/amd/test_qwen35_gdn_packed_fp8_in_proj.py b/test/registered/amd/test_qwen35_gdn_packed_fp8_in_proj.py new file mode 100644 index 000000000000..45f4d8e86ed9 --- /dev/null +++ b/test/registered/amd/test_qwen35_gdn_packed_fp8_in_proj.py @@ -0,0 +1,142 @@ +"""Quark FP8 GDN packing: actual TP widths, scales, padding and dispatch.""" + +import os +import types +import unittest +from unittest.mock import patch + +import torch + +from sglang.test.ci.ci_register import register_amd_ci +from sglang.test.test_utils import CustomTestCase + +register_amd_ci(est_time=90, suite="stage-b-test-1-gpu-small-amd-mi35x") + + +class _FP8Linear(torch.nn.Module): + def __init__(self, rows): + super().__init__() + from sglang.srt.layers.quantization.quark.schemes.quark_w8a8_fp8 import ( + QuarkW8A8Fp8, + ) + + self.weight = torch.nn.Parameter( + torch.randn(rows, 4096, device="cuda").to(torch.float8_e4m3fn), + requires_grad=False, + ) + self.weight_scale = torch.nn.Parameter( + torch.linspace(0.001, 0.1, rows, device="cuda", dtype=torch.float32), + requires_grad=False, + ) + self.scheme = QuarkW8A8Fp8( + {"qscheme": "per_channel"}, + {"qscheme": "per_channel", "is_dynamic": True}, + ) + self.scheme.process_weights_after_loading(self) + self.bias = None + self.quant_method = object() + + def forward(self, x): + return self.scheme.apply_weights(self, x), None + + +@unittest.skipUnless(torch.cuda.is_available() and torch.version.hip, "ROCm FP8 GEMM") +class TestQwen35GDNPackedFP8InProj(CustomTestCase): + @classmethod + def setUpClass(cls): + os.environ.setdefault("SGLANG_USE_AITER", "1") + from sglang.srt.models import qwen3_5 + + if not qwen3_5._use_aiter: + raise unittest.SkipTest("SGLANG_USE_AITER was disabled at import") + cls.module = qwen3_5 + cls.gdn_cls = qwen3_5.Qwen3_5GatedDeltaNet + + def make_gdn(self, tp, pp_size=1, model_type="qwen3_5_moe_text"): + gdn = types.SimpleNamespace( + config=types.SimpleNamespace(model_type=model_type), + in_proj_qkvz=_FP8Linear(20480 // tp), + in_proj_ba=_FP8Linear(128 // tp), + _fused_in_proj_weight=None, + _fused_in_proj_scale=None, + _fused_in_proj_sources=None, + alt_stream=None, + _fused_input_proj_cpu_enabled=types.SimpleNamespace(value=False), + ) + for name in ( + "_finalize_fused_fp8_in_proj", + "_fused_fp8_in_proj_sources_valid", + "_forward_input_proj_fused_quant_amd", + ): + setattr(gdn, name, types.MethodType(getattr(self.gdn_cls, name), gdn)) + with ( + patch.object( + self.module, + "get_parallel", + return_value=types.SimpleNamespace(pp_size=pp_size), + ), + patch.object( + self.module, + "get_lora", + return_value=types.SimpleNamespace(enable_lora=False, lora_paths=None), + ), + patch.dict(os.environ, {"SGLANG_QWEN35_PACKED_IN_PROJ": "1"}), + ): + self.gdn_cls.finalize_fused_in_proj(gdn) + return gdn + + def test_fp8_dispatch_matches_separate_projections(self): + torch.manual_seed(42) + for tp in (2, 4): + gdn = self.make_gdn(tp) + self.assertIsNotNone(gdn._fused_in_proj_scale) + self.assertEqual(gdn._fused_in_proj_weight.shape[0] % 64, 0) + self.assertTrue(gdn._fused_fp8_in_proj_sources_valid()) + for m in (1, 4, 33, 64, 65, 300): + with self.subTest(tp=tp, tokens=m): + x = torch.randn(m, 4096, device="cuda", dtype=torch.bfloat16) + expected = (gdn.in_proj_qkvz(x)[0], gdn.in_proj_ba(x)[0]) + for hidden in (x, (x, None, None)): + got = self.gdn_cls._forward_input_proj(gdn, hidden) + for actual, reference in zip(got, expected): + self.assertEqual(actual.shape, reference.shape) + error = (actual.float() - reference.float()).norm() + self.assertLess( + (error / reference.float().norm()).item(), 0.008 + ) + if m > 64: + self.assertTrue( + all(torch.equal(a, b) for a, b in zip(got, expected)) + ) + + def test_replaced_parameters_disable_cached_projection(self): + gdn = self.make_gdn(4) + gdn.in_proj_ba.weight = torch.nn.Parameter( + gdn.in_proj_ba.weight.clone(), requires_grad=False + ) + self.assertFalse(gdn._fused_fp8_in_proj_sources_valid()) + x = torch.randn(4, 4096, device="cuda", dtype=torch.bfloat16) + got = self.gdn_cls._forward_input_proj(gdn, x) + expected = (gdn.in_proj_qkvz(x)[0], gdn.in_proj_ba(x)[0]) + self.assertTrue(all(torch.equal(a, b) for a, b in zip(got, expected))) + + def test_pipeline_parallel_keeps_original_projections(self): + gdn = self.make_gdn(4, pp_size=2) + self.assertIsNone(gdn._fused_in_proj_weight) + + def test_fp8_packing_is_limited_to_qwen35(self): + for model_type, supported in ( + ("qwen3_5_text", True), + ("qwen3_5_moe_text", True), + ("qwen4_exp_text", False), + ("other_text_model", False), + ): + with self.subTest(model_type=model_type): + gdn = self.make_gdn(4, model_type=model_type) + self.assertEqual(gdn._fused_in_proj_weight is not None, supported) + self.assertEqual(gdn._fused_in_proj_scale is not None, supported) + self.assertEqual(hasattr(gdn, "_derived_weight_cache_error"), supported) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/amd/test_qwen35_gdn_packed_in_proj.py b/test/registered/amd/test_qwen35_gdn_packed_in_proj.py new file mode 100644 index 000000000000..2037b2901910 --- /dev/null +++ b/test/registered/amd/test_qwen35_gdn_packed_in_proj.py @@ -0,0 +1,214 @@ +"""Packed BF16 GDN input projection for Qwen3.5 decode on ROCm. + +``finalize_fused_in_proj`` stacks ``in_proj_qkvz`` and ``in_proj_ba`` into one +weight and re-points the module weights at row views of it; on ROCm +``_forward_input_proj`` then runs one aiter GEMM for verify-sized batches and +splits the result. Guards: the views alias the packed buffer and keep their +values, and the packed and separate paths agree on both sides of the token gate. +""" + +import os +import types +import unittest +from unittest.mock import Mock, patch + +import torch + +from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod +from sglang.test.ci.ci_register import register_amd_ci +from sglang.test.test_utils import CustomTestCase + +register_amd_ci(est_time=60, suite="stage-b-test-1-gpu-small-amd-mi35x") + +HIDDEN = 4096 +# Qwen3.5-397B at TP4: 4 K heads x 128 + 16 V heads x 128. +QKVZ = 2 * 4 * 128 + 2 * 16 * 128 +BA = 2 * 16 +# bf16 carries ~8 mantissa bits, so one ULP is ~4e-3 relative. +TOL = 8e-3 + + +class _Linear(torch.nn.Module): + """Stand-in for the unquantized column-parallel projections.""" + + def __init__(self, rows, device): + super().__init__() + self.weight = torch.nn.Parameter( + torch.randn(rows, HIDDEN, dtype=torch.bfloat16, device=device) * 0.02, + requires_grad=False, + ) + self.bias = None + self.quant_method = UnquantizedLinearMethod() + + def forward(self, x): + return torch.nn.functional.linear(x, self.weight), None + + +class _GDN: + """Only the attributes Qwen3_5GatedDeltaNet's input projection touches.""" + + def __init__(self, device): + self.config = types.SimpleNamespace(model_type="qwen3_5_moe_text") + self.in_proj_qkvz = _Linear(QKVZ, device) + self.in_proj_ba = _Linear(BA, device) + self._fused_in_proj_weight = None + self._fused_in_proj_qkvz_width = None + self._fused_in_proj_ba_width = 0 + self._fused_in_proj_scale = None + self._fused_in_proj_sources = None + self.alt_stream = None + self._fused_input_proj_cpu_enabled = types.SimpleNamespace(value=False) + + +@unittest.skipUnless(torch.cuda.is_available() and torch.version.hip, "ROCm aiter GEMM") +class TestQwen35GDNPackedInProj(CustomTestCase): + @classmethod + def setUpClass(cls): + os.environ.setdefault("SGLANG_USE_AITER", "1") + from sglang.srt.models import qwen3_5 + + if not qwen3_5._use_aiter: + raise unittest.SkipTest("qwen3_5 was imported without SGLANG_USE_AITER") + lora = patch.object( + qwen3_5, + "get_lora", + return_value=types.SimpleNamespace(enable_lora=False, lora_paths=None), + ) + lora.start() + cls.addClassCleanup(lora.stop) + cls.module = qwen3_5 + cls.gdn_cls = qwen3_5.Qwen3_5GatedDeltaNet + torch.manual_seed(0) + cls.gdn = _GDN(torch.device("cuda", 0)) + for name in ( + "_finalize_fused_fp8_in_proj", + "_fused_fp8_in_proj_sources_valid", + "_forward_input_proj_fused_quant_amd", + ): + setattr( + cls.gdn, name, types.MethodType(getattr(cls.gdn_cls, name), cls.gdn) + ) + cls.pre = [ + m.weight.data.clone() for m in (cls.gdn.in_proj_qkvz, cls.gdn.in_proj_ba) + ] + cls.gdn_cls.finalize_fused_in_proj(cls.gdn) + + @staticmethod + def _scope_gdn(model_type): + def linear(rows): + return types.SimpleNamespace( + weight=torch.nn.Parameter( + torch.arange(rows * 4, dtype=torch.bfloat16).reshape(rows, 4), + requires_grad=False, + ), + bias=None, + quant_method=UnquantizedLinearMethod(), + ) + + return types.SimpleNamespace( + config=types.SimpleNamespace(model_type=model_type), + in_proj_qkvz=linear(8), + in_proj_ba=linear(4), + _fused_in_proj_weight=None, + _finalize_fused_fp8_in_proj=lambda: False, + ) + + def test_rocm_packing_is_limited_to_qwen35(self): + for model_type, supported in ( + ("qwen3_5_text", True), + ("qwen3_5_moe_text", True), + ("qwen4_exp_text", False), + ("other_text_model", False), + ): + with self.subTest(model_type=model_type): + gdn = self._scope_gdn(model_type) + original = (gdn.in_proj_qkvz.weight, gdn.in_proj_ba.weight) + pointers = [weight.data_ptr() for weight in original] + self.gdn_cls.finalize_fused_in_proj(gdn) + self.assertEqual(gdn._fused_in_proj_weight is not None, supported) + if not supported: + self.assertEqual( + [weight.data_ptr() for weight in original], pointers + ) + + def test_cuda_packing_keeps_qwen4_support(self): + gdn = self._scope_gdn("qwen4_exp_text") + with ( + patch.object(self.module, "_is_cuda", True), + patch.object(self.module, "_use_aiter", False), + ): + self.gdn_cls.finalize_fused_in_proj(gdn) + self.assertIsNotNone(gdn._fused_in_proj_weight) + self.assertEqual( + gdn.in_proj_qkvz.weight.data_ptr(), gdn._fused_in_proj_weight.data_ptr() + ) + + def test_rocm_prepare_hook_is_limited_to_qwen35(self): + for model_type, supported in ( + ("qwen3_5_text", True), + ("qwen3_5_moe_text", True), + ("qwen4_exp_text", False), + ("other_text_model", False), + ): + with self.subTest(model_type=model_type): + gdn = Mock(spec=self.gdn_cls) + gdn._fused_in_proj_weight = None + model = types.SimpleNamespace( + config=types.SimpleNamespace(model_type=model_type), + modules=Mock(return_value=[gdn]), + flashinfer_mnnvl_cutedsl_fusion=None, + ) + self.module.Qwen3_5ForCausalLM.prepare_before_cuda_graph_capture( + model, None + ) + self.assertEqual(model.modules.call_count, int(supported)) + self.assertEqual(gdn.finalize_fused_in_proj.call_count, int(supported)) + + def test_packed_views_alias_and_preserve_values(self): + gdn = self.gdn + fused = gdn._fused_in_proj_weight + self.assertIsNotNone(fused) + self.assertEqual(tuple(fused.shape), (QKVZ + BA, HIDDEN)) + self.assertTrue(fused.is_contiguous()) + self.assertEqual(gdn._fused_in_proj_qkvz_width, QKVZ) + self.assertEqual(gdn.in_proj_qkvz.weight.data_ptr(), fused.data_ptr()) + self.assertEqual( + gdn.in_proj_ba.weight.data_ptr(), + fused.data_ptr() + QKVZ * HIDDEN * fused.element_size(), + ) + for got, want in zip((gdn.in_proj_qkvz, gdn.in_proj_ba), self.pre): + self.assertTrue(torch.equal(got.weight.data, want)) + # finalize is idempotent: a second call keeps the packed buffer + self.gdn_cls.finalize_fused_in_proj(gdn) + self.assertIs(gdn._fused_in_proj_weight, fused) + + def test_packed_and_separate_paths_agree(self): + gdn = self.gdn + for tokens in (1, 4, 33, 64, 65, 300): + with self.subTest(tokens=tokens): + x = torch.randn( + tokens, + HIDDEN, + dtype=torch.bfloat16, + device=gdn._fused_in_proj_weight.device, + ) + want_qkvz = torch.nn.functional.linear(x, self.pre[0]) + want_ba = torch.nn.functional.linear(x, self.pre[1]) + # the fused AR+RMSNorm path hands the bf16 side in a tuple + for hidden in (x, (x, None, None)): + qkvz, ba = self.gdn_cls._forward_input_proj(gdn, hidden) + self.assertEqual(tuple(qkvz.shape), (tokens, QKVZ)) + self.assertEqual(tuple(ba.shape), (tokens, BA)) + for name, got, want in ( + ("qkvz", qkvz, want_qkvz), + ("ba", ba, want_ba), + ): + scale = want.float().abs().max().clamp_min(1e-6) + rel = ((got.float() - want.float()).abs().max() / scale).item() + self.assertLess( + rel, TOL, f"{name} rel err {rel:.2e} at {tokens} tokens" + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/model_executor/test_derived_weight_cache.py b/test/registered/unit/model_executor/test_derived_weight_cache.py new file mode 100644 index 000000000000..aca13a294fea --- /dev/null +++ b/test/registered/unit/model_executor/test_derived_weight_cache.py @@ -0,0 +1,59 @@ +"""Online updates must reject model-owned derived buffers before writing.""" + +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.srt.model_executor.model_runner_components.weight_updater import ( + WeightUpdater, + _unsupported_derived_weight_cache_error, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") + + +class TestDerivedWeightCache(unittest.TestCase): + def test_nested_model_cache_rejects_updates(self): + model = torch.nn.Sequential( + torch.nn.Linear(2, 2, bias=False, device="cpu"), + torch.nn.Sequential(torch.nn.Module()), + ) + model[1][0]._derived_weight_cache_error = "derived scales require restart" + original = model[0].weight.detach().clone() + updater = WeightUpdater( + tp_rank=0, + device="cpu", + gpu_id=0, + model_config=None, + custom_weight_loaders={}, + get_model=lambda: model, + update_model_fields=lambda *args, **kwargs: None, + recapture_cuda_graph=lambda: None, + get_model_runner=lambda: None, + ) + with patch( + "sglang.srt.model_executor.model_runner_components.weight_updater.get_model", + return_value=SimpleNamespace(weight_cache_mode="off"), + ): + self.assertEqual( + updater.update_weights_from_tensor( + [("0.weight", torch.zeros_like(original))], load_format="direct" + ), + (False, "derived scales require restart"), + ) + self.assertTrue(torch.equal(model[0].weight, original)) + + def test_model_without_derived_cache_keeps_updates_enabled(self): + model = torch.nn.Module() + with patch( + "sglang.kernels.ops.attention.dsv4.gemm.hpc_bf16xfp32_gemm_enabled", + return_value=False, + ): + self.assertIsNone(_unsupported_derived_weight_cache_error(model)) + + +if __name__ == "__main__": + unittest.main()