[diffusion] refactor: refactor ComponentLoader and support loading native models from diffusers and transformers (#13205)

This commit is contained in:
Mick
2025-11-28 14:17:32 +08:00
committed by GitHub
parent 3543a04a48
commit 951028968c
3 changed files with 256 additions and 162 deletions
@@ -7,7 +7,6 @@ from typing import Any
import torch import torch
from sglang.multimodal_gen.runtime.managers.forward_context import get_forward_context from sglang.multimodal_gen.runtime.managers.forward_context import get_forward_context
from sglang.srt.layers.attention.flashattention_backend import FlashAttentionMetadata
try: try:
from sgl_kernel.flash_attn import flash_attn_varlen_func from sgl_kernel.flash_attn import flash_attn_varlen_func
@@ -10,7 +10,7 @@ import time
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from collections.abc import Generator, Iterable from collections.abc import Generator, Iterable
from copy import deepcopy from copy import deepcopy
from typing import cast from typing import Any, cast
import torch import torch
import torch.distributed as dist import torch.distributed as dist
@@ -20,7 +20,7 @@ from torch.distributed import init_device_mesh
from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from sglang.multimodal_gen.configs.models import EncoderConfig from sglang.multimodal_gen.configs.models import EncoderConfig, ModelConfig
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.fsdp_load import ( from sglang.multimodal_gen.runtime.loader.fsdp_load import (
maybe_load_fsdp_model, maybe_load_fsdp_model,
@@ -38,7 +38,8 @@ from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
get_config, get_config,
get_diffusers_config, get_diffusers_component_config,
get_hf_config,
) )
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
@@ -60,24 +61,152 @@ class skip_init_modules:
cls.reset_parameters = orig cls.reset_parameters = orig
def _normalize_module_type(module_type: str) -> str:
"""Normalize module types like 'text_encoder_2' -> 'text_encoder'."""
if module_type.endswith("_2"):
return module_type[:-2]
return module_type
def _clean_hf_config_inplace(model_config: dict) -> None:
"""Remove common extraneous HF fields if present."""
for key in (
"_name_or_path",
"transformers_version",
"model_type",
"tokenizer_class",
"torch_dtype",
):
model_config.pop(key, None)
def _list_safetensors_files(model_path: str) -> list[str]:
"""List all .safetensors files under a directory."""
return sorted(glob.glob(os.path.join(str(model_path), "*.safetensors")))
def load_native(library, component_module_path: str, server_args: ServerArgs):
if library == "transformers":
from transformers import AutoModel
config = get_hf_config(
component_module_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
)
return AutoModel.from_pretrained(
component_module_path,
config=config,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
)
elif library == "diffusers":
import diffusers
config = get_diffusers_component_config(model_path=component_module_path)
class_name = config.pop("_class_name", None)
if class_name:
cls = getattr(diffusers, class_name)
return cls.from_pretrained(
component_module_path, revision=server_args.revision, **config
)
else:
raise ValueError("Cannot determine class name for generic diffusers loader")
else:
raise ValueError(f"Unsupported library: {library}")
class ComponentLoader(ABC): class ComponentLoader(ABC):
"""Base class for loading a specific type of model component.""" """Base class for loading a specific type of model component."""
def __init__(self, device=None) -> None: def __init__(self, device=None) -> None:
self.device = device self.device = device
def should_offload(self, server_args, model_config: ModelConfig | None = None):
raise NotImplementedError()
def target_device(self, should_offload):
if should_offload:
return (
torch.device("mps")
if current_platform.is_mps()
else torch.device("cpu")
)
else:
return get_local_torch_device()
def load(
self,
component_model_path: str,
server_args: ServerArgs,
module_name: str,
transformers_or_diffusers: str,
):
"""
Template method that standardizes logging around the core load implementation.
The priority of loading method is:
1. load customized module
2. load native diffusers/transformers module
If all of the above methods failed, an error will be thrown
"""
logger.info("Loading %s from %s", module_name, component_model_path)
try:
component = self.load_customized(
component_model_path, server_args, module_name
)
source = "customized"
except Exception as _e:
# fallback to native version
component = self.load_native(
component_model_path, server_args, transformers_or_diffusers
)
should_offload = self.should_offload(server_args)
target_device = self.target_device(should_offload)
component = component.to(device=target_device)
source = "native"
logger.warning(
"Native module %s: %s is loaded, performance may be sub-optimal",
module_name,
component.__class__.__name__,
)
if component is None:
logger.warning("Loaded %s returned None", module_name)
else:
logger.info(
f"Loaded %s: %s from: {source}",
module_name,
component.__class__.__name__,
)
return component
def load_native(
self,
component_model_path: str,
server_args: ServerArgs,
transformers_or_diffusers: str,
):
"""
Load the component using the native library (transformers/diffusers).
"""
return load_native(transformers_or_diffusers, component_model_path, server_args)
def load_customized(
self, component_model_path: str, server_args: ServerArgs, module_name: str
):
"""
Load the customized version component, implemented and optimized in SGL-diffusion
"""
raise NotImplementedError(
f"load_customized not implemented for {self.__class__.__name__}"
)
@abstractmethod @abstractmethod
def load(self, model_path: str, server_args: ServerArgs, module_name: str): def load_customized(
""" self, model_path: str, server_args: ServerArgs, module_name: str
Load the component based on the model path, architecture, and inference args. ) -> Any:
"""Implement the minimal core load logic in subclasses."""
Args:
model_path: Path to the component model
server_args: ServerArgs
Returns:
The loaded component
"""
raise NotImplementedError raise NotImplementedError
@classmethod @classmethod
@@ -95,15 +224,13 @@ class ComponentLoader(ABC):
A component loader for the specified module type A component loader for the specified module type
""" """
# Map of module types to their loader classes and expected library # Map of module types to their loader classes and expected library
module_type = _normalize_module_type(module_type)
module_loaders = { module_loaders = {
"scheduler": (SchedulerLoader, "diffusers"), "scheduler": (SchedulerLoader, "diffusers"),
"transformer": (TransformerLoader, "diffusers"), "transformer": (TransformerLoader, "diffusers"),
"transformer_2": (TransformerLoader, "diffusers"),
"vae": (VAELoader, "diffusers"), "vae": (VAELoader, "diffusers"),
"text_encoder": (TextEncoderLoader, "transformers"), "text_encoder": (TextEncoderLoader, "transformers"),
"text_encoder_2": (TextEncoderLoader, "transformers"),
"tokenizer": (TokenizerLoader, "transformers"), "tokenizer": (TokenizerLoader, "transformers"),
"tokenizer_2": (TokenizerLoader, "transformers"),
"image_processor": (ImageProcessorLoader, "transformers"), "image_processor": (ImageProcessorLoader, "transformers"),
"image_encoder": (ImageEncoderLoader, "transformers"), "image_encoder": (ImageEncoderLoader, "transformers"),
"processor": (AutoProcessorLoader, "transformers"), "processor": (AutoProcessorLoader, "transformers"),
@@ -147,6 +274,12 @@ class TextEncoderLoader(ComponentLoader):
counter_before_loading_weights: float = 0.0 counter_before_loading_weights: float = 0.0
counter_after_loading_weights: float = 0.0 counter_after_loading_weights: float = 0.0
def should_offload(self, server_args, model_config: ModelConfig | None = None):
should_offload = server_args.text_encoder_cpu_offload
fsdp_shard_conditions = getattr(model_config, "_fsdp_shard_conditions", [])
use_cpu_offload = should_offload and len(fsdp_shard_conditions) > 0
return use_cpu_offload
def _prepare_weights( def _prepare_weights(
self, self,
model_name_or_path: str, model_name_or_path: str,
@@ -238,7 +371,9 @@ class TextEncoderLoader(ComponentLoader):
for source in secondary_weights: for source in secondary_weights:
yield from self._get_weights_iterator(source, to_cpu) yield from self._get_weights_iterator(source, to_cpu)
def load(self, model_path: str, server_args: ServerArgs, module_name: str): def load_customized(
self, component_model_path: str, server_args: ServerArgs, module_name: str
):
"""Load the text encoders based on the model path, and inference args.""" """Load the text encoders based on the model path, and inference args."""
# model_config: PretrainedConfig = get_hf_config( # model_config: PretrainedConfig = get_hf_config(
# model=model_path, # model=model_path,
@@ -246,13 +381,11 @@ class TextEncoderLoader(ComponentLoader):
# revision=server_args.revision, # revision=server_args.revision,
# model_override_args=None, # model_override_args=None,
# ) # )
diffusers_pretrained_config = get_config(model_path, trust_remote_code=True) diffusers_pretrained_config = get_config(
model_config = get_diffusers_config(model=model_path) component_model_path, trust_remote_code=True
model_config.pop("_name_or_path", None) )
model_config.pop("transformers_version", None) model_config = get_diffusers_component_config(model_path=component_model_path)
model_config.pop("model_type", None) _clean_hf_config_inplace(model_config)
model_config.pop("tokenizer_class", None)
model_config.pop("torch_dtype", None)
logger.info("HF model config: %s", model_config) logger.info("HF model config: %s", model_config)
def is_not_first_encoder(module_name): def is_not_first_encoder(module_name):
@@ -270,12 +403,10 @@ class TextEncoderLoader(ComponentLoader):
encoder_config = server_args.pipeline_config.text_encoder_configs[1] encoder_config = server_args.pipeline_config.text_encoder_configs[1]
encoder_config.update_model_arch(model_config) encoder_config.update_model_arch(model_config)
encoder_dtype = server_args.pipeline_config.text_encoder_precisions[1] encoder_dtype = server_args.pipeline_config.text_encoder_precisions[1]
target_device = get_local_torch_device()
# TODO(will): add support for other dtypes # TODO(will): add support for other dtypes
return self.load_model( return self.load_model(
model_path, component_model_path,
encoder_config, encoder_config,
target_device,
server_args, server_args,
encoder_dtype, encoder_dtype,
) )
@@ -284,31 +415,23 @@ class TextEncoderLoader(ComponentLoader):
self, self,
model_path: str, model_path: str,
model_config: EncoderConfig, model_config: EncoderConfig,
target_device: torch.device,
server_args: ServerArgs, server_args: ServerArgs,
dtype: str = "fp16", dtype: str = "fp16",
cpu_offload_flag: bool | None = None,
): ):
use_cpu_offload = ( # Determine CPU offload behavior and target device
server_args.text_encoder_cpu_offload
and len(getattr(model_config, "_fsdp_shard_conditions", [])) > 0
)
if server_args.text_encoder_cpu_offload:
target_device = (
torch.device("mps")
if current_platform.is_mps()
else torch.device("cpu")
)
local_torch_device = get_local_torch_device()
should_offload = self.should_offload(server_args, model_config)
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]): with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
with target_device, skip_init_modules(): with local_torch_device, skip_init_modules():
architectures = getattr(model_config, "architectures", []) architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures) model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
model = model_cls(model_config) model = model_cls(model_config)
weights_to_load = {name for name, _ in model.named_parameters()} weights_to_load = {name for name, _ in model.named_parameters()}
loaded_weights = model.load_weights( loaded_weights = model.load_weights(
self._get_all_weights(model, model_path, to_cpu=use_cpu_offload) self._get_all_weights(model, model_path, to_cpu=should_offload)
) )
self.counter_after_loading_weights = time.perf_counter() self.counter_after_loading_weights = time.perf_counter()
logger.info( logger.info(
@@ -318,9 +441,9 @@ class TextEncoderLoader(ComponentLoader):
) )
# Explicitly move model to target device after loading weights # Explicitly move model to target device after loading weights
model = model.to(target_device) model = model.to(local_torch_device)
if use_cpu_offload: if should_offload:
# Disable FSDP for MPS as it's not compatible # Disable FSDP for MPS as it's not compatible
if current_platform.is_mps(): if current_platform.is_mps():
logger.info( logger.info(
@@ -355,8 +478,15 @@ class TextEncoderLoader(ComponentLoader):
class ImageEncoderLoader(TextEncoderLoader): class ImageEncoderLoader(TextEncoderLoader):
def should_offload(self, server_args, model_config: ModelConfig | None = None):
should_offload = server_args.image_encoder_cpu_offload
fsdp_shard_conditions = getattr(model_config, "_fsdp_shard_conditions", [])
use_cpu_offload = should_offload and len(fsdp_shard_conditions) > 0
return use_cpu_offload
def load(self, model_path: str, server_args: ServerArgs, *args): def load_customized(
self, component_model_path: str, server_args: ServerArgs, *args
):
"""Load the text encoders based on the model path, and inference args.""" """Load the text encoders based on the model path, and inference args."""
# model_config: PretrainedConfig = get_hf_config( # model_config: PretrainedConfig = get_hf_config(
# model=model_path, # model=model_path,
@@ -364,91 +494,73 @@ class ImageEncoderLoader(TextEncoderLoader):
# revision=server_args.revision, # revision=server_args.revision,
# model_override_args=None, # model_override_args=None,
# ) # )
with open(os.path.join(model_path, "config.json")) as f: with open(os.path.join(component_model_path, "config.json")) as f:
model_config = json.load(f) model_config = json.load(f)
model_config.pop("_name_or_path", None) _clean_hf_config_inplace(model_config)
model_config.pop("transformers_version", None)
model_config.pop("torch_dtype", None)
model_config.pop("model_type", None)
logger.info("HF model config: %s", model_config) logger.info("HF model config: %s", model_config)
encoder_config = server_args.pipeline_config.image_encoder_config encoder_config = server_args.pipeline_config.image_encoder_config
encoder_config.update_model_arch(model_config) encoder_config.update_model_arch(model_config)
if server_args.image_encoder_cpu_offload: # Always start with local device; load_model will adjust for offload if needed
target_device = ( should_offload = self.should_offload(server_args)
torch.device("mps")
if current_platform.is_mps()
else torch.device("cpu")
)
else:
target_device = get_local_torch_device()
# TODO(will): add support for other dtypes # TODO(will): add support for other dtypes
return self.load_model( return self.load_model(
model_path, component_model_path,
encoder_config, encoder_config,
target_device,
server_args, server_args,
server_args.pipeline_config.image_encoder_precision, server_args.pipeline_config.image_encoder_precision,
cpu_offload_flag=server_args.image_encoder_cpu_offload,
) )
class ImageProcessorLoader(ComponentLoader): class ImageProcessorLoader(ComponentLoader):
"""Loader for image processor.""" """Loader for image processor."""
def load(self, model_path: str, server_args: ServerArgs, *args): def load_customized(
"""Load the image processor based on the model path, and inference args.""" self, component_model_path: str, server_args: ServerArgs, module_name: str
logger.info("Loading image processor from %s", model_path) ) -> Any:
return AutoImageProcessor.from_pretrained(component_model_path, use_fast=True)
image_processor = AutoImageProcessor.from_pretrained(model_path, use_fast=True)
logger.info("Loaded image processor: %s", image_processor.__class__.__name__)
return image_processor
class AutoProcessorLoader(ComponentLoader): class AutoProcessorLoader(ComponentLoader):
"""Loader for auto processor.""" """Loader for auto processor."""
def load(self, model_path: str, server_args: ServerArgs, *args): def load_customized(
"""Load the image processor based on the model path, and inference args.""" self, component_model_path: str, server_args: ServerArgs, module_name: str
logger.info("Loading auto processor from %s", model_path) ) -> Any:
return AutoProcessor.from_pretrained(component_model_path)
processor = AutoProcessor.from_pretrained(
model_path,
)
logger.info("Loaded auto processor: %s", processor.__class__.__name__)
return processor
class TokenizerLoader(ComponentLoader): class TokenizerLoader(ComponentLoader):
"""Loader for tokenizers.""" """Loader for tokenizers."""
def load(self, model_path: str, server_args: ServerArgs, *args): def load_customized(
"""Load the tokenizer based on the model path, and inference args.""" self, component_model_path: str, server_args: ServerArgs, module_name: str
logger.info("Loading tokenizer from %s", model_path) ) -> Any:
return AutoTokenizer.from_pretrained(
tokenizer = AutoTokenizer.from_pretrained( component_model_path,
model_path, # "<path to model>/tokenizer"
# in v0, this was same string as encoder_name "ClipTextModel"
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
# other method of config?
padding_size="right", padding_size="right",
) )
logger.info("Loaded tokenizer: %s", tokenizer.__class__.__name__)
return tokenizer
class VAELoader(ComponentLoader): class VAELoader(ComponentLoader):
"""Loader for VAE.""" """Loader for VAE."""
def load(self, model_path: str, server_args: ServerArgs, *args): def should_offload(self, server_args, cpu_offload_flag, model_config):
return True
def load_customized(
self, component_model_path: str, server_args: ServerArgs, *args
):
"""Load the VAE based on the model path, and inference args.""" """Load the VAE based on the model path, and inference args."""
config = get_diffusers_config(model=model_path) config = get_diffusers_component_config(model_path=component_model_path)
class_name = config.pop("_class_name") class_name = config.pop("_class_name")
assert ( assert (
class_name is not None class_name is not None
), "Model config does not contain a _class_name attribute. Only diffusers format is supported." ), "Model config does not contain a _class_name attribute. Only diffusers format is supported."
server_args.model_paths["vae"] = model_path server_args.model_paths["vae"] = component_model_path
# TODO: abstract these logics # TODO: abstract these logics
logger.info("HF model config: %s", config) logger.info("HF model config: %s", config)
@@ -458,14 +570,7 @@ class VAELoader(ComponentLoader):
# NOTE: some post init logics are only available after updated with config # NOTE: some post init logics are only available after updated with config
vae_config.post_init() vae_config.post_init()
if server_args.vae_cpu_offload: target_device = self.target_device(server_args.vae_cpu_offload)
target_device = (
torch.device("mps")
if current_platform.is_mps()
else torch.device("cpu")
)
else:
target_device = get_local_torch_device()
with set_default_torch_dtype( with set_default_torch_dtype(
PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision] PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
@@ -474,11 +579,11 @@ class VAELoader(ComponentLoader):
vae = vae_cls(vae_config).to(target_device) vae = vae_cls(vae_config).to(target_device)
# Find all safetensors files # Find all safetensors files
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors")) safetensors_list = _list_safetensors_files(component_model_path)
# TODO(PY) # TODO(PY)
assert ( assert (
len(safetensors_list) == 1 len(safetensors_list) == 1
), f"Found {len(safetensors_list)} safetensors files in {model_path}" ), f"Found {len(safetensors_list)} safetensors files in {component_model_path}"
loaded = safetensors_load_file(safetensors_list[0]) loaded = safetensors_load_file(safetensors_list[0])
vae.load_state_dict( vae.load_state_dict(
loaded, strict=False loaded, strict=False
@@ -490,9 +595,11 @@ class VAELoader(ComponentLoader):
class TransformerLoader(ComponentLoader): class TransformerLoader(ComponentLoader):
"""Loader for transformer.""" """Loader for transformer."""
def load(self, model_path: str, server_args: ServerArgs, *args): def load_customized(
self, component_model_path: str, server_args: ServerArgs, *args
):
"""Load the transformer based on the model path, and inference args.""" """Load the transformer based on the model path, and inference args."""
config = get_diffusers_config(model=model_path) config = get_diffusers_component_config(model_path=component_model_path)
hf_config = deepcopy(config) hf_config = deepcopy(config)
cls_name = config.pop("_class_name") cls_name = config.pop("_class_name")
if cls_name is None: if cls_name is None:
@@ -506,7 +613,7 @@ class TransformerLoader(ComponentLoader):
cls_name = server_args.override_transformer_cls_name cls_name = server_args.override_transformer_cls_name
logger.info("Overriding transformer cls_name to %s", cls_name) logger.info("Overriding transformer cls_name to %s", cls_name)
server_args.model_paths["transformer"] = model_path server_args.model_paths["transformer"] = component_model_path
# Config from Diffusers supersedes sgl_diffusion's model config # Config from Diffusers supersedes sgl_diffusion's model config
dit_config = server_args.pipeline_config.dit_config dit_config = server_args.pipeline_config.dit_config
@@ -515,9 +622,9 @@ class TransformerLoader(ComponentLoader):
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name) model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
# Find all safetensors files # Find all safetensors files
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors")) safetensors_list = _list_safetensors_files(component_model_path)
if not safetensors_list: if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}") raise ValueError(f"No safetensors files found in {component_model_path}")
# Check if we should use custom initialization weights # Check if we should use custom initialization weights
custom_weights_path = getattr( custom_weights_path = getattr(
@@ -533,9 +640,7 @@ class TransformerLoader(ComponentLoader):
custom_weights_path is not None custom_weights_path is not None
), "Custom initialization weights must be provided" ), "Custom initialization weights must be provided"
if os.path.isdir(custom_weights_path): if os.path.isdir(custom_weights_path):
safetensors_list = glob.glob( safetensors_list = _list_safetensors_files(custom_weights_path)
os.path.join(str(custom_weights_path), "*.safetensors")
)
else: else:
assert custom_weights_path.endswith( assert custom_weights_path.endswith(
".safetensors" ".safetensors"
@@ -584,9 +689,11 @@ class TransformerLoader(ComponentLoader):
class SchedulerLoader(ComponentLoader): class SchedulerLoader(ComponentLoader):
"""Loader for scheduler.""" """Loader for scheduler."""
def load(self, model_path: str, server_args: ServerArgs, *args): def load_customized(
self, component_model_path: str, server_args: ServerArgs, *args
):
"""Load the scheduler based on the model path, and inference args.""" """Load the scheduler based on the model path, and inference args."""
config = get_diffusers_config(model=model_path) config = get_diffusers_component_config(model_path=component_model_path)
class_name = config.pop("_class_name") class_name = config.pop("_class_name")
assert ( assert (
@@ -610,36 +717,6 @@ class GenericComponentLoader(ComponentLoader):
super().__init__() super().__init__()
self.library = library self.library = library
def load(self, model_path: str, server_args: ServerArgs, *args):
"""Load a generic component based on the model path, and inference args."""
logger.warning(
"Using generic loader for %s with library %s", model_path, self.library
)
if self.library == "transformers":
from transformers import AutoModel
model = AutoModel.from_pretrained(
model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
)
logger.info(
"Loaded generic transformers model: %s", model.__class__.__name__
)
return model
elif self.library == "diffusers":
logger.warning(
"Generic loading for diffusers components is not fully implemented"
)
model_config = get_diffusers_config(model=model_path)
logger.info("Diffusers Model config: %s", model_config)
# This is a placeholder - in a real implementation, you'd need to handle this properly
return None
else:
raise ValueError(f"Unsupported library: {self.library}")
class PipelineComponentLoader: class PipelineComponentLoader:
""" """
@@ -677,7 +754,12 @@ class PipelineComponentLoader:
try: try:
# Load the module # Load the module
return loader.load(component_model_path, server_args, module_name) return loader.load(
component_model_path,
server_args,
module_name,
transformers_or_diffusers,
)
except Exception as e: except Exception as e:
logger.error( logger.error(
f"Error while loading component: {module_name}, {component_model_path=}" f"Error while loading component: {module_name}, {component_model_path=}"
@@ -23,6 +23,7 @@ import hashlib
import json import json
import os import os
import tempfile import tempfile
from functools import reduce
from pathlib import Path from pathlib import Path
from typing import Any, Optional, cast from typing import Any, Optional, cast
@@ -60,24 +61,27 @@ def download_from_hf(model_path: str):
def get_hf_config( def get_hf_config(
model: str, component_model_path: str,
trust_remote_code: bool, trust_remote_code: bool,
revision: str | None = None, revision: str | None = None,
model_override_args: dict | None = None, model_override_args: dict | None = None,
**kwargs, **kwargs,
): ) -> PretrainedConfig:
is_gguf = check_gguf_file(model) is_gguf = check_gguf_file(component_model_path)
if is_gguf: if is_gguf:
raise NotImplementedError("GGUF models are not supported.") raise NotImplementedError("GGUF models are not supported.")
config = AutoConfig.from_pretrained( config = AutoConfig.from_pretrained(
model, trust_remote_code=trust_remote_code, revision=revision, **kwargs component_model_path,
trust_remote_code=trust_remote_code,
revision=revision,
**kwargs,
) )
if config.model_type in _CONFIG_REGISTRY: if config.model_type in _CONFIG_REGISTRY:
config_class = _CONFIG_REGISTRY[config.model_type] config_class = _CONFIG_REGISTRY[config.model_type]
config = config_class.from_pretrained(model, revision=revision) config = config_class.from_pretrained(component_model_path, revision=revision)
# NOTE(HandH1998): Qwen2VL requires `_name_or_path` attribute in `config`. # NOTE(HandH1998): Qwen2VL requires `_name_or_path` attribute in `config`.
config._name_or_path = model config._name_or_path = component_model_path
if model_override_args: if model_override_args:
config.update(model_override_args) config.update(model_override_args)
@@ -125,30 +129,39 @@ def load_dict(file_path):
) from e ) from e
def get_diffusers_config( def get_diffusers_component_config(
model: str, model_path: str,
) -> dict[str, Any]: ) -> dict[str, Any]:
"""Gets a configuration for the given diffusers model. """Gets a configuration of a submodule for the given diffusers model.
Args: Args:
model: The model name or path. model_path: the path of the submodule
Returns: Returns:
The loaded configuration. The loaded configuration.
""" """
config_name = "config.json"
if "scheduler" in model:
config_name = "scheduler_config.json"
# Check if the model path exists # Check if the model path exists
if os.path.exists(model): if os.path.exists(model_path):
config_file = os.path.join(model, config_name) # tokenizer
config_dict = load_dict(config_file) config_names = ["generation_config.json"]
generation_config_file = os.path.join(model, "generation_config.json") # By default, we load config.json, but scheduler_config.json for scheduler
generation_config_dict = load_dict(generation_config_file) if "scheduler" in model_path:
return config_dict | generation_config_dict config_names.append("scheduler_config.json")
else:
config_names.append("config.json")
config_file_paths = [
os.path.join(model_path, config_name) for config_name in config_names
]
combined_config = reduce(
lambda acc, path: acc | load_dict(path), config_file_paths, {}
)
return combined_config
else: else:
raise RuntimeError(f"Diffusers config file not found at {model}") raise RuntimeError(f"Diffusers config file not found at {model_path}")
# Models don't use the same configuration key for determining the maximum # Models don't use the same configuration key for determining the maximum