Automated PR - 2026-05-11

This commit is contained in:
github-actions[bot]
2026-05-11 13:14:05 +00:00
parent 41d9243716
commit 7df34dfa83
72 changed files with 3299 additions and 911 deletions
+6
View File
@@ -23,6 +23,7 @@ Inference pipelines for LTX-2 audio-video generation. Depends on `ltx-core` for
| `KeyframeInterpolationPipeline` | `keyframe_interpolation.py` | 2 | Full + distilled LoRA | Euler | Keyframe interpolation |
| `DistilledPipeline` | `distilled.py` | 2 | Distilled only | Euler | Fastest inference |
| `ICLoraPipeline` | `ic_lora.py` | 2 | Distilled only | Euler | Video-to-video with IC-LoRA control |
| `LipDubPipeline` | `lipdub.py` | 2 | Distilled only | Euler | Lip dubbing with IC-LoRA + audio ref conditioning |
| `RetakePipeline` | `retake.py` | 1 | Full or distilled | Euler | Video region regeneration |
## Guidance
@@ -65,6 +66,10 @@ Inference pipelines for LTX-2 audio-video generation. Depends on `ltx-core` for
- `GuidedDenoiser` -- CFG/STG with static `MultiModalGuider` instances (HQ, A2Vid, Retake non-distilled).
- `FactoryGuidedDenoiser` -- per-step guider creation via factory (OneStageTI2Vid, TwoStagesTI2Vid, Keyframe).
All denoisers return a `(video_result, audio_result)` tuple of `DenoisedLatentResult` (defined in `utils/types.py`), either element may be `None` for absent modalities. `DenoisedLatentResult.denoised` is the final blended tensor. Guided denoisers additionally populate per-pass fields (`.cond`, `.uncond`, `.ptb`, `.mod`) on each result; `SimpleDenoiser` leaves these `None`.
`GuidedDenoiser` and `FactoryGuidedDenoiser` accept `force_uncond_pass=True` to run the uncond pass even when `cfg_scale=1.0` (required by CFG++ when the guidance scale is 1 but the uncond prediction is still needed for the ODE derivative). Requires `negative_context` to be set on the guider. When enabled, `DenoisedLatentResult.uncond` will be a tensor instead of `None`.
Guided denoisers batch all guidance passes into a **single transformer call**: states are repeated along the batch dimension, contexts concatenated, and a `BatchedPerturbationConfig` controls which attention ops are skipped per sample. Pass count is dynamic: B=2 for CFG-only, up to B=4 with CFG+STG+modality isolation. Results are split back and blended by the guider.
## Per-pipeline unique features
@@ -72,6 +77,7 @@ Guided denoisers batch all guidance passes into a **single transformer call**: s
- **HQ**: Res2s second-order sampler for **both** stages, latent-dependent sigma schedule, distilled LoRA on both stages with separate strengths.
- **A2Vid**: Audio frozen in both stages (`frozen=True, noise_scale=0.0`). Returns original audio (not VAE-decoded); no `AudioDecoder`.
- **IC-LoRA**: `VideoConditionByReferenceLatent`, `reference_downscale_factor` from LoRA metadata, `skip_stage_2`, attention mask downsampling. Stage 2 is LoRA-free and uses `combined_image_conditionings` (no IC-LoRA conditioning).
- **LipDub**: Standalone pipeline; IC reference **video** helpers in `iclora_utils.py`, LipDub-only **audio** patchify/negative positions in `lipdub.py`. Appends frozen audio-reference tokens via `AudioConditionByReferenceLatent` (ltx-core), matching video token order (`[target | ref]`) while keeping reference RoPE positions negative (training-compatible). Single IC-LoRA on both stages; full IC-LoRA video conditioning at stage 1 and 2; stage-2 audio is frozen with S1 latent as initial state and uses S1-derived ref. Final audio decoded from stage 1 latent. The LipDub CLI does not expose `--conditioning-attention-mask`; use `ic_lora.py` if you need spatial IC attention masking.
- **Keyframe**: Uses `image_conditionings_by_adding_guiding_latent` in both stages (all frames as keyframe guidance, no replacement) -- unlike TI2Vid which uses `combined_image_conditionings` (frame_idx=0 replaces, others guide).
- **Retake**: `TemporalRegionMask` for selective time-window regeneration. `regenerate_video`/`regenerate_audio` flags. Conditional distilled/full behavior.
- **Distilled**: Single `self.stage` reused for both stages (not `stage_1`/`stage_2`).
+16
View File
@@ -64,6 +64,7 @@ Available pipeline modules:
- `ltx_pipelines.a2vid_two_stage` - Audio-to-video generation conditioned on an input audio.
- `ltx_pipelines.retake` - Regenerate a time region of an existing video.
- `ltx_pipelines.hdr_ic_lora` - Video-to-video with HDR output (linear float via LogC3 inverse decode).
- `ltx_pipelines.lipdub` - Lip dubbing / re-voicing with IC-LoRA and audio reference conditioning.
Use `--help` with any pipeline module to see all available options and parameters.
@@ -113,6 +114,7 @@ Do you need to condition on existing images/videos?
| **A2VidPipelineTwoStage** | 2 | ✅ | ✅ | Audio + Image | Audio-driven video generation |
| **RetakePipeline** | 1 | ✅ | ❌ | Source Video | Regenerating a time region of a video |
| **HDRICLoraPipeline** | 2 | ❌ | ✅ | Video | HDR video-to-video (linear float output for EXR) |
| **LipDubPipeline** | 2 | ✅ | ✅ | Video + Audio | Lip dubbing with audio ref conditioning |
---
@@ -238,6 +240,20 @@ Two-stage video-to-video on the distilled model with an HDR IC-LoRA. Decoded lat
---
### 10. LipDubPipeline
**Best for:** Lip dubbing, rephrasing while keeping the same speaker identity and matching lip movements to new audio.
**Source**: [`src/ltx_pipelines/lipdub.py`](src/ltx_pipelines/lipdub.py)
Uses IC-LoRA on a **distilled** checkpoint with a **single** lip-dub IC-LoRA applied in **both** stages. The reference clip provides video and audio reference tokens whose VAE latents are appended to the target audio sequence as frozen reference tokens. The frame count and frame rate are derived from the reference video (frame count is silently snapped to the nearest `8k+1`), so the CLI does not accept `--num-frames` or `--frame-rate`. Required: `--reference-video`. Optional: `--reference-strength`. LoRA: [`Lightricks/LTX-2.3-22b-IC-LoRA-LipDub`](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-LipDub).
**Note:** Requires a distilled model checkpoint and one lip-dub IC-LoRA (`--lora` exactly once).
**Use when:** Dubbing, rephrasing with matched lips and speaker identity.
---
## 🎨 Conditioning Types
Pipelines use different conditioning methods from [`ltx-core`](../ltx-core/) for controlling generation. See the [ltx-core conditioning documentation](../ltx-core/README.md#conditioning--control) for details.
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "ltx-pipelines"
version = "1.1.2"
version = "1.1.3"
description = "Pipelines implementation for Lightricks' LTX-2 model"
readme = "README.md"
requires-python = ">=3.10"
@@ -5,6 +5,7 @@ This package provides ready-to-use pipelines for video generation:
- TI2VidTwoStagesPipeline: Two-stage generation with upsampling
- DistilledPipeline: Fast distilled two-stage generation
- ICLoraPipeline: Image/video conditioning with distilled LoRA
- LipDubPipeline: Lip dubbing with IC-LoRA and audio conditioning
- KeyframeInterpolationPipeline: Keyframe-based video interpolation
- RetakePipeline: Regenerate a time region (retake) of an existing video
For more detailed components and utilities, import from specific submodules
@@ -15,6 +16,7 @@ from ltx_pipelines.a2vid_two_stage import A2VidPipelineTwoStage
from ltx_pipelines.distilled import DistilledPipeline
from ltx_pipelines.ic_lora import ICLoraPipeline
from ltx_pipelines.keyframe_interpolation import KeyframeInterpolationPipeline
from ltx_pipelines.lipdub import LipDubPipeline
from ltx_pipelines.retake import RetakePipeline
from ltx_pipelines.ti2vid_one_stage import TI2VidOneStagePipeline
from ltx_pipelines.ti2vid_two_stages import TI2VidTwoStagesPipeline
@@ -24,6 +26,7 @@ __all__ = [
"DistilledPipeline",
"ICLoraPipeline",
"KeyframeInterpolationPipeline",
"LipDubPipeline",
"RetakePipeline",
"TI2VidOneStagePipeline",
"TI2VidTwoStagesPipeline",
@@ -48,7 +48,7 @@ from ltx_core.model.video_vae import TilingConfig, VideoEncoder
from ltx_core.quantization import QuantizationPolicy
from ltx_core.tiling import DimensionTilingConfig, TileCountConfig
from ltx_core.tools import VideoLatentTools
from ltx_core.types import VideoLatentShape
from ltx_core.types import VideoLatentShape, VideoPixelShape
from ltx_pipelines.utils.blocks import (
DiffusionStage,
ImageConditioner,
@@ -412,7 +412,22 @@ class HDRICLoraPipeline:
high_quality_hdr=high_quality_hdr,
)
)
with self.stage_2.model_context() as transformer:
# video_tools is required by TiledDataParallelBuilder when stage_2 is
# wrapped for multi-GPU
stage2_video_tools = VideoLatentTools(
VideoLatentPatchifier(patch_size=1),
VideoLatentShape.from_pixel_shape(
VideoPixelShape(
batch=1,
frames=gen_num_frames,
height=gen_h,
width=gen_w,
fps=frame_rate,
)
),
frame_rate,
)
with self.stage_2.model_context(video_tools=stage2_video_tools) as transformer:
phase_latent = upscaled_video_latent
for phase_idx, (tiling, sigmas_list, use_ic) in enumerate(
zip(stage2_tilings, stage2_sigmas, stage2_use_ic_lora, strict=True)
@@ -542,10 +557,10 @@ class HDRICLoraPipeline:
"""
# Cast to float32 so tiled-decode accumulation buffers and blending
# masks run in full precision, avoiding bfloat16 seam artifacts.
# Request float32 [0, 1] output — apply_hdr_decode_postprocess expects it.
# apply_hdr_decode_postprocess expects float32 [0, 1].
latent = latent.float()
decoded = torch.cat(
list(self.video_decoder(latent, tiling_config, generator, output_dtype=torch.float32)),
[chunk.float() for chunk in self.video_decoder(latent, tiling_config, generator)],
dim=0,
)
decoded = rearrange(decoded, "f h w c -> 1 c f h w")
@@ -2,20 +2,18 @@ import logging
from collections.abc import Iterator
import torch
from einops import rearrange
from safetensors import safe_open
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.conditioning import (
ConditioningItem,
ConditioningItemAttentionStrengthWrapper,
VideoConditionByReferenceLatent,
)
from ltx_core.conditioning import ConditioningItem
from ltx_core.loader import LoraPathStrengthAndSDOps
from ltx_core.loader.registry import Registry
from ltx_core.model.video_vae import TilingConfig, VideoEncoder, get_video_chunks_number
from ltx_core.quantization import QuantizationPolicy
from ltx_core.types import Audio, VideoLatentShape, VideoPixelShape
from ltx_core.types import Audio, VideoPixelShape
from ltx_pipelines.iclora_utils import (
append_ic_lora_reference_video_conditionings,
read_lora_reference_downscale_factor,
)
from ltx_pipelines.utils.args import (
ImageConditioningInput,
VideoConditioningAction,
@@ -108,7 +106,7 @@ class ICLoraPipeline:
# so inference can resize reference videos to match training conditions.
self.reference_downscale_factor = 1
for lora in loras:
scale = _read_lora_reference_downscale_factor(lora.path)
scale = read_lora_reference_downscale_factor(lora.path)
if scale != 1:
if self.reference_downscale_factor not in (1, scale):
raise ValueError(
@@ -309,104 +307,26 @@ class ICLoraPipeline:
device=self.device,
)
# Calculate scaled dimensions for reference video conditioning.
# IC-LoRAs trained with downscaled reference videos expect the same ratio at inference.
scale = self.reference_downscale_factor
if scale != 1 and (height % scale != 0 or width % scale != 0):
raise ValueError(
f"Output dimensions ({height}x{width}) must be divisible by reference_downscale_factor ({scale})"
)
ref_height = height // scale
ref_width = width // scale
for video_path, strength in video_conditioning:
# Load video at scaled-down resolution (if scale > 1)
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)
# Build attention_mask for ConditioningItemAttentionStrengthWrapper
if conditioning_attention_mask is not None:
# Downsample pixel-space mask to latent space, then scale by strength
latent_mask = self._downsample_mask_to_latent(
mask=conditioning_attention_mask,
target_latent_shape=reference_video_shape,
)
attn_mask = latent_mask * conditioning_attention_strength
elif conditioning_attention_strength < 1.0:
# Use scalar strength only
attn_mask = conditioning_attention_strength
else:
attn_mask = None
cond = VideoConditionByReferenceLatent(
latent=encoded_video,
downscale_factor=scale,
strength=strength,
)
if attn_mask is not None:
cond = ConditioningItemAttentionStrengthWrapper(cond, attention_mask=attn_mask)
conditionings.append(cond)
append_ic_lora_reference_video_conditionings(
conditionings,
video_conditioning,
height=height,
width=width,
num_frames=num_frames,
video_encoder=video_encoder,
dtype=self.dtype,
device=self.device,
reference_downscale_factor=self.reference_downscale_factor,
conditioning_attention_strength=conditioning_attention_strength,
conditioning_attention_mask=conditioning_attention_mask,
tiling_config=None,
)
if video_conditioning:
logging.info(f"[IC-LoRA] Added {len(video_conditioning)} video conditioning(s)")
logging.info("[IC-LoRA] Added %d video conditioning(s)", len(video_conditioning))
return conditionings
@staticmethod
def _downsample_mask_to_latent(
mask: torch.Tensor,
target_latent_shape: VideoLatentShape,
) -> torch.Tensor:
"""
Downsample a pixel-space mask to latent space using VAE scale factors.
Handles causal temporal downsampling: the first frame is kept separately
(temporal scale factor = 1 for the first frame), while the remaining
frames are downsampled by the VAE's temporal scale factor.
Args:
mask: Pixel-space mask of shape (B, 1, F_pixel, H_pixel, W_pixel).
Values in [0, 1].
target_latent_shape: Expected latent shape after VAE encoding.
Used to determine the target (F_latent, H_latent, W_latent).
Returns:
Flattened latent-space mask of shape (B, F_lat * H_lat * W_lat),
matching the patchifier's token ordering (f, h, w).
"""
b = mask.shape[0]
f_lat = target_latent_shape.frames
h_lat = target_latent_shape.height
w_lat = target_latent_shape.width
# Step 1: Spatial downsampling (area interpolation per frame)
f_pix = mask.shape[2]
spatial_down = torch.nn.functional.interpolate(
rearrange(mask, "b 1 f h w -> (b f) 1 h w"),
size=(h_lat, w_lat),
mode="area",
)
spatial_down = rearrange(spatial_down, "(b f) 1 h w -> b 1 f h w", b=b)
# Step 2: Causal temporal downsampling
# First frame: kept as-is (causal VAE encodes first frame independently)
first_frame = spatial_down[:, :, :1, :, :] # (B, 1, 1, H_lat, W_lat)
if f_pix > 1 and f_lat > 1:
# Remaining frames: downsample by temporal factor via group-mean
t = (f_pix - 1) // (f_lat - 1) # temporal downscale factor
assert (f_pix - 1) % (f_lat - 1) == 0, (
f"Pixel frames ({f_pix}) not compatible with latent frames ({f_lat}): "
f"(f_pix - 1) must be divisible by (f_lat - 1)"
)
rest = rearrange(spatial_down[:, :, 1:, :, :], "b 1 (f t) h w -> b 1 f t h w", t=t)
rest = rest.mean(dim=3) # (B, 1, F_lat-1, H_lat, W_lat)
latent_mask = torch.cat([first_frame, rest], dim=2) # (B, 1, F_lat, H_lat, W_lat)
else:
latent_mask = first_frame
# Flatten to (B, F_lat * H_lat * W_lat) matching patchifier token order (f, h, w)
return rearrange(latent_mask, "b 1 f h w -> b (f h w)")
@torch.inference_mode()
def main() -> None:
@@ -523,26 +443,5 @@ def _load_mask_video(
return mask.clamp(0.0, 1.0)
def _read_lora_reference_downscale_factor(lora_path: str) -> int:
"""Read reference_downscale_factor from LoRA safetensors metadata.
Some IC-LoRA models are trained with reference videos at lower resolution than
the target output. This allows for more efficient training and can improve
generalization. The downscale factor indicates the ratio between target and
reference resolutions (e.g., factor=2 means reference is half the resolution).
Args:
lora_path: Path to the LoRA .safetensors file
Returns:
The reference downscale factor (1 if not specified in metadata, meaning
reference and target have the same resolution)
"""
try:
with safe_open(lora_path, framework="pt") as f:
metadata = f.metadata() or {}
return int(metadata.get("reference_downscale_factor", 1))
except Exception as e:
logging.warning(f"Failed to read metadata from LoRA file '{lora_path}': {e}")
return 1
if __name__ == "__main__":
main()
@@ -0,0 +1,120 @@
"""Shared IC-LoRA helpers: LoRA metadata, mask downsampling, reference-video conditioning.
Used by ``ic_lora`` and ``lipdub`` (video reference path only). LipDub audio helpers live in ``lipdub.py``.
"""
from __future__ import annotations
import logging
import torch
from einops import rearrange
from safetensors import safe_open
from ltx_core.conditioning import (
ConditioningItem,
ConditioningItemAttentionStrengthWrapper,
VideoConditionByReferenceLatent,
)
from ltx_core.model.video_vae import TilingConfig, VideoEncoder
from ltx_core.types import VideoLatentShape
from ltx_pipelines.utils.media_io import decode_video_by_frame, video_preprocess
def read_lora_reference_downscale_factor(lora_path: str) -> int:
"""Read ``reference_downscale_factor`` from LoRA safetensors metadata (default 1)."""
try:
with safe_open(lora_path, framework="pt") as f:
metadata = f.metadata() or {}
return int(metadata.get("reference_downscale_factor", 1))
except Exception as e:
logging.warning("Failed to read metadata from LoRA file '%s': %s", lora_path, e)
return 1
def downsample_mask_video_to_latent(
mask: torch.Tensor,
target_latent_shape: VideoLatentShape,
) -> torch.Tensor:
"""Downsample a pixel-space mask video to flattened latent token weights."""
b = mask.shape[0]
f_lat = target_latent_shape.frames
h_lat = target_latent_shape.height
w_lat = target_latent_shape.width
f_pix = mask.shape[2]
spatial_down = torch.nn.functional.interpolate(
rearrange(mask, "b 1 f h w -> (b f) 1 h w"),
size=(h_lat, w_lat),
mode="area",
)
spatial_down = rearrange(spatial_down, "(b f) 1 h w -> b 1 f h w", b=b)
first_frame = spatial_down[:, :, :1, :, :]
if f_pix > 1 and f_lat > 1:
t = (f_pix - 1) // (f_lat - 1)
assert (f_pix - 1) % (f_lat - 1) == 0, (
f"Pixel frames ({f_pix}) not compatible with latent frames ({f_lat}): "
f"(f_pix - 1) must be divisible by (f_lat - 1)"
)
rest = rearrange(spatial_down[:, :, 1:, :, :], "b 1 (f t) h w -> b 1 f t h w", t=t)
rest = rest.mean(dim=3)
latent_mask = torch.cat([first_frame, rest], dim=2)
else:
latent_mask = first_frame
return rearrange(latent_mask, "b 1 f h w -> b (f h w)")
def append_ic_lora_reference_video_conditionings( # noqa: PLR0913
conditionings: list[ConditioningItem],
video_conditioning: list[tuple[str, float]],
*,
height: int,
width: int,
num_frames: int,
video_encoder: VideoEncoder,
dtype: torch.dtype,
device: torch.device,
reference_downscale_factor: int,
conditioning_attention_strength: float,
conditioning_attention_mask: torch.Tensor | None,
tiling_config: TilingConfig | None = None,
) -> None:
"""Append :class:`VideoConditionByReferenceLatent` items for each reference path."""
scale = reference_downscale_factor
if scale != 1 and (height % scale != 0 or width % scale != 0):
raise ValueError(
f"Output dimensions ({height}x{width}) must be divisible by reference_downscale_factor ({scale})"
)
ref_height = height // scale
ref_width = width // scale
for video_path, strength in video_conditioning:
frame_gen = decode_video_by_frame(path=video_path, frame_cap=num_frames, device=device)
video = video_preprocess(frame_gen, ref_height, ref_width, dtype, device)
if tiling_config is not None:
encoded_video = video_encoder.tiled_encode(video, tiling_config)
else:
encoded_video = video_encoder(video)
reference_video_shape = VideoLatentShape.from_torch_shape(encoded_video.shape)
if conditioning_attention_mask is not None:
latent_mask = downsample_mask_video_to_latent(
mask=conditioning_attention_mask,
target_latent_shape=reference_video_shape,
)
attn_mask = latent_mask * conditioning_attention_strength
elif conditioning_attention_strength < 1.0:
attn_mask = conditioning_attention_strength
else:
attn_mask = None
cond = VideoConditionByReferenceLatent(
latent=encoded_video,
downscale_factor=scale,
strength=strength,
)
if attn_mask is not None:
cond = ConditioningItemAttentionStrengthWrapper(cond, attention_mask=attn_mask)
conditionings.append(cond)
@@ -0,0 +1,334 @@
"""Two-stage lip-dubbing pipeline with IC-LoRA and appended audio reference conditioning."""
from __future__ import annotations
import logging
from collections.abc import Iterator
import torch
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.patchifiers import AudioPatchifier
from ltx_core.conditioning import AudioConditionByReferenceLatent
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.video_vae import TilingConfig, VideoEncoder, get_video_chunks_number
from ltx_core.quantization import QuantizationPolicy
from ltx_core.types import Audio, AudioLatentShape, SpatioTemporalScaleFactors, VideoPixelShape
from ltx_pipelines.iclora_utils import (
append_ic_lora_reference_video_conditionings,
read_lora_reference_downscale_factor,
)
from ltx_pipelines.utils.args import (
ImageConditioningInput,
detect_checkpoint_path,
lipdub_arg_parser,
)
from ltx_pipelines.utils.blocks import (
AudioConditioner,
AudioDecoder,
DiffusionStage,
ImageConditioner,
PromptEncoder,
VideoDecoder,
VideoUpsampler,
)
from ltx_pipelines.utils.constants import DISTILLED_SIGMAS, STAGE_2_DISTILLED_SIGMAS, detect_params
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_audio_from_file, encode_video, get_videostream_metadata
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
def _snap_frames_to_8k1(frames: int) -> int:
"""Round ``frames`` down to the nearest ``8k+1`` (the model's required frame count)."""
time_scale = SpatioTemporalScaleFactors.default().time
return ((frames - 1) // time_scale) * time_scale + 1
class LipDubPipeline:
"""Two-stage lip-dubbing with IC-LoRA video reference and appended audio reference tokens."""
def __init__(
self,
distilled_checkpoint_path: str,
spatial_upsampler_path: str,
gemma_root: str,
ic_lora: LoraPathStrengthAndSDOps,
device: torch.device | None = None,
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
offload_mode: OffloadMode = OffloadMode.NONE,
) -> None:
self.device = device or get_device()
self.dtype = torch.bfloat16
self.ic_lora = ic_lora
loras = (ic_lora,)
self.prompt_encoder = PromptEncoder(
distilled_checkpoint_path,
gemma_root,
self.dtype,
self.device,
registry=registry,
offload_mode=offload_mode,
)
self.image_conditioner = ImageConditioner(distilled_checkpoint_path, self.dtype, self.device, registry=registry)
self.audio_conditioner = AudioConditioner(
distilled_checkpoint_path,
self.dtype,
self.device,
registry=registry,
)
self.stage = DiffusionStage(
distilled_checkpoint_path,
self.dtype,
self.device,
loras=loras,
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
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)
self.reference_downscale_factor = read_lora_reference_downscale_factor(ic_lora.path)
def _create_stage_conditionings(
self,
images: list[ImageConditioningInput],
reference_video_path: str,
reference_strength: float,
height: int,
width: int,
num_frames: int,
video_encoder: VideoEncoder,
encode_tiling: TilingConfig | None,
) -> list:
conditionings = combined_image_conditionings(
images=images,
height=height,
width=width,
video_encoder=video_encoder,
dtype=self.dtype,
device=self.device,
)
append_ic_lora_reference_video_conditionings(
conditionings,
[(reference_video_path, reference_strength)],
height=height,
width=width,
num_frames=num_frames,
video_encoder=video_encoder,
dtype=self.dtype,
device=self.device,
reference_downscale_factor=self.reference_downscale_factor,
conditioning_attention_strength=1.0,
conditioning_attention_mask=None,
tiling_config=encode_tiling,
)
return conditionings
def _encode_reference_audio_vae_latent(self, video_path: str) -> torch.Tensor:
audio = decode_audio_from_file(video_path, self.device)
if audio is None:
msg = f"No audio stream found in {video_path}"
raise ValueError(msg)
return self.audio_conditioner(lambda enc: vae_encode_audio(audio, enc, None))
@torch.inference_mode()
def __call__( # noqa: PLR0913
self,
prompt: str,
seed: int,
height: int,
width: int,
images: list[ImageConditioningInput],
reference_video_path: str,
reference_strength: float = 1.0,
enhance_prompt: bool = False,
tiling_config: TilingConfig | None = None,
stage_1_sigmas: torch.Tensor = DISTILLED_SIGMAS,
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
) -> tuple[Iterator[torch.Tensor], Audio]:
assert_resolution(height=height, width=width, is_two_stage=True)
meta = get_videostream_metadata(reference_video_path)
num_frames = _snap_frames_to_8k1(meta.frames)
frame_rate = float(meta.fps)
generator = torch.Generator(device=self.device).manual_seed(seed)
noiser = GaussianNoiser(generator=generator)
(ctx_p,) = self.prompt_encoder(
[prompt],
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
)
video_context, audio_context = ctx_p.video_encoding, ctx_p.audio_encoding
stage_1_output_shape = VideoPixelShape(
batch=1,
frames=num_frames,
width=width // 2,
height=height // 2,
fps=frame_rate,
)
encode_tiling = TilingConfig.default()
def build_image_conditionings(output_shape: VideoPixelShape) -> list:
return self.image_conditioner(
lambda enc: self._create_stage_conditionings(
images=images,
reference_video_path=reference_video_path,
reference_strength=reference_strength,
height=output_shape.height,
width=output_shape.width,
num_frames=num_frames,
video_encoder=enc,
encode_tiling=encode_tiling,
)
)
def build_audio_ref_conditioning(audio_latent: torch.Tensor) -> AudioConditionByReferenceLatent:
ref_patch, ref_pos = patchify_lipdub_audio_reference_latent(
audio_latent,
negative_positions=True,
device=self.device,
)
return AudioConditionByReferenceLatent(ref_patch, ref_pos, strength=1.0)
stage_1_conditionings = build_image_conditionings(stage_1_output_shape)
ref_vae = self._encode_reference_audio_vae_latent(reference_video_path)
audio_conditionings = [build_audio_ref_conditioning(ref_vae)]
stage_1_sigmas_tensor = stage_1_sigmas.to(dtype=torch.float32, device=self.device)
video_state, audio_state = self.stage(
denoiser=SimpleDenoiser(video_context, audio_context),
sigmas=stage_1_sigmas_tensor,
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,
conditionings=audio_conditionings,
),
)
s1_audio_latent = audio_state.latent.clone()
upscaled_video_latent = self.upsampler(video_state.latent[:1])
stage_2_sigmas_tensor = stage_2_sigmas.to(dtype=torch.float32, device=self.device)
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = build_image_conditionings(stage_2_output_shape)
stage_2_audio_conditionings = [build_audio_ref_conditioning(s1_audio_latent)]
video_state, _audio_unused = self.stage(
denoiser=SimpleDenoiser(video_context, audio_context),
sigmas=stage_2_sigmas_tensor,
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_tensor[0].item(),
initial_latent=upscaled_video_latent,
),
audio=ModalitySpec(
context=audio_context,
conditionings=stage_2_audio_conditionings,
frozen=True,
noise_scale=0.0,
initial_latent=s1_audio_latent,
),
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
decoded_audio = self.audio_decoder(s1_audio_latent)
return decoded_video, decoded_audio
def patchify_lipdub_audio_reference_latent(
vae_latents: torch.Tensor,
*,
negative_positions: bool,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Patchify audio VAE latents and build RoPE positions (optional negative shift for reference)."""
patchifier = AudioPatchifier(patch_size=1)
patchified = patchifier.patchify(vae_latents)
b, c, _t, mel_bins = vae_latents.shape
seq_len = patchified.shape[1]
latent_coords = patchifier.get_patch_grid_bounds(
output_shape=AudioLatentShape(batch=b, channels=c, frames=seq_len, mel_bins=mel_bins),
device=device,
)
positions = latent_coords.to(dtype=torch.float32)
if negative_positions:
aud_dur = positions[:, :, -1, 1].max().item()
positions = positions - aud_dur - 0.04
return patchified, positions
@torch.inference_mode()
def main() -> None:
logging.getLogger().setLevel(logging.INFO)
checkpoint_path = detect_checkpoint_path(distilled=True)
params = detect_params(checkpoint_path)
parser = lipdub_arg_parser(params=params)
args = parser.parse_args()
if not args.lora or len(args.lora) != 1:
raise ValueError("LipDub requires exactly one --lora (the lip-dub IC-LoRA).")
pipeline = LipDubPipeline(
distilled_checkpoint_path=args.distilled_checkpoint_path,
spatial_upsampler_path=args.spatial_upsampler_path,
gemma_root=args.gemma_root,
ic_lora=args.lora[0],
quantization=args.quantization,
torch_compile=args.compile,
offload_mode=args.offload_mode,
)
tiling_config = TilingConfig.default()
src = get_videostream_metadata(args.reference_video)
video_chunks_number = get_video_chunks_number(_snap_frames_to_8k1(src.frames), tiling_config)
video, audio = pipeline(
prompt=args.prompt,
seed=args.seed,
height=args.height,
width=args.width,
images=[],
reference_video_path=args.reference_video,
reference_strength=args.reference_strength,
tiling_config=tiling_config,
enhance_prompt=args.enhance_prompt,
)
encode_video(
video=video,
fps=int(src.fps),
audio=audio,
output_path=args.output_path,
video_chunks_number=video_chunks_number,
)
if __name__ == "__main__":
main()
@@ -307,7 +307,7 @@ def main() -> None:
gemma_root=args.gemma_root,
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
distilled=args.distilled,
distilled=True,
torch_compile=args.compile,
offload_mode=args.offload_mode,
)
@@ -16,15 +16,17 @@ from ltx_pipelines.utils.helpers import (
image_conditionings_by_adding_guiding_latent,
)
from ltx_pipelines.utils.samplers import (
euler_cfg_pp_denoising_loop,
euler_denoising_loop,
gradient_estimating_euler_denoising_loop,
res2s_audio_video_denoising_loop,
)
from ltx_pipelines.utils.types import Denoiser, ModalitySpec
from ltx_pipelines.utils.types import DenoisedLatentResult, Denoiser, ModalitySpec
__all__ = [
"AudioConditioner",
"AudioDecoder",
"DenoisedLatentResult",
"Denoiser",
"DiffusionStage",
"FactoryGuidedDenoiser",
@@ -38,6 +40,7 @@ __all__ = [
"assert_resolution",
"cleanup_memory",
"combined_image_conditionings",
"euler_cfg_pp_denoising_loop",
"euler_denoising_loop",
"get_device",
"gradient_estimating_euler_denoising_loop",
@@ -1,4 +1,5 @@
import argparse
from collections.abc import Sequence
from pathlib import Path
from typing import NamedTuple
@@ -115,35 +116,34 @@ def resolve_path(path: str) -> str:
QUANTIZATION_POLICIES = ("fp8-cast", "fp8-scaled-mm")
class QuantizationAction(argparse.Action):
def __call__(
self,
parser: argparse.ArgumentParser, # noqa: ARG002
namespace: argparse.Namespace,
values: list[str],
option_string: str | None = None,
) -> None:
if len(values) > 2:
msg = (
f"{option_string} accepts at most 2 arguments (POLICY and optional AMAX_PATH), got {len(values)} values"
def _resolve_quantization(namespace: argparse.Namespace) -> None:
# Resolution is deferred until after parse_args because fp8-scaled-mm needs the
# checkpoint path, which isn't on the namespace when the --quantization argument
# is parsed.
name = getattr(namespace, "quantization", None)
if name is None or isinstance(name, QuantizationPolicy):
return
if name == "fp8-cast":
namespace.quantization = QuantizationPolicy.fp8_cast()
return
if name == "fp8-scaled-mm":
ckpt = getattr(namespace, "checkpoint_path", None) or getattr(namespace, "distilled_checkpoint_path", None)
if ckpt is None:
raise SystemExit(
"--quantization fp8-scaled-mm requires --checkpoint-path (or --distilled-checkpoint-path)."
)
raise argparse.ArgumentError(self, msg)
namespace.quantization = QuantizationPolicy.fp8_scaled_mm(ckpt)
policy_name = values[0]
if policy_name not in QUANTIZATION_POLICIES:
msg = f"Unknown quantization policy '{policy_name}'. Choose from: {', '.join(QUANTIZATION_POLICIES)}"
raise argparse.ArgumentError(self, msg)
if policy_name == "fp8-cast":
if len(values) > 1:
msg = f"{option_string} fp8-cast does not accept additional arguments"
raise argparse.ArgumentError(self, msg)
policy = QuantizationPolicy.fp8_cast()
elif policy_name == "fp8-scaled-mm":
amax_path = resolve_path(values[1]) if len(values) > 1 else None
policy = QuantizationPolicy.fp8_scaled_mm(amax_path)
setattr(namespace, self.dest, policy)
class _PipelineArgumentParser(argparse.ArgumentParser):
def parse_args( # type: ignore[override]
self,
args: Sequence[str] | None = None,
namespace: argparse.Namespace | None = None,
) -> argparse.Namespace:
ns = super().parse_args(args, namespace)
_resolve_quantization(ns)
return ns
def detect_checkpoint_path(distilled: bool = False) -> str:
@@ -159,7 +159,7 @@ def basic_arg_parser(
params: PipelineParams = LTX_2_3_PARAMS,
distilled: bool = False,
) -> argparse.ArgumentParser:
parser = argparse.ArgumentParser()
parser = _PipelineArgumentParser()
if distilled:
parser.add_argument(
"--distilled-checkpoint-path",
@@ -264,16 +264,14 @@ def basic_arg_parser(
parser.add_argument(
"--quantization",
dest="quantization",
action=QuantizationAction,
nargs="+",
metavar=("POLICY", "AMAX_PATH"),
choices=QUANTIZATION_POLICIES,
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"
"fp8-scaled-mm uses FP8 scaled matrix multiplication; the layer set is auto-discovered "
"from the checkpoint's .weight_scale tensors. "
"Example: --quantization fp8-cast or --quantization fp8-scaled-mm"
),
)
parser.add_argument(
@@ -348,6 +346,53 @@ def video_editing_arg_parser(
return parser
def lipdub_arg_parser(
params: PipelineParams = LTX_2_3_PARAMS,
) -> argparse.ArgumentParser:
"""Argument parser for the lip-dub pipeline.
Frame count and frame rate are derived from the reference video at runtime (the frame count
is silently snapped down to the nearest 8k+1), so this parser intentionally omits
--num-frames, --frame-rate, and --image. Distilled checkpoint only.
"""
parser = basic_arg_parser(params=params, distilled=True)
parser.add_argument(
"--height",
type=int,
default=params.stage_2_height,
help=(
f"Height of the generated video in pixels, should be divisible by 64 (default: {params.stage_2_height})."
),
)
parser.add_argument(
"--width",
type=int,
default=params.stage_2_width,
help=f"Width of the generated video in pixels, should be divisible by 64 (default: {params.stage_2_width}).",
)
parser.add_argument(
"--spatial-upsampler-path",
type=resolve_path,
required=True,
help=(
"Path to the spatial upsampler model used to increase the resolution "
"of the generated video in the latent space."
),
)
parser.add_argument(
"--reference-video",
type=resolve_path,
required=True,
help="Reference video file (video + audio track used for IC-LoRA and audio identity).",
)
parser.add_argument(
"--reference-strength",
type=float,
default=1.0,
help="Strength for IC-LoRA video reference conditioning (default: 1.0).",
)
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
@@ -21,7 +21,8 @@ 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.loader import SDOps
from ltx_core.loader.primitives import LoraPathStrengthAndSDOps
from ltx_core.loader.module_ops import ModuleOps
from ltx_core.loader.primitives import BuilderProtocol, LoraPathStrengthAndSDOps, ModelBuilderProtocol
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 (
@@ -37,12 +38,14 @@ from ltx_core.model.audio_vae import (
)
from ltx_core.model.transformer import (
LTXV_MODEL_COMFY_RENAMING_MAP,
LTXModel,
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 (
MEMORY_EFFICIENT_DECODE,
VAE_DECODER_COMFY_KEYS_FILTER,
VAE_ENCODER_COMFY_KEYS_FILTER,
TilingConfig,
@@ -59,10 +62,11 @@ from ltx_core.text_encoders.gemma import (
GemmaTextEncoderConfigurator,
module_ops_from_gemma_root,
)
from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessorOutput
from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessor, 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.multigpu.delegating_builder import DelegatingBuilder
from ltx_pipelines.utils.gpu_model import gpu_model
from ltx_pipelines.utils.helpers import (
cleanup_memory,
@@ -83,6 +87,20 @@ _M = TypeVar("_M", bound=torch.nn.Module)
# ---------------------------------------------------------------------------
def _chain_quantization(
sd_ops: SDOps,
module_ops: tuple[ModuleOps, ...],
quantization: QuantizationPolicy,
) -> tuple[SDOps, tuple[ModuleOps, ...]]:
chained_sd_ops = sd_ops
if quantization.sd_ops is not None:
chained_sd_ops = SDOps(
name=f"sd_ops_chain_{sd_ops.name}+{quantization.sd_ops.name}",
mapping=(*sd_ops.mapping, *quantization.sd_ops.mapping),
)
return chained_sd_ops, (*module_ops, *quantization.module_ops)
@contextmanager
def _streaming_model(
builder: StreamingModelBuilder,
@@ -154,16 +172,43 @@ class DiffusionStage:
registry: Registry | None = None,
torch_compile: bool = False,
offload_mode: OffloadMode = OffloadMode.NONE,
transformer_builder: ModelBuilderProtocol[LTXModel] | DelegatingBuilder[LTXModel] | None = None,
) -> None:
self._dtype = dtype
self._device = device
self._quantization = quantization
self._torch_compile = torch_compile
self._offload_mode = offload_mode
if transformer_builder is not None:
self._transformer_builder = transformer_builder
else:
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(),
)
if offload_mode != OffloadMode.NONE:
if torch_compile:
raise ValueError("torch.compile is not supported with layer streaming")
streaming_sd_ops: SDOps = LTXV_MODEL_COMFY_RENAMING_MAP
streaming_module_ops: tuple[ModuleOps, ...] = ()
if quantization is not None:
raise ValueError("quantization is not supported with layer streaming")
if quantization.kind != QuantizationPolicy.Kind.FP8_CAST:
raise ValueError(
f"Layer streaming supports only QuantizationPolicy.fp8_cast(); "
f"got kind={quantization.kind!r} which produces heterogeneous block layouts."
)
streaming_sd_ops, streaming_module_ops = _chain_quantization(
streaming_sd_ops, streaming_module_ops, quantization
)
self._streaming_builder = StreamingModelBuilder(
model_class_configurator=LTXModelConfigurator,
model_path=checkpoint_path,
model_sd_ops=LTXV_MODEL_COMFY_RENAMING_MAP,
model_sd_ops=streaming_sd_ops,
module_ops=streaming_module_ops,
loras=tuple(loras),
registry=registry or DummyRegistry(),
blocks_attr="velocity_model.transformer_blocks",
@@ -172,19 +217,6 @@ class DiffusionStage:
model_wrapper=lambda m: X0Model(m).eval(),
)
self._dtype = dtype
self._device = device
self._quantization = quantization
self._torch_compile = torch_compile
self._offload_mode = offload_mode
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
@@ -198,18 +230,12 @@ class DiffusionStage:
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
),
modify_sd_ops_for_compilation(lora.sd_ops, 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),
)
sd_ops, module_ops = _chain_quantization(sd_ops, module_ops, self._quantization)
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()
@@ -359,31 +385,40 @@ class PromptEncoder:
device: torch.device,
registry: Registry | None = None,
offload_mode: OffloadMode = OffloadMode.NONE,
text_encoder_builder: BuilderProtocol | None = None,
) -> None:
self._dtype = dtype
self._device = device
self._offload_mode = offload_mode
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._streaming_text_encoder_builder = StreamingModelBuilder(
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(),
blocks_attr="model.model.language_model.layers",
blocks_prefix="model.model.language_model.layers",
)
if text_encoder_builder is not None:
if offload_mode != OffloadMode.NONE:
raise ValueError(
"text_encoder_builder cannot be used with offload_mode != OffloadMode.NONE "
"because no streaming text encoder builder is available."
)
self._text_encoder_builder = text_encoder_builder
self._streaming_text_encoder_builder = None
else:
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._streaming_text_encoder_builder = StreamingModelBuilder(
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(),
blocks_attr="model.model.language_model.layers",
blocks_prefix="model.model.language_model.layers",
)
self._embeddings_processor_builder = Builder(
model_path=checkpoint_path,
model_class_configurator=EmbeddingsProcessorConfigurator,
@@ -391,10 +426,18 @@ class PromptEncoder:
registry=registry or DummyRegistry(),
)
def _build_text_encoder(self) -> torch.nn.Module:
"""Build the Gemma text encoder (non-streaming path)."""
return self._text_encoder_builder.build(device=self._device, dtype=self._dtype).eval()
def _build_embeddings_processor(self) -> EmbeddingsProcessor:
"""Build the embeddings processor on the target device."""
return self._embeddings_processor_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
def _text_encoder_ctx(self) -> AbstractContextManager:
if self._offload_mode != OffloadMode.NONE:
return _streaming_model(self._streaming_text_encoder_builder, self._offload_mode, self._device, self._dtype)
return gpu_model(self._text_encoder_builder.build(device=self._device, dtype=self._dtype).eval())
return gpu_model(self._build_text_encoder())
def __call__(
self,
@@ -413,9 +456,7 @@ class PromptEncoder:
)
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:
with gpu_model(self._build_embeddings_processor()) as embeddings_processor:
return [embeddings_processor.process_hidden_states(hs, mask) for hs, mask in raw_outputs]
@@ -513,32 +554,31 @@ class VideoDecoder:
dtype: torch.dtype,
device: torch.device,
registry: Registry | None = None,
memory_efficient: bool = True,
decoder_builder: BuilderProtocol | 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(),
)
if decoder_builder is not None:
self._decoder_builder = decoder_builder
else:
self._decoder_builder = Builder(
model_path=checkpoint_path,
model_class_configurator=VideoDecoderConfigurator,
model_sd_ops=VAE_DECODER_COMFY_KEYS_FILTER,
registry=registry or DummyRegistry(),
module_ops=(MEMORY_EFFICIENT_DECODE,) if memory_efficient else (),
)
def __call__(
self,
latent: torch.Tensor,
tiling_config: TilingConfig | None = None,
generator: torch.Generator | None = None,
*,
output_dtype: torch.dtype = torch.uint8,
) -> Iterator[torch.Tensor]:
"""Decode *latent* to pixel-space video chunks. Decoder freed after exhaustion.
Args:
output_dtype: Target dtype for output tensors. ``torch.uint8``
(default) maps to ``[0, 255]``. Any floating dtype returns
``[0, 1]`` cast to that dtype.
"""
"""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, output_dtype=output_dtype), decoder)
return _cleanup_iter(decoder.decode_video(latent, tiling_config, generator), decoder)
# ---------------------------------------------------------------------------
@@ -0,0 +1,224 @@
"""Color space conversion utilities for video encoding.
Provides GPU-accelerated RGB to YUV420 conversion that runs between the
VAE decoder (which yields float RGB chunks) and ``encode_video``, bypassing
pyav's CPU-side libswscale conversion. The ``FrameConverter`` also carries
the codec metadata (pixel format, colour space, colour range) that
``encode_video`` needs to tag the output stream.
"""
from __future__ import annotations
import enum
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
class ColorSpace(enum.Enum):
"""YUV color space standard."""
BT_709 = "bt709"
BT_2020_NCL = "bt2020ncl"
@property
def av_colorspace(self) -> int:
"""FFmpeg ``AVCOL_SPC_*`` constant for ``codec_context.colorspace``."""
return _AV_COLORSPACE[self]
class ColorRange(enum.Enum):
"""YUV color range."""
MPEG = "mpeg"
JPEG = "jpeg"
@property
def av_color_range(self) -> int:
"""FFmpeg ``AVCOL_RANGE_*`` constant for ``codec_context.color_range``."""
return _AV_COLOR_RANGE[self]
class PixelFormat(enum.Enum):
"""Pixel format for video frames."""
RGB24 = "rgb24"
YUV420P = "yuv420p"
@property
def av_format(self) -> str:
"""PyAV format string for ``VideoFrame.from_ndarray``."""
return self.value
_AV_COLORSPACE = {
ColorSpace.BT_709: 1, # AVCOL_SPC_BT709
ColorSpace.BT_2020_NCL: 9, # AVCOL_SPC_BT2020_NCL
}
_AV_COLOR_RANGE = {
ColorRange.MPEG: 1, # AVCOL_RANGE_MPEG (limited)
ColorRange.JPEG: 2, # AVCOL_RANGE_JPEG (full)
}
# BT.709 RGB->YUV matrix (row-major: each row produces one of Y, U, V)
_BT709_MATRIX = torch.tensor(
[
[0.2126, 0.7152, 0.0722],
[-0.1146, -0.3854, 0.5],
[0.5, -0.4542, -0.0458],
],
dtype=torch.float32,
)
# BT.2020 NCL RGB->YUV matrix
_KR_2020 = 0.2627
_KG_2020 = 0.6780
_KB_2020 = 0.0593
_BT2020_MATRIX = torch.tensor(
[
[_KR_2020, _KG_2020, _KB_2020],
[-_KR_2020 / 1.8814, -_KG_2020 / 1.8814, 0.5],
[0.5, -_KG_2020 / 1.4746, -_KB_2020 / 1.4746],
],
dtype=torch.float32,
)
_COLOR_SPACE_MATRICES = {
ColorSpace.BT_709: _BT709_MATRIX,
ColorSpace.BT_2020_NCL: _BT2020_MATRIX,
}
@dataclass(frozen=True)
class FrameConverter:
"""Converts ``[*, C, H, W]`` float ``[0, 1]`` frames to uint8.
Carries encoding metadata so ``encode_video`` can derive pixel format,
color space, and color range from the converter itself.
The ``fn_`` callable **may mutate its input** (PyTorch trailing-underscore
convention). Callers that need to keep the original ``frames`` afterwards
must pass ``frames.clone()``. Inside ``encode_video``'s per-chunk
generator each chunk is consumed once, so direct passthrough is safe.
"""
pixel_format: PixelFormat
fn_: Callable[[torch.Tensor], torch.Tensor] = field(repr=False)
color_space: ColorSpace | None = None
color_range: ColorRange | None = None
def __call__(self, frames: torch.Tensor) -> torch.Tensor:
return self.fn_(frames)
def rgb_to_yuv(image: torch.Tensor, color_space: ColorSpace) -> torch.Tensor:
"""Convert an RGB image to YUV.
The image data is assumed to be in the range of ``[0, 1]``.
Uses a single matrix multiply for better memory locality.
Args:
image: RGB image with shape ``(*, 3, H, W)``.
color_space: Color space standard for the conversion matrix.
Returns:
YUV image with shape ``(*, 3, H, W)``.
"""
if len(image.shape) < 3 or image.shape[-3] != 3:
raise ValueError(f"Input size must have a shape of (*, 3, H, W). Got {image.shape}")
mat = _COLOR_SPACE_MATRICES[color_space].to(device=image.device, dtype=image.dtype)
# [*, 3, H, W] -> [*, H, W, 3] @ [3, 3]^T -> [*, H, W, 3] -> [*, 3, H, W]
pixels = image.movedim(-3, -1) # [*, H, W, 3]
yuv = pixels @ mat.T # [*, H, W, 3]
return yuv.movedim(-1, -3) # [*, 3, H, W]
def apply_color_range_(y: torch.Tensor, uv: torch.Tensor, color_range: ColorRange) -> tuple[torch.Tensor, torch.Tensor]:
"""Scale Y and UV planes to the specified color range, in-place.
Args:
y: Luma plane in ``[0, 1]``.
uv: Chroma planes centered at 0.
color_range: Target color range.
Returns:
Scaled ``(Y, UV)`` tensors (modified in-place).
"""
if color_range == ColorRange.MPEG:
y.mul_(219).add_(16)
uv.mul_(224).add_(128)
elif color_range == ColorRange.JPEG:
y.mul_(255)
uv.add_(0.5).mul_(255)
else:
raise ValueError(f"Unsupported color range: {color_range}")
return y, uv
def rgb_to_yuv420(
image: torch.Tensor, color_space: ColorSpace, color_range: ColorRange
) -> tuple[torch.Tensor, torch.Tensor]:
"""Convert an RGB image to YUV 4:2:0 with chroma subsampling.
Chroma is subsampled by averaging 2x2 pixel blocks (chroma siting
``(128, 128)``).
Args:
image: RGB image with shape ``(*, 3, H, W)`` in ``[0, 1]``.
H and W must be divisible by 2.
color_space: Color space standard.
color_range: Color range for the output.
Returns:
``(Y, UV)`` where Y has shape ``(*, 1, H, W)`` and UV has shape
``(*, 2, H//2, W//2)``.
"""
if len(image.shape) < 3 or image.shape[-3] != 3:
raise ValueError(f"Input size must have a shape of (*, 3, H, W). Got {image.shape}")
if image.shape[-2] % 2 != 0 or image.shape[-1] % 2 != 0:
raise ValueError(f"Input H and W must be divisible by 2. Got {image.shape}")
yuv = rgb_to_yuv(image, color_space)
y = yuv[..., :1, :, :]
# Subsample chroma: average 2x2 blocks via avg_pool2d (contiguous, fused kernel)
uv_full = yuv[..., 1:3, :, :].contiguous()
# Flatten leading dims for avg_pool2d which expects [N, C, H, W]
lead = uv_full.shape[:-3]
uv_flat = uv_full.reshape(-1, 2, uv_full.shape[-2], uv_full.shape[-1])
uv = torch.nn.functional.avg_pool2d(uv_flat, kernel_size=2, stride=2)
uv = uv.reshape(*lead, 2, uv.shape[-2], uv.shape[-1])
return apply_color_range_(y, uv, color_range)
def pack_i420(y: torch.Tensor, uv: torch.Tensor) -> torch.Tensor:
"""Pack Y and UV planes into I420 layout for pyav.
I420 packs the three planes into a single 2D array of height ``H * 3 // 2``
and width ``W``. The Y plane occupies the first ``H`` rows. The UV tensor
``(*, 2, H//2, W//2)`` is reshaped to ``(*, H//2, W)`` -- U rows packed
two-by-two followed by V rows packed two-by-two -- and appended below.
Args:
y: Luma with shape ``(*, 1, H, W)``.
uv: Chroma with shape ``(*, 2, H//2, W//2)``.
Returns:
Packed tensor with shape ``(*, H*3//2, W)`` uint8.
"""
y_plane = y[..., 0, :, :] # [*, H, W]
uv_packed = uv.reshape(*uv.shape[:-3], uv.shape[-2], uv.shape[-1] * 2) # [*, H//2, W]
packed = torch.cat([y_plane, uv_packed], dim=-2) # [*, H*3//2, W]
return packed.clamp_(0, 255).to(torch.uint8)
def _rgb_uint8_fn_(frames: torch.Tensor) -> torch.Tensor:
"""In-place: mutates ``frames`` via ``clamp_`` + ``mul_``, returns a uint8 view."""
return frames.clamp_(0.0, 1.0).mul_(255.0).to(torch.uint8).movedim(-3, -1)
rgb_uint8_converter_ = FrameConverter(pixel_format=PixelFormat.RGB24, fn_=_rgb_uint8_fn_)
"""``(*, 3, H, W)`` float ``[0, 1]`` to ``(*, H, W, 3)`` uint8. Mutates input."""
def _yuv420p_bt709_fn_(frames: torch.Tensor) -> torch.Tensor:
y, uv = rgb_to_yuv420(frames, ColorSpace.BT_709, ColorRange.MPEG)
return pack_i420(y, uv)
yuv420p_bt709_converter_ = FrameConverter(
pixel_format=PixelFormat.YUV420P,
fn_=_yuv420p_bt709_fn_,
color_space=ColorSpace.BT_709,
color_range=ColorRange.MPEG,
)
"""``(*, 3, H, W)`` float ``[0, 1]`` to ``(*, H*3//2, W)`` uint8 YUV420p BT.709 MPEG."""
@@ -20,6 +20,7 @@ from ltx_core.guidance.perturbations import (
from ltx_core.model.transformer import X0Model
from ltx_core.types import LatentState
from ltx_pipelines.utils.helpers import modality_from_latent_state
from ltx_pipelines.utils.types import DenoisedLatentResult
_POSITIVE_ONLY_GUIDER = MultiModalGuider(
params=MultiModalGuiderParams(cfg_scale=1.0, stg_scale=0.0, modality_scale=1.0),
@@ -53,7 +54,7 @@ def _repeat_state(state: LatentState, n: int) -> LatentState:
)
def _guided_denoise( # noqa: PLR0913
def _guided_denoise( # noqa: PLR0913,PLR0915
transformer: X0Model,
video_state: LatentState | None,
audio_state: LatentState | None,
@@ -66,7 +67,8 @@ def _guided_denoise( # noqa: PLR0913
last_denoised_video: torch.Tensor | None,
last_denoised_audio: torch.Tensor | None,
step_index: int,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
force_uncond_pass: bool = False,
) -> tuple[DenoisedLatentResult | None, DenoisedLatentResult | 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
@@ -80,7 +82,9 @@ def _guided_denoise( # noqa: PLR0913
a_skip = audio_guider.should_skip_step(step_index)
if v_skip and a_skip:
return last_denoised_video, last_denoised_audio
video_result = DenoisedLatentResult.result_or_none(denoised=last_denoised_video)
audio_result = DenoisedLatentResult.result_or_none(denoised=last_denoised_audio)
return video_result, audio_result
if video_state is not None and v_context is None:
raise ValueError("v_context is required when video_state is provided")
@@ -91,10 +95,12 @@ def _guided_denoise( # noqa: PLR0913
_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:
v_needs_neg = video_guider.do_unconditional_generation() or (force_uncond_pass and video_state is not None)
a_needs_neg = audio_guider.do_unconditional_generation() or (force_uncond_pass and audio_state is not None)
if v_needs_neg or a_needs_neg:
if v_needs_neg 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:
if a_needs_neg 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
@@ -172,7 +178,14 @@ def _guided_denoise( # noqa: PLR0913
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
return (
DenoisedLatentResult.result_or_none(
denoised=denoised_video, uncond=uncond_v, cond=cond_v, ptb=ptb_v, mod=mod_v
),
DenoisedLatentResult.result_or_none(
denoised=denoised_audio, uncond=uncond_a, cond=cond_a, ptb=ptb_a, mod=mod_a
),
)
class SimpleDenoiser:
@@ -195,11 +208,15 @@ class SimpleDenoiser:
audio_state: LatentState | None,
sigmas: torch.Tensor,
step_index: int,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
) -> tuple[DenoisedLatentResult | None, DenoisedLatentResult | 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)
denoised_video, denoised_audio = transformer(video=pos_video, audio=pos_audio, perturbations=None)
return (
DenoisedLatentResult.result_or_none(denoised=denoised_video),
DenoisedLatentResult.result_or_none(denoised=denoised_audio),
)
class GuidedDenoiser:
@@ -214,11 +231,13 @@ class GuidedDenoiser:
a_context: torch.Tensor | None,
video_guider: MultiModalGuider | None = None,
audio_guider: MultiModalGuider | None = None,
force_uncond_pass: bool = False,
) -> None:
self.v_context = v_context
self.a_context = a_context
self.video_guider = video_guider
self.audio_guider = audio_guider
self.force_uncond_pass = force_uncond_pass
self._last_denoised_video: torch.Tensor | None = None
self._last_denoised_audio: torch.Tensor | None = None
@@ -229,8 +248,8 @@ class GuidedDenoiser:
audio_state: LatentState | None,
sigmas: torch.Tensor,
step_index: int,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
denoised_video, denoised_audio = _guided_denoise(
) -> tuple[DenoisedLatentResult | None, DenoisedLatentResult | None]:
guided_denoise_result_v, guided_denoise_result_a = _guided_denoise(
transformer=transformer,
video_state=video_state,
audio_state=audio_state,
@@ -242,10 +261,11 @@ class GuidedDenoiser:
last_denoised_video=self._last_denoised_video,
last_denoised_audio=self._last_denoised_audio,
step_index=step_index,
force_uncond_pass=self.force_uncond_pass,
)
self._last_denoised_video = denoised_video
self._last_denoised_audio = denoised_audio
return denoised_video, denoised_audio
self._last_denoised_video = guided_denoise_result_v.denoised
self._last_denoised_audio = guided_denoise_result_a.denoised
return guided_denoise_result_v, guided_denoise_result_a
class FactoryGuidedDenoiser:
@@ -257,11 +277,13 @@ class FactoryGuidedDenoiser:
a_context: torch.Tensor | None,
video_guider_factory: MultiModalGuiderFactory | None = None,
audio_guider_factory: MultiModalGuiderFactory | None = None,
force_uncond_pass: bool = False,
) -> 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.force_uncond_pass = force_uncond_pass
self._last_denoised_video: torch.Tensor | None = None
self._last_denoised_audio: torch.Tensor | None = None
self._sigma_vals_cached: list[float] | None = None
@@ -273,7 +295,7 @@ class FactoryGuidedDenoiser:
audio_state: LatentState | None,
sigmas: torch.Tensor,
step_index: int,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
) -> tuple[DenoisedLatentResult | None, DenoisedLatentResult | None]:
if self._sigma_vals_cached is None:
self._sigma_vals_cached = sigmas.detach().cpu().tolist()
sigma_val = self._sigma_vals_cached[step_index]
@@ -287,7 +309,7 @@ class FactoryGuidedDenoiser:
else None
)
denoised_video, denoised_audio = _guided_denoise(
guided_denoise_result_v, guided_denoise_result_a = _guided_denoise(
transformer=transformer,
video_state=video_state,
audio_state=audio_state,
@@ -299,7 +321,8 @@ class FactoryGuidedDenoiser:
last_denoised_video=self._last_denoised_video,
last_denoised_audio=self._last_denoised_audio,
step_index=step_index,
force_uncond_pass=self.force_uncond_pass,
)
self._last_denoised_video = denoised_video
self._last_denoised_audio = denoised_audio
return denoised_video, denoised_audio
self._last_denoised_video = guided_denoise_result_v.denoised
self._last_denoised_audio = guided_denoise_result_a.denoised
return guided_denoise_result_v, guided_denoise_result_a
@@ -1,10 +1,12 @@
import enum
import logging
import math
import threading
from collections.abc import Generator, Iterator
from fractions import Fraction
from io import BytesIO
from pathlib import Path
from queue import Queue
import av
import numpy as np
@@ -17,6 +19,7 @@ from tqdm import tqdm
from ltx_core.hdr import LogC3
from ltx_core.types import Audio, VideoPixelShape
from ltx_pipelines.utils.color_conversion import FrameConverter, PixelFormat, yuv420p_bt709_converter_
from ltx_pipelines.utils.constants import DEFAULT_IMAGE_CRF
logger = logging.getLogger(__name__)
@@ -86,8 +89,8 @@ def resize_and_center_crop(tensor: torch.Tensor, height: int, width: int) -> tor
return tensor
def normalize_latent(latent: torch.Tensor, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
return (latent / 127.5 - 1.0).to(device=device, dtype=dtype)
def normalize_images(images: torch.Tensor, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
return (images / 127.5 - 1.0).to(device=device, dtype=dtype)
def to_vae_range(x: torch.Tensor) -> torch.Tensor:
@@ -116,7 +119,7 @@ def load_image_and_preprocess(
image = preprocess(image=image, crf=crf)
image = torch.tensor(image, dtype=torch.float32, device=device)
image = resize_and_center_crop(image, height, width)
image = normalize_latent(image, device, dtype)
image = normalize_images(image, device, dtype)
return image
@@ -137,11 +140,13 @@ def video_preprocess(
Returns:
Tensor of shape (1, C, F, height, width) with values in [-1, 1].
"""
result = None
result: torch.Tensor | None = None
for f in frames:
frame = resize_and_center_crop(f.to(torch.float32), height, width)
frame = normalize_latent(frame, device, dtype)
frame = normalize_images(frame, device, dtype)
result = frame if result is None else torch.cat([result, frame], dim=2)
if result is None:
raise ValueError("video_preprocess received an empty frame generator; no frames were decoded from the source.")
return result
@@ -325,47 +330,120 @@ def encode_video(
audio: Audio | None,
output_path: str,
video_chunks_number: int,
frame_converter: FrameConverter = yuv420p_bt709_converter_,
crf: int = 19,
preset: str = "veryfast",
thread_count: int = 0,
) -> None:
if isinstance(video, torch.Tensor):
video = iter([video])
first_chunk = next(video)
def convert(chunk: torch.Tensor) -> torch.Tensor:
return frame_converter(chunk.movedim(-1, -3))
_, height, width, _ = first_chunk.shape
first_chunk = convert(next(video))
if frame_converter.pixel_format == PixelFormat.RGB24:
height, width = first_chunk.shape[-3], first_chunk.shape[-2]
else:
height = first_chunk.shape[-2] * 2 // 3
width = first_chunk.shape[-1]
container = av.open(output_path, mode="w")
stream = container.add_stream("libx264", rate=int(fps))
stream.width = width
stream.height = height
stream.pix_fmt = "yuv420p"
success = False
try:
stream = container.add_stream("libx264", rate=int(fps), options={"crf": str(crf), "preset": preset})
stream.width = width
stream.height = height
stream.pix_fmt = "yuv420p"
stream.codec_context.thread_count = thread_count
stream.codec_context.thread_type = "FRAME"
if frame_converter.color_space is not None:
stream.codec_context.colorspace = frame_converter.color_space.av_colorspace
if frame_converter.color_range is not None:
stream.codec_context.color_range = frame_converter.color_range.av_color_range
if audio is not None:
audio_stream = _prepare_audio_stream(container, audio.sampling_rate)
if audio is not None:
audio_stream = _prepare_audio_stream(container, audio.sampling_rate)
def all_tiles(
first_chunk: torch.Tensor, tiles_generator: Generator[tuple[torch.Tensor, int], None, None]
) -> Generator[tuple[torch.Tensor, int], None, None]:
yield first_chunk
yield from tiles_generator
av_format = frame_converter.pixel_format.av_format
for video_chunk in tqdm(all_tiles(first_chunk, video), total=video_chunks_number):
video_chunk_cpu = video_chunk.to("cpu").numpy()
for frame_array in video_chunk_cpu:
frame = av.VideoFrame.from_ndarray(frame_array, format="rgb24")
for packet in stream.encode(frame):
container.mux(packet)
def cpu_chunks() -> Generator[np.ndarray, None, None]:
yield first_chunk.to("cpu").numpy()
for chunk in video:
yield convert(chunk).to("cpu").numpy()
# Flush encoder
for packet in stream.encode():
container.mux(packet)
_encode_chunks_threaded(
container=container,
stream=stream,
av_format=av_format,
chunks=cpu_chunks(),
progress_total=video_chunks_number,
)
if audio is not None:
_write_audio(container, audio_stream, audio)
container.close()
if audio is not None:
_write_audio(container, audio_stream, audio)
success = True
finally:
container.close()
if not success:
Path(output_path).unlink(missing_ok=True)
logger.info(f"Video saved to {output_path}")
def _encode_chunks_threaded(
container: av.container.Container,
stream: av.video.stream.VideoStream,
av_format: str,
chunks: Iterator[np.ndarray],
progress_total: int,
) -> None:
"""Run libx264 frame.encode + container.mux on a background thread while
the caller produces numpy chunks on the current thread. The 1-slot queue
lets the producer get one chunk ahead (so the next VAE/gather chunk
overlaps with libx264 encoding the previous chunk) without buffering more
than one chunk in CPU memory.
"""
chunk_queue: Queue[np.ndarray | None] = Queue(maxsize=1)
encoder_error: list[BaseException] = []
def encoder_worker() -> None:
error: BaseException | None = None
while True:
arr = chunk_queue.get()
if arr is None:
break
if error is not None:
continue
try:
for frame_array in arr:
frame = av.VideoFrame.from_ndarray(frame_array, format=av_format)
for packet in stream.encode(frame):
container.mux(packet)
except Exception as e:
error = e
if error is None:
try:
for packet in stream.encode():
container.mux(packet)
except Exception as e:
error = e
if error is not None:
encoder_error.append(error)
encoder_thread = threading.Thread(target=encoder_worker, name="h264-encoder")
encoder_thread.start()
try:
for arr in tqdm(chunks, total=progress_total):
chunk_queue.put(arr)
finally:
chunk_queue.put(None)
encoder_thread.join()
if encoder_error:
raise encoder_error[0]
_INT_FORMAT_MAX: dict[str, float] = {
"u8": 128.0,
"u8p": 128.0,
@@ -6,7 +6,7 @@ from typing import Callable
import torch
from tqdm import tqdm
from ltx_core.components.diffusion_steps import Res2sDiffusionStep
from ltx_core.components.diffusion_steps import EulerCfgPpDiffusionStep, Res2sDiffusionStep
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.model.transformer import X0Model
from ltx_core.utils import to_denoised, to_velocity
@@ -60,13 +60,15 @@ def euler_denoising_loop(
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)``.
and must return a :class:`~ltx_pipelines.utils.types.DenoisedLatentResult`.
### Returns
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 = denoiser(transformer, video_state, audio_state, sigmas, step_idx)
video_result, audio_result = denoiser(transformer, video_state, audio_state, sigmas, step_idx)
denoised_video = video_result.denoised if video_result is not None else None
denoised_audio = audio_result.denoised if audio_result is not None else None
video_state = _step_state(video_state, denoised_video, stepper, sigmas, step_idx)
audio_state = _step_state(audio_state, denoised_audio, stepper, sigmas, step_idx)
@@ -110,7 +112,9 @@ def gradient_estimating_euler_denoising_loop(
return current_velocity, denoised_sample
for step_idx, _ in enumerate(tqdm(sigmas[:-1])):
denoised_video, denoised_audio = denoiser(transformer, video_state, audio_state, sigmas, step_idx)
video_result, audio_result = denoiser(transformer, video_state, audio_state, sigmas, step_idx)
denoised_video = video_result.denoised if video_result is not None else None
denoised_audio = audio_result.denoised if audio_result is not None else None
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)
@@ -143,6 +147,11 @@ def gradient_estimating_euler_denoising_loop(
return (video_state, audio_state)
def _get_plain_noise(x: torch.Tensor, generator: torch.Generator) -> torch.Tensor:
"""Draw standard Gaussian noise matching the shape, dtype, and device of ``x``."""
return torch.randn(x.shape, generator=generator, dtype=x.dtype, device=x.device)
def _channelwise_normalize(x: torch.Tensor) -> torch.Tensor:
return x.sub_(x.mean(dim=(-2, -1), keepdim=True)).div_(x.std(dim=(-2, -1), keepdim=True))
@@ -278,7 +287,9 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915,PLR0912
# ====================================================================
# STAGE 1: Evaluate at current point
# ====================================================================
denoised_video_1, denoised_audio_1 = denoiser(transformer, video_state, audio_state, sigmas, step_idx)
video_result, audio_result = denoiser(transformer, video_state, audio_state, sigmas, step_idx)
denoised_video_1 = video_result.denoised if video_result is not None else None
denoised_audio_1 = audio_result.denoised if audio_result is not None else None
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:
@@ -355,13 +366,15 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915,PLR0912
else None
)
denoised_video_2, denoised_audio_2 = denoiser(
video_result_2, audio_result_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 = video_result_2.denoised if video_result_2 is not None else None
denoised_audio_2 = audio_result_2.denoised if audio_result_2 is not None else None
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:
@@ -410,7 +423,9 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915,PLR0912
# Final step if we need to fully remove the noise
if sigmas[-1] == 0:
denoised_video_1, denoised_audio_1 = denoiser(transformer, video_state, audio_state, sigmas, n_full_steps)
video_result_final, audio_result_final = denoiser(transformer, video_state, audio_state, sigmas, n_full_steps)
denoised_video_1 = video_result_final.denoised if video_result_final is not None else None
denoised_audio_1 = audio_result_final.denoised if audio_result_final is not None else None
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))
@@ -419,3 +434,121 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915,PLR0912
audio_state = replace(audio_state, latent=denoised_audio_1.to(model_dtype))
return video_state, audio_state
def euler_cfg_pp_denoising_loop(
sigmas: torch.Tensor,
video_state: LatentState | None,
audio_state: LatentState | None,
stepper: EulerCfgPpDiffusionStep,
transformer: X0Model,
denoiser: Denoiser,
noise_seed: int = -1,
new_noise_fn: Callable[[torch.Tensor, torch.Generator], torch.Tensor] = _get_plain_noise,
model_dtype: torch.dtype = torch.bfloat16,
) -> tuple[LatentState | None, LatentState | None]:
"""
Joint audio-video denoising loop using the CFG++ corrected Euler sampler.
Applies the CFG++ update rule at each step: the ODE derivative is computed
from the unconditioned denoised prediction rather than the standard velocity,
and an ancestral DDIM noise injection is applied in the rescaled sigma space.
Requires a guided denoiser whose :class:`~ltx_pipelines.utils.types.DenoisedLatentResult`
carries ``uncond`` tensors (i.e. CFG must be enabled).
Either ``video_state`` or ``audio_state`` may be ``None`` for absent modalities.
When both are present, noise is drawn from the same seeded generator (video
first, audio second) to produce a consistent random sequence.
### Parameters
sigmas:
1-D tensor of noise levels defining the sampling schedule.
video_state:
Current video :class:`~ltx_core.types.LatentState`, or ``None``.
audio_state:
Current audio :class:`~ltx_core.types.LatentState`, or ``None``.
stepper:
:class:`~ltx_core.components.diffusion_steps.EulerCfgPpDiffusionStep`
instance carrying ``eta`` and ``s_noise`` parameters.
transformer:
The diffusion model passed to the denoiser at each step.
denoiser:
Callable implementing :class:`~ltx_pipelines.utils.types.Denoiser`.
noise_seed:
Integer seed for the noise generator. Default ``-1``.
new_noise_fn:
``(latent, generator) -> noise`` callable. Defaults to plain
``torch.randn`` (no channel-wise normalization). Pass
:func:`_get_new_noise` for the normalized variant used in res2s.
model_dtype:
Dtype for latent state updates. Default ``bfloat16``.
### Returns
tuple[LatentState | None, LatentState | None]
Final ``(video_state, audio_state)`` after the denoising loop.
"""
if not isinstance(stepper, EulerCfgPpDiffusionStep):
raise ValueError(f"stepper must be an instance of EulerCfgPpDiffusionStep, got {type(stepper).__name__}")
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")
generator = torch.Generator(device=present_state.latent.device).manual_seed(noise_seed)
draw_noise = stepper.eta > 0 and stepper.s_noise > 0
for step_idx, _ in enumerate(tqdm(sigmas[:-1])):
video_result, audio_result = denoiser(transformer, video_state, audio_state, sigmas, step_idx)
denoised_video = video_result.denoised if video_result is not None else None
denoised_audio = audio_result.denoised if audio_result is not None else None
uncond_video = video_result.uncond if video_result is not None else None
uncond_audio = audio_result.uncond if audio_result is not None else None
if video_state is not None and not isinstance(uncond_video, torch.Tensor):
raise ValueError(
"euler_cfg_pp_denoising_loop requires video DenoisedLatentResult.uncond to be a tensor. "
"Use GuidedDenoiser or FactoryGuidedDenoiser with cfg_scale != 1 "
"or force_uncond_pass=True and a negative_context."
)
if audio_state is not None and not isinstance(uncond_audio, torch.Tensor):
raise ValueError(
"euler_cfg_pp_denoising_loop requires audio DenoisedLatentResult.uncond to be a tensor. "
"Use GuidedDenoiser or FactoryGuidedDenoiser with cfg_scale != 1 "
"or force_uncond_pass=True and a negative_context."
)
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:
if video_state is not None and denoised_video is not None:
video_state = replace(video_state, latent=denoised_video.to(model_dtype))
if audio_state is not None and denoised_audio is not None:
audio_state = replace(audio_state, latent=denoised_audio.to(model_dtype))
return video_state, audio_state
# Draw noise consecutively from the same generator: video first, audio second.
noise_video = new_noise_fn(video_state.latent, generator) if (video_state is not None and draw_noise) else None
noise_audio = new_noise_fn(audio_state.latent, generator) if (audio_state is not None and draw_noise) else None
if video_state is not None and denoised_video is not None:
x_next = stepper.step(
sample=video_state.latent,
denoised_sample=denoised_video,
sigmas=sigmas,
step_index=step_idx,
uncond_denoised=uncond_video,
noise=noise_video,
)
video_state = replace(video_state, latent=x_next.to(model_dtype))
if audio_state is not None and denoised_audio is not None:
x_next = stepper.step(
sample=audio_state.latent,
denoised_sample=denoised_audio,
sigmas=sigmas,
step_index=step_idx,
uncond_denoised=uncond_audio,
noise=noise_audio,
)
audio_state = replace(audio_state, latent=x_next.to(model_dtype))
return video_state, audio_state
@@ -40,6 +40,36 @@ class PipelineComponents:
self.audio_patchifier = AudioPatchifier(patch_size=1)
@dataclass(frozen=True)
class DenoisedLatentResult:
"""Output of one denoiser call for a single modality.
``denoised`` is the final blended prediction for this modality.
The remaining fields carry the per-pass raw outputs from ``_guided_denoise``
(all ``None`` for ``SimpleDenoiser``). Denoisers return a
``(video_result, audio_result)`` tuple; either element may be ``None``
for absent modalities.
"""
denoised: torch.Tensor
uncond: torch.Tensor | None = None
cond: torch.Tensor | None = None
ptb: torch.Tensor | None = None
mod: torch.Tensor | None = None
@classmethod
def result_or_none(
cls,
denoised: torch.Tensor | None,
uncond: torch.Tensor | None = None,
cond: torch.Tensor | None = None,
ptb: torch.Tensor | None = None,
mod: torch.Tensor | None = None,
) -> DenoisedLatentResult | None:
if denoised is None:
return None
return cls(denoised=denoised, uncond=uncond, cond=cond, ptb=ptb, mod=mod)
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
@@ -51,7 +81,8 @@ class Denoiser(Protocol):
sigmas: 1-D tensor of sigma values for each diffusion step.
step_index: Index of the current denoising step.
Returns:
``(denoised_video, denoised_audio)`` tensors (either may be ``None``).
A ``(video_result, audio_result)`` tuple of :class:`DenoisedLatentResult`,
either may be ``None`` for absent modalities.
"""
def __call__(
@@ -61,7 +92,7 @@ class Denoiser(Protocol):
audio_state: LatentState | None,
sigmas: torch.Tensor,
step_index: int,
) -> tuple[torch.Tensor | None, torch.Tensor | None]: ...
) -> tuple[DenoisedLatentResult | None, DenoisedLatentResult | None]: ...
@dataclass(frozen=True)