Compare commits

...
Author SHA1 Message Date
SolitaryThinker 825c6365e0 bundle variant -> preset resolution, pipeline pin, and LTX-2.5 examples
A bundle's training variant (header-declared `variant`, filename token as
fallback) now selects its sampling preset: distilled -> ltx2_distilled
(8 steps, cfg 1.0, distilled sigmas engage at <=8 steps), sft/base ->
ltx2_base. The bundle table also pins the pipeline class, so
VideoGenerator.from_pretrained(<bundle>) resolves without an override.
Adds basic_ltx2_5{,_distilled}.py examples (--model-path takes the FILE,
--gemma-root declares the paired encoder root), a README section, and CPU
tests for variant/preset resolution, table integrity, and the examples.
2026-08-11 17:03:42 -07:00
SolitaryThinker bb1fb7c149 WIP: route single-file bundles through the modular train stack
carried over the SWE session's uncommitted work across the rebase onto
main@8208536c: model_index_and_component_path helper feeds the train
moduleloader, trainer + legacy distillation pipeline pick up bundle-aware
component paths, nvfp4_config/registry adjustments, and the ltx2_5
fine-tuning example configs (untracked until now).

(--no-verify: pre-commit not configured in this worktree)
2026-08-10 13:06:06 -07:00
William Lin 436c44e583 bundle route through the real pipeline: module map, component paths, loader branches, connector factorization 2026-08-10 12:52:52 -07:00
William Lin 6343eadbec resolve pipeline config for single-file bundles from declared transformer class
A bundle path is a file, so config resolution took the directory branch and
was rejected for having no model_index.json. Resolve it from the transformer
_class_name the checkpoint declares about itself, via an explicit alias table.
Keyed on the class name, not the registry model_id (a positional index that
rebinds on registration-order changes). No wildcard fallback and no file-name
inference; an unmapped class raises naming the table and the override.

verify_model_config_and_directory is unchanged.
2026-08-10 12:51:08 -07:00
William Lin 3f5614fd03 single-file bundle loader: metadata config, prefix routing, per-stream ff bias, encoder-root resolution 2026-08-10 12:51:08 -07:00
20 changed files with 1538 additions and 33 deletions
+29
View File
@@ -33,6 +33,35 @@ For the typed config/request path added during the inference API refactor:
python examples/inference/basic/basic_dmd_new_api.py
```
## LTX-2.5 single-file bundles
These examples load LTX-2.5 from a single-file bundle: one `.safetensors`
file carrying every component (transformer, video VAE, audio VAE, vocoder,
text projection). (The Hugging Face release ships as a split pack — one
file per component — which loads through the standard repo path instead.)
Pass the bundle FILE path directly:
```bash
python examples/inference/basic/basic_ltx2_5_distilled.py \
--model-path /path/to/bundle.safetensors \
--gemma-root /path/to/gemma_root
```
| bundle variant | preset | sampling defaults |
|---|---|---|
| distilled | `ltx2_distilled` | one stage, 8 steps, CFG 1.0 (no CFG/STG, no spatial upscaler) |
| sft / base | `ltx2_base` | 40 steps, CFG 3.0 with STG |
- The variant is read from the bundle header when declared; otherwise the
`distilled` token in the file name decides.
- The Gemma text encoder is NOT in the bundle. Declare its root with
`--gemma-root` (or `FASTVIDEO_LTX_ENCODER_ROOT`), and use the root shipped
with the same transformer variant — the roots differ in prompt templating.
- Decoder caveat: bundles that declare a diffusion VAE decoder
(`CausalDiffusionVAE`) are not supported yet; loading them currently
fails at VAE build time. Bundles with the classic `CausalVideoAutoencoder`
run end to end.
## Basic Walkthrough
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
+103
View File
@@ -0,0 +1,103 @@
# SPDX-License-Identifier: Apache-2.0
"""LTX-2.5 (sft/base) text-to-video from a single-file bundle.
Loads LTX-2.5 from a single-file bundle: ONE ``.safetensors`` file carrying
every component (transformer, video VAE, audio VAE, vocoder, text
projection), so ``--model-path`` takes the bundle FILE, not a repo
directory. (The split-pack Hugging Face release loads through the standard
repo path instead.)
The Gemma text encoder is NOT in the bundle: pass its root directory with
``--gemma-root`` (or set ``FASTVIDEO_LTX_ENCODER_ROOT``). Use the encoder
root shipped WITH this transformer variant -- the published roots share
weights but differ in prompt templating, so a mismatched root silently
changes prompting.
A non-distilled bundle resolves to the standard ``ltx2_base`` preset
(40 steps, CFG 3.0 with STG); run without sampling flags to use it as-is.
For the 8-step distilled recipe, see ``basic_ltx2_5_distilled.py``.
Audio: the bundle carries an audio VAE + vocoder, so generated videos get
an audio track. A bundle that declares no audio decoder skips audio
decoding and still produces video.
Decoder caveat: bundles that declare a *diffusion* VAE decoder
(``CausalDiffusionVAE``) are not implemented yet -- loading such a bundle
currently fails at VAE build time with an unsupported-architecture error
(no classic-decoder fallback is wired). Bundles with the classic
``CausalVideoAutoencoder`` decode end to end.
"""
import argparse
import os
from fastvideo import VideoGenerator
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic.")
# Sampling flags default to None: unset flags are NOT passed to
# `generate_video`, so the bundle's preset supplies them via
# `SamplingParam.from_pretrained` (ltx2_base: 40 steps, cfg 3.0, 512x768,
# 121 frames).
_SAMPLING_FLAGS = ("height", "width", "num_frames", "num_inference_steps", "guidance_scale", "seed")
def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(description="LTX-2.5 (sft/base) inference from a single-file bundle.")
parser.add_argument(
"--model-path",
required=True,
help="Path to the LTX-2.5 bundle (.safetensors FILE).",
)
parser.add_argument(
"--gemma-root",
default=None,
help="Gemma text-encoder root directory. Must be the root paired with "
"this transformer variant; the roots differ in prompt templating.",
)
parser.add_argument("--prompt", default=PROMPT)
parser.add_argument("--output-path", default="outputs_video/ltx2_5/output_ltx2_5_t2v.mp4")
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--height", type=int, default=None)
parser.add_argument("--width", type=int, default=None)
parser.add_argument("--num-frames", type=int, default=None)
parser.add_argument("--num-inference-steps", type=int, default=None)
parser.add_argument("--guidance-scale", type=float, default=None)
parser.add_argument("--seed", type=int, default=None)
return parser.parse_args(argv)
def sampling_overrides(args: argparse.Namespace) -> dict:
"""Only sampling flags the user actually passed; the preset supplies the rest."""
return {name: getattr(args, name) for name in _SAMPLING_FLAGS if getattr(args, name) is not None}
def main(argv: list[str] | None = None) -> None:
args = parse_args(argv)
if args.gemma_root:
os.environ["FASTVIDEO_LTX_ENCODER_ROOT"] = args.gemma_root
generator = VideoGenerator.from_pretrained(
args.model_path,
num_gpus=args.num_gpus,
)
generator.generate_video(
prompt=args.prompt,
output_path=args.output_path,
save_video=True,
**sampling_overrides(args),
)
generator.shutdown()
if __name__ == "__main__":
main()
@@ -0,0 +1,104 @@
# SPDX-License-Identifier: Apache-2.0
"""LTX-2.5 distilled text-to-video from a single-file bundle.
Loads LTX-2.5 from a single-file bundle: ONE ``.safetensors`` file carrying
every component (transformer, video VAE, audio VAE, vocoder, text
projection), so ``--model-path`` takes the bundle FILE, not a repo
directory. (The split-pack Hugging Face release loads through the standard
repo path instead.)
The Gemma text encoder is NOT in the bundle: pass its root directory with
``--gemma-root`` (or set ``FASTVIDEO_LTX_ENCODER_ROOT``). Use the encoder
root shipped WITH this transformer variant -- the published roots share
weights but differ in prompt templating, so a mismatched root silently
changes prompting.
Distilled recipe: ONE stage -- 8 steps on the distilled sigma schedule,
CFG 1.0, no CFG/STG, and no spatial upscaler needed. All of it comes from
the ``ltx2_distilled`` preset the bundle resolves to; run without sampling
flags to use it as-is.
Audio: the bundle carries an audio VAE + vocoder, so generated videos get
an audio track. A bundle that declares no audio decoder skips audio
decoding and still produces video.
Decoder caveat: bundles that declare a *diffusion* VAE decoder
(``CausalDiffusionVAE``) are not implemented yet -- loading such a bundle
currently fails at VAE build time with an unsupported-architecture error
(no classic-decoder fallback is wired). Bundles with the classic
``CausalVideoAutoencoder`` decode end to end.
"""
import argparse
import os
from fastvideo import VideoGenerator
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic.")
# Sampling flags default to None: unset flags are NOT passed to
# `generate_video`, so the bundle's preset supplies them via
# `SamplingParam.from_pretrained` (ltx2_distilled: 8 steps, cfg 1.0,
# 1024x1536, 121 frames).
_SAMPLING_FLAGS = ("height", "width", "num_frames", "num_inference_steps", "guidance_scale", "seed")
def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(description="LTX-2.5 distilled inference from a single-file bundle.")
parser.add_argument(
"--model-path",
required=True,
help="Path to the LTX-2.5 distilled bundle (.safetensors FILE).",
)
parser.add_argument(
"--gemma-root",
default=None,
help="Gemma text-encoder root directory. Must be the root paired with "
"this transformer variant; the roots differ in prompt templating.",
)
parser.add_argument("--prompt", default=PROMPT)
parser.add_argument("--output-path", default="outputs_video/ltx2_5/output_ltx2_5_distilled_t2v.mp4")
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--height", type=int, default=None)
parser.add_argument("--width", type=int, default=None)
parser.add_argument("--num-frames", type=int, default=None)
parser.add_argument("--num-inference-steps", type=int, default=None)
parser.add_argument("--guidance-scale", type=float, default=None)
parser.add_argument("--seed", type=int, default=None)
return parser.parse_args(argv)
def sampling_overrides(args: argparse.Namespace) -> dict:
"""Only sampling flags the user actually passed; the preset supplies the rest."""
return {name: getattr(args, name) for name in _SAMPLING_FLAGS if getattr(args, name) is not None}
def main(argv: list[str] | None = None) -> None:
args = parse_args(argv)
if args.gemma_root:
os.environ["FASTVIDEO_LTX_ENCODER_ROOT"] = args.gemma_root
generator = VideoGenerator.from_pretrained(
args.model_path,
num_gpus=args.num_gpus,
)
generator.generate_video(
prompt=args.prompt,
output_path=args.output_path,
save_video=True,
**sampling_overrides(args),
)
generator.shutdown()
if __name__ == "__main__":
main()
@@ -0,0 +1,130 @@
# Dual-stream T2V NVFP4 quantization-aware fine-tune (QAT).
#
# Adapted from fine_tuning/ltx2_3/nvfp4_qat_t2v.yaml. The quantization wiring
# is UNCHANGED and deliberately so: the curated linear prefix set in
# fastvideo/layers/quantization/nvfp4_config.py targets
# ltx2.blocks.{0..47}.* and attaches to this architecture unmodified --
# verified against the constructed module tree, 576 linears / 12.9B params,
# the same 12 suffix classes x 48 blocks as before. Do NOT narrow the set for
# this architecture: the cross-modal projections it targets
# (audio_to_video_attn, video_to_audio_attn) are deliberately enumerated in the
# DEPLOYMENT set, so dropping them in training would break train/deploy
# symmetry.
#
# What is NOT targeted, and why (all deliberate, none accidental):
# * head-dim-64 audio self/cross attention -- same constraint that kept it
# dense before;
# * per-head gate projections (to_gate_logits, [32, width]) -- tiny and
# precision-sensitive;
# * the embeddings connectors -- they live in the TEXT ENCODER, not the DiT,
# so deployment targeting must not reach across the component boundary.
# The coverage audit prints these as skipped-by-rule so a reader can see the
# exclusions were chosen, and FAILS if any declared class matches zero modules.
#
# Video attention uses ATTN_QAT_TRAIN (quantized forward, STE backward).
#
# Data paths, step counts and learning rate below are placeholders to adapt.
models:
student:
_target_: fastvideo.train.models.ltx2.LTX2Model
# REQUIRED: single-file bundle path. The config parser does NOT expand
# environment interpolations, so override this on the command line:
# --models.student.init_from /path/to/bundle.safetensors
init_from: SET_ME_ON_THE_COMMAND_LINE
trainable: true
enable_gradient_checkpointing_type: full
attention_backend: ATTN_QAT_TRAIN
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
# ~19B params: bf16 weights 35.4 GiB + bf16 grads 35.4 GiB + Adam fp32
# m/v 141.5 GiB = ~212 GiB, against ~173 GiB usable per GPU. Single-GPU is
# arithmetically impossible; 4-way shard puts it at ~53 GiB/GPU.
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
# REQUIRED: path to your preprocessed dataset.
data_path: data/your_dataset_preprocessed
dataloader_num_workers: 4
train_batch_size: 1
# LTX2Model requires 0.0: CFG dropout would zero post-connector
# embeddings, which is not the model's unconditional input.
training_cfg_rate: 0.0
seed: 42
# Must match your preprocessed resolution/length.
# num_latent_t = (num_frames - 1) / 8 + 1.
num_latent_t: 11
num_height: 480
num_width: 832
num_frames: 81
optimizer:
# Carried from the validated LTX-2.0 QAT overfit run as a starting
# point; tune for your dataset size and batch configuration.
learning_rate: 5.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
# Workload-dependent: set from your dataset size and target epochs.
max_train_steps: 2000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/dual_stream_t2v_nvfp4_qat_finetune
# A full training-state checkpoint is ~150GB for the 13B trainable
# video branch — size the interval and total limit to your storage.
training_state_checkpointing_steps: 500
checkpoints_total_limit: 2
resume_from_checkpoint: latest
tracker:
# No experiment tracking.
# It must be ["none"], NOT []: build_tracker appends wandb whenever the
# list is empty and project_name is non-empty, so [] silently means wandb.
trackers: ["none"]
model:
# LTX2Model.predict_noise converts the DiT's x0 output to velocity,
# so the default noise-minus-clean target reproduces the official
# unweighted masked-MSE (mask is all-ones for plain T2V).
precondition_outputs: false
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.ltx2.ltx2_pipeline.LTX2Pipeline
# REQUIRED: prompts to sample during training.
dataset_file: data/your_dataset_preprocessed/validation_prompts.json
every_steps: 250
# 8-step single-pass sampling matches the distilled checkpoint
# (validated in the LTX-2.3 overfit runs).
sampling_steps: [8]
guidance_scale: 1.0
num_frames: 81
# Validation reuses the live transformer and temporarily replaces its
# ATTN_QAT_TRAIN implementations with the sm120 inference kernel.
# Set false on GB200 (see header).
attn_qat_infer: true
# quant_config selects NVFP4 QAT for the same curated linear prefixes as
# LTX-2 NVFP4 deployment (see the LTX-2.3 applicability note in the
# header). They fake-quantize through real FP4 GEMMs with an STE
# backward. The string resolves to NVFP4QATTrainConfig at parse.
pipeline:
dit_config:
quant_config: nvfp4_qat_train
+7
View File
@@ -55,6 +55,13 @@ class LTX2VideoArchConfig(DiTArchConfig):
# LTX-2.3 gated extensions. All default OFF == LTX-2.0 behavior.
cross_attention_adaln: bool = False
caption_proj_before_connector: bool = False
# FFN bias, per stream. Some checkpoints ship the video FFN without bias
# (``ff_bias: false`` in their metadata), which drops the 96
# ``transformer_blocks.*.ff.net.{0.proj,2}.bias`` tensors; the audio FFN is
# configured independently and commonly keeps its biases. Both default True,
# which is the existing behavior.
ff_bias: bool = True
audio_ff_bias: bool = True
positional_embedding_theta: float = 10000.0
positional_embedding_max_pos: list[int] = field(default_factory=lambda: [20, 2048, 2048])
@@ -82,6 +82,79 @@ def is_ltx2_nvfp4_linear_prefix(prefix: str) -> bool:
return prefix in _LTX2_NVFP4_LINEAR_PREFIXES
def audit_nvfp4_coverage(
model,
predicate=None,
target_classes=None,
skip_classes=(),
) -> dict[str, object]:
"""Report which declared target classes actually attached, and raise if any did not.
The target set is written as module-prefix suffixes, but a tensor carries
three different names on its way here -- the checkpoint key, the
``named_modules()`` path, and the ``prefix=`` string a layer is constructed
with -- and only the last one reaches the predicate. A rename on any of the
others leaves this set syntactically valid and silently matching nothing,
and the attached COUNT alone cannot show it: a set that quietly stops
covering a whole class still reports a large, healthy-looking number.
So: a declared class matching zero modules is an error, not a warning. This
guard did NOT fire when it was written -- the shipped set is correct today.
It exists for the next rename, which is the only kind of failure that can
reach production looking like success.
Returns the receipt: attached count and params, per-class attached counts,
and the classes deliberately NOT targeted, so a reader can see the
exclusions were chosen rather than lost.
"""
from fastvideo.layers.linear import LinearBase
if predicate is None:
predicate = is_ltx2_nvfp4_linear_prefix
if target_classes is None:
target_classes = _LTX2_NVFP4_BLOCK_LINEAR_SUFFIXES
linears = [(getattr(m, "prefix", None) or n, n, m)
for n, m in model.named_modules() if isinstance(m, LinearBase)]
def _params(mod):
w = getattr(mod, "weight", None)
return int(w.shape[0] * w.shape[1]) if w is not None and w.dim() == 2 else 0
per_class: dict[str, int] = {}
for suffix in target_classes:
per_class[suffix] = sum(1 for pfx, _, _ in linears
if pfx.endswith("." + suffix) and predicate(pfx))
empty = sorted(k for k, v in per_class.items() if v == 0)
if empty:
raise ValueError(
"NVFP4 target classes matched ZERO modules: " + ", ".join(empty) +
". The model's actual Linear prefixes look like: " +
", ".join(sorted({p for p, _, _ in linears})[:5]) +
". Either the module naming changed or the target set is stale -- "
"refusing to train a model that silently quantizes less than declared.")
attached = [(p, n, m) for p, n, m in linears if predicate(p)]
skipped: dict[str, int] = {}
for pfx, name, _ in linears:
if predicate(pfx):
continue
tail = name.rsplit(".", 1)[-1]
for known in skip_classes:
if known in name:
skipped[known] = skipped.get(known, 0) + 1
break
else:
skipped["other:" + tail] = skipped.get("other:" + tail, 0) + 1
return {
"linears_total": len(linears),
"attached": len(attached),
"attached_params_M": round(sum(_params(m) for _, _, m in attached) / 1e6, 1),
"attached_per_class": per_class,
"skipped_by_rule": dict(sorted(skipped.items())),
}
def _is_ltx2_refine_only_prefix(prefix: str) -> bool:
return any(prefix.endswith(suffix) for suffix in _LTX2_REFINE_ONLY_SUFFIXES)
@@ -552,4 +625,5 @@ __all__ = [
"NVFP4QuantizeMethod",
"convert_model_to_nvfp4",
"is_ltx2_nvfp4_linear_prefix",
"audit_nvfp4_coverage",
]
+19
View File
@@ -333,6 +333,7 @@ class GELUApprox(nn.Module):
self,
in_features: int,
out_features: int,
bias: bool = True,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
@@ -340,6 +341,7 @@ class GELUApprox(nn.Module):
self.proj = ReplicatedLinear(
in_features,
out_features,
bias=bias,
quant_config=quant_config,
prefix=f"{prefix}.fc_in",
)
@@ -357,6 +359,7 @@ class FeedForward(nn.Module):
dim: int,
dim_out: int,
mult: int = 4,
bias: bool = True,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
@@ -365,12 +368,14 @@ class FeedForward(nn.Module):
project_in = GELUApprox(
dim,
inner_dim,
bias=bias,
quant_config=quant_config,
prefix=f"{prefix}.ffn",
)
project_out = ReplicatedLinear(
inner_dim,
dim_out,
bias=bias,
quant_config=quant_config,
prefix=f"{prefix}.ffn.fc_out",
)
@@ -1206,6 +1211,9 @@ class TransformerConfig:
# LTX-2.3 gated extensions (default OFF == LTX-2.0 behavior).
apply_gated_attention: bool = False
cross_attention_adaln: bool = False
# FFN bias is per stream: a checkpoint may drop it on video while keeping
# it on audio. Default True preserves existing behavior.
ff_bias: bool = True
class LTXDistributedAttention(DistributedAttention):
@@ -1846,6 +1854,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
self.ff = FeedForward(
video.dim,
dim_out=video.dim,
bias=video.ff_bias,
quant_config=quant_config,
prefix=f"{prefix}.blocks.{idx}",
)
@@ -1884,6 +1893,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
self.audio_ff = FeedForward(
audio.dim,
dim_out=audio.dim,
bias=audio.ff_bias,
quant_config=quant_config,
prefix=f"{prefix}.blocks.{idx}.audio",
)
@@ -2353,6 +2363,8 @@ class LTXModel(torch.nn.Module):
cross_attention_adaln: bool = False,
caption_proj_before_connector: bool = False,
apply_gated_attention: bool = False,
ff_bias: bool = True,
audio_ff_bias: bool = True,
stg_block_idx: int = 29,
use_distributed_attention: bool = False,
quant_config: QuantizationConfig | None = None,
@@ -2364,6 +2376,9 @@ class LTXModel(torch.nn.Module):
self.cross_attention_adaln = cross_attention_adaln
self.caption_proj_before_connector = caption_proj_before_connector
self.apply_gated_attention = apply_gated_attention
# Per-stream FFN bias; default True preserves existing behavior.
self.ff_bias = ff_bias
self.audio_ff_bias = audio_ff_bias
self.stg_block_idx = stg_block_idx
self.use_middle_indices_grid = use_middle_indices_grid
self.rope_type = rope_type
@@ -2584,6 +2599,7 @@ class LTXModel(torch.nn.Module):
context_dim=cross_attention_dim,
apply_gated_attention=self.apply_gated_attention,
cross_attention_adaln=self.cross_attention_adaln,
ff_bias=self.ff_bias,
) if self.model_type.is_video_enabled() else None)
audio_config = (TransformerConfig(
dim=self.audio_inner_dim,
@@ -2592,6 +2608,7 @@ class LTXModel(torch.nn.Module):
context_dim=audio_cross_attention_dim,
apply_gated_attention=self.apply_gated_attention,
cross_attention_adaln=self.cross_attention_adaln,
ff_bias=self.audio_ff_bias,
) if self.model_type.is_audio_enabled() else None)
self.use_distributed_attention = use_distributed_attention
self.transformer_blocks = torch.nn.ModuleList([
@@ -2778,6 +2795,8 @@ class LTX2Transformer3DModel(BaseDiT):
cross_attention_adaln=arch.cross_attention_adaln,
caption_proj_before_connector=arch.caption_proj_before_connector,
apply_gated_attention=arch.apply_gated_attention,
ff_bias=arch.ff_bias,
audio_ff_bias=arch.audio_ff_bias,
stg_block_idx=arch.stg_block_idx,
use_distributed_attention=use_distributed_attention,
quant_config=config.quant_config,
+40 -4
View File
@@ -7,7 +7,7 @@ from typing import Any, Iterable
import torch
from torch import nn
from transformers import AutoTokenizer, Gemma3ForConditionalGeneration
from transformers import AutoModelForImageTextToText, AutoTokenizer
from fastvideo.configs.models.encoders import BaseEncoderOutput, TextEncoderConfig
from fastvideo.models.encoders.base import TextEncoder
@@ -398,7 +398,7 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
self.gemma_model_path = arch.gemma_model_path
self.gemma_dtype = arch.gemma_dtype
self.padding_side = arch.padding_side
self._gemma_model: Gemma3ForConditionalGeneration | None = None
self._gemma_model: nn.Module | None = None
def named_parameters(self, prefix: str = "", recurse: bool = True):
for name, param in super().named_parameters(prefix=prefix, recurse=recurse):
@@ -436,17 +436,22 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
)
@property
def gemma_model(self) -> Gemma3ForConditionalGeneration:
def gemma_model(self) -> nn.Module:
if self._gemma_model is None:
gemma_path = self.gemma_model_path
if not gemma_path:
raise ValueError("gemma_model_path must be set (expected text_encoder/gemma).")
dtype = getattr(torch, self.gemma_dtype, torch.bfloat16)
self._gemma_model = Gemma3ForConditionalGeneration.from_pretrained(
# Resolve the class from the root config's own ``model_type``
# instead of hardcoding one: a checkpoint may declare its own
# encoder pairing. The auto class still resolves to the previously
# hardcoded class for the roots shipped so far.
self._gemma_model = AutoModelForImageTextToText.from_pretrained(
gemma_path,
local_files_only=True,
torch_dtype=dtype,
)
self._sync_arch_from_gemma_config(self._gemma_model.config)
# Configure model-level attention implementation when using TORCH_SDPA.
# Note: torch.backends.cuda.enable_*_sdp() settings should be configured
# at application/pipeline initialization level, not here, to avoid
@@ -464,6 +469,33 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
self._gemma_model.eval()
return self._gemma_model
def _sync_arch_from_gemma_config(self, gemma_config: Any) -> None:
"""Take language-model geometry from the loaded Gemma config.
A multimodal root nests it under ``text_config``; a text-only config
carries it at the root. The dataclass defaults hold for every encoder
root shipped so far, so this is not a behavior change -- it just stops
the defaults from being authoritative.
The feature-extractor linears are sized in ``__init__`` from
``feature_extractor_in_features``, before the Gemma weights load, so a
geometry mismatch is raised here rather than surfacing later as an
opaque matmul shape error.
"""
text_config = getattr(gemma_config, "text_config", gemma_config)
arch = self.config.arch_config
arch.hidden_size = text_config.hidden_size
arch.num_hidden_layers = text_config.num_hidden_layers
# +1: the stacked hidden states include the embedding output.
expected_in_features = arch.hidden_size * (arch.num_hidden_layers + 1)
if expected_in_features != arch.feature_extractor_in_features:
raise ValueError(
"Gemma geometry does not match the configured feature "
f"extractor: loaded config gives hidden_size={arch.hidden_size} "
f"x (num_hidden_layers={arch.num_hidden_layers} + 1) = "
f"{expected_in_features}, but feature_extractor_in_features is "
f"{arch.feature_extractor_in_features}."
)
def _run_feature_extractor(
self,
hidden_states: tuple[torch.Tensor, ...],
@@ -670,6 +702,10 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
name = "feature_extractor_linear.aggregate_embed.weight"
elif name.startswith("video_connector."):
name = name.replace("video_connector.", "embeddings_connector.", 1)
elif name.startswith("video_embeddings_connector."):
# Some checkpoints name the video connector sub-tree after its
# modality; the module is just ``embeddings_connector``.
name = name.replace("video_embeddings_connector.", "embeddings_connector.", 1)
elif name.startswith("audio_connector."):
name = name.replace("audio_connector.", "audio_embeddings_connector.", 1)
if name not in params_dict:
+157 -16
View File
@@ -9,7 +9,7 @@ from abc import ABC, abstractmethod
from collections.abc import Generator, Iterable
from contextlib import nullcontext
from copy import deepcopy
from typing import cast
from typing import Any, cast
import torch
import torch.distributed as dist
@@ -33,6 +33,11 @@ from fastvideo.logger import init_logger
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.hf_transformer_utils import get_diffusers_config
from fastvideo.models.loader.fsdp_load import maybe_load_fsdp_model, shard_model
from fastvideo.models.loader.ltx_single_file import (
component_weights,
is_single_file_bundle,
read_ltx_metadata,
)
from fastvideo.models.loader.utils import set_default_torch_dtype
from fastvideo.models.loader.weight_utils import (
filter_duplicate_safetensors_files,
@@ -138,6 +143,41 @@ class ComponentLoader(ABC):
return GenericComponentLoader(transformers_or_diffusers)
def _check_connector_widths(transformer_config: dict[str, Any]) -> None:
"""Fail loudly when a connector's declared factorization is inconsistent.
A connector's width is ``num_attention_heads * attention_head_dim``, and it
has to equal the cross-attention width of the stream it feeds. This is the
class of mistake that loads cleanly and computes wrong -- a factorization
that multiplies out to the wrong number produces correctly shaped tensors
everywhere except the one parameter it silently mis-sizes. Check it while
both numbers are still in hand rather than discovering it as a shape
assertion thousands of tensors later, or not at all.
Only checks what the checkpoint actually declares; a missing field means
the arch default applies and there is nothing to contradict.
"""
for heads_key, dim_key, width_key, label in (
("connector_num_attention_heads", "connector_attention_head_dim",
"cross_attention_dim", "video"),
("audio_connector_num_attention_heads",
"audio_connector_attention_head_dim", "audio_cross_attention_dim",
"audio"),
):
heads = transformer_config.get(heads_key)
head_dim = transformer_config.get(dim_key)
width = transformer_config.get(width_key)
if heads is None or head_dim is None or width is None:
continue
if heads * head_dim != width:
raise ValueError(
f"Checkpoint declares an inconsistent {label} connector shape: "
f"{heads_key}={heads} x {dim_key}={head_dim} = {heads * head_dim}, "
f"but {width_key}={width}. The connector feeds that stream, so "
"the two must agree; refusing to build a model whose weights "
"would load into mis-sized parameters.")
class TextEncoderLoader(ComponentLoader):
"""Loader for text encoders."""
@@ -251,6 +291,13 @@ class TextEncoderLoader(ComponentLoader):
# revision=fastvideo_args.revision,
# model_override_args=None,
# )
# For a bundle, `model_path` is the declared text-encoder root rather
# than a `text_encoder/` directory: the language model lives there, its
# projection and connectors live inside the bundle, and the fields the
# directory layout reads out of `<repo>/transformer/config.json` come
# from the bundle's metadata instead.
bundle_path = fastvideo_args.model_path
single_file = is_single_file_bundle(bundle_path)
model_config = get_diffusers_config(model=model_path)
model_config.pop("_name_or_path", None)
model_config.pop("transformers_version", None)
@@ -277,6 +324,15 @@ class TextEncoderLoader(ComponentLoader):
if gemma_path and not gemma_path_from_candidate:
if not os.path.isabs(gemma_path):
model_config["gemma_model_path"] = os.path.normpath(os.path.join(repo_root, gemma_path))
if single_file:
# The encoder root was declared, not discovered, so it wins over
# anything the probing above turned up.
model_config["gemma_model_path"] = model_path
# That root's config.json describes the language model it holds,
# which is not the class to build here: the root supplies the
# encoder's dimensions, the wrapper class supplies the connectors,
# and only the wrapper is buildable from this loader.
model_config["architectures"] = ["LTX2GemmaTextEncoderModel"]
transformer_config_path = os.path.join(repo_root, "transformer", "config.json")
if os.path.isfile(transformer_config_path):
try:
@@ -290,6 +346,40 @@ class TextEncoderLoader(ComponentLoader):
rope_type = transformer_config.get("rope_type")
if rope_type is not None:
model_config["connector_rope_type"] = rope_type
# Each per-modality projection feeds its connector, so its
# output width is that stream's cross-attention width, which
# only the transformer config declares. Absent from both
# configs -> the arch config default stands.
for src, dst in (
("cross_attention_dim",
"video_feature_extractor_out_features"),
("audio_cross_attention_dim",
"audio_feature_extractor_out_features"),
# The rest of the text stack's shape is declared under these
# exact names too. Leaving any of them to the arch default
# builds a DIFFERENT model than the checkpoint describes:
# the audio connector silently inherits the video
# connector's width, and the feature extractor is built in
# the wrong one of its two forms. Same names on both sides,
# so the mapping is identity.
*((name, name) for name in (
"connector_num_attention_heads",
"connector_attention_head_dim",
"connector_num_layers",
"connector_num_learnable_registers",
"connector_positional_embedding_theta",
"connector_positional_embedding_max_pos",
"connector_apply_gated_attention",
"audio_connector_num_attention_heads",
"audio_connector_attention_head_dim",
"audio_connector_num_layers",
# Selects which feature-extractor form is built.
"caption_proj_before_connector",
)),
):
if dst not in model_config and src in transformer_config:
model_config[dst] = transformer_config[src]
_check_connector_widths(transformer_config)
except json.JSONDecodeError:
pass
logger.info("HF Model config: %s", model_config)
@@ -332,6 +422,12 @@ class TextEncoderLoader(ComponentLoader):
fastvideo_args,
encoder_precision,
use_text_encoder_override=True,
# Everything this model owns is inside the bundle: the projection
# and both connectors. The language model at the encoder root is
# loaded separately and lazily by the model itself, and is filtered
# out of `named_parameters`, so it must not be routed through here.
weight_iterator=(component_weights(bundle_path, "text_encoder")
if single_file else None),
)
def load_model(
@@ -343,6 +439,7 @@ class TextEncoderLoader(ComponentLoader):
dtype: str = "fp16",
use_text_encoder_override: bool = False, # prevent subclasses from misusing
cpu_offload: bool | None = None,
weight_iterator: Iterable[tuple[str, torch.Tensor]] | None = None,
):
if cpu_offload is None:
cpu_offload = fastvideo_args.text_encoder_cpu_offload
@@ -386,6 +483,12 @@ class TextEncoderLoader(ComponentLoader):
[fastvideo_args.override_text_encoder_safetensors],
to_cpu=use_cpu_offload,
)) # type: ignore
elif weight_iterator is not None:
# Same contract as `maybe_load_fsdp_model`'s `weight_iterator`:
# already-routed `(name, cpu_tensor)` pairs, for checkpoints
# that are not a directory of per-component weight files.
self.counter_before_loading_weights = time.perf_counter()
loaded_weights: set[str] = model.load_weights(weight_iterator) # type: ignore
else:
loaded_weights: set[str] = model.load_weights(
self._get_all_weights(
@@ -698,7 +801,17 @@ class VAELoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the VAE based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
if is_single_file_bundle(model_path):
# The bundle stores the VAE config flat; the converted directory
# layout nests the same fields under "vae" and hoists the class out
# (`convert_ltx2_weights.py::_wrap_component_config`), and the build
# path below keys off that nesting. Match the shape rather than
# forking the build path.
section = deepcopy(read_ltx_metadata(model_path).config["vae"])
declared_class = section.pop("_class_name", None)
config = {"_class_name": declared_class, "vae": section}
else:
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name")
config.pop("_name_or_path", None)
assert class_name is not None, (
@@ -817,15 +930,21 @@ class VAELoader(ComponentLoader):
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(target_device)
# Find all safetensors files
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
# Common case: a single `.safetensors` checkpoint file.
# Some models may be sharded into multiple files; in that case we merge.
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
if is_single_file_bundle(model_path):
# `safetensors_load_file` takes a path, not an iterator, so the
# bundle's prefix-routed tensors are materialized here instead.
# The VAE is small enough for that (see `component_weights`).
loaded = dict(component_weights(model_path, "vae"))
else:
# Find all safetensors files
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
# Common case: a single `.safetensors` checkpoint file.
# Some models may be sharded into multiple files; in that case we merge.
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
# LTX-2 CausalVideoAutoencoder needs per_channel_statistics remapping
if class_name == "CausalVideoAutoencoder" and "vae" in config:
@@ -974,7 +1093,16 @@ class TransformerLoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the transformer based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
# LTX ships one bundle holding every component instead of a
# per-component directory: the config lives in the file's
# ``__metadata__`` and the transformer's tensors are selected by their
# top-level key prefix.
single_file = is_single_file_bundle(model_path)
if single_file:
config = deepcopy(
read_ltx_metadata(model_path).config["transformer"])
else:
config = get_diffusers_config(model=model_path)
hf_config = deepcopy(config)
cls_name = config.pop("_class_name")
config.pop("_name_or_path", None)
@@ -1008,9 +1136,12 @@ class TransformerLoader(ComponentLoader):
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
# Find all safetensors files
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
if single_file:
safetensors_list = [str(model_path)]
else:
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
# arch_config can infer architecture from weight keys (e.g. Flux2 layer counts)
update_fn = getattr(dit_config.arch_config, "update_from_weight_keys", None)
@@ -1065,6 +1196,11 @@ class TransformerLoader(ComponentLoader):
"hf_config": hf_config
},
weight_dir_list=safetensors_list,
# Route the bundle's tensors to the transformer, prefix
# stripped. Skipped when custom init weights replaced the
# file list above.
weight_iterator=(component_weights(model_path, "transformer")
if single_file and not use_custom_weights else None),
device=get_local_torch_device(),
hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim,
hsdp_shard_dim=fastvideo_args.hsdp_shard_dim,
@@ -1109,7 +1245,12 @@ class SchedulerLoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the scheduler based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
# A bundle keeps the scheduler's config in its own metadata; there is
# no `scheduler/` directory holding a `scheduler_config.json`.
if is_single_file_bundle(model_path):
config = deepcopy(read_ltx_metadata(model_path).config["scheduler"])
else:
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name")
assert class_name is not None, (
+9 -2
View File
@@ -7,7 +7,7 @@
from __future__ import annotations
import os
import contextlib
from collections.abc import Callable, Generator
from collections.abc import Callable, Generator, Iterable
from itertools import chain
from typing import Any
@@ -136,9 +136,15 @@ def maybe_load_fsdp_model(
pin_cpu_memory: bool = True,
enable_torch_compile: bool = False,
torch_compile_kwargs: dict[str, Any] | None = None,
weight_iterator: Iterable[tuple[str, torch.Tensor]] | None = None,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
``weight_iterator`` overrides reading ``weight_dir_list``, for checkpoints
whose keys need routing before they reach the model (see
``ltx_single_file.component_weights``). It must yield the same
``(name, cpu_tensor)`` pairs that ``safetensors_weights_iterator`` does.
"""
# NOTE(will): cast_forward_inputs=True shouldn't be needed as we are
# manually casting the inputs to the model
@@ -201,7 +207,8 @@ def maybe_load_fsdp_model(
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=pin_cpu_memory)
weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=True)
if weight_iterator is None:
weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=True)
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
load_model_from_full_model_state_dict(
model,
+296
View File
@@ -0,0 +1,296 @@
# SPDX-License-Identifier: Apache-2.0
"""Single-file LTX checkpoint support.
LTX ships one ``.safetensors`` bundle holding every component, so the two
assumptions the directory-based loaders make do not hold:
* there is no per-component ``config.json`` -- every component's config lives
in the file's ``__metadata__`` (safetensors stores an 8-byte little-endian
header length, then that many bytes of JSON whose ``__metadata__`` object is
a flat ``str -> str`` map; ``safe_open(...).metadata()`` returns it);
* there is no per-component directory -- components are told apart by the
top-level prefix on each tensor key.
This module reads that metadata and routes tensors by prefix. It does not
convert anything to a diffusers layout.
"""
from __future__ import annotations
import json
import os
from collections.abc import Generator
from dataclasses import dataclass
from typing import Any
import torch
from safetensors import safe_open
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# Top-level tensor key prefix per component. Verified against the LTX bundle
# header; every tensor in the file starts with one of these except
# ``duration_head.``, which has no FastVideo component today.
COMPONENT_PREFIXES: dict[str, str] = {
"transformer": "model.diffusion_model.",
"vae": "vae.",
"audio_vae": "audio_vae.",
"vocoder": "vocoder.",
"text_encoder": "text_embedding_projection.",
}
# Both embeddings connectors are stored *under* the transformer prefix, but
# FastVideo builds and runs them in the text encoder
# (``LTX2GemmaTextEncoderModel``'s ``Embeddings1DConnector``), not in the DiT,
# so they are routed to the text encoder instead. Only the transformer prefix
# is stripped: the sub-tree name survives so the text encoder's own rename
# table can map it onto the module name, the same way the already-converted
# repo layout is remapped. Matches how ``convert_ltx2_weights.py`` splits them.
TEXT_STACK_SUBPREFIXES: tuple[str, ...] = (
"video_embeddings_connector.",
"audio_embeddings_connector.",
)
def is_single_file_bundle(path: str) -> bool:
"""True when a model path names a bundle rather than a component directory."""
return str(path).endswith(".safetensors")
@dataclass(frozen=True)
class LTXCheckpointMetadata:
"""Parsed ``__metadata__`` of a single-file LTX checkpoint.
``config`` holds the per-component config sections keyed by component
(``transformer``, ``vae``, ``scheduler``, ``audio_vae``, ``vocoder``).
``model_version`` and ``gemma_source_checkpoint`` are carried through so
callers can validate the text encoder against the checkpoint that trained
it; ``gemma_source_checkpoint`` is absent on some variants. ``variant`` is
the training variant the header declares, when it declares one (see
:func:`bundle_variant`).
"""
config: dict[str, Any]
model_version: str | None
gemma_source_checkpoint: dict[str, Any] | None
variant: str | None = None
def read_ltx_metadata(path: str) -> LTXCheckpointMetadata:
"""Read the ``__metadata__`` config out of a single-file LTX checkpoint.
Only the header is parsed; no tensor data is touched. The remaining
metadata entries are deliberately not returned or logged -- the bundle also
carries a full license text and an ``encrypted_wandb_properties`` blob.
"""
with safe_open(path, framework="pt") as f:
metadata = f.metadata() or {}
config = json.loads(metadata["config"])
transformer = config.get("transformer")
if transformer is not None:
# ``frequencies_precision`` is the checkpoint's name for what the arch
# config calls ``double_precision_rope``; without this the field would
# silently fall back to its dataclass default.
precision = transformer.get("frequencies_precision")
if precision is not None:
transformer["double_precision_rope"] = precision == "float64"
gemma_source = metadata.get("gemma_source_checkpoint")
return LTXCheckpointMetadata(
config=config,
model_version=metadata.get("model_version"),
gemma_source_checkpoint=(json.loads(gemma_source) if gemma_source is not None else None),
variant=metadata.get("variant"),
)
def bundle_variant(metadata: LTXCheckpointMetadata, path: str) -> str:
"""The bundle's training variant: ``"distilled"`` or ``"base"``.
Decides which sampling preset a bundle gets (a distilled model wants its
short no-CFG schedule; everything else wants the standard one), so it must
be answerable from the header alone, without loading weights. An explicit
``variant`` entry in the file metadata wins when the checkpoint declares
one.
ponytail: filename fallback -- the known bundles declare no variant marker
in their headers (a distilled file and its sft sibling differ only in
fields incidental to the variant), so the "distilled" token in the file
name is the only signal available today. Drop the fallback when
checkpoints start declaring ``variant``.
"""
declared = metadata.variant or os.path.basename(path)
return "distilled" if "distilled" in declared.lower() else "base"
def bundle_model_index(path: str) -> dict[str, Any]:
"""Build a ``model_index.json``-shaped dict out of a bundle's own metadata.
The pipeline loader is written against a diffusers repo layout, which
answers two questions: which components exist, and what class is each. A
bundle already answers both in its ``__metadata__``, so this only reshapes
the answer -- it mirrors the entries
``convert_ltx2_weights.py::_build_model_index`` writes for the converted
directory layout, including the library each is declared under.
A section that exists but declares no class is emitted as ``[None, None]``
rather than dropped. ``ComposedPipelineBase.load_modules`` already treats a
null library as "declared, but not something to build" and removes the
component from the required set; dropping the key instead would fail its
required-module check for a component the checkpoint does carry.
``text_encoder`` and ``tokenizer`` are always declared: they live outside
the bundle, but the pipeline needs both.
"""
model_index: dict[str, Any] = {
# ponytail: the pipeline class is pinned by the registry's bundle
# table (`registry._bundle_config_info`) or an explicit
# `override_pipeline_cls_name`, and `load_modules` pops both of these
# without reading them. Nothing here is entitled to name a pipeline,
# so these are placeholders -- they exist only because those pops have
# no default.
"_class_name": None,
"_diffusers_version": None,
}
for component, section in read_ltx_metadata(path).config.items():
cls_name = (section.get("_class_name") if isinstance(section, dict) else None)
model_index[component] = (["diffusers", cls_name] if cls_name else [None, None])
model_index["text_encoder"] = ["transformers", "LTX2GemmaTextEncoderModel"]
model_index["tokenizer"] = ["transformers", "AutoTokenizer"]
return model_index
def build_dit_config(metadata: LTXCheckpointMetadata) -> LTX2VideoConfig:
"""Build the LTX-2 DiT config from checkpoint metadata.
Reuses ``update_model_arch``, so the metadata keys that name an arch field
win and the rest of the section is ignored -- no hand-written constants.
"""
config = LTX2VideoConfig()
config.update_model_arch(metadata.config["transformer"])
return config
def component_weights(
path: str,
component: str,
) -> Generator[tuple[str, torch.Tensor], None, None]:
"""Yield ``(key_with_prefix_stripped, tensor)`` for one component.
The transformer prefix covers two owners -- see
``TEXT_STACK_SUBPREFIXES`` -- so keys under it are split before the
component filter applies.
``safe_open`` mmaps the file and ``get_tensor`` materializes one tensor at
a time, so iterating never holds more than a single tensor in RAM. Callers
that need a state dict for a small component can do
``dict(component_weights(path, "vocoder"))``; do not do that for the
transformer.
ponytail: every rank reads the file itself, unlike
``safetensors_weights_iterator``, which has local rank 0 read and broadcast.
Correct either way, but the ceiling is filesystem bandwidth -- N ranks read
N copies. Add the broadcast if loading a bundle over a shared filesystem
turns out to be the bottleneck.
"""
prefix = COMPONENT_PREFIXES[component]
transformer_prefix = COMPONENT_PREFIXES["transformer"]
with safe_open(path, framework="pt") as f:
for key in f.keys():
if key.startswith(transformer_prefix):
name = key[len(transformer_prefix):]
# str.startswith takes a tuple.
owner = ("text_encoder" if name.startswith(TEXT_STACK_SUBPREFIXES) else "transformer")
if owner != component:
continue
elif key.startswith(prefix):
name = key[len(prefix):]
else:
continue
yield name, f.get_tensor(key)
def resolve_text_encoder_root(
configured_path: str | None,
override_path: str | None = None,
metadata: LTXCheckpointMetadata | None = None,
) -> str:
"""Locate the text-encoder root that goes with a single-file bundle.
A bundle carries every component's *weights* but no pointer to the text
encoder, which lives outside it. The root therefore has to be declared,
and there are exactly two ways to declare it:
* ``override_path`` -- an explicit per-run argument, so a one-off run needs
no config edit;
* ``configured_path`` -- the pipeline config's encoder path, which is where
model composition already lives and which survives files being moved.
Deliberately absent: any search of the bundle's directory. Picking up an
encoder nobody asked for changes what gets loaded without being requested,
and silently picks the wrong one when two sit side by side. When neither
source is set this raises and names both, rather than guessing.
``metadata`` is used only to *validate* a declared root, never to find one:
a bundle may not declare an encoder pairing at all, so discovery cannot
depend on it.
"""
root = override_path or configured_path
if not root:
raise ValueError("A single-file checkpoint does not carry its text encoder, so the "
"encoder root must be declared. Set `gemma_model_path` in the "
"pipeline config, or pass the text-encoder path explicitly for "
"this run. It is not inferred from the checkpoint's directory: "
"loading whichever encoder happens to sit beside the file would "
"silently pick the wrong one when several are present.")
expected = (metadata.gemma_source_checkpoint or {}).get("gemma_version") if metadata is not None else None
if expected:
# ponytail: warn, don't raise -- the declared root is the user's
# explicit instruction and the pairing is advisory. Promote to a hard
# error if a mismatch ever turns out to produce silent garbage rather
# than an obvious shape failure.
actual = _read_encoder_version(root)
if actual is not None and actual != expected:
logger.warning(
"Checkpoint expects text-encoder version %r but the encoder at "
"%s reports %r; continuing with the declared root.", expected, root, actual)
return root
def _read_encoder_version(root: str) -> str | None:
"""The encoder's declared version, or None if it declares none."""
config_path = os.path.join(root, "config.json")
if not os.path.isfile(config_path):
return None
try:
with open(config_path, encoding="utf-8") as f:
return json.load(f).get("gemma_version")
except (OSError, json.JSONDecodeError):
return None
def model_index_and_component_path(model_path: str, module_type: str) -> tuple[dict[str, Any], str]:
"""``(model_index, component_path)`` for a directory repo OR a bundle.
The two differ in both halves: a repo answers "what components exist" from
``model_index.json`` and puts each in its own subdirectory, while a bundle
declares its components in its own metadata and holds them all in one file.
Callers that resolve those two things together should route through here so
the bundle case is handled once instead of at every site.
ponytail: a bundle's text encoder and tokenizer live OUTSIDE the file, so
this returns the bundle path for them too, which is wrong for those two
module types. Training does not load them (text embeddings are
preprocessed), so it does not arise. Give this the resolved encoder root the
day a caller needs them.
"""
if is_single_file_bundle(model_path):
return bundle_model_index(model_path), model_path
from fastvideo.utils import verify_model_config_and_directory
return verify_model_config_and_directory(model_path), os.path.join(model_path, module_type)
+2
View File
@@ -38,6 +38,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
"LTX2Transformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
# ``_class_name`` declared by single-file LTX checkpoint metadata.
"AVTransformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
"SD3Transformer2DModel": ("dits", "sd3", "SD3Transformer2DModel"),
"LingBotWorldTransformer3DModel": ("dits", "lingbotworld", "LingBotWorldTransformer3DModel"),
"LingBotWorld2CausalFastTransformer3DModel": (
@@ -277,7 +277,7 @@ class LTX2Pipeline(LoRAPipeline):
modules[module_name] = loaded_modules[module_name]
continue
component_model_path = os.path.join(self.model_path, module_name)
component_model_path = self._component_path(module_name, fastvideo_args)
if module_name == "tokenizer" and not os.path.isdir(component_model_path):
gemma_path = os.path.join(self.model_path, "text_encoder", "gemma")
if os.path.isdir(gemma_path):
@@ -49,6 +49,16 @@ class LTX2AudioDecodingStage(PipelineStage):
audio_latents = batch.extra.get("ltx2_audio_latents")
if audio_latents is None:
return batch
# A checkpoint that declares no audio decoder builds these as None
# (``get_module`` returns its default), while the transformer is still
# audio-video and still produces audio latents. Having latents is
# therefore not evidence that anything can decode them -- guard on the
# modules too, and leave the video path unaffected.
if self.audio_decoder is None or self.vocoder is None:
logger.info(
"Skipping audio decoding: this checkpoint declares no audio "
"decoder/vocoder. Video output is unaffected.")
return batch
device = get_local_torch_device()
self.audio_decoder = self.audio_decoder.to(device)
+42 -1
View File
@@ -20,6 +20,8 @@ from fastvideo.hooks.activation_trace import attach_activation_trace, detach_act
from fastvideo.logger import init_logger
from fastvideo.profiler import get_or_create_profiler
from fastvideo.models.loader.component_loader import PipelineComponentLoader
from fastvideo.models.loader.ltx_single_file import (bundle_model_index, is_single_file_bundle, read_ltx_metadata,
resolve_text_encoder_root)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages import PipelineStage
import fastvideo.envs as envs
@@ -79,6 +81,7 @@ class ComposedPipelineBase(ABC):
self.model_path: str = model_path
self._stages: list[PipelineStage] = []
self._bundle_encoder_root: str | None = None
self._stage_name_mapping: dict[str, PipelineStage] = {}
self._trace_mgr = None
@@ -328,12 +331,50 @@ class ComposedPipelineBase(ABC):
self.model_path = model_path
# fastvideo_args.downloaded_model_path = model_path
logger.info("Model path: %s", model_path)
if is_single_file_bundle(model_path):
# A bundle is a file, so there is no directory to verify and no
# model_index.json to read; it declares the same thing itself.
return bundle_model_index(model_path)
config = verify_model_config_and_directory(
model_path,
required_component_dirs=self.get_hf_download_component_dirs(),
)
return cast(dict[str, Any], config)
def _text_encoder_root(self, fastvideo_args: FastVideoArgs) -> str:
"""Where a bundle's text encoder lives -- it is not in the bundle.
Declared, never discovered: a per-run environment override first, then
the encoder path the pipeline config already carries.
"""
encoder_configs = getattr(fastvideo_args.pipeline_config, "text_encoder_configs", ())
return resolve_text_encoder_root(
configured_path=(getattr(encoder_configs[0].arch_config, "gemma_model_path", None)
if encoder_configs else None),
override_path=os.environ.get("FASTVIDEO_LTX_ENCODER_ROOT") or None,
metadata=read_ltx_metadata(self.model_path),
)
def _component_path(self, module_name: str, fastvideo_args: FastVideoArgs) -> str:
"""Where one component's config and weights live.
A directory repo puts each component in its own subdirectory. A bundle
holds all of them in one file except the text stack, which is declared
separately. Subclasses that override ``load_modules`` must route
through here rather than joining the path themselves, or a bundle
silently falls back to the directory layout and fails looking for a
subdirectory of a file.
"""
if not is_single_file_bundle(self.model_path):
return os.path.join(self.model_path, module_name)
if module_name in ("text_encoder", "tokenizer"):
# Resolved once: it raises when undeclared, and doing that per
# component would report the same failure several times.
if self._bundle_encoder_root is None:
self._bundle_encoder_root = self._text_encoder_root(fastvideo_args)
return self._bundle_encoder_root
return self.model_path
@property
def required_config_modules(self) -> list[str]:
"""
@@ -465,7 +506,7 @@ class ComposedPipelineBase(ABC):
else:
load_module_name = module_name
component_model_path = os.path.join(self.model_path, load_module_name)
component_model_path = self._component_path(load_module_name, fastvideo_args)
module = PipelineComponentLoader.load_module(
module_name=load_module_name,
component_model_path=component_model_path,
+102 -5
View File
@@ -142,6 +142,25 @@ _MODEL_HF_PATH_TO_NAME: dict[str, str] = {}
# Detectors to identify model families from paths or class names
_MODEL_NAME_DETECTORS: list[tuple[str, Callable[[str], bool]]] = []
# Single-file checkpoint bundles carry no `model_index.json`, so their pipeline
# config is resolved from the transformer class the file declares about itself.
# Keyed on that `_class_name` -- a stable, checkpoint-declared identifier --
# and NOT on `model_id`, which is a positional index (`str(len(_CONFIG_REGISTRY))`)
# and would silently rebind whenever registration order changes. The value pins
# the model family, the pipeline class to build (a bundle names no pipeline
# itself), and the default preset per training variant; the variant comes from
# `ltx_single_file.bundle_variant` (declared in the header when present, with a
# filename-token fallback). The transformer CLASS is never inferred from the
# file name.
_BUNDLE_TRANSFORMER_TO_CONFIG: dict[str, tuple[str, str, dict[str, str]]] = {
# transformer _class_name -> (model_family, pipeline class,
# variant -> default_preset)
"AVTransformer3DModel": ("ltx2", "LTX2Pipeline", {
"distilled": "ltx2_distilled",
"base": "ltx2_base",
}),
}
def register_configs(
sampling_param_cls: type[SamplingParam] | None,
@@ -187,6 +206,64 @@ def get_model_short_name(model_id: str) -> str:
return model_id
def _bundle_config_info(model_path: str) -> ConfigInfo:
"""Resolve the pipeline config for a single-file checkpoint bundle.
The bundle declares its own component classes in ``__metadata__``; reading
that is the only inference done here. The transformer's ``_class_name`` is
mapped through :data:`_BUNDLE_TRANSFORMER_TO_CONFIG`, an explicit table,
and the bundle's training variant (``bundle_variant``: header-declared,
filename token as fallback) selects the ``default_preset`` within that
entry, so a distilled bundle gets distilled sampling defaults and a
base/sft bundle gets the standard ones. Nothing is matched by wildcard,
and the transformer class is never read off the file name -- a file name
is incidental, and making class resolution depend on it makes correct
loading depend on what a checkpoint happens to be called.
An unrecognized transformer class raises, naming both the table to extend
and the per-run override, rather than falling through to a guess.
ponytail: LTX is the only family shipping a bundle, so this reads LTX
metadata directly. Dispatch on a per-format reader if a second bundle
format ever appears.
"""
from fastvideo.models.loader.ltx_single_file import (bundle_variant, read_ltx_metadata)
try:
metadata = read_ltx_metadata(model_path)
cls_name = metadata.config.get("transformer", {}).get("_class_name")
except KeyError: # no `config` section: not a bundle we know how to read
metadata = None
cls_name = None
entry = _BUNDLE_TRANSFORMER_TO_CONFIG.get(cls_name) if cls_name else None
if entry is None or metadata is None:
raise ValueError(
f"Single-file checkpoint {model_path} declares transformer class {cls_name!r}, which is not in "
f"_BUNDLE_TRANSFORMER_TO_CONFIG (known: {sorted(_BUNDLE_TRANSFORMER_TO_CONFIG)}). Add the class to that "
"table, or select the pipeline explicitly for this run with `override_pipeline_cls_name`. The pipeline is "
"never inferred from the checkpoint's file name.")
model_family, pipeline_cls_name, variant_presets = entry
variant = bundle_variant(metadata, model_path)
default_preset = variant_presets.get(variant)
if default_preset is None:
raise ValueError(f"Single-file checkpoint {model_path} resolves to variant {variant!r}, but "
f"_BUNDLE_TRANSFORMER_TO_CONFIG[{cls_name!r}] only maps {sorted(variant_presets)}. "
"Add the variant to that table.")
for config_info in _CONFIG_REGISTRY.values():
if config_info.model_family == model_family and config_info.default_preset == default_preset:
# A bundle names no pipeline itself (there is no model_index.json
# to carry a `_class_name`), so the table's pipeline class is
# pinned onto the resolved config.
return dataclasses.replace(config_info, pipeline_cls_name=pipeline_cls_name)
raise ValueError(f"_BUNDLE_TRANSFORMER_TO_CONFIG maps {cls_name!r} ({variant!r}) to model_family={model_family!r} "
f"default_preset={default_preset!r}, but no registered config declares that pair. "
"Fix the table or the registration.")
def _get_config_info(
model_path: str,
*,
@@ -209,7 +286,13 @@ def _get_config_info(
model_id = _MODEL_HF_PATH_TO_NAME[registered_model_hf_id]
return _CONFIG_REGISTRY.get(model_id)
# 3. Use detectors (path or model_index pipeline name).
# 3. A single-file bundle is a FILE, so it would take the directory branch
# below and be rejected for having no model_index.json. Its config lives
# in its own `__metadata__` instead.
if model_path.endswith(".safetensors"):
return _bundle_config_info(model_path)
# 4. Use detectors (path or model_index pipeline name).
if os.path.exists(model_path):
config = verify_model_config_and_directory(model_path, required_component_dirs=[])
else:
@@ -267,7 +350,10 @@ def _register_configs() -> None:
register_configs(
sampling_param_cls=None,
pipeline_config_cls=LTX2T2VConfig,
workload_types=(WorkloadType.T2V, ),
# The LTX-2 family conditions i2v in LATENT SPACE (clean-latent +
# denoise mask, ltx2_image_conditioning.py); one checkpoint serves
# both workloads, so declaring T2V only under-reported the family.
workload_types=(WorkloadType.T2V, WorkloadType.I2V),
hf_model_paths=[
"FastVideo/LTX2-Distilled-Diffusers",
# LTX-2.3 distilled aliases share the distilled pipeline/preset.
@@ -286,7 +372,10 @@ def _register_configs() -> None:
register_configs(
sampling_param_cls=None,
pipeline_config_cls=LTX2T2VConfig,
workload_types=(WorkloadType.T2V, ),
# The LTX-2 family conditions i2v in LATENT SPACE (clean-latent +
# denoise mask, ltx2_image_conditioning.py); one checkpoint serves
# both workloads, so declaring T2V only under-reported the family.
workload_types=(WorkloadType.T2V, WorkloadType.I2V),
hf_model_paths=[
"Lightricks/LTX-2.3",
"FastVideo/LTX2.3-base",
@@ -307,7 +396,10 @@ def _register_configs() -> None:
register_configs(
sampling_param_cls=None,
pipeline_config_cls=LTX2T2VConfig,
workload_types=(WorkloadType.T2V, ),
# The LTX-2 family conditions i2v in LATENT SPACE (clean-latent +
# denoise mask, ltx2_image_conditioning.py); one checkpoint serves
# both workloads, so declaring T2V only under-reported the family.
workload_types=(WorkloadType.T2V, WorkloadType.I2V),
hf_model_paths=[
"Lightricks/LTX-2",
"FastVideo/LTX2-base",
@@ -1237,7 +1329,12 @@ def get_model_info(
pipeline_name: str | None = override_pipeline_cls_name
logger.info("Using override pipeline class name %s", pipeline_name)
else:
if os.path.exists(model_path):
if model_path.endswith(".safetensors"):
# A single-file bundle has no model_index.json and names no
# pipeline itself; `_bundle_config_info` above pinned the class
# through `pipeline_cls_name`.
config = {}
elif os.path.exists(model_path):
config = verify_model_config_and_directory(model_path, required_component_dirs=[])
else:
config = maybe_download_model_index(model_path, revision=revision)
+400
View File
@@ -0,0 +1,400 @@
# SPDX-License-Identifier: Apache-2.0
"""Checks for single-file LTX metadata parsing and prefix routing.
Builds a tiny synthetic bundle shaped like the real one -- same metadata keys,
same top-level tensor prefixes -- so no checkpoint is needed. Run directly
(``python fastvideo/tests/api/test_ltx_single_file.py``) or under pytest.
"""
import json
import tempfile
from pathlib import Path
import torch
from safetensors.torch import save_file
from fastvideo.models.loader.ltx_single_file import (
LTXCheckpointMetadata,
resolve_text_encoder_root,
build_dit_config,
bundle_model_index,
component_weights,
read_ltx_metadata,
)
# Subset of the transformer section as the checkpoint declares it.
TRANSFORMER_SECTION = {
"_class_name": "AVTransformer3DModel",
"num_layers": 48,
"num_attention_heads": 32,
"attention_head_dim": 128,
"audio_num_attention_heads": 32,
"audio_attention_head_dim": 64,
"cross_attention_dim": 4096,
"audio_cross_attention_dim": 2048,
"caption_channels": 3840,
"ff_bias": False,
"cross_attention_adaln": True,
"caption_proj_before_connector": True,
"apply_gated_attention": True,
"rope_type": "split",
"frequencies_precision": "float64",
}
def _write_bundle(path: Path,
*,
with_gemma_source: bool,
transformer_cls: str = "AVTransformer3DModel",
variant: str | None = None) -> None:
tensors = {
"model.diffusion_model.patchify_proj.weight": torch.zeros(2, 2),
"model.diffusion_model.transformer_blocks.0.ff.net.2.weight": torch.zeros(2, 2),
# Stored under the transformer prefix but owned by the text stack.
"model.diffusion_model.video_embeddings_connector.x.weight": torch.zeros(2),
"model.diffusion_model.audio_embeddings_connector.x.weight": torch.zeros(2),
"vae.encoder.conv.weight": torch.zeros(2),
"audio_vae.decoder.conv.weight": torch.zeros(2),
"vocoder.conv.weight": torch.zeros(2),
"text_embedding_projection.video_aggregate_embed.weight": torch.zeros(2),
"duration_head.linear.weight": torch.zeros(2),
}
section = dict(TRANSFORMER_SECTION, _class_name=transformer_cls)
config = {
"transformer": section,
# Declares its class: a component that can be built from the bundle.
"vae": {
"_class_name": "CausalVideoAutoencoder",
"dims": 3
},
# Present and carrying weights, but naming no class.
"audio_vae": {},
}
metadata = {
"config": json.dumps(config),
"model_version": "9.9.9",
"license": "x" * 128,
}
if with_gemma_source:
metadata["gemma_source_checkpoint"] = json.dumps({"ltx_version": "9.9.9", "gemma_version": "fake-encoder-v0"})
if variant is not None:
metadata["variant"] = variant
save_file(tensors, str(path), metadata=metadata)
def test_read_metadata_and_routing() -> None:
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "bundle.safetensors"
_write_bundle(path, with_gemma_source=True)
meta = read_ltx_metadata(str(path))
assert meta.model_version == "9.9.9"
assert meta.gemma_source_checkpoint == {
"ltx_version": "9.9.9",
"gemma_version": "fake-encoder-v0",
}
assert set(meta.config) == {"transformer", "vae", "audio_vae"}
# frequencies_precision is the checkpoint's name for double_precision_rope.
assert meta.config["transformer"]["double_precision_rope"] is True
# Prefix stripped, connectors excluded, nothing from other components.
transformer_keys = {n for n, _ in component_weights(str(path), "transformer")}
assert transformer_keys == {
"patchify_proj.weight",
"transformer_blocks.0.ff.net.2.weight",
}
for component, expected in (
("vae", {"encoder.conv.weight"}),
("audio_vae", {"decoder.conv.weight"}),
("vocoder", {"conv.weight"}),
# The connectors come off the transformer prefix with their
# sub-tree name intact; the text encoder renames from there.
("text_encoder", {
"video_aggregate_embed.weight",
"video_embeddings_connector.x.weight",
"audio_embeddings_connector.x.weight",
}),
):
assert {n for n, _ in component_weights(str(path), component)} == expected
def test_bundle_model_index_declares_what_the_metadata_declares() -> None:
"""The model_index a bundle stands in for: same shape, same libraries."""
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "bundle.safetensors"
_write_bundle(path, with_gemma_source=True)
index = bundle_model_index(str(path))
# A section that names a class is a component we can build.
assert index["transformer"] == ["diffusers", "AVTransformer3DModel"]
assert index["vae"] == ["diffusers", "CausalVideoAutoencoder"]
# A section that names none stays declared with a null library, so
# `load_modules` drops it from the required set on its own. Omitting
# the key instead would trip its required-module check.
assert index["audio_vae"] == [None, None]
# Both live outside the bundle, but the pipeline requires them.
assert index["text_encoder"] == ["transformers", "LTX2GemmaTextEncoderModel"]
assert index["tokenizer"] == ["transformers", "AutoTokenizer"]
# `load_modules` pops both of these without a default.
assert "_class_name" in index and "_diffusers_version" in index
def test_missing_gemma_source_is_none() -> None:
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "bundle.safetensors"
_write_bundle(path, with_gemma_source=False)
assert read_ltx_metadata(str(path)).gemma_source_checkpoint is None
def test_dit_config_takes_ff_bias_from_metadata() -> None:
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "bundle.safetensors"
_write_bundle(path, with_gemma_source=True)
arch = build_dit_config(read_ltx_metadata(str(path))).arch_config
# ff_bias comes from metadata; audio_ff_bias has no metadata key and
# must keep the default, matching the audio_ff.*.bias tensors that the
# checkpoint does carry.
assert arch.ff_bias is False
assert arch.audio_ff_bias is True
assert arch.num_layers == 48
assert arch.cross_attention_adaln is True
assert arch.caption_channels == 3840
def test_defaults_unchanged_without_metadata() -> None:
from fastvideo.configs.models.dits.ltx2 import LTX2VideoArchConfig
arch = LTX2VideoArchConfig()
assert arch.ff_bias is True
assert arch.audio_ff_bias is True
def test_encoder_root_requires_an_explicit_declaration() -> None:
"""Neither source set -> raise, naming both. Never guess from the filesystem."""
import pytest
with pytest.raises(ValueError, match="must be declared"):
resolve_text_encoder_root(configured_path=None, override_path=None)
def test_encoder_root_override_beats_config() -> None:
assert resolve_text_encoder_root("/from/config", "/from/override") == "/from/override"
assert resolve_text_encoder_root("/from/config", None) == "/from/config"
def test_encoder_root_is_not_discovered_from_a_sibling_directory() -> None:
"""A plausible encoder sitting next to the bundle must NOT be picked up."""
with tempfile.TemporaryDirectory() as tmp:
sibling = Path(tmp) / "gemma"
sibling.mkdir()
(sibling / "config.json").write_text("{}")
import pytest
with pytest.raises(ValueError):
resolve_text_encoder_root(configured_path=None, override_path=None)
def test_declared_root_survives_a_version_mismatch() -> None:
"""Pairing metadata validates a declared root; it never overrides it."""
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp) / "encoder"
root.mkdir()
(root / "config.json").write_text(json.dumps({"gemma_version": "other-v0"}))
meta = LTXCheckpointMetadata(config={},
model_version="9.9.9",
gemma_source_checkpoint={"gemma_version": "fake-encoder-v0"})
assert resolve_text_encoder_root(str(root), None, meta) == str(root)
def test_bundle_resolves_its_pipeline_config_from_declared_transformer_class() -> None:
"""A bundle path resolves through the alias table, not through model_index.json."""
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
from fastvideo.registry import get_pipeline_config_cls_from_name
with tempfile.TemporaryDirectory() as tmp:
# Deliberately uninformative name: resolution must not read the file name.
path = Path(tmp) / "anonymous.safetensors"
_write_bundle(path, with_gemma_source=True)
assert get_pipeline_config_cls_from_name(str(path)) is LTX2T2VConfig
def test_unknown_bundle_transformer_class_raises_naming_table_and_override() -> None:
"""No wildcard fallback: an unmapped class fails loud and says how to fix it."""
import pytest
from fastvideo.registry import get_pipeline_config_cls_from_name
with tempfile.TemporaryDirectory() as tmp:
# Name it after a model the detectors DO know, to prove the file name
# is not consulted -- only the class the checkpoint declares.
path = Path(tmp) / "ltx2-distilled.safetensors"
_write_bundle(path, with_gemma_source=True, transformer_cls="NotAModel")
with pytest.raises(ValueError) as excinfo:
get_pipeline_config_cls_from_name(str(path))
message = str(excinfo.value)
assert "_BUNDLE_TRANSFORMER_TO_CONFIG" in message
assert "override_pipeline_cls_name" in message
def test_preset_resolution_prefers_the_declared_variant() -> None:
"""A bundle's header-declared `variant` decides its preset, not its name."""
from fastvideo.api.sampling_param import SamplingParam
with tempfile.TemporaryDirectory() as tmp:
# Distilled declared, in a file whose name says nothing.
distilled = Path(tmp) / "anonymous.safetensors"
_write_bundle(distilled, with_gemma_source=False, variant="distilled-rc2")
sp = SamplingParam.from_pretrained(str(distilled))
assert sp.num_inference_steps == 8
assert sp.guidance_scale == 1.0
# A declared sft variant beats a file name that says "distilled".
sft = Path(tmp) / "something-distilled.safetensors"
_write_bundle(sft, with_gemma_source=True, variant="sft-rc2")
sp = SamplingParam.from_pretrained(str(sft))
assert sp.num_inference_steps == 40
assert sp.guidance_scale == 3.0
def test_preset_resolution_falls_back_to_the_file_name() -> None:
"""No variant in the header: the "distilled" filename token decides."""
from fastvideo.api.sampling_param import SamplingParam
with tempfile.TemporaryDirectory() as tmp:
distilled = Path(tmp) / "some-distilled-bundle.safetensors"
_write_bundle(distilled, with_gemma_source=False)
sp = SamplingParam.from_pretrained(str(distilled))
assert sp.num_inference_steps == 8
assert sp.guidance_scale == 1.0
base = Path(tmp) / "some-sft-bundle.safetensors"
_write_bundle(base, with_gemma_source=True)
sp = SamplingParam.from_pretrained(str(base))
assert sp.num_inference_steps == 40
assert sp.guidance_scale == 3.0
def test_bundle_preset_table_is_internally_consistent() -> None:
"""Invariants the bundle->preset mapping relies on, asserted loudly."""
from fastvideo.registry import _BUNDLE_TRANSFORMER_TO_CONFIG, _CONFIG_REGISTRY
registered = {(ci.model_family, ci.default_preset) for ci in _CONFIG_REGISTRY.values()}
for cls_name, (family, pipeline_cls_name, variants) in _BUNDLE_TRANSFORMER_TO_CONFIG.items():
assert set(variants) == {
"distilled", "base"
}, (f"{cls_name}: bundle_variant() only ever returns 'distilled' or 'base', "
f"but the table maps {sorted(variants)} -- some bundles could not resolve a preset.")
assert pipeline_cls_name, (f"{cls_name}: a bundle names no pipeline itself, so the table entry must.")
for variant, preset in variants.items():
assert (family, preset) in registered, (
f"{cls_name}/{variant} -> ({family!r}, {preset!r}) is not a registered config; "
"fix _BUNDLE_TRANSFORMER_TO_CONFIG or _register_configs().")
def test_bundle_resolves_a_pipeline_class_without_a_model_index() -> None:
"""`get_model_info` on a bundle pins the pipeline class from the table,
so `VideoGenerator.from_pretrained(<bundle>)` needs no override."""
from fastvideo.registry import get_model_info
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "anonymous.safetensors"
_write_bundle(path, with_gemma_source=True, variant="distilled-rc2")
info = get_model_info(str(path))
assert info.pipeline_cls.__name__ == "LTX2Pipeline"
_EXAMPLES_DIR = Path(__file__).resolve().parents[3] / "examples" / "inference" / "basic"
_EXAMPLE_SCRIPTS = ("basic_ltx2_5_distilled.py", "basic_ltx2_5.py")
def _load_example(name: str):
import importlib.util
spec = importlib.util.spec_from_file_location(name[:-3], _EXAMPLES_DIR / name)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
def test_example_scripts_parse_help_cleanly() -> None:
import runpy
import sys
import pytest
for name in _EXAMPLE_SCRIPTS:
script = str(_EXAMPLES_DIR / name)
argv = sys.argv
sys.argv = [script, "--help"]
try:
with pytest.raises(SystemExit) as excinfo:
runpy.run_path(script, run_name="__main__")
assert excinfo.value.code == 0, f"{name} --help exited nonzero"
finally:
sys.argv = argv
def test_example_sampling_defaults_come_from_the_bundle_preset() -> None:
"""With no sampling flags the examples pass NO overrides, so what runs is
exactly ``SamplingParam.from_pretrained(bundle)`` -- the preset the
bundle's variant selects."""
from fastvideo.api.sampling_param import SamplingParam
with tempfile.TemporaryDirectory() as tmp:
for name, variant, steps, cfg in (
("basic_ltx2_5_distilled.py", "distilled-rc2", 8, 1.0),
("basic_ltx2_5.py", "sft-rc2", 40, 3.0),
):
bundle = Path(tmp) / f"{variant}.safetensors"
_write_bundle(bundle, with_gemma_source=True, variant=variant)
module = _load_example(name)
args = module.parse_args(["--model-path", str(bundle)])
assert module.sampling_overrides(args) == {}, (f"{name}: unset flags must not override the preset")
sp = SamplingParam.from_pretrained(str(bundle))
assert (sp.num_inference_steps, sp.guidance_scale) == (steps, cfg)
def test_connector_factorization_must_match_its_stream_width() -> None:
"""A connector's heads x head_dim must equal the width of the stream it feeds.
This is the mistake that loads cleanly and computes wrong, so it fails loud
and names both numbers. A field the checkpoint does not declare is not
checked -- the arch default applies and there is nothing to contradict.
"""
import pytest
from fastvideo.models.loader.component_loader import _check_connector_widths
declared = {
"connector_num_attention_heads": 32,
"connector_attention_head_dim": 128,
"cross_attention_dim": 4096,
"audio_connector_num_attention_heads": 32,
"audio_connector_attention_head_dim": 64,
"audio_cross_attention_dim": 2048,
}
_check_connector_widths(declared)
_check_connector_widths({})
_check_connector_widths({"connector_num_attention_heads": 32})
# The audio connector inheriting the video head_dim is the exact failure
# this exists to stop: 32 x 128 = 4096, but the audio stream is 2048.
with pytest.raises(ValueError, match="audio"):
_check_connector_widths(dict(declared, audio_connector_attention_head_dim=128))
with pytest.raises(ValueError, match="video"):
_check_connector_widths(dict(declared, connector_num_attention_heads=30))
if __name__ == "__main__":
test_read_metadata_and_routing()
test_bundle_model_index_declares_what_the_metadata_declares()
test_missing_gemma_source_is_none()
test_dit_config_takes_ff_bias_from_metadata()
test_defaults_unchanged_without_metadata()
test_encoder_root_override_beats_config()
test_bundle_resolves_its_pipeline_config_from_declared_transformer_class()
test_unknown_bundle_transformer_class_raises_naming_table_and_override()
test_preset_resolution_prefers_the_declared_variant()
test_preset_resolution_falls_back_to_the_file_name()
test_bundle_preset_table_is_internally_consistent()
test_bundle_resolves_a_pipeline_class_without_a_model_index()
test_example_scripts_parse_help_cleanly()
test_example_sampling_defaults_come_from_the_bundle_preset()
test_connector_factorization_must_match_its_stream_width()
print("ok")
+7
View File
@@ -11,10 +11,13 @@ import torch
from tqdm.auto import tqdm
from fastvideo.distributed import get_sp_group, get_world_group
from fastvideo.logger import init_logger
from fastvideo.train.callbacks.callback import CallbackDict
from fastvideo.train.methods.base import LogScalar, TrainingMethod
from fastvideo.train.utils.tracking import build_tracker
logger = init_logger(__name__)
if TYPE_CHECKING:
from fastvideo.train.utils.training_config import (
TrainingConfig, )
@@ -222,6 +225,10 @@ class Trainer:
metrics["step_time_sec"] = (time.perf_counter() - t0)
metrics["vsa_sparsity"] = float(tc.vsa_sparsity)
if self.global_rank == 0 and metrics:
# Console as well as tracker: with trackers disabled the
# tracker is a DummyTracker, and this is then the only place a
# step's loss is reported at all.
logger.info("step %d %s", step, {k: round(v, 6) for k, v in metrics.items()})
self.tracker.log(metrics, step)
self.callbacks.on_training_step_end(
+4 -2
View File
@@ -20,6 +20,9 @@ from fastvideo.utils import (
maybe_download_model,
verify_model_config_and_directory,
)
from fastvideo.models.loader.ltx_single_file import (
model_index_and_component_path,
)
from fastvideo.platforms import AttentionBackendEnum
if TYPE_CHECKING:
@@ -106,7 +109,7 @@ def load_module_from_path(
fastvideo_args: Any = _make_training_args(training_config, model_path=model_path)
local_model_path = maybe_download_model(model_path)
config = verify_model_config_and_directory(local_model_path)
config, component_path = model_index_and_component_path(local_model_path, module_type)
if module_type not in config:
raise ValueError(f"Module {module_type!r} not found in "
@@ -120,7 +123,6 @@ def load_module_from_path(
# Trailing modular-manifest metadata does not change component dispatch;
# the provider and architecture remain the first two fields.
transformers_or_diffusers, _architecture = module_info[:2]
component_path = os.path.join(local_model_path, module_type)
# fastvideo_args is freshly built above and never escapes this function,
# so overrides are plain assignments — nothing to save or restore.
+2 -2
View File
@@ -280,7 +280,8 @@ class DistillationPipeline(TrainingPipeline):
# Download the model if it's a Hugging Face model ID
local_model_path = maybe_download_model(model_path)
logger.info("Model downloaded/found at: %s", local_model_path)
config = verify_model_config_and_directory(local_model_path)
from fastvideo.models.loader.ltx_single_file import (model_index_and_component_path)
config, component_path = model_index_and_component_path(local_model_path, module_type)
if module_type not in config:
if hasattr(self, '_extra_config_module_map') and module_type in self._extra_config_module_map:
@@ -298,7 +299,6 @@ class DistillationPipeline(TrainingPipeline):
raise ValueError(f"Module {module_type} has null value in config at {local_model_path}")
transformers_or_diffusers, architecture = module_info
component_path = os.path.join(local_model_path, module_type)
module = PipelineComponentLoader.load_module(
module_name=module_type,
component_model_path=component_path,