Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
825c6365e0 | ||
|
|
bb1fb7c149 | ||
|
|
436c44e583 | ||
|
|
6343eadbec | ||
|
|
3f5614fd03 |
@@ -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!
|
||||
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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, (
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
@@ -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(
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user