[diffusion] model: GLM-Image (#16894)
Co-authored-by: jianyingzhu <53300651@qq.com>
This commit is contained in:
@@ -266,6 +266,7 @@ class ComponentLoader(ABC):
|
||||
"image_processor": (ImageProcessorLoader, "transformers"),
|
||||
"image_encoder": (ImageEncoderLoader, "transformers"),
|
||||
"processor": (AutoProcessorLoader, "transformers"),
|
||||
"vision_language_encoder": (VisionLanguageEncoderLoader, "transformers"),
|
||||
}
|
||||
|
||||
if module_type in module_loaders:
|
||||
@@ -779,6 +780,36 @@ class GenericComponentLoader(ComponentLoader):
|
||||
self.library = library
|
||||
|
||||
|
||||
class VisionLanguageEncoderLoader(ComponentLoader):
|
||||
"""Loader for vision language encoder (typically Causal LM or Vision2Seq)."""
|
||||
|
||||
def load_customized(
|
||||
self,
|
||||
component_model_path: str,
|
||||
server_args: ServerArgs,
|
||||
transformers_or_diffusers: str = "vision_language_encoder",
|
||||
) -> Any:
|
||||
if transformers_or_diffusers == "vision_language_encoder":
|
||||
from transformers import GlmImageForConditionalGeneration
|
||||
|
||||
config = get_hf_config(
|
||||
component_model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
revision=server_args.revision,
|
||||
)
|
||||
model = GlmImageForConditionalGeneration.from_pretrained(
|
||||
component_model_path,
|
||||
config=config,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
revision=server_args.revision,
|
||||
).to(get_local_torch_device())
|
||||
return model
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported library for VisionLanguageEncoder: {transformers_or_diffusers}"
|
||||
)
|
||||
|
||||
|
||||
class PipelineComponentLoader:
|
||||
"""
|
||||
Utility class for loading pipeline components.
|
||||
|
||||
Reference in New Issue
Block a user