[diffusion] refactor: unify model loading and offloading behavior (#15923)

This commit is contained in:
Mick
2025-12-27 16:18:24 +08:00
committed by GitHub
parent 171912a9e3
commit aa89c6a7e2
9 changed files with 97 additions and 38 deletions
@@ -96,9 +96,11 @@ class ComponentLoader(ABC):
def __init__(self, device=None) -> None:
self.device = device
def should_offload(self, server_args, model_config: ModelConfig | None = None):
# offload by default
return True
def should_offload(
self, server_args: ServerArgs, model_config: ModelConfig | None = None
):
# not offload by default
return False
def target_device(self, should_offload):
if should_offload:
@@ -276,6 +278,8 @@ class TextEncoderLoader(ComponentLoader):
def should_offload(self, server_args, model_config: ModelConfig | None = None):
should_offload = server_args.text_encoder_cpu_offload
if not should_offload:
return False
# _fsdp_shard_conditions is in arch_config, not directly on model_config
arch_config = (
getattr(model_config, "arch_config", model_config) if model_config else None
@@ -486,6 +490,8 @@ class TextEncoderLoader(ComponentLoader):
class ImageEncoderLoader(TextEncoderLoader):
def should_offload(self, server_args, model_config: ModelConfig | None = None):
should_offload = server_args.image_encoder_cpu_offload
if not should_offload:
return False
# _fsdp_shard_conditions is in arch_config, not directly on model_config
arch_config = (
getattr(model_config, "arch_config", model_config) if model_config else None
@@ -558,8 +564,10 @@ class TokenizerLoader(ComponentLoader):
class VAELoader(ComponentLoader):
"""Loader for VAE."""
def should_offload(self, server_args, cpu_offload_flag, model_config):
return True
def should_offload(
self, server_args: ServerArgs, model_config: ModelConfig | None = None
):
return server_args.vae_cpu_offload
def load_customized(
self, component_model_path: str, server_args: ServerArgs, *args
@@ -580,7 +588,8 @@ class VAELoader(ComponentLoader):
# NOTE: some post init logics are only available after updated with config
vae_config.post_init()
target_device = self.target_device(server_args.vae_cpu_offload)
should_offload = self.should_offload(server_args)
target_device = self.target_device(should_offload)
# Check for auto_map first (custom VAE classes)
auto_map = config.get("auto_map", {})