[diffusion] improve: skip loading vision module for text encoders (#16304)
This commit is contained in:
@@ -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()}
|
||||
|
||||
Reference in New Issue
Block a user