diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 934579fd3..1d1a35c37 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -994,6 +994,11 @@ class DefaultModelLoader(BaseModelLoader): @staticmethod def load_weights_and_postprocess(model, weights, target_device): + DefaultModelLoader.load_weights_only(model, weights, target_device) + DefaultModelLoader.postprocess_weights(model, target_device) + + @staticmethod + def load_weights_only(model, weights, target_device): # Used in tests to verify memory savings when using online quantization. if is_cuda_alike(): peak_memory = torch.cuda.max_memory_allocated() @@ -1049,6 +1054,8 @@ class DefaultModelLoader(BaseModelLoader): f"{memory_start - memory_end:.3f}", ) + @staticmethod + def postprocess_weights(model, target_device): for _, module in model.named_modules(): quant_method = getattr(module, "quant_method", None) if quant_method is not None: diff --git a/test/registered/unit/model_loader/test_default_model_loader.py b/test/registered/unit/model_loader/test_default_model_loader.py new file mode 100644 index 000000000..da5b3676a --- /dev/null +++ b/test/registered/unit/model_loader/test_default_model_loader.py @@ -0,0 +1,120 @@ +import unittest +from contextlib import nullcontext +from types import SimpleNamespace +from unittest.mock import Mock, patch + +import torch + +import sglang.srt.model_loader.loader as loader_mod +from sglang.srt.model_loader.loader import DefaultModelLoader +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +class TestDefaultModelLoader(CustomTestCase): + def test_load_weights_only_precedes_postprocessing(self): + events = [] + model = Mock() + model.quant_config = None + module = Mock() + module.quant_method.process_weights_after_loading.side_effect = lambda _: ( + events.append("postprocess") + ) + model.load_weights.side_effect = lambda _: events.append("load") + model.named_modules.return_value = [("layer", module)] + + with ( + patch.object(loader_mod, "is_cuda_alike", return_value=False), + patch.object( + loader_mod, + "device_loading_context", + side_effect=lambda *_: nullcontext(), + ), + ): + DefaultModelLoader.load_weights_and_postprocess( + model, + iter(()), + torch.device("cpu"), + ) + + self.assertEqual(events, ["load", "postprocess"]) + + def test_boot_paths_preserve_load_order_and_custom_override(self): + class CustomModelLoader(DefaultModelLoader): + def load_weights_and_postprocess(self, model, weights, target_device): + events.append("override") + DefaultModelLoader.load_weights_and_postprocess( + model, weights, target_device + ) + + for loader_class in (DefaultModelLoader, CustomModelLoader): + with self.subTest(loader=loader_class.__name__): + events = [] + loader = object.__new__(loader_class) + loader.load_config = object() + model = Mock(quant_config=None) + model.eval.return_value = model + model.load_weights.side_effect = lambda weights: events.append( + ("load", list(weights)) + ) + module = Mock() + module.quant_method.process_weights_after_loading.side_effect = ( + lambda _: events.append("postprocess") + ) + model.named_modules.return_value = [("layer", module)] + model_config = SimpleNamespace(modelopt_quant=None, dtype=torch.float32) + resolved_source = SimpleNamespace(source=object()) + + with ( + patch.object(loader_mod, "is_cuda_alike", return_value=False), + patch.object( + loader_mod, "_get_quantization_config", return_value=None + ), + patch.object(loader_mod, "_initialize_model", return_value=model), + patch.object( + loader_mod, + "device_loading_context", + side_effect=lambda *_: nullcontext(), + ), + patch.object( + loader, "_get_all_weights", return_value=iter([("direct", 1)]) + ), + patch.object( + loader, + "_get_weights_iterator", + return_value=iter([("deferred", 2)]), + ) as get_weights, + ): + self.assertIs( + loader.load_model( + model_config=model_config, + device_config=SimpleNamespace(device="cpu"), + ), + model, + ) + loader.commit_model_weights( + model=model, + model_config=model_config, + resolved_sources=(resolved_source,), + target_device=torch.device("cpu"), + startup_prefetch_active=True, + ) + + expected = [] + for weights in ([("direct", 1)], [("deferred", 2)]): + if loader_class is CustomModelLoader: + expected.append("override") + expected.extend([("load", weights), "postprocess"]) + self.assertEqual(events, expected) + get_weights.assert_called_once_with( + resolved_source.source, + resolved_source=resolved_source, + startup_prefetch_started=True, + startup_prefetch_active=True, + ) + + +if __name__ == "__main__": + unittest.main()