The Lightricks/LTX-2.3-fp8 checkpoints store float8 weights with per-tensor weight_scale and input_scale companions (verified from the file headers: 1496 F8_E4M3 tensors, 2924 scale tensors). A plain dtype cast would silently corrupt every quantized layer, so the 2.3 assembly now detects the companions and raises with a pointer to the GGUF quants, which offer comparable fidelity through the supported path. Dequantizing the scaled fp8 layout is a possible follow-up.
560 lines
22 KiB
Python
560 lines
22 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
|
|
|
"""LTX-2.3 pipeline assembly for diffusers 0.39.
|
|
|
|
diffusers 0.39 ships every LTX-2.3 model class (the transformer's
|
|
``cross_attn_mod``/``gated_attn`` flags, the per-modality ``LTX2TextConnectors``,
|
|
the deeper 2.3 video VAE decoder, and ``LTX2VocoderWithBWE``) but its single-file
|
|
loader maps every LTX-2 checkpoint to the 2.0 config, so 2.3 checkpoints fail a
|
|
shape check at load. On top of that, the community transformer-only GGUFs carry
|
|
the DiT and the embedding connectors but NOT the text projections, VAEs, or
|
|
vocoder, which 2.3 moved out of the transformer.
|
|
|
|
This module assembles the full 2.3 pipeline:
|
|
|
|
- transformer: from the GGUF/single-file checkpoint via ``from_single_file`` with
|
|
the 2.3 config overrides (the loader merges signature-matching kwargs into the
|
|
fetched 2.0 config) and the ``prompt_adaln_single`` keys pre-renamed, since the
|
|
library converter does not know them.
|
|
- connectors: from the same checkpoint's ``*_embeddings_connector`` keys plus the
|
|
4 ``text_embedding_projection`` tensors, fetched from the companion file in
|
|
``unsloth/LTX-2.3-GGUF`` when the checkpoint does not bundle them.
|
|
- video VAE, audio VAE, vocoder: from the checkpoint when bundled (the official
|
|
Lightricks single files), else from the companion files in the same repo.
|
|
- scheduler, text encoder (Gemma3), tokenizer: from the LTX-2.0 base repo, which
|
|
2.3 shares.
|
|
|
|
Every config and key-rename table below mirrors diffusers' own
|
|
``scripts/convert_ltx2_to_diffusers.py`` (the authoritative 2.3 mapping, which
|
|
the library loader has not absorbed yet). The pipeline is assembled through the
|
|
constructor rather than ``from_pretrained`` because the vocoder class differs
|
|
from the base repo's pin (``LTX2VocoderWithBWE`` vs ``LTX2Vocoder``) and the
|
|
component type gate would reject it.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from pathlib import Path
|
|
from typing import Any, Optional
|
|
|
|
from loggers import get_logger
|
|
|
|
logger = get_logger(__name__)
|
|
|
|
# The companion files (text projections, VAEs incl. vocoder) live next to the
|
|
# quants in unsloth's GGUF repo; they are the official Lightricks weights split
|
|
# out of the combined checkpoint. Keyed by checkpoint variant.
|
|
LTX23_EXTRAS_REPO = "unsloth/LTX-2.3-GGUF"
|
|
_EXTRAS_TEXT_PROJ = "text_encoders/ltx-2.3-22b-{variant}_embeddings_connectors.safetensors"
|
|
_EXTRAS_VIDEO_VAE = "vae/ltx-2.3-22b-{variant}_video_vae.safetensors"
|
|
_EXTRAS_AUDIO_VAE = "vae/ltx-2.3-22b-{variant}_audio_vae.safetensors"
|
|
|
|
# ── configs + rename tables, verbatim from scripts/convert_ltx2_to_diffusers.py ──
|
|
|
|
# from_single_file config overrides on top of the base repo's 2.0 transformer
|
|
# config (version == "2.3" in get_ltx2_transformer_config).
|
|
LTX_2_3_TRANSFORMER_CONFIG_OVERRIDES: dict[str, Any] = {
|
|
"gated_attn": True,
|
|
"cross_attn_mod": True,
|
|
"audio_gated_attn": True,
|
|
"audio_cross_attn_mod": True,
|
|
"use_prompt_embeddings": False,
|
|
"perturbed_attn": True,
|
|
}
|
|
|
|
# Keys the library's 2.0-era converter does not know; renamed before handing the
|
|
# state dict to from_single_file. Ordered so the audio prefix matches first.
|
|
_TRANSFORMER_PRERENAME = (
|
|
("audio_prompt_adaln_single.", "audio_prompt_adaln."),
|
|
("prompt_adaln_single.", "prompt_adaln."),
|
|
)
|
|
|
|
_CONNECTOR_KEY_PREFIXES = (
|
|
"video_embeddings_connector",
|
|
"audio_embeddings_connector",
|
|
"transformer_1d_blocks",
|
|
"text_embedding_projection",
|
|
"connectors.",
|
|
"video_connector",
|
|
"audio_connector",
|
|
"text_proj_in",
|
|
)
|
|
|
|
_CONNECTORS_RENAME = {
|
|
"connectors.": "",
|
|
"video_embeddings_connector": "video_connector",
|
|
"audio_embeddings_connector": "audio_connector",
|
|
"transformer_1d_blocks": "transformer_blocks",
|
|
"text_embedding_projection.audio_aggregate_embed": "audio_text_proj_in",
|
|
"text_embedding_projection.video_aggregate_embed": "video_text_proj_in",
|
|
"q_norm": "norm_q",
|
|
"k_norm": "norm_k",
|
|
}
|
|
|
|
_CONNECTORS_CONFIG: dict[str, Any] = {
|
|
"caption_channels": 3840,
|
|
"text_proj_in_factor": 49,
|
|
"video_connector_num_attention_heads": 32,
|
|
"video_connector_attention_head_dim": 128,
|
|
"video_connector_num_layers": 8,
|
|
"video_connector_num_learnable_registers": 128,
|
|
"video_gated_attn": True,
|
|
"audio_connector_num_attention_heads": 32,
|
|
"audio_connector_attention_head_dim": 64,
|
|
"audio_connector_num_layers": 8,
|
|
"audio_connector_num_learnable_registers": 128,
|
|
"audio_gated_attn": True,
|
|
"connector_rope_base_seq_len": 4096,
|
|
"rope_theta": 10000.0,
|
|
"rope_double_precision": True,
|
|
"causal_temporal_positioning": False,
|
|
"rope_type": "split",
|
|
"per_modality_projections": True,
|
|
"video_hidden_dim": 4096,
|
|
"audio_hidden_dim": 2048,
|
|
"proj_bias": True,
|
|
}
|
|
|
|
_VIDEO_VAE_RENAME = {
|
|
# Encoder
|
|
"down_blocks.0": "down_blocks.0",
|
|
"down_blocks.1": "down_blocks.0.downsamplers.0",
|
|
"down_blocks.2": "down_blocks.1",
|
|
"down_blocks.3": "down_blocks.1.downsamplers.0",
|
|
"down_blocks.4": "down_blocks.2",
|
|
"down_blocks.5": "down_blocks.2.downsamplers.0",
|
|
"down_blocks.6": "down_blocks.3",
|
|
"down_blocks.7": "down_blocks.3.downsamplers.0",
|
|
"down_blocks.8": "mid_block",
|
|
# Decoder (2.3 adds up_blocks.7/8: a 4th decoder stage)
|
|
"up_blocks.0": "mid_block",
|
|
"up_blocks.1": "up_blocks.0.upsamplers.0",
|
|
"up_blocks.2": "up_blocks.0",
|
|
"up_blocks.3": "up_blocks.1.upsamplers.0",
|
|
"up_blocks.4": "up_blocks.1",
|
|
"up_blocks.5": "up_blocks.2.upsamplers.0",
|
|
"up_blocks.6": "up_blocks.2",
|
|
"up_blocks.7": "up_blocks.3.upsamplers.0",
|
|
"up_blocks.8": "up_blocks.3",
|
|
"last_time_embedder": "time_embedder",
|
|
"last_scale_shift_table": "scale_shift_table",
|
|
# Common
|
|
"res_blocks": "resnets",
|
|
"per_channel_statistics.mean-of-means": "latents_mean",
|
|
"per_channel_statistics.std-of-means": "latents_std",
|
|
}
|
|
|
|
_VIDEO_VAE_REMOVE_SUFFIXES = (
|
|
"per_channel_statistics.channel",
|
|
"per_channel_statistics.mean-of-stds",
|
|
)
|
|
|
|
_VIDEO_VAE_CONFIG: dict[str, Any] = {
|
|
"in_channels": 3,
|
|
"out_channels": 3,
|
|
"latent_channels": 128,
|
|
"block_out_channels": (256, 512, 1024, 1024),
|
|
"down_block_types": (
|
|
"LTX2VideoDownBlock3D",
|
|
"LTX2VideoDownBlock3D",
|
|
"LTX2VideoDownBlock3D",
|
|
"LTX2VideoDownBlock3D",
|
|
),
|
|
"decoder_block_out_channels": (256, 512, 512, 1024),
|
|
"layers_per_block": (4, 6, 4, 2, 2),
|
|
"decoder_layers_per_block": (4, 6, 4, 2, 2),
|
|
"spatio_temporal_scaling": (True, True, True, True),
|
|
"decoder_spatio_temporal_scaling": (True, True, True, True),
|
|
"decoder_inject_noise": (False, False, False, False, False),
|
|
"downsample_type": ("spatial", "temporal", "spatiotemporal", "spatiotemporal"),
|
|
"upsample_type": ("spatiotemporal", "spatiotemporal", "temporal", "spatial"),
|
|
"upsample_residual": (False, False, False, False),
|
|
"upsample_factor": (2, 2, 1, 2),
|
|
"timestep_conditioning": False,
|
|
"patch_size": 4,
|
|
"patch_size_t": 1,
|
|
"resnet_norm_eps": 1e-6,
|
|
"encoder_causal": True,
|
|
"decoder_causal": False,
|
|
"encoder_spatial_padding_mode": "zeros",
|
|
"decoder_spatial_padding_mode": "zeros",
|
|
"spatial_compression_ratio": 32,
|
|
"temporal_compression_ratio": 8,
|
|
}
|
|
|
|
_AUDIO_VAE_RENAME = {
|
|
"per_channel_statistics.mean-of-means": "latents_mean",
|
|
"per_channel_statistics.std-of-means": "latents_std",
|
|
}
|
|
|
|
# Same config as LTX-2.0 (upstream's comment); the weights are still 2.3-specific.
|
|
_AUDIO_VAE_CONFIG: dict[str, Any] = {
|
|
"base_channels": 128,
|
|
"output_channels": 2,
|
|
"ch_mult": (1, 2, 4),
|
|
"num_res_blocks": 2,
|
|
"attn_resolutions": None,
|
|
"in_channels": 2,
|
|
"resolution": 256,
|
|
"latent_channels": 8,
|
|
"norm_type": "pixel",
|
|
"causality_axis": "height",
|
|
"dropout": 0.0,
|
|
"mid_block_add_attention": False,
|
|
"sample_rate": 16000,
|
|
"mel_hop_length": 160,
|
|
"is_causal": True,
|
|
"mel_bins": 64,
|
|
"double_z": True,
|
|
}
|
|
|
|
_VOCODER_RENAME = {
|
|
"resblocks": "resnets",
|
|
"conv_pre": "conv_in",
|
|
"conv_post": "conv_out",
|
|
"act_post": "act_out",
|
|
"downsample.lowpass": "downsample",
|
|
}
|
|
|
|
_VOCODER_CONFIG: dict[str, Any] = {
|
|
"in_channels": 128,
|
|
"hidden_channels": 1536,
|
|
"out_channels": 2,
|
|
"upsample_kernel_sizes": [11, 4, 4, 4, 4, 4],
|
|
"upsample_factors": [5, 2, 2, 2, 2, 2],
|
|
"resnet_kernel_sizes": [3, 7, 11],
|
|
"resnet_dilations": [[1, 3, 5], [1, 3, 5], [1, 3, 5]],
|
|
"act_fn": "snakebeta",
|
|
"leaky_relu_negative_slope": 0.1,
|
|
"antialias": True,
|
|
"antialias_ratio": 2,
|
|
"antialias_kernel_size": 12,
|
|
"final_act_fn": None,
|
|
"final_bias": False,
|
|
"bwe_in_channels": 128,
|
|
"bwe_hidden_channels": 512,
|
|
"bwe_out_channels": 2,
|
|
"bwe_upsample_kernel_sizes": [12, 11, 4, 4, 4],
|
|
"bwe_upsample_factors": [6, 5, 2, 2, 2],
|
|
"bwe_resnet_kernel_sizes": [3, 7, 11],
|
|
"bwe_resnet_dilations": [[1, 3, 5], [1, 3, 5], [1, 3, 5]],
|
|
"bwe_act_fn": "snakebeta",
|
|
"bwe_leaky_relu_negative_slope": 0.1,
|
|
"bwe_antialias": True,
|
|
"bwe_antialias_ratio": 2,
|
|
"bwe_antialias_kernel_size": 12,
|
|
"bwe_final_act_fn": None,
|
|
"bwe_final_bias": False,
|
|
"filter_length": 512,
|
|
"hop_length": 80,
|
|
"window_length": 512,
|
|
"num_mel_channels": 64,
|
|
"input_sampling_rate": 16000,
|
|
"output_sampling_rate": 48000,
|
|
}
|
|
|
|
_DIT_PREFIX = "model.diffusion_model."
|
|
|
|
|
|
# ── checkpoint inspection ────────────────────────────────────────────────────
|
|
|
|
|
|
def read_checkpoint_header(checkpoint_path: Path | str) -> dict[str, tuple[int, ...]]:
|
|
"""Tensor name -> shape from the checkpoint HEADER only (no weight data).
|
|
|
|
GGUF shapes come back in GGML (reversed) order; callers must not assume a
|
|
dimension position and should membership-test instead.
|
|
"""
|
|
names_shapes: dict[str, tuple[int, ...]] = {}
|
|
path = str(checkpoint_path)
|
|
if path.lower().endswith(".gguf"):
|
|
from gguf import GGUFReader
|
|
|
|
for tensor in GGUFReader(path).tensors:
|
|
names_shapes[str(tensor.name)] = tuple(int(x) for x in tensor.shape)
|
|
else:
|
|
from safetensors import safe_open
|
|
|
|
with safe_open(path, framework = "pt") as handle:
|
|
for name in handle.keys():
|
|
names_shapes[name] = tuple(handle.get_slice(name).get_shape())
|
|
return names_shapes
|
|
|
|
|
|
def is_ltx23_checkpoint(checkpoint_path: Path | str) -> bool:
|
|
"""True when the checkpoint carries the 9-row LTX-2.3 modulation tables.
|
|
|
|
2.0 checkpoints have 6-row per-block scale/shift tables; 2.3 widens them to
|
|
9 (the cross-attn adaln rows). An unreadable header returns False so the
|
|
caller falls back to the stock 2.0 load path and surfaces its own error.
|
|
"""
|
|
try:
|
|
header = read_checkpoint_header(checkpoint_path)
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.warning("video.ltx2_header_probe_failed: %s", exc)
|
|
return False
|
|
for name, shape in header.items():
|
|
if name.endswith("transformer_blocks.0.scale_shift_table"):
|
|
return 9 in shape
|
|
return False
|
|
|
|
|
|
# ── state-dict plumbing ──────────────────────────────────────────────────────
|
|
|
|
|
|
def _apply_rename(state: dict[str, Any], rename: dict[str, str]) -> dict[str, Any]:
|
|
out: dict[str, Any] = {}
|
|
for key, value in state.items():
|
|
new_key = key
|
|
for old, new in rename.items():
|
|
new_key = new_key.replace(old, new)
|
|
out[new_key] = value
|
|
return out
|
|
|
|
|
|
def _to_plain_dtype(state: dict[str, Any], torch_dtype: Any) -> dict[str, Any]:
|
|
"""Materialise every tensor as a plain torch tensor in torch_dtype.
|
|
|
|
GGUF-sourced tensors arrive as GGUFParameter (block-quantized); the small
|
|
non-DiT components run dense, so dequantize them here. Their fidelity is
|
|
whatever the GGUF holds -- identical numbers to what native GGUF runners use.
|
|
"""
|
|
import torch
|
|
|
|
try:
|
|
from diffusers.quantizers.gguf.utils import GGUFParameter, dequantize_gguf_tensor
|
|
except Exception: # noqa: BLE001 -- gguf support not installed; plain tensors only
|
|
GGUFParameter, dequantize_gguf_tensor = (), None
|
|
|
|
out: dict[str, Any] = {}
|
|
for key, value in state.items():
|
|
if dequantize_gguf_tensor is not None and isinstance(value, GGUFParameter):
|
|
value = dequantize_gguf_tensor(value)
|
|
out[key] = value.to(torch_dtype) if isinstance(value, torch.Tensor) else value
|
|
return out
|
|
|
|
|
|
def _split_checkpoint(state: dict[str, Any]) -> dict[str, dict[str, Any]]:
|
|
"""Partition a combined LTX checkpoint into per-component state dicts.
|
|
|
|
Handles both layouts: the official combined single file (``vae.*``,
|
|
``audio_vae.*``, ``vocoder.*``, ``model.diffusion_model.*`` plus top-level
|
|
``text_embedding_projection.*``) and transformer-only GGUFs (bare DiT +
|
|
connector keys, nothing else).
|
|
"""
|
|
groups: dict[str, dict[str, Any]] = {
|
|
"dit": {}, "connectors": {}, "vae": {}, "audio_vae": {}, "vocoder": {}
|
|
}
|
|
for key, value in state.items():
|
|
bare = key[len(_DIT_PREFIX):] if key.startswith(_DIT_PREFIX) else key
|
|
if bare.startswith("vae."):
|
|
groups["vae"][bare[len("vae."):]] = value
|
|
elif bare.startswith("audio_vae."):
|
|
groups["audio_vae"][bare[len("audio_vae."):]] = value
|
|
elif bare.startswith("vocoder."):
|
|
groups["vocoder"][bare[len("vocoder."):]] = value
|
|
elif bare.startswith(_CONNECTOR_KEY_PREFIXES):
|
|
groups["connectors"][bare] = value
|
|
else:
|
|
groups["dit"][bare] = value
|
|
return groups
|
|
|
|
|
|
def _load_extras_file(filename: str, hf_token: Optional[str]) -> dict[str, Any]:
|
|
from safetensors.torch import load_file
|
|
|
|
from utils.hf_xet_fallback import hf_hub_download_with_xet_fallback
|
|
|
|
path = hf_hub_download_with_xet_fallback(LTX23_EXTRAS_REPO, filename, hf_token)
|
|
return load_file(path)
|
|
|
|
|
|
def checkpoint_variant(checkpoint_path: Path | str) -> str:
|
|
"""Which companion-weight set a checkpoint pairs with ("dev"/"distilled").
|
|
|
|
The distilled-1.1 refresh only retrained the DiT, so it shares the distilled
|
|
companions.
|
|
"""
|
|
return "dev" if "dev" in Path(checkpoint_path).name.lower() else "distilled"
|
|
|
|
|
|
# ── component builders ───────────────────────────────────────────────────────
|
|
|
|
|
|
def _build_from_config(model_cls: Any, config: dict[str, Any], state: dict[str, Any],
|
|
rename: dict[str, str], torch_dtype: Any,
|
|
remove_suffixes: tuple[str, ...] = ()) -> Any:
|
|
from accelerate import init_empty_weights
|
|
|
|
state = _apply_rename(_to_plain_dtype(state, torch_dtype), rename)
|
|
for key in [k for k in state if k.endswith(remove_suffixes)] if remove_suffixes else []:
|
|
state.pop(key)
|
|
with init_empty_weights():
|
|
model = model_cls.from_config(config)
|
|
model.load_state_dict(state, strict = True, assign = True)
|
|
return model.to(torch_dtype)
|
|
|
|
|
|
def load_ltx23_transformer(dit_state: dict[str, Any], *, base_repo: str, torch_dtype: Any,
|
|
is_gguf: bool, hf_token: Optional[str]) -> Any:
|
|
import diffusers
|
|
from diffusers import LTX2VideoTransformer3DModel
|
|
|
|
# Pre-rename the 2.3-only keys the library converter does not know, then let
|
|
# from_single_file do the rest: it merges the config overrides below into the
|
|
# base repo's 2.0 transformer config and runs the stock 2.0 key conversion.
|
|
for old, new in _TRANSFORMER_PRERENAME:
|
|
for key in [k for k in dit_state if k.startswith(old)]:
|
|
dit_state[new + key[len(old):]] = dit_state.pop(key)
|
|
kwargs: dict[str, Any] = {
|
|
"config": base_repo,
|
|
"subfolder": "transformer",
|
|
"torch_dtype": torch_dtype,
|
|
"token": hf_token,
|
|
**LTX_2_3_TRANSFORMER_CONFIG_OVERRIDES,
|
|
}
|
|
if is_gguf:
|
|
kwargs["quantization_config"] = diffusers.GGUFQuantizationConfig(
|
|
compute_dtype = torch_dtype
|
|
)
|
|
return LTX2VideoTransformer3DModel.from_single_file(dit_state, **kwargs)
|
|
|
|
|
|
def load_ltx23_connectors(connector_state: dict[str, Any], *, variant: str,
|
|
torch_dtype: Any, hf_token: Optional[str]) -> Any:
|
|
from diffusers.pipelines.ltx2.connectors import LTX2TextConnectors
|
|
|
|
# Transformer-only checkpoints carry the connector stacks but not the huge
|
|
# per-modality text projections; fetch those from the companion file.
|
|
if not any(k.startswith("text_embedding_projection") for k in connector_state):
|
|
connector_state = dict(connector_state)
|
|
connector_state.update(_load_extras_file(
|
|
_EXTRAS_TEXT_PROJ.format(variant = variant), hf_token
|
|
))
|
|
return _build_from_config(
|
|
LTX2TextConnectors, _CONNECTORS_CONFIG, connector_state, _CONNECTORS_RENAME,
|
|
torch_dtype,
|
|
)
|
|
|
|
|
|
def load_ltx23_vae(vae_state: dict[str, Any], *, variant: str, torch_dtype: Any,
|
|
hf_token: Optional[str]) -> Any:
|
|
from diffusers import AutoencoderKLLTX2Video
|
|
|
|
if not vae_state:
|
|
vae_state = _load_extras_file(_EXTRAS_VIDEO_VAE.format(variant = variant), hf_token)
|
|
return _build_from_config(
|
|
AutoencoderKLLTX2Video, _VIDEO_VAE_CONFIG, vae_state, _VIDEO_VAE_RENAME,
|
|
torch_dtype, remove_suffixes = _VIDEO_VAE_REMOVE_SUFFIXES,
|
|
)
|
|
|
|
|
|
def load_ltx23_audio_vae_and_vocoder(
|
|
audio_vae_state: dict[str, Any], vocoder_state: dict[str, Any], *, variant: str,
|
|
torch_dtype: Any, hf_token: Optional[str],
|
|
) -> tuple[Any, Any]:
|
|
from diffusers import AutoencoderKLLTX2Audio
|
|
from diffusers.pipelines.ltx2.vocoder import LTX2VocoderWithBWE
|
|
|
|
if not audio_vae_state or not vocoder_state:
|
|
combined = _load_extras_file(_EXTRAS_AUDIO_VAE.format(variant = variant), hf_token)
|
|
audio_vae_state = {
|
|
k[len("audio_vae."):]: v for k, v in combined.items() if k.startswith("audio_vae.")
|
|
}
|
|
vocoder_state = {
|
|
k[len("vocoder."):]: v for k, v in combined.items() if k.startswith("vocoder.")
|
|
}
|
|
audio_vae = _build_from_config(
|
|
AutoencoderKLLTX2Audio, _AUDIO_VAE_CONFIG, audio_vae_state, _AUDIO_VAE_RENAME,
|
|
torch_dtype,
|
|
)
|
|
# The 2.3 vocoder is a composite (base vocoder + bandwidth-extension stack +
|
|
# mel STFT buffers); keys line up module-for-module after the renames.
|
|
vocoder_state = _apply_rename(_to_plain_dtype(vocoder_state, torch_dtype), _VOCODER_RENAME)
|
|
for key in [k for k in vocoder_state if ".ups." in k]:
|
|
vocoder_state[key.replace(".ups.", ".upsamplers.")] = vocoder_state.pop(key)
|
|
from accelerate import init_empty_weights
|
|
|
|
with init_empty_weights():
|
|
vocoder = LTX2VocoderWithBWE.from_config(_VOCODER_CONFIG)
|
|
vocoder.load_state_dict(vocoder_state, strict = True, assign = True)
|
|
return audio_vae, vocoder.to(torch_dtype)
|
|
|
|
|
|
# ── pipeline assembly ────────────────────────────────────────────────────────
|
|
|
|
|
|
def load_ltx23_pipeline(checkpoint_path: Path | str, *, base_repo: str, torch_dtype: Any,
|
|
is_gguf: bool, hf_token: Optional[str] = None) -> Any:
|
|
"""Full LTX-2.3 pipeline from a single-file/GGUF checkpoint.
|
|
|
|
Assembled per-component (constructor, not from_pretrained) because the base
|
|
repo's model_index pins LTX2Vocoder while 2.3 needs LTX2VocoderWithBWE, and
|
|
the from_pretrained component type gate rejects the substitution.
|
|
"""
|
|
import transformers
|
|
from diffusers import LTX2Pipeline
|
|
from diffusers.loaders.single_file_utils import load_single_file_checkpoint
|
|
|
|
variant = checkpoint_variant(checkpoint_path)
|
|
logger.info(
|
|
"video.ltx23_assembly: variant=%s gguf=%s extras=%s",
|
|
variant, is_gguf, LTX23_EXTRAS_REPO,
|
|
)
|
|
state = load_single_file_checkpoint(str(checkpoint_path))
|
|
groups = _split_checkpoint(state)
|
|
del state
|
|
|
|
# The Lightricks fp8 single files store SCALED float8 weights (per-tensor
|
|
# .weight_scale/.input_scale companions). Casting those without applying the
|
|
# scales silently corrupts every quantized layer, so refuse loudly. GGUF
|
|
# Q8_0 offers comparable fidelity at similar size through the supported path.
|
|
if any(k.endswith((".weight_scale", ".input_scale")) for k in groups["dit"]):
|
|
raise ValueError(
|
|
"This LTX checkpoint stores scaled fp8 weights, which this loader does "
|
|
"not dequantize yet. Use the GGUF quants from unsloth/LTX-2.3-GGUF "
|
|
"instead (Q8_0 for the highest fidelity) or the official bf16 checkpoint."
|
|
)
|
|
|
|
transformer = load_ltx23_transformer(
|
|
groups["dit"], base_repo = base_repo, torch_dtype = torch_dtype,
|
|
is_gguf = is_gguf, hf_token = hf_token,
|
|
)
|
|
connectors = load_ltx23_connectors(
|
|
groups["connectors"], variant = variant, torch_dtype = torch_dtype,
|
|
hf_token = hf_token,
|
|
)
|
|
vae = load_ltx23_vae(
|
|
groups["vae"], variant = variant, torch_dtype = torch_dtype, hf_token = hf_token
|
|
)
|
|
audio_vae, vocoder = load_ltx23_audio_vae_and_vocoder(
|
|
groups["audio_vae"], groups["vocoder"], variant = variant,
|
|
torch_dtype = torch_dtype, hf_token = hf_token,
|
|
)
|
|
|
|
# Shared 2.0/2.3 components from the base repo, resolved through model_index
|
|
# so class renames upstream break loudly here rather than silently drifting.
|
|
index = LTX2Pipeline.load_config(base_repo, token = hf_token)
|
|
|
|
def _sub(name: str, **extra: Any) -> Any:
|
|
library, class_name = index[name]
|
|
module = transformers if library == "transformers" else __import__("diffusers")
|
|
return getattr(module, class_name).from_pretrained(
|
|
base_repo, subfolder = name, token = hf_token, **extra
|
|
)
|
|
|
|
scheduler = _sub("scheduler")
|
|
tokenizer = _sub("tokenizer")
|
|
text_encoder = _sub("text_encoder", torch_dtype = torch_dtype)
|
|
|
|
return LTX2Pipeline(
|
|
scheduler = scheduler,
|
|
text_encoder = text_encoder,
|
|
tokenizer = tokenizer,
|
|
connectors = connectors,
|
|
transformer = transformer,
|
|
vae = vae,
|
|
audio_vae = audio_vae,
|
|
vocoder = vocoder,
|
|
)
|