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
+3 -1
View File
@@ -39,7 +39,7 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https://
**Temporal Upscaler** - Supported by the model and will be required for future pipeline implementations
* [`ltx-2.3-temporal-upscaler-x2-1.0.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-temporal-upscaler-x2-1.0.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-temporal-upscaler-x2-1.0.safetensors)
**Distilled LoRA** - Required for current two-stage pipeline implementations in this repository (except DistilledPipeline and ICLoraPipeline)
**Distilled LoRA** - Required for current two-stage pipeline implementations in this repository (except DistilledPipeline, ICLoraPipeline, and LipDubPipeline)
* [`ltx-2.3-22b-distilled-lora-384-1.1.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-22b-distilled-lora-384-1.1.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-22b-distilled-lora-384-1.1.safetensors)
**Gemma Text Encoder** (download all assets from the repository)
@@ -58,6 +58,7 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https://
* [`LTX-2-19b-LoRA-Camera-Control-Jib-Up`](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Jib-Up) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Jib-Up/resolve/main/ltx-2-19b-lora-camera-control-jib-up.safetensors)
* [`LTX-2-19b-LoRA-Camera-Control-Static`](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Static) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Static/resolve/main/ltx-2-19b-lora-camera-control-static.safetensors)
* [`LTX-2.3-22b-IC-LoRA-HDR`](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-HDR) - HDR IC-LoRA and pre-computed text embeddings for `HDRICLoraPipeline`
* [`LTX-2.3-22b-IC-LoRA-LipDub`](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-LipDub) - [Download](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-LipDub/resolve/main/ltx-2.3-22b-ic-lora-lipdub-0.9.safetensors)
### Available Pipelines
@@ -70,6 +71,7 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https://
* **[A2VidPipelineTwoStage](packages/ltx-pipelines/src/ltx_pipelines/a2vid_two_stage.py)** - Audio-to-video generation conditioned on an input audio file
* **[RetakePipeline](packages/ltx-pipelines/src/ltx_pipelines/retake.py)** - Regenerate a specific time region of an existing video
* **[HDRICLoraPipeline](packages/ltx-pipelines/src/ltx_pipelines/hdr_ic_lora.py)** - Video-to-video with HDR output (linear float frames via LogC3 inverse decode, suitable for EXR export and tonemapping)
* **[LipDubPipeline](packages/ltx-pipelines/src/ltx_pipelines/lipdub.py)** - Lip dubbing, rephrasing, matching speaker identity (distilled model, single IC-LoRA, Two stages).
### ⚡ Optimization Tips
+7 -3
View File
@@ -77,13 +77,17 @@ model = builder.build(device=torch.device("cuda"))
Use the `.lora()` method to attach one or more LoRA adapters before calling `.build()`:
```python
from ltx_core.loader import SDOps
lora_sd_ops = SDOps(name="identity").with_matching() # or a model-specific key-renaming SDOps
builder = (
SingleGPUModelBuilder(
model_class_configurator=MyModelConfigurator,
model_path="/path/to/model.safetensors",
)
.lora("/path/to/lora_a.safetensors", strength=0.8)
.lora("/path/to/lora_b.safetensors", strength=0.5)
.lora("/path/to/lora_a.safetensors", 0.8, lora_sd_ops)
.lora("/path/to/lora_b.safetensors", 0.5, lora_sd_ops)
)
model = builder.build(device=torch.device("cuda"))
```
@@ -103,7 +107,7 @@ builder = SingleGPUModelBuilder(
model_class_configurator=MyModelConfigurator,
model_path="/path/to/model.safetensors",
lora_load_device=torch.device("cuda"),
).lora("/path/to/lora.safetensors", strength=1.0)
).lora("/path/to/lora.safetensors", 1.0, lora_sd_ops)
model = builder.build(device=torch.device("cuda"))
```
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "ltx-core"
version = "1.1.2"
version = "1.1.3"
description = "Core implementation of Lightricks' LTX-2 model"
readme = "README.md"
requires-python = ">=3.10"
@@ -7,16 +7,17 @@ from collections.abc import Callable
from dataclasses import dataclass, field, replace
from typing import Generic
import safetensors
import torch
from torch import nn
from ltx_core.block_streaming.disk import DiskBlockReader, DiskTensorReader, LoraSource
from ltx_core.block_streaming.pool import BlockLayout, WeightPool
from ltx_core.block_streaming.pool import WeightPool
from ltx_core.block_streaming.provider import WeightsProvider
from ltx_core.block_streaming.source import DiskWeightSource, PinnedWeightSource, WeightSource
from ltx_core.block_streaming.utils import build_pool_layout, resolve_attr
from ltx_core.block_streaming.utils import allocate_layout_views, derive_layout, make_block_key, resolve_attr
from ltx_core.block_streaming.wrapper import BlockStreamingWrapper
from ltx_core.loader.fuse_loras import apply_loras
from ltx_core.loader.fuse_loras import aggregate_lora_products, fuse_lora_weights
from ltx_core.loader.helpers import create_meta_model, load_state_dict, read_model_config
from ltx_core.loader.module_ops import ModuleOps
from ltx_core.loader.primitives import (
@@ -53,8 +54,8 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
``"velocity_model.transformer_blocks"``).
blocks_prefix: State-dict key prefix for block weights
(e.g. ``"transformer_blocks"``).
state_dict_prefix: Key prefix for non-block weights
(e.g. ``"velocity_model."``).
state_dict_prefix: Wrapper offset prepended to keys when loading into
the meta model (e.g. ``"velocity_model."`` when wrapped by ``X0Model``).
model_wrapper: Optional callable wrapping the model
(e.g. ``X0Model``).
"""
@@ -110,7 +111,6 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
if not self.blocks_prefix:
raise ValueError("blocks_prefix must be non-empty for streaming")
# 1. Create meta model (no weights allocated).
config = read_model_config(self.model_path, self.model_loader)
meta_model: nn.Module = create_meta_model(self.model_class_configurator, config, self.module_ops)
if self.model_wrapper is not None:
@@ -118,22 +118,29 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
meta_model.eval()
blocks = resolve_attr(meta_model, self.blocks_attr)
layout = build_pool_layout(blocks[0], dtype)
# 2. Determine slot counts.
checkpoint_paths = list(self.model_path) if isinstance(self.model_path, tuple) else [self.model_path]
block_key_map, non_block_keys = _scan_checkpoint_keys(checkpoint_paths, self.model_sd_ops, self.blocks_prefix)
cpu_slots_count = cpu_slots_count if cpu_slots_count is not None else len(blocks)
gpu_slots_count = gpu_slots_count if gpu_slots_count is not None else _DEFAULT_GPU_SLOTS
# 3. Build source and load non-block weights.
if cpu_slots_count >= len(blocks):
source, lora_sources = self._build_pinned_source(meta_model, target_device, dtype, cpu_slots_count)
source, lora_sources = self._build_pinned_source(
meta_model, target_device, dtype, cpu_slots_count, block_key_map, non_block_keys
)
else:
source, lora_sources = self._build_disk_source(meta_model, layout, target_device, dtype, cpu_slots_count)
reader = DiskTensorReader(checkpoint_paths)
source, lora_sources = self._build_disk_source(
meta_model, target_device, dtype, cpu_slots_count, reader, block_key_map, non_block_keys
)
# 4. Create provider and wrapper.
copy_stream = torch.cuda.Stream(device=target_device)
gpu_pool = WeightPool(
layout, gpu_slots_count, target_device, reuse_barrier=lambda event: copy_stream.wait_event(event)
source.block_layout,
gpu_slots_count,
target_device,
reuse_barrier=lambda event: copy_stream.wait_event(event),
)
provider = WeightsProvider(gpu_pool, copy_stream, target_device, source, lora_sources, self.blocks_prefix)
return BlockStreamingWrapper(
@@ -149,89 +156,90 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
target_device: torch.device,
dtype: torch.dtype,
cpu_slots_count: int,
block_key_map: dict[int, list[tuple[str, str]]],
non_block_keys: list[tuple[str, str]],
) -> tuple[WeightSource, list[LoraSource]]:
"""Pre-load all blocks into pinned CPU buffers with LoRA fusion."""
model_sd = load_state_dict(
self.model_path, self.model_loader, self.registry, torch.device("cpu"), self.model_sd_ops
)
if self.loras:
lora_sds = [
load_state_dict([lora.path], self.model_loader, self.registry, torch.device("cpu"), lora.sd_ops)
for lora in self.loras
]
lora_sd_and_strengths = [
LoraStateDictWithStrength(sd, lora.strength) for sd, lora in zip(lora_sds, self.loras, strict=True)
]
model_sd = apply_loras(
model_sd=model_sd,
lora_sd_and_strengths=lora_sd_and_strengths,
dtype=dtype,
destination_sd=model_sd if isinstance(self.registry, DummyRegistry) else None,
lora_sd_and_strengths = [
LoraStateDictWithStrength(
load_state_dict([lora.path], self.model_loader, self.registry, torch.device("cpu"), lora.sd_ops),
lora.strength,
)
for lora in self.loras
]
# Partition: non-block weights go to GPU, block weights go directly
# to pinned buffers. This avoids holding the full state dict and
# pinned copies simultaneously.
non_block_sd: dict[str, torch.Tensor] = {}
block_tensors: dict[int, dict[str, torch.Tensor]] = {}
prefix_dot = self.blocks_prefix + "."
for block_idx in block_key_map:
if block_idx >= cpu_slots_count:
raise ValueError(
f"Pinned source requires one CPU slot per block; "
f"got block index {block_idx} with only {cpu_slots_count} slots."
)
for key, tensor in model_sd.sd.items():
if key.startswith(prefix_dot):
rest = key[len(prefix_dot) :]
idx_str, _, param_name = rest.partition(".")
try:
block_idx = int(idx_str)
except ValueError:
non_block_sd[self.state_dict_prefix + key] = tensor.to(device=target_device, dtype=dtype)
continue
block_tensors.setdefault(block_idx, {})[param_name] = tensor
blocks = resolve_attr(meta_model, self.blocks_attr)
block_tensors: dict[str, torch.Tensor] = {}
for block_idx, entries in block_key_map.items():
block_params = dict(blocks[block_idx].named_parameters())
for _sft_key, param_name in entries:
key = make_block_key(self.blocks_prefix, block_idx, param_name)
block_tensors[key] = block_params[param_name]
blocks_layout = derive_layout(block_tensors, dtype)
pinned_blocks = allocate_layout_views(blocks_layout, pin_memory=True)
should_sync = False
for key, fused in fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype=None, preserve_input_device=False):
if key in pinned_blocks:
pinned_blocks[key].copy_(fused, non_blocking=True)
model_sd.sd[key] = None
should_sync = True
else:
non_block_sd[self.state_dict_prefix + key] = tensor.to(device=target_device, dtype=dtype)
model_sd.sd[key] = fused
if should_sync:
torch.cuda.synchronize()
# Fill remaining pinned keys from the source state dict.
for key in blocks_layout:
if model_sd.sd[key] is None:
continue
pinned_blocks[key].copy_(model_sd.sd[key])
model_sd.sd[key] = None
pinned: dict[int, dict[str, torch.Tensor]] = {
block_idx: {
param_name: pinned_blocks[make_block_key(self.blocks_prefix, block_idx, param_name)]
for _sft_key, param_name in entries
}
for block_idx, entries in block_key_map.items()
}
non_block_sd: dict[str, torch.Tensor] = {
self.state_dict_prefix + model_key: model_sd.sd[model_key].to(device=target_device, dtype=dtype)
for _sft_key, model_key in non_block_keys
}
meta_model.load_state_dict(non_block_sd, strict=False, assign=True)
del model_sd, non_block_sd
# Pin block weights one block at a time, freeing the source tensors as we go.
pinned: dict[int, dict[str, torch.Tensor]] = {}
for idx in range(cpu_slots_count):
src = block_tensors.pop(idx)
pinned[idx] = {name: tensor.to(dtype=dtype).pin_memory() for name, tensor in src.items()}
return PinnedWeightSource(pinned), []
def _build_disk_source(
self,
meta_model: nn.Module,
layout: BlockLayout,
target_device: torch.device,
dtype: torch.dtype,
cpu_slots_count: int,
reader: DiskTensorReader,
block_key_map: dict[int, list[tuple[str, str]]],
non_block_keys: list[tuple[str, str]],
) -> tuple[WeightSource, list[LoraSource]]:
"""Create a DiskWeightSource backed by a DiskBlockReader for lazy loading."""
"""Create a DiskWeightSource backed by a DiskBlockReader for lazy loading.
Derives the shared pool layout from the meta model's block 0 — this
relies on module_ops (e.g. fp8_cast) leaving the meta param dtype in
sync with the post-sd_ops checkpoint dtype.
"""
lora_sources = [LoraSource(lora.path, lora.sd_ops, lora.strength) for lora in self.loras]
checkpoint_paths = list(self.model_path) if isinstance(self.model_path, tuple) else [self.model_path]
reader = DiskTensorReader(checkpoint_paths)
block_key_map: dict[int, list[tuple[str, str]]] = {}
non_block_keys: list[tuple[str, str]] = []
for sft_key in reader.keys(): # noqa: SIM118
model_key = self.model_sd_ops.apply_to_key(sft_key) if self.model_sd_ops else sft_key
if model_key is None:
continue
if model_key.startswith(self.blocks_prefix + "."):
rest = model_key[len(self.blocks_prefix) + 1 :]
idx_str, _, param_name = rest.partition(".")
try:
block_idx = int(idx_str)
except ValueError:
non_block_keys.append((sft_key, model_key))
continue
block_key_map.setdefault(block_idx, []).append((sft_key, param_name))
else:
non_block_keys.append((sft_key, model_key))
self._load_non_block_weights(
reader,
@@ -242,9 +250,11 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
sd_ops=self.model_sd_ops,
key_prefix=self.state_dict_prefix,
lora_sources=lora_sources,
matmul_device=target_device,
)
blocks = resolve_attr(meta_model, self.blocks_attr)
layout = derive_layout(dict(blocks[0].named_parameters()), dtype)
cpu_pool = WeightPool(
layout,
cpu_slots_count,
@@ -252,7 +262,12 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
reuse_barrier=lambda event: event.synchronize(),
pin_memory=True,
)
block_reader = DiskBlockReader(reader=reader, block_key_map=block_key_map, dtype=dtype)
block_reader = DiskBlockReader(
reader=reader,
block_key_map=block_key_map,
sd_ops=self.model_sd_ops,
blocks_prefix=self.blocks_prefix,
)
source = DiskWeightSource(cpu_pool, block_reader)
return source, lora_sources
@@ -265,17 +280,17 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
model_key: str,
tensor: torch.Tensor,
lora_sources: list[LoraSource],
matmul_device: torch.device | None = None,
) -> torch.Tensor:
"""Add all matching LoRA deltas to *tensor* in-place."""
"""Add all matching LoRA deltas to *tensor* in-place via ``addmm_``."""
if not lora_sources or not model_key.endswith(".weight"):
return tensor
prefix = model_key[: -len(".weight")]
device = tensor.device if tensor.device.type == "cuda" else matmul_device
for source in lora_sources:
delta = source.get_delta(prefix, device=device)
if delta is not None:
tensor = tensor.add_(delta.to(device=tensor.device, dtype=tensor.dtype))
products = (
ab
for ab in (s.get_ab(prefix, device=tensor.device, dtype=tensor.dtype) for s in lora_sources)
if ab is not None
)
aggregate_lora_products(products, out=tensor)
return tensor
@staticmethod
@@ -289,17 +304,48 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
sd_ops: SDOps | None = None,
key_prefix: str = "",
lora_sources: list[LoraSource] | None = None,
matmul_device: torch.device | None = None,
) -> None:
"""Load non-block weights into *model* on *device*."""
state_dict: dict[str, torch.Tensor] = {}
sources = lora_sources or []
for sft_key, model_key in non_block_keys:
tensor = reader.get_tensor(sft_key).to(device=device, dtype=dtype)
tensor = StreamingModelBuilder._fuse_lora_delta(model_key, tensor, sources, matmul_device)
tensor = StreamingModelBuilder._fuse_lora_delta(model_key, tensor, sources)
if sd_ops is not None:
for kv in sd_ops.apply_to_key_value(model_key, tensor):
state_dict[key_prefix + kv.new_key] = kv.new_value
continue
state_dict[key_prefix + model_key] = tensor
model.load_state_dict(state_dict, strict=False, assign=True)
def _scan_checkpoint_keys(
checkpoint_paths: list[str],
sd_ops: SDOps | None,
blocks_prefix: str,
) -> tuple[dict[int, list[tuple[str, str]]], list[tuple[str, str]]]:
"""Partition checkpoint keys into per-block and non-block lists.
Opens the safetensors files for header-only key enumeration; no tensor data
is read.
"""
block_key_map: dict[int, list[tuple[str, str]]] = {}
non_block_keys: list[tuple[str, str]] = []
prefix_dot = blocks_prefix + "."
for path in checkpoint_paths:
with safetensors.safe_open(path, framework="pt", device="cpu") as handle:
for sft_key in handle.keys(): # noqa: SIM118
model_key = sd_ops.apply_to_key(sft_key) if sd_ops else sft_key
if model_key is None:
continue
if model_key.startswith(prefix_dot):
rest = model_key[len(prefix_dot) :]
idx_str, _, param_name = rest.partition(".")
try:
block_idx = int(idx_str)
except ValueError:
non_block_keys.append((sft_key, model_key))
continue
block_key_map.setdefault(block_idx, []).append((sft_key, param_name))
else:
non_block_keys.append((sft_key, model_key))
return block_key_map, non_block_keys
@@ -2,11 +2,22 @@
from __future__ import annotations
from collections.abc import Iterator
import safetensors
import torch
from ltx_core.block_streaming.utils import allocate_layout_views, make_block_key
from ltx_core.loader.fuse_loras import LoraProduct
from ltx_core.loader.sd_ops import SDOps
_SAFETENSORS_DTYPE_TO_TORCH: dict[str, torch.dtype] = {
"F64": torch.float64,
"F32": torch.float32,
"F16": torch.float16,
"BF16": torch.bfloat16,
}
class DiskTensorReader:
"""Key-based tensor accessor over one or more safetensors files."""
@@ -21,9 +32,6 @@ class DiskTensorReader:
for sft_key in handle.keys(): # noqa: SIM118
self._key_to_handle_idx[sft_key] = handle_idx
def keys(self) -> list[str]:
return list(self._key_to_handle_idx.keys())
def get_tensor(self, key: str) -> torch.Tensor:
return self._handles[self._key_to_handle_idx[key]].get_tensor(key)
@@ -31,49 +39,58 @@ class DiskTensorReader:
self._handles.clear()
self._key_to_handle_idx.clear()
def __contains__(self, key: str) -> bool:
return key in self._key_to_handle_idx
def __iter__(self) -> Iterator[str]:
return iter(self._key_to_handle_idx)
class DiskBlockReader:
"""Reads one block at a time from safetensors into provided buffers.
Maps block indices to safetensors keys via a pre-computed key map.
"""
"""Reads one block at a time from safetensors into provided buffers."""
def __init__(
self,
reader: DiskTensorReader,
block_key_map: dict[int, list[tuple[str, str]]],
dtype: torch.dtype,
sd_ops: SDOps | None = None,
blocks_prefix: str = "",
) -> None:
self._reader = reader
self._block_key_map = block_key_map
self._dtype = dtype
self._sd_ops = sd_ops
self._blocks_prefix = blocks_prefix
def read_into(self, target: dict[str, torch.Tensor], block_idx: int) -> None:
block_prefix = make_block_key(self._blocks_prefix, block_idx, "")
for sft_key, param_name in self._block_key_map[block_idx]:
tensor = self._reader.get_tensor(sft_key)
if tensor.dtype != self._dtype:
tensor = tensor.to(self._dtype)
target[param_name].copy_(tensor)
if self._sd_ops is None:
target[param_name].copy_(tensor)
continue
full_key = make_block_key(self._blocks_prefix, block_idx, param_name)
for result in self._sd_ops.apply_to_key_value(full_key, tensor):
if not result.new_key.startswith(block_prefix):
raise ValueError(
f"SDOps output key '{result.new_key}' is outside block {block_idx} "
f"(expected prefix '{block_prefix}'); cannot route to a per-block buffer."
)
target[result.new_key[len(block_prefix) :]].copy_(result.new_value)
def cleanup(self) -> None:
self._reader.close()
class LoraSource:
"""Pinned-memory cache of LoRA A/B matrices for on-the-fly fusion.
At init, loads all matched A/B pairs into pinned CPU memory.
:meth:`get_delta` computes ``(B * strength) @ A`` on the given device.
"""
"""Pinned-memory cache of matched LoRA A/B factors backed by a single buffer."""
def __init__(self, path: str, sd_ops: SDOps | None, strength: float) -> None:
self.strength = strength
# param_prefix -> (pinned_a, pinned_b)
self._pinned_ab: dict[str, tuple[torch.Tensor, torch.Tensor]] = {}
a_keys: dict[str, str] = {}
b_keys: dict[str, str] = {}
with safetensors.safe_open(path, framework="pt", device="cpu") as handle:
# First pass: build key map.
for sft_key in handle.keys(): # noqa: SIM118
model_key = sd_ops.apply_to_key(sft_key) if sd_ops is not None else sft_key
if model_key is None:
@@ -83,24 +100,49 @@ class LoraSource:
elif model_key.endswith(".lora_B.weight"):
b_keys[model_key[: -len(".lora_B.weight")]] = sft_key
# Second pass: load and pin matched A+B pairs (orphans silently skipped).
for prefix in a_keys.keys() & b_keys.keys():
self._pinned_ab[prefix] = (
handle.get_tensor(a_keys[prefix]).pin_memory(),
handle.get_tensor(b_keys[prefix]).pin_memory(),
matched_prefixes = list(a_keys.keys() & b_keys.keys())
# Build the layout from safetensors header metadata only — no tensor data is read.
layout: dict[str, tuple[torch.Size, torch.dtype]] = {}
for prefix in matched_prefixes:
a_slice_view = handle.get_slice(a_keys[prefix])
b_slice_view = handle.get_slice(b_keys[prefix])
layout[f"{prefix}.A"] = (
torch.Size(a_slice_view.get_shape()),
_SAFETENSORS_DTYPE_TO_TORCH[a_slice_view.get_dtype()],
)
layout[f"{prefix}.B"] = (
torch.Size(b_slice_view.get_shape()),
_SAFETENSORS_DTYPE_TO_TORCH[b_slice_view.get_dtype()],
)
def get_delta(self, param_prefix: str, device: torch.device | None = None) -> torch.Tensor | None:
"""Return ``(B * strength) @ A`` for *param_prefix*, or ``None``."""
all_views = allocate_layout_views(layout, pin_memory=True)
for prefix in matched_prefixes:
a_view = all_views[f"{prefix}.A"]
b_view = all_views[f"{prefix}.B"]
a_view.copy_(handle.get_tensor(a_keys[prefix]))
b_view.copy_(handle.get_tensor(b_keys[prefix]))
self._pinned_ab[prefix] = (a_view, b_view)
def get_ab(
self,
param_prefix: str,
device: torch.device | None = None,
dtype: torch.dtype | None = None,
) -> LoraProduct | None:
"""Return the :class:`LoraProduct` for *param_prefix*, or ``None``."""
pair = self._pinned_ab.get(param_prefix)
if pair is None:
return None
a, b = pair
if device is not None and device.type == "cuda":
a = a.to(device=device)
b = b.to(device=device)
delta = torch.matmul(b * self.strength, a)
return delta
a = a.to(device=device, non_blocking=True)
b = b.to(device=device, non_blocking=True)
if dtype is not None:
a = a.to(dtype=dtype)
b = b.to(dtype=dtype)
return LoraProduct(a, b, self.strength)
def cleanup(self) -> None:
self._pinned_ab.clear()
@@ -7,20 +7,16 @@ from typing import Callable
import torch
from ltx_core.block_streaming.utils import allocate_buffer
# Type alias for the buffer layout used by slot allocation.
BlockLayout = dict[str, tuple[torch.Size, torch.dtype]]
from ltx_core.block_streaming.utils import allocate_layout_views
from ltx_core.loader.primitives import TensorLayout
class WeightPool:
"""Fixed pool of pre-allocated weight buffers with event-based reuse safety.
Buffers are allocated once at construction. :meth:`acquire` pops a
free buffer (waiting any pending event first). :meth:`release`
returns it, optionally attaching an event that must complete before
the buffer can be reused.
"""Fixed pool of pre-allocated weight buffers with event-based reuse.
All slots share a single buffer (CPU or GPU); each slot is a
contiguous slice carved out of it via :func:`allocate_layout_views`.
Args:
layout: ``{name: (shape, dtype)}`` for each buffer.
buffer_layout: ``{name: (shape, dtype)}`` for each buffer.
capacity: Number of buffers to pre-allocate.
device: Device for allocation.
reuse_barrier: Called with the pending event before a buffer is reused.
@@ -29,23 +25,34 @@ class WeightPool:
def __init__(
self,
layout: BlockLayout,
buffer_layout: TensorLayout,
capacity: int,
device: torch.device,
reuse_barrier: Callable[[torch.cuda.Event], None],
pin_memory: bool = False,
) -> None:
self._buffer_layout = buffer_layout
self._capacity = capacity
self._free: deque[dict[str, torch.Tensor]] = deque()
self._events: dict[int, torch.cuda.Event] = {}
self._reuse_barrier = reuse_barrier
for _ in range(capacity):
self._free.append(allocate_buffer(layout, device, pin_memory))
memory_layout = {
_make_key(slot, name): (shape, dtype)
for slot in range(capacity)
for name, (shape, dtype) in buffer_layout.items()
}
all_views = allocate_layout_views(memory_layout, device=device, pin_memory=pin_memory)
for slot in range(capacity):
self._free.append({name: all_views[_make_key(slot, name)] for name in buffer_layout})
@property
def capacity(self) -> int:
return self._capacity
@property
def buffer_layout(self) -> TensorLayout:
return self._buffer_layout
def acquire(self) -> dict[str, torch.Tensor]:
"""Take a free buffer, waiting any pending event before returning."""
weights = self._free.popleft()
@@ -62,3 +69,7 @@ class WeightPool:
if event is not None:
self._events[id(weights)] = event
self._free.append(weights)
def _make_key(slot: int, name: str) -> str:
return f"{slot}/{name}"
@@ -9,6 +9,29 @@ import torch
from ltx_core.block_streaming.disk import LoraSource
from ltx_core.block_streaming.pool import WeightPool
from ltx_core.block_streaming.source import WeightSource
from ltx_core.block_streaming.utils import FP8_DTYPES
from ltx_core.loader.fuse_loras import aggregate_lora_products, fuse_cast_fp8_weight
def _contiguous_byte_view(weights: dict[str, torch.Tensor]) -> torch.Tensor | None:
"""Return a ``uint8`` view spanning every tensor in *weights*, or ``None`` if
they don't share one contiguous storage region."""
tensors = list(weights.values())
if not tensors:
return None
storage = tensors[0].untyped_storage()
storage_ptr = storage.data_ptr()
start = end = tensors[0].storage_offset() * tensors[0].element_size()
for t in tensors:
if t.untyped_storage().data_ptr() != storage_ptr or not t.is_contiguous():
return None
offset = t.storage_offset() * t.element_size()
nbytes = t.numel() * t.element_size()
start = min(start, offset)
end = max(end, offset + nbytes)
view = torch.empty(0, dtype=torch.uint8, device=tensors[0].device)
view.set_(storage, start, (end - start,), (1,))
return view
class WeightsProvider:
@@ -70,8 +93,13 @@ class WeightsProvider:
instrumentation regions wrapping it -- observe the full transfer time.
"""
with torch.cuda.stream(self._copy_stream):
for name, gpu_tensor in gpu_weights.items():
gpu_tensor.copy_(cpu_weights[name], non_blocking=True)
gpu_view = _contiguous_byte_view(gpu_weights)
cpu_view = _contiguous_byte_view(cpu_weights)
if gpu_view is not None and cpu_view is not None and gpu_view.numel() == cpu_view.numel():
gpu_view.copy_(cpu_view, non_blocking=True)
else:
for name, gpu_tensor in gpu_weights.items():
gpu_tensor.copy_(cpu_weights[name], non_blocking=True)
if self._lora_sources:
self._fuse_block_loras(idx, gpu_weights)
h2d_event = torch.cuda.Event()
@@ -102,9 +130,18 @@ class WeightsProvider:
for name, tensor in weights.items():
if not name.endswith(".weight"):
continue
full_key = f"{self._blocks_prefix}.{idx}.{name}"
prefix = full_key[: -len(".weight")]
for source in self._lora_sources:
delta = source.get_delta(prefix, device=self._target_device)
if delta is not None:
tensor.add_(delta.to(dtype=tensor.dtype))
prefix = f"{self._blocks_prefix}.{idx}.{name}".removesuffix(".weight")
is_fp8 = tensor.dtype in FP8_DTYPES
agg_dtype = torch.bfloat16 if is_fp8 else tensor.dtype
products = (
ab
for ab in (s.get_ab(prefix, device=self._target_device, dtype=agg_dtype) for s in self._lora_sources)
if ab is not None
)
aggregated = aggregate_lora_products(products, agg_dtype)
if aggregated is None:
continue
if is_fp8:
tensor.copy_(fuse_cast_fp8_weight(aggregated, tensor, tensor.dtype))
else:
tensor.add_(aggregated)
@@ -9,10 +9,18 @@ import torch
from ltx_core.block_streaming.disk import DiskBlockReader
from ltx_core.block_streaming.pool import WeightPool
from ltx_core.loader.primitives import TensorLayout
class WeightSource(Protocol):
"""Provides pinned CPU weights for a given block index."""
"""Provides pinned CPU weights for a given block index.
Assumes all buffers share an identical layout across all block indices.
"""
@property
def block_layout(self) -> TensorLayout:
"""Shared per-block buffer layout (shape + dtype for each param)."""
...
def get(self, idx: int) -> dict[str, torch.Tensor]:
"""Return CPU weights for block *idx*."""
@@ -36,6 +44,10 @@ class DiskWeightSource(WeightSource):
self._events: dict[int, torch.cuda.Event] = {}
self._reader = reader
@property
def block_layout(self) -> TensorLayout:
return self._pool.buffer_layout
def get(self, idx: int) -> dict[str, torch.Tensor]:
"""Return CPU weights for block *idx*. Reads from disk on miss."""
if idx in self._cache:
@@ -68,8 +80,15 @@ class PinnedWeightSource(WeightSource):
"""Pre-loaded pinned CPU weights."""
def __init__(self, weights: dict[int, dict[str, torch.Tensor]]) -> None:
if not weights:
raise ValueError("PinnedWeightSource requires at least one block")
self._weights = weights
@property
def block_layout(self) -> TensorLayout:
first_block = self._weights[min(self._weights)]
return {name: (t.shape, t.dtype) for name, t in first_block.items()}
def get(self, idx: int) -> dict[str, torch.Tensor]:
return self._weights[idx]
@@ -2,14 +2,24 @@
from __future__ import annotations
import itertools
from typing import TYPE_CHECKING, Any
import math
import weakref
from dataclasses import dataclass
from typing import Any
import torch
from torch import nn
if TYPE_CHECKING:
from ltx_core.block_streaming.pool import BlockLayout
from ltx_core.loader.primitives import TensorLayout
FP8_DTYPES = frozenset({torch.float8_e4m3fn, torch.float8_e5m2})
_BUFFER_ALIGN = 16
def make_block_key(blocks_prefix: str, block_idx: int, param_name: str) -> str:
"""Return the state-dict key for *param_name* under block *block_idx*."""
return f"{blocks_prefix}.{block_idx}.{param_name}"
def resolve_attr(module: nn.Module, dotted_path: str) -> nn.ModuleList:
@@ -40,21 +50,85 @@ def assign_tensor_to_module(root: nn.Module, dotted_name: str, tensor: torch.Ten
raise AttributeError(f"{leaf} is not a parameter or buffer of {type(parent).__name__}")
def build_pool_layout(block: nn.Module, dtype: torch.dtype) -> BlockLayout:
"""Derive a buffer layout from a block's parameters and buffers.
Works on meta-device blocks (shapes are valid regardless of device).
The *dtype* argument overrides each tensor's dtype so the pool matches
the target inference precision.
def derive_layout(tensors: dict[str, torch.Tensor], dtype: torch.dtype | None = None) -> TensorLayout:
"""Derive a layout from a ``{name: tensor}`` dict.
If ``dtype`` is given, non-FP8 dtypes are coerced to it (FP8 preserved). If
``None``, the source dtype is preserved as-is.
"""
layout: BlockLayout = {}
for name, tensor in itertools.chain(block.named_parameters(), block.named_buffers()):
layout[name] = (tensor.shape, dtype)
return layout
def allocate_buffer(layout: BlockLayout, device: torch.device, pin_memory: bool = False) -> dict[str, torch.Tensor]:
"""Allocate a single buffer dict matching *layout*."""
return {
name: torch.empty(shape, dtype=dtype, device=device, pin_memory=pin_memory)
for name, (shape, dtype) in layout.items()
name: (t.shape, t.dtype if dtype is None or t.dtype in FP8_DTYPES else dtype) for name, t in tensors.items()
}
def _align_up(offset: int, alignment: int) -> int:
return (offset + alignment - 1) & ~(alignment - 1)
def _alloc_pinned_exact(nbytes: int) -> torch.Tensor | None:
"""Allocate exactly ``nbytes`` of pinned host memory via ``cudaHostRegister``.
Bypasses PyTorch's ``CachingHostAllocator``, which rounds every
``pin_memory=True`` request up to ``PowerOf2Ceil(N)`` (see
``aten/src/ATen/core/CachingHostAllocator.h``). Returns ``None`` if
registration fails. The unregister hook is bound to the storage (not the
tensor) so views of the buffer keep the registration alive until the
memory is actually freed. Caller is responsible for ensuring CUDA is
available.
"""
cudart = torch.cuda.cudart()
buf = torch.empty(nbytes, dtype=torch.uint8)
ptr = buf.data_ptr()
err = int(cudart.cudaHostRegister(ptr, nbytes, 0))
if err != 0:
return None
weakref.finalize(buf.untyped_storage(), lambda p=ptr: cudart.cudaHostUnregister(p))
return buf
def _alloc_buffer(nbytes: int, device: torch.device | None, pin_memory: bool) -> torch.Tensor:
"""Allocate one ``uint8`` buffer for :func:`allocate_layout_views`.
For pinned host buffers, prefer ``cudaHostRegister`` to dodge the caching
allocator's power-of-2 rounding. Falls back to the caching allocator if
registration fails. Raises if pinning is requested without a CUDA runtime,
since pinning is fundamentally a CUDA driver operation.
"""
if pin_memory and (device is None or torch.device(device).type == "cpu"):
if not torch.cuda.is_available():
raise RuntimeError("pin_memory=True requires CUDA, which is not available")
buf = _alloc_pinned_exact(nbytes)
if buf is not None:
return buf
return torch.empty(nbytes, dtype=torch.uint8, device=device, pin_memory=pin_memory)
@dataclass(frozen=True)
class _TensorSlice:
"""Location of a single tensor view within the buffer."""
offset: int
shape: torch.Size
dtype: torch.dtype
def size(self) -> int:
return math.prod(self.shape) * self.dtype.itemsize
def allocate_layout_views(
layout: TensorLayout,
device: torch.device | None = None,
pin_memory: bool = False,
) -> dict[str, torch.Tensor]:
"""Allocate a single ``uint8`` buffer and return per-key tensor views into it.
All keys in *layout* live in one contiguous allocation; each returned
tensor is a non-overlapping slice of that buffer reinterpreted at the
requested shape and dtype. The views keep the underlying storage alive
via PyTorch refcounting — drop them all to release the memory.
"""
slices: dict[str, _TensorSlice] = {}
cursor = 0
for key, (shape, dtype) in layout.items():
cursor = _align_up(cursor, _BUFFER_ALIGN)
slices[key] = _TensorSlice(offset=cursor, shape=shape, dtype=dtype)
cursor += slices[key].size()
# Allocate at least one byte so empty layouts still produce a valid buffer.
buffer = _alloc_buffer(max(_align_up(cursor, _BUFFER_ALIGN), 1), device, pin_memory)
return {key: buffer[s.offset : s.offset + s.size()].view(s.dtype).view(s.shape) for key, s in slices.items()}
@@ -4,6 +4,24 @@ from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.utils import to_velocity
def _get_ancestral_step(
sigma_from: torch.Tensor,
sigma_to: torch.Tensor,
eta: float = 1.0,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Compute ``(sigma_down, sigma_up)`` for one DDIM ancestral sampling step.
Both inputs are in the rescaled parameterization ``sigma / alpha``.
Returns ``sigma_down`` (deterministic component) and ``sigma_up``
(stochastic component) in the same rescaled space.
"""
if not eta:
return sigma_to, torch.zeros_like(sigma_to)
variance = sigma_to**2 * (sigma_from**2 - sigma_to**2).clamp(min=0) / sigma_from**2
sigma_up = (eta * variance**0.5).clamp(max=sigma_to)
sigma_down = (sigma_to**2 - sigma_up**2).clamp(min=0) ** 0.5
return sigma_down, sigma_up
class EulerDiffusionStep(DiffusionStepProtocol):
"""
First-order Euler method for diffusion sampling.
@@ -104,3 +122,65 @@ class Res2sDiffusionStep(DiffusionStepProtocol):
# Mix deterministic and stochastic components
x_noised = alpha_ratio * (denoised_next + sigma_down * eps_next) + sigma_up * noise
return x_noised.to(output_dtype)
class EulerCfgPpDiffusionStep(DiffusionStepProtocol):
"""Euler step using the CFG++ correction for the ODE derivative.
Instead of the standard velocity formula, the ODE derivative is computed
from the unconditioned prediction, keeping the conditioned prediction as
the target denoised state. Ancestral (DDIM) noise injection is applied
in the rescaled sigma parameterization (sigma / alpha).
All diffusion quantities (alpha, ODE derivative, ancestral coefficients)
are computed internally from ``sigmas`` and ``uncond_denoised``.
Reference: CFG++ (https://arxiv.org/abs/2406.08070).
"""
def __init__(self, eta: float = 1.0, s_noise: float = 1.0) -> None:
self.eta = eta
self.s_noise = s_noise
def step(
self,
sample: torch.Tensor,
denoised_sample: torch.Tensor,
sigmas: torch.Tensor,
step_index: int,
uncond_denoised: torch.Tensor,
noise: torch.Tensor | None = None,
**_kwargs,
) -> torch.Tensor:
"""Advance one CFG++ Euler step.
Args:
sample: Current noisy latent x_t.
denoised_sample: Conditioned denoised prediction x_0^cond.
sigmas: Full sigma schedule tensor.
step_index: Current step index.
uncond_denoised: Unconditioned denoised prediction x_0^uncond,
used to compute the ODE derivative direction.
noise: Noise tensor for stochastic injection; ignored when
``eta=0`` or ``s_noise=0``.
Returns:
Updated latent x_{t-1}.
"""
sigma_s = sigmas[step_index].to(torch.float32)
sigma_t = sigmas[step_index + 1].to(torch.float32)
_eps = torch.finfo(torch.float32).eps
# Clamp to avoid division by zero when sigma == 1.0 exactly.
alpha_s = (1.0 - sigma_s).clamp(min=_eps)
alpha_t = (1.0 - sigma_t).clamp(min=_eps)
x = sample.to(torch.float32)
denoised = denoised_sample.to(torch.float32)
uncond = uncond_denoised.to(torch.float32)
# ODE derivative: direction toward noise using uncond prediction (CFG++ correction)
d = (x - alpha_s * uncond) / sigma_s
# Ancestral step in rescaled sigma space (sigma / alpha)
sigma_down, sigma_up = _get_ancestral_step(sigma_s / alpha_s, sigma_t / alpha_t, eta=self.eta)
sigma_down = alpha_t * sigma_down
x_next = alpha_t * denoised + sigma_down * d
if noise is not None and self.eta > 0 and self.s_noise > 0:
x_next = x_next + alpha_t * noise.to(torch.float32) * self.s_noise * sigma_up
return x_next.to(sample.dtype)
@@ -151,17 +151,22 @@ def get_pixel_coords(
that treat frame zero differently still yield non-negative timestamps.
"""
# Broadcast the VAE scale factors so they align with the `(batch, axis, patch, bound)` layout.
# Axis 1 of `latent_coords` is ordered (frame/time, height, width) — match that explicitly by
# pulling fields from the NamedTuple rather than relying on tuple iteration order.
broadcast_shape = [1] * latent_coords.ndim
broadcast_shape[1] = -1 # axis dimension corresponds to (frame/time, height, width)
scale_tensor = torch.tensor(scale_factors, device=latent_coords.device).view(*broadcast_shape)
scale_tensor = torch.tensor(
[scale_factors.time, scale_factors.height, scale_factors.width],
device=latent_coords.device,
).view(*broadcast_shape)
# Apply per-axis scaling to convert latent bounds into pixel-space coordinates.
pixel_coords = latent_coords * scale_tensor
if causal_fix:
# VAE temporal stride for the very first frame is 1 instead of `scale_factors[0]`.
# VAE temporal stride for the very first frame is 1 instead of `scale_factors.time`.
# Shift and clamp to keep the first-frame timestamps causal and non-negative.
pixel_coords[:, 0, ...] = (pixel_coords[:, 0, ...] + 1 - scale_factors[0]).clamp(min=0)
pixel_coords[:, 0, ...] = (pixel_coords[:, 0, ...] + 1 - scale_factors.time).clamp(min=0)
return pixel_coords
@@ -3,6 +3,7 @@
from ltx_core.conditioning.exceptions import ConditioningError
from ltx_core.conditioning.item import ConditioningItem
from ltx_core.conditioning.types import (
AudioConditionByReferenceLatent,
ConditioningItemAttentionStrengthWrapper,
VideoConditionByKeyframeIndex,
VideoConditionByLatentIndex,
@@ -10,6 +11,7 @@ from ltx_core.conditioning.types import (
)
__all__ = [
"AudioConditionByReferenceLatent",
"ConditioningError",
"ConditioningItem",
"ConditioningItemAttentionStrengthWrapper",
@@ -3,9 +3,11 @@
from ltx_core.conditioning.types.attention_strength_wrapper import ConditioningItemAttentionStrengthWrapper
from ltx_core.conditioning.types.keyframe_cond import VideoConditionByKeyframeIndex
from ltx_core.conditioning.types.latent_cond import VideoConditionByLatentIndex
from ltx_core.conditioning.types.reference_audio_cond import AudioConditionByReferenceLatent
from ltx_core.conditioning.types.reference_video_cond import VideoConditionByReferenceLatent
__all__ = [
"AudioConditionByReferenceLatent",
"ConditioningItemAttentionStrengthWrapper",
"VideoConditionByKeyframeIndex",
"VideoConditionByLatentIndex",
@@ -0,0 +1,59 @@
"""Audio reference conditioning items."""
from __future__ import annotations
import torch
from ltx_core.conditioning.mask_utils import update_attention_mask
from ltx_core.tools import LatentTools
from ltx_core.types import LatentState
class AudioConditionByReferenceLatent:
"""Append patchified reference audio tokens after the target audio sequence.
Mirrors :class:`ltx_core.conditioning.types.reference_video_cond.VideoConditionByReferenceLatent`
but for audio. The reference tokens are appended so the target audio tokens stay
in the first ``num_noisy_tokens`` positions and can be kept by
:meth:`ltx_core.tools.LatentTools.clear_conditioning`.
Args:
patchified: Patchified reference latent ``[B, T_ref, C]``.
positions: RoPE positions for reference tokens, ``[B, 1, T_ref, 2]``.
strength: 1.0 keeps reference clean; 0.0 would fully denoise it.
"""
def __init__(
self,
patchified: torch.Tensor,
positions: torch.Tensor,
strength: float = 1.0,
) -> None:
self.patchified = patchified
self.positions = positions.to(dtype=torch.float32)
self.strength = strength
def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState:
tokens = self.patchified
denoise_mask = torch.full(
size=(*tokens.shape[:2], 1),
fill_value=1.0 - self.strength,
device=tokens.device,
dtype=tokens.dtype,
)
new_attention_mask = update_attention_mask(
latent_state=latent_state,
attention_mask=None,
num_noisy_tokens=latent_tools.patchifier.get_token_count(latent_tools.target_shape),
num_new_tokens=tokens.shape[1],
batch_size=tokens.shape[0],
device=tokens.device,
dtype=tokens.dtype,
)
return LatentState(
latent=torch.cat([latent_state.latent, tokens], dim=1),
denoise_mask=torch.cat([latent_state.denoise_mask, denoise_mask], dim=1),
positions=torch.cat([latent_state.positions, self.positions], dim=2),
clean_latent=torch.cat([latent_state.clean_latent, tokens], dim=1),
attention_mask=new_attention_mask,
)
@@ -1,26 +1,81 @@
from collections.abc import Iterator
from collections.abc import Iterable, Iterator
from typing import NamedTuple
import torch
from ltx_core.loader.kernels import TRITON_AVAILABLE
from ltx_core.loader.primitives import LoraStateDictWithStrength, StateDict
from ltx_core.quantization.fp8_cast import _fused_add_round_launch
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
class LoraProduct(NamedTuple):
"""A LoRA's ``A``, ``B`` factors and its strength scalar."""
a: torch.Tensor
b: torch.Tensor
strength: float
def _get_device() -> torch.device:
if torch.cuda.is_available():
return torch.device("cuda", torch.cuda.current_device())
return torch.device("cpu")
def aggregate_lora_products(
products: Iterable[LoraProduct],
dtype: torch.dtype | None = None,
*,
out: torch.Tensor | None = None,
) -> torch.Tensor | None:
"""Accumulate ``sum((B * strength) @ A)`` across :class:`LoraProduct` items.
If ``out`` is provided, ``addmm_`` accumulates directly into it — caller
ensures A/B dtypes and devices match ``out``. Otherwise the first product
materializes the ``(out, in)``-shape aggregator at ``dtype``; subsequent
products use ``addmm_`` to avoid allocating the full intermediate delta.
Returns ``out`` (or the new aggregator), or ``None`` if ``products`` was empty
and ``out`` was not given.
"""
aggregated = out
for product in products:
if aggregated is None:
aggregated = torch.matmul(product.b * product.strength, product.a).to(dtype=dtype)
else:
aggregated.addmm_(product.b, product.a, alpha=product.strength)
return aggregated
def fuse_cast_fp8_weight(
delta_bf16: torch.Tensor,
weight_fp8: torch.Tensor,
target_dtype: torch.dtype,
) -> torch.Tensor:
"""Return ``(delta_bf16 + dequantize(weight_fp8)).to(target_dtype)``.
CUDA with Triton uses stochastic rounding; otherwise uses a deterministic bf16 add.
``delta_bf16`` is the bf16 accumulator and is mutated in place.
"""
if delta_bf16.dtype != torch.bfloat16:
raise ValueError(f"delta_bf16 must be bfloat16, got {delta_bf16.dtype}")
if str(weight_fp8.device).startswith("cuda") and TRITON_AVAILABLE:
fused_add_round_launch(delta_bf16, weight_fp8, seed=0)
else:
delta_bf16.add_(weight_fp8.to(dtype=torch.bfloat16))
return delta_bf16.to(dtype=target_dtype)
def fuse_lora_weights(
model_sd: StateDict,
lora_sd_and_strengths: list[LoraStateDictWithStrength],
dtype: torch.dtype | None = None,
preserve_input_device: bool = True,
) -> 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.
When ``preserve_input_device`` is False, fused tensors are yielded on the device
used for fusion; caller is responsible for moving them to their final
destination.
"""
for key, original_weight in model_sd.sd.items():
if original_weight is None or key.endswith(".weight_scale"):
@@ -30,7 +85,7 @@ def fuse_lora_weights(
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)
deltas = _aggregate_deltas(lora_sd_and_strengths, key, deltas_dtype, weight.device)
if deltas is None:
continue
@@ -41,14 +96,15 @@ def fuse_lora_weights(
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)
fused = {key: fuse_cast_fp8_weight(deltas, weight, target_dtype)}
elif weight.dtype == torch.bfloat16:
fused = _fuse_delta_with_bfloat16(deltas, weight, key, target_dtype)
deltas.add_(weight)
fused = {key: deltas.to(dtype=target_dtype)}
else:
raise ValueError(f"Unsupported dtype: {weight.dtype}")
for k, v in fused.items():
yield k, v.to(device=original_device)
yield k, v.to(device=original_device) if preserve_input_device else v
def apply_loras(
@@ -57,10 +113,12 @@ def apply_loras(
dtype: torch.dtype | None = None,
destination_sd: StateDict | None = None,
) -> StateDict:
"""Fuse LoRAs into ``model_sd`` and place the results in ``destination_sd``.
When ``destination_sd`` is provided, the fused tensors are placed directly into it.
"""
if destination_sd is not None:
sd = destination_sd.sd
for key, tensor in fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype):
sd[key] = tensor
for key, fused in fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype):
destination_sd.sd[key] = fused
return destination_sd
fused = dict(fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype))
@@ -68,26 +126,22 @@ def apply_loras(
return StateDict(sd, model_sd.device, model_sd.size, model_sd.dtype)
def _prepare_deltas(
def _aggregate_deltas(
lora_sd_and_strengths: list[LoraStateDictWithStrength], key: str, dtype: torch.dtype, device: torch.device
) -> torch.Tensor | None:
deltas = []
prefix = key[: -len(".weight")]
key_a = f"{prefix}.lora_A.weight"
key_b = f"{prefix}.lora_B.weight"
for lsd, coef in lora_sd_and_strengths:
if key_a not in lsd.sd or key_b not in lsd.sd:
continue
a = lsd.sd[key_a].to(device=device)
b = lsd.sd[key_b].to(device=device)
product = torch.matmul(b * coef, a)
del a, b
deltas.append(product.to(dtype=dtype))
if len(deltas) == 0:
return None
elif len(deltas) == 1:
return deltas[0]
return torch.sum(torch.stack(deltas, dim=0), dim=0)
def _ab_products() -> Iterator[LoraProduct]:
for lsd, coef in lora_sd_and_strengths:
if key_a not in lsd.sd or key_b not in lsd.sd:
continue
a = lsd.sd[key_a].to(device=device, dtype=dtype, non_blocking=True)
b = lsd.sd[key_b].to(device=device, dtype=dtype, non_blocking=True)
yield LoraProduct(a, b, coef)
return aggregate_lora_products(_ab_products(), dtype)
def _fuse_delta_with_scaled_fp8(
@@ -100,34 +154,9 @@ def _fuse_delta_with_scaled_fp8(
"""Dequantize scaled FP8 weight, add LoRA delta, and re-quantize."""
weight_scale = model_sd.sd[scale_key]
original_weight = weight.t().to(torch.float32) * weight_scale
original_weight = weight.to(torch.float32) * weight_scale
new_weight = original_weight + deltas.to(torch.float32)
new_fp8_weight, new_weight_scale = quantize_weight_to_fp8_per_tensor(new_weight)
return {key: new_fp8_weight, scale_key: new_weight_scale}
def _fuse_delta_with_cast_fp8(
deltas: torch.Tensor,
weight: torch.Tensor,
key: str,
target_dtype: torch.dtype,
) -> dict[str, torch.Tensor]:
"""Fuse LoRA delta with cast-only FP8 weight (no scale factor)."""
if str(weight.device).startswith("cuda"):
_fused_add_round_launch(deltas, weight, seed=0)
else:
deltas.add_(weight.to(dtype=deltas.dtype))
return {key: deltas.to(dtype=target_dtype)}
def _fuse_delta_with_bfloat16(
deltas: torch.Tensor,
weight: torch.Tensor,
key: str,
target_dtype: torch.dtype,
) -> dict[str, torch.Tensor]:
"""Fuse LoRA delta with bfloat16 weight."""
deltas.add_(weight)
return {key: deltas.to(dtype=target_dtype)}
@@ -1,72 +1,79 @@
# ruff: noqa: ANN001, ANN201, ERA001, N803, N806
import triton
import triton.language as tl
try:
import triton
import triton.language as tl
TRITON_AVAILABLE = True
except (ImportError, OSError):
TRITON_AVAILABLE = False
@triton.jit
def fused_add_round_kernel(
x_ptr,
output_ptr, # contents will be added to the output
seed,
n_elements,
EXPONENT_BIAS,
MANTISSA_BITS,
BLOCK_SIZE: tl.constexpr,
):
"""
A kernel to upcast 8bit quantized weights to bfloat16 with stochastic rounding
and add them to bfloat16 output weights. Might be used to upcast original model weights
and to further add them to precalculated deltas coming from LoRAs.
"""
# Get program ID and compute offsets
pid = tl.program_id(axis=0)
block_start = pid * BLOCK_SIZE
offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
if TRITON_AVAILABLE:
# Load data
x = tl.load(x_ptr + offsets, mask=mask)
rand_vals = tl.rand(seed, offsets) - 0.5
@triton.jit
def fused_add_round_kernel(
x_ptr,
output_ptr, # contents will be added to the output
seed,
n_elements,
EXPONENT_BIAS,
MANTISSA_BITS,
BLOCK_SIZE: tl.constexpr,
):
"""
A kernel to upcast 8bit quantized weights to bfloat16 with stochastic rounding
and add them to bfloat16 output weights. Might be used to upcast original model weights
and to further add them to precalculated deltas coming from LoRAs.
"""
# Get program ID and compute offsets
pid = tl.program_id(axis=0)
block_start = pid * BLOCK_SIZE
offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
x = tl.cast(x, tl.float16)
delta = tl.load(output_ptr + offsets, mask=mask)
delta = tl.cast(delta, tl.float16)
x = x + delta
# Load data
x = tl.load(x_ptr + offsets, mask=mask)
rand_vals = tl.rand(seed, offsets) - 0.5
x_bits = tl.cast(x, tl.int16, bitcast=True)
x = tl.cast(x, tl.float16)
delta = tl.load(output_ptr + offsets, mask=mask)
delta = tl.cast(delta, tl.float16)
x = x + delta
# Calculate the exponent. Unbiased fp16 exponent is ((x_bits & 0x7C00) >> 10) - 15 for
# normal numbers and -14 for subnormals.
fp16_exponent_bits = (x_bits & 0x7C00) >> 10
fp16_normals = fp16_exponent_bits > 0
fp16_exponent = tl.where(fp16_normals, fp16_exponent_bits - 15, -14)
x_bits = tl.cast(x, tl.int16, bitcast=True)
# Add the target dtype's exponent bias and clamp to the target dtype's exponent range.
exponent = fp16_exponent + EXPONENT_BIAS
MAX_EXPONENT = 2 * EXPONENT_BIAS + 1
exponent = tl.where(exponent > MAX_EXPONENT, MAX_EXPONENT, exponent)
exponent = tl.where(exponent < 0, 0, exponent)
# Calculate the exponent. Unbiased fp16 exponent is ((x_bits & 0x7C00) >> 10) - 15 for
# normal numbers and -14 for subnormals.
fp16_exponent_bits = (x_bits & 0x7C00) >> 10
fp16_normals = fp16_exponent_bits > 0
fp16_exponent = tl.where(fp16_normals, fp16_exponent_bits - 15, -14)
# Normal ULP exponent, expressed as an fp16 exponent field:
# (exponent - EXPONENT_BIAS - MANTISSA_BITS) + 15
# Simplifies to: fp16_exponent - MANTISSA_BITS + 15
# See https://en.wikipedia.org/wiki/Unit_in_the_last_place
eps_exp = tl.maximum(0, tl.minimum(31, exponent - EXPONENT_BIAS - MANTISSA_BITS + 15))
# Add the target dtype's exponent bias and clamp to the target dtype's exponent range.
exponent = fp16_exponent + EXPONENT_BIAS
MAX_EXPONENT = 2 * EXPONENT_BIAS + 1
exponent = tl.where(exponent > MAX_EXPONENT, MAX_EXPONENT, exponent)
exponent = tl.where(exponent < 0, 0, exponent)
# Calculate epsilon in the target dtype
eps_normal = tl.cast(tl.cast(eps_exp << 10, tl.int16), tl.float16, bitcast=True)
# Normal ULP exponent, expressed as an fp16 exponent field:
# (exponent - EXPONENT_BIAS - MANTISSA_BITS) + 15
# Simplifies to: fp16_exponent - MANTISSA_BITS + 15
# See https://en.wikipedia.org/wiki/Unit_in_the_last_place
eps_exp = tl.maximum(0, tl.minimum(31, exponent - EXPONENT_BIAS - MANTISSA_BITS + 15))
# Subnormal ULP: 2^(1 - EXPONENT_BIAS - MANTISSA_BITS) ->
# fp16 exponent bits: (1 - EXPONENT_BIAS - MANTISSA_BITS) + 15 =
# 16 - EXPONENT_BIAS - MANTISSA_BITS
eps_subnormal = tl.cast(tl.cast((16 - EXPONENT_BIAS - MANTISSA_BITS) << 10, tl.int16), tl.float16, bitcast=True)
eps = tl.where(exponent > 0, eps_normal, eps_subnormal)
# Calculate epsilon in the target dtype
eps_normal = tl.cast(tl.cast(eps_exp << 10, tl.int16), tl.float16, bitcast=True)
# Apply zero mask to epsilon
eps = tl.where(x == 0, 0.0, eps)
# Subnormal ULP: 2^(1 - EXPONENT_BIAS - MANTISSA_BITS) ->
# fp16 exponent bits: (1 - EXPONENT_BIAS - MANTISSA_BITS) + 15 =
# 16 - EXPONENT_BIAS - MANTISSA_BITS
eps_subnormal = tl.cast(tl.cast((16 - EXPONENT_BIAS - MANTISSA_BITS) << 10, tl.int16), tl.float16, bitcast=True)
eps = tl.where(exponent > 0, eps_normal, eps_subnormal)
# Apply stochastic rounding
output = tl.cast(x + rand_vals * eps, tl.bfloat16)
# Apply zero mask to epsilon
eps = tl.where(x == 0, 0.0, eps)
# Store the result
tl.store(output_ptr + offsets, output, mask=mask)
# Apply stochastic rounding
output = tl.cast(x + rand_vals * eps, tl.bfloat16)
# Store the result
tl.store(output_ptr + offsets, output, mask=mask)
@@ -13,6 +13,10 @@ if TYPE_CHECKING:
from ltx_core.loader.registry import Registry
# Per-key shape and dtype description for a flat collection of tensors.
TensorLayout = dict[str, tuple[torch.Size, torch.dtype]]
@dataclass(frozen=True)
class StateDict:
"""
@@ -52,7 +56,15 @@ class StateDictLoader(Protocol):
"""
class ModelBuilderProtocol(Protocol[ModelType]):
class BuilderProtocol(Protocol[ModelType]):
"""Protocol for model builders that produce a model via ``build()``."""
def build(
self, device: torch.device | None = None, dtype: torch.dtype | None = None, **kwargs: object
) -> ModelType: ...
class ModelBuilderProtocol(BuilderProtocol[ModelType], Protocol[ModelType]):
"""
Protocol for building PyTorch models from configuration dictionaries.
Implementations must provide:
@@ -102,7 +102,7 @@ class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType],
registry: Registry = field(default_factory=DummyRegistry)
lora_load_device: torch.device = field(default_factory=lambda: torch.device("cpu"))
def lora(self, lora_path: str, strength: float = 1.0, sd_ops: SDOps | None = None) -> "SingleGPUModelBuilder":
def lora(self, lora_path: str, strength: float, sd_ops: SDOps) -> "SingleGPUModelBuilder":
return replace(self, loras=(*self.loras, LoraPathStrengthAndSDOps(lora_path, strength, sd_ops)))
def with_sd_ops(self, sd_ops: SDOps | None) -> "SingleGPUModelBuilder":
@@ -144,7 +144,7 @@ class Attention(torch.nn.Module):
heads: int = 8,
dim_head: int = 64,
norm_eps: float = 1e-6,
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
rope_type: LTXRopeType = LTXRopeType.SPLIT,
attention_function: AttentionCallable | AttentionFunction = AttentionFunction.DEFAULT,
apply_gated_attention: bool = False,
) -> None:
@@ -57,7 +57,7 @@ class LTXModel(torch.nn.Module):
audio_cross_attention_dim: int = 2048,
audio_positional_embedding_max_pos: list[int] | None = None,
av_ca_timestep_scale_multiplier: int = 1,
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
rope_type: LTXRopeType = LTXRopeType.SPLIT,
double_precision_rope: bool = False,
apply_gated_attention: bool = False,
caption_projection: torch.nn.Module | None = None,
@@ -62,7 +62,7 @@ class LTXModelConfigurator(ModelConfigurator[LTXModel]):
audio_cross_attention_dim=config.get("audio_cross_attention_dim", 2048),
audio_positional_embedding_max_pos=config.get("audio_positional_embedding_max_pos", [20]),
av_ca_timestep_scale_multiplier=config.get("av_ca_timestep_scale_multiplier", 1),
rope_type=LTXRopeType(config.get("rope_type", "interleaved")),
rope_type=LTXRopeType(config.get("rope_type", "split")),
double_precision_rope=config.get("frequencies_precision", False) == "float64",
apply_gated_attention=config.get("apply_gated_attention", False),
caption_projection=caption_projection,
@@ -114,7 +114,7 @@ class LTXVideoOnlyModelConfigurator(ModelConfigurator[LTXModel]):
positional_embedding_max_pos=config.get("positional_embedding_max_pos", [20, 2048, 2048]),
timestep_scale_multiplier=config.get("timestep_scale_multiplier", 1000),
use_middle_indices_grid=config.get("use_middle_indices_grid", True),
rope_type=LTXRopeType(config.get("rope_type", "interleaved")),
rope_type=LTXRopeType(config.get("rope_type", "split")),
double_precision_rope=config.get("frequencies_precision", False) == "float64",
apply_gated_attention=config.get("apply_gated_attention", False),
caption_projection=caption_projection,
@@ -16,9 +16,10 @@ class LTXRopeType(Enum):
def apply_rotary_emb(
input_tensor: torch.Tensor,
freqs_cis: Tuple[torch.Tensor, torch.Tensor],
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
rope_type: LTXRopeType = LTXRopeType.SPLIT,
) -> torch.Tensor:
if rope_type == LTXRopeType.INTERLEAVED:
# Note: INTERLEAVED rope is a legacy mode. Prefer SPLIT instead.
return apply_interleaved_rotary_emb(input_tensor, *freqs_cis)
elif rope_type == LTXRopeType.SPLIT:
return apply_split_rotary_emb(input_tensor, *freqs_cis)
@@ -45,6 +46,11 @@ def apply_split_rotary_emb(
needs_reshape = False
if input_tensor.ndim != 4 and cos_freqs.ndim == 4:
b, h, t, _ = cos_freqs.shape
if input_tensor.shape[0] != b:
raise ValueError(
f"apply_split_rotary_emb: input_tensor batch ({input_tensor.shape[0]}) "
f"must equal cos_freqs batch ({b})."
)
input_tensor = input_tensor.reshape(b, t, h, -1).swapaxes(1, 2)
needs_reshape = True
@@ -183,7 +189,7 @@ def precompute_freqs_cis(
max_pos: list[int] | None = None,
use_middle_indices_grid: bool = False,
num_attention_heads: int = 32,
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
rope_type: LTXRopeType = LTXRopeType.SPLIT,
freq_grid_generator: Callable[[float, int, int, torch.device], torch.Tensor] = generate_freq_grid_pytorch,
) -> tuple[torch.Tensor, torch.Tensor]:
if max_pos is None:
@@ -27,7 +27,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
idx: int,
video: TransformerConfig | None = None,
audio: TransformerConfig | None = None,
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
rope_type: LTXRopeType = LTXRopeType.SPLIT,
norm_eps: float = 1e-6,
attention_function: AttentionFunction | AttentionCallable = AttentionFunction.DEFAULT,
):
@@ -1,5 +1,6 @@
"""Video VAE package."""
from ltx_core.model.video_vae.memory_efficient_decode import MEMORY_EFFICIENT_DECODE
from ltx_core.model.video_vae.model_configurator import (
VAE_DECODER_COMFY_KEYS_FILTER,
VAE_ENCODER_COMFY_KEYS_FILTER,
@@ -10,6 +11,7 @@ from ltx_core.model.video_vae.tiling import SpatialTilingConfig, TemporalTilingC
from ltx_core.model.video_vae.video_vae import VideoDecoder, VideoEncoder, get_video_chunks_number
__all__ = [
"MEMORY_EFFICIENT_DECODE",
"VAE_DECODER_COMFY_KEYS_FILTER",
"VAE_ENCODER_COMFY_KEYS_FILTER",
"SpatialTilingConfig",
@@ -0,0 +1,612 @@
"""Memory-efficient VAE decoder operations.
Reduces peak VRAM usage during video decoding through in-place operations
and workspace buffer reuse. The main optimizations are:
1. **Workspace buffers** -- Pre-allocated tensors with temporal padding replace
dynamic padding (``F.pad`` / ``concatenate``) in ``CausalConv3d``. A
workspace of shape ``[B, C, T+2, H, W]`` holds the data in positions
``[1:-1]`` with replicate padding at ``[0]`` and ``[-1]``.
2. **In-place temporal-chunked Conv3d** *(non-causal only)* -- The convolution
output is written back into the workspace buffer, avoiding a separate
output allocation. Temporal chunking with boundary save/restore ensures
correct reads despite in-place writes.
3. **In-place normalization and affine transforms** -- PixelNorm, scale/shift,
and SiLU are applied in-place on workspace views.
4. **Free-before-conv** -- For ``DepthToSpaceUpsample`` blocks the input
tensor is freed before the convolution runs so that peak VRAM never holds
input *and* output simultaneously.
Both causal and non-causal modes are supported. Non-causal mode benefits
from all four optimizations. Causal mode benefits from optimizations 1, 3,
and 4; in-place conv (2) is skipped because the asymmetric causal padding
layout prevents clean in-place overwrites.
Usage via the ``ModuleOps`` pattern (preferred)::
from ltx_core.model.video_vae import MEMORY_EFFICIENT_DECODE
builder = decoder_builder.with_module_ops(
(*decoder_builder.module_ops, MEMORY_EFFICIENT_DECODE)
)
Or applied directly to an existing decoder::
from ltx_core.model.video_vae.memory_efficient_decode import (
enable_memory_efficient_decode,
)
enable_memory_efficient_decode(decoder)
"""
from __future__ import annotations
import math
from typing import TYPE_CHECKING
import torch
from einops import rearrange
from torch import nn
from torch.nn import functional as F
from ltx_core.loader.module_ops import ModuleOps
from ltx_core.model.common.normalization import PixelNorm
from ltx_core.model.video_vae.convolution import CausalConv3d
from ltx_core.model.video_vae.ops import unpatchify
from ltx_core.model.video_vae.resnet import ResnetBlock3D, UNetMidBlock3D
from ltx_core.model.video_vae.sampling import DepthToSpaceUpsample
if TYPE_CHECKING:
from ltx_core.model.video_vae.video_vae import VideoDecoder
# ---------------------------------------------------------------------------
# Low-level helpers
# ---------------------------------------------------------------------------
def _find_temporal_split_size(num_frames: int) -> int:
"""Find chunk size for in-place temporal convolution.
The chunk size ensures the last chunk has at least 3 frames
(the temporal kernel size), avoiding degenerate chunks.
"""
for s in range(16, 2, -1):
remainder = num_frames % s
if remainder == 0 or remainder >= 3:
return s
raise ValueError(
f"Unable to find a valid temporal split size for num_frames={num_frames}. "
"Expected a split size between 3 and 16 such that the final chunk is "
"either exact or has at least 3 frames."
)
def _pad_workspace_temporal(workspace: torch.Tensor) -> None:
"""Apply non-causal replicate padding to temporal boundaries.
Sets ``workspace[:, :, 0]`` to a copy of ``workspace[:, :, 1]`` and
``workspace[:, :, -1]`` to a copy of ``workspace[:, :, -2]``.
"""
workspace[:, :, 0, :, :].copy_(workspace[:, :, 1, :, :])
workspace[:, :, -1, :, :].copy_(workspace[:, :, -2, :, :])
# ---------------------------------------------------------------------------
# In-place Conv3d (non-causal only)
# ---------------------------------------------------------------------------
def inplace_conv3d_temporal_chunked(workspace: torch.Tensor, conv: nn.Conv3d) -> None:
"""Run a 3x3x3 Conv3d in-place on a temporally-padded workspace.
The workspace has shape ``[B, C, T+2, H, W]`` where positions ``[1:-1]``
hold the real data and positions ``[0]`` and ``[-1]`` are padding slots.
The convolution must have ``kernel_size=(3,3,3)``, ``stride=(1,1,1)``,
``padding=(0,1,1)`` -- no temporal padding, symmetric spatial padding.
The output (T frames) overwrites positions ``[1:-1]``. Temporal chunking
with boundary save/restore ensures each chunk reads unmodified input even
though earlier chunks already wrote to the same buffer.
Only valid for **non-causal** mode (symmetric replicate padding).
Args:
workspace: Tensor ``[B, max(C_in, C_out), T+2, H, W]``.
Modified in-place; after the call ``workspace[:, :C_out, 1:-1]``
holds the convolution result.
conv: ``nn.Conv3d`` with the constraints above.
"""
if conv.kernel_size != (3, 3, 3):
raise ValueError(f"Expected kernel_size=(3,3,3), got {conv.kernel_size}")
if conv.stride != (1, 1, 1):
raise ValueError(f"Expected stride=(1,1,1), got {conv.stride}")
if conv.padding != (0, 1, 1):
raise ValueError(f"Expected padding=(0,1,1), got {conv.padding}")
_pad_workspace_temporal(workspace)
total_frames = workspace.shape[2]
out_channels = conv.out_channels
in_channels = conv.in_channels
if total_frames > 16:
split_size = _find_temporal_split_size(total_frames)
num_splits = (total_frames + split_size - 1) // split_size
else:
split_size = total_frames - 1
num_splits = 1
# 1-frame buffers for saving / restoring boundary frames across chunks.
x_buf = torch.empty(
workspace.shape[0],
workspace.shape[1],
1,
workspace.shape[3],
workspace.shape[4],
device=workspace.device,
dtype=workspace.dtype,
)
o_buf = torch.empty_like(x_buf)
# Helper: extract a chunk and make it contiguous. Workspace views can
# inherit strides > 2^31 from the full buffer, which makes Conv3d's
# reflect-padding path (F.pad) crash with "input tensor must fit into
# 32-bit index math". A small .clone() per chunk avoids this.
needs_clone = workspace.untyped_storage().nbytes() > (2**31 - 1) * workspace.element_size()
def _chunk(t_start: int, t_end: int) -> torch.Tensor:
s = workspace[:, :in_channels, t_start:t_end]
return s.clone() if needs_clone else s
# --- First chunk ---
if num_splits > 1:
# Save the boundary now so the loop below can restore it. Skipped
# when there is only one chunk: the loop never runs, and the save
# would be a wasted full HW slice copy.
x_buf[:, :, 0] = workspace[:, :, split_size - 1].clone()
workspace[:, :out_channels, 1:split_size] = conv(_chunk(0, split_size + 1))
# --- Remaining chunks ---
for i in range(1, num_splits):
start = i * split_size
end = min((i + 1) * split_size, total_frames - 1)
# Save the value at start-1 (now holds previous chunk's output).
o_buf[:, :, 0] = workspace[:, :, start - 1].clone()
# Restore the original input value needed by this chunk's conv.
workspace[:, :, start - 1] = x_buf[:, :, 0]
# Save the boundary for the *next* chunk before we overwrite it.
x_buf[:, :, 0] = workspace[:, :, end - 1].clone()
workspace[:, :out_channels, start:end] = conv(_chunk(start - 1, end + 1))
# Put back the previous chunk's output at the boundary.
workspace[:, :, start - 1] = o_buf[:, :, 0]
# ---------------------------------------------------------------------------
# Causal conv helper (free-before-conv)
# ---------------------------------------------------------------------------
def _causal_pad(x: torch.Tensor, pad_size: int) -> torch.Tensor:
"""Build a causal-padded buffer of shape ``[B, C, T+pad_size, H, W]``.
Copies ``x`` into ``padded[:, :, pad_size:]`` and replicates the first
real frame into the leading ``pad_size`` slots. The caller still owns
``x`` after this returns.
"""
padded = torch.empty(
x.shape[0],
x.shape[1],
x.shape[2] + pad_size,
x.shape[3],
x.shape[4],
device=x.device,
dtype=x.dtype,
)
padded[:, :, pad_size:].copy_(x)
for i in range(pad_size):
padded[:, :, i] = padded[:, :, pad_size]
return padded
def _causal_pad_free_and_conv(x: torch.Tensor, causal_conv: CausalConv3d) -> torch.Tensor:
"""Causal-pad *x*, free it, then run the raw ``nn.Conv3d``.
This avoids the peak where both the original and padded tensors are
live simultaneously (as happens inside ``CausalConv3d.forward``).
Args:
x: Input ``[B, C_in, T, H, W]``. **Deleted** inside this function;
the caller must not use it afterwards.
Returns:
Convolution output ``[B, C_out, T, H, W]``.
"""
padded = _causal_pad(x, causal_conv.time_kernel_size - 1)
del x
result = causal_conv.conv(padded)
del padded
return result
# ---------------------------------------------------------------------------
# In-place normalization
# ---------------------------------------------------------------------------
def _pixel_norm_inplace(x: torch.Tensor, eps: float = 1e-8) -> None:
"""In-place RMS (pixel) normalization along the channel dimension."""
rms = torch.sqrt(torch.mean(x**2, dim=1, keepdim=True) + eps)
x.div_(rms)
def _norm_inplace(norm: nn.Module, x: torch.Tensor) -> None:
"""Apply *norm* in-place, using an optimised path for ``PixelNorm``."""
if isinstance(norm, PixelNorm):
_pixel_norm_inplace(x, eps=norm.eps)
else:
# GroupNorm or other -- fall back to allocating a temporary.
result = norm(x)
x.copy_(result)
del result
# ---------------------------------------------------------------------------
# Per-block efficient forwards
# ---------------------------------------------------------------------------
def _resnet_block_forward_inplace(
resnet: ResnetBlock3D,
workspace: torch.Tensor,
causal: bool,
timestep: torch.Tensor | None,
generator: torch.Generator | None,
) -> None:
"""Run a ``ResnetBlock3D`` in-place on a workspace buffer.
The workspace has shape ``[B, C, T+2, H, W]`` with real data in
``[1:-1]``. After this call ``workspace[:, :, 1:-1]`` holds the
residual-branch output ``F(x)`` (without the skip connection --
the caller adds it back to the hidden state).
Only valid when ``in_channels == out_channels`` (true for all
``ResnetBlock3D`` instances inside a ``UNetMidBlock3D``).
"""
if resnet.in_channels != resnet.out_channels:
raise ValueError(
"In-place resnet forward requires in_channels == out_channels, "
f"got {resnet.in_channels} != {resnet.out_channels}"
)
interior = workspace[:, :, 1:-1]
# --- norm1 + [ada scaling] + SiLU + conv1 ---
_norm_inplace(resnet.norm1, interior)
if resnet.timestep_conditioning and timestep is not None:
ada = resnet.scale_shift_table[None, ..., None, None, None].to(
device=interior.device, dtype=interior.dtype
) + timestep.reshape(
interior.shape[0],
4,
-1,
timestep.shape[-3],
timestep.shape[-2],
timestep.shape[-1],
)
shift1, scale1, shift2, scale2 = ada.unbind(dim=1)
interior.mul_(1 + scale1).add_(shift1)
F.silu(interior, inplace=True)
if causal:
result = resnet.conv1(interior, causal=True)
interior.copy_(result)
del result
else:
inplace_conv3d_temporal_chunked(workspace, resnet.conv1.conv)
if resnet.inject_noise:
spatial_shape = interior.shape[-2:]
scale = resnet.per_channel_scale1.to(device=interior.device, dtype=interior.dtype)
noise = torch.randn(spatial_shape, device=interior.device, dtype=interior.dtype, generator=generator)
interior.add_((noise * scale)[None, :, None, ...])
# --- norm2 + [ada scaling] + SiLU + conv2 ---
_norm_inplace(resnet.norm2, interior)
if resnet.timestep_conditioning and timestep is not None:
interior.mul_(1 + scale2).add_(shift2) # type: ignore[possibly-undefined]
F.silu(interior, inplace=True)
# dropout is always 0.0 during inference -- skip.
if causal:
result = resnet.conv2(interior, causal=True)
interior.copy_(result)
del result
else:
inplace_conv3d_temporal_chunked(workspace, resnet.conv2.conv)
if resnet.inject_noise:
spatial_shape = interior.shape[-2:]
scale = resnet.per_channel_scale2.to(device=interior.device, dtype=interior.dtype)
noise = torch.randn(spatial_shape, device=interior.device, dtype=interior.dtype, generator=generator)
interior.add_((noise * scale)[None, :, None, ...])
def _midblock_forward_efficient(
block: UNetMidBlock3D,
hidden_states: torch.Tensor,
causal: bool,
timestep: torch.Tensor | None,
generator: torch.Generator | None,
) -> torch.Tensor:
"""Memory-efficient ``UNetMidBlock3D`` forward.
Allocates a single workspace buffer that is reused across all
``ResnetBlock3D`` iterations. For each block the workspace is
populated with the current hidden state, processed in-place, and
the result is added back (residual connection).
"""
timestep_embed = None
if block.timestep_conditioning:
if timestep is None:
raise ValueError("'timestep' required when timestep_conditioning=True")
batch_size = hidden_states.shape[0]
timestep_embed = block.time_embedder(
timestep=timestep.flatten(),
hidden_dtype=hidden_states.dtype,
)
timestep_embed = timestep_embed.view(batch_size, timestep_embed.shape[-1], 1, 1, 1)
workspace = torch.empty(
hidden_states.shape[0],
hidden_states.shape[1],
hidden_states.shape[2] + 2,
hidden_states.shape[3],
hidden_states.shape[4],
device=hidden_states.device,
dtype=hidden_states.dtype,
)
for resnet in block.res_blocks:
workspace[:, :, 1:-1].copy_(hidden_states)
_resnet_block_forward_inplace(resnet, workspace, causal, timestep_embed, generator)
hidden_states.add_(workspace[:, :, 1:-1])
del workspace
return hidden_states
def _upsample_forward_efficient(
block: DepthToSpaceUpsample,
x: torch.Tensor,
causal: bool,
) -> torch.Tensor:
"""Memory-efficient ``DepthToSpaceUpsample`` forward.
For non-causal mode the input is copied into a workspace and the
convolution runs in-place. For causal mode the input is manually
padded and freed before the convolution runs. Both paths avoid
the peak where input *and* output coexist.
"""
if block.residual:
x_in = rearrange(
x,
"b (c p1 p2 p3) d h w -> b c (d p1) (h p2) (w p3)",
p1=block.stride[0],
p2=block.stride[1],
p3=block.stride[2],
)
num_repeat = math.prod(block.stride) // block.out_channels_reduction_factor
x_in = x_in.repeat(1, num_repeat, 1, 1, 1)
if block.stride[0] == 2:
x_in = x_in[:, :, 1:, :, :]
conv = block.conv.conv # underlying nn.Conv3d inside CausalConv3d
in_channels = x.shape[1]
out_channels = conv.out_channels
if causal:
x = _causal_pad_free_and_conv(x, block.conv)
else:
workspace = torch.empty(
x.shape[0],
max(in_channels, out_channels),
x.shape[2] + 2,
x.shape[3],
x.shape[4],
device=x.device,
dtype=x.dtype,
)
workspace[:, :in_channels, 1:-1].copy_(x)
del x
inplace_conv3d_temporal_chunked(workspace, conv)
x = workspace[:, :out_channels, 1:-1].contiguous()
del workspace
x = rearrange(
x,
"b (c p1 p2 p3) d h w -> b c (d p1) (h p2) (w p3)",
p1=block.stride[0],
p2=block.stride[1],
p3=block.stride[2],
)
if block.stride[0] == 2:
x = x[:, :, 1:, :, :]
if block.residual:
x = x + x_in
del x_in
return x
# ---------------------------------------------------------------------------
# Final norm + conv_out
# ---------------------------------------------------------------------------
def _final_norm_and_conv_out(
decoder: VideoDecoder,
sample: torch.Tensor,
causal: bool,
scaled_timestep: torch.Tensor | None,
batch_size: int,
) -> torch.Tensor:
"""Workspace-based final norm + [ada] + SiLU + conv_out + unpatchify."""
conv_out_mod: CausalConv3d = decoder.conv_out # type: ignore[assignment]
conv_out = conv_out_mod.conv
feature_channels = sample.shape[1]
workspace = torch.empty(
sample.shape[0],
max(feature_channels, conv_out.out_channels),
sample.shape[2] + 2,
sample.shape[3],
sample.shape[4],
device=sample.device,
dtype=sample.dtype,
)
workspace[:, :feature_channels, 1:-1].copy_(sample)
del sample
interior = workspace[:, :feature_channels, 1:-1]
_norm_inplace(decoder.conv_norm_out, interior)
if decoder.timestep_conditioning:
embedded_timestep = decoder.last_time_embedder(
timestep=scaled_timestep.flatten(),
hidden_dtype=interior.dtype,
)
embedded_timestep = embedded_timestep.view(batch_size, embedded_timestep.shape[-1], 1, 1, 1)
ada_values = decoder.last_scale_shift_table[None, ..., None, None, None].to(
device=interior.device, dtype=interior.dtype
) + embedded_timestep.reshape(
batch_size,
2,
-1,
embedded_timestep.shape[-3],
embedded_timestep.shape[-2],
embedded_timestep.shape[-1],
)
shift, scale = ada_values.unbind(dim=1)
interior.mul_(1 + scale).add_(shift)
F.silu(interior, inplace=True)
if causal:
# Causal: build padded tensor directly from the interior view,
# then free the workspace before running the conv.
padded = _causal_pad(interior, conv_out_mod.time_kernel_size - 1)
del workspace, interior
result = conv_out(padded)
del padded
else:
inplace_conv3d_temporal_chunked(workspace, conv_out)
result = workspace[:, : conv_out.out_channels, 1:-1].contiguous()
del workspace, interior
return unpatchify(result, patch_size_hw=decoder.patch_size, patch_size_t=1)
# ---------------------------------------------------------------------------
# Top-level efficient decoder forward
# ---------------------------------------------------------------------------
def _memory_efficient_forward(
decoder: VideoDecoder,
sample: torch.Tensor,
timestep: torch.Tensor | None = None,
generator: torch.Generator | None = None,
) -> torch.Tensor:
"""Full memory-efficient ``VideoDecoder.forward`` replacement.
Orchestrates the entire decode through workspace-based operations:
``UNetMidBlock3D`` and ``DepthToSpaceUpsample`` blocks use efficient
paths; standalone ``ResnetBlock3D`` blocks fall back to the standard
forward. The final norm + ada + SiLU + conv_out is also workspace-based.
"""
causal = decoder.causal
batch_size = sample.shape[0]
sample = sample.to(next(decoder.parameters()).dtype)
# --- Noise injection and de-normalisation (identical to standard path) ---
if decoder.timestep_conditioning:
noise = (
torch.randn(sample.size(), generator=generator, dtype=sample.dtype, device=sample.device)
* decoder.decode_noise_scale
)
sample = noise + (1.0 - decoder.decode_noise_scale) * sample
sample = decoder.per_channel_statistics.un_normalize(sample)
if timestep is None and decoder.timestep_conditioning:
timestep = torch.full((batch_size,), decoder.decode_timestep, device=sample.device, dtype=sample.dtype)
# --- conv_in (latent tensor is small -- standard path is fine) ---
sample = decoder.conv_in(sample, causal=causal)
upscale_dtype = next(iter(decoder.up_blocks.parameters())).dtype
sample = sample.to(upscale_dtype)
scaled_timestep = None
if decoder.timestep_conditioning:
if timestep is None:
raise ValueError("'timestep' required when timestep_conditioning=True")
scaled_timestep = timestep * decoder.timestep_scale_multiplier.to(sample)
# --- Up blocks (dispatch to efficient path per block type) ---
for up_block in decoder.up_blocks:
if isinstance(up_block, UNetMidBlock3D):
sample = _midblock_forward_efficient(
up_block,
sample,
causal=causal,
timestep=scaled_timestep if decoder.timestep_conditioning else None,
generator=generator,
)
elif isinstance(up_block, DepthToSpaceUpsample):
sample = _upsample_forward_efficient(up_block, sample, causal=causal)
elif isinstance(up_block, ResnetBlock3D):
sample = up_block(sample, causal=causal, generator=generator)
else:
sample = up_block(sample, causal=causal)
return _final_norm_and_conv_out(decoder, sample, causal, scaled_timestep, batch_size)
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
def enable_memory_efficient_decode(decoder: nn.Module) -> nn.Module:
"""Patch a ``VideoDecoder`` to use the memory-efficient forward path.
The original ``forward`` is saved as ``decoder._original_forward`` so
that it can be restored later with :func:`disable_memory_efficient_decode`.
"""
# Import here to avoid circular dependency at module level.
from ltx_core.model.video_vae.video_vae import VideoDecoder # noqa: PLC0415
if not isinstance(decoder, VideoDecoder):
raise TypeError(f"Expected VideoDecoder, got {type(decoder).__name__}")
if hasattr(decoder, "_original_forward"):
return decoder
original_forward = decoder.forward
def efficient_forward(
sample: torch.Tensor,
timestep: torch.Tensor | None = None,
generator: torch.Generator | None = None,
) -> torch.Tensor:
return _memory_efficient_forward(decoder, sample, timestep, generator)
decoder._original_forward = original_forward # type: ignore[attr-defined]
decoder.forward = efficient_forward # type: ignore[assignment]
return decoder
def disable_memory_efficient_decode(decoder: nn.Module) -> nn.Module:
"""Restore the original ``forward`` method on a patched ``VideoDecoder``."""
if hasattr(decoder, "_original_forward"):
decoder.forward = decoder._original_forward # type: ignore[attr-defined]
del decoder._original_forward # type: ignore[attr-defined]
return decoder
def _is_video_decoder(model: nn.Module) -> bool:
"""Matcher for the ``MEMORY_EFFICIENT_DECODE`` module op."""
from ltx_core.model.video_vae.video_vae import VideoDecoder # noqa: PLC0415
return isinstance(model, VideoDecoder)
MEMORY_EFFICIENT_DECODE = ModuleOps(
name="memory_efficient_vae_decode",
matcher=_is_video_decoder,
mutator=enable_memory_efficient_decode,
)
@@ -64,6 +64,6 @@ class TilingConfig:
@classmethod
def default(cls) -> "TilingConfig":
return cls(
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),
spatial_config=SpatialTilingConfig(tile_size_in_pixels=768, tile_overlap_in_pixels=64),
temporal_config=TemporalTilingConfig(tile_size_in_frames=80, tile_overlap_in_frames=24),
)
@@ -259,6 +259,7 @@ class VideoEncoder(nn.Module):
Args:
sample: Input video (B, C, F, H, W). F should be 1 + 8*k (e.g., 1, 9, 17, 25, 33...).
If not, the encoder crops the last frames to the nearest valid length.
Should be normalized to [-1, 1] range before encoding.
Returns:
Normalized latent means (B, 128, F', H', W') where F' = 1+(F-1)/8, H' = H/32, W' = W/32.
Example: (B, 3, 33, 512, 512) -> (B, 128, 5, 16, 16).
@@ -605,8 +606,8 @@ class VideoDecoder(nn.Module):
# many video frames and pixels correspond to a single latent cell.
self.video_downscale_factors = SpatioTemporalScaleFactors(
time=8,
width=32,
height=32,
width=32,
)
self.patch_size = patch_size
@@ -905,34 +906,22 @@ class VideoDecoder(nn.Module):
latent: torch.Tensor,
tiling_config: TilingConfig | None = None,
generator: torch.Generator | None = None,
*,
output_dtype: torch.dtype = torch.uint8,
) -> Iterator[torch.Tensor]:
"""Decode a video latent tensor, yielding chunks ``[f, h, w, c]``.
"""Decode a video latent tensor, yielding float chunks ``[f, h, w, c]`` in ``[0, 1]``.
Subclasses (e.g. ``DistributedVideoDecoder``) may override this to
control eagerness or distribution across ranks.
Args:
output_dtype: Target dtype for output tensors. ``torch.uint8``
(default) maps the decoder's ``[-1, 1]`` output to
``[0, 255]``. Any floating dtype returns ``[0, 1]`` cast
to that dtype.
"""
def _convert(frames: torch.Tensor) -> torch.Tensor:
# rearrange materializes a new contiguous tensor for this permutation,
# so in-place ops below do not mutate the caller's data.
def to_rgb(frames: torch.Tensor) -> torch.Tensor:
video = rearrange(frames[0], "c f h w -> f h w c")
video.add_(1.0).mul_(0.5).clamp_(0.0, 1.0)
if output_dtype == torch.uint8:
return video.mul_(255.0).to(torch.uint8)
return video.to(output_dtype)
return video.add_(1.0).mul_(0.5).clamp_(0.0, 1.0)
if tiling_config is not None:
for frames in self.tiled_decode(latent, tiling_config, generator=generator):
yield _convert(frames)
yield to_rgb(frames)
else:
decoded = self(latent, generator=generator)
yield _convert(decoded)
yield to_rgb(decoded)
def _group_tiles_by_temporal_slice(self, tiles: List[Tile]) -> List[List[Tile]]:
"""Group tiles by their temporal output slice."""
@@ -3,12 +3,9 @@ from ltx_core.quantization.fp8_cast import (
UPCAST_DURING_INFERENCE,
UpcastWithStochasticRounding,
)
from ltx_core.quantization.fp8_scaled_mm import FP8_PREPARE_MODULE_OPS, FP8_TRANSPOSE_SD_OPS
from ltx_core.quantization.policy import QuantizationPolicy
__all__ = [
"FP8_PREPARE_MODULE_OPS",
"FP8_TRANSPOSE_SD_OPS",
"TRANSFORMER_LINEAR_DOWNCAST_MAP",
"UPCAST_DURING_INFERENCE",
"QuantizationPolicy",
@@ -1,5 +1,6 @@
import torch
from ltx_core.loader.kernels import TRITON_AVAILABLE
from ltx_core.loader.module_ops import ModuleOps
from ltx_core.loader.sd_ops import KeyValueOperationResult, SDOps
from ltx_core.model.transformer.model import LTXModel
@@ -7,8 +8,13 @@ from ltx_core.model.transformer.model import LTXModel
BLOCK_SIZE = 1024
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
def fused_add_round_launch(target_weight: torch.Tensor, original_weight: torch.Tensor, seed: int) -> torch.Tensor:
if not TRITON_AVAILABLE:
raise RuntimeError(
"fused_add_round_launch requires Triton, which is not available on this platform. "
"Callers should gate on ltx_core.loader.kernels.TRITON_AVAILABLE and use a "
"deterministic-rounding fallback instead."
)
import triton # noqa: PLC0415
from ltx_core.loader.kernels import fused_add_round_kernel # noqa: PLC0415
@@ -53,10 +59,13 @@ def _upcast_and_round(
"""
Upcast the weight to the given dtype and optionally apply stochastic rounding.
Input weight needs to have float8_e4m3fn or float8_e5m2 dtype.
Stochastic rounding is implemented via a Triton kernel. When Triton is not
available (e.g., on Windows), this falls back to deterministic (nearest)
rounding via ``weight.to(dtype)``.
"""
if not with_stochastic_rounding:
if not with_stochastic_rounding or not TRITON_AVAILABLE or weight.device.type != "cuda":
return weight.to(dtype)
return _fused_add_round_launch(torch.zeros_like(weight, dtype=dtype), weight, seed)
return fused_add_round_launch(torch.zeros_like(weight, dtype=dtype), weight, seed)
class Fp8CastLinear(torch.nn.Linear):
@@ -82,14 +91,26 @@ class Fp8CastLinear(torch.nn.Linear):
def _replace_fwd_with_upcast(layer: torch.nn.Linear, with_stochastic_rounding: bool = False, seed: int = 0) -> None:
"""
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.
instances. Forward remains defined at the class level, which is required for
torch.compile compatibility — instance-level closure monkey-patches cause
graph breaks.
Also retypes ``weight`` and ``bias`` to fp8 so the meta param dtype matches
the post-load tensor dtype (sd_ops downcasts checkpoint bf16 -> fp8 at load).
Block streaming relies on this to derive pool buffer layout from the meta
model without an eager checkpoint read.
"""
layer.__class__ = Fp8CastLinear
layer._with_stochastic_rounding = with_stochastic_rounding
layer._seed = seed
layer.weight = torch.nn.Parameter(
torch.empty(layer.weight.shape, dtype=torch.float8_e4m3fn, device=layer.weight.device),
requires_grad=layer.weight.requires_grad,
)
if layer.bias is not None:
layer.bias = torch.nn.Parameter(
torch.empty(layer.bias.shape, dtype=torch.float8_e4m3fn, device=layer.bias.device),
requires_grad=layer.bias.requires_grad,
)
def _amend_forward_with_upcast(
@@ -1,11 +1,21 @@
import json
import struct
from typing import Callable
import torch
from torch import nn
from ltx_core.loader.module_ops import ModuleOps
from ltx_core.loader.sd_ops import KeyValueOperationResult, SDOps
from ltx_core.model.transformer import LTXModel
from ltx_core.quantization.trtllm_scaled_usable import trtllm_scaled_mm_usable
def _read_safetensors_dtypes(path: str) -> dict[str, str]:
"""Return ``{tensor_name: dtype_string}`` from the safetensors header."""
with open(path, "rb") as f:
header_size = struct.unpack("<Q", f.read(8))[0]
header = json.loads(f.read(header_size).decode("utf-8"))
return {k: v["dtype"] for k, v in header.items() if k != "__metadata__"}
class FP8Linear(nn.Module):
@@ -25,11 +35,8 @@ class FP8Linear(nn.Module):
self.in_features = in_features
self.out_features = out_features
fp8_shape = (in_features, out_features)
self.weight = nn.Parameter(torch.empty(fp8_shape, dtype=torch.float8_e4m3fn, device=device))
# Weight scale for FP8 dequantization (shape matches checkpoint format)
self.weight = nn.Parameter(torch.empty((out_features, in_features), dtype=torch.float8_e4m3fn, device=device))
self.weight_scale = nn.Parameter(torch.empty((), dtype=torch.float32, device=device))
# Input scale for static quantization (pre-quantized checkpoints)
self.input_scale = nn.Parameter(torch.empty((), dtype=torch.float32, device=device))
if bias:
@@ -40,31 +47,38 @@ class FP8Linear(nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor:
origin_shape = x.shape
# Static quantization: use pre-computed scale
qinput, cur_input_scale = torch.ops.tensorrt_llm.static_quantize_e4m3_per_tensor(x, self.input_scale)
if trtllm_scaled_mm_usable():
qinput, cur_input_scale = torch.ops.tensorrt_llm.static_quantize_e4m3_per_tensor(x, self.input_scale)
if qinput.dim() == 3:
qinput = qinput.reshape(-1, qinput.shape[-1])
output = torch.ops.trtllm.cublas_scaled_mm(
qinput,
self.weight.t(),
scale_a=cur_input_scale,
scale_b=self.weight_scale,
bias=None,
out_dtype=x.dtype,
)
else:
# Clamp before cast: out-of-range values cast to NaN/saturated FP8, which
# produces black-screen output on some checkpoints (e.g. ltx-2-19b-dev-fp8).
fp8_min = torch.finfo(torch.float8_e4m3fn).min
fp8_max = torch.finfo(torch.float8_e4m3fn).max
qinput = torch.clamp(x * self.input_scale.reciprocal(), fp8_min, fp8_max).to(torch.float8_e4m3fn)
if qinput.dim() == 3:
qinput = qinput.reshape(-1, qinput.shape[-1])
output = torch._scaled_mm(
qinput,
self.weight.t(),
scale_a=self.input_scale,
scale_b=self.weight_scale,
out_dtype=x.dtype,
use_fast_accum=True,
)
# Flatten to 2D for matmul
if qinput.dim() == 3:
qinput = qinput.reshape(-1, qinput.shape[-1])
# FP8 scaled matmul
output = torch.ops.trtllm.cublas_scaled_mm(
qinput,
self.weight,
scale_a=cur_input_scale,
scale_b=self.weight_scale,
bias=None,
out_dtype=x.dtype,
)
# Add bias
if self.bias is not None:
bias = self.bias
if bias.dtype != output.dtype:
bias = bias.to(output.dtype)
output = output + bias
output = output + self.bias.to(output.dtype)
# Restore original shape
if output.dim() != len(origin_shape):
output_shape = list(origin_shape)
output_shape[-1] = output.shape[-1]
@@ -74,15 +88,7 @@ class FP8Linear(nn.Module):
def quantize_weight_to_fp8_per_tensor(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
Quantize a weight tensor to FP8 (float8_e4m3fn) using per-tensor scaling.
Args:
weight: The weight tensor to quantize (any dtype, will be cast to float32)
Returns:
Tuple of (quantized_weight, weight_scale):
- quantized_weight: FP8 tensor, transposed for cublas_scaled_mm
- weight_scale: Per-tensor scale factor (reciprocal of quantization scale)
"""
"""Quantize a weight tensor to ``float8_e4m3fn`` with a per-tensor scale."""
weight_fp32 = weight.to(torch.float32)
fp8_min = torch.finfo(torch.float8_e4m3fn).min
@@ -96,7 +102,6 @@ def quantize_weight_to_fp8_per_tensor(weight: torch.Tensor) -> tuple[torch.Tenso
weight_fp32: torch.Tensor, scale: torch.Tensor, fp8_min: torch.Tensor, fp8_max: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
quantized_weight = torch.clamp(weight_fp32 * scale, min=fp8_min, max=fp8_max).to(torch.float8_e4m3fn)
quantized_weight = quantized_weight.t()
weight_scale = scale.reciprocal()
return quantized_weight, weight_scale
@@ -104,36 +109,8 @@ def quantize_weight_to_fp8_per_tensor(weight: torch.Tensor) -> tuple[torch.Tenso
return quantized_weight, weight_scale
def _should_skip_layer(layer_name: str, excluded_layer_substrings: tuple[str, ...]) -> bool:
return any(substring in layer_name for substring in excluded_layer_substrings)
EXCLUDED_LAYER_SUBSTRINGS = (
"patchify_proj",
"adaln_single",
"av_ca_video_scale_shift_adaln_single",
"av_ca_a2v_gate_adaln_single",
"caption_projection",
"proj_out",
"audio_patchify_proj",
"audio_adaln_single",
"av_ca_audio_scale_shift_adaln_single",
"av_ca_v2a_gate_adaln_single",
"audio_caption_projection",
"audio_proj_out",
"transformer_blocks.0.",
*[f"transformer_blocks.{i}." for i in range(43, 48)],
)
def _linear_to_fp8linear(layer: nn.Linear) -> FP8Linear:
"""
Create an FP8Linear layer from an nn.Linear layer.
Args:
layer: The nn.Linear layer to convert (typically on meta device)
Returns:
A new FP8Linear with the same configuration
"""
"""Create an ``FP8Linear`` matching the shape/bias of *layer*."""
return FP8Linear(
in_features=layer.in_features,
out_features=layer.out_features,
@@ -142,15 +119,14 @@ def _linear_to_fp8linear(layer: nn.Linear) -> FP8Linear:
)
def _apply_fp8_prepare_to_model(model: nn.Module, excluded_layer_substrings: tuple[str, ...]) -> nn.Module:
"""Replace nn.Linear layers with FP8Linear in the module tree."""
def _swap_linears_to_fp8(model: nn.Module, should_swap: Callable[[str], bool]) -> nn.Module:
"""Replace nn.Linear layers with FP8Linear where ``should_swap(name)`` returns True."""
replacements: list[tuple[nn.Module, str, nn.Linear]] = []
for name, module in model.named_modules():
if not isinstance(module, nn.Linear) or isinstance(module, FP8Linear):
continue
if _should_skip_layer(name, excluded_layer_substrings):
if not should_swap(name):
continue
if "." in name:
@@ -168,40 +144,32 @@ def _apply_fp8_prepare_to_model(model: nn.Module, excluded_layer_substrings: tup
return model
def _create_transpose_kv_operation(
excluded_layer_substrings: tuple[str, ...],
) -> Callable[[str, torch.Tensor], list[KeyValueOperationResult]]:
def transpose_if_matches(key: str, value: torch.Tensor) -> list[KeyValueOperationResult]:
# Only process .weight keys
if not key.endswith(".weight"):
return [KeyValueOperationResult(key, value)]
def get_fp8_swap_module_ops(checkpoint_path: str) -> tuple[ModuleOps, ...]:
"""Return the FP8 swap ``ModuleOps`` for layers whose ``.weight`` is ``F8_E4M3``
and which have a sibling ``.weight_scale`` tensor in the checkpoint.
Raises ``ValueError`` if no such layers are found — that combination is ambiguous
(a BF16 checkpoint with this policy would load as a no-op).
"""
dtypes = _read_safetensors_dtypes(checkpoint_path)
fp8_scale_paths = frozenset(
key.removesuffix(".weight_scale")
for key in dtypes
if key.endswith(".weight_scale") and dtypes.get(key.removesuffix(".weight_scale") + ".weight") == "F8_E4M3"
)
if not fp8_scale_paths:
raise ValueError(
f"fp8_scaled_mm requires a pre-quantized checkpoint with F8_E4M3 .weight + .weight_scale "
f"tensors, but {checkpoint_path!r} has none. Use QuantizationPolicy.fp8_cast() for BF16 checkpoints."
)
# Only transpose 2D FP8 tensors (Linear weights)
if value.dim() != 2 or value.dtype != torch.float8_e4m3fn:
return [KeyValueOperationResult(key, value)]
def _should_swap(name: str) -> bool:
suffix = "." + name
return any(p == name or p.endswith(suffix) for p in fp8_scale_paths)
# Check if the layer is excluded
layer_name = key.rsplit(".weight", 1)[0]
if _should_skip_layer(layer_name, excluded_layer_substrings):
return [KeyValueOperationResult(key, value)]
# Transpose to cuBLAS layout (in, out)
transposed_weight = value.t()
return [KeyValueOperationResult(key, transposed_weight)]
return transpose_if_matches
FP8_TRANSPOSE_SD_OPS = SDOps("fp8_transpose_weights").with_kv_operation(
_create_transpose_kv_operation(EXCLUDED_LAYER_SUBSTRINGS),
key_prefix="transformer_blocks.",
key_suffix=".weight",
)
FP8_PREPARE_MODULE_OPS = ModuleOps(
name="fp8_prepare_for_loading",
matcher=lambda model: isinstance(model, LTXModel),
mutator=lambda model: _apply_fp8_prepare_to_model(model, EXCLUDED_LAYER_SUBSTRINGS),
)
return (
ModuleOps(
name="fp8_swap_linears",
matcher=lambda model: isinstance(model, LTXModel),
mutator=lambda model: _swap_linears_to_fp8(model, _should_swap),
),
)
@@ -1,39 +1,48 @@
from dataclasses import dataclass
from enum import Enum
from ltx_core.loader.module_ops import ModuleOps
from ltx_core.loader.sd_ops import SDOps
from ltx_core.quantization.fp8_cast import TRANSFORMER_LINEAR_DOWNCAST_MAP, UPCAST_DURING_INFERENCE
from ltx_core.quantization.fp8_scaled_mm import FP8_PREPARE_MODULE_OPS, FP8_TRANSPOSE_SD_OPS
from ltx_core.quantization.fp8_scaled_mm import get_fp8_swap_module_ops
@dataclass(frozen=True)
class QuantizationPolicy:
"""Configuration for model quantization during loading.
Attributes:
sd_ops: State dict operations for weight transformation.
module_ops: Post-load module transformations.
kind: Discriminator for the policy variant.
sd_ops: State-dict operations applied to each tensor during load.
module_ops: Post-load module transformations applied to the meta model.
"""
class Kind(str, Enum):
FP8_CAST = "fp8_cast"
FP8_SCALED_MM = "fp8_scaled_mm"
kind: Kind
sd_ops: SDOps | None = None
module_ops: tuple[ModuleOps, ...] = ()
@classmethod
def fp8_cast(cls) -> "QuantizationPolicy":
"""Create policy using FP8 casting with upcasting during inference."""
"""FP8 casting with upcasting during inference."""
return cls(
kind=cls.Kind.FP8_CAST,
sd_ops=TRANSFORMER_LINEAR_DOWNCAST_MAP,
module_ops=(UPCAST_DURING_INFERENCE,),
)
@classmethod
def fp8_scaled_mm(cls) -> "QuantizationPolicy":
"""Create policy using FP8 scaled matrix multiplication."""
try:
import tensorrt_llm # noqa: F401, PLC0415
except ImportError as e:
raise ImportError("tensorrt_llm is not installed, skipping FP8 scaled MM quantization") from e
def fp8_scaled_mm(cls, checkpoint_path: str) -> "QuantizationPolicy":
"""FP8 scaled matmul for checkpoints pre-quantized with per-tensor scales.
The set of layers to swap to ``FP8Linear`` is discovered from the
checkpoint's ``.weight_scale`` tensors via suffix-matching against the
model's named modules. Requires a pre-quantized checkpoint; for BF16
checkpoints, use :meth:`fp8_cast` instead.
"""
return cls(
sd_ops=FP8_TRANSPOSE_SD_OPS,
module_ops=(FP8_PREPARE_MODULE_OPS,),
kind=cls.Kind.FP8_SCALED_MM,
sd_ops=None,
module_ops=get_fp8_swap_module_ops(checkpoint_path),
)
@@ -0,0 +1,37 @@
"""Runtime detection of TensorRT-LLM FP8 scaled-matmul availability.
When the TRT-LLM ops are usable on the current host (Linux + Hopper-class CUDA
+ tensorrt_llm wheel installed) we use them since they outperform the PyTorch-native
``torch._scaled_mm`` path. Otherwise we fall back to the native implementation,
which is portable across platforms (Windows, macOS, AMD GPUs).
The check runs once and is cached.
"""
from __future__ import annotations
import platform
from functools import cache
import torch
@cache
def trtllm_scaled_mm_usable() -> bool:
if platform.system() != "Linux":
return False
if not torch.cuda.is_available():
return False
major, minor = torch.cuda.get_device_capability()
sm = major * 10 + minor
if sm < 90 or sm >= 120:
return False
# The import is load-bearing — registers the trtllm torch ops as a side effect.
try:
import tensorrt_llm # noqa: F401, PLC0415
except Exception:
return False
return True
@@ -18,7 +18,7 @@ class _BasicTransformerBlock1D(torch.nn.Module):
dim: int,
heads: int,
dim_head: int,
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
rope_type: LTXRopeType = LTXRopeType.SPLIT,
apply_gated_attention: bool = False,
):
super().__init__()
@@ -39,7 +39,7 @@ class _BasicTransformerBlock1D(torch.nn.Module):
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
additive_attention_mask: torch.Tensor | None = None,
pe: torch.Tensor | None = None,
) -> torch.Tensor:
# Notice that normalization is always applied before the real computation in the following blocks.
@@ -49,8 +49,8 @@ class _BasicTransformerBlock1D(torch.nn.Module):
norm_hidden_states = norm_hidden_states.squeeze(1)
# 2. Self-Attention
attn_output = self.attn1(norm_hidden_states, mask=attention_mask, pe=pe)
# 2. Self-Attention — `mask` is the kernel-boundary name for the additive mask.
attn_output = self.attn1(norm_hidden_states, mask=additive_attention_mask, pe=pe)
hidden_states = attn_output + hidden_states
if hidden_states.ndim == 4:
@@ -84,7 +84,7 @@ class Embeddings1DConnector(torch.nn.Module):
causal_temporal_positioning (bool): If True, uses causal attention (default=False).
num_learnable_registers (int | None): Number of learnable registers to replace padded tokens. If None, disables
register replacement. (default=128)
rope_type (LTXRopeType): The RoPE variant to use (default=DEFAULT_ROPE_TYPE).
rope_type (LTXRopeType): The RoPE variant to use.
double_precision_rope (bool): Use double precision rope calculation (default=False).
"""
@@ -99,7 +99,7 @@ class Embeddings1DConnector(torch.nn.Module):
positional_embedding_max_pos: list[int] | None = None,
causal_temporal_positioning: bool = False,
num_learnable_registers: int | None = 128,
rope_type: LTXRopeType = LTXRopeType.INTERLEAVED,
rope_type: LTXRopeType = LTXRopeType.SPLIT,
double_precision_rope: bool = False,
apply_gated_attention: bool = False,
):
@@ -133,51 +133,40 @@ class Embeddings1DConnector(torch.nn.Module):
)
def _replace_padded_with_learnable_registers(
self, hidden_states: torch.Tensor, attention_mask: torch.Tensor
self, hidden_states: torch.Tensor, additive_attention_mask: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
assert hidden_states.shape[1] % self.num_learnable_registers == 0, (
f"Hidden states sequence length {hidden_states.shape[1]} must be divisible by num_learnable_registers "
f"{self.num_learnable_registers}."
)
batch_size, seq_len, _ = hidden_states.shape
num_registers_duplications = hidden_states.shape[1] // self.num_learnable_registers
learnable_registers = torch.tile(self.learnable_registers, (num_registers_duplications, 1))
attention_mask_binary = (attention_mask.squeeze(1).squeeze(1).unsqueeze(-1) >= -9000.0).int()
assert seq_len % self.num_learnable_registers == 0
non_zero_hidden_states = hidden_states[:, attention_mask_binary.squeeze().bool(), :]
non_zero_nums = non_zero_hidden_states.shape[1]
pad_length = hidden_states.shape[1] - non_zero_nums
adjusted_hidden_states = torch.nn.functional.pad(non_zero_hidden_states, pad=(0, 0, 0, pad_length), value=0)
flipped_mask = torch.flip(attention_mask_binary, dims=[1])
hidden_states = flipped_mask * adjusted_hidden_states + (1 - flipped_mask) * learnable_registers
registers = self.learnable_registers.repeat(seq_len // self.num_learnable_registers, 1).to(hidden_states.dtype)
registers = registers.unsqueeze(0).expand(batch_size, -1, -1) # (B, seq_len, hidden_dim)
binary_mask = additive_attention_mask[:, 0, 0, :].unsqueeze(-1) >= 0
binary_mask = binary_mask.to(hidden_states.dtype)
hidden_states = binary_mask * hidden_states + (1 - binary_mask) * registers
attention_mask = torch.full_like(
attention_mask,
0.0,
dtype=attention_mask.dtype,
device=attention_mask.device,
)
return hidden_states, attention_mask
return hidden_states, torch.zeros_like(additive_attention_mask)
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
additive_attention_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Forward pass of Embeddings1DConnector.
"""Forward pass of Embeddings1DConnector.
Args:
hidden_states (torch.Tensor): Input tensor of embeddings (shape [batch, seq_len, feature_dim]).
attention_mask (torch.Tensor|None): Optional mask for valid tokens (shape compatible with hidden_states).
hidden_states: (B, S, D) input embeddings.
additive_attention_mask: optional additive mask of shape (B, 1, 1, S), where
valid = 0.0 and padding = -torch.finfo(dtype).max.
Returns:
tuple[torch.Tensor, torch.Tensor]: Processed features and the corresponding (possibly modified) mask.
(hidden_states, additive_attention_mask)
"""
if self.num_learnable_registers:
hidden_states, attention_mask = self._replace_padded_with_learnable_registers(hidden_states, attention_mask)
hidden_states, additive_attention_mask = self._replace_padded_with_learnable_registers(
hidden_states, additive_attention_mask
)
indices_grid = torch.arange(hidden_states.shape[1], dtype=torch.float32, device=hidden_states.device)
indices_grid = indices_grid[None, None, :]
indices_grid = indices_grid[None, None, :].expand(hidden_states.shape[0], -1, -1)
freq_grid_generator = generate_freq_grid_np if self.double_precision_rope else generate_freq_grid_pytorch
freqs_cis = precompute_freqs_cis(
indices_grid=indices_grid,
@@ -191,11 +180,11 @@ class Embeddings1DConnector(torch.nn.Module):
)
for block in self.transformer_1d_blocks:
hidden_states = block(hidden_states, attention_mask=attention_mask, pe=freqs_cis)
hidden_states = block(hidden_states, additive_attention_mask=additive_attention_mask, pe=freqs_cis)
hidden_states = rms_norm(hidden_states)
return hidden_states, attention_mask
return hidden_states, additive_attention_mask
class Embeddings1DConnectorConfigurator(ModelConfigurator[Embeddings1DConnector]):
@@ -204,7 +193,7 @@ class Embeddings1DConnectorConfigurator(ModelConfigurator[Embeddings1DConnector]
@classmethod
def from_config(cls: type[Embeddings1DConnector], config: dict) -> Embeddings1DConnector:
transformer_config = config.get("transformer", {})
rope_type = LTXRopeType(transformer_config.get("rope_type", "interleaved"))
rope_type = LTXRopeType(transformer_config.get("rope_type", "split"))
double_precision_rope = transformer_config.get("frequencies_precision", False) == "float64"
pe_max_pos = transformer_config.get("connector_positional_embedding_max_pos", [1])
@@ -231,7 +220,7 @@ class AudioEmbeddings1DConnectorConfigurator(ModelConfigurator[Embeddings1DConne
@classmethod
def from_config(cls: type[Embeddings1DConnector], config: dict) -> Embeddings1DConnector:
transformer_config = config.get("transformer", {})
rope_type = LTXRopeType(transformer_config.get("rope_type", "interleaved"))
rope_type = LTXRopeType(transformer_config.get("rope_type", "split"))
double_precision_rope = transformer_config.get("frequencies_precision", False) == "float64"
pe_max_pos = transformer_config.get("connector_positional_embedding_max_pos", [1])
@@ -19,12 +19,32 @@ def convert_to_additive_mask(attention_mask: torch.Tensor, dtype: torch.dtype) -
) * torch.finfo(dtype).max
def _to_binary_mask(encoded: torch.Tensor, encoded_mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Convert connector output mask to binary mask and apply to encoded tensor."""
binary_mask = (encoded_mask < 0.000001).to(torch.int64)
binary_mask = binary_mask.reshape([encoded.shape[0], encoded.shape[1], 1])
encoded = encoded * binary_mask
return encoded, binary_mask
def _compute_right_pad_order(additive_mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Compute the index permutation that places valid tokens before pads in each row.
Stable sort: valid tokens keep their relative order. Idempotent for inputs already
right-padded. The sort and reordered mask depend only on the mask, so they can be
computed once and reused across multiple feature tensors that share the mask.
Args:
additive_mask: (B, 1, 1, S) additive mask, ``0.0`` for valid, ``-finfo.max`` for pad.
Returns:
``(sort_idx, reordered_additive_mask)``: ``sort_idx`` is (B, S); the reordered mask
has the same shape as the input.
"""
binary = (additive_mask[:, 0, 0, :] >= 0).to(torch.int32) # (B, S)
sort_idx = torch.argsort(binary, dim=-1, descending=True, stable=True) # (B, S)
new_binary = torch.gather(binary, 1, sort_idx)
new_additive = (new_binary.to(additive_mask.dtype) - 1) * torch.finfo(additive_mask.dtype).max
return sort_idx, new_additive[:, None, None, :]
def _apply_right_pad_order(features: torch.Tensor, sort_idx: torch.Tensor) -> torch.Tensor:
"""Apply a precomputed right-pad permutation (from ``_compute_right_pad_order``) to features."""
return torch.gather(features, 1, sort_idx.unsqueeze(-1).expand_as(features))
def _to_binary_mask(encoded_mask: torch.Tensor, lead_shape: tuple[int, int]) -> torch.Tensor:
"""Convert connector output mask to a binary (0/1) mask shaped ``(B, S, 1)`` for broadcasting."""
return (encoded_mask < 0.000001).to(torch.int64).reshape([lead_shape[0], lead_shape[1], 1])
class EmbeddingsProcessor(nn.Module):
@@ -57,12 +77,19 @@ class EmbeddingsProcessor(nn.Module):
if self.audio_connector is None and audio_features is not None:
raise ValueError("Audio features were provided but no audio connector is configured.")
video_encoded, video_mask = self.video_connector(video_features, additive_attention_mask)
video_encoded, binary_mask = _to_binary_mask(video_encoded, video_mask)
# Connectors expect right-padded input ([valid, pad]). Normalize layout here so the
# upstream tokenizer can keep using either side without coupling to the connector.
# The sort index depends only on the mask, so compute it once and reuse for audio.
sort_idx, mask_for_connector = _compute_right_pad_order(additive_attention_mask)
video_features = _apply_right_pad_order(video_features, sort_idx)
video_encoded, video_mask = self.video_connector(video_features, mask_for_connector)
binary_mask = _to_binary_mask(video_mask, video_encoded.shape[:2])
video_encoded = video_encoded * binary_mask
audio_encoded = None
if self.audio_connector is not None:
audio_encoded, _ = self.audio_connector(audio_features, additive_attention_mask)
audio_features = _apply_right_pad_order(audio_features, sort_idx)
audio_encoded, _ = self.audio_connector(audio_features, mask_for_connector)
return video_encoded, audio_encoded, binary_mask.squeeze(-1)
@@ -11,37 +11,25 @@ from torch import nn
def _norm_and_concat_padded_batch(
encoded_text: torch.Tensor,
sequence_lengths: torch.Tensor,
padding_side: str = "right",
attention_mask: torch.Tensor,
) -> torch.Tensor:
"""Normalize and flatten multi-layer hidden states, respecting padding.
Performs per-batch, per-layer normalization using masked mean and range,
then concatenates across the layer dimension.
then concatenates across the layer dimension. Padding-side agnostic: the
binary ``attention_mask`` already encodes which positions are valid.
Args:
encoded_text: Hidden states of shape [batch, seq_len, hidden_dim, num_layers].
sequence_lengths: Number of valid (non-padded) tokens per batch item.
padding_side: Whether padding is on "left" or "right".
attention_mask: Binary mask of shape [batch, seq_len], 1 for valid tokens, 0 for padding.
Returns:
Normalized tensor of shape [batch, seq_len, hidden_dim * num_layers],
with padded positions zeroed out.
"""
b, t, d, l = encoded_text.shape # noqa: E741
device = encoded_text.device
token_indices = torch.arange(t, device=device)[None, :]
if padding_side == "right":
mask = token_indices < sequence_lengths[:, None]
elif padding_side == "left":
start_indices = t - sequence_lengths[:, None]
mask = token_indices >= start_indices
else:
raise ValueError(f"padding_side must be 'left' or 'right', got {padding_side}")
mask = rearrange(mask, "b t -> b t 1 1")
b, _, d, l = encoded_text.shape # noqa: E741
eps = 1e-6
sequence_lengths = attention_mask.sum(dim=-1)
mask = rearrange(attention_mask.bool(), "b t -> b t 1 1")
masked = encoded_text.masked_fill(~mask, 0.0)
denom = (sequence_lengths * d).view(b, 1, 1, 1)
mean = masked.sum(dim=(1, 2), keepdim=True) / (denom + eps)
@@ -51,12 +39,10 @@ def _norm_and_concat_padded_batch(
range_ = x_max - x_min
normed = 8 * (encoded_text - mean) / (range_ + eps)
normed = normed.reshape(b, t, -1)
normed = normed.reshape(b, -1, d * l)
mask_flattened = rearrange(mask, "b t 1 1 -> b t 1").expand(-1, -1, d * l)
normed = normed.masked_fill(~mask_flattened, 0.0)
return normed
return normed.masked_fill(~mask_flattened, 0.0)
def norm_and_concat_per_token_rms(
@@ -97,12 +83,14 @@ class FeatureExtractorV1(nn.Module):
self.is_av = is_av
def forward(
self, hidden_states: torch.Tensor, attention_mask: torch.Tensor, padding_side: str = "left"
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
padding_side: str = "left", # noqa: ARG002 — kept for API stability; norm is layout-agnostic
) -> tuple[torch.Tensor, torch.Tensor | None]:
encoded = torch.stack(hidden_states, dim=-1) if isinstance(hidden_states, (list, tuple)) else hidden_states
dtype = encoded.dtype
sequence_lengths = attention_mask.sum(dim=-1)
normed = _norm_and_concat_padded_batch(encoded, sequence_lengths, padding_side)
normed = _norm_and_concat_padded_batch(encoded, attention_mask)
features = self.aggregate_embed(normed.to(dtype))
if self.is_av:
return features, features
@@ -1,6 +1,13 @@
from enum import Enum
from transformers import AutoTokenizer
class PaddingSide(str, Enum):
LEFT = "left"
RIGHT = "right"
class LTXVGemmaTokenizer:
"""
Tokenizer wrapper for Gemma models compatible with LTXV processes.
@@ -8,18 +15,18 @@ class LTXVGemmaTokenizer:
ensuring correct settings and output formatting for downstream consumption.
"""
def __init__(self, tokenizer_path: str, max_length: int = 256):
def __init__(self, tokenizer_path: str, max_length: int = 256, padding_side: PaddingSide = PaddingSide.LEFT):
"""
Initialize the tokenizer.
Args:
tokenizer_path (str): Path to the pretrained tokenizer files or model directory.
max_length (int, optional): Max sequence length for encoding. Defaults to 256.
padding_side (PaddingSide, optional): Side to pad on. Defaults to ``PaddingSide.LEFT``.
"""
self.tokenizer = AutoTokenizer.from_pretrained(
tokenizer_path, local_files_only=True, model_max_length=max_length
)
# Gemma expects left padding for chat-style prompts; for plain text it doesn't matter much.
self.tokenizer.padding_side = "left"
self.tokenizer.padding_side = padding_side.value
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
+1 -1
View File
@@ -138,7 +138,7 @@ class VideoLatentTools(LatentTools):
LatentState(
latent=initial_latent,
denoise_mask=denoise_mask,
positions=positions.to(dtype),
positions=positions,
clean_latent=clean_latent,
)
)
+7 -5
View File
@@ -20,15 +20,17 @@ class SpatioTemporalScaleFactors(NamedTuple):
"""
Describes the spatiotemporal downscaling between decoded video space and
the corresponding VAE latent grid.
Field order matches the (frame/time, height, width) axis layout used by
latent tensors and meshgrid coordinates elsewhere in the codebase.
"""
time: int
width: int
height: int
width: int
@classmethod
def default(cls) -> "SpatioTemporalScaleFactors":
return cls(time=8, width=32, height=32)
return cls(time=8, height=32, width=32)
VIDEO_SCALE_FACTORS = SpatioTemporalScaleFactors.default()
@@ -74,9 +76,9 @@ class VideoLatentShape(NamedTuple):
latent_channels: int = 128,
scale_factors: SpatioTemporalScaleFactors = VIDEO_SCALE_FACTORS,
) -> "VideoLatentShape":
frames = (shape.frames - 1) // scale_factors[0] + 1
height = shape.height // scale_factors[1]
width = shape.width // scale_factors[2]
frames = (shape.frames - 1) // scale_factors.time + 1
height = shape.height // scale_factors.height
width = shape.width // scale_factors.width
return VideoLatentShape(
batch=shape.batch,
+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)
@@ -166,6 +166,11 @@ acceleration:
# Useful when GPU memory is limited
load_text_encoder_in_8bit: false
# Offload optimizer state to CPU during validation video sampling.
# Helps avoid OOM when VAE decoder + transformer + optimizer state can't coexist
# on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP.
offload_optimizer_during_validation: false
# -----------------------------------------------------------------------------
# Data Configuration
# -----------------------------------------------------------------------------
@@ -178,6 +178,11 @@ acceleration:
# Useful when GPU memory is limited
load_text_encoder_in_8bit: true
# Offload optimizer state to CPU during validation video sampling.
# Helps avoid OOM when VAE decoder + transformer + optimizer state can't coexist
# on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP.
offload_optimizer_during_validation: true
# -----------------------------------------------------------------------------
# Data Configuration
# -----------------------------------------------------------------------------
@@ -166,6 +166,11 @@ acceleration:
# Useful when GPU memory is limited
load_text_encoder_in_8bit: false
# Offload optimizer state to CPU during validation video sampling.
# Helps avoid OOM when VAE decoder + transformer + optimizer state can't coexist
# on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP.
offload_optimizer_during_validation: false
# -----------------------------------------------------------------------------
# Data Configuration
# -----------------------------------------------------------------------------
@@ -215,18 +215,20 @@ Hardware acceleration and compute optimization settings.
```yaml
acceleration:
mixed_precision_mode: "bf16" # "no", "fp16", or "bf16"
quantization: null # Quantization options
load_text_encoder_in_8bit: false # Load text encoder in 8-bit
mixed_precision_mode: "bf16" # "no", "fp16", or "bf16"
quantization: null # Quantization options
load_text_encoder_in_8bit: false # Load text encoder in 8-bit
offload_optimizer_during_validation: false # Offload optimizer state to CPU during validation
```
**Key parameters:**
| Parameter | Description |
|-----------------------------|------------------------------------------------------------------------------------|
| `mixed_precision_mode` | Precision mode - `"bf16"` recommended for modern GPUs |
| `quantization` | Model quantization: `null`, `"int8-quanto"`, `"int4-quanto"`, `"fp8-quanto"`, etc. |
| `load_text_encoder_in_8bit` | Load the Gemma text encoder in 8-bit to save GPU memory |
| Parameter | Description |
|---------------------------------------|------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| `mixed_precision_mode` | Precision mode - `"bf16"` recommended for modern GPUs |
| `quantization` | Model quantization: `null`, `"int8-quanto"`, `"int4-quanto"`, `"fp8-quanto"`, etc. |
| `load_text_encoder_in_8bit` | Load the Gemma text encoder in 8-bit to save GPU memory |
| `offload_optimizer_during_validation` | Move optimizer state to CPU before validation video sampling and back afterwards. Useful when validation OOMs because VAE decoder + transformer + optimizer state can't coexist on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP. |
### DataConfig
@@ -50,17 +50,20 @@ This will create a `dataset.json` file containing video paths and their captions
**Captioning options:**
| Option | Description |
|--------|-------------|
| `--captioner-type` | `qwen_omni` (default, local) or `gemini_flash` (API) |
| `--use-8bit` | Enable 8-bit quantization for lower VRAM usage |
| `--no-audio` | Disable audio processing (video-only captions) |
| `--override` | Re-caption files that already have captions |
| `--api-key` | API key for Gemini Flash (or set `GOOGLE_API_KEY` env var) |
| Option | Description |
| ------------------ | ---------------------------------------------------------- |
| `--captioner-type` | `qwen_omni` (default, local) or `gemini_flash` (API) |
| `--use-8bit` | Enable 8-bit quantization for lower VRAM usage |
| `--no-audio` | Disable audio processing (video-only captions) |
| `--override` | Re-caption files that already have captions |
| `--api-key` | API key for Gemini Flash (or set `GOOGLE_API_KEY` env var) |
**Caption format:**
The captioner produces structured captions with sections for:
- **Visual content**: People, objects, actions, settings, colors, movements
- **Speech transcription**: Word-for-word transcription of spoken content
- **Sounds**: Music, ambient sounds, sound effects
@@ -106,15 +109,57 @@ uv run python scripts/process_dataset.py dataset.json \
--with-audio
```
### 🚀 Multi-GPU Preprocessing
Preprocessing large datasets can take a while. To run it across multiple GPUs in parallel, wrap the command with
`accelerate launch` (for example `--num_processes 4`). Each process handles an interleaved slice of the dataset.
The same approach applies to `process_videos.py` and `process_captions.py` when you run them standalone.
```bash
uv run accelerate launch --num_processes 4 scripts/process_dataset.py dataset.json \
--resolution-buckets "960x544x49" \
--model-path /path/to/ltx-2-model.safetensors \
--text-encoder-path /path/to/gemma-model
```
Outputs are written atomically (via a per-process temporary file, then renamed), so an interrupted run leaves no
corrupt files. By default a rerun **resumes** — items whose output `.pt` already exists are skipped.
> [!IMPORTANT]
> Pass `**--overwrite`** when rerunning with changed parameters (different model checkpoint, resolution buckets,
> text encoder, `--lora-trigger`, etc.). Without it the script keeps the stale outputs from the previous run.
>
> ```bash
> uv run accelerate launch --num_processes 4 scripts/process_dataset.py dataset.json \
> --resolution-buckets "960x544x49" \
> --model-path /path/to/ltx-2.3-model.safetensors \
> --text-encoder-path /path/to/gemma-model \
> --overwrite
> ```
### 📊 Dataset Format
The trainer supports either videos or single images.
Note that your dataset must be homogeneous - either all videos or all images, mixing is not supported.
The trainer supports videos, single images, or a mix of both in the same dataset.
> [!TIP]
> **Image Datasets:** When using images, follow the same preprocessing steps and format requirements as with videos,
> but use `1` for the frame count in the resolution bucket (e.g., `960x544x1`).
> [!NOTE]
> **Mixed image + video datasets:** Mixing stills and videos in a single dataset is supported, but requires some care:
>
> - Preprocess with **multiple resolution buckets** covering both frame counts — e.g.
> `--resolution-buckets "960x544x1;960x544x49"`. Images are automatically assigned to the `F=1` bucket and
> videos to an `F>1` bucket.
> - You **must** set `optimization.batch_size: 1` in your training config (see the warning under
> [Resolution Buckets](#-resolution-buckets)), since samples with different shapes cannot be collated into a
> single batch. Use `gradient_accumulation_steps` if you need a larger effective batch.
> - Per-step cost differs substantially between a single-frame sample and a many-frame sample, which can lead to
> uneven gradient magnitudes across steps. Consider weighting the two subsets or tuning the learning rate if
> you observe instability.
> - If you prefer a fully officially-supported path, train two separate LoRAs (one on stills, one on video) and
> stack them at inference.
The dataset must be a CSV, JSON, or JSONL metadata file with columns for captions and video paths:
**JSON format example:**
@@ -197,6 +242,7 @@ uv run python scripts/process_dataset.py dataset.json \
> ```
>
> Where:
>
> - H = Height of video
> - W = Width of video
> - F = Number of frames
@@ -204,6 +250,7 @@ uv run python scripts/process_dataset.py dataset.json \
> - 8 = VAE's temporal downsampling factor
>
> For example, a 768×448×89 video would have sequence length:
>
> ```
> (768/32) * (448/32) * ((89-1)/8 + 1) = 24 * 14 * 12 = 4,032
> ```
@@ -268,7 +315,6 @@ uv run python scripts/process_dataset.py dataset.json \
This will create an additional `reference_latents/` directory containing the preprocessed reference video latents.
### Generating Reference Videos
**Dataset Requirements for IC-LoRA:**
@@ -277,7 +323,7 @@ This will create an additional `reference_latents/` directory containing the pre
- Reference and target videos must have *identical* resolution and length
- Both reference and target videos should be preprocessed together using the same resolution buckets
We provide an example script, [`scripts/compute_reference.py`](../scripts/compute_reference.py), to generate reference
We provide an example script, `[scripts/compute_reference.py](../scripts/compute_reference.py)`, to generate reference
videos for a given dataset. The default implementation generates Canny edge reference videos.
```bash
@@ -293,7 +339,6 @@ If you want to generate a different type of condition (depth maps, pose skeleton
For reference, see our **[Canny Control Dataset](https://huggingface.co/datasets/Lightricks/Canny-Control-Dataset)** which demonstrates proper IC-LoRA dataset structure with paired videos and Canny edge maps.
## 🎯 LoRA Trigger Words
When training a LoRA, you can specify a trigger token that will be prepended to all captions:
@@ -84,6 +84,20 @@ optimization:
optimizer_type: "adamw8bit"
```
#### 7. Offload Optimizer State During Validation
If you OOM specifically during validation video sampling — typically in
full fine-tunes or high-rank LoRA runs where AdamW state and the VAE decoder
can't coexist on the GPU — offload optimizer state to CPU during sampling:
```yaml
acceleration:
offload_optimizer_during_validation: true
```
The offload + reload happens once per validation interval, not per step.
No effect for FSDP (sharded state).
---
## ⚠️ Common Usage Issues
@@ -143,6 +143,27 @@ uv run python scripts/process_dataset.py dataset.json \
> [!NOTE]
> When training with multiple resolution buckets, set `optimization.batch_size: 1`.
**Multi-GPU preprocessing.** Launch with `accelerate launch` to shard the dataset across processes. Reruns resume
by default (existing `.pt` outputs are skipped); writes are atomic so interrupted runs are safe. Pass `--overwrite`
when rerunning with changed parameters (different model, resolution buckets, text encoder, `--lora-trigger`, etc.)
so stale outputs are replaced. Use the same `accelerate launch` pattern (and `--overwrite` when needed) with
`process_videos.py` or `process_captions.py` when you run those scripts standalone.
```bash
# Multi-GPU preprocessing
uv run accelerate launch --num_processes 4 scripts/process_dataset.py dataset.json \
--resolution-buckets "960x544x49" \
--model-path /path/to/ltx-2-model.safetensors \
--text-encoder-path /path/to/gemma-model
# Force re-encoding of all items (e.g. after switching model or resolution)
uv run accelerate launch --num_processes 4 scripts/process_dataset.py dataset.json \
--resolution-buckets "960x544x49" \
--model-path /path/to/ltx-2.3-model.safetensors \
--text-encoder-path /path/to/gemma-model \
--overwrite
```
For detailed usage, see the [Dataset Preparation Guide](dataset-preparation.md).
### Reference Video Generation
+2 -2
View File
@@ -1,6 +1,6 @@
[project]
name = "ltx-trainer"
version = "1.1.2"
version = "1.1.3"
description = "LTX-2 training, democratized."
readme = "README.md"
authors = [
@@ -48,7 +48,7 @@ build-backend = "hatchling.build"
[tool.ruff]
target-version = "1.1.2"
target-version = "1.1.3"
line-length = 120
[tool.ruff.lint]
@@ -13,12 +13,14 @@ Can be used as a standalone script:
import json
import os
from collections.abc import Callable
from pathlib import Path
from typing import Any
import pandas as pd
import torch
import typer
from accelerate import PartialState
from rich.console import Console
from rich.progress import (
BarColumn,
@@ -30,7 +32,7 @@ from rich.progress import (
TimeElapsedColumn,
TimeRemainingColumn,
)
from torch.utils.data import DataLoader, Dataset
from torch.utils.data import DataLoader, Dataset, Subset
from transformers.utils.logging import disable_progress_bar
from ltx_trainer import logger
@@ -232,9 +234,14 @@ def compute_captions_embeddings( # noqa: PLR0913
batch_size: int = 8,
device: str = "cuda",
load_in_8bit: bool = False,
overwrite: bool = False,
) -> None:
"""
Process captions and save text embeddings.
Under ``accelerate launch``, each process handles an interleaved shard of
the dataset (rank/world read from ``accelerate.PartialState``). Already-
computed ``.pt`` outputs are skipped unless ``overwrite=True``; writes are
atomic so an interrupted run is safe to resume.
Args:
dataset_file: Path to metadata file (CSV/JSON/JSONL) containing captions and media paths
output_dir: Directory to save embeddings
@@ -247,11 +254,12 @@ def compute_captions_embeddings( # noqa: PLR0913
batch_size: Batch size for processing
device: Device to use for computation
load_in_8bit: Whether to load the Gemma text encoder in 8-bit precision
overwrite: Re-encode every item even if its output exists. Use when rerunning with
changed parameters (different text encoder, lora_trigger, etc.) so stale
outputs are replaced.
"""
console = Console()
# Create dataset
dataset = CaptionsDataset(
dataset_file=dataset_file,
caption_column=caption_column,
@@ -264,6 +272,24 @@ def compute_captions_embeddings( # noqa: PLR0913
output_path = Path(output_dir)
output_path.mkdir(parents=True, exist_ok=True)
# TODO(batch-tokenization): The current Gemma tokenizer doesn't support batched tokenization.
if batch_size > 1:
logger.warning(
"Batch size greater than 1 is not currently supported with the Gemma tokenizer. "
"Overriding batch_size to 1. This will be fixed in a future update."
)
batch_size = 1
dataloader = _build_sharded_dataloader(
dataset,
batch_size=batch_size,
num_workers=2,
is_done=lambda idx: (output_path / dataset.output_paths[idx]).is_file(),
overwrite=overwrite,
)
if dataloader is None:
return
# Load text encoder and embeddings processor
with console.status("[bold]Loading Gemma text encoder...", spinner="dots"):
text_encoder = load_text_encoder(
@@ -279,21 +305,7 @@ def compute_captions_embeddings( # noqa: PLR0913
)
logger.info("Text encoder and embeddings processor loaded successfully")
# TODO(batch-tokenization): The current Gemma tokenizer doesn't support batched tokenization.
if batch_size > 1:
logger.warning(
"Batch size greater than 1 is not currently supported with the Gemma tokenizer. "
"Overriding batch_size to 1. This will be fixed in a future update."
)
batch_size = 1
# Create dataloader
dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=False, num_workers=2)
# Process batches
total_batches = len(dataloader)
logger.info(f"Processing captions in {total_batches:,} batches...")
logger.info(f"Processing captions in {len(dataloader):,} batches...")
with Progress(
SpinnerColumn(),
@@ -333,11 +345,44 @@ def compute_captions_embeddings( # noqa: PLR0913
embedding_data["audio_prompt_embeds"] = audio_prompt_embeds[0].cpu().contiguous()
output_file = output_path / output_rel_path
torch.save(embedding_data, output_file)
_atomic_save(embedding_data, output_file)
progress.advance(task)
logger.info(f"Processed {len(dataset):,} captions. Embeddings saved to {output_path}")
logger.info(f"Processed {len(dataloader.dataset):,} captions -> {output_path}") # type: ignore[arg-type]
def _atomic_save(data: Any, out: Path) -> None: # noqa: ANN401
"""Save to ``out`` atomically via per-PID temp file + replace.
Crash mid-write leaves an orphan ``.tmp.<pid>`` file that the skip logic
ignores. The per-PID suffix makes concurrent writes from multiple ranks
collision-free.
"""
tmp = out.with_suffix(f"{out.suffix}.tmp.{os.getpid()}")
torch.save(data, tmp)
tmp.replace(out)
def _build_sharded_dataloader(
dataset: Dataset,
*,
batch_size: int,
num_workers: int,
is_done: Callable[[int], bool],
overwrite: bool,
) -> DataLoader | None:
"""Return a DataLoader over this rank's interleaved shard of ``dataset``.
When ``overwrite`` is False, items whose outputs already exist (per
``is_done``) are filtered out. Returns ``None`` if this rank has nothing
to do, so the caller can early-return without loading any models.
"""
state = PartialState()
todo = [i for i in range(state.process_index, len(dataset), state.num_processes) if overwrite or not is_done(i)]
if not todo:
logger.info(f"Rank {state.process_index}/{state.num_processes}: nothing to do")
return None
logger.info(f"Rank {state.process_index}/{state.num_processes}: processing {len(todo):,} of {len(dataset):,} items")
return DataLoader(Subset(dataset, todo), batch_size=batch_size, shuffle=False, num_workers=num_workers)
@app.command()
@@ -387,8 +432,15 @@ def main( # noqa: PLR0913
default=False,
help="Load the Gemma text encoder in 8-bit precision to save GPU memory (requires bitsandbytes)",
),
overwrite: bool = typer.Option(
default=False,
help="Re-encode every caption even if its output exists. Use when rerunning with "
"changed parameters (different text encoder, lora_trigger, etc.) so stale outputs are replaced.",
),
) -> None:
"""Process text captions and save embeddings for video generation training.
For multi-GPU preprocessing, invoke under ``accelerate launch`` - each process
will handle an interleaved shard of the dataset.
This script processes captions from metadata files and saves text embeddings
that can be used for training video generation models. The output embeddings
will maintain the same folder structure and naming as the corresponding media files.
@@ -428,6 +480,7 @@ def main( # noqa: PLR0913
batch_size=batch_size,
device=device,
load_in_8bit=load_text_encoder_in_8bit,
overwrite=overwrite,
)
@@ -50,6 +50,7 @@ def preprocess_dataset( # noqa: PLR0913
reference_downscale_factor: int = 1,
with_audio: bool = False,
load_text_encoder_in_8bit: bool = False,
overwrite: bool = False,
) -> None:
"""Run the preprocessing pipeline with the given arguments."""
# Validate dataset file
@@ -77,6 +78,7 @@ def preprocess_dataset( # noqa: PLR0913
batch_size=batch_size,
device=device,
load_in_8bit=load_text_encoder_in_8bit,
overwrite=overwrite,
)
# Process videos using the dedicated function
@@ -97,6 +99,7 @@ def preprocess_dataset( # noqa: PLR0913
vae_tiling=vae_tiling,
with_audio=with_audio,
audio_output_dir=str(audio_latents_dir) if audio_latents_dir else None,
overwrite=overwrite,
)
# Process reference videos if reference_column is provided
@@ -133,6 +136,7 @@ def preprocess_dataset( # noqa: PLR0913
batch_size=batch_size,
device=device,
vae_tiling=vae_tiling,
overwrite=overwrite,
)
# Handle decoding if requested (for verification)
@@ -252,8 +256,15 @@ def main( # noqa: PLR0913
help="Downscale factor for reference video resolution. When > 1, reference videos are processed at "
"1/n resolution (e.g., 2 means half resolution). Used for efficient IC-LoRA training.",
),
overwrite: bool = typer.Option(
default=False,
help="Re-compute every item even if its output exists. Use when rerunning with "
"changed parameters (different model, resolution, etc.) so stale outputs are replaced.",
),
) -> None:
"""Preprocess a video dataset by computing and saving latents and text embeddings.
For multi-GPU preprocessing, invoke under ``accelerate launch`` - each process
will handle an interleaved shard of the dataset.
The dataset must be a CSV, JSON, or JSONL file with columns for captions and video paths.
This script is designed for LTX-2 models which use the Gemma text encoder.
Examples:
@@ -310,6 +321,7 @@ def main( # noqa: PLR0913
reference_downscale_factor=reference_downscale_factor,
with_audio=with_audio,
load_text_encoder_in_8bit=load_text_encoder_in_8bit,
overwrite=overwrite,
)
+79 -18
View File
@@ -14,6 +14,8 @@ Can be used as a standalone script:
import json
import math
import os
from collections.abc import Callable
from pathlib import Path
from typing import Any
@@ -22,6 +24,7 @@ import pandas as pd
import torch
import torchaudio
import typer
from accelerate import PartialState
from pillow_heif import register_heif_opener
from rich.console import Console
from rich.progress import (
@@ -34,7 +37,7 @@ from rich.progress import (
TimeElapsedColumn,
TimeRemainingColumn,
)
from torch.utils.data import DataLoader, Dataset
from torch.utils.data import DataLoader, Dataset, Subset
from torchvision import transforms
from torchvision.transforms import InterpolationMode
from torchvision.transforms.functional import crop, resize, to_tensor
@@ -444,9 +447,14 @@ def compute_latents( # noqa: PLR0913, PLR0915
vae_tiling: bool = False,
with_audio: bool = False,
audio_output_dir: str | None = None,
overwrite: bool = False,
) -> None:
"""
Process videos and save latent representations.
Under ``accelerate launch``, each process handles an interleaved shard of
the dataset (rank/world read from ``accelerate.PartialState``). Already-
computed ``.pt`` outputs are skipped unless ``overwrite=True``; writes are
atomic so an interrupted run is safe to resume.
Args:
dataset_file: Path to metadata file (CSV/JSON/JSONL) containing video paths
video_column: Column name for video paths in the metadata file
@@ -460,15 +468,15 @@ def compute_latents( # noqa: PLR0913, PLR0915
vae_tiling: Whether to enable VAE tiling
with_audio: Whether to extract and encode audio from videos
audio_output_dir: Directory to save audio latents (required if with_audio=True)
overwrite: Re-process every item even if its output exists. Use when rerunning with
changed parameters (different model, resolution, etc.) so stale outputs are replaced.
"""
# Validate audio parameters
if with_audio and audio_output_dir is None:
raise ValueError("audio_output_dir must be provided when with_audio=True")
console = Console()
torch_device = torch.device(device)
# Create dataset
dataset = MediaDataset(
dataset_file=dataset_file,
main_media_column=main_media_column or video_column,
@@ -481,13 +489,34 @@ def compute_latents( # noqa: PLR0913, PLR0915
output_path = Path(output_dir)
output_path.mkdir(parents=True, exist_ok=True)
# Set up audio output directory if needed
audio_output_path = None
audio_output_path: Path | None = None
if with_audio:
audio_output_path = Path(audio_output_dir)
audio_output_path.mkdir(parents=True, exist_ok=True)
# Audio processing requires batch_size=1; must be applied before the dataloader is built.
if with_audio and batch_size > 1:
logger.warning("Audio processing requires batch_size=1. Overriding batch_size to 1.")
batch_size = 1
data_root = dataset.dataset_file.parent
def _is_done(idx: int) -> bool:
rel = dataset.main_media_paths[idx].relative_to(data_root).with_suffix(".pt")
if not (output_path / rel).is_file():
return False
return audio_output_path is None or (audio_output_path / rel).is_file()
dataloader = _build_sharded_dataloader(
dataset,
batch_size=batch_size,
num_workers=4,
is_done=_is_done,
overwrite=overwrite,
)
if dataloader is None:
return
# Load video VAE encoder
with console.status(f"[bold]Loading video VAE encoder from [cyan]{model_path}[/]...", spinner="dots"):
vae = load_video_vae_encoder(model_path, device=torch_device, dtype=torch.bfloat16)
@@ -510,14 +539,6 @@ def compute_latents( # noqa: PLR0913, PLR0915
n_fft=audio_vae_encoder.n_fft,
).to(torch_device)
# Create dataloader
# Note: batch_size=1 required when with_audio because audio extraction can fail for some videos,
# and the default collate function can't handle mixed None/dict values across a batch.
if with_audio and batch_size > 1:
logger.warning("Audio processing requires batch_size=1. Overriding batch_size to 1.")
batch_size = 1
dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=False, num_workers=4)
# Track audio statistics
audio_success_count = 0
audio_skip_count = 0
@@ -560,7 +581,7 @@ def compute_latents( # noqa: PLR0913, PLR0915
"fps": batch["video_metadata"]["fps"][i].item(),
}
torch.save(latent_data, output_file)
_atomic_save(latent_data, output_file)
# Process audio if enabled (audio is already extracted by the dataset)
if with_audio:
@@ -588,7 +609,7 @@ def compute_latents( # noqa: PLR0913, PLR0915
"duration": audio_latents["duration"],
}
torch.save(audio_save_data, audio_output_file)
_atomic_save(audio_save_data, audio_output_file)
audio_success_count += 1
else:
# Video has no audio track
@@ -596,8 +617,7 @@ def compute_latents( # noqa: PLR0913, PLR0915
progress.advance(task)
# Log summary
logger.info(f"Processed {len(dataset)} videos. Latents saved to {output_path}")
logger.info(f"Processed {len(dataloader.dataset)} videos -> {output_path}") # type: ignore[arg-type]
if with_audio:
logger.info(
f"Audio processing: {audio_success_count} videos with audio, "
@@ -935,6 +955,39 @@ def compute_scaled_resolution_buckets(
return scaled_buckets
def _atomic_save(data: Any, out: Path) -> None: # noqa: ANN401
"""Save to ``out`` atomically via per-PID temp file + replace.
Crash mid-write leaves an orphan ``.tmp.<pid>`` file that the skip logic
ignores. The per-PID suffix makes concurrent writes from multiple ranks
collision-free.
"""
tmp = out.with_suffix(f"{out.suffix}.tmp.{os.getpid()}")
torch.save(data, tmp)
tmp.replace(out)
def _build_sharded_dataloader(
dataset: Dataset,
*,
batch_size: int,
num_workers: int,
is_done: Callable[[int], bool],
overwrite: bool,
) -> DataLoader | None:
"""Return a DataLoader over this rank's interleaved shard of ``dataset``.
When ``overwrite`` is False, items whose outputs already exist (per
``is_done``) are filtered out. Returns ``None`` if this rank has nothing
to do, so the caller can early-return without loading any models.
"""
state = PartialState()
todo = [i for i in range(state.process_index, len(dataset), state.num_processes) if overwrite or not is_done(i)]
if not todo:
logger.info(f"Rank {state.process_index}/{state.num_processes}: nothing to do")
return None
logger.info(f"Rank {state.process_index}/{state.num_processes}: processing {len(todo):,} of {len(dataset):,} items")
return DataLoader(Subset(dataset, todo), batch_size=batch_size, shuffle=False, num_workers=num_workers)
@app.command()
def main( # noqa: PLR0913
dataset_file: str = typer.Argument(
@@ -981,8 +1034,15 @@ def main( # noqa: PLR0913
default=None,
help="Output directory for audio latents (required if --with-audio is set)",
),
overwrite: bool = typer.Option(
default=False,
help="Re-encode every item even if its output exists. Use when rerunning with "
"changed parameters (different model, resolution, etc.) so stale outputs are replaced.",
),
) -> None:
"""Process videos/images and save latent representations for video generation training.
For multi-GPU preprocessing, invoke under ``accelerate launch`` - each process
will handle an interleaved shard of the dataset.
This script processes videos and images from metadata files and saves latent representations
that can be used for training video generation models. The output latents will maintain
the same folder structure and naming as the corresponding media files.
@@ -1032,6 +1092,7 @@ def main( # noqa: PLR0913
vae_tiling=vae_tiling,
with_audio=with_audio,
audio_output_dir=audio_output_dir,
overwrite=overwrite,
)
@@ -168,6 +168,15 @@ class AccelerationConfig(ConfigBaseModel):
description="Whether to load the text encoder in 8-bit precision to save memory",
)
offload_optimizer_during_validation: bool = Field(
default=False,
description="Offload optimizer state to CPU before validation video sampling and reload "
"it afterwards, to free VRAM for inference. Useful when optimizer state is large "
"(e.g. AdamW for full fine-tuning or high-rank LoRA) and validation OOMs because the "
"VAE decoder + transformer + optimizer state cannot coexist on the GPU. Has no effect "
"for FSDP (sharded state). Disabled by default.",
)
class DataConfig(ConfigBaseModel):
"""Configuration for data loading and processing"""
@@ -85,6 +85,7 @@ def print_config(config: LtxTrainerConfig) -> None:
("Mixed Precision", accel.mixed_precision_mode or "[dim]—[/]"),
("Quantization", str(accel.quantization) if accel.quantization else "[dim]—[/]"),
("Text Encoder 8bit", fmt(accel.load_text_encoder_in_8bit)),
("Optimizer CPU Offload", fmt(accel.offload_optimizer_during_validation)),
],
),
(
@@ -12,6 +12,7 @@ Example usage:
from __future__ import annotations
import logging
import os
from collections.abc import Generator
from contextlib import contextmanager
from pathlib import Path
@@ -22,7 +23,11 @@ from ltx_core.text_encoders.gemma.encoders.base_encoder import GemmaTextEncoder
from ltx_core.text_encoders.gemma.tokenizer import LTXVGemmaTokenizer
def load_8bit_gemma(gemma_model_path: str | Path, dtype: torch.dtype = torch.bfloat16) -> GemmaTextEncoder:
def load_8bit_gemma(
gemma_model_path: str | Path,
dtype: torch.dtype = torch.bfloat16,
device: torch.device | str | int | None = None,
) -> GemmaTextEncoder:
"""Load the Gemma text encoder in 8-bit precision using bitsandbytes.
Only the Gemma LLM backbone is loaded here. The embeddings processor
(feature extractor + connectors) should be loaded separately via
@@ -30,6 +35,10 @@ def load_8bit_gemma(gemma_model_path: str | Path, dtype: torch.dtype = torch.bfl
Args:
gemma_model_path: Path to Gemma model directory
dtype: Data type for non-quantized model weights
device: Device to place the quantized model on. When ``None`` (default),
the device is inferred from ``LOCAL_RANK`` if CUDA is available, so
multi-process launches put each rank's encoder on its own GPU
instead of all colliding on ``cuda:0``.
Returns:
GemmaTextEncoder with 8-bit quantized Gemma backbone
Raises:
@@ -46,13 +55,23 @@ def load_8bit_gemma(gemma_model_path: str | Path, dtype: torch.dtype = torch.bfl
gemma_path = _find_gemma_subpath(gemma_model_path, "model*.safetensors")
tokenizer_path = _find_gemma_subpath(gemma_model_path, "tokenizer.model")
# Pin the entire model to a single device. `device_map="auto"` collides on cuda:0
# in multi-process launches because every rank picks the same default device.
device_map: str | dict[str, int | str | torch.device]
if device is not None:
device_map = {"": device}
elif torch.cuda.is_available():
device_map = {"": int(os.environ.get("LOCAL_RANK", "0"))}
else:
device_map = "auto"
quantization_config = BitsAndBytesConfig(load_in_8bit=True)
with _suppress_accelerate_memory_warnings():
gemma_model = Gemma3ForConditionalGeneration.from_pretrained(
gemma_path,
quantization_config=quantization_config,
torch_dtype=torch.bfloat16,
device_map="auto",
device_map=device_map,
local_files_only=True,
)
@@ -199,8 +199,6 @@ def load_text_encoder(
device: Device to load model on
dtype: Data type for model weights
load_in_8bit: Whether to load the Gemma model in 8-bit precision using bitsandbytes.
When True, the model is loaded with device_map="auto" and the device argument
is ignored for the Gemma backbone.
Returns:
Loaded GemmaTextEncoder
"""
@@ -211,7 +209,7 @@ def load_text_encoder(
if load_in_8bit:
from ltx_trainer.gemma_8bit import load_8bit_gemma
return load_8bit_gemma(gemma_model_path, dtype)
return load_8bit_gemma(gemma_model_path, dtype, device=device)
# Standard loading path
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
+126 -43
View File
@@ -1,7 +1,10 @@
import contextlib
import math
import os
import re
import time
import warnings
from collections.abc import Iterator
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable
@@ -9,8 +12,8 @@ from typing import Any, Callable
import torch
import wandb
import yaml
from accelerate import Accelerator, DistributedType
from accelerate.utils import set_seed
from accelerate import Accelerator, DistributedDataParallelKwargs, DistributedType
from accelerate.utils import gather_object, set_seed
from peft import LoraConfig, get_peft_model, get_peft_model_state_dict, set_peft_model_state_dict
from peft.tuners.tuners_utils import BaseTunerLayer
from peft.utils import ModulesToSaveWrapper
@@ -63,7 +66,7 @@ if not IS_MAIN_PROCESS:
disable_progress_bar()
StepCallback = Callable[[int, int, list[Path]], None] # (step, total, list[sampled_video_path]) -> None
StepCallback = Callable[[int, int, list[Path] | None], None] # (step, total, sampled paths or None) -> None
MEMORY_CHECK_INTERVAL = 200
@@ -186,9 +189,8 @@ class LtxvTrainer:
with progress:
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:
self._log_validation_samples(sampled_videos_paths, cfg.validation.prompts)
with self._offloaded_optimizer_state():
sampled_videos_paths = self._run_distributed_validation(progress)
self._accelerator.wait_for_everyone()
@@ -228,16 +230,8 @@ class LtxvTrainer:
and self._global_step % cfg.validation.interval == 0
and is_optimization_step
):
if self._accelerator.distributed_type == DistributedType.FSDP:
# FSDP: All processes must participate in validation
sampled_videos_paths = self._sample_videos(progress)
if IS_MAIN_PROCESS and sampled_videos_paths and self._config.wandb.log_validation_videos:
self._log_validation_samples(sampled_videos_paths, cfg.validation.prompts)
# DDP: Only main process runs validation
elif IS_MAIN_PROCESS:
sampled_videos_paths = self._sample_videos(progress)
if sampled_videos_paths and self._config.wandb.log_validation_videos:
self._log_validation_samples(sampled_videos_paths, cfg.validation.prompts)
with self._offloaded_optimizer_state():
sampled_videos_paths = self._run_distributed_validation(progress)
# Save checkpoint if needed
if (
@@ -398,11 +392,14 @@ class LtxvTrainer:
# 3. If validation prompts are configured, computes and caches their embeddings
# 4. Unloads the Gemma model entirely, keeps the embeddings processor for training
# Load text encoder (pure Gemma LLM) on GPU
# Load text encoder (pure Gemma LLM) on GPU — LOCAL_RANK before Accelerator exists
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
init_device = torch.device(f"cuda:{local_rank}" if torch.cuda.is_available() else "cpu")
logger.debug("Loading text encoder...")
text_encoder = load_text_encoder(
gemma_model_path=self._config.model.text_encoder_path,
device="cuda",
device=init_device,
dtype=torch.bfloat16,
load_in_8bit=self._config.acceleration.load_text_encoder_in_8bit,
)
@@ -411,7 +408,7 @@ class LtxvTrainer:
logger.debug("Loading embeddings processor...")
self._embeddings_processor = load_embeddings_processor(
checkpoint_path=self._config.model.model_path,
device="cuda",
device=init_device,
dtype=torch.bfloat16,
)
@@ -788,6 +785,41 @@ class LtxvTrainer:
# noinspection PyTypeChecker
self._optimizer, self._lr_scheduler = self._accelerator.prepare(optimizer, lr_scheduler)
@contextlib.contextmanager
def _offloaded_optimizer_state(self) -> Iterator[None]:
"""Context manager that offloads optimizer state to CPU during validation.
Opt-in via `acceleration.offload_optimizer_during_validation`. Frees VRAM for
validation video generation when optimizer state is large (e.g. full fine-tune
AdamW, high-rank LoRA). No-op for FSDP (sharded state -- manual `.cpu()` breaks
metadata).
"""
enabled = (
self._config.acceleration.offload_optimizer_during_validation
and self._accelerator.distributed_type != DistributedType.FSDP
)
# Track exactly which tensors we move so we don't promote ones that were
# intentionally on CPU (e.g. AdamW's `step` scalar on recent PyTorch).
offloaded: list[tuple[dict, str]] = []
if enabled:
offloaded_bytes = 0
for state in self._optimizer.state.values():
for k, v in state.items():
if isinstance(v, torch.Tensor) and v.is_cuda:
offloaded.append((state, k))
offloaded_bytes += v.nbytes
if offloaded:
logger.info(f"Offloading optimizer state to CPU ({offloaded_bytes / 1e9:.1f} GB)")
for state, k in offloaded:
state[k] = state[k].cpu()
try:
yield
finally:
device = self._accelerator.device
for state, k in offloaded:
state[k] = state[k].to(device)
def _create_scheduler(self, optimizer: torch.optim.Optimizer) -> LRScheduler | None:
"""Create learning rate scheduler based on config."""
scheduler_type = self._config.optimization.scheduler_type
@@ -844,11 +876,18 @@ class LtxvTrainer:
def _setup_accelerator(self) -> None:
"""Initialize the Accelerator with the appropriate settings."""
# find_unused_parameters=True keeps DDP happy when LoRA targets a branch the forward
# pass skips (e.g. audio LoRA with `with_audio: false`, or short module patterns like
# "to_k" that match the audio branch unintentionally). It's a no-op for FSDP and
# single-GPU runs. The probing cost is paid only on the first step.
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
# All distributed setup (DDP/FSDP, number of processes, etc.) is controlled by
# the user's Accelerate configuration (accelerate config / accelerate launch).
self._accelerator = Accelerator(
mixed_precision=self._config.acceleration.mixed_precision_mode,
gradient_accumulation_steps=self._config.optimization.gradient_accumulation_steps,
kwargs_handlers=[ddp_kwargs],
)
if self._accelerator.num_processes > 1:
@@ -881,11 +920,42 @@ class LtxvTrainer:
"Monitor training stability and consider disabling quantization if issues arise."
)
def _run_distributed_validation(self, progress: TrainingProgress) -> list[Path]:
"""Run validation across all ranks and log gathered results on rank 0.
Each rank generates only its assigned subset of prompts (see `_sample_videos`),
so all GPUs stay busy and no rank idles long enough to trigger NCCL timeouts.
Paths are gathered across ranks so rank 0 has the full list for W&B logging.
Note: Multi-node training requires a shared filesystem so rank 0 can read
videos written by other ranks.
"""
sampled = self._sample_videos(progress)
if self._accelerator.num_processes > 1:
# gather_object returns a flat list from all ranks
sampled = sorted(gather_object(sampled), key=lambda x: x[0])
paths = [p for _, p in sampled]
if self._accelerator.is_main_process and paths:
self._log_validation_samples(paths, self._config.validation.prompts)
# Non-main ranks must not reach checkpoint collectives while main is still logging to W&B.
self._accelerator.wait_for_everyone()
return paths
# Note: Use @torch.no_grad() instead of @torch.inference_mode() to avoid FSDP inplace update errors after validation
@torch.no_grad()
@free_gpu_memory_context(after=True)
def _sample_videos(self, progress: TrainingProgress) -> list[Path] | None:
"""Run validation by generating videos from validation prompts."""
def _sample_videos(self, progress: TrainingProgress) -> list[tuple[int, Path]]:
"""Run validation by generating videos from this rank's share of the validation prompts.
Prompts are split round-robin across ranks via `process_index` / `num_processes`,
which collapses to "all prompts" when running on a single GPU. Returns
(prompt_idx, path) tuples so the caller can reconstruct global order without
relying on filename conventions.
Under FSDP with multiple processes, ranks pad with extra generate passes (same prompt,
no disk write) so every rank runs the same number of forwards avoids collective mismatch.
"""
use_images = self._config.validation.images is not None
use_reference_videos = self._config.validation.reference_videos is not None
generate_audio = self._config.validation.generate_audio
@@ -895,13 +965,24 @@ class LtxvTrainer:
self._optimizer.zero_grad(set_to_none=True)
free_gpu_memory()
# Start sampling progress tracking
prompts = self._config.validation.prompts
rank = self._accelerator.process_index
world_size = self._accelerator.num_processes
rank_indices = list(range(rank, len(prompts), world_size))
# FSDP: every rank must run the same number of forwards; pad with duplicate generates (no save).
work: list[tuple[int, bool]] = [(i, True) for i in rank_indices]
if self._accelerator.distributed_type == DistributedType.FSDP and world_size > 1:
max_per_rank = math.ceil(len(prompts) / world_size)
pad_seed = rank_indices[-1] if rank_indices else 0
work += [(pad_seed, False)] * (max_per_rank - len(work))
sampling_ctx = progress.start_sampling(
num_prompts=len(self._config.validation.prompts),
num_prompts=len(work),
num_steps=inference_steps,
)
# Create validation sampler with loaded models and progress tracking
# Create a validation sampler with loaded models and progress tracking
sampler = ValidationSampler(
transformer=self._transformer,
vae_decoder=self._vae_decoder,
@@ -915,12 +996,12 @@ class LtxvTrainer:
output_dir = Path(self._config.output_dir) / "samples"
output_dir.mkdir(exist_ok=True, parents=True)
video_paths = []
results: list[tuple[int, Path]] = []
width, height, num_frames = self._config.validation.video_dims
for prompt_idx, prompt in enumerate(self._config.validation.prompts):
# Update progress to show current video
sampling_ctx.start_video(prompt_idx)
for local_i, (prompt_idx, save_output) in enumerate(work):
prompt = prompts[prompt_idx]
sampling_ctx.start_video(local_i)
# Load conditioning image if provided
condition_image = None
@@ -972,28 +1053,30 @@ class LtxvTrainer:
device=self._accelerator.device,
)
if not save_output:
continue
# Save output (image for single frame, video otherwise)
if IS_MAIN_PROCESS:
ext = "png" if num_frames == 1 else "mp4"
output_path = output_dir / f"step_{self._global_step:06d}_{prompt_idx + 1}.{ext}"
if num_frames == 1:
save_image(video, output_path)
else:
save_video(
video_tensor=video,
output_path=output_path,
fps=self._config.validation.frame_rate,
audio=audio,
audio_sample_rate=self._vocoder.output_sampling_rate if audio is not None else None,
)
video_paths.append(output_path)
ext = "png" if num_frames == 1 else "mp4"
output_path = output_dir / f"step_{self._global_step:06d}_{prompt_idx + 1:02d}.{ext}"
if num_frames == 1:
save_image(video, output_path)
else:
save_video(
video_tensor=video,
output_path=output_path,
fps=self._config.validation.frame_rate,
audio=audio,
audio_sample_rate=self._vocoder.output_sampling_rate if audio is not None else None,
)
results.append((prompt_idx, output_path))
# Clean up progress tasks
sampling_ctx.cleanup()
rel_outputs_path = output_dir.relative_to(self._config.output_dir)
logger.info(f"🎥 Validation samples for step {self._global_step} saved in {rel_outputs_path}")
return video_paths
return results
@staticmethod
def _log_training_stats(stats: TrainingStats) -> None:
Generated
+4 -4
View File
@@ -2063,7 +2063,7 @@ wheels = [
[[package]]
name = "ltx-core"
version = "1.1.2"
version = "1.1.3"
source = { editable = "packages/ltx-core" }
dependencies = [
{ name = "accelerate" },
@@ -2121,7 +2121,7 @@ dev = [{ name = "scikit-image", specifier = ">=0.25.2" }]
[[package]]
name = "ltx-pipelines"
version = "1.1.2"
version = "1.1.3"
source = { editable = "packages/ltx-pipelines" }
dependencies = [
{ name = "av" },
@@ -2144,7 +2144,7 @@ requires-dist = [
[[package]]
name = "ltx-trainer"
version = "1.1.2"
version = "1.1.3"
source = { editable = "packages/ltx-trainer" }
dependencies = [
{ name = "accelerate" },
@@ -7105,7 +7105,7 @@ dependencies = [
{ name = "torch", version = "2.9.1", source = { registry = "https://pypi.org/simple" } },
]
wheels = [
{ url = "https://download.pytorch.org/whl/cu129/xformers-0.0.33%2B5d4b92a5.d20251029-cp39-abi3-linux_x86_64.whl", upload-time = "2025-10-30T00:15:46Z" },
{ url = "https://download.pytorch.org/whl/cu129/xformers-0.0.33%2B5d4b92a5.d20251029-cp39-abi3-linux_x86_64.whl", hash = "sha256:4e4f2dea153b60ca4f21cc44b82c072b97b096a16f1920c7539c91ec8ffd7ba4", upload-time = "2026-04-27T19:01:32Z" },
]
[[package]]