[diffusion] chore: minor code cleanups and improve logging (#14916)

This commit is contained in:
Mick
2025-12-12 18:48:07 +08:00
committed by GitHub
parent 94e1251131
commit 82e33170e1
9 changed files with 40 additions and 78 deletions

View File

@@ -87,33 +87,6 @@ def _list_safetensors_files(model_path: str) -> list[str]:
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":
from diffusers import AutoModel
return AutoModel.from_pretrained(
component_module_path,
revision=server_args.revision,
trust_remote_code=server_args.trust_remote_code,
)
else:
raise ValueError(f"Unsupported library: {library}")
class ComponentLoader(ABC):
"""Base class for loading a specific type of model component."""
@@ -193,7 +166,30 @@ class ComponentLoader(ABC):
"""
Load the component using the native library (transformers/diffusers).
"""
return load_native(transformers_or_diffusers, component_model_path, server_args)
if transformers_or_diffusers == "transformers":
from transformers import AutoModel
config = get_hf_config(
component_model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
)
return AutoModel.from_pretrained(
component_model_path,
config=config,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
)
elif transformers_or_diffusers == "diffusers":
from diffusers import AutoModel
return AutoModel.from_pretrained(
component_model_path,
revision=server_args.revision,
trust_remote_code=server_args.trust_remote_code,
)
else:
raise ValueError(f"Unsupported library: {transformers_or_diffusers}")
def load_customized(
self, component_model_path: str, server_args: ServerArgs, module_name: str
@@ -621,7 +617,6 @@ class TransformerLoader(ComponentLoader):
"Only diffusers format is supported."
)
logger.info("transformer cls_name: %s", cls_name)
if server_args.override_transformer_cls_name is not None:
cls_name = server_args.override_transformer_cls_name
logger.info("Overriding transformer cls_name to %s", cls_name)
@@ -660,16 +655,16 @@ class TransformerLoader(ComponentLoader):
), "Custom initialization weights must be a safetensors file"
safetensors_list = [custom_weights_path]
logger.info(
"Loading model from %s safetensors files: %s",
len(safetensors_list),
safetensors_list,
)
default_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision]
logger.info(
"Loading %s from %s safetensors files, default_dtype: %s",
cls_name,
len(safetensors_list),
default_dtype,
)
# Load the model using FSDP loader
logger.info("Loading %s, default_dtype: %s", cls_name, default_dtype)
assert server_args.hsdp_shard_dim is not None
model = maybe_load_fsdp_model(
model_cls=model_cls,
@@ -760,12 +755,6 @@ class PipelineComponentLoader:
Returns:
The loaded module
"""
logger.info(
"Loading %s using %s from %s",
module_name,
transformers_or_diffusers,
component_model_path,
)
# Get the appropriate loader for this module type
loader = ComponentLoader.for_module_type(module_name, transformers_or_diffusers)