[diffusion] model: LTX-2 (1/2) (#17495)
Co-authored-by: FlamingoPg <1106310035@qq.com>
This commit is contained in:
@@ -269,6 +269,17 @@ class ComponentLoader(ABC):
|
||||
"vision_language_encoder": (VisionLanguageEncoderLoader, "transformers"),
|
||||
}
|
||||
|
||||
# Loaders for audio/video specific components that might vary
|
||||
av_module_loaders = {
|
||||
"audio_vae": (VAELoader, "diffusers"),
|
||||
"vocoder": (VocoderLoader, "diffusers"),
|
||||
"connectors": (AdapterLoader, "diffusers"),
|
||||
}
|
||||
|
||||
# NOTE(FlamingoPg): special for LTX-2 models
|
||||
if module_type == "vocoder" or module_type == "connectors":
|
||||
transformers_or_diffusers = "diffusers"
|
||||
|
||||
if module_type in module_loaders:
|
||||
loader_cls, expected_library = module_loaders[module_type]
|
||||
# Assert that the library matches what's expected for this module type
|
||||
@@ -277,6 +288,11 @@ class ComponentLoader(ABC):
|
||||
), f"{module_type} must be loaded from {expected_library}, got {transformers_or_diffusers}"
|
||||
return loader_cls()
|
||||
|
||||
if module_type in av_module_loaders:
|
||||
loader_cls, expected_library = av_module_loaders[module_type]
|
||||
if transformers_or_diffusers == expected_library:
|
||||
return loader_cls()
|
||||
|
||||
# For unknown module types, use a generic loader
|
||||
logger.warning(
|
||||
"No specific loader found for module type: %s. Using generic loader.",
|
||||
@@ -484,6 +500,10 @@ class TextEncoderLoader(ComponentLoader):
|
||||
self._get_all_weights(model, model_path, to_cpu=should_offload)
|
||||
)
|
||||
|
||||
# Explicitly move model to target device after loading weights
|
||||
if not should_offload:
|
||||
model = model.to(local_torch_device)
|
||||
|
||||
if should_offload:
|
||||
# Disable FSDP for MPS as it's not compatible
|
||||
if current_platform.is_mps():
|
||||
@@ -604,7 +624,7 @@ class VAELoader(ComponentLoader):
|
||||
return server_args.vae_cpu_offload
|
||||
|
||||
def load_customized(
|
||||
self, component_model_path: str, server_args: ServerArgs, *args
|
||||
self, component_model_path: str, server_args: ServerArgs, module_name: str
|
||||
):
|
||||
"""Load the VAE based on the model path, and inference args."""
|
||||
config = get_diffusers_component_config(model_path=component_model_path)
|
||||
@@ -613,10 +633,16 @@ class VAELoader(ComponentLoader):
|
||||
class_name is not None
|
||||
), "Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
||||
|
||||
server_args.model_paths["vae"] = component_model_path
|
||||
server_args.model_paths[module_name] = component_model_path
|
||||
|
||||
logger.debug("HF model config: %s", config)
|
||||
vae_config = server_args.pipeline_config.vae_config
|
||||
if module_name == "audio_vae":
|
||||
vae_config = server_args.pipeline_config.audio_vae_config
|
||||
vae_precision = server_args.pipeline_config.audio_vae_precision
|
||||
else:
|
||||
vae_config = server_args.pipeline_config.vae_config
|
||||
vae_precision = server_args.pipeline_config.vae_precision
|
||||
|
||||
vae_config.update_model_arch(config)
|
||||
|
||||
# NOTE: some post init logics are only available after updated with config
|
||||
@@ -635,7 +661,7 @@ class VAELoader(ComponentLoader):
|
||||
custom_module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(custom_module)
|
||||
vae_cls = getattr(custom_module, cls_name)
|
||||
vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
||||
vae_dtype = PRECISION_TO_TYPE[vae_precision]
|
||||
with set_default_torch_dtype(vae_dtype):
|
||||
vae = vae_cls.from_pretrained(
|
||||
component_model_path,
|
||||
@@ -647,9 +673,7 @@ class VAELoader(ComponentLoader):
|
||||
|
||||
# Load from ModelRegistry (standard VAE classes)
|
||||
with (
|
||||
set_default_torch_dtype(
|
||||
PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision]
|
||||
),
|
||||
set_default_torch_dtype(PRECISION_TO_TYPE[vae_precision]),
|
||||
skip_init_modules(),
|
||||
):
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
@@ -664,6 +688,71 @@ class VAELoader(ComponentLoader):
|
||||
return vae.eval()
|
||||
|
||||
|
||||
class VocoderLoader(ComponentLoader):
|
||||
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, module_name: str
|
||||
):
|
||||
config = get_diffusers_component_config(model_path=component_model_path)
|
||||
class_name = config.pop("_class_name", None)
|
||||
assert (
|
||||
class_name is not None
|
||||
), "Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
||||
|
||||
server_args.model_paths[module_name] = component_model_path
|
||||
|
||||
from sglang.multimodal_gen.configs.models.vocoder.ltx_vocoder import (
|
||||
LTXVocoderConfig,
|
||||
)
|
||||
|
||||
vocoder_config = LTXVocoderConfig()
|
||||
vocoder_config.update_model_arch(config)
|
||||
|
||||
try:
|
||||
vocoder_precision = server_args.pipeline_config.audio_vae_precision
|
||||
except AttributeError:
|
||||
vocoder_precision = "fp32"
|
||||
vocoder_dtype = PRECISION_TO_TYPE[vocoder_precision]
|
||||
|
||||
should_offload = self.should_offload(server_args)
|
||||
target_device = self.target_device(should_offload)
|
||||
|
||||
with set_default_torch_dtype(vocoder_dtype), skip_init_modules():
|
||||
vocoder_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
vocoder = vocoder_cls(vocoder_config).to(target_device)
|
||||
|
||||
safetensors_list = _list_safetensors_files(component_model_path)
|
||||
assert (
|
||||
len(safetensors_list) == 1
|
||||
), f"Found {len(safetensors_list)} safetensors files in {component_model_path}"
|
||||
loaded = safetensors_load_file(safetensors_list[0])
|
||||
incompatible = vocoder.load_state_dict(loaded, strict=False)
|
||||
missing_keys = []
|
||||
unexpected_keys = []
|
||||
try:
|
||||
missing_keys = incompatible.missing_keys
|
||||
unexpected_keys = incompatible.unexpected_keys
|
||||
except AttributeError:
|
||||
# Best-effort fallback in case older torch returns a tuple-like.
|
||||
try:
|
||||
missing_keys = incompatible[0]
|
||||
unexpected_keys = incompatible[1]
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if missing_keys or unexpected_keys:
|
||||
logger.warning(
|
||||
"Loaded vocoder with missing_keys=%d unexpected_keys=%d",
|
||||
len(missing_keys),
|
||||
len(unexpected_keys),
|
||||
)
|
||||
return vocoder.eval()
|
||||
|
||||
|
||||
class TransformerLoader(ComponentLoader):
|
||||
"""Loader for transformer."""
|
||||
|
||||
@@ -755,6 +844,58 @@ class TransformerLoader(ComponentLoader):
|
||||
return model
|
||||
|
||||
|
||||
class AdapterLoader(ComponentLoader):
|
||||
"""Loader for small adapter-style modules (e.g., LTX-2 connectors).
|
||||
|
||||
This loader intentionally avoids FSDP sharding and just:
|
||||
1) Instantiates the module from `config.json`.
|
||||
2) Loads a single safetensors state_dict.
|
||||
"""
|
||||
|
||||
def load_customized(
|
||||
self, component_model_path: str, server_args: ServerArgs, *args
|
||||
):
|
||||
config = get_diffusers_component_config(model_path=component_model_path)
|
||||
|
||||
cls_name = config.pop("_class_name", None)
|
||||
if cls_name is None:
|
||||
raise ValueError(
|
||||
"Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported."
|
||||
)
|
||||
|
||||
config.pop("_diffusers_version", None)
|
||||
config.pop("_name_or_path", None)
|
||||
|
||||
server_args.model_paths["connectors"] = component_model_path
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||
|
||||
target_device = get_local_torch_device()
|
||||
default_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.dit_precision]
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
with set_default_torch_dtype(default_dtype), skip_init_modules():
|
||||
connector_cfg = SimpleNamespace(**config)
|
||||
model = model_cls(connector_cfg).to(
|
||||
device=target_device, dtype=default_dtype
|
||||
)
|
||||
|
||||
safetensors_list = _list_safetensors_files(component_model_path)
|
||||
if not safetensors_list:
|
||||
raise ValueError(f"No safetensors files found in {component_model_path}")
|
||||
if len(safetensors_list) != 1:
|
||||
raise ValueError(
|
||||
f"Found {len(safetensors_list)} safetensors files in {component_model_path}, expected 1"
|
||||
)
|
||||
|
||||
loaded = safetensors_load_file(safetensors_list[0])
|
||||
model.load_state_dict(loaded, strict=False)
|
||||
|
||||
return model.eval()
|
||||
|
||||
|
||||
class SchedulerLoader(ComponentLoader):
|
||||
"""Loader for scheduler."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user