[diffusion] improve: skip loading vision module for text encoders (#16304)

This commit is contained in:
Mick
2026-01-03 19:30:45 +08:00
committed by GitHub
parent 38d48de93d
commit 2c09de343e
3 changed files with 38 additions and 54 deletions
@@ -23,6 +23,9 @@ from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from sglang.multimodal_gen.configs.models import EncoderConfig, ModelConfig
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig,
)
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.fsdp_load import (
maybe_load_fsdp_model,
@@ -462,6 +465,14 @@ class TextEncoderLoader(ComponentLoader):
with local_torch_device, skip_init_modules():
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
enable_image_understanding = (
True
if isinstance(
server_args.pipeline_config, QwenImageEditPipelineConfig
)
else False
)
model_config.enable_image_understanding = enable_image_understanding
model = model_cls(model_config)
weights_to_load = {name for name, _ in model.named_parameters()}