From 4360279036cec783bb887e7a3d4bacb6a70e7b82 Mon Sep 17 00:00:00 2001 From: zyksir Date: Thu, 20 Nov 2025 23:14:24 -0800 Subject: [PATCH] [diffusion] server: use meta to avoid Linear init for TextEncoder (#13564) Co-authored-by: Mick --- .../runtime/loader/component_loader.py | 18 ++++++++++++++++-- 1 file changed, 16 insertions(+), 2 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loader.py index bdd3c4822..9cf0c1b92 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loader.py @@ -46,6 +46,20 @@ from sglang.multimodal_gen.utils import PRECISION_TO_TYPE logger = init_logger(__name__) +class skip_init_modules: + def __enter__(self): + # Save originals + self._orig_reset = {} + for cls in (nn.Linear, nn.Conv1d, nn.Conv2d, nn.Conv3d): + self._orig_reset[cls] = cls.reset_parameters + cls.reset_parameters = lambda self: None # skip init + + def __exit__(self, exc_type, exc_value, traceback): + # Restore originals + for cls, orig in self._orig_reset.items(): + cls.reset_parameters = orig + + class ComponentLoader(ABC): """Base class for loading a specific type of model component.""" @@ -287,7 +301,7 @@ class TextEncoderLoader(ComponentLoader): ) with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]): - with target_device: + with target_device, skip_init_modules(): architectures = getattr(model_config, "architectures", []) model_cls, _ = ModelRegistry.resolve_model_cls(architectures) model = model_cls(model_config) @@ -454,7 +468,7 @@ class VAELoader(ComponentLoader): with set_default_torch_dtype( PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision] - ): + ), skip_init_modules(): vae_cls, _ = ModelRegistry.resolve_model_cls(class_name) vae = vae_cls(vae_config).to(target_device)