[diffusion] server: use meta to avoid Linear init for TextEncoder (#13564)

Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
zyksir
2025-11-20 23:14:24 -08:00
committed by GitHub
parent b537ac0d1f
commit 4360279036

View File

@@ -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)