Automated PR - 2026-03-30

This commit is contained in:
github-actions[bot]
2026-03-30 17:59:34 +00:00
parent ae855f8538
commit f4d0c1ec0e
48 changed files with 8429 additions and 6644 deletions
+2 -2
View File
@@ -2,7 +2,7 @@
[![Website](https://img.shields.io/badge/Website-LTX-181717?logo=google-chrome)](https://ltx.io)
[![Model](https://img.shields.io/badge/HuggingFace-Model-orange?logo=huggingface)](https://huggingface.co/Lightricks/LTX-2.3)
[![Demo](https://img.shields.io/badge/Demo-Try%20Now-brightgreen?logo=vercel)](https://app.ltx.studio/ltx-2-playground/i2v)
[![Demo](https://img.shields.io/badge/Demo-Try%20Now-brightgreen?logo=data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAABQAAAAUCAYAAACNiR0NAAAAAXNSR0IArs4c6QAAAERlWElmTU0AKgAAAAgAAYdpAAQAAAABAAAAGgAAAAAAA6ABAAMAAAABAAEAAKACAAQAAAABAAAAFKADAAQAAAABAAAAFAAAAACy3fD9AAACmElEQVQ4Ea1VP2haYRA/fRo0mESRIIqb2IwxuNUl0CGFQBC6OAWcikMottCpqYtDQIgdQsBFhAjZqiQhbhmySJBOgmNU0EGCg9r61Bivd0ffoykE0iQH37/77n7f3e/uqQFIPB7P/N3d3QeDwfAFEedZ91ghnyH5JM1m87dWq6UavF6vdTKZfDcajW/p4rE49+wIFMj33Gq1vlNo+kxg758KpiETqP/29vaXweVyqaS0aBfPXEfGFwTjWCwM+KBQoWA4HAJx/KDNvxcmTTGbzYAH8SljOp2C2+2GjY0NqNfrcHFxAXNzc2LDfCuKIq78KBdFOwsgGzidTnA4HHBzcwO9Xg8sFgtsbm7C3t4eVCoVaDQa0O12YXl5GUwmk5z5cZ/PB6PRCNrttgADFQUXFhbw8PAQVVXF3d1dJAeMx+P0zn0Jh8OYz+eRADCRSGAqlcLxeIz7+/u4tLSEjKUDZrNZ8U4mk0jR4fr6Op6enoru+voa0+k0rq2tYTAYxE6ng9QiSLRgrVZDv9+PFLkA6kUhT+GEC8C8XF5ewtHRkejICShiaDabwPvj42NJm3k7ODiQdDl9Fr0ocqJpdXUVIpEIdz7Y7XZRr6ysQDQahXK5LORvbW1p5rC9vQ2UifAooBqHuVxO0vt72tnZwWq1qqtisRgWCgU5ZzIZPDk50fdUUEmZvxTmAgKBgAxunT/fJpRKJWmhUCgEVDi4uroSG46kWCzC4uKitNVgMICzszOhSgA5fiJZhp4Lbbh1KARpbF65D/lx3vMdP05Vlkf5zKIDyukFJi7N6AVwNAhVsdlsM+LsjaZ56sq8kyQUqs4P6rsAKV49B4x4Padf7Y9Kv9+fEmiBQH8S4Gsa5v8EHpL9VwL7xH8BvwEcd4ccVf02KQAAAABJRU5ErkJggg==)](https://console.ltx.video/playground)
[![Paper](https://img.shields.io/badge/Paper-PDF-EC1C24?logo=adobeacrobatreader&logoColor=white)](https://arxiv.org/abs/2601.03233)
[![Discord](https://img.shields.io/badge/Join-Discord-5865F2?logo=discord)](https://discord.gg/ltxplatform)
@@ -72,7 +72,7 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https://
### ⚡ Optimization Tips
* **Use DistilledPipeline** - Fastest inference with only 8 predefined sigmas (8 steps stage 1, 4 steps stage 2)
* **Enable FP8 quantization** - Enables lower memory footprint: `--quantization fp8-cast` (CLI) or `quantization=QuantizationPolicy.fp8_cast()` (Python). For Hopper GPUs with TensorRT-LLM, use `--quantization fp8-scaled-mm` for FP8 scaled matrix multiplication.
* **Enable FP8 quantization** - Enables lower memory footprint: `--quantization fp8-cast` (CLI) or `quantization=QuantizationPolicy.fp8_cast()` (Python). Fp8-cast should be used with bf16 checkpoints, it shall downcast them on the fly. For Hopper GPUs with TensorRT-LLM, use `--quantization fp8-scaled-mm` for FP8 scaled matrix multiplication. Fp8-scaled-mm should be used with fp8 checkpoints.
* **Install attention optimizations** - Use xFormers (`uv sync --extra xformers`) or [Flash Attention 3](https://github.com/Dao-AILab/flash-attention) for Hopper GPUs
* **Use gradient estimation** - Reduce inference steps from 40 to 20-30 while maintaining quality (see [pipeline documentation](packages/ltx-pipelines/README.md#denoising-loop-optimization))
* **Skip memory cleanup** - If you have sufficient VRAM, disable automatic memory cleanup between stages for faster processing
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "ltx-core"
version = "1.0.0"
version = "1.1.0"
description = "Core implementation of Lightricks' LTX-2 model"
readme = "README.md"
requires-python = ">=3.10"
@@ -0,0 +1,95 @@
"""Batch-splitting adapter for the transformer.
Wraps an ``X0Model`` (or ``LayerStreamingWrapper``) and splits batched inputs
into smaller chunks before forwarding, then concatenates the results. This
controls peak activation memory at the cost of more forward passes.
The adapter is transparent — it has the same ``forward`` signature as
``X0Model`` and proxies attribute access to the wrapped model.
Example
-------
>>> from ltx_core.batch_split import BatchSplitAdapter
>>> adapter = BatchSplitAdapter(model, max_batch_size=1)
>>> # Receives B=4, runs 4xB=1 internally, returns B=4
>>> denoised_video, denoised_audio = adapter(video=v_b4, audio=a_b4, perturbations=ptb)
"""
from __future__ import annotations
from typing import Any
import torch
from torch import nn
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
from ltx_core.model.transformer.modality import Modality
def _split_perturbations(config: BatchedPerturbationConfig, sizes: list[int]) -> list[BatchedPerturbationConfig]:
"""Split a ``BatchedPerturbationConfig`` along the batch dimension."""
it = iter(config.perturbations)
return [BatchedPerturbationConfig([next(it) for _ in range(s)]) for s in sizes]
def _merge_tensors(tensors: list[torch.Tensor | None]) -> torch.Tensor | None:
"""Concatenate tensors along batch dim, or return None if all are None."""
non_none = [t for t in tensors if t is not None]
if not non_none:
return None
return torch.cat(non_none, dim=0)
class BatchSplitAdapter(nn.Module):
"""Wraps a model and splits batched forward calls into smaller chunks.
Has the same ``forward`` signature as ``X0Model``:
``(video, audio, perturbations) -> (denoised_video, denoised_audio)``.
Args:
model: The model to wrap (``X0Model``, ``LayerStreamingWrapper``, etc.).
max_batch_size: Maximum batch size per forward pass. Input batches
larger than this are split into sequential chunks.
"""
def __init__(self, model: nn.Module, max_batch_size: int) -> None:
if max_batch_size < 1:
raise ValueError(f"max_batch_size must be >= 1, got {max_batch_size}")
super().__init__()
self._model = model
self._max_batch_size = max_batch_size
def _get_chunk_sizes(self, batch_size: int) -> list[int]:
full, remainder = divmod(batch_size, self._max_batch_size)
sizes = [self._max_batch_size] * full
if remainder:
sizes.append(remainder)
return sizes
def forward(
self,
video: Modality | None,
audio: Modality | None,
perturbations: BatchedPerturbationConfig,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
batch_size = (video or audio).latent.shape[0]
if batch_size <= self._max_batch_size:
return self._model(video=video, audio=audio, perturbations=perturbations)
sizes = self._get_chunk_sizes(batch_size)
n = len(sizes)
v_chunks = video.split(sizes) if video is not None else [None] * n
a_chunks = audio.split(sizes) if audio is not None else [None] * n
p_chunks = _split_perturbations(perturbations, sizes)
chunk_results = [
self._model(video=vc, audio=ac, perturbations=pc)
for vc, ac, pc in zip(v_chunks, a_chunks, p_chunks, strict=True)
]
results_v, results_a = zip(*chunk_results, strict=True)
return _merge_tensors(list(results_v)), _merge_tensors(list(results_a))
def __getattr__(self, name: str) -> Any: # noqa: ANN401
"""Proxy attribute access to the wrapped model."""
try:
return super().__getattr__(name)
except AttributeError:
return getattr(self._model, name)
@@ -77,11 +77,22 @@ class Res2sDiffusionStep(DiffusionStepProtocol):
sigmas: torch.Tensor,
step_index: int,
noise: torch.Tensor,
eta: float = 0.5,
) -> torch.Tensor:
"""Advance one step with SDE noise injection via get_sde_coeff."""
"""Advance one step with SDE noise injection via get_sde_coeff.
Args:
sample: Current noisy sample.
denoised_sample: Denoised prediction from the model.
sigmas: Noise schedule tensor.
step_index: Current step index in the schedule.
noise: Random noise tensor for stochastic injection.
eta: Controls stochastic noise injection strength (0=deterministic, 1=maximum). Default 0.5.
Returns:
Next sample with SDE noise injection applied.
"""
sigma = sigmas[step_index]
sigma_next = sigmas[step_index + 1]
alpha_ratio, sigma_down, sigma_up = self.get_sde_coeff(sigma_next, sigma_up=sigma_next * 0.5)
alpha_ratio, sigma_down, sigma_up = self.get_sde_coeff(sigma_next, sigma_up=sigma_next * eta)
output_dtype = denoised_sample.dtype
if torch.any(sigma_up == 0) or torch.any(sigma_next == 0):
return denoised_sample
@@ -0,0 +1,45 @@
from dataclasses import dataclass
from ltx_core.components.patchifiers import get_pixel_coords
from ltx_core.conditioning.item import ConditioningItem
from ltx_core.tools import LatentTools, SpatioTemporalScaleFactors
from ltx_core.types import AudioLatentShape, LatentState, VideoLatentShape
@dataclass(frozen=True)
class TemporalRegionMask(ConditioningItem):
"""Conditioning item that sets ``denoise_mask = 0`` outside a time range
and ``1`` inside, so only the specified temporal region is regenerated.
Uses ``start_time`` and ``end_time`` in seconds. Works in *patchified*
(token) space using the patchifier's ``get_patch_grid_bounds``: for video
coords are latent frame indices (converted from seconds via ``fps``), for
audio coords are already in seconds.
"""
start_time: float # seconds, inclusive
end_time: float # seconds, exclusive
fps: float
def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState:
coords = latent_tools.patchifier.get_patch_grid_bounds(
latent_tools.target_shape, device=latent_state.denoise_mask.device
)
if isinstance(latent_tools.target_shape, AudioLatentShape):
# Audio: patchifier get_patch_grid_bounds returns seconds
t_boundaries = coords[:, 0]
elif isinstance(latent_tools.target_shape, VideoLatentShape):
# Video: patchifier get_patch_grid_bounds returns latent bounds, converting to frame numbers & pixel bounds
scale_factors = getattr(latent_tools, "scale_factors", SpatioTemporalScaleFactors.default())
pixel_bounds = get_pixel_coords(coords, scale_factors, causal_fix=getattr(latent_tools, "causal_fix", True))
# converting frame numbers to seconds
t_boundaries = pixel_bounds[:, 0] / self.fps
else:
raise ValueError("Unsupported LatentShape type, expected AudioLatentShape or VideoLatentShape")
t_start, t_end = t_boundaries.unbind(dim=-1) # [B, N]
in_region = (t_end > self.start_time) & (t_start < self.end_time)
state = latent_state.clone()
mask_val = in_region.to(state.denoise_mask.dtype)
if state.denoise_mask.dim() == 3:
mask_val = mask_val.unsqueeze(-1)
state.denoise_mask.copy_(mask_val)
return state
@@ -0,0 +1,324 @@
"""Layer streaming wrapper for memory-efficient inference.
Keeps most transformer/decoder layers on CPU pinned memory and streams them
to GPU on demand, using a secondary CUDA stream to prefetch upcoming layers
so that data transfer overlaps with compute.
General-purpose: works with any ``nn.Module`` whose forward iterates over a
``nn.ModuleList`` attribute (e.g. ``transformer_blocks``, ``layers``).
Each layer is evicted back to CPU immediately after its forward completes,
and prefetch uses modular indexing so the last layer's prefetch wraps around
to prepare early layers for the next forward pass.
Example
-------
>>> model = build_my_model(device=torch.device("cpu"))
>>> model = LayerStreamingWrapper(
... model,
... layers_attr="transformer_blocks",
... target_device=torch.device("cuda:0"),
... prefetch_count=2,
... )
>>> out = model(inputs) # hooks handle layer streaming
>>> model.teardown() # move everything back to CPU
"""
from __future__ import annotations
import functools
import itertools
import logging
from typing import Any
import torch
from torch import nn
logger = logging.getLogger(__name__)
def _resolve_attr(module: nn.Module, dotted_path: str) -> nn.ModuleList:
"""Resolve a dotted attribute path like ``'model.language_model.layers'``."""
obj: Any = module
for part in dotted_path.split("."):
obj = getattr(obj, part)
if not isinstance(obj, nn.ModuleList):
raise TypeError(f"Expected nn.ModuleList at '{dotted_path}', got {type(obj).__name__}")
return obj
class _LayerStore:
"""Manages on-demand pinning of layer parameters for GPU streaming.
Stores references to each layer's source data (which may be file-backed
mmap views or in-memory tensors). When a layer needs to be transferred
to GPU, its source data is pinned on demand and copied; on eviction the
pinned copy is freed and the source data is restored.
"""
def __init__(self, layers: nn.ModuleList, target_device: torch.device) -> None:
self.target_device = target_device
self.num_layers = len(layers)
self._on_gpu: set[int] = set()
# Keep a reference to the source data for each layer so we can pin it
# on demand and restore it after eviction.
self._source_data: list[dict[str, torch.Tensor]] = []
for layer in layers:
source: dict[str, torch.Tensor] = {}
for name, tensor in itertools.chain(layer.named_parameters(), layer.named_buffers()):
source[name] = tensor.data
self._source_data.append(source)
# Hold pinned tensors alive until the H2D transfer completes.
# Without this, the CachingHostAllocator can reclaim a pinned tensor
# as soon as its Python reference is dropped, even if an async H2D
# transfer is still reading from it.
self._pinned_in_flight: dict[int, list[torch.Tensor]] = {}
def _check_idx(self, idx: int) -> None:
if idx < 0 or idx >= self.num_layers:
raise IndexError(f"Layer index {idx} out of range [0, {self.num_layers})")
def is_on_gpu(self, idx: int) -> bool:
return idx in self._on_gpu
def move_to_gpu(self, idx: int, layer: nn.Module, *, non_blocking: bool = False) -> None:
"""Pin layer *idx* on demand, then transfer to GPU."""
self._check_idx(idx)
if idx in self._on_gpu:
return
source = self._source_data[idx]
pinned_refs: list[torch.Tensor] = []
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
pinned = source[name].pin_memory()
param.data = pinned.to(self.target_device, non_blocking=non_blocking)
pinned_refs.append(pinned)
# Keep pinned tensors alive until eviction — the async H2D transfer
# may still be reading from them.
self._pinned_in_flight[idx] = pinned_refs
self._on_gpu.add(idx)
def evict_to_cpu(self, idx: int, layer: nn.Module) -> None:
"""Restore source data, freeing the GPU and pinned copies."""
self._check_idx(idx)
if idx not in self._on_gpu:
return
source = self._source_data[idx]
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
param.data = source[name]
# Release pinned tensors — the H2D transfer is complete by now
# (the compute stream waited on the prefetch event before using
# the layer, and we only evict after compute finishes).
self._pinned_in_flight.pop(idx, None)
self._on_gpu.discard(idx)
def cleanup(self) -> None:
"""Release all source data and in-flight pinned references.
After this call, the source tensors can be garbage-collected once
the layer parameters (which still reference them via ``.data``) are
also released (e.g. via ``.to("meta")``).
"""
for source_dict in self._source_data:
source_dict.clear()
self._source_data.clear()
self._pinned_in_flight.clear()
class _AsyncPrefetcher:
"""Issues H2D transfers on a dedicated CUDA stream.
Uses per-layer CUDA events so that the compute stream only waits for the
specific layer it needs, not all pending transfers.
"""
def __init__(self, store: _LayerStore, layers: nn.ModuleList) -> None:
self._store = store
self._layers = layers
self._stream = torch.cuda.Stream(device=store.target_device)
self._events: dict[int, torch.cuda.Event] = {}
def prefetch(self, idx: int) -> None:
"""Begin async transfer of layer *idx* to GPU (no-op if already there)."""
if self._store.is_on_gpu(idx) or idx in self._events:
return
with torch.cuda.stream(self._stream):
self._store.move_to_gpu(idx, self._layers[idx], non_blocking=True)
event = torch.cuda.Event()
event.record(self._stream)
self._events[idx] = event
def wait(self, idx: int) -> None:
"""Block the compute stream until layer *idx* transfer is complete."""
event = self._events.pop(idx, None)
if event is not None:
torch.cuda.current_stream(self._store.target_device).wait_event(event)
def cleanup(self) -> None:
"""Drain pending work and release CUDA stream/event resources."""
self._events.clear()
self._stream = None
self._layers = None
self._store = None
class LayerStreamingWrapper(nn.Module):
"""Wraps a model to stream its sequential layers between CPU and GPU.
Each layer is evicted immediately after its forward completes, and
prefetch wraps around using modular indexing so the end of one forward
pass prepares early layers for the next.
Parameters
----------
model:
The model to wrap, with all parameters on **CPU**.
layers_attr:
Dotted attribute path to the ``nn.ModuleList`` of sequential layers
(e.g. ``"transformer_blocks"`` or ``"model.language_model.layers"``).
target_device:
The GPU device to use for compute.
prefetch_count:
How many layers ahead to prefetch. The maximum number of layers on
GPU at once is ``1 + prefetch_count``. Must be >= 1.
"""
def __init__(
self,
model: nn.Module,
layers_attr: str,
target_device: torch.device,
prefetch_count: int = 2,
) -> None:
if prefetch_count < 1:
raise ValueError("prefetch_count must be >= 1")
super().__init__()
# Store the wrapped model as a submodule so parameters are discoverable.
self._model = model
self._layers = _resolve_attr(model, layers_attr)
self._target_device = target_device
# Clamp: no point prefetching more than num_layers - 1 (the rest are evicted).
self._prefetch_count = min(prefetch_count, len(self._layers) - 1)
self._hooks: list[torch.utils.hooks.RemovableHandle] = []
self._setup()
# ------------------------------------------------------------------
# Setup / teardown
# ------------------------------------------------------------------
def _setup(self) -> None:
# 1. Build the pinned CPU store (copies all layer tensors to pinned memory).
self._store = _LayerStore(self._layers, self._target_device)
# 2. Move all NON-layer params/buffers to GPU.
layer_tensor_ids: set[int] = set()
for layer in self._layers:
for t in itertools.chain(layer.parameters(), layer.buffers()):
layer_tensor_ids.add(id(t))
for p in self._model.parameters():
if id(p) not in layer_tensor_ids:
p.data = p.data.to(self._target_device)
for b in self._model.buffers():
if id(b) not in layer_tensor_ids:
b.data = b.data.to(self._target_device)
# 3. Pre-load the first (1 + prefetch_count) layers synchronously.
for idx in range(min(self._prefetch_count + 1, len(self._layers))):
self._store.move_to_gpu(idx, self._layers[idx])
# 4. Create the async prefetcher and register hooks.
self._prefetcher = _AsyncPrefetcher(self._store, self._layers)
self._register_hooks()
def _register_hooks(self) -> None:
idx_map: dict[int, int] = {id(layer): idx for idx, layer in enumerate(self._layers)}
num_layers = len(self._layers)
compute_stream = torch.cuda.current_stream(self._target_device)
def _pre_hook(
module: nn.Module,
_args: Any, # noqa: ANN401
*,
idx: int,
) -> None:
# Wait only for THIS layer's H2D transfer (not all pending ones).
self._prefetcher.wait(idx)
if not self._store.is_on_gpu(idx):
self._store.move_to_gpu(idx, module)
# Record that the compute stream will read these weight tensors.
# They were allocated on the prefetch stream, so without this the
# caching allocator would allow the prefetch stream to reuse their
# memory immediately after eviction — even if the compute kernel
# that reads them hasn't finished yet.
for param in itertools.chain(module.parameters(), module.buffers()):
param.data.record_stream(compute_stream)
# Kick off prefetch for upcoming layers (wraps around for next pass).
for offset in range(1, self._prefetch_count + 1):
self._prefetcher.prefetch((idx + offset) % num_layers)
def _post_hook(
module: nn.Module,
_args: Any, # noqa: ANN401
_output: Any, # noqa: ANN401
*,
idx: int,
) -> None:
# Evict this layer immediately — its computation is done.
self._store.evict_to_cpu(idx, module)
for layer in self._layers:
idx = idx_map[id(layer)]
h1 = layer.register_forward_pre_hook(functools.partial(_pre_hook, idx=idx))
h2 = layer.register_forward_hook(functools.partial(_post_hook, idx=idx))
self._hooks.extend([h1, h2])
def teardown(self) -> None:
"""Remove hooks, release resources, and move parameters back to CPU.
After this call the wrapper is inert: hooks are removed, the prefetch
stream is drained and destroyed, all parameters reside on CPU, and the
``_LayerStore`` source data references are cleared. Callers should
still follow up with ``.to("meta")`` to release the CPU copies if the
model is no longer needed.
"""
for h in self._hooks:
h.remove()
self._hooks.clear()
# Drain all in-flight async H2D copies, then release stream resources.
# Without the synchronize, clearing the stream/events can trigger
# use-after-free at the CUDA driver level.
torch.cuda.synchronize(device=self._target_device)
if self._prefetcher is not None:
self._prefetcher.cleanup()
self._prefetcher = None
# Move everything to CPU.
for idx, layer in enumerate(self._layers):
self._store.evict_to_cpu(idx, layer)
for p in self._model.parameters():
p.data = p.data.to("cpu")
for b in self._model.buffers():
b.data = b.data.to("cpu")
# Release source data references. After evict_to_cpu() the layer
# params point to the source data. The caller is expected to follow
# up with .to("meta") to drop the param refs; cleanup() drops the
# store's refs.
self._store.cleanup()
# ------------------------------------------------------------------
# Forward and attribute delegation
# ------------------------------------------------------------------
def forward(self, *args: Any, **kwargs: Any) -> Any: # noqa: ANN401
return self._model(*args, **kwargs)
def __getattr__(self, name: str) -> Any: # noqa: ANN401
"""Proxy attribute access to the wrapped model.
This allows calling methods like ``encode()`` on a wrapped
GemmaTextEncoder without the caller needing to know about the wrapper.
``nn.Module.__getattr__`` is only called when normal attribute lookup
fails, so ``_model``, ``_store``, etc. are found first via ``__dict__``.
"""
try:
return super().__getattr__(name)
except AttributeError:
return getattr(self._model, name)
@@ -1,46 +1,71 @@
from collections.abc import Iterator
import torch
from ltx_core.loader.primitives import LoraStateDictWithStrength, StateDict
from ltx_core.quantization.fp8_cast import calculate_weight_float8
from ltx_core.quantization.fp8_cast import _fused_add_round_launch
from ltx_core.quantization.fp8_scaled_mm import quantize_weight_to_fp8_per_tensor
def _get_device() -> torch.device:
if torch.cuda.is_available():
return torch.device("cuda", torch.cuda.current_device())
return torch.device("cpu")
def fuse_lora_weights(
model_sd: StateDict,
lora_sd_and_strengths: list[LoraStateDictWithStrength],
dtype: torch.dtype | None = None,
) -> Iterator[tuple[str, torch.Tensor]]:
"""Yield ``(key, fused_tensor)`` for each weight modified by at least one LoRA.
For scaled-FP8 weights, this includes both the updated ``.weight`` tensor
and its corresponding ``.weight_scale`` tensor.
"""
for key, original_weight in model_sd.sd.items():
if original_weight is None or key.endswith(".weight_scale"):
continue
original_device = original_weight.device
weight = original_weight.to(device=_get_device())
target_dtype = dtype if dtype is not None else weight.dtype
deltas_dtype = target_dtype if target_dtype not in [torch.float8_e4m3fn, torch.float8_e5m2] else torch.bfloat16
deltas = _prepare_deltas(lora_sd_and_strengths, key, deltas_dtype, weight.device)
if deltas is None:
continue
scale_key = key.replace(".weight", ".weight_scale") if key.endswith(".weight") else None
is_scaled_fp8 = scale_key is not None and scale_key in model_sd.sd
if weight.dtype == torch.float8_e4m3fn:
if is_scaled_fp8:
fused = _fuse_delta_with_scaled_fp8(deltas, weight, key, scale_key, model_sd)
else:
fused = _fuse_delta_with_cast_fp8(deltas, weight, key, target_dtype)
elif weight.dtype == torch.bfloat16:
fused = _fuse_delta_with_bfloat16(deltas, weight, key, target_dtype)
else:
raise ValueError(f"Unsupported dtype: {weight.dtype}")
for k, v in fused.items():
yield k, v.to(device=original_device)
def apply_loras(
model_sd: StateDict,
lora_sd_and_strengths: list[LoraStateDictWithStrength],
dtype: torch.dtype | None = None,
destination_sd: StateDict | None = None,
) -> StateDict:
sd = {}
if destination_sd is not None:
sd = destination_sd.sd
size = 0
device = torch.device("meta")
inner_dtypes = set()
for key, weight in model_sd.sd.items():
if weight is None:
continue
# Skip scale keys - they are handled together with their weight keys
if key.endswith(".weight_scale"):
continue
device = weight.device
target_dtype = dtype if dtype is not None else weight.dtype
deltas_dtype = target_dtype if target_dtype not in [torch.float8_e4m3fn, torch.float8_e5m2] else torch.bfloat16
scale_key = key.replace(".weight", ".weight_scale") if key.endswith(".weight") else None
is_scaled_fp8 = scale_key is not None and scale_key in model_sd.sd
deltas = _prepare_deltas(lora_sd_and_strengths, key, deltas_dtype, device)
fused = _fuse_deltas(deltas, weight, key, sd, target_dtype, device, is_scaled_fp8, scale_key, model_sd)
sd.update(fused)
for tensor in fused.values():
inner_dtypes.add(tensor.dtype)
size += tensor.nbytes
if destination_sd is not None:
for key, tensor in fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype):
sd[key] = tensor
return destination_sd
return StateDict(sd, device, size, inner_dtypes)
fused = dict(fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype))
sd = {k: (fused[k] if k in fused else v.clone()) for k, v in model_sd.sd.items()}
return StateDict(sd, model_sd.device, model_sd.size, model_sd.dtype)
def _prepare_deltas(
@@ -65,50 +90,6 @@ def _prepare_deltas(
return torch.sum(torch.stack(deltas, dim=0), dim=0)
def _fuse_deltas(
deltas: torch.Tensor | None,
weight: torch.Tensor,
key: str,
sd: dict[str, torch.Tensor],
target_dtype: torch.dtype,
device: torch.device,
is_scaled_fp8: bool,
scale_key: str | None,
model_sd: StateDict,
) -> dict[str, torch.Tensor]:
if deltas is None:
if key in sd:
return {}
fused = _copy_weight_without_lora(weight, key, target_dtype, device, is_scaled_fp8, scale_key, model_sd)
elif weight.dtype == torch.float8_e4m3fn:
if is_scaled_fp8:
fused = _fuse_delta_with_scaled_fp8(deltas, weight, key, scale_key, model_sd)
else:
fused = _fuse_delta_with_cast_fp8(deltas, weight, key, target_dtype, device)
elif weight.dtype == torch.bfloat16:
fused = _fuse_delta_with_bfloat16(deltas, weight, key, target_dtype)
else:
raise ValueError(f"Unsupported dtype: {weight.dtype}")
return fused
def _copy_weight_without_lora(
weight: torch.Tensor,
key: str,
target_dtype: torch.dtype,
device: torch.device,
is_scaled_fp8: bool,
scale_key: str | None,
model_sd: StateDict,
) -> dict[str, torch.Tensor]:
"""Copy original weight (and scale if applicable) when no LoRA affects this key."""
result = {key: weight.clone().to(dtype=target_dtype, device=device)}
if is_scaled_fp8:
result[scale_key] = model_sd.sd[scale_key].clone()
return result
def _fuse_delta_with_scaled_fp8(
deltas: torch.Tensor,
weight: torch.Tensor,
@@ -132,13 +113,12 @@ def _fuse_delta_with_cast_fp8(
weight: torch.Tensor,
key: str,
target_dtype: torch.dtype,
device: torch.device,
) -> dict[str, torch.Tensor]:
"""Fuse LoRA delta with cast-only FP8 weight (no scale factor)."""
if str(device).startswith("cuda"):
deltas = calculate_weight_float8(deltas, weight)
if str(weight.device).startswith("cuda"):
_fused_add_round_launch(deltas, weight, seed=0)
else:
deltas.add_(weight.to(dtype=deltas.dtype, device=device))
deltas.add_(weight.to(dtype=deltas.dtype))
return {key: deltas.to(dtype=target_dtype)}
@@ -1,5 +1,7 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import NamedTuple, Protocol
from typing import TYPE_CHECKING, NamedTuple, Protocol
import torch
@@ -7,6 +9,9 @@ from ltx_core.loader.module_ops import ModuleOps
from ltx_core.loader.sd_ops import SDOps
from ltx_core.model.model_protocol import ModelType
if TYPE_CHECKING:
from ltx_core.loader.registry import Registry
@dataclass(frozen=True)
class StateDict:
@@ -55,6 +60,11 @@ class ModelBuilderProtocol(Protocol[ModelType]):
- build: Create and initialize a model from state dictionary and apply dtype transformations
"""
model_sd_ops: SDOps | None
module_ops: tuple[ModuleOps, ...]
loras: tuple["LoraPathStrengthAndSDOps", ...]
registry: "Registry"
def meta_model(self, config: dict, module_ops: list[ModuleOps] | None = None) -> ModelType:
"""
Create a model on the meta device from a configuration dictionary.
@@ -68,16 +78,43 @@ class ModelBuilderProtocol(Protocol[ModelType]):
"""
...
def build(self, dtype: torch.dtype | None = None) -> ModelType:
def with_sd_ops(self, sd_ops: SDOps | None) -> "ModelBuilderProtocol[ModelType]":
"""Return a copy of this builder with the given state-dict key remapping ops."""
...
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> "ModelBuilderProtocol[ModelType]":
"""Return a copy of this builder with the given module operations (e.g. quantization)."""
...
def with_loras(self, loras: tuple["LoraPathStrengthAndSDOps", ...]) -> "ModelBuilderProtocol[ModelType]":
"""Return a copy of this builder with the given LoRAs to fuse at build time."""
...
def with_registry(self, registry: "Registry") -> "ModelBuilderProtocol[ModelType]":
"""Return a copy of this builder using the given weight registry for allocation."""
...
def with_lora_load_device(self, device: torch.device) -> "ModelBuilderProtocol[ModelType]":
"""Return a copy of this builder that loads LoRA weights onto the given device."""
...
def build(
self, device: torch.device | None = None, dtype: torch.dtype | None = None, **kwargs: object
) -> ModelType:
"""
Build the model
Args:
device: Target device for the model
dtype: Target dtype for the model, if None, uses the dtype of the model_path model
Returns:
Model instance
"""
...
def model_config(self) -> dict:
"""Return the model configuration dictionary extracted from the checkpoint metadata."""
...
class LoRAAdaptableProtocol(Protocol):
"""
@@ -64,6 +64,7 @@ class SDOps:
mapping: tuple[
ContentReplacement | ContentMatching | SDKeyValueOperation, ...
] = () # Immutable tuple of (key, value) pairs
allowed_keys: frozenset[str] | None = None
def with_replacement(self, content: str, replacement: str) -> "SDOps":
"""Create a new SDOps instance with the specified replacement added to the mapping."""
@@ -77,6 +78,13 @@ class SDOps:
new_mapping = (*self.mapping, ContentMatching(prefix, suffix))
return replace(self, mapping=new_mapping)
def with_additional_allowed_keys(self, keys: frozenset[str]) -> "SDOps":
"""Create a new SDOps instance that only passes keys present in *keys* (post-replacement).
If allowed_keys already exists, the sets are merged via union.
"""
merged = frozenset(keys) | self.allowed_keys if self.allowed_keys is not None else frozenset(keys)
return replace(self, allowed_keys=merged)
def with_kv_operation(
self,
operation: KeyValueOperation,
@@ -101,6 +109,10 @@ class SDOps:
continue
if replacement.content in key:
key = key.replace(replacement.content, replacement.replacement)
if self.allowed_keys is not None and key not in self.allowed_keys:
return None
return key
def apply_to_key_value(self, key: str, value: torch.Tensor) -> list[KeyValueOperationResult]:
@@ -53,6 +53,21 @@ class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType],
def lora(self, lora_path: str, strength: float = 1.0, sd_ops: SDOps | None = None) -> "SingleGPUModelBuilder":
return replace(self, loras=(*self.loras, LoraPathStrengthAndSDOps(lora_path, strength, sd_ops)))
def with_sd_ops(self, sd_ops: SDOps | None) -> "SingleGPUModelBuilder":
return replace(self, model_sd_ops=sd_ops)
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> "SingleGPUModelBuilder":
return replace(self, module_ops=module_ops)
def with_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> "SingleGPUModelBuilder":
return replace(self, loras=loras)
def with_registry(self, registry: Registry) -> "SingleGPUModelBuilder":
return replace(self, registry=registry)
def with_lora_load_device(self, device: torch.device) -> "SingleGPUModelBuilder":
return replace(self, lora_load_device=device)
def model_config(self) -> dict:
first_shard_path = self.model_path[0] if isinstance(self.model_path, tuple) else self.model_path
return self.model_loader.metadata(first_shard_path)
@@ -83,7 +98,12 @@ class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType],
retval = meta_model.to(device)
return retval
def build(self, device: torch.device | None = None, dtype: torch.dtype | None = None) -> ModelType:
def build(
self,
device: torch.device | None = None,
dtype: torch.dtype | None = None,
**kwargs: object, # noqa: ARG002
) -> ModelType:
device = torch.device("cuda") if device is None else device
config = self.model_config()
meta_model = self.meta_model(config, self.module_ops)
@@ -0,0 +1,222 @@
"""Video modality tiling helpers.
Provides :class:`VideoModalityTilingHelper` — a stateless helper that
tiles and blends video :class:`Modality` token sequences by
spatial/temporal region. Tile geometry is represented by the existing
:class:`Tile` NamedTuple from :mod:`ltx_core.tiling`; no distributed
primitives are required.
"""
from __future__ import annotations
from dataclasses import dataclass, replace
import torch
from ltx_core.model.transformer.modality import Modality
from ltx_core.tiling import Tile, TileCountConfig, create_tiles, identity_mapping_operation, split_by_count
from ltx_core.tools import VideoLatentTools
from ltx_core.types import VideoLatentShape
@dataclass(frozen=True)
class TilingContext:
"""Opaque context produced by :meth:`VideoModalityTilingHelper.tile_modality`.
Carries the token-level keep mask and per-conditioning-token blend
weights needed by :meth:`~VideoModalityTilingHelper.blend`.
"""
keep_mask: torch.Tensor
cond_blend_weights: torch.Tensor | None
"""``(num_kept_cond,)`` — weight for each kept conditioning token,
equal to ``1 / num_tiles_that_keep_this_token``. ``None`` when
there are no conditioning tokens."""
class VideoModalityTilingHelper:
"""Stateless helper that tiles and blends video :class:`Modality` sequences.
Constructed once with a :class:`TileCountConfig` and
:class:`VideoLatentTools`. Tiles are computed at construction and
available via the :attr:`tiles` property. Use :meth:`tile_modality`
and :meth:`blend` with any tile from that list.
Usage::
helper = VideoModalityTilingHelper(tiling, video_tools)
for tile in helper.tiles:
tiled_mod, ctx = helper.tile_modality(modality, tile)
result = run_model(tiled_mod)
helper.blend(result, tile, ctx, output=output)
"""
def __init__(self, tiling: TileCountConfig, video_tools: VideoLatentTools) -> None:
self._patchifier = video_tools.patchifier
self._latent_shape = video_tools.target_shape
self._num_generated_tokens = self._patchifier.get_token_count(self._latent_shape)
self._tiles = create_tiles(
torch.Size([self._latent_shape.frames, self._latent_shape.height, self._latent_shape.width]),
splitters=[
split_by_count(tiling.frames.num_tiles, tiling.frames.overlap),
split_by_count(tiling.height.num_tiles, tiling.height.overlap),
split_by_count(tiling.width.num_tiles, tiling.width.overlap),
],
mappers=[identity_mapping_operation] * 3,
)
@property
def tiles(self) -> list[Tile]:
"""All tiles for the configured tiling layout."""
return self._tiles
# -- tile modality -----------------------------------------------------
def tile_modality(self, modality: Modality, tile: Tile) -> tuple[Modality, TilingContext]:
"""Slice *modality* to the tokens covered by *tile*.
Selects generated tokens belonging to the tile's spatial region
and conditioning tokens that overlap with the tile (or have
negative time coordinates).
Returns:
A ``(tiled_modality, context)`` tuple. Pass *context* to
:meth:`blend` together with the model output.
"""
keep_mask = self._keep_mask(modality, tile)
tile_attention_mask = None
if modality.attention_mask is not None:
keep_indices = keep_mask.nonzero(as_tuple=False).squeeze(1)
tile_attention_mask = modality.attention_mask[:, keep_indices, :][:, :, keep_indices]
tiled = replace(
modality,
latent=modality.latent[:, keep_mask, :],
timesteps=modality.timesteps[:, keep_mask],
positions=modality.positions[:, :, keep_mask, :],
attention_mask=tile_attention_mask,
)
cond_blend_weights = None
num_total = modality.latent.shape[1]
if num_total > self._num_generated_tokens:
cond_keep = keep_mask[self._num_generated_tokens :]
# Count how many tiles keep each conditioning token.
cond_counts = torch.zeros(cond_keep.sum(), dtype=torch.float32)
for t in self._tiles:
other_mask = self._keep_mask(modality, t)
other_cond = other_mask[self._num_generated_tokens :]
# Map other tile's kept cond tokens into this tile's kept subset.
cond_counts += other_cond[cond_keep].float()
cond_blend_weights = 1.0 / cond_counts
return tiled, TilingContext(keep_mask=keep_mask, cond_blend_weights=cond_blend_weights)
# -- blend -------------------------------------------------------------
def blend(
self,
tile_to_blend: torch.Tensor,
tile: Tile,
context: TilingContext,
output: torch.Tensor | None = None,
) -> torch.Tensor:
"""Blend-weight tile results and accumulate into the full token space.
Premultiplied (blend-weighted) data is **added** to *output*,
allowing multiple tiles to be accumulated into the same buffer.
Args:
tile_to_blend: Denoised tile tensor ``(B, num_tile_tokens, D)``,
where the first ``_tile_generated_token_count(tile)``
entries are generated tokens and the remainder are
conditioning tokens.
tile: The :class:`Tile` that was used in :meth:`tile_modality`.
context: The :class:`TilingContext` returned by :meth:`tile_modality`.
output: Optional pre-allocated output tensor. When provided
its shape must be ``(B, num_total_tokens, D)`` and the
blended tile is **added** into it. When ``None`` a new
zero-filled tensor is created.
Returns:
The output tensor with the blended tile added at the correct
positions.
"""
batch, _, dim = tile_to_blend.shape
num_tile_gen = self._tile_generated_token_count(tile)
gen_indices = self._generated_token_indices(tile)
num_total_tokens = context.keep_mask.shape[0]
expected_shape = (batch, num_total_tokens, dim)
if output is not None:
if output.shape != expected_shape:
raise ValueError(f"Expected output shape {expected_shape}, got {output.shape}")
result = output
else:
result = torch.zeros(*expected_shape, device=tile_to_blend.device, dtype=tile_to_blend.dtype)
# Blend mask is (tile_F, tile_H, tile_W) — one weight per token in row-major order.
blend_weights = tile.blend_mask.reshape(-1).to(device=tile_to_blend.device, dtype=tile_to_blend.dtype)
tile_gen = tile_to_blend[:, :num_tile_gen, :] * blend_weights[None, :, None]
result[:, gen_indices, :] += tile_gen
# Scatter kept conditioning tokens, weighted by 1/N where N is
# the number of tiles that keep each token (so they sum to 1).
if num_total_tokens > self._num_generated_tokens and context.cond_blend_weights is not None:
cond_keep = context.keep_mask[self._num_generated_tokens :]
cond_indices = self._num_generated_tokens + cond_keep.nonzero(as_tuple=False).squeeze(1)
weights = context.cond_blend_weights.to(device=tile_to_blend.device, dtype=tile_to_blend.dtype)
result[:, cond_indices, :] += tile_to_blend[:, num_tile_gen:, :] * weights[None, :, None]
return result
# -- private -----------------------------------------------------------
def _tile_generated_token_count(self, tile: Tile) -> int:
"""Number of generated tokens in *tile*."""
frame_slice, height_slice, width_slice = tile.in_coords
tile_shape = VideoLatentShape(
batch=self._latent_shape.batch,
channels=self._latent_shape.channels,
frames=frame_slice.stop - frame_slice.start,
height=height_slice.stop - height_slice.start,
width=width_slice.stop - width_slice.start,
)
return self._patchifier.get_token_count(tile_shape)
def _generated_token_indices(self, tile: Tile) -> torch.Tensor:
"""Flat token indices of *tile*'s generated tokens in the full sequence."""
frame_slice, height_slice, width_slice = tile.in_coords
f = torch.arange(frame_slice.start, frame_slice.stop)
h = torch.arange(height_slice.start, height_slice.stop)
w = torch.arange(width_slice.start, width_slice.stop)
return (
f[:, None, None] * self._latent_shape.height * self._latent_shape.width
+ h[None, :, None] * self._latent_shape.width
+ w[None, None, :]
).reshape(-1)
def _keep_mask(self, modality: Modality, tile: Tile) -> torch.Tensor:
"""Boolean mask ``(num_total_tokens,)`` — True for tokens the tile processes.
Generated tokens are selected by grid position. Conditioning
tokens are kept when their ``[start, end)`` intervals overlap
the tile in all three dimensions, or when they have a negative
time coordinate (reference tokens).
"""
num_total = modality.latent.shape[1]
mask = torch.zeros(num_total, dtype=torch.bool)
gen_indices = self._generated_token_indices(tile)
mask[gen_indices] = True
if num_total > self._num_generated_tokens:
gen_positions = modality.positions[:, :, gen_indices, :] # (B, 3, num_tile_gen, 2)
tile_start = gen_positions[..., 0].amin(dim=2) # (B, 3)
tile_end = gen_positions[..., 1].amax(dim=2) # (B, 3)
cond_positions = modality.positions[:, :, self._num_generated_tokens :, :] # (B, 3, num_cond, 2)
overlaps = (cond_positions[..., 0] < tile_end.unsqueeze(2)) & (
cond_positions[..., 1] > tile_start.unsqueeze(2)
) # (B, 3, num_cond)
overlaps_all_dims = overlaps.all(dim=1) # (B, num_cond)
has_negative_time = cond_positions[:, 0, :, 0] < 0 # (B, num_cond)
keep_cond = (overlaps_all_dims | has_negative_time).any(dim=0) # (num_cond,)
mask[self._num_generated_tokens :] = keep_cond
return mask
@@ -500,6 +500,8 @@ class VocoderWithBWE(nn.Module):
to a higher sample rate. The BWE computes a mel spectrogram from the
vocoder output, runs it through a second generator to predict a residual,
and adds it to a sinc-resampled skip connection.
The forward pass runs in fp32 via autocast to avoid bfloat16 accumulation
errors that degrade spectral metrics by 40-90%.
"""
def __init__(
@@ -548,28 +550,45 @@ class VocoderWithBWE(nn.Module):
def forward(self, mel_spec: torch.Tensor) -> torch.Tensor:
"""Run the full vocoder + BWE forward pass.
Runs in float32 regardless of weight or input dtype. bfloat16 arithmetic
causes 40-90% spectral metric degradation due to accumulation errors
compounding through 108 sequential convolutions in the BigVGAN v2 architecture.
Args:
mel_spec: Mel spectrogram of shape (B, 2, T, mel_bins) for stereo
or (B, T, mel_bins) for mono. Same format as Vocoder.forward.
Returns:
Waveform tensor of shape (B, out_channels, T_out) clipped to [-1, 1].
"""
x = self.vocoder(mel_spec)
_, _, length_low_rate = x.shape
output_length = length_low_rate * self.output_sampling_rate // self.input_sampling_rate
input_dtype = mel_spec.dtype
# Run the entire forward pass in fp32. bfloat16 accumulation errors
# compound through 108 sequential convolutions and degrade spectral
# metrics (mel_l1, MRSTFT) by 40-90% while perceptual quality (CDPAM)
# is unaffected. fp32 eliminates this degradation.
# We use autocast(dtype=float32) rather than self.float() because it
# upcasts bf16 weights per-op at kernel level, avoiding the temporary
# memory spike of self.float() / self.to(original_dtype).
# Benchmarked on H100 (128.5M-param model):
# autocast fp32: +70 MB peak VRAM, 123 ms (vs 482 MB / 95 ms for bf16)
# model.float(): +324 MB peak VRAM, 149 ms
# Tested: both approaches produce bit-identical output.
# Pad to multiple of hop_length for exact mel frame count
remainder = length_low_rate % self.hop_length
if remainder != 0:
x = F.pad(x, (0, self.hop_length - remainder))
with torch.autocast(device_type=mel_spec.device.type, dtype=torch.float32):
x = self.vocoder(mel_spec.float())
_, _, length_low_rate = x.shape
output_length = length_low_rate * self.output_sampling_rate // self.input_sampling_rate
# Compute mel spectrogram from vocoder output: (B, C, n_mels, T_frames)
mel = self._compute_mel(x)
# Pad to multiple of hop_length for exact mel frame count
remainder = length_low_rate % self.hop_length
if remainder != 0:
x = F.pad(x, (0, self.hop_length - remainder))
# Vocoder.forward expects (B, C, T, mel_bins) — transpose before calling bwe_generator
mel_for_bwe = mel.transpose(2, 3) # (B, C, T_frames, mel_bins)
residual = self.bwe_generator(mel_for_bwe)
skip = self.resampler(x)
assert residual.shape == skip.shape, f"residual {residual.shape} != skip {skip.shape}"
# Compute mel spectrogram from vocoder output: (B, C, n_mels, T_frames)
mel = self._compute_mel(x)
return torch.clamp(residual + skip, -1, 1)[..., :output_length]
# Vocoder.forward expects (B, C, T, mel_bins) — transpose before calling bwe_generator
mel_for_bwe = mel.transpose(2, 3) # (B, C, T_frames, mel_bins)
residual = self.bwe_generator(mel_for_bwe)
skip = self.resampler(x)
assert residual.shape == skip.shape, f"residual {residual.shape} != skip {skip.shape}"
return torch.clamp(residual + skip, -1, 1)[..., :output_length].to(input_dtype)
@@ -122,22 +122,18 @@ class AttentionFunction(Enum):
FLASH_ATTENTION_3 = "flash_attention_3"
DEFAULT = "default"
def __call__(
self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, heads: int, mask: torch.Tensor | None = None
) -> torch.Tensor:
def to_callable(self) -> AttentionCallable:
"""Resolve to a concrete callable. Use this at module init time so that
torch.compile can trace through the attention call without graph breaks."""
if self is AttentionFunction.PYTORCH:
return PytorchAttention()(q, k, v, heads, mask)
return PytorchAttention()
elif self is AttentionFunction.XFORMERS:
return XFormersAttention()(q, k, v, heads, mask)
return XFormersAttention()
elif self is AttentionFunction.FLASH_ATTENTION_3:
return FlashAttention3()(q, k, v, heads, mask)
return FlashAttention3()
else:
# Default behavior: XFormers if installed else - PyTorch
return (
XFormersAttention()(q, k, v, heads, mask)
if memory_efficient_attention is not None
else PytorchAttention()(q, k, v, heads, mask)
)
return XFormersAttention() if memory_efficient_attention is not None else PytorchAttention()
class Attention(torch.nn.Module):
@@ -154,7 +150,11 @@ class Attention(torch.nn.Module):
) -> None:
super().__init__()
self.rope_type = rope_type
self.attention_function = attention_function
self.attention_function = (
attention_function.to_callable()
if isinstance(attention_function, AttentionFunction)
else attention_function
)
inner_dim = dim_head * heads
context_dim = query_dim if context_dim is None else context_dim
@@ -0,0 +1,37 @@
import torch
from ltx_core.loader.module_ops import ModuleOps
from ltx_core.loader.sd_ops import SDOps
from ltx_core.model.transformer.model import LTXModel
def compile_transformer(model: LTXModel) -> LTXModel:
model.transformer_blocks = torch.nn.ModuleList(torch.compile(m) for m in model.transformer_blocks)
def patched_dynamo_forward(*args, **kwargs) -> tuple[torch.Tensor, torch.Tensor]:
with (
torch._inductor.config.patch(unsafe_skip_cache_dynamic_shape_guards=True),
torch._dynamo.config.patch( # type: ignore[attr-defined]
inline_inbuilt_nn_modules=True, cache_size_limit=256, allow_unspec_int_on_nn_module=True
),
):
return model.forward_without_compilation(*args, **kwargs)
model.forward_without_compilation = model.forward
model.forward = patched_dynamo_forward
return model
COMPILE_TRANSFORMER = ModuleOps(
name="compile_transformer",
matcher=lambda model: isinstance(model, LTXModel),
mutator=lambda model: compile_transformer(model),
)
def modify_sd_ops_for_compilation(original_sd_ops: SDOps, number_of_blocks: int = 48) -> SDOps:
for i in range(number_of_blocks):
original_sd_ops = original_sd_ops.with_replacement(
f"transformer_blocks.{i}.", f"transformer_blocks.{i}._orig_mod."
)
return original_sd_ops
@@ -1,3 +1,6 @@
from __future__ import annotations
import dataclasses
from dataclasses import dataclass
import torch
@@ -38,3 +41,17 @@ class Modality:
enabled: bool = True
context_mask: torch.Tensor | None = None
attention_mask: torch.Tensor | None = None
def split(self, sizes: list[int]) -> list[Modality]:
"""Split along the batch dimension into chunks of the given sizes."""
n = len(sizes)
split_fields: dict[str, list[torch.Tensor | None] | list[bool]] = {}
for f in dataclasses.fields(self):
value = getattr(self, f.name)
if isinstance(value, torch.Tensor):
split_fields[f.name] = list(value.split(sizes, dim=0))
elif value is None or isinstance(value, bool):
split_fields[f.name] = [value] * n
else:
raise TypeError(f"Cannot split field {f.name!r}: unsupported type {type(value)}")
return [Modality(**{name: parts[i] for name, parts in split_fields.items()}) for i in range(n)]
@@ -7,7 +7,7 @@ from ltx_core.model.video_vae.model_configurator import (
VideoEncoderConfigurator,
)
from ltx_core.model.video_vae.tiling import SpatialTilingConfig, TemporalTilingConfig, TilingConfig
from ltx_core.model.video_vae.video_vae import VideoDecoder, VideoEncoder, decode_video, get_video_chunks_number
from ltx_core.model.video_vae.video_vae import VideoDecoder, VideoEncoder, get_video_chunks_number
__all__ = [
"VAE_DECODER_COMFY_KEYS_FILTER",
@@ -19,6 +19,5 @@ __all__ = [
"VideoDecoderConfigurator",
"VideoEncoder",
"VideoEncoderConfigurator",
"decode_video",
"get_video_chunks_number",
]
@@ -1,72 +1,4 @@
import itertools
from dataclasses import dataclass
from typing import Callable, List, NamedTuple, Tuple
import torch
def compute_trapezoidal_mask_1d(
length: int,
ramp_left: int,
ramp_right: int,
left_starts_from_0: bool = False,
) -> torch.Tensor:
"""
Generate a 1D trapezoidal blending mask with linear ramps.
Args:
length: Output length of the mask.
ramp_left: Fade-in length on the left.
ramp_right: Fade-out length on the right.
left_starts_from_0: Whether the ramp starts from 0 or first non-zero value.
Useful for temporal tiles where the first tile is causal.
Returns:
A 1D tensor of shape `(length,)` with values in [0, 1].
"""
if length <= 0:
raise ValueError("Mask length must be positive.")
ramp_left = max(0, min(ramp_left, length))
ramp_right = max(0, min(ramp_right, length))
mask = torch.ones(length)
if ramp_left > 0:
interval_length = ramp_left + 1 if left_starts_from_0 else ramp_left + 2
fade_in = torch.linspace(0.0, 1.0, interval_length)[:-1]
if not left_starts_from_0:
fade_in = fade_in[1:]
mask[:ramp_left] *= fade_in
if ramp_right > 0:
fade_out = torch.linspace(1.0, 0.0, steps=ramp_right + 2)[1:-1]
mask[-ramp_right:] *= fade_out
return mask.clamp_(0, 1)
def compute_rectangular_mask_1d(
length: int,
left_ramp: int,
right_ramp: int,
) -> torch.Tensor:
"""
Generate a 1D rectangular (pulse) mask.
Args:
length: Output length of the mask.
left_ramp: Number of elements at the start of the mask to set to 0.
right_ramp: Number of elements at the end of the mask to set to 0.
Returns:
A 1D tensor of shape `(length,)` with values 0 or 1.
"""
if length <= 0:
raise ValueError("Mask length must be positive.")
mask = torch.ones(length)
if left_ramp > 0:
mask[:left_ramp] = 0
if right_ramp > 0:
mask[-right_ramp:] = 0
return mask
@dataclass(frozen=True)
@@ -135,157 +67,3 @@ class TilingConfig:
spatial_config=SpatialTilingConfig(tile_size_in_pixels=512, tile_overlap_in_pixels=64),
temporal_config=TemporalTilingConfig(tile_size_in_frames=64, tile_overlap_in_frames=24),
)
@dataclass(frozen=True)
class DimensionIntervals:
"""Defines how a single dimension is split into overlapping intervals (tiles).
Each list has length N where N is the number of intervals. The i-th element
of each list describes the i-th interval.
Attributes:
starts: Start index of each interval (inclusive).
ends: End index of each interval (exclusive).
left_ramps: Length of the left blend ramp for each interval.
Used to create masks that fade in from 0 to 1.
right_ramps: Length of the right blend ramp for each interval.
Used to create masks that fade out from 1 to 0.
"""
starts: List[int]
ends: List[int]
left_ramps: List[int]
right_ramps: List[int]
@dataclass(frozen=True)
class TensorTilingSpec:
"""Specifies how a tensor of a given shape is split into intervals (tiles) along each dimension.
Attributes:
original_shape: Shape of the tensor being tiled.
dimension_intervals: Per-dimension intervals (starts, ends, ramps) for each axis.
"""
original_shape: torch.Size
dimension_intervals: Tuple[DimensionIntervals, ...]
# Operation to split a single dimension of the tensor into intervals based on the length along the dimension.
SplitOperation = Callable[[int], DimensionIntervals]
# Operation to map the intervals in input dimension to slices and masks along a corresponding output dimension.
MappingOperation = Callable[[DimensionIntervals], tuple[list[slice], list[torch.Tensor | None]]]
def default_split_operation(length: int) -> DimensionIntervals:
return DimensionIntervals(starts=[0], ends=[length], left_ramps=[0], right_ramps=[0])
DEFAULT_SPLIT_OPERATION: SplitOperation = default_split_operation
def default_mapping_operation(
_intervals: DimensionIntervals,
) -> tuple[list[slice], list[torch.Tensor | None]]:
return [slice(0, None)], [None]
DEFAULT_MAPPING_OPERATION: MappingOperation = default_mapping_operation
class Tile(NamedTuple):
"""
Represents a single tile.
Attributes:
in_coords:
Tuple of slices specifying where to cut the tile from the INPUT tensor.
out_coords:
Tuple of slices specifying where this tile's OUTPUT should be placed in the reconstructed OUTPUT tensor.
masks_1d:
Per-dimension masks in OUTPUT units.
These are used to create all-dimensional blending mask.
Methods:
blend_mask:
Create a single N-D mask from the per-dimension masks.
"""
in_coords: Tuple[slice, ...]
out_coords: Tuple[slice, ...]
masks_1d: Tuple[Tuple[torch.Tensor, ...]]
@property
def blend_mask(self) -> torch.Tensor:
num_dims = len(self.out_coords)
per_dimension_masks: List[torch.Tensor] = []
for dim_idx in range(num_dims):
mask_1d = self.masks_1d[dim_idx]
view_shape = [1] * num_dims
if mask_1d is None:
# Broadcast mask along this dimension (length 1).
one = torch.ones(1)
view_shape[dim_idx] = 1
per_dimension_masks.append(one.view(*view_shape))
continue
# Reshape (L,) -> (1, ..., L, ..., 1) so masks across dimensions broadcast-multiply.
view_shape[dim_idx] = mask_1d.shape[0]
per_dimension_masks.append(mask_1d.view(*view_shape))
# Multiply per-dimension masks to form the full N-D mask (separable blending window).
combined_mask = per_dimension_masks[0]
for mask in per_dimension_masks[1:]:
combined_mask = combined_mask * mask
return combined_mask
def create_tiles_from_intervals_and_mappers(
intervals: TensorTilingSpec,
mappers: List[MappingOperation],
) -> List[Tile]:
full_dim_input_slices = []
full_dim_output_slices = []
full_dim_masks_1d = []
for axis_index in range(len(intervals.original_shape)):
dimension_intervals = intervals.dimension_intervals[axis_index]
starts = dimension_intervals.starts
ends = dimension_intervals.ends
input_slices = [slice(s, e) for s, e in zip(starts, ends, strict=True)]
output_slices, masks_1d = mappers[axis_index](dimension_intervals)
full_dim_input_slices.append(input_slices)
full_dim_output_slices.append(output_slices)
full_dim_masks_1d.append(masks_1d)
tiles = []
tile_in_coords = list(itertools.product(*full_dim_input_slices))
tile_out_coords = list(itertools.product(*full_dim_output_slices))
tile_mask_1ds = list(itertools.product(*full_dim_masks_1d))
for in_coord, out_coord, mask_1d in zip(tile_in_coords, tile_out_coords, tile_mask_1ds, strict=True):
tiles.append(
Tile(
in_coords=in_coord,
out_coords=out_coord,
masks_1d=mask_1d,
)
)
return tiles
def create_tiles(
tensor_shape: torch.Size,
splitters: List[SplitOperation],
mappers: List[MappingOperation],
) -> List[Tile]:
if len(splitters) != len(tensor_shape):
raise ValueError(
f"Number of splitters must be equal to number of dimensions in tensor shape, "
f"got {len(splitters)} and {len(tensor_shape)}"
)
if len(mappers) != len(tensor_shape):
raise ValueError(
f"Number of mappers must be equal to number of dimensions in tensor shape, "
f"got {len(mappers)} and {len(tensor_shape)}"
)
intervals = [splitter(length) for splitter, length in zip(splitters, tensor_shape, strict=True)]
tiling_spec = TensorTilingSpec(original_shape=tensor_shape, dimension_intervals=tuple(intervals))
return create_tiles_from_intervals_and_mappers(tiling_spec, mappers)
@@ -1,5 +1,4 @@
import logging
from dataclasses import replace
from typing import Any, Callable, Iterator, List, Tuple
import torch
@@ -13,17 +12,23 @@ from ltx_core.model.video_vae.enums import LogVarianceType, NormLayerType, Paddi
from ltx_core.model.video_vae.ops import PerChannelStatistics, patchify, unpatchify
from ltx_core.model.video_vae.resnet import ResnetBlock3D, UNetMidBlock3D
from ltx_core.model.video_vae.sampling import DepthToSpaceUpsample, SpaceToDepthDownsample
from ltx_core.model.video_vae.tiling import (
from ltx_core.model.video_vae.tiling import TilingConfig
from ltx_core.tiling import (
DEFAULT_MAPPING_OPERATION,
DEFAULT_SPLIT_OPERATION,
DimensionIntervals,
MappingOperation,
SplitOperation,
Tile,
TilingConfig,
compute_rectangular_mask_1d,
compute_trapezoidal_mask_1d,
create_tiles,
split_temporal,
)
from ltx_core.tiling import (
split_by_size as split_in_spatial,
)
from ltx_core.tiling import (
split_temporal_causal as split_in_temporal,
)
from ltx_core.types import VIDEO_SCALE_FACTORS, SpatioTemporalScaleFactors, VideoLatentShape
@@ -444,12 +449,12 @@ def prepare_tiles_for_encoding(
# Define split and map operations for the spatial dimensions
# Height axis (H)
splitters[3] = split_with_symmetric_overlaps(tile_size_px, overlap_px)
mappers[3] = make_mapping_operation(map_spatial_interval_to_latent, scale=VIDEO_SCALE_FACTORS.height)
splitters[3] = split_in_spatial(tile_size_px, overlap_px)
mappers[3] = to_mapping_operation(map_spatial_interval_to_latent, scale=VIDEO_SCALE_FACTORS.height)
# Width axis (W)
splitters[4] = split_with_symmetric_overlaps(tile_size_px, overlap_px)
mappers[4] = make_mapping_operation(map_spatial_interval_to_latent, scale=VIDEO_SCALE_FACTORS.width)
splitters[4] = split_in_spatial(tile_size_px, overlap_px)
mappers[4] = to_mapping_operation(map_spatial_interval_to_latent, scale=VIDEO_SCALE_FACTORS.width)
if tiling_config is not None and tiling_config.temporal_config is not None:
cfg = tiling_config.temporal_config
@@ -460,8 +465,8 @@ def prepare_tiles_for_encoding(
logger.warning(f"Overlap frames {overlap_frames} is less than 16, setting to minimum required 16")
overlap_frames = minimum_temporal_overlap_frames
splitters[2] = split_temporal_frames(tile_size_frames, overlap_frames)
mappers[2] = make_mapping_operation(map_temporal_interval_to_latent, scale=VIDEO_SCALE_FACTORS.time)
splitters[2] = split_temporal(tile_size_frames, overlap_frames)
mappers[2] = to_mapping_operation(map_temporal_interval_to_latent, scale=VIDEO_SCALE_FACTORS.time)
return create_tiles(video.shape, splitters, mappers)
@@ -784,8 +789,8 @@ class VideoDecoder(nn.Module):
axis_length = latent.shape[axis_idx]
lower_threshold = max(2, overlap + 1)
tile_size = max(lower_threshold, round(size * axis_length / long_side))
splitters[axis_idx] = split_with_symmetric_overlaps(tile_size, overlap)
mappers[axis_idx] = make_mapping_operation(map_spatial_interval_to_pixel, scale=factor)
splitters[axis_idx] = split_in_spatial(tile_size, overlap)
mappers[axis_idx] = to_mapping_operation(map_spatial_slice, scale=factor)
enable_on_axis(3, self.video_downscale_factors.height)
enable_on_axis(4, self.video_downscale_factors.width)
@@ -794,8 +799,8 @@ class VideoDecoder(nn.Module):
cfg = tiling_config.temporal_config
tile_size = cfg.tile_size_in_frames // self.video_downscale_factors.time
overlap = cfg.tile_overlap_in_frames // self.video_downscale_factors.time
splitters[2] = split_temporal_latents(tile_size, overlap)
mappers[2] = make_mapping_operation(map_temporal_interval_to_frame, scale=self.video_downscale_factors.time)
splitters[2] = split_in_temporal(tile_size, overlap)
mappers[2] = to_mapping_operation(map_temporal_slice, scale=self.video_downscale_factors.time)
return create_tiles(latent.shape, splitters, mappers)
@@ -892,6 +897,29 @@ class VideoDecoder(nn.Module):
previous_weights = previous_weights.clamp(min=1e-8)
yield previous_chunk / previous_weights
def decode_video(
self,
latent: torch.Tensor,
tiling_config: TilingConfig | None = None,
generator: torch.Generator | None = None,
) -> Iterator[torch.Tensor]:
"""Decode a video latent tensor, yielding uint8 chunks ``[f, h, w, c]``.
Subclasses (e.g. ``DistributedVideoDecoder``) may override this to
control eagerness or distribution across ranks.
"""
def convert_to_uint8(frames: torch.Tensor) -> torch.Tensor:
frames = (((frames + 1.0) / 2.0).clamp(0.0, 1.0) * 255.0).to(torch.uint8)
frames = rearrange(frames[0], "c f h w -> f h w c")
return frames
if tiling_config is not None:
for frames in self.tiled_decode(latent, tiling_config, generator=generator):
yield convert_to_uint8(frames)
else:
decoded = self(latent, generator=generator)
yield convert_to_uint8(decoded)
def _group_tiles_by_temporal_slice(self, tiles: List[Tile]) -> List[List[Tile]]:
"""Group tiles by their temporal output slice."""
if not tiles:
@@ -963,36 +991,6 @@ class VideoDecoder(nn.Module):
return weights
def decode_video(
latent: torch.Tensor,
video_decoder: VideoDecoder,
tiling_config: TilingConfig | None = None,
generator: torch.Generator | None = None,
) -> Iterator[torch.Tensor]:
"""
Decode a video latent tensor with the given decoder.
Args:
latent: Tensor [c, f, h, w]
video_decoder: Decoder module.
tiling_config: Optional tiling settings.
generator: Optional random generator for deterministic decoding.
Yields:
Decoded chunk [f, h, w, c], uint8 in [0, 255].
"""
def convert_to_uint8(frames: torch.Tensor) -> torch.Tensor:
frames = (((frames + 1.0) / 2.0).clamp(0.0, 1.0) * 255.0).to(torch.uint8)
frames = rearrange(frames[0], "c f h w -> f h w c")
return frames
if tiling_config is not None:
for frames in video_decoder.tiled_decode(latent, tiling_config, generator=generator):
yield convert_to_uint8(frames)
else:
decoded_video = video_decoder(latent, generator=generator)
yield convert_to_uint8(decoded_video)
def get_video_chunks_number(num_frames: int, tiling_config: TilingConfig | None = None) -> int:
"""
Get the number of video chunks for a given number of frames and tiling configuration.
@@ -1009,82 +1007,7 @@ def get_video_chunks_number(num_frames: int, tiling_config: TilingConfig | None
return (num_frames - 1 + frame_stride - 1) // frame_stride
def split_with_symmetric_overlaps(size: int, overlap: int) -> SplitOperation:
def split(dimension_size: int) -> DimensionIntervals:
if dimension_size <= size:
return DEFAULT_SPLIT_OPERATION(dimension_size)
amount = (dimension_size + size - 2 * overlap - 1) // (size - overlap)
starts = [i * (size - overlap) for i in range(amount)]
ends = [start + size for start in starts]
ends[-1] = dimension_size
left_ramps = [0] + [overlap] * (amount - 1)
right_ramps = [overlap] * (amount - 1) + [0]
return DimensionIntervals(starts=starts, ends=ends, left_ramps=left_ramps, right_ramps=right_ramps)
return split
def split_temporal_latents(size: int, overlap: int) -> SplitOperation:
"""Split a temporal axis into overlapping tiles with causal handling.
Example with size=24, overlap=8 (units are whatever axis you split):
Non-causal split would produce:
Tile 0: [0, 24), left_ramp=0, right_ramp=8
Tile 1: [16, 40), left_ramp=8, right_ramp=8
Tile 2: [32, 56), left_ramp=8, right_ramp=0
Causal split produces:
Tile 0: [0, 24), left_ramp=0, right_ramp=8 (unchanged - starts at anchor)
Tile 1: [15, 40), left_ramp=9, right_ramp=8 (shifted back 1, ramp +1)
Tile 2: [31, 56), left_ramp=9, right_ramp=0 (shifted back 1, ramp +1)
This ensures each tile can causally depend on frames from previous tiles while maintaining
proper temporal continuity through the blend ramps.
Args:
size: Tile size in *axis units* (latent steps for LTX time tiling)
overlap: Overlap between tiles in the same units
Returns:
Split operation that divides temporal dimension with causal handling
"""
non_causal_split = split_with_symmetric_overlaps(size, overlap)
def split(dimension_size: int) -> DimensionIntervals:
if dimension_size <= size:
return DEFAULT_SPLIT_OPERATION(dimension_size)
intervals = non_causal_split(dimension_size)
starts = intervals.starts
starts[1:] = [s - 1 for s in starts[1:]]
# Extend blend ramps by 1 for non-first tiles to blend over the extra frame
left_ramps = intervals.left_ramps
left_ramps[1:] = [r + 1 for r in left_ramps[1:]]
return replace(intervals, starts=starts, left_ramps=left_ramps)
return split
def split_temporal_frames(tile_size_frames: int, overlap_frames: int) -> SplitOperation:
"""Split a temporal axis in video frame space into overlapping tiles.
Args:
tile_size_frames: Tile length in frames.
overlap_frames: Overlap between consecutive tiles in frames.
Returns:
Split operation that takes frame count and returns DimensionIntervals in frame indices.
"""
non_causal_split = split_with_symmetric_overlaps(tile_size_frames, overlap_frames)
def split(dimension_size: int) -> DimensionIntervals:
if dimension_size <= tile_size_frames:
return DEFAULT_SPLIT_OPERATION(dimension_size)
intervals = non_causal_split(dimension_size)
ends = intervals.ends
ends[:-1] = [e + 1 for e in ends[:-1]]
right_ramps = [0] * len(intervals.right_ramps)
return replace(intervals, ends=ends, right_ramps=right_ramps)
return split
def make_mapping_operation(
def to_mapping_operation(
map_func: Callable[[int, int, int, int, int], Tuple[slice, torch.Tensor | None]],
scale: int,
) -> MappingOperation:
@@ -1102,13 +1025,10 @@ def make_mapping_operation(
def map_op(intervals: DimensionIntervals) -> tuple[list[slice], list[torch.Tensor | None]]:
output_slices: list[slice] = []
masks_1d: list[torch.Tensor | None] = []
number_of_slices = len(intervals.starts)
for i in range(number_of_slices):
start = intervals.starts[i]
end = intervals.ends[i]
left_ramp = intervals.left_ramps[i]
right_ramp = intervals.right_ramps[i]
output_slice, mask_1d = map_func(start, end, left_ramp, right_ramp, scale)
for interval in intervals.intervals:
output_slice, mask_1d = map_func(
interval.start, interval.end, interval.left_ramp, interval.right_ramp, scale
)
output_slices.append(output_slice)
masks_1d.append(mask_1d)
return output_slices, masks_1d
@@ -1116,31 +1036,13 @@ def make_mapping_operation(
return map_op
def map_temporal_interval_to_frame(
begin: int,
end: int,
left_ramp: int,
right_ramp: int,
scale: int,
) -> Tuple[slice, torch.Tensor]:
"""Map temporal interval in latent space to video frame space.
Args:
begin: Start position in latent space
end: End position in latent space
left_ramp: Left ramp size in latent space
right_ramp: Right ramp size in latent space
scale: Scale factor for transformation
Returns:
Tuple of (output_slice, blend_mask)
"""
def map_temporal_slice(begin: int, end: int, left_ramp: int, right_ramp: int, scale: int) -> Tuple[slice, torch.Tensor]:
start = begin * scale
stop = 1 + (end - 1) * scale
left_ramp = 0 if left_ramp == 0 else 1 + (left_ramp - 1) * scale
right_ramp = right_ramp * scale
left_ramp_frames = 0 if left_ramp == 0 else 1 + (left_ramp - 1) * scale
right_ramp_frames = right_ramp * scale
mask_1d = compute_trapezoidal_mask_1d(stop - start, left_ramp_frames, right_ramp_frames, True)
return slice(start, stop), mask_1d
return slice(start, stop), compute_trapezoidal_mask_1d(stop - start, left_ramp, right_ramp, True)
def map_temporal_interval_to_latent(
@@ -1171,25 +1073,13 @@ def map_temporal_interval_to_latent(
return slice(start, stop), mask_1d
def map_spatial_interval_to_pixel(
begin: int,
end: int,
left_ramp: int,
right_ramp: int,
scale: int,
) -> Tuple[slice, torch.Tensor]:
"""Map spatial interval in latent space to pixel space.
Args:
begin: Start position in latent space
end: End position in latent space
left_ramp: Left ramp size in latent space
right_ramp: Right ramp size in latent space
scale: Scale factor for transformation
"""
def map_spatial_slice(begin: int, end: int, left_ramp: int, right_ramp: int, scale: int) -> Tuple[slice, torch.Tensor]:
start = begin * scale
stop = end * scale
mask_1d = compute_trapezoidal_mask_1d(stop - start, left_ramp * scale, right_ramp * scale, False)
return slice(start, stop), mask_1d
left_ramp = left_ramp * scale
right_ramp = right_ramp * scale
return slice(start, stop), compute_trapezoidal_mask_1d(stop - start, left_ramp, right_ramp, False)
def map_spatial_interval_to_latent(
@@ -7,12 +7,6 @@ from ltx_core.model.transformer.model import LTXModel
BLOCK_SIZE = 1024
def calculate_weight_float8(target_weights: torch.Tensor, original_weights: torch.Tensor) -> torch.Tensor:
result = _fused_add_round_launch(target_weights, original_weights, seed=0).to(target_weights.dtype)
target_weights.copy_(result, non_blocking=True)
return target_weights
def _fused_add_round_launch(target_weight: torch.Tensor, original_weight: torch.Tensor, seed: int) -> torch.Tensor:
# Lazy import triton - only available on CUDA platforms
import triton # noqa: PLC0415
@@ -65,34 +59,44 @@ def _upcast_and_round(
return _fused_add_round_launch(torch.zeros_like(weight, dtype=dtype), weight, seed)
class Fp8CastLinear(torch.nn.Linear):
"""nn.Linear storing weights in fp8, upcasting to input dtype during forward.
Used via __class__ reassignment (not subclassing) so existing weight tensors
are preserved in-place. Class-level forward is required for torch.compile
compatibility — instance-level closure monkey-patches cause graph breaks.
"""
_with_stochastic_rounding: bool
_seed: int
def forward(self, input: torch.Tensor) -> torch.Tensor: # noqa: A002, type: ignore[override]
w_up = _upcast_and_round(self.weight, input.dtype, self._with_stochastic_rounding, self._seed)
b_up = (
_upcast_and_round(self.bias, input.dtype, self._with_stochastic_rounding, self._seed)
if self.bias is not None
else None
)
return torch.nn.functional.linear(input, w_up, b_up)
def _replace_fwd_with_upcast(layer: torch.nn.Linear, with_stochastic_rounding: bool = False, seed: int = 0) -> None:
"""
Replace linear.forward and rms_norm.forward with a version that:
- upcasts weight and bias to input's dtype
- returns F.linear or F.rms_norm calculated in that dtype
Intended to be applied via __class__ reassignment to existing nn.Linear
instances so that their parameter and buffer tensors are preserved in-place,
avoiding re-instantiation. Forward remains defined at the class level, which
is required for torch.compile compatibility — instance-level closure
monkey-patches cause graph breaks.
"""
layer.original_forward = layer.forward
def new_linear_forward(*args, **_kwargs) -> torch.Tensor:
# assume first arg is the input tensor
x = args[0]
w_up = _upcast_and_round(layer.weight, x.dtype, with_stochastic_rounding, seed)
b_up = None
if layer.bias is not None:
b_up = _upcast_and_round(layer.bias, x.dtype, with_stochastic_rounding, seed)
return torch.nn.functional.linear(x, w_up, b_up)
layer.forward = new_linear_forward
layer.__class__ = Fp8CastLinear
layer._with_stochastic_rounding = with_stochastic_rounding
layer._seed = seed
def _amend_forward_with_upcast(
model: torch.nn.Module, with_stochastic_rounding: bool = False, seed: int = 0
) -> torch.nn.Module:
"""
Replace the forward method of the model's Linear and RMSNorm layers to forward
Replace the forward method of the model's Linear layers to forward
with upcast and optional stochastic rounding.
"""
for m in model.modules():
+464
View File
@@ -0,0 +1,464 @@
from __future__ import annotations
import itertools
from dataclasses import dataclass, replace
from typing import Callable, NamedTuple
import torch
def compute_trapezoidal_mask_1d(
length: int,
ramp_left: int,
ramp_right: int,
left_starts_from_0: bool = False,
) -> torch.Tensor:
"""
Generate a 1D trapezoidal blending mask with linear ramps.
Args:
length: Output length of the mask.
ramp_left: Fade-in length on the left.
ramp_right: Fade-out length on the right.
left_starts_from_0: Whether the ramp starts from 0 or first non-zero value.
Useful for temporal tiles where the first tile is causal.
Returns:
A 1D tensor of shape `(length,)` with values in [0, 1].
"""
if length <= 0:
raise ValueError("Mask length must be positive.")
ramp_left = max(0, min(ramp_left, length))
ramp_right = max(0, min(ramp_right, length))
mask = torch.ones(length)
if ramp_left > 0:
interval_length = ramp_left + 1 if left_starts_from_0 else ramp_left + 2
fade_in = torch.linspace(0.0, 1.0, interval_length)[:-1]
if not left_starts_from_0:
fade_in = fade_in[1:]
mask[:ramp_left] *= fade_in
if ramp_right > 0:
fade_out = torch.linspace(1.0, 0.0, steps=ramp_right + 2)[1:-1]
mask[-ramp_right:] *= fade_out
return mask.clamp_(0, 1)
def compute_rectangular_mask_1d(
length: int,
left_ramp: int,
right_ramp: int,
) -> torch.Tensor:
"""
Generate a 1D rectangular (pulse) mask.
Args:
length: Output length of the mask.
left_ramp: Number of elements at the start of the mask to set to 0.
right_ramp: Number of elements at the end of the mask to set to 0.
Returns:
A 1D tensor of shape `(length,)` with values 0 or 1.
"""
if length <= 0:
raise ValueError("Mask length must be positive.")
mask = torch.ones(length)
if left_ramp > 0:
mask[:left_ramp] = 0
if right_ramp > 0:
mask[-right_ramp:] = 0
return mask
@dataclass(frozen=True)
class DimensionInterval:
start: int
end: int
left_ramp: int
right_ramp: int
@dataclass(frozen=True)
class DimensionIntervals:
"""Intervals which a single dimension of the latent space is split into.
Each interval is defined by its start, end, left ramp, and right ramp.
The start and end are the indices of the first and last element (exclusive) in the interval.
Ramps are regions of the interval where the value of the mask tensor is
interpolated between 0 and 1 for blending with neighboring intervals.
The left ramp and right ramp values are the lengths of the left and right ramps.
"""
intervals: list[DimensionInterval]
@dataclass(frozen=True)
class LatentIntervals:
"""Intervals which the latent tensor of given shape is split into.
Each dimension of the latent space is split into intervals based on the length along said dimension.
"""
original_shape: torch.Size
dimension_intervals: tuple[DimensionIntervals, ...]
# Operation to split a single dimension of the tensor into intervals based on the length along the dimension.
SplitOperation = Callable[[int], DimensionIntervals]
# Operation to map the intervals in input dimension to slices and masks along a corresponding output dimension.
MappingOperation = Callable[[DimensionIntervals], tuple[list[slice], list[torch.Tensor | None]]]
def default_split_operation(length: int) -> DimensionIntervals:
return DimensionIntervals(intervals=[DimensionInterval(start=0, end=length, left_ramp=0, right_ramp=0)])
DEFAULT_SPLIT_OPERATION: SplitOperation = default_split_operation
def default_mapping_operation(
_intervals: DimensionIntervals,
) -> tuple[list[slice], list[torch.Tensor | None]]:
return [slice(0, None)], [None]
DEFAULT_MAPPING_OPERATION: MappingOperation = default_mapping_operation
# ---------------------------------------------------------------------------
# Split functions
# ---------------------------------------------------------------------------
def split_by_size(size: int, overlap: int) -> SplitOperation:
"""Split a dimension into overlapping tiles of a given size.
Tiles are sized ``size`` with ``overlap`` shared elements between
consecutive tiles. The last tile may be shorter if the dimension
doesn't divide evenly.
Args:
size: Target tile size (in axis units).
overlap: Overlap between consecutive tiles.
Returns:
A split operation that divides a dimension into tiles.
"""
if size <= 0:
raise ValueError(f"size must be > 0, got {size}")
if overlap < 0 or overlap >= size:
raise ValueError(f"overlap must satisfy 0 <= overlap < size, got overlap={overlap}, size={size}")
def split(dimension_size: int) -> DimensionIntervals:
if dimension_size <= size:
return DEFAULT_SPLIT_OPERATION(dimension_size)
amount = (dimension_size + size - 2 * overlap - 1) // (size - overlap)
intervals = [
DimensionInterval(start=0, end=size, left_ramp=0, right_ramp=overlap),
*(
DimensionInterval(
start=i * (size - overlap),
end=i * (size - overlap) + size,
left_ramp=overlap,
right_ramp=overlap,
)
for i in range(1, amount - 1)
),
DimensionInterval(
start=(amount - 1) * (size - overlap), end=dimension_size, left_ramp=overlap, right_ramp=0
),
]
return DimensionIntervals(intervals=intervals)
return split
def split_temporal_causal(size: int, overlap: int) -> SplitOperation:
"""Split a temporal axis into overlapping tiles with causal handling.
Each tile after the first is shifted back by 1 and its left ramp is
increased by 1, ensuring causal continuity through the blend ramps.
Args:
size: Tile size in axis units.
overlap: Overlap between tiles in the same units.
Returns:
Split operation that divides temporal dimension with causal handling.
"""
non_causal_split = split_by_size(size, overlap)
def split(dimension_size: int) -> DimensionIntervals:
if dimension_size <= size:
return DEFAULT_SPLIT_OPERATION(dimension_size)
dim_intervals = non_causal_split(dimension_size)
modified_intervals = [dim_intervals.intervals[0]] + [
replace(interval, start=interval.start - 1, left_ramp=interval.left_ramp + 1)
for interval in dim_intervals.intervals[1:]
]
return DimensionIntervals(intervals=modified_intervals)
return split
def split_temporal(tile_size_frames: int, overlap_frames: int) -> SplitOperation:
"""Split a temporal axis in video frame space into overlapping tiles.
Args:
tile_size_frames: Tile length in frames.
overlap_frames: Overlap between consecutive tiles in frames.
Returns:
Split operation that takes frame count and returns DimensionIntervals in frame indices.
"""
non_causal_split = split_by_size(tile_size_frames, overlap_frames)
def split(dimension_size: int) -> DimensionIntervals:
if dimension_size <= tile_size_frames:
return DEFAULT_SPLIT_OPERATION(dimension_size)
dim_intervals = non_causal_split(dimension_size)
modified_intervals = [
replace(interval, end=interval.end + 1, right_ramp=0) for interval in dim_intervals.intervals[:-1]
] + [replace(dim_intervals.intervals[-1], right_ramp=0)]
return DimensionIntervals(intervals=modified_intervals)
return split
def split_by_count_temporal_causal(num_tiles: int, overlap: int = 0) -> SplitOperation:
"""Split a temporal dimension by count with causal handling.
Wraps :func:`split_by_count` with the same causal adjustment as
:func:`split_temporal_causal`: each tile after the first is shifted
back by 1 and its left ramp is increased by 1.
Args:
num_tiles: Number of tiles. Must be >= 1.
overlap: Overlap between adjacent tiles (default 0).
Returns:
A split operation that divides a temporal dimension into tiles.
"""
non_causal_split = split_by_count(num_tiles, overlap)
def split(dimension_size: int) -> DimensionIntervals:
dim_intervals = non_causal_split(dimension_size)
if len(dim_intervals.intervals) <= 1:
return dim_intervals
modified_intervals = [dim_intervals.intervals[0]] + [
replace(interval, start=interval.start - 1, left_ramp=interval.left_ramp + 1)
for interval in dim_intervals.intervals[1:]
]
return DimensionIntervals(intervals=modified_intervals)
return split
def split_by_count(num_tiles: int, overlap: int = 0) -> SplitOperation:
"""Split a dimension into a given number of tiles with overlap.
Computes the tile size as
``(dim_size + overlap * (num_tiles - 1)) // num_tiles`` so that
``num_tiles`` tiles of that size with ``overlap`` shared elements
cover the dimension evenly. Delegates to :func:`split_by_size` for
the actual interval construction.
When the total ``dim_size + overlap * (num_tiles - 1)`` is not evenly
divisible by ``num_tiles``, the first ``remainder`` tiles each absorb
one extra unit.
Args:
num_tiles: Number of tiles. Must be >= 1.
overlap: Overlap between adjacent tiles (default 0). Must be >= 0
and less than the computed tile size.
Returns:
A split operation that divides a dimension into tiles.
"""
if num_tiles < 1:
raise ValueError(f"num_tiles must be >= 1, got {num_tiles}")
if overlap < 0:
raise ValueError(f"overlap must be >= 0, got {overlap}")
def split(dim_size: int) -> DimensionIntervals:
if num_tiles > dim_size:
raise ValueError(
f"num_tiles ({num_tiles}) exceeds dim_size ({dim_size}). Cannot assign at least 1 unit per tile."
)
if num_tiles == 1:
return DEFAULT_SPLIT_OPERATION(dim_size)
total = dim_size + overlap * (num_tiles - 1)
tile_size = total // num_tiles
remainder = total % num_tiles
base_intervals = split_by_size(tile_size, overlap)(dim_size - remainder).intervals
# First `remainder` tiles each absorb 1 extra unit; shift subsequent boundaries.
intervals: list[DimensionInterval] = []
for i, iv in enumerate(base_intervals):
shift = min(i, remainder)
grow = 1 if i < remainder else 0
intervals.append(replace(iv, start=iv.start + shift, end=iv.end + shift + grow))
return DimensionIntervals(intervals=intervals)
return split
# ---------------------------------------------------------------------------
# Mapping operations
# ---------------------------------------------------------------------------
def identity_mapping_operation(intervals: DimensionIntervals) -> tuple[list[slice], list[torch.Tensor | None]]:
"""Map each DimensionInterval to an output region at the same position, with trapezoidal blend masks.
For every interval the output start/end matches the input start/end and a
1-D blending mask is built from the interval's left_ramp and right_ramp.
"""
out_slices: list[slice] = []
masks: list[torch.Tensor | None] = []
for iv in intervals.intervals:
out_slices.append(slice(iv.start, iv.end))
masks.append(compute_trapezoidal_mask_1d(iv.end - iv.start, iv.left_ramp, iv.right_ramp))
return out_slices, masks
class Tile(NamedTuple):
"""
Represents a single tile.
Attributes:
in_coords:
Tuple of slices specifying where to cut the tile from the INPUT tensor.
out_coords:
Tuple of slices specifying where this tile's OUTPUT should be placed in the reconstructed OUTPUT tensor.
masks_1d:
Per-dimension masks in OUTPUT units.
These are used to create all-dimensional blending mask.
Methods:
blend_mask:
Create a single N-D mask from the per-dimension masks.
"""
in_coords: tuple[slice, ...]
out_coords: tuple[slice, ...]
masks_1d: tuple[torch.Tensor | None, ...]
@property
def blend_mask(self) -> torch.Tensor:
num_dims = len(self.out_coords)
per_dimension_masks: list[torch.Tensor] = []
for dim_idx in range(num_dims):
mask_1d = self.masks_1d[dim_idx]
view_shape = [1] * num_dims
if mask_1d is None:
# Broadcast mask along this dimension (length 1).
one = torch.ones(1)
view_shape[dim_idx] = 1
per_dimension_masks.append(one.view(*view_shape))
continue
# Reshape (L,) -> (1, ..., L, ..., 1) so masks across dimensions broadcast-multiply.
view_shape[dim_idx] = mask_1d.shape[0]
per_dimension_masks.append(mask_1d.view(*view_shape))
# Multiply per-dimension masks to form the full N-D mask (separable blending window).
combined_mask = per_dimension_masks[0]
for mask in per_dimension_masks[1:]:
combined_mask = combined_mask * mask
return combined_mask
def create_tiles_from_intervals_and_mappers(
intervals: LatentIntervals,
mappers: list[MappingOperation],
) -> list[Tile]:
full_dim_input_slices: list[list[slice]] = []
full_dim_output_slices: list[list[slice]] = []
full_dim_masks_1d: list[list[torch.Tensor | None]] = []
for axis_index in range(len(intervals.original_shape)):
dimension_intervals = intervals.dimension_intervals[axis_index]
input_slices = [slice(interval.start, interval.end) for interval in dimension_intervals.intervals]
output_slices, masks_1d = mappers[axis_index](dimension_intervals)
n_intervals = len(input_slices)
if len(output_slices) != n_intervals or len(masks_1d) != n_intervals:
raise ValueError(
f"Axis {axis_index}: mapper produced {len(output_slices)} output slices and "
f"{len(masks_1d)} masks for {n_intervals} input intervals"
)
full_dim_input_slices.append(input_slices)
full_dim_output_slices.append(output_slices)
full_dim_masks_1d.append(masks_1d)
return [
Tile(in_coords=in_coord, out_coords=out_coord, masks_1d=mask_1d)
for in_coord, out_coord, mask_1d in zip(
itertools.product(*full_dim_input_slices),
itertools.product(*full_dim_output_slices),
itertools.product(*full_dim_masks_1d),
strict=True,
)
]
def create_tiles(
latent_shape: torch.Size,
splitters: list[SplitOperation],
mappers: list[MappingOperation],
) -> list[Tile]:
if len(splitters) != len(latent_shape):
raise ValueError(
f"Number of splitters must be equal to number of dimensions in latent shape, "
f"got {len(splitters)} and {len(latent_shape)}"
)
if len(mappers) != len(latent_shape):
raise ValueError(
f"Number of mappers must be equal to number of dimensions in latent shape, "
f"got {len(mappers)} and {len(latent_shape)}"
)
intervals = [splitter(length) for splitter, length in zip(splitters, latent_shape, strict=True)]
latent_intervals = LatentIntervals(original_shape=latent_shape, dimension_intervals=tuple(intervals))
return create_tiles_from_intervals_and_mappers(latent_intervals, mappers)
# ---------------------------------------------------------------------------
# Video-grid tiling configs
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class DimensionTilingConfig:
"""Tiling parameters for a single dimension of the patchified grid.
Attributes:
num_tiles: Number of tiles along this dimension.
overlap: Overlap between adjacent tiles, in latent grid units.
Adjacent tiles share ``overlap`` grid cells at their
boundary, producing an overlap zone blended with
trapezoidal masks.
"""
num_tiles: int
overlap: int = 0
def __post_init__(self) -> None:
if self.num_tiles < 1:
raise ValueError(f"num_tiles must be >= 1, got {self.num_tiles}")
if self.overlap < 0:
raise ValueError(f"overlap must be >= 0, got {self.overlap}")
@classmethod
def from_tile_size(cls, dim_size: int, tile_size: int, overlap: int = 0) -> DimensionTilingConfig:
"""Create config by computing ``num_tiles`` from dimension size and tile size.
Args:
dim_size: Total length of the dimension.
tile_size: Desired tile size.
overlap: Overlap between consecutive tiles.
Returns:
A ``DimensionTilingConfig`` with the computed ``num_tiles``.
"""
split_op = split_by_size(tile_size, overlap)
intervals = split_op(dim_size)
return cls(num_tiles=len(intervals.intervals), overlap=overlap)
@dataclass(frozen=True)
class TileCountConfig:
"""Tiling layout for a ``(F, H, W)`` grid.
Specifies tile *counts* per dimension (as opposed to tile *sizes*
which are used by the single-GPU VAE ``TilingConfig``).
Attributes:
frames: Tiling along the temporal (frames) dimension.
height: Tiling along the latent height dimension.
width: Tiling along the latent width dimension.
"""
frames: DimensionTilingConfig = DimensionTilingConfig(num_tiles=1, overlap=0)
height: DimensionTilingConfig = DimensionTilingConfig(num_tiles=1, overlap=0)
width: DimensionTilingConfig = DimensionTilingConfig(num_tiles=1, overlap=0)
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "ltx-pipelines"
version = "1.0.0"
version = "1.1.0"
description = "Pipelines implementation for Lightricks' LTX-2 model"
readme = "README.md"
requires-python = ">=3.10"
@@ -7,7 +7,6 @@ This package provides ready-to-use pipelines for video generation:
- ICLoraPipeline: Image/video conditioning with distilled LoRA
- KeyframeInterpolationPipeline: Keyframe-based video interpolation
- RetakePipeline: Regenerate a time region (retake) of an existing video
- ModelLedger: Central coordinator for loading and building models
For more detailed components and utilities, import from specific submodules
like `ltx_pipelines.utils.media_io` or `ltx_pipelines.utils.constants`.
"""
@@ -3,38 +3,35 @@ from collections.abc import Iterator
import torch
from ltx_core.components.diffusion_steps import EulerDiffusionStep
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.components.schedulers import LTX2Scheduler
from ltx_core.loader import LoraPathStrengthAndSDOps
from ltx_core.loader.registry import Registry
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
from ltx_core.model.upsampler import upsample_video
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy
from ltx_core.types import Audio, AudioLatentShape, LatentState, VideoPixelShape
from ltx_pipelines.utils import ModelLedger
from ltx_core.types import Audio, AudioLatentShape, VideoPixelShape
from ltx_pipelines.utils.args import default_2_stage_arg_parser
from ltx_pipelines.utils.blocks import (
AudioConditioner,
DiffusionStage,
ImageConditioner,
PromptEncoder,
VideoDecoder,
VideoUpsampler,
)
from ltx_pipelines.utils.constants import (
STAGE_2_DISTILLED_SIGMA_VALUES,
)
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
from ltx_pipelines.utils.helpers import (
assert_resolution,
cleanup_memory,
combined_image_conditionings,
denoise_video_only,
encode_prompts,
get_device,
multi_modal_guider_denoising_func,
simple_denoising_func,
)
from ltx_pipelines.utils.media_io import decode_audio_from_file, encode_video
from ltx_pipelines.utils.samplers import euler_denoising_loop
from ltx_pipelines.utils.types import PipelineComponents
device = get_device()
from ltx_pipelines.utils.types import ModalitySpec
class A2VidPipelineTwoStage:
@@ -52,30 +49,40 @@ class A2VidPipelineTwoStage:
spatial_upsampler_path: str,
gemma_root: str,
loras: list[LoraPathStrengthAndSDOps],
device: torch.device = device,
device: torch.device | None = None,
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
):
self.device = device
self.device = device or get_device()
self.dtype = torch.bfloat16
self.stage_1_model_ledger = ModelLedger(
dtype=self.dtype,
device=device,
checkpoint_path=checkpoint_path,
gemma_root_path=gemma_root,
spatial_upsampler_path=spatial_upsampler_path,
loras=loras,
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
self.audio_conditioner = AudioConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
self.stage_1 = DiffusionStage(
checkpoint_path,
self.dtype,
self.device,
loras=tuple(loras),
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
self.stage_2_model_ledger = self.stage_1_model_ledger.with_additional_loras(
loras=distilled_lora,
stage_2_loras = (*tuple(loras), *tuple(distilled_lora))
self.stage_2 = DiffusionStage(
checkpoint_path,
self.dtype,
self.device,
loras=stage_2_loras,
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
self.pipeline_components = PipelineComponents(
dtype=self.dtype,
device=device,
self.upsampler = VideoUpsampler(
checkpoint_path, spatial_upsampler_path, self.dtype, self.device, registry=registry
)
self.video_decoder = VideoDecoder(checkpoint_path, self.dtype, self.device, registry=registry)
def __call__( # noqa: PLR0913
self,
@@ -94,31 +101,35 @@ class A2VidPipelineTwoStage:
audio_max_duration: float | None = None,
tiling_config: TilingConfig | None = None,
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
) -> tuple[Iterator[torch.Tensor], Audio]:
assert_resolution(height=height, width=width, is_two_stage=True)
generator = torch.Generator(device=self.device).manual_seed(seed)
noiser = GaussianNoiser(generator=generator)
stepper = EulerDiffusionStep()
dtype = torch.bfloat16
ctx_p, ctx_n = encode_prompts(
ctx_p, ctx_n = self.prompt_encoder(
[prompt, negative_prompt],
self.stage_1_model_ledger,
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_n, _ = ctx_n.video_encoding, ctx_n.audio_encoding
# Encode audio.
decoded_audio = decode_audio_from_file(audio_path, self.device, audio_start_time, audio_max_duration)
encoded_audio_latent = vae_encode_audio(decoded_audio, self.stage_1_model_ledger.audio_encoder())
if decoded_audio is None:
raise ValueError(f"Failed to decode audio from {audio_path}. Please check the file and try again.")
encoded_audio_latent = self.audio_conditioner(lambda enc: vae_encode_audio(decoded_audio, enc, None))
audio_shape = AudioLatentShape.from_duration(batch=1, duration=num_frames / frame_rate, channels=8, mel_bins=16)
encoded_audio_latent = encoded_audio_latent[:, :, : audio_shape.frames]
# Stage 1: encode image conditionings with the VAE encoder, then free it
# before loading the transformer to reduce peak VRAM.
# Stage 1: encode image conditionings with the VAE encoder, then denoise
# video-only (audio frozen).
stage_1_output_shape = VideoPixelShape(
batch=1,
frames=num_frames,
@@ -126,122 +137,91 @@ class A2VidPipelineTwoStage:
height=height // 2,
fps=frame_rate,
)
video_encoder = self.stage_1_model_ledger.video_encoder()
stage_1_conditionings = combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
stage_1_conditionings = self.image_conditioner(
lambda enc: combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=enc,
dtype=dtype,
device=self.device,
)
)
torch.cuda.synchronize()
del video_encoder
cleanup_memory()
transformer = self.stage_1_model_ledger.transformer()
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
def first_stage_denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return euler_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=multi_modal_guider_denoising_func(
video_guider=MultiModalGuider(
params=video_guider_params,
negative_context=v_context_n,
),
audio_guider=MultiModalGuider(
params=MultiModalGuiderParams(),
),
v_context=v_context_p,
a_context=a_context_p,
transformer=transformer, # noqa: F821
video_state, _ = self.stage_1(
denoiser=GuidedDenoiser(
v_context=v_context_p,
a_context=a_context_p,
video_guider=MultiModalGuider(
params=video_guider_params,
negative_context=v_context_n,
),
)
video_state = denoise_video_only(
output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings,
noiser=noiser,
audio_guider=MultiModalGuider(
params=MultiModalGuiderParams(),
),
),
sigmas=sigmas,
stepper=stepper,
denoising_loop_fn=first_stage_denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
initial_audio_latent=encoded_audio_latent,
noiser=noiser,
width=stage_1_output_shape.width,
height=stage_1_output_shape.height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(
context=v_context_p,
conditionings=stage_1_conditionings,
),
audio=ModalitySpec(
context=a_context_p,
frozen=True,
noise_scale=0.0,
initial_latent=encoded_audio_latent,
),
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
torch.cuda.synchronize()
del transformer
cleanup_memory()
# Stage 2: Upsample and refine the video at higher resolution with distilled LoRA.
video_encoder = self.stage_1_model_ledger.video_encoder()
upscaled_video_latent = upsample_video(
latent=video_state.latent[:1],
video_encoder=video_encoder,
upsampler=self.stage_2_model_ledger.spatial_upsampler(),
)
upscaled_video_latent = self.upsampler(video_state.latent[:1])
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = combined_image_conditionings(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
del video_encoder
torch.cuda.synchronize()
cleanup_memory()
transformer = self.stage_2_model_ledger.transformer()
distilled_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
def second_stage_denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return euler_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=simple_denoising_func(
video_context=v_context_p,
audio_context=a_context_p,
transformer=transformer, # noqa: F821
),
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = self.image_conditioner(
lambda enc: combined_image_conditionings(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=enc,
dtype=dtype,
device=self.device,
)
)
video_state = denoise_video_only(
output_shape=stage_2_output_shape,
conditionings=stage_2_conditionings,
noiser=noiser,
video_state, _ = self.stage_2(
denoiser=SimpleDenoiser(v_context_p, a_context_p),
sigmas=distilled_sigmas,
stepper=stepper,
denoising_loop_fn=second_stage_denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
noise_scale=distilled_sigmas[0],
initial_video_latent=upscaled_video_latent,
initial_audio_latent=encoded_audio_latent,
noiser=noiser,
width=width,
height=height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(
context=v_context_p,
conditionings=stage_2_conditionings,
noise_scale=distilled_sigmas[0].item(),
initial_latent=upscaled_video_latent,
),
audio=ModalitySpec(
context=a_context_p,
frozen=True,
noise_scale=0.0,
initial_latent=encoded_audio_latent,
),
streaming_prefetch_count=streaming_prefetch_count,
)
torch.cuda.synchronize()
del transformer
cleanup_memory()
decoded_video = vae_decode_video(
video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config, generator
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
# Return the original input audio instead of VAE-decoded audio to preserve fidelity.
# decode_audio_from_file already returns normalised [-1, 1] float values.
@@ -280,6 +260,7 @@ def main() -> None:
gemma_root=args.gemma_root,
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -308,6 +289,8 @@ def main() -> None:
audio_max_duration=args.audio_max_duration
if args.audio_max_duration is not None
else args.num_frames / args.frame_rate,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
encode_video(
@@ -3,40 +3,38 @@ from collections.abc import Iterator
import torch
from ltx_core.components.diffusion_steps import EulerDiffusionStep
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.loader import LoraPathStrengthAndSDOps
from ltx_core.model.audio_vae import decode_audio as vae_decode_audio
from ltx_core.model.upsampler import upsample_video
from ltx_core.loader.registry import Registry
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy
from ltx_core.types import Audio, LatentState, VideoPixelShape
from ltx_pipelines.utils import ModelLedger, euler_denoising_loop
from ltx_core.types import Audio
from ltx_pipelines.utils.args import (
ImageConditioningInput,
default_2_stage_distilled_arg_parser,
detect_checkpoint_path,
)
from ltx_pipelines.utils.blocks import (
AudioDecoder,
DiffusionStage,
ImageConditioner,
PromptEncoder,
VideoDecoder,
VideoUpsampler,
)
from ltx_pipelines.utils.constants import (
DISTILLED_SIGMA_VALUES,
STAGE_2_DISTILLED_SIGMA_VALUES,
detect_params,
)
from ltx_pipelines.utils.denoisers import SimpleDenoiser
from ltx_pipelines.utils.helpers import (
assert_resolution,
cleanup_memory,
combined_image_conditionings,
denoise_audio_video,
encode_prompts,
get_device,
simple_denoising_func,
)
from ltx_pipelines.utils.media_io import encode_video
from ltx_pipelines.utils.types import PipelineComponents
device = get_device()
from ltx_pipelines.utils.types import ModalitySpec
class DistilledPipeline:
@@ -52,26 +50,32 @@ class DistilledPipeline:
gemma_root: str,
spatial_upsampler_path: str,
loras: list[LoraPathStrengthAndSDOps],
device: torch.device = device,
device: torch.device | None = None,
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
):
self.device = device
self.device = device or get_device()
self.dtype = torch.bfloat16
self.model_ledger = ModelLedger(
dtype=self.dtype,
device=device,
checkpoint_path=distilled_checkpoint_path,
spatial_upsampler_path=spatial_upsampler_path,
gemma_root_path=gemma_root,
loras=loras,
self.prompt_encoder = PromptEncoder(
distilled_checkpoint_path, gemma_root, self.dtype, self.device, registry=registry
)
self.image_conditioner = ImageConditioner(distilled_checkpoint_path, self.dtype, self.device, registry=registry)
self.stage = DiffusionStage(
distilled_checkpoint_path,
self.dtype,
self.device,
loras=tuple(loras),
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
self.pipeline_components = PipelineComponents(
dtype=self.dtype,
device=device,
self.upsampler = VideoUpsampler(
distilled_checkpoint_path, spatial_upsampler_path, self.dtype, self.device, registry=registry
)
self.video_decoder = VideoDecoder(distilled_checkpoint_path, self.dtype, self.device, registry=registry)
self.audio_decoder = AudioDecoder(distilled_checkpoint_path, self.dtype, self.device, registry=registry)
def __call__(
self,
@@ -84,114 +88,88 @@ class DistilledPipeline:
images: list[ImageConditioningInput],
tiling_config: TilingConfig | None = None,
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
) -> tuple[Iterator[torch.Tensor], Audio]:
assert_resolution(height=height, width=width, is_two_stage=True)
generator = torch.Generator(device=self.device).manual_seed(seed)
noiser = GaussianNoiser(generator=generator)
stepper = EulerDiffusionStep()
dtype = torch.bfloat16
(ctx_p,) = encode_prompts(
(ctx_p,) = self.prompt_encoder(
[prompt],
self.model_ledger,
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
streaming_prefetch_count=streaming_prefetch_count,
)
video_context, audio_context = ctx_p.video_encoding, ctx_p.audio_encoding
# Stage 1: Initial low resolution video generation.
video_encoder = self.model_ledger.video_encoder()
transformer = self.model_ledger.transformer()
stage_1_sigmas = torch.Tensor(DISTILLED_SIGMA_VALUES).to(self.device)
def denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return euler_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=simple_denoising_func(
video_context=video_context,
audio_context=audio_context,
transformer=transformer, # noqa: F821
),
stage_1_w, stage_1_h = width // 2, height // 2
stage_1_conditionings = self.image_conditioner(
lambda enc: combined_image_conditionings(
images=images,
height=stage_1_h,
width=stage_1_w,
video_encoder=enc,
dtype=dtype,
device=self.device,
)
stage_1_output_shape = VideoPixelShape(
batch=1,
frames=num_frames,
width=width // 2,
height=height // 2,
fps=frame_rate,
)
stage_1_conditionings = combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
video_state, audio_state = denoise_audio_video(
output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings,
noiser=noiser,
video_state, audio_state = self.stage(
denoiser=SimpleDenoiser(video_context, audio_context),
sigmas=stage_1_sigmas,
stepper=stepper,
denoising_loop_fn=denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
noiser=noiser,
width=stage_1_w,
height=stage_1_h,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(context=video_context, conditionings=stage_1_conditionings),
audio=ModalitySpec(context=audio_context),
streaming_prefetch_count=streaming_prefetch_count,
)
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
upscaled_video_latent = upsample_video(
latent=video_state.latent[:1], video_encoder=video_encoder, upsampler=self.model_ledger.spatial_upsampler()
)
torch.cuda.synchronize()
cleanup_memory()
upscaled_video_latent = self.upsampler(video_state.latent[:1])
stage_2_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = combined_image_conditionings(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
stage_2_conditionings = self.image_conditioner(
lambda enc: combined_image_conditionings(
images=images,
height=height,
width=width,
video_encoder=enc,
dtype=dtype,
device=self.device,
)
)
video_state, audio_state = denoise_audio_video(
output_shape=stage_2_output_shape,
conditionings=stage_2_conditionings,
noiser=noiser,
video_state, audio_state = self.stage(
denoiser=SimpleDenoiser(video_context, audio_context),
sigmas=stage_2_sigmas,
stepper=stepper,
denoising_loop_fn=denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
noise_scale=stage_2_sigmas[0],
initial_video_latent=upscaled_video_latent,
initial_audio_latent=audio_state.latent,
noiser=noiser,
width=width,
height=height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(
context=video_context,
conditionings=stage_2_conditionings,
noise_scale=stage_2_sigmas[0].item(),
initial_latent=upscaled_video_latent,
),
audio=ModalitySpec(
context=audio_context,
noise_scale=stage_2_sigmas[0].item(),
initial_latent=audio_state.latent,
),
streaming_prefetch_count=streaming_prefetch_count,
)
torch.cuda.synchronize()
del transformer
del video_encoder
cleanup_memory()
decoded_video = vae_decode_video(
video_state.latent, self.model_ledger.video_decoder(), tiling_config, generator
)
decoded_audio = vae_decode_audio(
audio_state.latent, self.model_ledger.audio_decoder(), self.model_ledger.vocoder()
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
decoded_audio = self.audio_decoder(audio_state.latent)
return decoded_video, decoded_audio
@@ -208,6 +186,7 @@ def main() -> None:
gemma_root=args.gemma_root,
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -221,6 +200,7 @@ def main() -> None:
images=args.images,
tiling_config=tiling_config,
enhance_prompt=args.enhance_prompt,
streaming_prefetch_count=args.streaming_prefetch_count,
)
encode_video(
@@ -5,32 +5,17 @@ import torch
from einops import rearrange
from safetensors import safe_open
from ltx_core.components.diffusion_steps import EulerDiffusionStep
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.conditioning import (
ConditioningItem,
ConditioningItemAttentionStrengthWrapper,
VideoConditionByReferenceLatent,
)
from ltx_core.loader import LoraPathStrengthAndSDOps
from ltx_core.model.audio_vae import decode_audio as vae_decode_audio
from ltx_core.model.upsampler import upsample_video
from ltx_core.loader.registry import Registry
from ltx_core.model.video_vae import TilingConfig, VideoEncoder, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy
from ltx_core.types import Audio, LatentState, VideoLatentShape, VideoPixelShape
from ltx_pipelines.utils import (
ModelLedger,
assert_resolution,
cleanup_memory,
combined_image_conditionings,
denoise_audio_video,
encode_prompts,
euler_denoising_loop,
get_device,
simple_denoising_func,
)
from ltx_core.types import Audio, VideoLatentShape, VideoPixelShape
from ltx_pipelines.utils.args import (
ImageConditioningInput,
VideoConditioningAction,
@@ -38,15 +23,23 @@ from ltx_pipelines.utils.args import (
default_2_stage_distilled_arg_parser,
detect_checkpoint_path,
)
from ltx_pipelines.utils.blocks import (
AudioDecoder,
DiffusionStage,
ImageConditioner,
PromptEncoder,
VideoDecoder,
VideoUpsampler,
)
from ltx_pipelines.utils.constants import (
DISTILLED_SIGMA_VALUES,
STAGE_2_DISTILLED_SIGMA_VALUES,
detect_params,
)
from ltx_pipelines.utils.media_io import encode_video, load_video_conditioning
from ltx_pipelines.utils.types import PipelineComponents
device = get_device()
from ltx_pipelines.utils.denoisers import SimpleDenoiser
from ltx_pipelines.utils.helpers import assert_resolution, combined_image_conditionings, get_device
from ltx_pipelines.utils.media_io import decode_video_by_frame, encode_video, video_preprocess
from ltx_pipelines.utils.types import ModalitySpec
class ICLoraPipeline:
@@ -66,33 +59,41 @@ class ICLoraPipeline:
spatial_upsampler_path: str,
gemma_root: str,
loras: list[LoraPathStrengthAndSDOps],
device: torch.device = device,
device: torch.device | None = None,
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
):
self.device = device or get_device()
self.dtype = torch.bfloat16
self.stage_1_model_ledger = ModelLedger(
dtype=self.dtype,
device=device,
checkpoint_path=distilled_checkpoint_path,
spatial_upsampler_path=spatial_upsampler_path,
gemma_root_path=gemma_root,
loras=loras,
self.prompt_encoder = PromptEncoder(
distilled_checkpoint_path, gemma_root, self.dtype, self.device, registry=registry
)
self.image_conditioner = ImageConditioner(distilled_checkpoint_path, self.dtype, self.device, registry=registry)
self.stage_1 = DiffusionStage(
distilled_checkpoint_path,
self.dtype,
self.device,
loras=tuple(loras),
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
self.stage_2_model_ledger = ModelLedger(
dtype=self.dtype,
device=device,
checkpoint_path=distilled_checkpoint_path,
spatial_upsampler_path=spatial_upsampler_path,
gemma_root_path=gemma_root,
loras=[],
self.stage_2 = DiffusionStage(
distilled_checkpoint_path,
self.dtype,
self.device,
loras=(),
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
self.pipeline_components = PipelineComponents(
dtype=self.dtype,
device=device,
self.upsampler = VideoUpsampler(
distilled_checkpoint_path, spatial_upsampler_path, self.dtype, self.device, registry=registry
)
self.device = device
self.video_decoder = VideoDecoder(distilled_checkpoint_path, self.dtype, self.device, registry=registry)
self.audio_decoder = AudioDecoder(distilled_checkpoint_path, self.dtype, self.device, registry=registry)
# Read reference downscale factor from LoRA metadata.
# IC-LoRAs trained with low-resolution reference videos store this factor
@@ -124,6 +125,7 @@ class ICLoraPipeline:
conditioning_attention_strength: float = 1.0,
skip_stage_2: bool = False,
conditioning_attention_mask: torch.Tensor | None = None,
streaming_prefetch_count: int | None = None,
) -> tuple[Iterator[torch.Tensor], Audio]:
"""
Generate video with IC-LoRA conditioning.
@@ -165,15 +167,13 @@ class ICLoraPipeline:
generator = torch.Generator(device=self.device).manual_seed(seed)
noiser = GaussianNoiser(generator=generator)
stepper = EulerDiffusionStep()
dtype = torch.bfloat16
(ctx_p,) = encode_prompts(
(ctx_p,) = self.prompt_encoder(
[prompt],
self.stage_1_model_ledger,
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
video_context, audio_context = ctx_p.video_encoding, ctx_p.audio_encoding
@@ -186,130 +186,87 @@ class ICLoraPipeline:
fps=frame_rate,
)
# Encode conditionings before loading transformer to reduce peak VRAM
video_encoder = self.stage_1_model_ledger.video_encoder()
stage_1_conditionings = self._create_conditionings(
images=images,
video_conditioning=video_conditioning,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=video_encoder,
num_frames=num_frames,
conditioning_attention_strength=conditioning_attention_strength,
conditioning_attention_mask=conditioning_attention_mask,
# Encode conditionings using the video encoder block
stage_1_conditionings = self.image_conditioner(
lambda enc: self._create_conditionings(
images=images,
video_conditioning=video_conditioning,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=enc,
num_frames=num_frames,
conditioning_attention_strength=conditioning_attention_strength,
conditioning_attention_mask=conditioning_attention_mask,
)
)
transformer = self.stage_1_model_ledger.transformer()
stage_1_sigmas = torch.Tensor(DISTILLED_SIGMA_VALUES).to(self.device)
def first_stage_denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return euler_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=simple_denoising_func(
video_context=video_context,
audio_context=audio_context,
transformer=transformer, # noqa: F821
),
)
video_state, audio_state = denoise_audio_video(
output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings,
noiser=noiser,
video_state, audio_state = self.stage_1(
denoiser=SimpleDenoiser(video_context, audio_context),
sigmas=stage_1_sigmas,
stepper=stepper,
denoising_loop_fn=first_stage_denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
noiser=noiser,
width=stage_1_output_shape.width,
height=stage_1_output_shape.height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(
context=video_context,
conditionings=stage_1_conditionings,
),
audio=ModalitySpec(
context=audio_context,
),
streaming_prefetch_count=streaming_prefetch_count,
)
torch.cuda.synchronize()
del transformer
cleanup_memory()
if skip_stage_2:
# Skip Stage 2: Decode directly from Stage 1 output at half resolution
logging.info("[IC-LoRA] Skipping Stage 2 (--skip-stage-2 enabled)")
decoded_video = vae_decode_video(
video_state.latent, self.stage_1_model_ledger.video_decoder(), tiling_config, generator
)
decoded_audio = vae_decode_audio(
audio_state.latent, self.stage_1_model_ledger.audio_decoder(), self.stage_1_model_ledger.vocoder()
)
del video_encoder
cleanup_memory()
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
decoded_audio = self.audio_decoder(audio_state.latent)
return decoded_video, decoded_audio
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
upscaled_video_latent = upsample_video(
latent=video_state.latent[:1],
video_encoder=video_encoder,
upsampler=self.stage_2_model_ledger.spatial_upsampler(),
)
upscaled_video_latent = self.upsampler(video_state.latent[:1])
torch.cuda.synchronize()
cleanup_memory()
transformer = self.stage_2_model_ledger.transformer()
distilled_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
def second_stage_denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return euler_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=simple_denoising_func(
video_context=video_context,
audio_context=audio_context,
transformer=transformer, # noqa: F821
),
)
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = combined_image_conditionings(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=video_encoder,
dtype=self.dtype,
device=self.device,
stage_2_conditionings = self.image_conditioner(
lambda enc: combined_image_conditionings(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=enc,
dtype=self.dtype,
device=self.device,
)
)
video_state, audio_state = denoise_audio_video(
output_shape=stage_2_output_shape,
conditionings=stage_2_conditionings,
noiser=noiser,
video_state, audio_state = self.stage_2(
denoiser=SimpleDenoiser(video_context, audio_context),
sigmas=distilled_sigmas,
stepper=stepper,
denoising_loop_fn=second_stage_denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
noise_scale=distilled_sigmas[0],
initial_video_latent=upscaled_video_latent,
initial_audio_latent=audio_state.latent,
noiser=noiser,
width=width,
height=height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(
context=video_context,
conditionings=stage_2_conditionings,
noise_scale=distilled_sigmas[0].item(),
initial_latent=upscaled_video_latent,
),
audio=ModalitySpec(
context=audio_context,
noise_scale=distilled_sigmas[0].item(),
initial_latent=audio_state.latent,
),
streaming_prefetch_count=streaming_prefetch_count,
)
torch.cuda.synchronize()
del transformer
del video_encoder
cleanup_memory()
decoded_video = vae_decode_video(
video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config, generator
)
decoded_audio = vae_decode_audio(
audio_state.latent, self.stage_2_model_ledger.audio_decoder(), self.stage_2_model_ledger.vocoder()
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
decoded_audio = self.audio_decoder(audio_state.latent)
return decoded_video, decoded_audio
def _create_conditionings(
@@ -358,14 +315,8 @@ class ICLoraPipeline:
for video_path, strength in video_conditioning:
# Load video at scaled-down resolution (if scale > 1)
video = load_video_conditioning(
video_path=video_path,
height=ref_height,
width=ref_width,
frame_cap=num_frames,
dtype=self.dtype,
device=self.device,
)
frame_gen = decode_video_by_frame(path=video_path, frame_cap=num_frames, device=self.device)
video = video_preprocess(frame_gen, ref_height, ref_width, self.dtype, self.device)
encoded_video = video_encoder(video)
reference_video_shape = VideoLatentShape.from_torch_shape(encoded_video.shape)
@@ -509,6 +460,7 @@ def main() -> None:
gemma_root=args.gemma_root,
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -525,6 +477,7 @@ def main() -> None:
conditioning_attention_strength=conditioning_attention_strength,
skip_stage_2=args.skip_stage_2,
conditioning_attention_mask=conditioning_attention_mask,
streaming_prefetch_count=args.streaming_prefetch_count,
)
encode_video(
@@ -553,17 +506,12 @@ def _load_mask_video(
Returns:
Tensor of shape ``(1, 1, F, H, W)`` with values in ``[0, 1]``.
"""
mask_video = load_video_conditioning(
video_path=mask_path,
height=height,
width=width,
frame_cap=num_frames,
dtype=torch.bfloat16,
device=device,
)
device = get_device()
frame_gen = decode_video_by_frame(path=mask_path, frame_cap=num_frames, device=device)
mask_video = video_preprocess(frame_gen, height, width, torch.bfloat16, device)
# mask_video shape: (1, C, F, H, W) — take mean over channels for grayscale
mask = mask_video.mean(dim=1, keepdim=True) # (1, 1, F, H, W)
# Normalise to [0, 1] — load_video_conditioning applies normalize_latent,
# Normalise to [0, 1] — video_preprocess applies normalize_latent,
# so undo that: values are in [-1, 1], remap to [0, 1]
mask = (mask + 1.0) / 2.0
return mask.clamp(0.0, 1.0)
@@ -3,40 +3,39 @@ from collections.abc import Iterator
import torch
from ltx_core.components.diffusion_steps import EulerDiffusionStep
from ltx_core.components.guiders import (
MultiModalGuiderFactory,
MultiModalGuiderParams,
create_multimodal_guider_factory,
)
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.components.schedulers import LTX2Scheduler
from ltx_core.loader import LoraPathStrengthAndSDOps
from ltx_core.model.audio_vae import decode_audio as vae_decode_audio
from ltx_core.model.upsampler import upsample_video
from ltx_core.loader.registry import Registry
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy
from ltx_core.types import Audio, LatentState, VideoPixelShape
from ltx_pipelines.utils import ModelLedger
from ltx_core.types import Audio, VideoPixelShape
from ltx_pipelines.utils.args import ImageConditioningInput, default_2_stage_arg_parser, detect_checkpoint_path
from ltx_pipelines.utils.constants import STAGE_2_DISTILLED_SIGMA_VALUES, detect_params
from ltx_pipelines.utils.blocks import (
AudioDecoder,
DiffusionStage,
ImageConditioner,
PromptEncoder,
VideoDecoder,
VideoUpsampler,
)
from ltx_pipelines.utils.constants import (
STAGE_2_DISTILLED_SIGMA_VALUES,
detect_params,
)
from ltx_pipelines.utils.denoisers import FactoryGuidedDenoiser, SimpleDenoiser
from ltx_pipelines.utils.helpers import (
assert_resolution,
cleanup_memory,
denoise_audio_video,
encode_prompts,
get_device,
image_conditionings_by_adding_guiding_latent,
multi_modal_guider_factory_denoising_func,
simple_denoising_func,
)
from ltx_pipelines.utils.media_io import encode_video
from ltx_pipelines.utils.samplers import euler_denoising_loop
from ltx_pipelines.utils.types import PipelineComponents
device = get_device()
from ltx_pipelines.utils.types import ModalitySpec
class KeyframeInterpolationPipeline:
@@ -56,27 +55,40 @@ class KeyframeInterpolationPipeline:
spatial_upsampler_path: str,
gemma_root: str,
loras: list[LoraPathStrengthAndSDOps],
device: torch.device = device,
device: torch.device | None = None,
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
):
self.device = device
self.device = device or get_device()
self.dtype = torch.bfloat16
self.stage_1_model_ledger = ModelLedger(
dtype=self.dtype,
device=device,
checkpoint_path=checkpoint_path,
spatial_upsampler_path=spatial_upsampler_path,
gemma_root_path=gemma_root,
loras=loras,
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
self.stage_1 = DiffusionStage(
checkpoint_path,
self.dtype,
self.device,
loras=tuple(loras),
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
self.stage_2_model_ledger = self.stage_1_model_ledger.with_additional_loras(
loras=distilled_lora,
stage_2_loras = (*tuple(loras), *tuple(distilled_lora))
self.stage_2 = DiffusionStage(
checkpoint_path,
self.dtype,
self.device,
loras=stage_2_loras,
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
self.pipeline_components = PipelineComponents(
dtype=self.dtype,
device=device,
self.upsampler = VideoUpsampler(
checkpoint_path, spatial_upsampler_path, self.dtype, self.device, registry=registry
)
self.video_decoder = VideoDecoder(checkpoint_path, self.dtype, self.device, registry=registry)
self.audio_decoder = AudioDecoder(checkpoint_path, self.dtype, self.device, registry=registry)
def __call__( # noqa: PLR0913
self,
@@ -93,52 +105,28 @@ class KeyframeInterpolationPipeline:
images: list[ImageConditioningInput],
tiling_config: TilingConfig | None = None,
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
) -> tuple[Iterator[torch.Tensor], Audio]:
assert_resolution(height=height, width=width, is_two_stage=True)
generator = torch.Generator(device=self.device).manual_seed(seed)
noiser = GaussianNoiser(generator=generator)
stepper = EulerDiffusionStep()
dtype = torch.bfloat16
ctx_p, ctx_n = encode_prompts(
ctx_p, ctx_n = self.prompt_encoder(
[prompt, negative_prompt],
self.stage_1_model_ledger,
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
# Stage 1: Initial low resolution video generation.
video_encoder = self.stage_1_model_ledger.video_encoder()
transformer = self.stage_1_model_ledger.transformer()
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
def first_stage_denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return euler_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=multi_modal_guider_factory_denoising_func(
video_guider_factory=create_multimodal_guider_factory(
params=video_guider_params,
negative_context=v_context_n,
),
audio_guider_factory=create_multimodal_guider_factory(
params=audio_guider_params,
negative_context=a_context_n,
),
v_context=v_context_p,
a_context=a_context_p,
transformer=transformer, # noqa: F821
),
)
stage_1_output_shape = VideoPixelShape(
batch=1,
frames=num_frames,
@@ -146,93 +134,90 @@ class KeyframeInterpolationPipeline:
height=height // 2,
fps=frame_rate,
)
stage_1_conditionings = image_conditionings_by_adding_guiding_latent(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
video_state, audio_state = denoise_audio_video(
output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings,
noiser=noiser,
sigmas=sigmas,
stepper=stepper,
denoising_loop_fn=first_stage_denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
stage_1_conditionings = self.image_conditioner(
lambda enc: image_conditionings_by_adding_guiding_latent(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=enc,
dtype=dtype,
device=self.device,
)
)
torch.cuda.synchronize()
del transformer
cleanup_memory()
video_guider_factory = create_multimodal_guider_factory(
params=video_guider_params,
negative_context=v_context_n,
)
audio_guider_factory = create_multimodal_guider_factory(
params=audio_guider_params,
negative_context=a_context_n,
)
video_state, audio_state = self.stage_1(
denoiser=FactoryGuidedDenoiser(
v_context=v_context_p,
a_context=a_context_p,
video_guider_factory=video_guider_factory,
audio_guider_factory=audio_guider_factory,
),
sigmas=sigmas,
noiser=noiser,
width=stage_1_output_shape.width,
height=stage_1_output_shape.height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(
context=v_context_p,
conditionings=stage_1_conditionings,
),
audio=ModalitySpec(
context=a_context_p,
),
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
upscaled_video_latent = upsample_video(
latent=video_state.latent[:1],
video_encoder=video_encoder,
upsampler=self.stage_2_model_ledger.spatial_upsampler(),
)
upscaled_video_latent = self.upsampler(video_state.latent[:1])
torch.cuda.synchronize()
cleanup_memory()
transformer = self.stage_2_model_ledger.transformer()
distilled_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
def second_stage_denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return euler_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=simple_denoising_func(
video_context=v_context_p,
audio_context=a_context_p,
transformer=transformer, # noqa: F821
),
)
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = image_conditionings_by_adding_guiding_latent(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
stage_2_conditionings = self.image_conditioner(
lambda enc: image_conditionings_by_adding_guiding_latent(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=enc,
dtype=dtype,
device=self.device,
)
)
video_state, audio_state = denoise_audio_video(
output_shape=stage_2_output_shape,
conditionings=stage_2_conditionings,
noiser=noiser,
video_state, audio_state = self.stage_2(
denoiser=SimpleDenoiser(v_context_p, a_context_p),
sigmas=distilled_sigmas,
stepper=stepper,
denoising_loop_fn=second_stage_denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
noise_scale=distilled_sigmas[0],
initial_video_latent=upscaled_video_latent,
initial_audio_latent=audio_state.latent,
noiser=noiser,
width=width,
height=height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(
context=v_context_p,
conditionings=stage_2_conditionings,
noise_scale=distilled_sigmas[0].item(),
initial_latent=upscaled_video_latent,
),
audio=ModalitySpec(
context=a_context_p,
noise_scale=distilled_sigmas[0].item(),
initial_latent=audio_state.latent,
),
streaming_prefetch_count=streaming_prefetch_count,
)
torch.cuda.synchronize()
del transformer
del video_encoder
cleanup_memory()
decoded_video = vae_decode_video(
video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config, generator
)
decoded_audio = vae_decode_audio(
audio_state.latent, self.stage_2_model_ledger.audio_decoder(), self.stage_2_model_ledger.vocoder()
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
decoded_audio = self.audio_decoder(audio_state.latent)
return decoded_video, decoded_audio
@@ -250,6 +235,7 @@ def main() -> None:
gemma_root=args.gemma_root,
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -280,6 +266,8 @@ def main() -> None:
),
images=args.images,
tiling_config=tiling_config,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
encode_video(
+148 -304
View File
@@ -1,153 +1,42 @@
from __future__ import annotations
import argparse
import logging
from collections.abc import Iterator
from dataclasses import dataclass
import torch
from ltx_core.components.diffusion_steps import EulerDiffusionStep
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.patchifiers import get_pixel_coords
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.components.schedulers import LTX2Scheduler
from ltx_core.conditioning import ConditioningItem
from ltx_core.conditioning.types.noise_mask_cond import TemporalRegionMask
from ltx_core.loader import LoraPathStrengthAndSDOps
from ltx_core.model.audio_vae import decode_audio as vae_decode_audio
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
from ltx_core.loader.registry import Registry
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy
from ltx_core.tools import LatentTools
from ltx_core.types import (
Audio,
AudioLatentShape,
LatentState,
SpatioTemporalScaleFactors,
VideoPixelShape,
)
from ltx_pipelines.utils import ModelLedger
from ltx_pipelines.utils.args import QuantizationAction
from ltx_pipelines.utils.args import video_editing_arg_parser
from ltx_pipelines.utils.blocks import (
AudioConditioner,
AudioDecoder,
DiffusionStage,
ImageConditioner,
PromptEncoder,
VideoDecoder,
)
from ltx_pipelines.utils.constants import DISTILLED_SIGMA_VALUES, detect_params
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
from ltx_pipelines.utils.helpers import (
cleanup_memory,
encode_prompts,
audio_latent_from_file,
get_device,
multi_modal_guider_denoising_func,
noise_audio_state,
noise_video_state,
simple_denoising_func,
video_latent_from_file,
)
from ltx_pipelines.utils.media_io import (
decode_audio_from_file,
encode_video,
get_videostream_metadata,
load_video_conditioning,
)
from ltx_pipelines.utils.samplers import euler_denoising_loop
from ltx_pipelines.utils.types import PipelineComponents
device = get_device()
def _encode_video_for_retake(
video_encoder: torch.nn.Module,
video_path: str,
output_shape: VideoPixelShape,
dtype: torch.dtype,
device: torch.device,
) -> torch.Tensor:
"""Load video and encode to latents."""
pixel_video = load_video_conditioning(
video_path=video_path,
height=output_shape.height,
width=output_shape.width,
frame_cap=output_shape.frames,
dtype=dtype,
device=device,
) # (1, C, F, H, W)
return video_encoder(pixel_video)
def _encode_audio_for_retake(
audio_encoder: torch.nn.Module,
waveform: torch.Tensor,
waveform_sr: int,
output_shape: VideoPixelShape,
dtype: torch.dtype,
) -> torch.Tensor:
"""Encode audio to latents and trim/pad to match output_shape."""
waveform_batch = waveform.unsqueeze(0) if waveform.dim() == 2 else waveform
initial_audio_latent = vae_encode_audio(
Audio(waveform=waveform_batch.to(dtype), sampling_rate=waveform_sr), audio_encoder, None
)
expected_audio_shape = AudioLatentShape.from_video_pixel_shape(output_shape)
expected_frames = expected_audio_shape.frames
actual_frames = initial_audio_latent.shape[2]
if actual_frames > expected_frames:
initial_audio_latent = initial_audio_latent[:, :, :expected_frames, :]
elif actual_frames < expected_frames:
pad = torch.zeros(
initial_audio_latent.shape[0],
initial_audio_latent.shape[1],
expected_frames - actual_frames,
initial_audio_latent.shape[3],
device=initial_audio_latent.device,
dtype=initial_audio_latent.dtype,
)
initial_audio_latent = torch.cat([initial_audio_latent, pad], dim=2)
return initial_audio_latent
# ---------------------------------------------------------------------------
# Custom conditioning item: temporal region mask
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class TemporalRegionMask:
"""Conditioning item that sets ``denoise_mask = 0`` outside a time range
and ``1`` inside, so only the specified temporal region is regenerated.
Uses ``start_time`` and ``end_time`` in seconds. Works in *patchified*
(token) space using the patchifier's ``get_patch_grid_bounds``: for video
coords are latent frame indices (converted from seconds via ``fps``), for
audio coords are already in seconds.
"""
start_time: float # seconds, inclusive
end_time: float # seconds, exclusive
fps: float
def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState:
coords = latent_tools.patchifier.get_patch_grid_bounds(
latent_tools.target_shape, device=latent_state.denoise_mask.device
)
# coords: [B, 3, N, 2] (video) or [B, 1, N, 2] (audio); temporal dim is index 0
if coords.shape[1] == 1:
# Audio: patchifier returns seconds
t_start = coords[:, 0, :, 0] # [B, N]
t_end = coords[:, 0, :, 1] # [B, N]
in_region = (t_end > self.start_time) & (t_start < self.end_time)
else:
# Video: get pixel bounds per patch, find patches for start/end frame, read latent from coords.
scale_factors = getattr(latent_tools, "scale_factors", SpatioTemporalScaleFactors.default())
pixel_bounds = get_pixel_coords(coords, scale_factors, causal_fix=getattr(latent_tools, "causal_fix", True))
timestamp_bounds = pixel_bounds[0, 0] / self.fps
t_start, t_end = timestamp_bounds.unbind(dim=-1)
in_region = (t_end > self.start_time) & (t_start < self.end_time)
state = latent_state.clone()
mask_val = in_region.to(state.denoise_mask.dtype)
if state.denoise_mask.dim() == 3:
mask_val = mask_val.unsqueeze(-1)
state.denoise_mask.copy_(mask_val)
return state
# ---------------------------------------------------------------------------
# Pipeline
# ---------------------------------------------------------------------------
from ltx_pipelines.utils.types import ModalitySpec
class RetakePipeline:
@@ -168,6 +57,11 @@ class RetakePipeline:
Target device (default: CUDA if available).
quantization : QuantizationPolicy | None
Optional quantization policy for the transformer.
distilled : bool
Set to ``True`` if using distilled model or passing distillation
lora with full model. If set to ``True``, distilled sigma schedule
(``DISTILLED_SIGMA_VALUES``) and a simple (non-guided) denoising
function will be used during ``__call__``.
"""
def __init__(
@@ -175,29 +69,61 @@ class RetakePipeline:
checkpoint_path: str,
gemma_root: str,
loras: list[LoraPathStrengthAndSDOps],
device: torch.device = device,
device: torch.device | None = None,
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
distilled: bool = True,
torch_compile: bool = False,
):
self.device = device
self.device = device or get_device()
self.dtype = torch.bfloat16
self.model_ledger = ModelLedger(
dtype=self.dtype,
device=device,
self.distilled = distilled
self.prompt_encoder = PromptEncoder(
checkpoint_path=checkpoint_path,
gemma_root_path=gemma_root,
loras=loras,
quantization=quantization,
)
self.pipeline_components = PipelineComponents(
gemma_root=gemma_root,
dtype=self.dtype,
device=device,
device=self.device,
registry=registry,
)
self.image_conditioner = ImageConditioner(
checkpoint_path=checkpoint_path,
dtype=self.dtype,
device=self.device,
registry=registry,
)
self.audio_conditioner = AudioConditioner(
checkpoint_path=checkpoint_path,
dtype=self.dtype,
device=self.device,
registry=registry,
)
self.stage = DiffusionStage(
checkpoint_path=checkpoint_path,
dtype=self.dtype,
device=self.device,
loras=tuple(loras),
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
self.video_decoder = VideoDecoder(
checkpoint_path=checkpoint_path,
dtype=self.dtype,
device=self.device,
registry=registry,
)
self.audio_decoder = AudioDecoder(
checkpoint_path=checkpoint_path,
dtype=self.dtype,
device=self.device,
registry=registry,
)
# --------------------------------------------------------------------- #
# Public entry point #
# --------------------------------------------------------------------- #
def __call__( # noqa: PLR0913, PLR0915
def __call__( # noqa: PLR0913
self,
video_path: str,
prompt: str,
@@ -212,8 +138,9 @@ class RetakePipeline:
regenerate_video: bool = True,
regenerate_audio: bool = True,
enhance_prompt: bool = False,
distilled: bool = False,
tiling_config: TilingConfig | None = None,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
) -> tuple[Iterator[torch.Tensor], torch.Tensor]:
"""Regenerate ``[start_time, end_time]`` of the source video (retake).
Parameters
@@ -235,19 +162,13 @@ class RetakePipeline:
Guidance parameters for video and audio modalities. Ignored in
distilled mode.
regenerate_video : bool
If ``True`` (default), preserve video outside ``[start_time, end_time]``
and only regenerate the masked region. If ``False``, fully regenerate
all video frames (the encoded video is still used as the initial latent
but with ``denoise_mask = 1`` everywhere).
If ``True`` (default), regenerate video inside ``[start_time, end_time]``.
If ``False``, video is preserved as-is (no regeneration).
regenerate_audio : bool
If True, regenerate audio in the [start_time, end_time] window; if False,
audio is preserved as-is (no regeneration).
enhance_prompt : bool
Whether to enhance the prompt via the text encoder.
distilled : bool
If ``True``, use the distilled sigma schedule
(``DISTILLED_SIGMA_VALUES``) and a simple (non-guided) denoising
function. The model checkpoint must be the distilled variant.
Returns
-------
tuple[Iterator[torch.Tensor], torch.Tensor]
@@ -256,95 +177,66 @@ class RetakePipeline:
if start_time >= end_time:
raise ValueError(f"start_time ({start_time}) must be less than end_time ({end_time})")
effective_seed = torch.randint(0, 2**31, (1,), device=self.device).item() if seed < 0 else seed
generator = torch.Generator(device=self.device).manual_seed(effective_seed)
generator = torch.Generator(device=self.device).manual_seed(seed)
noiser = GaussianNoiser(generator=generator)
stepper = EulerDiffusionStep()
dtype = self.dtype
video_encoder = self.model_ledger.video_encoder()
# Use av to get metadata
fps, num_pixel_frames, src_width, src_height = get_videostream_metadata(video_path)
output_shape = VideoPixelShape(
batch=1,
frames=num_pixel_frames,
width=src_width,
height=src_height,
fps=fps,
)
initial_video_latent = _encode_video_for_retake(
video_encoder=video_encoder,
video_path=video_path,
output_shape=output_shape,
dtype=dtype,
device=self.device,
)
video_conditionings: list[ConditioningItem] = [
TemporalRegionMask(
start_time=start_time if regenerate_video else 0.0,
end_time=end_time if regenerate_video else 0.0,
fps=fps,
)
]
del video_encoder
cleanup_memory()
initial_audio_latent: torch.Tensor | None = None
audio_conditionings: list[ConditioningItem] = []
audio_in = decode_audio_from_file(video_path, self.device)
audio_encoder = self.model_ledger.audio_encoder()
if audio_in is not None:
waveform = audio_in.waveform.squeeze(0)
waveform_sr = audio_in.sampling_rate
else:
waveform, waveform_sr = None, None
if waveform is not None:
initial_audio_latent = _encode_audio_for_retake(
audio_encoder=audio_encoder,
waveform=waveform,
waveform_sr=waveform_sr,
output_shape = get_videostream_metadata(video_path)
initial_video_latent = self.image_conditioner(
lambda enc: video_latent_from_file(
video_encoder=enc,
file_path=video_path,
output_shape=output_shape,
dtype=dtype,
device=self.device,
)
audio_conditionings = [
TemporalRegionMask(
start_time=start_time if regenerate_audio else 0.0,
end_time=end_time if regenerate_audio else 0.0,
fps=fps,
)
]
)
del audio_encoder
cleanup_memory()
initial_audio_latent = self.audio_conditioner(
lambda enc: audio_latent_from_file(
audio_encoder=enc,
file_path=video_path,
output_shape=output_shape,
dtype=dtype,
device=self.device,
)
)
prompts_to_encode = [prompt] if distilled else [prompt, negative_prompt]
contexts = encode_prompts(
prompts_to_encode = [prompt] if self.distilled else [prompt, negative_prompt]
contexts = self.prompt_encoder(
prompts_to_encode,
self.model_ledger,
enhance_first_prompt=enhance_prompt,
enhance_prompt_seed=effective_seed,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = contexts[0].video_encoding, contexts[0].audio_encoding
if not distilled:
v_context_n, a_context_n = contexts[1].video_encoding, contexts[1].audio_encoding
transformer = self.model_ledger.transformer()
sigmas = (
torch.tensor(DISTILLED_SIGMA_VALUES) if distilled else LTX2Scheduler().execute(steps=num_inference_steps)
).to(dtype=torch.float32, device=self.device)
if distilled:
denoise_fn = simple_denoising_func(
video_context=v_context_p,
audio_context=a_context_p,
transformer=transformer,
video_modality_spec = ModalitySpec(
context=v_context_p,
conditionings=[TemporalRegionMask(start_time=start_time, end_time=end_time, fps=output_shape.fps)]
if regenerate_video
else [],
initial_latent=initial_video_latent,
frozen=not regenerate_video,
)
audio_modality_spec = ModalitySpec(
context=a_context_p,
conditionings=[TemporalRegionMask(start_time=start_time, end_time=end_time, fps=output_shape.fps)]
if (initial_audio_latent is not None and regenerate_audio)
else [],
initial_latent=initial_audio_latent,
frozen=initial_audio_latent is not None and not regenerate_audio,
)
# Build denoiser
if self.distilled:
sigmas = torch.tensor(DISTILLED_SIGMA_VALUES).to(dtype=torch.float32, device=self.device)
denoiser = SimpleDenoiser(
v_context=v_context_p,
a_context=a_context_p,
)
else:
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
v_context_n, a_context_n = contexts[1].video_encoding, contexts[1].audio_encoding
video_guider = MultiModalGuider(
params=video_guider_params,
negative_context=v_context_n,
@@ -353,66 +245,31 @@ class RetakePipeline:
params=audio_guider_params,
negative_context=a_context_n,
)
denoise_fn = multi_modal_guider_denoising_func(
video_guider,
audio_guider,
denoiser = GuidedDenoiser(
v_context=v_context_p,
a_context=a_context_p,
transformer=transformer,
video_guider=video_guider,
audio_guider=audio_guider,
)
def denoising_loop(
sigmas: torch.Tensor,
video_state: LatentState,
audio_state: LatentState,
stepper: DiffusionStepProtocol,
) -> tuple[LatentState, LatentState]:
return euler_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=denoise_fn,
)
# Build noised states with the encoded latents as initial values and
# the temporal masks applied via conditionings.
video_state, video_tools = noise_video_state(
output_shape=output_shape,
# Run diffusion stage
video_state, audio_state = self.stage(
denoiser=denoiser,
sigmas=sigmas,
noiser=noiser,
conditionings=video_conditionings,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
initial_latent=initial_video_latent,
)
audio_state, audio_tools = noise_audio_state(
output_shape=output_shape,
noiser=noiser,
conditionings=audio_conditionings,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
initial_latent=initial_audio_latent,
width=output_shape.width,
height=output_shape.height,
frames=output_shape.frames,
fps=output_shape.fps,
video=video_modality_spec,
audio=audio_modality_spec,
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
video_state, audio_state = denoising_loop(sigmas, video_state, audio_state, stepper)
video_state = video_tools.clear_conditioning(video_state)
video_state = video_tools.unpatchify(video_state)
audio_state = audio_tools.clear_conditioning(audio_state)
audio_state = audio_tools.unpatchify(audio_state)
torch.cuda.synchronize()
del transformer
cleanup_memory()
decoded_video = vae_decode_video(
video_state.latent, self.model_ledger.video_decoder(), tiling_config, generator
)
decoded_audio = vae_decode_audio(
audio_state.latent, self.model_ledger.audio_decoder(), self.model_ledger.vocoder()
)
# Decode
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
decoded_audio = self.audio_decoder(audio_state.latent)
return decoded_video, decoded_audio
@@ -421,25 +278,8 @@ class RetakePipeline:
def main() -> None:
"""CLI entry point for retake (regenerate a time region)."""
logging.getLogger().setLevel(logging.INFO)
parser = argparse.ArgumentParser(description="Retake: regenerate a time region of a video with LTX-2.")
parser.add_argument("--video-path", type=str, required=True, help="Path to the source video.")
parser.add_argument("--prompt", type=str, required=True, help="Text prompt for the regenerated region.")
parser.add_argument("--start-time", type=float, required=True, help="Start time of the region to regenerate (s).")
parser.add_argument("--end-time", type=float, required=True, help="End time of the region to regenerate (s).")
parser.add_argument("--output-path", type=str, required=True, help="Path for the output video.")
parser.add_argument("--checkpoint-path", type=str, required=True, help="Path to the LTX-2 checkpoint.")
parser.add_argument("--gemma-root", type=str, required=True, help="Path to Gemma text encoder weights.")
parser.add_argument("--seed", type=int, default=42, help="Random seed. Use -1 for a random seed.")
parser.add_argument("--loras", nargs="*", default=[], help="LoRA paths (optional).")
parser.add_argument(
"--quantization",
dest="quantization",
action=QuantizationAction,
nargs="+",
metavar=("POLICY", "AMAX_PATH"),
default=None,
help="Quantization policy: fp8-cast or fp8-scaled-mm [AMAX_PATH].",
)
parser = video_editing_arg_parser(distilled=True)
parser.description = "Retake: regenerate a time region of a video with LTX-2."
args = parser.parse_args()
if args.start_time >= args.end_time:
@@ -447,22 +287,24 @@ def main() -> None:
# Validate frame count (8k+1) and resolution (multiples of 32) at CLI stage
video_scale = SpatioTemporalScaleFactors.default()
fps, num_frames, width, height = get_videostream_metadata(args.video_path)
if (num_frames - 1) % video_scale.time != 0:
snapped = ((num_frames - 1) // video_scale.time) * video_scale.time + 1
src = get_videostream_metadata(args.video_path)
if (src.frames - 1) % video_scale.time != 0:
snapped = ((src.frames - 1) // video_scale.time) * video_scale.time + 1
raise ValueError(
f"Video frame count must satisfy 8k+1 (e.g. 97, 193). Got {num_frames}; use a video with {snapped} frames."
f"Video frame count must satisfy 8k+1 (e.g. 97, 193). Got {src.frames}; use a video with {snapped} frames."
)
if width % 32 != 0 or height % 32 != 0:
raise ValueError(f"Video width and height must be multiples of 32. Got {width}x{height}.")
if src.width % 32 != 0 or src.height % 32 != 0:
raise ValueError(f"Video width and height must be multiples of 32. Got {src.width}x{src.height}.")
pipeline = RetakePipeline(
checkpoint_path=args.checkpoint_path,
checkpoint_path=args.distilled_checkpoint_path,
gemma_root=args.gemma_root,
loras=tuple(args.loras) if args.loras else (),
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
distilled=args.distilled,
torch_compile=args.compile,
)
params = detect_params(args.checkpoint_path)
params = detect_params(args.distilled_checkpoint_path)
tiling_config = TilingConfig.default()
video_iter, audio = pipeline(
video_path=args.video_path,
@@ -473,11 +315,13 @@ def main() -> None:
video_guider_params=params.video_guider_params,
audio_guider_params=params.audio_guider_params,
tiling_config=tiling_config,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
video_chunks_number = get_video_chunks_number(num_frames, tiling_config)
video_chunks_number = get_video_chunks_number(src.frames, tiling_config)
encode_video(
video=video_iter,
fps=int(fps),
fps=int(src.fps),
audio=audio,
output_path=args.output_path,
video_chunks_number=video_chunks_number,
@@ -3,37 +3,35 @@ from collections.abc import Iterator
import torch
from ltx_core.components.diffusion_steps import EulerDiffusionStep
from ltx_core.components.guiders import (
MultiModalGuiderFactory,
MultiModalGuiderParams,
create_multimodal_guider_factory,
)
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.components.schedulers import LTX2Scheduler
from ltx_core.loader import LoraPathStrengthAndSDOps
from ltx_core.model.audio_vae import decode_audio as vae_decode_audio
from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.loader.registry import Registry
from ltx_core.model.video_vae.tiling import TilingConfig
from ltx_core.quantization import QuantizationPolicy
from ltx_core.types import Audio, LatentState, VideoPixelShape
from ltx_core.types import Audio
from ltx_pipelines.utils import (
ModelLedger,
assert_resolution,
cleanup_memory,
combined_image_conditionings,
denoise_audio_video,
encode_prompts,
euler_denoising_loop,
get_device,
multi_modal_guider_factory_denoising_func,
)
from ltx_pipelines.utils.args import ImageConditioningInput, default_1_stage_arg_parser, detect_checkpoint_path
from ltx_pipelines.utils.blocks import (
AudioDecoder,
DiffusionStage,
ImageConditioner,
PromptEncoder,
VideoDecoder,
)
from ltx_pipelines.utils.constants import detect_params
from ltx_pipelines.utils.denoisers import FactoryGuidedDenoiser
from ltx_pipelines.utils.media_io import encode_video
from ltx_pipelines.utils.types import PipelineComponents
device = get_device()
from ltx_pipelines.utils.types import ModalitySpec
class TI2VidOneStagePipeline:
@@ -50,22 +48,46 @@ class TI2VidOneStagePipeline:
checkpoint_path: str,
gemma_root: str,
loras: list[LoraPathStrengthAndSDOps],
device: torch.device = device,
device: torch.device | None = None,
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
):
self.dtype = torch.bfloat16
self.device = device
self.model_ledger = ModelLedger(
dtype=self.dtype,
device=device,
self.device = device or get_device()
self.prompt_encoder = PromptEncoder(
checkpoint_path=checkpoint_path,
gemma_root_path=gemma_root,
loras=loras,
quantization=quantization,
)
self.pipeline_components = PipelineComponents(
gemma_root=gemma_root,
dtype=self.dtype,
device=device,
device=self.device,
registry=registry,
)
self.image_conditioner = ImageConditioner(
checkpoint_path=checkpoint_path,
dtype=self.dtype,
device=self.device,
registry=registry,
)
self.stage = DiffusionStage(
checkpoint_path=checkpoint_path,
dtype=self.dtype,
device=self.device,
loras=tuple(loras),
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
self.video_decoder = VideoDecoder(
checkpoint_path=checkpoint_path,
dtype=self.dtype,
device=self.device,
registry=registry,
)
self.audio_decoder = AudioDecoder(
checkpoint_path=checkpoint_path,
dtype=self.dtype,
device=self.device,
registry=registry,
)
def __call__( # noqa: PLR0913
@@ -82,41 +104,37 @@ class TI2VidOneStagePipeline:
audio_guider_params: MultiModalGuiderParams | MultiModalGuiderFactory,
images: list[ImageConditioningInput],
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
tiling_config: TilingConfig | None = None,
max_batch_size: int = 1,
) -> tuple[Iterator[torch.Tensor], Audio]:
assert_resolution(height=height, width=width, is_two_stage=False)
generator = torch.Generator(device=self.device).manual_seed(seed)
noiser = GaussianNoiser(generator=generator)
stepper = EulerDiffusionStep()
dtype = torch.bfloat16
ctx_p, ctx_n = encode_prompts(
ctx_p, ctx_n = self.prompt_encoder(
[prompt, negative_prompt],
self.model_ledger,
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
# Encode image conditionings with the VAE encoder, then free it
# before loading the transformer to reduce peak VRAM.
stage_1_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
video_encoder = self.model_ledger.video_encoder()
stage_1_conditionings = combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
stage_1_conditionings = self.image_conditioner(
lambda enc: combined_image_conditionings(
images=images,
height=height,
width=width,
video_encoder=enc,
dtype=dtype,
device=self.device,
)
)
torch.cuda.synchronize()
del video_encoder
cleanup_memory()
transformer = self.model_ledger.transformer()
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
video_guider_factory = create_multimodal_guider_factory(
@@ -128,43 +146,32 @@ class TI2VidOneStagePipeline:
negative_context=a_context_n,
)
def first_stage_denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return euler_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=multi_modal_guider_factory_denoising_func(
video_guider_factory=video_guider_factory,
audio_guider_factory=audio_guider_factory,
v_context=v_context_p,
a_context=a_context_p,
transformer=transformer, # noqa: F821
),
)
video_state, audio_state = denoise_audio_video(
output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings,
noiser=noiser,
video_state, audio_state = self.stage(
denoiser=FactoryGuidedDenoiser(
v_context=v_context_p,
a_context=a_context_p,
video_guider_factory=video_guider_factory,
audio_guider_factory=audio_guider_factory,
),
sigmas=sigmas,
stepper=stepper,
denoising_loop_fn=first_stage_denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
noiser=noiser,
width=width,
height=height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(
context=v_context_p,
conditionings=stage_1_conditionings,
),
audio=ModalitySpec(
context=a_context_p,
),
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
torch.cuda.synchronize()
del transformer
cleanup_memory()
decoded_video = vae_decode_video(video_state.latent, self.model_ledger.video_decoder(), generator=generator)
decoded_audio = vae_decode_audio(
audio_state.latent, self.model_ledger.audio_decoder(), self.model_ledger.vocoder()
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator=generator)
decoded_audio = self.audio_decoder(audio_state.latent)
return decoded_video, decoded_audio
@@ -180,6 +187,7 @@ def main() -> None:
gemma_root=args.gemma_root,
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
)
video, audio = pipeline(
prompt=args.prompt,
@@ -207,6 +215,8 @@ def main() -> None:
stg_blocks=args.audio_stg_blocks,
),
images=args.images,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
encode_video(
@@ -3,40 +3,39 @@ from collections.abc import Iterator
import torch
from ltx_core.components.diffusion_steps import EulerDiffusionStep
from ltx_core.components.guiders import (
MultiModalGuiderFactory,
MultiModalGuiderParams,
create_multimodal_guider_factory,
)
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.components.schedulers import LTX2Scheduler
from ltx_core.loader import LoraPathStrengthAndSDOps
from ltx_core.model.audio_vae import decode_audio as vae_decode_audio
from ltx_core.model.upsampler import upsample_video
from ltx_core.loader.registry import Registry
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy
from ltx_core.types import Audio, LatentState, VideoPixelShape
from ltx_pipelines.utils import (
ModelLedger,
assert_resolution,
cleanup_memory,
combined_image_conditionings,
denoise_audio_video,
encode_prompts,
euler_denoising_loop,
get_device,
multi_modal_guider_factory_denoising_func,
simple_denoising_func,
)
from ltx_core.types import Audio, VideoPixelShape
from ltx_pipelines.utils.args import ImageConditioningInput, default_2_stage_arg_parser, detect_checkpoint_path
from ltx_pipelines.utils.constants import STAGE_2_DISTILLED_SIGMA_VALUES, detect_params
from ltx_pipelines.utils.blocks import (
AudioDecoder,
DiffusionStage,
ImageConditioner,
PromptEncoder,
VideoDecoder,
VideoUpsampler,
)
from ltx_pipelines.utils.constants import (
STAGE_2_DISTILLED_SIGMA_VALUES,
detect_params,
)
from ltx_pipelines.utils.denoisers import FactoryGuidedDenoiser, SimpleDenoiser
from ltx_pipelines.utils.helpers import (
assert_resolution,
combined_image_conditionings,
get_device,
)
from ltx_pipelines.utils.media_io import encode_video
from ltx_pipelines.utils.types import PipelineComponents
device = get_device()
from ltx_pipelines.utils.types import ModalitySpec
class TI2VidTwoStagesPipeline:
@@ -55,28 +54,39 @@ class TI2VidTwoStagesPipeline:
spatial_upsampler_path: str,
gemma_root: str,
loras: list[LoraPathStrengthAndSDOps],
device: torch.device = device,
device: torch.device | None = None,
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
):
self.device = device
self.device = device or get_device()
self.dtype = torch.bfloat16
self.stage_1_model_ledger = ModelLedger(
dtype=self.dtype,
device=device,
checkpoint_path=checkpoint_path,
gemma_root_path=gemma_root,
spatial_upsampler_path=spatial_upsampler_path,
loras=loras,
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
self.upsampler = VideoUpsampler(
checkpoint_path, spatial_upsampler_path, self.dtype, self.device, registry=registry
)
self.video_decoder = VideoDecoder(checkpoint_path, self.dtype, self.device, registry=registry)
self.audio_decoder = AudioDecoder(checkpoint_path, self.dtype, self.device, registry=registry)
self.stage_1 = DiffusionStage(
checkpoint_path,
self.dtype,
self.device,
loras=tuple(loras),
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
self.stage_2_model_ledger = self.stage_1_model_ledger.with_additional_loras(
loras=distilled_lora,
)
self.pipeline_components = PipelineComponents(
dtype=self.dtype,
device=device,
self.stage_2 = DiffusionStage(
checkpoint_path,
self.dtype,
self.device,
loras=(*tuple(loras), *distilled_lora),
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
def __call__( # noqa: PLR0913
@@ -94,26 +104,26 @@ class TI2VidTwoStagesPipeline:
images: list[ImageConditioningInput],
tiling_config: TilingConfig | None = None,
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
) -> tuple[Iterator[torch.Tensor], Audio]:
assert_resolution(height=height, width=width, is_two_stage=True)
generator = torch.Generator(device=self.device).manual_seed(seed)
noiser = GaussianNoiser(generator=generator)
stepper = EulerDiffusionStep()
dtype = torch.bfloat16
ctx_p, ctx_n = encode_prompts(
ctx_p, ctx_n = self.prompt_encoder(
[prompt, negative_prompt],
self.stage_1_model_ledger,
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
# Stage 1: encode image conditionings with the VAE encoder, then free it
# before loading the transformer to reduce peak VRAM.
# Stage 1: Generate video at half resolution with CFG guidance.
stage_1_output_shape = VideoPixelShape(
batch=1,
frames=num_frames,
@@ -121,125 +131,83 @@ class TI2VidTwoStagesPipeline:
height=height // 2,
fps=frame_rate,
)
video_encoder = self.stage_1_model_ledger.video_encoder()
stage_1_conditionings = combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
stage_1_conditionings = self.image_conditioner(
lambda enc: combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=enc,
dtype=dtype,
device=self.device,
)
)
torch.cuda.synchronize()
del video_encoder
cleanup_memory()
transformer = self.stage_1_model_ledger.transformer()
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
def first_stage_denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return euler_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=multi_modal_guider_factory_denoising_func(
video_guider_factory=create_multimodal_guider_factory(
params=video_guider_params,
negative_context=v_context_n,
),
audio_guider_factory=create_multimodal_guider_factory(
params=audio_guider_params,
negative_context=a_context_n,
),
v_context=v_context_p,
a_context=a_context_p,
transformer=transformer, # noqa: F821
video_state, audio_state = self.stage_1(
denoiser=FactoryGuidedDenoiser(
v_context=v_context_p,
a_context=a_context_p,
video_guider_factory=create_multimodal_guider_factory(
params=video_guider_params,
negative_context=v_context_n,
),
)
video_state, audio_state = denoise_audio_video(
output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings,
noiser=noiser,
audio_guider_factory=create_multimodal_guider_factory(
params=audio_guider_params,
negative_context=a_context_n,
),
),
sigmas=sigmas,
stepper=stepper,
denoising_loop_fn=first_stage_denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
)
torch.cuda.synchronize()
del transformer
cleanup_memory()
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
video_encoder = self.stage_1_model_ledger.video_encoder()
upscaled_video_latent = upsample_video(
latent=video_state.latent[:1],
video_encoder=video_encoder,
upsampler=self.stage_2_model_ledger.spatial_upsampler(),
)
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = combined_image_conditionings(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
del video_encoder
torch.cuda.synchronize()
cleanup_memory()
transformer = self.stage_2_model_ledger.transformer()
distilled_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
def second_stage_denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return euler_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=simple_denoising_func(
video_context=v_context_p,
audio_context=a_context_p,
transformer=transformer, # noqa: F821
),
)
video_state, audio_state = denoise_audio_video(
output_shape=stage_2_output_shape,
conditionings=stage_2_conditionings,
noiser=noiser,
width=stage_1_output_shape.width,
height=stage_1_output_shape.height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(context=v_context_p, conditionings=stage_1_conditionings),
audio=ModalitySpec(context=a_context_p),
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
# Stage 2: Upsample and refine the video at higher resolution with distilled LoRA.
upscaled_video_latent = self.upsampler(video_state.latent[:1])
distilled_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
stage_2_conditionings = self.image_conditioner(
lambda enc: combined_image_conditionings(
images=images,
height=height,
width=width,
video_encoder=enc,
dtype=dtype,
device=self.device,
)
)
video_state, audio_state = self.stage_2(
denoiser=SimpleDenoiser(v_context=v_context_p, a_context=a_context_p),
sigmas=distilled_sigmas,
stepper=stepper,
denoising_loop_fn=second_stage_denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
noise_scale=distilled_sigmas[0],
initial_video_latent=upscaled_video_latent,
initial_audio_latent=audio_state.latent,
noiser=noiser,
width=width,
height=height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(
context=v_context_p,
conditionings=stage_2_conditionings,
noise_scale=distilled_sigmas[0].item(),
initial_latent=upscaled_video_latent,
),
audio=ModalitySpec(
context=a_context_p,
noise_scale=distilled_sigmas[0].item(),
initial_latent=audio_state.latent,
),
streaming_prefetch_count=streaming_prefetch_count,
)
torch.cuda.synchronize()
del transformer
cleanup_memory()
decoded_video = vae_decode_video(
video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config, generator
)
decoded_audio = vae_decode_audio(
audio_state.latent, self.stage_2_model_ledger.audio_decoder(), self.stage_2_model_ledger.vocoder()
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
decoded_audio = self.audio_decoder(audio_state.latent)
return decoded_video, decoded_audio
@@ -257,6 +225,7 @@ def main() -> None:
gemma_root=args.gemma_root,
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -287,6 +256,8 @@ def main() -> None:
),
images=args.images,
tiling_config=tiling_config,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
encode_video(
@@ -6,34 +6,34 @@ import torch
from ltx_core.components.diffusion_steps import Res2sDiffusionStep
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.components.schedulers import LTX2Scheduler
from ltx_core.loader import LoraPathStrengthAndSDOps
from ltx_core.model.audio_vae import decode_audio as vae_decode_audio
from ltx_core.model.upsampler import upsample_video
from ltx_core.loader.registry import Registry
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy
from ltx_core.tools import VideoLatentShape
from ltx_core.types import Audio, LatentState, VideoPixelShape
from ltx_pipelines.utils import (
ModelLedger,
assert_resolution,
cleanup_memory,
combined_image_conditionings,
denoise_audio_video,
encode_prompts,
get_device,
multi_modal_guider_denoising_func,
res2s_audio_video_denoising_loop,
simple_denoising_func,
)
from ltx_core.types import Audio, VideoLatentShape, VideoPixelShape
from ltx_pipelines.utils.args import ImageConditioningInput, hq_2_stage_arg_parser
from ltx_pipelines.utils.constants import LTX_2_3_HQ_PARAMS, STAGE_2_DISTILLED_SIGMA_VALUES
from ltx_pipelines.utils.blocks import (
AudioDecoder,
DiffusionStage,
ImageConditioner,
PromptEncoder,
VideoDecoder,
VideoUpsampler,
)
from ltx_pipelines.utils.constants import (
LTX_2_3_HQ_PARAMS,
STAGE_2_DISTILLED_SIGMA_VALUES,
)
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
from ltx_pipelines.utils.helpers import (
assert_resolution,
combined_image_conditionings,
get_device,
)
from ltx_pipelines.utils.media_io import encode_video
from ltx_pipelines.utils.types import PipelineComponents
device = get_device()
from ltx_pipelines.utils.samplers import res2s_audio_video_denoising_loop
from ltx_pipelines.utils.types import ModalitySpec
class TI2VidTwoStagesHQPipeline:
@@ -48,7 +48,7 @@ class TI2VidTwoStagesHQPipeline:
the images parameter.
"""
def __init__(
def __init__( # noqa: PLR0913
self,
checkpoint_path: str,
distilled_lora: list[LoraPathStrengthAndSDOps],
@@ -57,11 +57,14 @@ class TI2VidTwoStagesHQPipeline:
spatial_upsampler_path: str,
gemma_root: str,
loras: tuple[LoraPathStrengthAndSDOps, ...],
device: str = device,
device: torch.device | None = None,
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
):
self.device = device
self.device = device or get_device()
self.dtype = torch.bfloat16
distilled_lora_stage_1 = LoraPathStrengthAndSDOps(
path=distilled_lora[0].path,
strength=distilled_lora_strength_stage_1,
@@ -72,23 +75,32 @@ class TI2VidTwoStagesHQPipeline:
strength=distilled_lora_strength_stage_2,
sd_ops=distilled_lora[0].sd_ops,
)
self.stage_1_model_ledger = ModelLedger(
dtype=self.dtype,
device=device,
checkpoint_path=checkpoint_path,
gemma_root_path=gemma_root,
spatial_upsampler_path=spatial_upsampler_path,
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
self.upsampler = VideoUpsampler(
checkpoint_path, spatial_upsampler_path, self.dtype, self.device, registry=registry
)
self.video_decoder = VideoDecoder(checkpoint_path, self.dtype, self.device, registry=registry)
self.audio_decoder = AudioDecoder(checkpoint_path, self.dtype, self.device, registry=registry)
self.stage_1 = DiffusionStage(
checkpoint_path,
self.dtype,
self.device,
loras=(*loras, distilled_lora_stage_1),
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
self.stage_2_model_ledger = self.stage_1_model_ledger.with_loras(
self.stage_2 = DiffusionStage(
checkpoint_path,
self.dtype,
self.device,
loras=(*loras, distilled_lora_stage_2),
)
self.pipeline_components = PipelineComponents(
dtype=self.dtype,
device=device,
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
)
@torch.inference_mode()
@@ -107,6 +119,8 @@ class TI2VidTwoStagesHQPipeline:
images: list[ImageConditioningInput],
tiling_config: TilingConfig | None = None,
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
) -> tuple[Iterator[torch.Tensor], Audio]:
assert_resolution(height=height, width=width, is_two_stage=True)
@@ -114,18 +128,17 @@ class TI2VidTwoStagesHQPipeline:
noiser = GaussianNoiser(generator=generator)
dtype = torch.bfloat16
ctx_p, ctx_n = encode_prompts(
ctx_p, ctx_n = self.prompt_encoder(
[prompt, negative_prompt],
self.stage_1_model_ledger,
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
# Stage 1: encode image conditionings with the VAE encoder, then free it
# before loading the transformer to reduce peak VRAM.
# Stage 1: Generate video at half resolution with CFG guidance using res2s sampler.
stage_1_output_shape = VideoPixelShape(
batch=1,
frames=num_frames,
@@ -133,20 +146,16 @@ class TI2VidTwoStagesHQPipeline:
height=height // 2,
fps=frame_rate,
)
video_encoder = self.stage_1_model_ledger.video_encoder()
stage_1_conditionings = combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
stage_1_conditionings = self.image_conditioner(
lambda enc: combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=enc,
dtype=dtype,
device=self.device,
)
)
torch.cuda.synchronize()
del video_encoder
cleanup_memory()
transformer = self.stage_1_model_ledger.transformer()
empty_latent = torch.empty(VideoLatentShape.from_pixel_shape(stage_1_output_shape).to_torch_shape())
stepper = Res2sDiffusionStep()
@@ -156,109 +165,75 @@ class TI2VidTwoStagesHQPipeline:
.to(dtype=torch.float32, device=self.device)
)
def first_stage_denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return res2s_audio_video_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=multi_modal_guider_denoising_func(
video_guider=MultiModalGuider(
params=video_guider_params,
negative_context=v_context_n,
),
audio_guider=MultiModalGuider(
params=audio_guider_params,
negative_context=a_context_n,
),
v_context=v_context_p,
a_context=a_context_p,
transformer=transformer, # noqa: F821
video_state, audio_state = self.stage_1(
denoiser=GuidedDenoiser(
v_context=v_context_p,
a_context=a_context_p,
video_guider=MultiModalGuider(
params=video_guider_params,
negative_context=v_context_n,
),
)
video_state, audio_state = denoise_audio_video(
output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings,
noiser=noiser,
audio_guider=MultiModalGuider(
params=audio_guider_params,
negative_context=a_context_n,
),
),
sigmas=sigmas,
stepper=stepper,
denoising_loop_fn=first_stage_denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
)
torch.cuda.synchronize()
del transformer
cleanup_memory()
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
video_encoder = self.stage_1_model_ledger.video_encoder()
upscaled_video_latent = upsample_video(
latent=video_state.latent[:1],
video_encoder=video_encoder,
upsampler=self.stage_2_model_ledger.spatial_upsampler(),
)
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = combined_image_conditionings(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
torch.cuda.synchronize()
del video_encoder
cleanup_memory()
transformer = self.stage_2_model_ledger.transformer()
distilled_sigmas = torch.tensor(STAGE_2_DISTILLED_SIGMA_VALUES, device=self.device)
def second_stage_denoising_loop(
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
) -> tuple[LatentState, LatentState]:
return res2s_audio_video_denoising_loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
denoise_fn=simple_denoising_func(
video_context=v_context_p,
audio_context=a_context_p,
transformer=transformer, # noqa: F821
),
)
video_state, audio_state = denoise_audio_video(
output_shape=stage_2_output_shape,
conditionings=stage_2_conditionings,
noiser=noiser,
sigmas=distilled_sigmas,
stepper=stepper,
denoising_loop_fn=second_stage_denoising_loop,
components=self.pipeline_components,
dtype=dtype,
device=self.device,
noise_scale=distilled_sigmas[0],
initial_video_latent=upscaled_video_latent,
initial_audio_latent=audio_state.latent,
width=stage_1_output_shape.width,
height=stage_1_output_shape.height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(context=v_context_p, conditionings=stage_1_conditionings),
audio=ModalitySpec(context=a_context_p),
loop=res2s_audio_video_denoising_loop,
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
torch.cuda.synchronize()
del transformer
cleanup_memory()
# Stage 2: Upsample and refine the video at higher resolution with distilled LoRA.
upscaled_video_latent = self.upsampler(video_state.latent[:1])
decoded_video = vae_decode_video(
video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config, generator
distilled_sigmas = torch.tensor(STAGE_2_DISTILLED_SIGMA_VALUES, device=self.device)
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = self.image_conditioner(
lambda enc: combined_image_conditionings(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=enc,
dtype=dtype,
device=self.device,
)
)
decoded_audio = vae_decode_audio(
audio_state.latent, self.stage_2_model_ledger.audio_decoder(), self.stage_2_model_ledger.vocoder()
video_state, audio_state = self.stage_2(
denoiser=SimpleDenoiser(v_context=v_context_p, a_context=a_context_p),
sigmas=distilled_sigmas,
noiser=noiser,
stepper=stepper,
width=width,
height=height,
frames=num_frames,
fps=frame_rate,
video=ModalitySpec(
context=v_context_p,
conditionings=stage_2_conditionings,
noise_scale=distilled_sigmas[0].item(),
initial_latent=upscaled_video_latent,
),
audio=ModalitySpec(
context=a_context_p,
noise_scale=distilled_sigmas[0].item(),
initial_latent=audio_state.latent,
),
loop=res2s_audio_video_denoising_loop,
streaming_prefetch_count=streaming_prefetch_count,
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
decoded_audio = self.audio_decoder(audio_state.latent)
return decoded_video, decoded_audio
@@ -276,6 +251,7 @@ def main() -> None:
gemma_root=args.gemma_root,
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -306,6 +282,8 @@ def main() -> None:
),
images=args.images,
tiling_config=tiling_config,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
encode_video(
@@ -1,35 +1,46 @@
from ltx_pipelines.utils.blocks import (
AudioConditioner,
AudioDecoder,
DiffusionStage,
ImageConditioner,
PromptEncoder,
VideoDecoder,
VideoUpsampler,
)
from ltx_pipelines.utils.denoisers import FactoryGuidedDenoiser, GuidedDenoiser, SimpleDenoiser
from ltx_pipelines.utils.helpers import (
assert_resolution,
cleanup_memory,
combined_image_conditionings,
denoise_audio_video,
encode_prompts,
generate_enhanced_prompt,
get_device,
multi_modal_guider_denoising_func,
multi_modal_guider_factory_denoising_func,
simple_denoising_func,
image_conditionings_by_adding_guiding_latent,
)
from ltx_pipelines.utils.model_ledger import ModelLedger
from ltx_pipelines.utils.samplers import (
euler_denoising_loop,
gradient_estimating_euler_denoising_loop,
res2s_audio_video_denoising_loop,
)
from ltx_pipelines.utils.types import Denoiser, ModalitySpec
__all__ = [
"ModelLedger",
"AudioConditioner",
"AudioDecoder",
"Denoiser",
"DiffusionStage",
"FactoryGuidedDenoiser",
"GuidedDenoiser",
"ImageConditioner",
"ModalitySpec",
"PromptEncoder",
"SimpleDenoiser",
"VideoDecoder",
"VideoUpsampler",
"assert_resolution",
"cleanup_memory",
"combined_image_conditionings",
"denoise_audio_video",
"encode_prompts",
"euler_denoising_loop",
"generate_enhanced_prompt",
"get_device",
"gradient_estimating_euler_denoising_loop",
"multi_modal_guider_denoising_func",
"multi_modal_guider_factory_denoising_func",
"image_conditionings_by_adding_guiding_latent",
"res2s_audio_video_denoising_loop",
"simple_denoising_func",
]
@@ -173,6 +173,15 @@ def basic_arg_parser(
required=True,
help="Path to LTX-2 model checkpoint (.safetensors file).",
)
parser.add_argument(
"--num-inference-steps",
type=int,
default=params.num_inference_steps,
help=(
f"Number of denoising steps in the diffusion sampling process. "
f"Higher values improve quality but increase generation time (default: {params.num_inference_steps})."
),
)
parser.add_argument(
"--gemma-root",
type=resolve_path,
@@ -197,6 +206,85 @@ def basic_arg_parser(
default=params.seed,
help=f"Random seed for reproducible generation (default: {params.seed}).",
)
parser.add_argument(
"--lora",
dest="lora",
action=LoraAction,
nargs="+", # Accept 1-2 arguments per use (path and optional strength); validation is handled in LoraAction
metavar=("PATH", "STRENGTH"),
default=[],
help=(
"LoRA (Low-Rank Adaptation) model: path to model file and optional strength "
f"(default strength: {DEFAULT_LORA_STRENGTH}). Can be specified multiple times. "
"Example: --lora path/to/lora1.safetensors 0.8 --lora path/to/lora2.safetensors"
),
)
parser.add_argument("--enhance-prompt", action="store_true")
def _positive_int(value: str) -> int:
try:
int_value = int(value)
if int_value < 1:
raise argparse.ArgumentTypeError("must be >= 1")
return int_value
except ValueError as e:
raise argparse.ArgumentTypeError(f"must be an integer, got {value}") from e
# Layer streaming
parser.add_argument(
"--streaming-prefetch-count",
type=_positive_int,
default=None,
metavar="N",
help=(
"Enable layer streaming prefetching N layers ahead. "
"At most 1 + N layers reside on GPU at once. "
"Must be >= 1. Example: --streaming-prefetch-count 2"
),
)
parser.add_argument(
"--max-batch-size",
type=_positive_int,
default=1,
metavar="N",
help=(
"Maximum batch size per transformer forward pass. "
"Guided denoisers batch up to 4 guidance passes into a single call. "
"Default 1 runs passes sequentially. Set to 4 to batch all passes "
"together, which reduces layer-streaming PCIe transfers. "
"Example: --max-batch-size 4"
),
)
parser.add_argument(
"--quantization",
dest="quantization",
action=QuantizationAction,
nargs="+",
metavar=("POLICY", "AMAX_PATH"),
default=None,
help=(
f"Quantization policy: {', '.join(QUANTIZATION_POLICIES)}. "
"fp8-cast uses FP8 casting with upcasting during inference. "
"fp8-scaled-mm uses FP8 scaled matrix multiplication (optionally provide amax calibration file path). "
"Example: --quantization fp8-cast or --quantization fp8-scaled-mm /path/to/amax.json"
),
)
parser.add_argument(
"--compile",
action="store_true",
help="Enable torch.compile for transformer blocks to optimize performance.",
)
return parser
def new_video_gen_arg_parser(
params: PipelineParams = LTX_2_3_PARAMS,
distilled: bool = False,
) -> argparse.ArgumentParser:
parser = basic_arg_parser(params=params, distilled=distilled)
parser.add_argument(
"--height",
type=int,
@@ -222,15 +310,6 @@ def basic_arg_parser(
default=params.frame_rate,
help=f"Frame rate of the generated video (fps) (default: {params.frame_rate}).",
)
parser.add_argument(
"--num-inference-steps",
type=int,
default=params.num_inference_steps,
help=(
f"Number of denoising steps in the diffusion sampling process. "
f"Higher values improve quality but increase generation time (default: {params.num_inference_steps})."
),
)
parser.add_argument(
"--image",
dest="images",
@@ -247,42 +326,28 @@ def basic_arg_parser(
"--image path/to/image2.jpg 160 0.9 0"
),
)
parser.add_argument(
"--lora",
dest="lora",
action=LoraAction,
nargs="+", # Accept 1-2 arguments per use (path and optional strength); validation is handled in LoraAction
metavar=("PATH", "STRENGTH"),
default=[],
help=(
"LoRA (Low-Rank Adaptation) model: path to model file and optional strength "
f"(default strength: {DEFAULT_LORA_STRENGTH}). Can be specified multiple times. "
"Example: --lora path/to/lora1.safetensors 0.8 --lora path/to/lora2.safetensors"
),
)
parser.add_argument("--enhance-prompt", action="store_true")
parser.add_argument(
"--quantization",
dest="quantization",
action=QuantizationAction,
nargs="+",
metavar=("POLICY", "AMAX_PATH"),
default=None,
help=(
f"Quantization policy: {', '.join(QUANTIZATION_POLICIES)}. "
"fp8-cast uses FP8 casting with upcasting during inference. "
"fp8-scaled-mm uses FP8 scaled matrix multiplication (optionally provide amax calibration file path). "
"Example: --quantization fp8-cast or --quantization fp8-scaled-mm /path/to/amax.json"
),
)
return parser
def video_editing_arg_parser(
distilled: bool = True,
) -> argparse.ArgumentParser:
"""Base argument parser for video-editing pipelines (retake, extension, inpainting, sticker movement).
Uses the same actions and conventions as basic_arg_parser but only the args needed for editing
(no height/width/num-frames; resolution comes from input video). Default is distilled checkpoint only.
"""
parser = basic_arg_parser(distilled=distilled)
parser.add_argument("--video-path", type=resolve_path, required=True, help="Path to the source video.")
parser.add_argument("--start-time", type=float, required=True, help="Start time of the region to regenerate (s).")
parser.add_argument("--end-time", type=float, required=True, help="End time of the region to regenerate (s).")
return parser
def default_1_stage_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argparse.ArgumentParser:
video_guider = params.video_guider_params
audio_guider = params.audio_guider_params
parser = basic_arg_parser(params=params)
parser = new_video_gen_arg_parser(params=params)
parser.add_argument(
"--negative-prompt",
type=str,
@@ -476,7 +541,7 @@ def hq_2_stage_arg_parser(params: PipelineParams = LTX_2_3_HQ_PARAMS) -> argpars
def default_2_stage_distilled_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argparse.ArgumentParser:
parser = basic_arg_parser(params=params, distilled=True)
parser = new_video_gen_arg_parser(params=params, distilled=True)
parser.set_defaults(height=params.stage_2_height, width=params.stage_2_width)
# Update help text to reflect 2-stage defaults
for action in parser._actions:
@@ -0,0 +1,574 @@
"""Pipeline blocks — each block owns its model lifecycle.
Blocks build a model on each ``__call__``, use it, then free GPU memory.
This eliminates manual ``del model; cleanup_memory()`` in pipelines and
removes the need for :class:`ModelLedger`.
"""
from __future__ import annotations
import logging
from collections.abc import Iterator
from contextlib import AbstractContextManager, contextmanager
from dataclasses import replace
from typing import Callable, TypeVar
import torch
from ltx_core.batch_split import BatchSplitAdapter
from ltx_core.components.diffusion_steps import EulerDiffusionStep
from ltx_core.components.noisers import Noiser
from ltx_core.components.patchifiers import AudioPatchifier, VideoLatentPatchifier
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.layer_streaming import LayerStreamingWrapper
from ltx_core.loader import SDOps
from ltx_core.loader.primitives import LoraPathStrengthAndSDOps
from ltx_core.loader.registry import DummyRegistry, Registry
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder
from ltx_core.model.audio_vae import (
AUDIO_VAE_DECODER_COMFY_KEYS_FILTER,
AUDIO_VAE_ENCODER_COMFY_KEYS_FILTER,
VOCODER_COMFY_KEYS_FILTER,
AudioDecoderConfigurator,
AudioEncoderConfigurator,
VocoderConfigurator,
)
from ltx_core.model.audio_vae import (
decode_audio as vae_decode_audio,
)
from ltx_core.model.transformer import (
LTXV_MODEL_COMFY_RENAMING_MAP,
LTXModelConfigurator,
X0Model,
)
from ltx_core.model.transformer.compiling import COMPILE_TRANSFORMER, modify_sd_ops_for_compilation
from ltx_core.model.upsampler import LatentUpsamplerConfigurator, upsample_video
from ltx_core.model.video_vae import (
VAE_DECODER_COMFY_KEYS_FILTER,
VAE_ENCODER_COMFY_KEYS_FILTER,
TilingConfig,
VideoDecoderConfigurator,
VideoEncoder,
VideoEncoderConfigurator,
)
from ltx_core.quantization import QuantizationPolicy
from ltx_core.text_encoders.gemma import (
EMBEDDINGS_PROCESSOR_KEY_OPS,
GEMMA_LLM_KEY_OPS,
GEMMA_MODEL_OPS,
EmbeddingsProcessorConfigurator,
GemmaTextEncoderConfigurator,
module_ops_from_gemma_root,
)
from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessorOutput
from ltx_core.tools import AudioLatentTools, LatentTools, VideoLatentTools
from ltx_core.types import Audio, AudioLatentShape, LatentState, VideoLatentShape, VideoPixelShape
from ltx_core.utils import find_matching_file
from ltx_pipelines.utils.gpu_model import gpu_model
from ltx_pipelines.utils.helpers import (
cleanup_memory,
create_noised_state,
generate_enhanced_prompt,
)
from ltx_pipelines.utils.samplers import euler_denoising_loop
from ltx_pipelines.utils.types import Denoiser, ModalitySpec
logger = logging.getLogger(__name__)
T = TypeVar("T")
_M = TypeVar("_M", bound=torch.nn.Module)
# ---------------------------------------------------------------------------
# Internal helpers
# ---------------------------------------------------------------------------
@contextmanager
def _streaming_model(
model: _M,
layers_attr: str,
target_device: torch.device,
prefetch_count: int,
) -> Iterator[_M]:
"""Wrap *model* with :class:`LayerStreamingWrapper`, yield it, then tear down."""
wrapped = LayerStreamingWrapper(
model,
layers_attr=layers_attr,
target_device=target_device,
prefetch_count=prefetch_count,
)
try:
yield wrapped # type: ignore[misc]
finally:
wrapped.teardown()
wrapped.to("meta")
cleanup_memory()
# Flush the host (pinned) memory cache so that freed pinned pages are
# returned to the OS. Without this, sequential streaming models
# (e.g. text encoder then transformer) exhaust host memory because the
# CachingHostAllocator keeps freed blocks cached indefinitely.
torch.cuda.synchronize(device=target_device)
try:
if hasattr(torch._C, "_host_emptyCache"):
torch._C._host_emptyCache()
except Exception:
logger.warning("Host empty cache cleanup failed; ignoring.", exc_info=True)
def _build_state(
spec: ModalitySpec,
tools: LatentTools,
noiser: Noiser,
dtype: torch.dtype,
device: torch.device,
) -> LatentState:
"""Create a noised latent state from a modality spec and tools."""
state = create_noised_state(
tools=tools,
conditionings=spec.conditionings,
noiser=noiser,
dtype=dtype,
device=device,
noise_scale=spec.noise_scale,
initial_latent=spec.initial_latent,
)
if spec.frozen:
state = replace(state, denoise_mask=torch.zeros_like(state.denoise_mask))
return state
def _cleanup_iter(it: Iterator[torch.Tensor], model: torch.nn.Module) -> Iterator[torch.Tensor]:
"""Wrap an iterator to clean up *model* memory once it is exhausted or abandoned."""
with gpu_model(model):
yield from it
# ---------------------------------------------------------------------------
# DiffusionStage
# ---------------------------------------------------------------------------
class DiffusionStage:
"""Owns transformer lifecycle. Builds on each call, frees on exit.
Replaces the manual ``model_ledger.transformer()`` / ``del transformer``
pattern in every pipeline.
"""
def __init__(
self,
checkpoint_path: str,
dtype: torch.dtype,
device: torch.device,
loras: tuple[LoraPathStrengthAndSDOps, ...] = (),
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
) -> None:
self._dtype = dtype
self._device = device
self._quantization = quantization
self._torch_compile = torch_compile
self._transformer_builder = Builder(
model_path=checkpoint_path,
model_class_configurator=LTXModelConfigurator,
model_sd_ops=LTXV_MODEL_COMFY_RENAMING_MAP,
loras=tuple(loras),
registry=registry or DummyRegistry(),
)
def _build_transformer(self, *, device: torch.device | None = None, **kwargs: object) -> X0Model:
target = device or self._device
sd_ops = self._transformer_builder.model_sd_ops
module_ops = self._transformer_builder.module_ops
loras = self._transformer_builder.loras
if self._torch_compile:
module_ops = (*module_ops, COMPILE_TRANSFORMER)
number_of_layers = self._transformer_builder.model_config()["transformer"]["num_layers"]
sd_ops = modify_sd_ops_for_compilation(sd_ops, number_of_layers)
loras = tuple(
LoraPathStrengthAndSDOps(
lora.path,
lora.strength,
modify_sd_ops_for_compilation(
lora.sd_ops if lora.sd_ops is not None else SDOps(name="identity"), number_of_layers
),
)
for lora in loras
)
if self._quantization is not None:
module_ops = (*module_ops, *self._quantization.module_ops)
sd_ops = SDOps(
name=f"sd_ops_chain_{sd_ops.name}+{self._quantization.sd_ops.name}",
mapping=(*sd_ops.mapping, *self._quantization.sd_ops.mapping),
)
builder = self._transformer_builder.with_module_ops(module_ops).with_sd_ops(sd_ops).with_loras(loras)
return X0Model(builder.build(device=target, **kwargs)).to(target).eval()
def _transformer_ctx(
self,
streaming_prefetch_count: int | None,
**kwargs: object,
) -> AbstractContextManager:
if streaming_prefetch_count is not None:
return _streaming_model(
self._build_transformer(device=torch.device("cpu"), **kwargs),
layers_attr="velocity_model.transformer_blocks",
target_device=self._device,
prefetch_count=streaming_prefetch_count,
)
return gpu_model(self._build_transformer(**kwargs))
def __call__( # noqa: PLR0913
self,
denoiser: Denoiser,
sigmas: torch.Tensor,
noiser: Noiser,
width: int,
height: int,
frames: int,
fps: float,
video: ModalitySpec | None = None,
audio: ModalitySpec | None = None,
stepper: DiffusionStepProtocol | None = None,
loop: Callable[..., tuple[LatentState | None, LatentState | None]] | None = None,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
) -> tuple[LatentState | None, LatentState | None]:
"""Build transformer → run denoising loop → free transformer.
Args:
width: Output width in pixels.
height: Output height in pixels.
frames: Number of output frames.
fps: Frame rate.
loop: Denoising loop function. Must accept
``(sigmas, video_state, audio_state, stepper, transformer, denoiser)``
as the first six positional arguments. When ``None``, resolves to
:func:`euler_denoising_loop` at call time.
streaming_prefetch_count: When set, build the transformer on CPU and
wrap with :class:`LayerStreamingWrapper` for memory-efficient
inference, prefetching this many layers ahead.
max_batch_size: Maximum batch size per transformer forward pass.
Guided denoisers make up to 4 transformer calls per step.
When set to a value > 1, the transformer batches multiple
calls together, reducing layer-streaming PCIe transfers.
Default ``1`` preserves sequential behavior.
Returns ``(video_state | None, audio_state | None)`` with cleared
conditionings and unpatchified latents for present modalities.
"""
if video is None and audio is None:
raise ValueError("At least one of `video` or `audio` must be provided")
if loop is None:
loop = euler_denoising_loop
if stepper is None:
stepper = EulerDiffusionStep()
pixel_shape = VideoPixelShape(batch=1, frames=frames, height=height, width=width, fps=fps)
video_state: LatentState | None = None
video_tools: LatentTools | None = None
if video is not None:
v_shape = VideoLatentShape.from_pixel_shape(pixel_shape)
video_tools = VideoLatentTools(VideoLatentPatchifier(patch_size=1), v_shape, fps)
video_state = _build_state(video, video_tools, noiser, self._dtype, self._device)
audio_state: LatentState | None = None
audio_tools: LatentTools | None = None
if audio is not None:
a_shape = AudioLatentShape.from_video_pixel_shape(pixel_shape)
audio_tools = AudioLatentTools(AudioPatchifier(patch_size=1), a_shape)
audio_state = _build_state(audio, audio_tools, noiser, self._dtype, self._device)
with self._transformer_ctx(streaming_prefetch_count, video_tools=video_tools) as base_transformer:
transformer = BatchSplitAdapter(base_transformer, max_batch_size=max_batch_size)
video_state, audio_state = loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
transformer=transformer,
denoiser=denoiser,
)
# Post-process: clear conditionings and unpatchify
if video_state is not None and video_tools is not None:
video_state = video_tools.clear_conditioning(video_state)
video_state = video_tools.unpatchify(video_state)
if audio_state is not None and audio_tools is not None:
audio_state = audio_tools.clear_conditioning(audio_state)
audio_state = audio_tools.unpatchify(audio_state)
return video_state, audio_state
# ---------------------------------------------------------------------------
# PromptEncoder
# ---------------------------------------------------------------------------
class PromptEncoder:
"""Owns text encoder + embeddings processor lifecycle.
Loads Gemma, encodes prompts, frees Gemma, then loads the embeddings
processor to produce final outputs.
"""
def __init__(
self,
checkpoint_path: str,
gemma_root: str,
dtype: torch.dtype,
device: torch.device,
registry: Registry | None = None,
) -> None:
self._dtype = dtype
self._device = device
module_ops = module_ops_from_gemma_root(gemma_root)
model_folder = find_matching_file(gemma_root, "model*.safetensors").parent
weight_paths = [str(p) for p in model_folder.rglob("*.safetensors")]
self._text_encoder_builder = Builder(
model_path=tuple(weight_paths),
model_class_configurator=GemmaTextEncoderConfigurator,
model_sd_ops=GEMMA_LLM_KEY_OPS,
module_ops=(GEMMA_MODEL_OPS, *module_ops),
registry=registry or DummyRegistry(),
)
self._embeddings_processor_builder = Builder(
model_path=checkpoint_path,
model_class_configurator=EmbeddingsProcessorConfigurator,
model_sd_ops=EMBEDDINGS_PROCESSOR_KEY_OPS,
registry=registry or DummyRegistry(),
)
def _text_encoder_ctx(
self,
streaming_prefetch_count: int | None,
) -> AbstractContextManager:
if streaming_prefetch_count is not None:
return _streaming_model(
self._text_encoder_builder.build(device=torch.device("cpu"), dtype=self._dtype).eval(),
layers_attr="model.model.language_model.layers",
target_device=self._device,
prefetch_count=streaming_prefetch_count,
)
return gpu_model(self._text_encoder_builder.build(device=self._device, dtype=self._dtype).eval())
def __call__(
self,
prompts: list[str],
*,
enhance_first_prompt: bool = False,
enhance_prompt_image: str | None = None,
enhance_prompt_seed: int = 42,
streaming_prefetch_count: int | None = None,
) -> list[EmbeddingsProcessorOutput]:
"""Encode *prompts* through Gemma → embeddings processor, freeing each model after use."""
with self._text_encoder_ctx(streaming_prefetch_count) as text_encoder:
if enhance_first_prompt:
prompts = list(prompts)
prompts[0] = generate_enhanced_prompt(
text_encoder, prompts[0], enhance_prompt_image, seed=enhance_prompt_seed
)
raw_outputs = [text_encoder.encode(p) for p in prompts]
with gpu_model(
self._embeddings_processor_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
) as embeddings_processor:
return [embeddings_processor.process_hidden_states(hs, mask) for hs, mask in raw_outputs]
# ---------------------------------------------------------------------------
# ImageConditioner
# ---------------------------------------------------------------------------
class ImageConditioner:
"""Owns video encoder lifecycle.
Builds the encoder, passes it to the user-supplied callable, then frees it.
"""
def __init__(
self,
checkpoint_path: str,
dtype: torch.dtype,
device: torch.device,
registry: Registry | None = None,
) -> None:
self._dtype = dtype
self._device = device
self._encoder_builder = Builder(
model_path=checkpoint_path,
model_class_configurator=VideoEncoderConfigurator,
model_sd_ops=VAE_ENCODER_COMFY_KEYS_FILTER,
registry=registry or DummyRegistry(),
)
def _build_encoder(self) -> VideoEncoder:
return self._encoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
def __call__(self, fn: Callable[[VideoEncoder], T]) -> T:
"""Build video encoder → call *fn(encoder)* → free encoder."""
with gpu_model(self._build_encoder()) as encoder:
return fn(encoder)
# ---------------------------------------------------------------------------
# VideoUpsampler
# ---------------------------------------------------------------------------
class VideoUpsampler:
"""Owns video encoder + spatial upsampler lifecycle."""
def __init__(
self,
checkpoint_path: str,
upsampler_path: str,
dtype: torch.dtype,
device: torch.device,
registry: Registry | None = None,
) -> None:
self._dtype = dtype
self._device = device
self._encoder_builder = Builder(
model_path=checkpoint_path,
model_class_configurator=VideoEncoderConfigurator,
model_sd_ops=VAE_ENCODER_COMFY_KEYS_FILTER,
registry=registry or DummyRegistry(),
)
self._upsampler_builder = Builder(
model_path=upsampler_path,
model_class_configurator=LatentUpsamplerConfigurator,
registry=registry or DummyRegistry(),
)
def __call__(self, latent: torch.Tensor) -> torch.Tensor:
"""Upsample *latent* using video encoder + spatial upsampler, then free both."""
with (
gpu_model(
self._encoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
) as encoder,
gpu_model(
self._upsampler_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
) as upsampler,
):
return upsample_video(latent=latent, video_encoder=encoder, upsampler=upsampler)
# ---------------------------------------------------------------------------
# VideoDecoder
# ---------------------------------------------------------------------------
class VideoDecoder:
"""Owns video decoder lifecycle.
Returns an iterator that cleans up the decoder after all chunks are consumed.
"""
def __init__(
self,
checkpoint_path: str,
dtype: torch.dtype,
device: torch.device,
registry: Registry | None = None,
) -> None:
self._dtype = dtype
self._device = device
self._decoder_builder = Builder(
model_path=checkpoint_path,
model_class_configurator=VideoDecoderConfigurator,
model_sd_ops=VAE_DECODER_COMFY_KEYS_FILTER,
registry=registry or DummyRegistry(),
)
def __call__(
self,
latent: torch.Tensor,
tiling_config: TilingConfig | None = None,
generator: torch.Generator | None = None,
) -> Iterator[torch.Tensor]:
"""Decode *latent* to pixel-space video chunks. Decoder freed after exhaustion."""
decoder = self._decoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
return _cleanup_iter(decoder.decode_video(latent, tiling_config, generator), decoder)
# ---------------------------------------------------------------------------
# AudioDecoder
# ---------------------------------------------------------------------------
class AudioDecoder:
"""Owns audio decoder + vocoder lifecycle."""
def __init__(
self,
checkpoint_path: str,
dtype: torch.dtype,
device: torch.device,
registry: Registry | None = None,
) -> None:
self._dtype = dtype
self._device = device
self._decoder_builder = Builder(
model_path=checkpoint_path,
model_class_configurator=AudioDecoderConfigurator,
model_sd_ops=AUDIO_VAE_DECODER_COMFY_KEYS_FILTER,
registry=registry or DummyRegistry(),
)
self._vocoder_builder = Builder(
model_path=checkpoint_path,
model_class_configurator=VocoderConfigurator,
model_sd_ops=VOCODER_COMFY_KEYS_FILTER,
registry=registry or DummyRegistry(),
)
def __call__(self, latent: torch.Tensor) -> Audio:
"""Decode audio *latent* through VAE decoder + vocoder, then free both."""
with (
gpu_model(
self._decoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
) as decoder,
gpu_model(
self._vocoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
) as vocoder,
):
return vae_decode_audio(latent, decoder, vocoder)
# ---------------------------------------------------------------------------
# AudioEncoder
# ---------------------------------------------------------------------------
class AudioConditioner:
"""Owns audio encoder lifecycle.
Builds the encoder, passes it to the user-supplied callable, then frees it.
Mirrors :class:`ImageConditioner` for the audio modality.
"""
def __init__(
self,
checkpoint_path: str,
dtype: torch.dtype,
device: torch.device,
registry: Registry | None = None,
) -> None:
self._dtype = dtype
self._device = device
self._encoder_builder = Builder(
model_path=checkpoint_path,
model_class_configurator=AudioEncoderConfigurator,
model_sd_ops=AUDIO_VAE_ENCODER_COMFY_KEYS_FILTER,
registry=registry or DummyRegistry(),
)
def __call__(self, fn: Callable[[torch.nn.Module], T]) -> T:
"""Build audio encoder → call *fn(encoder)* → free encoder."""
with gpu_model(
self._encoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
) as encoder:
return fn(encoder)
@@ -0,0 +1,305 @@
"""Flat denoiser classes — transformer received at call time, not stored.
Three implementations of the :class:`~ltx_pipelines.utils.types.Denoiser` protocol:
* :class:`SimpleDenoiser` single transformer call, no guidance.
* :class:`GuidedDenoiser` static guiders, handles CFG + STG + isolated modality.
* :class:`FactoryGuidedDenoiser` resolves guiders per-step from sigma.
``GuidedDenoiser`` and ``FactoryGuidedDenoiser`` share the core multi-pass
logic via the module-level :func:`_guided_denoise` function, which batches
all guidance passes into a single transformer call.
"""
import torch
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderFactory, MultiModalGuiderParams
from ltx_core.guidance.perturbations import (
BatchedPerturbationConfig,
Perturbation,
PerturbationConfig,
PerturbationType,
)
from ltx_core.model.transformer import X0Model
from ltx_core.types import LatentState
from ltx_pipelines.utils.helpers import modality_from_latent_state
_POSITIVE_ONLY_GUIDER = MultiModalGuider(
params=MultiModalGuiderParams(cfg_scale=1.0, stg_scale=0.0, modality_scale=1.0),
)
"""Guider that only runs the conditioned pass and returns cond unchanged."""
def _ensure_guider(guider: MultiModalGuider | None) -> MultiModalGuider:
"""Return the guider as-is, or a positive-only guider for absent modalities."""
return guider if guider is not None else _POSITIVE_ONLY_GUIDER
def _repeat_state(state: LatentState, n: int) -> LatentState:
"""Repeat a ``LatentState`` *n* times along the batch dimension.
``(B, ...) (n*B, ...)`` by tiling the whole tensor n times, so the
ordering is ``[item0, item1, ..., item0, item1, ...]`` matching
``torch.cat`` of n per-pass contexts.
"""
def _repeat(t: torch.Tensor) -> torch.Tensor:
repeats = [1] * t.dim()
repeats[0] = n
return t.repeat(repeats)
return LatentState(
latent=_repeat(state.latent),
denoise_mask=_repeat(state.denoise_mask),
positions=_repeat(state.positions),
clean_latent=_repeat(state.clean_latent),
attention_mask=_repeat(state.attention_mask) if state.attention_mask is not None else None,
)
def _guided_denoise( # noqa: PLR0913
transformer: X0Model,
video_state: LatentState | None,
audio_state: LatentState | None,
sigma: torch.Tensor,
video_guider: MultiModalGuider,
audio_guider: MultiModalGuider,
v_context: torch.Tensor | None,
a_context: torch.Tensor | None,
*,
last_denoised_video: torch.Tensor | None,
last_denoised_audio: torch.Tensor | None,
step_index: int,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
"""Core guided denoising — batches all guidance passes into one transformer call.
Collects per-pass contexts first, then builds a single batched Modality
per present modality via :func:`modality_from_latent_state`. When wrapped
with :class:`~ltx_core.batch_split.BatchSplitAdapter`, the transformer may
split this batch into sequential chunks internally.
Guiders must not be ``None``. For absent modalities, callers should pass
:data:`_POSITIVE_ONLY_GUIDER` (via :func:`_ensure_guider`) so that only
the conditioned pass runs and ``calculate()`` returns cond unchanged.
"""
v_skip = video_guider.should_skip_step(step_index)
a_skip = audio_guider.should_skip_step(step_index)
if v_skip and a_skip:
return last_denoised_video, last_denoised_audio
if video_state is not None and v_context is None:
raise ValueError("v_context is required when video_state is provided")
if audio_state is not None and a_context is None:
raise ValueError("a_context is required when audio_state is provided")
# Define passes: (name, video_context, audio_context, perturbation_config).
# Context is None for absent modalities — filtered out during collection.
_pass = tuple[str, torch.Tensor | None, torch.Tensor | None, PerturbationConfig]
passes: list[_pass] = [("cond", v_context, a_context, PerturbationConfig.empty())]
if video_guider.do_unconditional_generation() or audio_guider.do_unconditional_generation():
if video_guider.do_unconditional_generation() and video_guider.negative_context is None:
raise ValueError("Negative context is required for unconditioned denoising")
if audio_guider.do_unconditional_generation() and audio_guider.negative_context is None:
raise ValueError("Negative context is required for unconditioned denoising")
v_neg = video_guider.negative_context if video_guider.negative_context is not None else v_context
a_neg = audio_guider.negative_context if audio_guider.negative_context is not None else a_context
passes.append(("uncond", v_neg, a_neg, PerturbationConfig.empty()))
stg_perturbations: list[Perturbation] = []
if video_guider.do_perturbed_generation():
stg_perturbations.append(
Perturbation(type=PerturbationType.SKIP_VIDEO_SELF_ATTN, blocks=video_guider.params.stg_blocks)
)
if audio_guider.do_perturbed_generation():
stg_perturbations.append(
Perturbation(type=PerturbationType.SKIP_AUDIO_SELF_ATTN, blocks=audio_guider.params.stg_blocks)
)
if stg_perturbations:
passes.append(("ptb", v_context, a_context, PerturbationConfig(stg_perturbations)))
if video_guider.do_isolated_modality_generation() or audio_guider.do_isolated_modality_generation():
passes.append(
(
"mod",
v_context,
a_context,
PerturbationConfig(
[
Perturbation(type=PerturbationType.SKIP_A2V_CROSS_ATTN, blocks=None),
Perturbation(type=PerturbationType.SKIP_V2A_CROSS_ATTN, blocks=None),
]
),
)
)
# Collect contexts, repeat states, and build batched modalities.
pass_names = [name for name, _, _, _ in passes]
ptb_configs = [ptb for _, _, _, ptb in passes]
n = len(passes)
def _batched_sigma(state: LatentState) -> torch.Tensor:
"""Expand scalar sigma to (n * B,) matching the repeated state."""
return sigma.expand(state.latent.shape[0] * n)
batched_video = None
if video_state is not None:
v_context = torch.cat([vc for _, vc, _, _ in passes], dim=0)
batched_video = modality_from_latent_state(
_repeat_state(video_state, n),
v_context,
_batched_sigma(video_state),
enabled=not v_skip,
)
batched_audio = None
if audio_state is not None:
a_context = torch.cat([ac for _, _, ac, _ in passes], dim=0)
batched_audio = modality_from_latent_state(
_repeat_state(audio_state, n),
a_context,
_batched_sigma(audio_state),
enabled=not a_skip,
)
all_v, all_a = transformer(
video=batched_video, audio=batched_audio, perturbations=BatchedPerturbationConfig(ptb_configs)
)
# Split results back and combine via guiders.
splits_v = list(all_v.chunk(n)) if all_v is not None else [0.0] * n
splits_a = list(all_a.chunk(n)) if all_a is not None else [0.0] * n
r = dict(zip(pass_names, zip(splits_v, splits_a, strict=True), strict=True))
cond_v, cond_a = r["cond"]
uncond_v, uncond_a = r.get("uncond", (0.0, 0.0))
ptb_v, ptb_a = r.get("ptb", (0.0, 0.0))
mod_v, mod_a = r.get("mod", (0.0, 0.0))
denoised_video = last_denoised_video if v_skip else video_guider.calculate(cond_v, uncond_v, ptb_v, mod_v)
denoised_audio = last_denoised_audio if a_skip else audio_guider.calculate(cond_a, uncond_a, ptb_a, mod_a)
return denoised_video, denoised_audio
class SimpleDenoiser:
"""Single transformer call, no guidance.
Passes ``None`` Modality for absent modalities.
"""
def __init__(
self,
v_context: torch.Tensor | None,
a_context: torch.Tensor | None,
) -> None:
self.v_context = v_context
self.a_context = a_context
def __call__(
self,
transformer: X0Model,
video_state: LatentState | None,
audio_state: LatentState | None,
sigmas: torch.Tensor,
step_index: int,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
sigma = sigmas[step_index]
pos_video = modality_from_latent_state(video_state, self.v_context, sigma) if video_state is not None else None
pos_audio = modality_from_latent_state(audio_state, self.a_context, sigma) if audio_state is not None else None
return transformer(video=pos_video, audio=pos_audio, perturbations=None)
class GuidedDenoiser:
"""Static guiders — handles CFG + STG + isolated modality.
Context/guider can be ``None`` for absent modalities (a positive-only
guider is substituted at call time).
"""
def __init__(
self,
v_context: torch.Tensor | None,
a_context: torch.Tensor | None,
video_guider: MultiModalGuider | None = None,
audio_guider: MultiModalGuider | None = None,
) -> None:
self.v_context = v_context
self.a_context = a_context
self.video_guider = video_guider
self.audio_guider = audio_guider
self._last_denoised_video: torch.Tensor | None = None
self._last_denoised_audio: torch.Tensor | None = None
def __call__(
self,
transformer: X0Model,
video_state: LatentState | None,
audio_state: LatentState | None,
sigmas: torch.Tensor,
step_index: int,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
denoised_video, denoised_audio = _guided_denoise(
transformer=transformer,
video_state=video_state,
audio_state=audio_state,
sigma=sigmas[step_index],
video_guider=_ensure_guider(self.video_guider),
audio_guider=_ensure_guider(self.audio_guider),
v_context=self.v_context,
a_context=self.a_context,
last_denoised_video=self._last_denoised_video,
last_denoised_audio=self._last_denoised_audio,
step_index=step_index,
)
self._last_denoised_video = denoised_video
self._last_denoised_audio = denoised_audio
return denoised_video, denoised_audio
class FactoryGuidedDenoiser:
"""Resolves guiders per-step from sigma, then delegates to shared guided logic."""
def __init__(
self,
v_context: torch.Tensor | None,
a_context: torch.Tensor | None,
video_guider_factory: MultiModalGuiderFactory | None = None,
audio_guider_factory: MultiModalGuiderFactory | None = None,
) -> None:
self.v_context = v_context
self.a_context = a_context
self.video_guider_factory = video_guider_factory
self.audio_guider_factory = audio_guider_factory
self._last_denoised_video: torch.Tensor | None = None
self._last_denoised_audio: torch.Tensor | None = None
self._sigma_vals_cached: list[float] | None = None
def __call__(
self,
transformer: X0Model,
video_state: LatentState | None,
audio_state: LatentState | None,
sigmas: torch.Tensor,
step_index: int,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
if self._sigma_vals_cached is None:
self._sigma_vals_cached = sigmas.detach().cpu().tolist()
sigma_val = self._sigma_vals_cached[step_index]
video_guider = _ensure_guider(
self.video_guider_factory.build_from_sigma(sigma_val) if self.video_guider_factory else None
)
audio_guider = _ensure_guider(
(self.audio_guider_factory or self.video_guider_factory).build_from_sigma(sigma_val)
if self.video_guider_factory or self.audio_guider_factory
else None
)
denoised_video, denoised_audio = _guided_denoise(
transformer=transformer,
video_state=video_state,
audio_state=audio_state,
sigma=sigmas[step_index],
video_guider=video_guider,
audio_guider=audio_guider,
v_context=self.v_context,
a_context=self.a_context,
last_denoised_video=self._last_denoised_video,
last_denoised_audio=self._last_denoised_audio,
step_index=step_index,
)
self._last_denoised_video = denoised_video
self._last_denoised_audio = denoised_audio
return denoised_video, denoised_audio
@@ -0,0 +1,30 @@
from collections.abc import Iterator
from contextlib import contextmanager
from typing import TypeVar
import torch
from ltx_pipelines.utils.helpers import cleanup_memory
_M = TypeVar("_M", bound=torch.nn.Module)
@contextmanager
def gpu_model(model: _M) -> Iterator[_M]:
"""Context manager that yields a model and releases its memory on exit.
Moves all parameters and buffers to ``meta`` device on exit, which
immediately releases the underlying storage on **both** GPU and CPU,
then runs ``cleanup_memory()`` to reclaim fragmented CUDA memory.
Usage::
with gpu_model(build_encoder()) as encoder:
... # use encoder — typed as the concrete class
# GPU + CPU memory freed automatically
"""
try:
yield model
finally:
torch.cuda.synchronize()
# .to("meta") releases storage for all parameters/buffers regardless
# of their original device (CUDA or CPU).
model.to("meta")
cleanup_memory()
@@ -1,41 +1,35 @@
import gc
import logging
from dataclasses import replace
import torch
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderFactory
from ltx_core.components.noisers import Noiser
from ltx_core.components.protocols import DiffusionStepProtocol, GuiderProtocol
from ltx_core.conditioning import (
ConditioningItem,
VideoConditionByKeyframeIndex,
VideoConditionByLatentIndex,
)
from ltx_core.guidance.perturbations import (
BatchedPerturbationConfig,
Perturbation,
PerturbationConfig,
PerturbationType,
)
from ltx_core.model.transformer import Modality, X0Model
from ltx_core.model.video_vae import VideoEncoder
from ltx_core.model.audio_vae import encode_audio
from ltx_core.model.transformer import Modality
from ltx_core.model.video_vae import TilingConfig, VideoEncoder
from ltx_core.text_encoders.gemma import GemmaTextEncoder
from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessorOutput
from ltx_core.tools import AudioLatentTools, LatentTools, VideoLatentTools
from ltx_core.tools import LatentTools
from ltx_core.types import AudioLatentShape, LatentState, VideoLatentShape, VideoPixelShape
from ltx_pipelines.utils.args import ImageConditioningInput
from ltx_pipelines.utils.media_io import decode_image, load_image_conditioning, resize_aspect_ratio_preserving
from ltx_pipelines.utils.types import (
DenoisingFunc,
DenoisingLoopFunc,
PipelineComponents,
from ltx_pipelines.utils.media_io import (
decode_audio_from_file,
decode_image,
decode_video_from_file,
get_videostream_fps,
load_image_and_preprocess,
resize_aspect_ratio_preserving,
video_preprocess,
)
def get_device() -> torch.device:
if torch.cuda.is_available():
return torch.device("cuda")
return torch.device("cuda", torch.cuda.current_device())
return torch.device("cpu")
@@ -45,45 +39,89 @@ def cleanup_memory() -> None:
torch.cuda.synchronize()
def encode_prompts(
prompts: list[str],
model_ledger: object,
*,
enhance_prompt_image: str | None = None,
enhance_prompt_seed: int = 42,
enhance_first_prompt: bool = False,
) -> list[EmbeddingsProcessorOutput]:
"""Encode prompts through Gemma → embeddings processor, freeing each after use.
Loads the text encoder from *model_ledger*, optionally enhances the first
prompt, encodes all *prompts*, frees the text encoder, then loads the
embeddings processor to produce the final outputs. Because the text encoder
is loaded and freed entirely within this function, there are no lingering
references that could prevent GPU memory reclamation.
Args:
prompts: Text prompts to encode.
model_ledger: ModelLedger instance (used to load text encoder and embeddings processor).
enhance_prompt_image: Optional image path for prompt enhancement.
enhance_prompt_seed: Seed for prompt enhancement (default 42).
enhance_first_prompt: If True, enhance ``prompts[0]`` before encoding.
Returns:
List of EmbeddingsProcessorOutput, one per prompt.
"""
text_encoder = model_ledger.text_encoder()
if enhance_first_prompt:
prompts = list(prompts)
prompts[0] = generate_enhanced_prompt(text_encoder, prompts[0], enhance_prompt_image, seed=enhance_prompt_seed)
raw_outputs = [text_encoder.encode(p) for p in prompts]
torch.cuda.synchronize()
del text_encoder
cleanup_memory()
def _conform_latent_length(latent: torch.Tensor, expected_frames_count: int) -> torch.Tensor:
actual_frames = latent.shape[2]
if actual_frames > expected_frames_count:
latent = latent[:, :, :expected_frames_count]
elif actual_frames < expected_frames_count:
shape_as_list = list(latent.shape)
shape_as_list[2] = expected_frames_count - actual_frames
pad = torch.zeros(
shape_as_list,
device=latent.device,
dtype=latent.dtype,
)
latent = torch.cat([latent, pad], dim=2)
return latent
embeddings_processor = model_ledger.gemma_embeddings_processor()
results: list[EmbeddingsProcessorOutput] = [
embeddings_processor.process_hidden_states(hs, mask) for hs, mask in raw_outputs
]
del embeddings_processor
cleanup_memory()
return results
def video_latent_from_file(
video_encoder: VideoEncoder,
file_path: str,
output_shape: VideoPixelShape,
device: torch.device,
dtype: torch.dtype,
start_time: float = 0.0,
max_duration: float | None = None,
tiling_config: TilingConfig | None = None,
) -> torch.Tensor | None:
"""Load video from a file, and construct the video latent conforming to video output shape.
Args:
video_encoder: Model used to encode pixel frames to latent space.
file_path: Path to the video file.
output_shape: Target pixel shape (height, width, frames, fps) for the conditioning.
device: Device to run the encoder and hold tensors on.
dtype: Dtype for the output latents.
start_time: Start time in seconds to begin reading the video (default 0.0).
max_duration: Maximum duration in seconds. If None, uses output_shape.frames at
output_shape.fps (default None).
tiling_config: Tiling configuration for the encoder. Defaults to TilingConfig.default().
Returns:
Encoded video latents of shape (1, C, T, H, W) with T = required_latent_frames, or
None (currently this function always returns a tensor).
"""
fps = get_videostream_fps(file_path)
if fps != output_shape.fps:
raise ValueError(f"Input video FPS {fps} does not match output FPS {output_shape.fps}, not supported")
max_duration = max_duration or output_shape.frames / fps
frame_gen = decode_video_from_file(path=file_path, device=device, start_time=start_time, max_duration=max_duration)
frames = video_preprocess(frame_gen, output_shape.height, output_shape.width, dtype, device)
latents = video_encoder.tiled_encode(frames, tiling_config or TilingConfig.default())
required_latent_frames = VideoLatentShape.from_pixel_shape(output_shape).frames
return _conform_latent_length(latents, required_latent_frames)
def audio_latent_from_file(
audio_encoder: torch.nn.Module,
file_path: str,
output_shape: VideoPixelShape,
device: torch.device,
dtype: torch.dtype,
start_time: float = 0.0,
max_duration: float | None = None,
) -> torch.Tensor | None:
"""Load audio from a file, and construct the audio latent conforming to video output shape.
Args:
audio_encoder: Model used to encode audio to latent space.
file_path: Path to the audio or video file containing an audio stream.
output_shape: Target video pixel shape; used to derive required latent frames
and, when max_duration is None, the audio duration (output_shape.frames / fps).
device: Device to run the encoder and hold tensors on.
dtype: Dtype for the output latents.
start_time: Start time in seconds to begin reading the audio (default 0.0).
max_duration: Maximum duration in seconds. If None, uses the full span implied
by output_shape (default None).
Returns:
Encoded audio latents of shape (1, C, T, ...) with T = required_latent_frames, or
None if the file has no audio stream.
"""
max_duration = max_duration or output_shape.frames / output_shape.fps
audio_in = decode_audio_from_file(file_path, device, start_time, max_duration)
if audio_in is None:
return None
latents = encode_audio(audio_in, audio_encoder, None).to(device, dtype)
required_latent_frames = AudioLatentShape.from_video_pixel_shape(output_shape).frames
return _conform_latent_length(latents, required_latent_frames)
def combined_image_conditionings(
@@ -98,7 +136,7 @@ def combined_image_conditionings(
and using other encoded images as the keyframe conditionings."""
conditionings = []
for img in images:
image = load_image_conditioning(
image = load_image_and_preprocess(
image_path=img.path,
height=height,
width=width,
@@ -133,7 +171,7 @@ def image_conditionings_by_replacing_latent(
) -> list[ConditioningItem]:
conditionings = []
for img in images:
image = load_image_conditioning(
image = load_image_and_preprocess(
image_path=img.path,
height=height,
width=width,
@@ -163,7 +201,7 @@ def image_conditionings_by_adding_guiding_latent(
) -> list[ConditioningItem]:
conditionings = []
for img in images:
image = load_image_conditioning(
image = load_image_and_preprocess(
image_path=img.path,
height=height,
width=width,
@@ -178,72 +216,6 @@ def image_conditionings_by_adding_guiding_latent(
return conditionings
def noise_video_state(
output_shape: VideoPixelShape,
noiser: Noiser,
conditionings: list[ConditioningItem],
components: PipelineComponents,
dtype: torch.dtype,
device: torch.device,
noise_scale: float = 1.0,
initial_latent: torch.Tensor | None = None,
) -> tuple[LatentState, VideoLatentTools]:
"""Initialize and noise a video latent state for the diffusion pipeline.
Creates a video latent state from the output shape, applies conditionings,
and adds noise using the provided noiser. Returns the noised state and
video latent tools for further processing. If initial_latent is provided, it will be used to create the initial
state, otherwise an empty initial state will be created.
"""
video_latent_shape = VideoLatentShape.from_pixel_shape(
shape=output_shape,
latent_channels=components.video_latent_channels,
scale_factors=components.video_scale_factors,
)
video_tools = VideoLatentTools(components.video_patchifier, video_latent_shape, output_shape.fps)
video_state = create_noised_state(
tools=video_tools,
conditionings=conditionings,
noiser=noiser,
dtype=dtype,
device=device,
noise_scale=noise_scale,
initial_latent=initial_latent,
)
return video_state, video_tools
def noise_audio_state(
output_shape: VideoPixelShape,
noiser: Noiser,
conditionings: list[ConditioningItem],
components: PipelineComponents,
dtype: torch.dtype,
device: torch.device,
noise_scale: float = 1.0,
initial_latent: torch.Tensor | None = None,
) -> tuple[LatentState, AudioLatentTools]:
"""Initialize and noise an audio latent state for the diffusion pipeline.
Creates an audio latent state from the output shape, applies conditionings,
and adds noise using the provided noiser. Returns the noised state and
audio latent tools for further processing. If initial_latent is provided, it will be used to create the initial
state, otherwise an empty initial state will be created.
"""
audio_latent_shape = AudioLatentShape.from_video_pixel_shape(output_shape)
audio_tools = AudioLatentTools(components.audio_patchifier, audio_latent_shape)
audio_state = create_noised_state(
tools=audio_tools,
conditionings=conditionings,
noiser=noiser,
dtype=dtype,
device=device,
noise_scale=noise_scale,
initial_latent=initial_latent,
)
return audio_state, audio_tools
def create_noised_state(
tools: LatentTools,
conditionings: list[ConditioningItem],
@@ -308,301 +280,14 @@ def timesteps_from_mask(denoise_mask: torch.Tensor, sigma: float | torch.Tensor)
"""Compute timesteps from a denoise mask and sigma value.
Multiplies the denoise mask by sigma to produce timesteps for each position
in the latent state. Areas where the mask is 0 will have zero timesteps.
When sigma is ``(B,)`` it is reshaped to ``(B, 1, ...)`` so the batch
dimension aligns correctly with ``denoise_mask``.
"""
if isinstance(sigma, torch.Tensor) and sigma.dim() == 1:
sigma = sigma.view(-1, *([1] * (denoise_mask.dim() - 1)))
return denoise_mask * sigma
def simple_denoising_func(
video_context: torch.Tensor, audio_context: torch.Tensor, transformer: X0Model
) -> DenoisingFunc:
def simple_denoising_step(
video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
) -> tuple[torch.Tensor, torch.Tensor]:
sigma = sigmas[step_index]
pos_video = modality_from_latent_state(video_state, video_context, sigma)
pos_audio = modality_from_latent_state(audio_state, audio_context, sigma)
denoised_video, denoised_audio = transformer(video=pos_video, audio=pos_audio, perturbations=None)
return denoised_video, denoised_audio
return simple_denoising_step
def guider_denoising_func(
guider: GuiderProtocol,
v_context_p: torch.Tensor,
v_context_n: torch.Tensor,
a_context_p: torch.Tensor,
a_context_n: torch.Tensor,
transformer: X0Model,
) -> DenoisingFunc:
def guider_denoising_step(
video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
) -> tuple[torch.Tensor, torch.Tensor]:
sigma = sigmas[step_index]
pos_video = modality_from_latent_state(video_state, v_context_p, sigma)
pos_audio = modality_from_latent_state(audio_state, a_context_p, sigma)
denoised_video, denoised_audio = transformer(video=pos_video, audio=pos_audio, perturbations=None)
if guider.enabled():
neg_video = modality_from_latent_state(video_state, v_context_n, sigma)
neg_audio = modality_from_latent_state(audio_state, a_context_n, sigma)
neg_denoised_video, neg_denoised_audio = transformer(video=neg_video, audio=neg_audio, perturbations=None)
denoised_video = denoised_video + guider.delta(denoised_video, neg_denoised_video)
denoised_audio = denoised_audio + guider.delta(denoised_audio, neg_denoised_audio)
return denoised_video, denoised_audio
return guider_denoising_step
def multi_modal_guider_denoising_func(
video_guider: MultiModalGuider,
audio_guider: MultiModalGuider,
v_context: torch.Tensor,
a_context: torch.Tensor,
transformer: X0Model,
*,
last_denoised_video: torch.Tensor | None = None,
last_denoised_audio: torch.Tensor | None = None,
) -> DenoisingFunc:
def guider_denoising_step(
video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
) -> tuple[torch.Tensor, torch.Tensor]:
nonlocal last_denoised_video, last_denoised_audio
if video_guider.should_skip_step(step_index) and audio_guider.should_skip_step(step_index):
return last_denoised_video, last_denoised_audio
sigma = sigmas[step_index]
pos_video_modality = modality_from_latent_state(
video_state, v_context, sigma, enabled=not video_guider.should_skip_step(step_index)
)
pos_audio_modality = modality_from_latent_state(
audio_state, a_context, sigma, enabled=not audio_guider.should_skip_step(step_index)
)
denoised_video, denoised_audio = transformer(
video=pos_video_modality, audio=pos_audio_modality, perturbations=None
)
neg_denoised_video, neg_denoised_audio = 0.0, 0.0
if video_guider.do_unconditional_generation() or audio_guider.do_unconditional_generation():
if video_guider.do_unconditional_generation() and video_guider.negative_context is None:
raise ValueError("Negative context is required for unconditioned denoising")
if audio_guider.do_unconditional_generation() and audio_guider.negative_context is None:
raise ValueError("Negative context is required for unconditioned denoising")
neg_video_modality = modality_from_latent_state(
video_state,
video_guider.negative_context
if video_guider.negative_context is not None
else pos_video_modality.context,
sigma,
)
neg_audio_modality = modality_from_latent_state(
audio_state,
audio_guider.negative_context
if audio_guider.negative_context is not None
else pos_audio_modality.context,
sigma,
)
neg_denoised_video, neg_denoised_audio = transformer(
video=neg_video_modality, audio=neg_audio_modality, perturbations=None
)
ptb_denoised_video, ptb_denoised_audio = 0.0, 0.0
if video_guider.do_perturbed_generation() or audio_guider.do_perturbed_generation():
perturbations = []
if video_guider.do_perturbed_generation():
perturbations.append(
Perturbation(type=PerturbationType.SKIP_VIDEO_SELF_ATTN, blocks=video_guider.params.stg_blocks)
)
if audio_guider.do_perturbed_generation():
perturbations.append(
Perturbation(type=PerturbationType.SKIP_AUDIO_SELF_ATTN, blocks=audio_guider.params.stg_blocks)
)
perturbation_config = PerturbationConfig(perturbations=perturbations)
ptb_denoised_video, ptb_denoised_audio = transformer(
video=pos_video_modality,
audio=pos_audio_modality,
perturbations=BatchedPerturbationConfig(perturbations=[perturbation_config]),
)
mod_denoised_video, mod_denoised_audio = 0.0, 0.0
if video_guider.do_isolated_modality_generation() or audio_guider.do_isolated_modality_generation():
perturbations = [
Perturbation(type=PerturbationType.SKIP_A2V_CROSS_ATTN, blocks=None),
Perturbation(type=PerturbationType.SKIP_V2A_CROSS_ATTN, blocks=None),
]
perturbation_config = PerturbationConfig(perturbations=perturbations)
mod_denoised_video, mod_denoised_audio = transformer(
video=pos_video_modality,
audio=pos_audio_modality,
perturbations=BatchedPerturbationConfig(perturbations=[perturbation_config]),
)
if video_guider.should_skip_step(step_index):
denoised_video = last_denoised_video
else:
denoised_video = video_guider.calculate(
denoised_video, neg_denoised_video, ptb_denoised_video, mod_denoised_video
)
if audio_guider.should_skip_step(step_index):
denoised_audio = last_denoised_audio
else:
denoised_audio = audio_guider.calculate(
denoised_audio, neg_denoised_audio, ptb_denoised_audio, mod_denoised_audio
)
last_denoised_video = denoised_video
last_denoised_audio = denoised_audio
return denoised_video, denoised_audio
return guider_denoising_step
def multi_modal_guider_factory_denoising_func(
video_guider_factory: MultiModalGuiderFactory,
audio_guider_factory: MultiModalGuiderFactory | None,
v_context: torch.Tensor,
a_context: torch.Tensor,
transformer: X0Model,
) -> DenoisingFunc:
"""Resolve guiders per step via factory.build_from_sigma, then multi_modal_guider_denoising_func."""
last_denoised_video: torch.Tensor | None = None
last_denoised_audio: torch.Tensor | None = None
sigma_vals_cached: list[float] | None = None
def guider_denoising_step(
video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
) -> tuple[torch.Tensor, torch.Tensor]:
nonlocal last_denoised_video, last_denoised_audio, sigma_vals_cached
if sigma_vals_cached is None:
sigma_vals_cached = sigmas.detach().cpu().tolist()
sigma_val = sigma_vals_cached[step_index]
video_guider = video_guider_factory.build_from_sigma(sigma_val)
audio_guider = (audio_guider_factory or video_guider_factory).build_from_sigma(sigma_val)
denoise_fn = multi_modal_guider_denoising_func(
video_guider,
audio_guider,
v_context,
a_context,
transformer,
last_denoised_video=last_denoised_video,
last_denoised_audio=last_denoised_audio,
)
denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_index)
last_denoised_video, last_denoised_audio = denoised_video, denoised_audio
return denoised_video, denoised_audio
return guider_denoising_step
def denoise_audio_video( # noqa: PLR0913
output_shape: VideoPixelShape,
conditionings: list[ConditioningItem],
noiser: Noiser,
sigmas: torch.Tensor,
stepper: DiffusionStepProtocol,
denoising_loop_fn: DenoisingLoopFunc,
components: PipelineComponents,
dtype: torch.dtype,
device: torch.device,
noise_scale: float = 1.0,
initial_video_latent: torch.Tensor | None = None,
initial_audio_latent: torch.Tensor | None = None,
) -> tuple[LatentState, LatentState]:
video_state, video_tools = noise_video_state(
output_shape=output_shape,
noiser=noiser,
conditionings=conditionings,
components=components,
dtype=dtype,
device=device,
noise_scale=noise_scale,
initial_latent=initial_video_latent,
)
audio_state, audio_tools = noise_audio_state(
output_shape=output_shape,
noiser=noiser,
conditionings=[],
components=components,
dtype=dtype,
device=device,
noise_scale=noise_scale,
initial_latent=initial_audio_latent,
)
video_state, audio_state = denoising_loop_fn(
sigmas,
video_state,
audio_state,
stepper,
)
video_state = video_tools.clear_conditioning(video_state)
video_state = video_tools.unpatchify(video_state)
audio_state = audio_tools.clear_conditioning(audio_state)
audio_state = audio_tools.unpatchify(audio_state)
return video_state, audio_state
def denoise_video_only( # noqa: PLR0913
output_shape: VideoPixelShape,
conditionings: list[ConditioningItem],
noiser: Noiser,
sigmas: torch.Tensor,
stepper: DiffusionStepProtocol,
denoising_loop_fn: DenoisingLoopFunc,
components: PipelineComponents,
dtype: torch.dtype,
device: torch.device,
noise_scale: float = 1.0,
initial_video_latent: torch.Tensor | None = None,
initial_audio_latent: torch.Tensor | None = None,
) -> LatentState:
video_state, video_tools = noise_video_state(
output_shape=output_shape,
noiser=noiser,
conditionings=conditionings,
components=components,
dtype=dtype,
device=device,
noise_scale=noise_scale,
initial_latent=initial_video_latent,
)
audio_state, _ = noise_audio_state(
output_shape=output_shape,
noiser=noiser,
conditionings=[],
components=components,
dtype=dtype,
device=device,
noise_scale=0.0,
initial_latent=initial_audio_latent,
)
audio_state = replace(audio_state, denoise_mask=torch.zeros_like(audio_state.denoise_mask))
video_state, audio_state = denoising_loop_fn(
sigmas,
video_state,
audio_state,
stepper,
)
video_state = video_tools.clear_conditioning(video_state)
video_state = video_tools.unpatchify(video_state)
return video_state
_UNICODE_REPLACEMENTS = str.maketrans("\u2018\u2019\u201c\u201d\u2014\u2013\u00a0\u2032\u2212", "''\"\"-- '-")
@@ -12,7 +12,7 @@ from PIL import Image
from torch._prims_common import DeviceLikeType
from tqdm import tqdm
from ltx_core.types import Audio
from ltx_core.types import Audio, VideoPixelShape
from ltx_pipelines.utils.constants import DEFAULT_IMAGE_CRF
logger = logging.getLogger(__name__)
@@ -79,7 +79,7 @@ def normalize_latent(latent: torch.Tensor, device: torch.device, dtype: torch.dt
return (latent / 127.5 - 1.0).to(device=device, dtype=dtype)
def load_image_conditioning(
def load_image_and_preprocess(
image_path: str,
height: int,
width: int,
@@ -99,14 +99,23 @@ def load_image_conditioning(
return image
def load_video_conditioning(
video_path: str, height: int, width: int, frame_cap: int, dtype: torch.dtype, device: torch.device
def video_preprocess(
frames: Generator[torch.Tensor],
height: int,
width: int,
dtype: torch.dtype,
device: torch.device,
) -> torch.Tensor:
"""Preprocesses a video frame generator for conditioning.
Args:
frames: Generator of video frames as tensors of shape (1, H, W, C), dtype uint8.
height: Target height in pixels.
width: Target width in pixels.
dtype: Target dtype for the output tensor.
device: Target device for the output tensor.
Returns:
Tensor of shape (1, C, F, height, width) with values in [-1, 1].
"""
Loads a video from a path and preprocesses it for conditioning.
Note: The video is resized to the nearest multiple of 2 for compatibility with video codecs.
"""
frames = decode_video_from_file(path=video_path, frame_cap=frame_cap, device=device)
result = None
for f in frames:
frame = resize_and_center_crop(f.to(torch.float32), height, width)
@@ -257,9 +266,23 @@ def _audio_frame_to_float(frame: av.AudioFrame) -> np.ndarray:
return arr
def get_videostream_metadata(path: str) -> tuple[float, int, int, int]:
"""Read video stream metadata: (fps, num_frames, width, height).
def get_videostream_fps(path: str) -> float:
"""Read video stream FPS."""
container = av.open(path)
try:
video_stream = next(s for s in container.streams if s.type == "video")
return float(video_stream.average_rate)
finally:
container.close()
def get_videostream_metadata(path: str) -> VideoPixelShape:
"""Read video stream metadata as a VideoPixelShape with batch=1.
If frame count is missing in the container, decodes the stream to count frames.
Args:
path: Path to the video file.
Returns:
VideoPixelShape with batch=1, frames, height, width, and fps populated from the stream.
"""
container = av.open(path)
try:
@@ -270,7 +293,7 @@ def get_videostream_metadata(path: str) -> tuple[float, int, int, int]:
num_frames = sum(1 for _ in container.decode(video_stream))
width = video_stream.codec_context.width
height = video_stream.codec_context.height
return fps, num_frames, width, height
return VideoPixelShape(batch=1, frames=num_frames, height=height, width=width, fps=fps)
finally:
container.close()
@@ -338,16 +361,85 @@ def decode_audio_from_file(
return Audio(waveform=waveform, sampling_rate=sample_rate)
def decode_video_from_file(path: str, frame_cap: int, device: DeviceLikeType) -> Generator[torch.Tensor]:
def decode_video_by_frame(
path: str,
device: DeviceLikeType,
starting_frame: int = 0,
frame_cap: int | None = None,
) -> Generator[torch.Tensor]:
"""Decodes video from a file by sequential frame index, without relying on pts.
Args:
path: Path to the video file.
device: Device to place the resulting tensors on.
starting_frame: Number of leading frames to skip (default 0).
frame_cap: Maximum number of frames to yield. If None, no frame limit (default None).
Yields:
Frames as tensors of shape (1, H, W, C), dtype uint8.
"""
container = av.open(path)
try:
video_stream = next(s for s in container.streams if s.type == "video")
for frame in container.decode(video_stream):
for index, frame in enumerate(container.decode(video_stream)):
if index < starting_frame:
continue
tensor = torch.tensor(frame.to_rgb().to_ndarray(), dtype=torch.uint8, device=device).unsqueeze(0)
yield tensor
frame_cap = frame_cap - 1
if frame_cap == 0:
if frame_cap is not None:
frame_cap -= 1
if frame_cap == 0:
break
finally:
container.close()
def decode_video_from_file(
path: str,
device: DeviceLikeType,
start_time: float = 0.0,
max_duration: float | None = None,
) -> Generator[torch.Tensor]:
"""Decodes video from a file using presentation timestamps for time-based trimming.
If a frame with no pts is encountered, falls back to :func:`decode_video_by_frame`
using FPS-derived frame indices.
Args:
path: Path to the video file.
device: Device to place the resulting tensors on.
start_time: Start time in seconds (default 0.0).
max_duration: Maximum duration in seconds to decode. If None, reads to end of
stream (default None).
Yields:
Frames as tensors of shape (1, H, W, C), dtype uint8.
"""
container = av.open(path)
try:
video_stream = next(s for s in container.streams if s.type == "video")
time_base = float(video_stream.time_base)
if start_time > 0:
container.seek(int(start_time / time_base), stream=video_stream)
end_time = start_time + max_duration if max_duration is not None else None
for frame in container.decode(video_stream):
# PyAV may leave pts unset when the demuxer does not expose per-frame
# timestamps (e.g. some raw/elementary streams, stripped or missing
# metadata, or certain remux paths). Without pts we cannot map frames to
# wall-clock time, so we fall back to sequential frame indices using the
# stream's average frame rate.
if frame.pts is None:
fps = float(video_stream.average_rate)
starting_frame = round(start_time * fps)
frame_cap = round(max_duration * fps) if max_duration is not None else None
yield from decode_video_by_frame(
path=path, device=device, starting_frame=starting_frame, frame_cap=frame_cap
)
return
frame_time = frame.pts * time_base
if frame_time < start_time:
continue
if end_time is not None and frame_time >= end_time:
break
yield torch.tensor(frame.to_rgb().to_ndarray(), dtype=torch.uint8, device=device).unsqueeze(0)
finally:
container.close()
@@ -1,304 +0,0 @@
from dataclasses import replace
import torch
from ltx_core.loader import SDOps
from ltx_core.loader.primitives import LoraPathStrengthAndSDOps
from ltx_core.loader.registry import DummyRegistry, Registry
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder
from ltx_core.model.audio_vae import (
AUDIO_VAE_DECODER_COMFY_KEYS_FILTER,
AUDIO_VAE_ENCODER_COMFY_KEYS_FILTER,
VOCODER_COMFY_KEYS_FILTER,
AudioDecoder,
AudioDecoderConfigurator,
AudioEncoder,
AudioEncoderConfigurator,
Vocoder,
VocoderConfigurator,
)
from ltx_core.model.transformer import (
LTXV_MODEL_COMFY_RENAMING_MAP,
LTXModelConfigurator,
X0Model,
)
from ltx_core.model.upsampler import LatentUpsampler, LatentUpsamplerConfigurator
from ltx_core.model.video_vae import (
VAE_DECODER_COMFY_KEYS_FILTER,
VAE_ENCODER_COMFY_KEYS_FILTER,
VideoDecoder,
VideoDecoderConfigurator,
VideoEncoder,
VideoEncoderConfigurator,
)
from ltx_core.quantization import QuantizationPolicy
from ltx_core.text_encoders.gemma import (
EMBEDDINGS_PROCESSOR_KEY_OPS,
GEMMA_LLM_KEY_OPS,
GEMMA_MODEL_OPS,
EmbeddingsProcessor,
EmbeddingsProcessorConfigurator,
GemmaTextEncoder,
GemmaTextEncoderConfigurator,
module_ops_from_gemma_root,
)
from ltx_core.utils import find_matching_file
class ModelLedger:
"""
Central coordinator for loading and building models used in an LTX pipeline.
The ledger wires together multiple model builders (transformer, video VAE encoder/decoder,
audio VAE decoder, vocoder, text encoder, and optional latent upsampler) and exposes
factory methods for constructing model instances.
### Model Building
Each model method (e.g. :meth:`transformer`, :meth:`video_decoder`, :meth:`text_encoder`)
constructs a new model instance on each call. The builder uses the
:class:`~ltx_core.loader.registry.Registry` to load weights from the checkpoint,
instantiates the model with the configured ``dtype``, and moves it to ``self.device``.
.. note::
Models are **not cached**. Each call to a model method creates a new instance.
Callers are responsible for storing references to models they wish to reuse
and for freeing GPU memory (e.g. by deleting references and calling
``torch.cuda.empty_cache()``).
### Constructor parameters
dtype:
Torch dtype used when constructing all models (e.g. ``torch.bfloat16``).
device:
Target device to which models are moved after construction (e.g. ``torch.device("cuda")``).
checkpoint_path:
Path to a checkpoint directory or file containing the core model weights
(transformer, video VAE, audio VAE, text encoder, vocoder). If ``None``, the
corresponding builders are not created and calling those methods will raise
a :class:`ValueError`.
gemma_root_path:
Base path to Gemma-compatible CLIP/text encoder weights. Required to
initialize the text encoder builder; if omitted, :meth:`text_encoder` cannot be used.
spatial_upsampler_path:
Optional path to a latent upsampler checkpoint. If provided, the
:meth:`spatial_upsampler` method becomes available; otherwise calling it raises
a :class:`ValueError`.
loras:
Tuple of LoRA configurations (path, strength, sd_ops) applied on top of the base
transformer weights. Use ``()`` for none.
registry:
Optional :class:`Registry` instance for weight caching across builders.
Defaults to :class:`DummyRegistry` which performs no cross-builder caching.
quantization:
Optional :class:`QuantizationPolicy` controlling how transformer weights
are stored and how matmul is executed. Defaults to None, which means no quantization.
### Creating Variants
Use :meth:`with_additional_loras` to create a new ``ModelLedger`` instance that
includes additional LoRA configurations or :meth:`with_loras` to replace existing
lora configurations while sharing the same registry for weight caching.
"""
def __init__(
self,
dtype: torch.dtype,
device: torch.device,
checkpoint_path: str | None = None,
gemma_root_path: str | None = None,
spatial_upsampler_path: str | None = None,
loras: tuple[LoraPathStrengthAndSDOps, ...] = (),
registry: Registry | None = None,
quantization: QuantizationPolicy | None = None,
):
self.dtype = dtype
self.device = device
self.checkpoint_path = checkpoint_path
self.gemma_root_path = gemma_root_path
self.spatial_upsampler_path = spatial_upsampler_path
self.loras = loras
self.registry = registry or DummyRegistry()
self.quantization = quantization
self.build_model_builders()
def build_model_builders(self) -> None:
if self.checkpoint_path is not None:
self.transformer_builder = Builder(
model_path=self.checkpoint_path,
model_class_configurator=LTXModelConfigurator,
model_sd_ops=LTXV_MODEL_COMFY_RENAMING_MAP,
loras=tuple(self.loras),
registry=self.registry,
)
self.vae_decoder_builder = Builder(
model_path=self.checkpoint_path,
model_class_configurator=VideoDecoderConfigurator,
model_sd_ops=VAE_DECODER_COMFY_KEYS_FILTER,
registry=self.registry,
)
self.vae_encoder_builder = Builder(
model_path=self.checkpoint_path,
model_class_configurator=VideoEncoderConfigurator,
model_sd_ops=VAE_ENCODER_COMFY_KEYS_FILTER,
registry=self.registry,
)
self.audio_encoder_builder = Builder[AudioEncoder](
model_path=self.checkpoint_path,
model_class_configurator=AudioEncoderConfigurator,
model_sd_ops=AUDIO_VAE_ENCODER_COMFY_KEYS_FILTER,
registry=self.registry,
)
self.audio_decoder_builder = Builder(
model_path=self.checkpoint_path,
model_class_configurator=AudioDecoderConfigurator,
model_sd_ops=AUDIO_VAE_DECODER_COMFY_KEYS_FILTER,
registry=self.registry,
)
self.vocoder_builder = Builder(
model_path=self.checkpoint_path,
model_class_configurator=VocoderConfigurator,
model_sd_ops=VOCODER_COMFY_KEYS_FILTER,
registry=self.registry,
)
# Embeddings processor only needs the LTX checkpoint (no Gemma weights)
self.embeddings_processor_builder = Builder(
model_path=self.checkpoint_path,
model_class_configurator=EmbeddingsProcessorConfigurator,
model_sd_ops=EMBEDDINGS_PROCESSOR_KEY_OPS,
registry=self.registry,
)
if self.gemma_root_path is not None:
module_ops = module_ops_from_gemma_root(self.gemma_root_path)
model_folder = find_matching_file(self.gemma_root_path, "model*.safetensors").parent
weight_paths = [str(p) for p in model_folder.rglob("*.safetensors")]
self.text_encoder_builder = Builder(
model_path=tuple(weight_paths),
model_class_configurator=GemmaTextEncoderConfigurator,
model_sd_ops=GEMMA_LLM_KEY_OPS,
registry=self.registry,
module_ops=(GEMMA_MODEL_OPS, *module_ops),
)
if self.spatial_upsampler_path is not None:
self.upsampler_builder = Builder(
model_path=self.spatial_upsampler_path,
model_class_configurator=LatentUpsamplerConfigurator,
registry=self.registry,
)
def _target_device(self) -> torch.device:
if isinstance(self.registry, DummyRegistry) or self.registry is None:
return self.device
else:
return torch.device("cpu")
def with_additional_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> "ModelLedger":
"""Add new lora configurations to the existing ones."""
return self.with_loras((*self.loras, *loras))
def with_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> "ModelLedger":
"""Replace existing lora configurations with new ones."""
return ModelLedger(
dtype=self.dtype,
device=self.device,
checkpoint_path=self.checkpoint_path,
gemma_root_path=self.gemma_root_path,
spatial_upsampler_path=self.spatial_upsampler_path,
loras=loras,
registry=self.registry,
quantization=self.quantization,
)
def transformer(self) -> X0Model:
if not hasattr(self, "transformer_builder"):
raise ValueError(
"Transformer not initialized. Please provide a checkpoint path to the ModelLedger constructor."
)
if self.quantization is None:
return (
X0Model(self.transformer_builder.build(device=self._target_device(), dtype=self.dtype))
.to(self.device)
.eval()
)
else:
sd_ops = self.transformer_builder.model_sd_ops
if self.quantization.sd_ops is not None:
sd_ops = SDOps(
name=f"sd_ops_chain_{sd_ops.name}+{self.quantization.sd_ops.name}",
mapping=(*sd_ops.mapping, *self.quantization.sd_ops.mapping),
)
builder = replace(
self.transformer_builder,
module_ops=(*self.transformer_builder.module_ops, *self.quantization.module_ops),
model_sd_ops=sd_ops,
)
return X0Model(builder.build(device=self._target_device())).to(self.device).eval()
def video_decoder(self) -> VideoDecoder:
if not hasattr(self, "vae_decoder_builder"):
raise ValueError(
"Video decoder not initialized. Please provide a checkpoint path to the ModelLedger constructor."
)
return self.vae_decoder_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval()
def video_encoder(self) -> VideoEncoder:
if not hasattr(self, "vae_encoder_builder"):
raise ValueError(
"Video encoder not initialized. Please provide a checkpoint path to the ModelLedger constructor."
)
return self.vae_encoder_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval()
def text_encoder(self) -> GemmaTextEncoder:
if not hasattr(self, "text_encoder_builder"):
raise ValueError(
"Text encoder not initialized. Please provide a checkpoint path and gemma root path to the "
"ModelLedger constructor."
)
return self.text_encoder_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval()
def gemma_embeddings_processor(self) -> EmbeddingsProcessor:
if not hasattr(self, "embeddings_processor_builder"):
raise ValueError(
"Embeddings processor not initialized. Please provide a checkpoint path to the ModelLedger constructor."
)
return (
self.embeddings_processor_builder.build(device=self._target_device(), dtype=self.dtype)
.to(self.device)
.eval()
)
def audio_encoder(self) -> AudioEncoder:
if not hasattr(self, "audio_encoder_builder"):
raise ValueError(
"Audio encoder not initialized. Please provide a checkpoint path to the ModelLedger constructor."
)
return self.audio_encoder_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval()
def audio_decoder(self) -> AudioDecoder:
if not hasattr(self, "audio_decoder_builder"):
raise ValueError(
"Audio decoder not initialized. Please provide a checkpoint path to the ModelLedger constructor."
)
return self.audio_decoder_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval()
def vocoder(self) -> Vocoder:
if not hasattr(self, "vocoder_builder"):
raise ValueError(
"Vocoder not initialized. Please provide a checkpoint path to the ModelLedger constructor."
)
return self.vocoder_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval()
def spatial_upsampler(self) -> LatentUpsampler:
if not hasattr(self, "upsampler_builder"):
raise ValueError("Upsampler not initialized. Please provide upsampler path to the ModelLedger constructor.")
return self.upsampler_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval()
@@ -8,86 +8,91 @@ from tqdm import tqdm
from ltx_core.components.diffusion_steps import Res2sDiffusionStep
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.model.transformer import X0Model
from ltx_core.utils import to_denoised, to_velocity
from ltx_pipelines.utils.helpers import post_process_latent, timesteps_from_mask
from ltx_pipelines.utils.res2s import get_res2s_coefficients
from ltx_pipelines.utils.types import DenoisingFunc, LatentState
from ltx_pipelines.utils.types import Denoiser, LatentState
logger = logging.getLogger(__name__)
def _step_state(
state: LatentState | None,
denoised: torch.Tensor | None,
stepper: DiffusionStepProtocol,
sigmas: torch.Tensor,
step_idx: int,
) -> LatentState | None:
"""Advance one diffusion step for a single modality, or return ``None`` if absent."""
if state is None or denoised is None:
return state
denoised = post_process_latent(denoised, state.denoise_mask, state.clean_latent)
return replace(state, latent=stepper.step(state.latent, denoised, sigmas, step_idx))
def euler_denoising_loop(
sigmas: torch.Tensor,
video_state: LatentState,
audio_state: LatentState,
video_state: LatentState | None,
audio_state: LatentState | None,
stepper: DiffusionStepProtocol,
denoise_fn: DenoisingFunc,
) -> tuple[LatentState, LatentState]:
transformer: X0Model,
denoiser: Denoiser,
) -> tuple[LatentState | None, LatentState | None]:
"""
Perform the joint audio-video denoising loop over a diffusion schedule.
This function iterates over all but the final value in ``sigmas`` and, at
each diffusion step, calls ``denoise_fn`` to obtain denoised video and
audio latents. The denoised latents are post-processed with their
respective denoise masks and clean latents, then passed to ``stepper`` to
advance the noisy latents one step along the diffusion schedule.
Either ``video_state`` or ``audio_state`` may be ``None`` for absent
modalities; the absent modality is passed through unchanged.
### Parameters
sigmas:
A 1D tensor of noise levels (diffusion sigmas) defining the sampling
schedule. All steps except the last element are iterated over.
video_state:
The current video :class:`LatentState`, containing the noisy latent,
its clean reference latent, and the denoising mask.
The current video :class:`LatentState`, or ``None`` if video is absent.
audio_state:
The current audio :class:`LatentState`, analogous to ``video_state``
but for the audio modality.
The current audio :class:`LatentState`, or ``None`` if audio is absent.
stepper:
An implementation of :class:`DiffusionStepProtocol` that updates a
latent given the current latent, its denoised estimate, the full
``sigmas`` schedule, and the current step index.
denoise_fn:
A callable implementing :class:`DenoisingFunc`. It is invoked as
``denoise_fn(video_state, audio_state, sigmas, step_index)`` and must
return a tuple ``(denoised_video, denoised_audio)``, where each element
is a tensor with the same shape as the corresponding latent.
transformer:
The diffusion model passed to the denoiser at each step.
denoiser:
A callable implementing :class:`Denoiser`. It is invoked as
``denoiser(transformer, video_state, audio_state, sigmas, step_index)``
and must return ``(denoised_video, denoised_audio)``.
### Returns
tuple[LatentState, LatentState]
A pair ``(video_state, audio_state)`` containing the final video and
audio latent states after completing the denoising loop.
tuple[LatentState | None, LatentState | None]
Final ``(video_state, audio_state)`` after the denoising loop.
"""
for step_idx, _ in enumerate(tqdm(sigmas[:-1])):
denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_idx)
denoised_video, denoised_audio = denoiser(transformer, video_state, audio_state, sigmas, step_idx)
denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent)
denoised_audio = post_process_latent(denoised_audio, audio_state.denoise_mask, audio_state.clean_latent)
video_state = replace(video_state, latent=stepper.step(video_state.latent, denoised_video, sigmas, step_idx))
audio_state = replace(audio_state, latent=stepper.step(audio_state.latent, denoised_audio, sigmas, step_idx))
video_state = _step_state(video_state, denoised_video, stepper, sigmas, step_idx)
audio_state = _step_state(audio_state, denoised_audio, stepper, sigmas, step_idx)
return (video_state, audio_state)
def gradient_estimating_euler_denoising_loop(
sigmas: torch.Tensor,
video_state: LatentState,
audio_state: LatentState,
video_state: LatentState | None,
audio_state: LatentState | None,
stepper: DiffusionStepProtocol,
denoise_fn: DenoisingFunc,
transformer: X0Model,
denoiser: Denoiser,
ge_gamma: float = 2.0,
) -> tuple[LatentState, LatentState]:
) -> tuple[LatentState | None, LatentState | None]:
"""
Perform the joint audio-video denoising loop using gradient-estimation sampling.
This function is similar to :func:`euler_denoising_loop`, but applies
gradient estimation to improve the denoised estimates by tracking velocity
changes across steps. See the referenced function for detailed parameter
documentation.
Same interface as :func:`euler_denoising_loop` with an additional
``ge_gamma`` parameter for velocity correction.
### Parameters
ge_gamma:
Gradient estimation coefficient controlling the velocity correction term.
Default is 2.0. Paper: https://openreview.net/pdf?id=o2ND9v0CeK
sigmas, video_state, audio_state, stepper, denoise_fn:
See :func:`euler_denoising_loop` for parameter descriptions.
### Returns
tuple[LatentState, LatentState]
tuple[LatentState | None, LatentState | None]
See :func:`euler_denoising_loop` for return value description.
"""
@@ -105,23 +110,35 @@ def gradient_estimating_euler_denoising_loop(
return current_velocity, denoised_sample
for step_idx, _ in enumerate(tqdm(sigmas[:-1])):
denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_idx)
denoised_video, denoised_audio = denoiser(transformer, video_state, audio_state, sigmas, step_idx)
denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent)
denoised_audio = post_process_latent(denoised_audio, audio_state.denoise_mask, audio_state.clean_latent)
if video_state is not None and denoised_video is not None:
denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent)
if audio_state is not None and denoised_audio is not None:
denoised_audio = post_process_latent(denoised_audio, audio_state.denoise_mask, audio_state.clean_latent)
if sigmas[step_idx + 1] == 0:
return replace(video_state, latent=denoised_video), replace(audio_state, latent=denoised_audio)
if video_state is not None and denoised_video is not None:
video_state = replace(video_state, latent=denoised_video)
if audio_state is not None and denoised_audio is not None:
audio_state = replace(audio_state, latent=denoised_audio)
return video_state, audio_state
previous_video_velocity, denoised_video = update_velocity_and_sample(
video_state.latent, denoised_video, sigmas[step_idx], previous_video_velocity
)
previous_audio_velocity, denoised_audio = update_velocity_and_sample(
audio_state.latent, denoised_audio, sigmas[step_idx], previous_audio_velocity
)
if video_state is not None and denoised_video is not None:
previous_video_velocity, denoised_video = update_velocity_and_sample(
video_state.latent, denoised_video, sigmas[step_idx], previous_video_velocity
)
video_state = replace(
video_state, latent=stepper.step(video_state.latent, denoised_video, sigmas, step_idx)
)
video_state = replace(video_state, latent=stepper.step(video_state.latent, denoised_video, sigmas, step_idx))
audio_state = replace(audio_state, latent=stepper.step(audio_state.latent, denoised_audio, sigmas, step_idx))
if audio_state is not None and denoised_audio is not None:
previous_audio_velocity, denoised_audio = update_velocity_and_sample(
audio_state.latent, denoised_audio, sigmas[step_idx], previous_audio_velocity
)
audio_state = replace(
audio_state, latent=stepper.step(audio_state.latent, denoised_audio, sigmas, step_idx)
)
return (video_state, audio_state)
@@ -146,6 +163,7 @@ def _inject_sde_noise(
sigmas: torch.Tensor,
step_idx: int,
legacy_mode: bool = False,
eta: float = 0.5,
) -> torch.Tensor:
sigmas_copy = sigmas.clone()
new_noise = new_noise_fn(state.latent, step_noise_generator)
@@ -160,6 +178,7 @@ def _inject_sde_noise(
sigmas=sigmas,
step_index=step_idx,
noise=new_noise,
eta=eta,
)
if legacy_mode:
@@ -168,20 +187,22 @@ def _inject_sde_noise(
return x_next
def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915
def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915,PLR0912
sigmas: torch.Tensor,
video_state: LatentState,
audio_state: LatentState,
video_state: LatentState | None,
audio_state: LatentState | None,
stepper: DiffusionStepProtocol,
denoise_fn: DenoisingFunc,
transformer: X0Model,
denoiser: Denoiser,
noise_seed: int = -1,
noise_seed_substep: int | None = None,
eta: float = 0.5,
bongmath: bool = True,
bongmath_max_iter: int = 100,
new_noise_fn: Callable[[torch.Tensor, torch.Generator], torch.Tensor] = _get_new_noise,
model_dtype: torch.dtype = torch.bfloat16,
legacy_mode: bool = True,
) -> tuple[LatentState, LatentState]:
) -> tuple[LatentState | None, LatentState | None]:
"""
Joint audio-video denoising loop using the res_2s second-order sampler.
Iterates over the diffusion schedule with a two-stage Runge-Kutta step:
@@ -189,46 +210,48 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915
noise), then combines both with RK coefficients. Supports anchor-point
refinement (bong iteration) and optional SDE noise injection. Requires
:class:`Res2sDiffusionStep` as ``stepper``.
Either modality may be ``None`` (absent).
### Parameters
sigmas:
A 1D tensor of noise levels defining the sampling schedule.
video_state:
Current video :class:`LatentState` (noisy latent, clean reference, mask).
audio_state:
Current audio :class:`LatentState`, same structure as ``video_state``.
stepper:
Must be an instance of :class:`Res2sDiffusionStep`; performs SDE step
with noise injection.
denoise_fn:
Callable ``(video_state, audio_state, sigmas, step_index)`` returning
``(denoised_video, denoised_audio)``.
transformer:
The diffusion model passed to the denoiser at each step.
denoiser:
Callable implementing :class:`Denoiser`.
noise_seed:
Seed for step-level SDE noise; substep seed defaults to ``noise_seed + 10000``.
noise_seed_substep:
Optional seed for substep SDE noise; if None, derived from ``noise_seed``.
eta:
Controls stochastic noise injection strength (0=deterministic, 1=maximum).
Applies to main diffusion steps; substeps always use 0.5. Default 0.5.
bongmath:
Whether to run iterative anchor refinement (bong iteration) when step size is small.
bongmath_max_iter:
Max iterations for bong refinement when enabled.
new_noise_fn:
Callable ``(latent, generator) -> noise`` for SDE injection; default
uses normalized channel-wise Gaussian noise.
Callable ``(latent, generator) -> noise`` for SDE injection.
model_dtype:
Dtype for latent state updates (e.g. bfloat16).
### Returns
tuple[LatentState, LatentState]
tuple[LatentState | None, LatentState | None]
Final ``(video_state, audio_state)`` after the denoising loop.
"""
# Determine device from whichever state is present
present_state = video_state or audio_state
if present_state is None:
raise ValueError("At least one of video_state or audio_state must be provided")
state_device = present_state.latent.device
# Initialize noise generators with different seeds
if noise_seed_substep is None:
noise_seed_substep = noise_seed + 10000 # Offset to ensure different seeds
step_noise_generator = torch.Generator(device=video_state.latent.device).manual_seed(noise_seed)
substep_noise_generator = torch.Generator(device=video_state.latent.device).manual_seed(noise_seed_substep)
step_noise_generator = torch.Generator(device=state_device).manual_seed(noise_seed)
substep_noise_generator = torch.Generator(device=state_device).manual_seed(noise_seed_substep)
sde_noise_injecting_fn = partial(
_inject_sde_noise, stepper=stepper, new_noise_fn=new_noise_fn, legacy_mode=legacy_mode
)
step_noise_injecting_fn = partial(sde_noise_injecting_fn, step_noise_generator=step_noise_generator)
substep_noise_injecting_fn = partial(sde_noise_injecting_fn, step_noise_generator=substep_noise_generator)
step_noise_injecting_fn = partial(sde_noise_injecting_fn, step_noise_generator=step_noise_generator, eta=eta)
# substep eta is always default 0.5 for compatibility with original implementation.
substep_noise_injecting_fn = partial(sde_noise_injecting_fn, step_noise_generator=substep_noise_generator, eta=0.5)
if not isinstance(stepper, Res2sDiffusionStep):
raise ValueError("stepper must be an instance of Res2sDiffusionStep")
@@ -241,26 +264,25 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915
hs = -torch.log(sigmas[1:].double().cpu() / (sigmas[:-1].double().cpu()))
# Initialize phi cache for reuse across loop iterations
# Cache key: (j, neg_h) where j is phi order and neg_h is negative step value
phi_cache = {}
c2 = 0.5 # Midpoint for res_2s
# Progress bar shows only full two-stage steps; final (sigma_next==0) step is done silently
for step_idx in tqdm(range(n_full_steps)):
sigma = sigmas[step_idx].double()
sigma_next = sigmas[step_idx + 1].double()
# Initialize anchor point
x_anchor_video = video_state.latent.clone().double()
x_anchor_audio = audio_state.latent.clone().double()
x_anchor_video = video_state.latent.clone().double() if video_state is not None else None
x_anchor_audio = audio_state.latent.clone().double() if audio_state is not None else None
# ====================================================================
# STAGE 1: Evaluate at current point
# ====================================================================
denoised_video_1, denoised_audio_1 = denoise_fn(video_state, audio_state, sigmas, step_idx)
denoised_video_1 = post_process_latent(denoised_video_1, video_state.denoise_mask, video_state.clean_latent)
denoised_audio_1 = post_process_latent(denoised_audio_1, audio_state.denoise_mask, audio_state.clean_latent)
denoised_video_1, denoised_audio_1 = denoiser(transformer, video_state, audio_state, sigmas, step_idx)
if video_state is not None and denoised_video_1 is not None:
denoised_video_1 = post_process_latent(denoised_video_1, video_state.denoise_mask, video_state.clean_latent)
if audio_state is not None and denoised_audio_1 is not None:
denoised_audio_1 = post_process_latent(denoised_audio_1, audio_state.denoise_mask, audio_state.clean_latent)
h = hs[step_idx].item()
@@ -273,91 +295,127 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915
# ====================================================================
# Compute substep x using RK coefficient a21
# ====================================================================
eps_1_video = denoised_video_1.double() - x_anchor_video
eps_1_audio = denoised_audio_1.double() - x_anchor_audio
if x_anchor_video is not None and denoised_video_1 is not None:
eps_1_video = denoised_video_1.double() - x_anchor_video
x_mid_video = x_anchor_video.double() + h * a21 * eps_1_video
else:
eps_1_video = None
x_mid_video = None
x_mid_video = x_anchor_video.double() + h * a21 * eps_1_video
x_mid_audio = x_anchor_audio.double() + h * a21 * eps_1_audio
if x_anchor_audio is not None and denoised_audio_1 is not None:
eps_1_audio = denoised_audio_1.double() - x_anchor_audio
x_mid_audio = x_anchor_audio.double() + h * a21 * eps_1_audio
else:
eps_1_audio = None
x_mid_audio = None
# ====================================================================
# SDE noise injection at substep
# ====================================================================
x_mid_video = substep_noise_injecting_fn(
state=video_state,
sample=x_anchor_video,
denoised_sample=x_mid_video,
sigmas=torch.stack([sigma, sub_sigma]),
step_idx=0,
)
x_mid_audio = substep_noise_injecting_fn(
state=audio_state,
sample=x_anchor_audio,
denoised_sample=x_mid_audio,
sigmas=torch.stack([sigma, sub_sigma]),
step_idx=0,
)
if x_mid_video is not None and video_state is not None:
x_mid_video = substep_noise_injecting_fn(
state=video_state,
sample=x_anchor_video,
denoised_sample=x_mid_video,
sigmas=torch.stack([sigma, sub_sigma]),
step_idx=0,
)
if x_mid_audio is not None and audio_state is not None:
x_mid_audio = substep_noise_injecting_fn(
state=audio_state,
sample=x_anchor_audio,
denoised_sample=x_mid_audio,
sigmas=torch.stack([sigma, sub_sigma]),
step_idx=0,
)
# ====================================================================
# ITERATIVE REFINEMENT (Bong Iteration) - Stabilize anchor point
# ITERATIVE REFINEMENT (Bong Iteration)
# ====================================================================
if bongmath and h < 0.5 and sigma > 0.03:
for _ in range(bongmath_max_iter):
x_anchor_video = x_mid_video - h * a21 * eps_1_video
eps_1_video = denoised_video_1.double() - x_anchor_video
x_anchor_audio = x_mid_audio - h * a21 * eps_1_audio
eps_1_audio = denoised_audio_1.double() - x_anchor_audio
if x_mid_video is not None and eps_1_video is not None:
x_anchor_video = x_mid_video - h * a21 * eps_1_video
eps_1_video = denoised_video_1.double() - x_anchor_video
if x_mid_audio is not None and eps_1_audio is not None:
x_anchor_audio = x_mid_audio - h * a21 * eps_1_audio
eps_1_audio = denoised_audio_1.double() - x_anchor_audio
# ====================================================================
# STAGE 2: Evaluate at substep point (WITH NOISE)
# ====================================================================
mid_video_state = replace(video_state, latent=x_mid_video.to(model_dtype))
mid_audio_state = replace(audio_state, latent=x_mid_audio.to(model_dtype))
mid_video_state = (
replace(video_state, latent=x_mid_video.to(model_dtype))
if video_state is not None and x_mid_video is not None
else None
)
mid_audio_state = (
replace(audio_state, latent=x_mid_audio.to(model_dtype))
if audio_state is not None and x_mid_audio is not None
else None
)
denoised_video_2, denoised_audio_2 = denoise_fn(
denoised_video_2, denoised_audio_2 = denoiser(
transformer,
video_state=mid_video_state,
audio_state=mid_audio_state,
sigmas=torch.stack([sub_sigma]).to(sigmas.device),
step_index=0,
)
denoised_video_2 = post_process_latent(denoised_video_2, video_state.denoise_mask, video_state.clean_latent)
denoised_audio_2 = post_process_latent(denoised_audio_2, audio_state.denoise_mask, audio_state.clean_latent)
if video_state is not None and denoised_video_2 is not None:
denoised_video_2 = post_process_latent(denoised_video_2, video_state.denoise_mask, video_state.clean_latent)
if audio_state is not None and denoised_audio_2 is not None:
denoised_audio_2 = post_process_latent(denoised_audio_2, audio_state.denoise_mask, audio_state.clean_latent)
# ====================================================================
# FINAL COMBINATION: Compute x_next using RK coefficients
# ====================================================================
eps_2_video = denoised_video_2.double() - x_anchor_video
eps_2_audio = denoised_audio_2.double() - x_anchor_audio
if x_anchor_video is not None and eps_1_video is not None and denoised_video_2 is not None:
eps_2_video = denoised_video_2.double() - x_anchor_video
x_next_video = x_anchor_video + h * (b1 * eps_1_video + b2 * eps_2_video)
else:
x_next_video = None
x_next_video = x_anchor_video + h * (b1 * eps_1_video + b2 * eps_2_video)
x_next_audio = x_anchor_audio + h * (b1 * eps_1_audio + b2 * eps_2_audio)
if x_anchor_audio is not None and eps_1_audio is not None and denoised_audio_2 is not None:
eps_2_audio = denoised_audio_2.double() - x_anchor_audio
x_next_audio = x_anchor_audio + h * (b1 * eps_1_audio + b2 * eps_2_audio)
else:
x_next_audio = None
# ====================================================================
# SDE NOISE INJECTION AT STEP LEVEL
# ====================================================================
x_next_video = step_noise_injecting_fn(
state=video_state,
sample=x_anchor_video,
denoised_sample=x_next_video,
sigmas=sigmas,
step_idx=step_idx,
)
x_next_audio = step_noise_injecting_fn(
state=audio_state,
sample=x_anchor_audio,
denoised_sample=x_next_audio,
sigmas=sigmas,
step_idx=step_idx,
)
if x_next_video is not None and video_state is not None:
x_next_video = step_noise_injecting_fn(
state=video_state,
sample=x_anchor_video,
denoised_sample=x_next_video,
sigmas=sigmas,
step_idx=step_idx,
)
if x_next_audio is not None and audio_state is not None:
x_next_audio = step_noise_injecting_fn(
state=audio_state,
sample=x_anchor_audio,
denoised_sample=x_next_audio,
sigmas=sigmas,
step_idx=step_idx,
)
# Update states
video_state = replace(video_state, latent=x_next_video.to(model_dtype))
audio_state = replace(audio_state, latent=x_next_audio.to(model_dtype))
if video_state is not None and x_next_video is not None:
video_state = replace(video_state, latent=x_next_video.to(model_dtype))
if audio_state is not None and x_next_audio is not None:
audio_state = replace(audio_state, latent=x_next_audio.to(model_dtype))
# Final step if we need to fully remove the noise
if sigmas[-1] == 0:
denoised_video_1, denoised_audio_1 = denoise_fn(video_state, audio_state, sigmas, n_full_steps)
denoised_video_1 = post_process_latent(denoised_video_1, video_state.denoise_mask, video_state.clean_latent)
denoised_audio_1 = post_process_latent(denoised_audio_1, audio_state.denoise_mask, audio_state.clean_latent)
video_state = replace(video_state, latent=denoised_video_1.to(model_dtype))
audio_state = replace(audio_state, latent=denoised_audio_1.to(model_dtype))
denoised_video_1, denoised_audio_1 = denoiser(transformer, video_state, audio_state, sigmas, n_full_steps)
if video_state is not None and denoised_video_1 is not None:
denoised_video_1 = post_process_latent(denoised_video_1, video_state.denoise_mask, video_state.clean_latent)
video_state = replace(video_state, latent=denoised_video_1.to(model_dtype))
if audio_state is not None and denoised_audio_1 is not None:
denoised_audio_1 = post_process_latent(denoised_audio_1, audio_state.denoise_mask, audio_state.clean_latent)
audio_state = replace(audio_state, latent=denoised_audio_1.to(model_dtype))
return video_state, audio_state
@@ -1,9 +1,11 @@
from dataclasses import dataclass, field
from typing import Protocol
import torch
from ltx_core.components.patchifiers import AudioPatchifier, VideoLatentPatchifier
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.conditioning import ConditioningItem
from ltx_core.model.transformer import X0Model
from ltx_core.types import LatentState
from ltx_pipelines.utils.constants import VIDEO_LATENT_CHANNELS, VIDEO_SCALE_FACTORS
@@ -35,39 +37,40 @@ class PipelineComponents:
self.audio_patchifier = AudioPatchifier(patch_size=1)
class DenoisingFunc(Protocol):
"""
Protocol for a denoising function used in the LTX pipeline.
class Denoiser(Protocol):
"""Protocol for a denoiser that receives the transformer at call time.
The transformer is not stored it is passed as the first argument so the
caller (a denoising loop or a pipeline block) controls its lifecycle.
Args:
video_state (LatentState): The current latent state for video.
audio_state (LatentState): The current latent state for audio.
sigmas (torch.Tensor): A 1D tensor of sigma values for each diffusion step.
step_index (int): Index of the current denoising step.
transformer: The diffusion model.
video_state: Current video latent state, or ``None`` if absent.
audio_state: Current audio latent state, or ``None`` if absent.
sigmas: 1-D tensor of sigma values for each diffusion step.
step_index: Index of the current denoising step.
Returns:
tuple[torch.Tensor, torch.Tensor]: The denoised video and audio tensors.
"""
def __call__(
self, video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
) -> tuple[torch.Tensor, torch.Tensor]: ...
class DenoisingLoopFunc(Protocol):
"""
Protocol for a denoising loop function used in the LTX pipeline.
Args:
sigmas (torch.Tensor): A 1D tensor of sigma values for each diffusion step.
video_state (LatentState): The current latent state for video.
audio_state (LatentState): The current latent state for audio.
stepper (DiffusionStepProtocol): The diffusion step protocol to use.
Returns:
tuple[LatentState, LatentState]: The denoised video and audio latent states.
``(denoised_video, denoised_audio)`` tensors (either may be ``None``).
"""
def __call__(
self,
transformer: X0Model,
video_state: LatentState | None,
audio_state: LatentState | None,
sigmas: torch.Tensor,
video_state: LatentState,
audio_state: LatentState,
stepper: DiffusionStepProtocol,
) -> tuple[torch.Tensor, torch.Tensor]: ...
step_index: int,
) -> tuple[torch.Tensor | None, torch.Tensor | None]: ...
@dataclass(frozen=True)
class ModalitySpec:
"""Specification for one modality passed to a diffusion stage.
Carries everything needed to build the initial noised latent state
and run the denoising loop for a single modality (video or audio).
Tools are created by ``DiffusionStage`` from pixel-space dimensions.
"""
context: torch.Tensor
conditionings: list[ConditioningItem] = field(default_factory=list)
noise_scale: float = 1.0
frozen: bool = False
initial_latent: torch.Tensor | None = None
+2 -2
View File
@@ -1,6 +1,6 @@
[project]
name = "ltx-trainer"
version = "1.0.0"
version = "1.1.0"
description = "LTX-2 training, democratized."
readme = "README.md"
authors = [
@@ -48,7 +48,7 @@ build-backend = "hatchling.build"
[tool.ruff]
target-version = "1.0.0"
target-version = "1.1.0"
line-length = 120
[tool.ruff.lint]
+2 -3
View File
@@ -7,6 +7,8 @@ either LoRA fine-tuning or full model fine-tuning. It loads configuration from
a YAML file and passes it to the trainer.
Basic usage:
python scripts/train.py CONFIG_PATH [--disable-progress-bars]
Resume is automatic when a training state file exists next to the loaded checkpoint.
To start fresh, set `checkpoints.no_resume: true` in the YAML config.
For multi-GPU/FSDP training, configure and launch via Accelerate:
accelerate config
accelerate launch scripts/train.py CONFIG_PATH
@@ -39,7 +41,6 @@ def main(
),
) -> None:
"""Train the model using the provided configuration file."""
# Load the configuration from the YAML file
config_path = Path(config_path)
if not config_path.exists():
typer.echo(f"Error: Configuration file {config_path} does not exist.")
@@ -48,14 +49,12 @@ def main(
with open(config_path, "r") as file:
config_data = yaml.safe_load(file)
# Convert the loaded data to the LtxTrainerConfig object
try:
trainer_config = LtxTrainerConfig(**config_data)
except Exception as e:
typer.echo(f"Error: Invalid configuration data: {e}")
raise typer.Exit(code=1) from e
# Initialize the training process
trainer = LtxvTrainer(trainer_config)
trainer.train(disable_progress_bars=disable_progress_bars)
@@ -133,6 +133,7 @@ class OptimizationConfig(ConfigBaseModel):
"cosine",
"cosine_with_restarts",
"polynomial",
"step",
] = Field(
default="linear",
description="Type of scheduler to use for training",
@@ -398,6 +399,21 @@ class CheckpointsConfig(ConfigBaseModel):
description="Precision to use when saving checkpoint weights. Options: 'bfloat16' or 'float32'.",
)
no_resume: bool = Field(
default=False,
description="When True, ignore any saved training state and start from step 0. "
"Model weights from load_checkpoint are still loaded, but optimizer/scheduler "
"state and step counter are reset.",
)
save_training_state: Literal["full", "minimal", "off"] = Field(
default="minimal",
description="Save training state alongside checkpoints for resume. "
"'full': optimizer + scheduler + RNG + step (~800MB for LoRA, much larger for full fine-tuning). "
"'minimal': scheduler + RNG + step only (~few KB, sufficient for LoRA). "
"'off': nothing saved, resume not possible.",
)
class HubConfig(ConfigBaseModel):
"""Configuration for Hugging Face Hub integration"""
+227 -20
View File
@@ -1,4 +1,5 @@
import os
import re
import time
import warnings
from pathlib import Path
@@ -39,6 +40,7 @@ from ltx_trainer.model_loader import load_model as load_ltx_model
from ltx_trainer.progress import TrainingProgress
from ltx_trainer.quantization import quantize_model
from ltx_trainer.timestep_samplers import SAMPLERS
from ltx_trainer.training_state import ConfigFingerprint, RngStates, TrainingState
from ltx_trainer.training_strategies import get_training_strategy
from ltx_trainer.utils import open_image_as_srgb, save_image
from ltx_trainer.validation_sampler import CachedPromptEmbeddings, GenerationConfig, ValidationSampler
@@ -85,11 +87,14 @@ class LtxvTrainer:
self._load_models()
self._setup_accelerator()
self._collect_trainable_params()
self._loaded_checkpoint_path: Path | None = None
self._load_checkpoint()
self._prepare_models_for_training()
self._dataset = None
self._global_step = -1
self._checkpoint_paths = []
self._checkpoint_paths: list[Path] = []
self._training_state_paths: list[Path] = []
self._training_state_size_warned = False
self._init_wandb()
def train( # noqa: PLR0912, PLR0915
@@ -99,6 +104,9 @@ class LtxvTrainer:
) -> tuple[Path, TrainingStats]:
"""
Start the training process.
Args:
disable_progress_bars: Disable Rich progress bars (useful for multi-process runs).
step_callback: Optional callback invoked after each optimization step.
Returns:
Tuple of (saved_model_path, training_stats)
"""
@@ -108,11 +116,18 @@ class LtxvTrainer:
train_start_time = time.time()
# Use the same seed for all processes and ensure deterministic operations
initial_step, training_state = self._resume_state
resuming = training_state is not None
set_seed(cfg.seed)
logger.debug(f"Process {self._accelerator.process_index} using seed: {cfg.seed}")
self._init_optimizer()
if training_state is not None and not self._restore_training_state(training_state):
initial_step = 0
resuming = False
self._init_dataloader()
data_iter = iter(self._dataloader)
self._init_timestep_sampler()
@@ -125,27 +140,36 @@ class LtxvTrainer:
# Save the training configuration as YAML
self._save_config()
logger.info("🚀 Starting training...")
remaining_steps = cfg.optimization.steps - initial_step
if remaining_steps <= 0:
raise ValueError(
f"No remaining training steps: initial_step={initial_step} >= "
f"target_steps={cfg.optimization.steps}. Nothing to train."
)
if resuming:
logger.info(f"🚀 Resuming training from step {initial_step}{cfg.optimization.steps}")
else:
logger.info("🚀 Starting training...")
# Create progress tracking (disabled for non-main processes or when explicitly disabled)
progress_enabled = IS_MAIN_PROCESS and not disable_progress_bars
progress = TrainingProgress(
enabled=progress_enabled,
total_steps=cfg.optimization.steps,
total_steps=remaining_steps,
)
if IS_MAIN_PROCESS and disable_progress_bars:
logger.warning("Progress bars disabled. Intermediate status messages will be logged instead.")
self._transformer.train()
self._global_step = 0
self._global_step = initial_step
peak_mem_during_training = start_mem
sampled_videos_paths = None
with progress:
# Initial validation before training starts
if cfg.validation.interval and not cfg.validation.skip_initial_validation:
sampled_videos_paths = self._sample_videos(progress)
if IS_MAIN_PROCESS and sampled_videos_paths and self._config.wandb.log_validation_videos:
@@ -153,7 +177,7 @@ class LtxvTrainer:
self._accelerator.wait_for_everyone()
for step in range(cfg.optimization.steps * cfg.optimization.gradient_accumulation_steps):
for step in range(remaining_steps * cfg.optimization.gradient_accumulation_steps):
# Get next batch, reset the dataloader if needed
try:
batch = next(data_iter)
@@ -242,9 +266,9 @@ class LtxvTrainer:
# Fallback logging when progress bars are disabled
if disable_progress_bars and IS_MAIN_PROCESS and self._global_step % 20 == 0:
elapsed = time.time() - train_start_time
progress_percentage = self._global_step / cfg.optimization.steps
if progress_percentage > 0:
total_estimated = elapsed / progress_percentage
steps_done = self._global_step - initial_step
if steps_done > 0:
total_estimated = elapsed / steps_done * remaining_steps
total_time = f"{total_estimated // 3600:.0f}h {(total_estimated % 3600) // 60:.0f}m"
else:
total_time = "calculating..."
@@ -266,7 +290,7 @@ class LtxvTrainer:
# Calculate steps/second over entire training
total_time_seconds = train_end_time - train_start_time
steps_per_second = cfg.optimization.steps / total_time_seconds
steps_per_second = remaining_steps / total_time_seconds
samples_per_second = steps_per_second * self._accelerator.num_processes * cfg.optimization.batch_size
@@ -499,15 +523,18 @@ class LtxvTrainer:
self._transformer = get_peft_model(self._transformer, lora_config)
def _load_checkpoint(self) -> None:
"""Load checkpoint if specified in config."""
"""Load checkpoint if specified in config, then resolve resume state."""
if not self._config.model.load_checkpoint:
self._resume_state: tuple[int, TrainingState | None] = (0, None)
return
checkpoint_path = self._find_checkpoint(self._config.model.load_checkpoint)
if not checkpoint_path:
logger.warning(f"⚠️ Could not find checkpoint at {self._config.model.load_checkpoint}")
self._resume_state = (0, None)
return
self._loaded_checkpoint_path = checkpoint_path
logger.info(f"📥 Loading checkpoint from {checkpoint_path}")
if self._config.model.training_mode == "full":
@@ -515,6 +542,8 @@ class LtxvTrainer:
else: # LoRA mode
self._load_lora_checkpoint(checkpoint_path)
self._resume_state = self._resolve_resume_state()
def _load_full_checkpoint(self, checkpoint_path: Path) -> None:
"""Load full model checkpoint."""
state_dict = load_file(checkpoint_path)
@@ -536,6 +565,98 @@ class LtxvTrainer:
logger.info("✅ LoRA checkpoint loaded successfully")
def _resolve_resume_state(self) -> tuple[int, TrainingState | None]:
"""Determine resume state by looking for a training state file next to the loaded checkpoint.
Returns (initial_step, TrainingState or None).
If no_resume config is set, no checkpoint loaded, or no state file found: returns (0, None).
"""
if self._config.checkpoints.no_resume or self._loaded_checkpoint_path is None:
return 0, None
state = self._load_training_state(self._loaded_checkpoint_path)
if state is None:
return 0, None
fp = state.config_fingerprint
cfg = self._config
mismatches: list[str] = []
if fp.optimizer_type != cfg.optimization.optimizer_type:
mismatches.append(f"optimizer_type: {fp.optimizer_type}{cfg.optimization.optimizer_type}")
if fp.scheduler_type != cfg.optimization.scheduler_type:
mismatches.append(f"scheduler_type: {fp.scheduler_type}{cfg.optimization.scheduler_type}")
if fp.training_mode != cfg.model.training_mode:
mismatches.append(f"training_mode: {fp.training_mode}{cfg.model.training_mode}")
if (
cfg.model.training_mode == "lora"
and cfg.lora is not None
and fp.lora_rank is not None
and fp.lora_rank != cfg.lora.rank
):
mismatches.append(f"lora_rank: {fp.lora_rank}{cfg.lora.rank}")
if mismatches:
logger.warning(
f"⚠️ Training state config mismatch ({', '.join(mismatches)}). "
"Starting from step 0. Set checkpoints.no_resume=true to silence this warning."
)
return 0, None
if state.global_step < 0:
logger.warning(f"⚠️ Training state has invalid global_step={state.global_step!r}. Starting from step 0.")
return 0, None
logger.info(f"📌 Resuming from step {state.global_step}")
return state.global_step, state
@staticmethod
def _load_training_state(checkpoint_path: Path) -> TrainingState | None:
"""Load training state file that corresponds to a checkpoint weights file."""
match = re.search(r"step_(\d+)", checkpoint_path.name)
if not match:
return None
step_str = match.group(1)
state_path = checkpoint_path.parent / f"training_state_step_{step_str}.pt"
if not state_path.exists():
return None
try:
raw: dict = torch.load(state_path, map_location="cpu", weights_only=False)
state = TrainingState.from_save_dict(raw)
logger.info(f"📥 Loaded training state from {state_path}")
return state
except Exception as e:
logger.warning(f"⚠️ Failed to load training state from {state_path}: {e}. Starting from step 0.")
return None
def _restore_training_state(self, training_state: TrainingState) -> bool:
"""Restore optimizer, scheduler, and RNG states from a loaded TrainingState.
Must be called after _init_optimizer() (which calls accelerator.prepare).
Returns True if restore succeeded, False if it failed (caller should fall back to step 0).
"""
try:
if training_state.optimizer_state_dict is not None:
self._optimizer.load_state_dict(training_state.optimizer_state_dict)
logger.debug("Restored optimizer state (full mode)")
if training_state.lr_scheduler_state_dict is not None and self._lr_scheduler is not None:
self._lr_scheduler.load_state_dict(training_state.lr_scheduler_state_dict)
logger.debug("Restored LR scheduler state")
except Exception as e:
logger.warning(f"⚠️ Failed to restore training state: {e}. Starting from step 0.")
return False
rng = training_state.rng_states
if self._accelerator.num_processes > 1:
logger.debug("Skipping RNG restore in multi-process mode (only main process state was saved)")
else:
if rng.torch_state is not None:
torch.random.set_rng_state(rng.torch_state)
if rng.cuda_state is not None and torch.cuda.is_available():
torch.cuda.set_rng_state(rng.cuda_state)
logger.debug("Restored RNG states")
return True
def _prepare_models_for_training(self) -> None:
"""Prepare models for training with Accelerate."""
@@ -643,7 +764,6 @@ class LtxvTrainer:
else:
raise ValueError(f"Unknown optimizer type: {opt_cfg.optimizer_type}")
# Add scheduler initialization
lr_scheduler = self._create_scheduler(optimizer)
# noinspection PyTypeChecker
@@ -676,8 +796,8 @@ class LtxvTrainer:
elif scheduler_type == "cosine_with_restarts":
scheduler = CosineAnnealingWarmRestarts(
optimizer,
T_0=params.pop("T_0", steps // 4), # First restart cycle length
T_mult=params.pop("T_mult", 1), # Multiplicative factor for cycle lengths
T_0=params.pop("T_0", steps // 4),
T_mult=params.pop("T_mult", 1),
eta_min=params.pop("eta_min", 5e-5),
**params,
)
@@ -924,9 +1044,11 @@ class LtxvTrainer:
rel_path = saved_weights_path.relative_to(self._config.output_dir)
logger.info(f"💾 {prefix.capitalize()} weights for step {self._global_step} saved in {rel_path}")
# Keep track of checkpoint paths, and cleanup old checkpoints if needed
self._checkpoint_paths.append(saved_weights_path)
self._cleanup_checkpoints()
self._save_training_state(save_dir)
return saved_weights_path
def _cleanup_checkpoints(self) -> None:
@@ -936,10 +1058,88 @@ class LtxvTrainer:
for old_checkpoint in checkpoints_to_remove:
if old_checkpoint.exists():
old_checkpoint.unlink()
logger.info(f"Removed old checkpoints: {old_checkpoint}")
# Update the list to only contain kept checkpoints
logger.info(f"Removed old checkpoint: {old_checkpoint}")
self._checkpoint_paths = self._checkpoint_paths[-self._config.checkpoints.keep_last_n :]
def _save_training_state(self, save_dir: Path) -> None:
"""Save training state alongside checkpoint for resume.
Respects checkpoints.save_training_state config:
- "full": optimizer + scheduler + RNG + step
- "minimal": scheduler + RNG + step only
- "off": skip entirely
"""
if not IS_MAIN_PROCESS:
return
mode = self._config.checkpoints.save_training_state
if mode == "off":
return
is_fsdp = self._accelerator.distributed_type == DistributedType.FSDP
optimizer_state = None
if mode == "full":
if is_fsdp:
logger.warning(
"⚠️ save_training_state='full' is not supported with FSDP. "
"Saving 'minimal' state (scheduler + RNG only)."
)
else:
optimizer_state = self._optimizer.state_dict()
state = TrainingState(
global_step=self._global_step,
config_fingerprint=ConfigFingerprint(
optimizer_type=self._config.optimization.optimizer_type,
scheduler_type=self._config.optimization.scheduler_type,
training_mode=self._config.model.training_mode,
lora_rank=self._config.lora.rank if self._config.lora is not None else None,
),
rng_states=RngStates(
torch_state=torch.random.get_rng_state(),
cuda_state=torch.cuda.get_rng_state() if torch.cuda.is_available() else None,
),
lr_scheduler_state_dict=self._lr_scheduler.state_dict() if self._lr_scheduler is not None else None,
optimizer_state_dict=optimizer_state,
)
state_path = save_dir / f"training_state_step_{self._global_step:05d}.pt"
tmp_path = state_path.with_suffix(".pt.tmp")
try:
torch.save(state.to_save_dict(), tmp_path)
except Exception:
if tmp_path.exists():
tmp_path.unlink()
raise
tmp_path.rename(state_path)
file_size_gb = state_path.stat().st_size / (1024**3)
if file_size_gb > 1.0 and not self._training_state_size_warned:
self._training_state_size_warned = True
logger.warning(
f"⚠️ Training state file is {file_size_gb:.1f} GB (full mode includes optimizer state). "
f'Set checkpoints.save_training_state="minimal" to save only scheduler/RNG/step (~few KB), '
f'or "off" to disable entirely.'
)
if not self._training_state_paths or self._training_state_paths[-1] != state_path:
self._training_state_paths.append(state_path)
self._cleanup_training_states()
rel_path = state_path.relative_to(self._config.output_dir)
logger.debug(f"Training state saved to {rel_path}")
def _cleanup_training_states(self) -> None:
"""Clean up old training state files, using the same keep_last_n as checkpoints."""
keep_n = self._config.checkpoints.keep_last_n
if 0 < keep_n < len(self._training_state_paths):
to_remove = self._training_state_paths[:-keep_n]
for old_state in to_remove:
if old_state.exists():
old_state.unlink()
logger.debug(f"Removed old training state: {old_state}")
self._training_state_paths = self._training_state_paths[-keep_n:]
def _build_checkpoint_metadata(self) -> dict[str, str]:
"""Build metadata dictionary for safetensors checkpoint.
Delegates to the training strategy to get strategy-specific metadata
@@ -994,7 +1194,14 @@ class LtxvTrainer:
# Determine if outputs are images or videos based on file extension
is_image = sample_paths and sample_paths[0].suffix.lower() in (".png", ".jpg", ".jpeg", ".heic", ".webp")
media_cls = wandb.Image if is_image else wandb.Video
samples = [media_cls(str(path), caption=prompt) for path, prompt in zip(sample_paths, prompts, strict=True)]
if is_image:
samples = [
wandb.Image(str(path), caption=prompt) for path, prompt in zip(sample_paths, prompts, strict=True)
]
else:
samples = [
wandb.Video(str(path), caption=prompt, format=path.suffix.lower().lstrip("."))
for path, prompt in zip(sample_paths, prompts, strict=True)
]
self._wandb_run.log({"validation_samples": samples}, step=self._global_step)
@@ -0,0 +1,51 @@
from __future__ import annotations
from typing import Any
import torch
from pydantic import BaseModel, ConfigDict
class ConfigFingerprint(BaseModel):
optimizer_type: str
scheduler_type: str
training_mode: str
lora_rank: int | None = None
class RngStates(BaseModel):
model_config = ConfigDict(arbitrary_types_allowed=True)
torch_state: torch.Tensor
cuda_state: torch.Tensor | None = None
class TrainingState(BaseModel):
model_config = ConfigDict(arbitrary_types_allowed=True)
global_step: int
config_fingerprint: ConfigFingerprint
rng_states: RngStates
lr_scheduler_state_dict: dict[str, Any] | None = None
optimizer_state_dict: dict[str, Any] | None = None
def to_save_dict(self) -> dict[str, Any]:
"""Build dict suitable for torch.save -- recurses BaseModel sub-models, passes tensors/dicts through."""
def _convert(value: object) -> object:
if isinstance(value, BaseModel):
return {k: _convert(v) for k, v in value if v is not None}
return value
return {k: _convert(v) for k, v in self if v is not None}
@classmethod
def from_save_dict(cls, data: dict[str, Any]) -> TrainingState:
"""Construct from torch.load output with Pydantic validation."""
return cls(
global_step=data["global_step"],
config_fingerprint=ConfigFingerprint(**data["config_fingerprint"]),
rng_states=RngStates(**data["rng_states"]),
lr_scheduler_state_dict=data.get("lr_scheduler_state_dict"),
optimizer_state_dict=data.get("optimizer_state_dict"),
)
@@ -767,8 +767,9 @@ class ValidationSampler:
def _decode_audio(self, audio_state: LatentState, device: torch.device) -> Tensor:
"""Decode audio latents to waveform."""
self._audio_decoder.to(device)
# Ensure latent is bfloat16 to match decoder weights
latent = audio_state.latent.to(dtype=torch.bfloat16)
first_param = next(self._audio_decoder.parameters(), None)
decoder_dtype = first_param.dtype if first_param is not None else audio_state.latent.dtype
latent = audio_state.latent.to(dtype=decoder_dtype, device=device)
decoded_audio = self._audio_decoder(latent)
self._audio_decoder.to("cpu")
+2
View File
@@ -11,6 +11,8 @@ url = "https://pypi.org/simple"
[dependency-groups]
dev = [
"google-cloud-storage>=2.0",
"matplotlib>=3.7",
"pre-commit>=4.3.0",
"ruff>=0.14.3",
"pytest~=9.0",
Generated
+4243 -3908
View File
File diff suppressed because it is too large Load Diff