[diffusion] model: LTX-2 (1/2) (#17495)

Co-authored-by: FlamingoPg <1106310035@qq.com>
This commit is contained in:
GMI Xiao Jin
2026-01-23 19:59:48 -08:00
committed by GitHub
parent 894928a951
commit 797a9811a2
9 changed files with 519 additions and 29 deletions

View File

@@ -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."""