Automated PR - 2026-04-23

This commit is contained in:
github-actions[bot]
2026-04-23 12:43:54 +00:00
parent a2c3f24078
commit b604d3fab3
49 changed files with 2664 additions and 568 deletions
+1
View File
@@ -7,3 +7,4 @@
*.jpeg filter=lfs diff=lfs merge=lfs -text
*.jpg filter=lfs diff=lfs merge=lfs -text
*.webp filter=lfs diff=lfs merge=lfs -text
*.exr filter=lfs diff=lfs merge=lfs -text
+5
View File
@@ -27,6 +27,7 @@ tmp
*.sft
# Media files
*.exr
*.gif
*.heic
*.heif
@@ -40,5 +41,9 @@ tmp
*.wav
*.webp
# HDR IC-LoRA e2e test baseline (checked in via Git LFS)
!packages/ltx-pipelines/tests/assets/expected_hdr_ic_lora_exr/frame_*.exr
!packages/ltx-pipelines/tests/assets/hdr_ic_lora_test_input.mp4
# Binary files
*.so
+2
View File
@@ -57,6 +57,7 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https://
* [`LTX-2-19b-LoRA-Camera-Control-Jib-Down`](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Jib-Down) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Jib-Down/resolve/main/ltx-2-19b-lora-camera-control-jib-down.safetensors)
* [`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`
### Available Pipelines
@@ -68,6 +69,7 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https://
* **[KeyframeInterpolationPipeline](packages/ltx-pipelines/src/ltx_pipelines/keyframe_interpolation.py)** - Interpolate between keyframe images
* **[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)
### ⚡ Optimization Tips
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "ltx-core"
version = "1.1.1"
version = "1.1.2"
description = "Core implementation of Lightricks' LTX-2 model"
readme = "README.md"
requires-python = ">=3.10"
@@ -1,5 +1,5 @@
"""Batch-splitting adapter for the transformer.
Wraps an ``X0Model`` (or ``LayerStreamingWrapper``) and splits batched inputs
Wraps an ``X0Model`` (or ``BlockStreamingWrapper``) and splits batched inputs
into smaller chunks before forwarding, then concatenates the results. This
controls peak activation memory at the cost of more forward passes.
The adapter is transparent — it has the same ``forward`` signature as
@@ -42,7 +42,7 @@ class BatchSplitAdapter(nn.Module):
Has the same ``forward`` signature as ``X0Model``:
``(video, audio, perturbations) -> (denoised_video, denoised_audio)``.
Args:
model: The model to wrap (``X0Model``, ``LayerStreamingWrapper``, etc.).
model: The model to wrap (``X0Model``, ``BlockStreamingWrapper``, etc.).
max_batch_size: Maximum batch size per forward pass. Input batches
larger than this are split into sequential chunks.
"""
@@ -0,0 +1,19 @@
"""Block streaming: memory-efficient sequential-block inference.
Streams transformer blocks from safetensors to GPU one at a time.
Block weights are provided by a :class:`WeightsProvider` which handles
CPU-to-GPU copies, caching, and stream synchronization. Two weight
source strategies are available:
- **RAM streaming** (default): all blocks pre-loaded into pinned CPU
buffers with LoRA fusion at build time. Fast, higher CPU memory.
- **Disk streaming** (``cpu_slots < num_blocks``): blocks read from
disk on demand with FIFO eviction. Slower, lower CPU memory.
"""
from ltx_core.block_streaming.builder import DISK_CPU_SLOTS, StreamingModelBuilder
from ltx_core.block_streaming.wrapper import BlockStreamingWrapper
__all__ = [
"DISK_CPU_SLOTS",
"BlockStreamingWrapper",
"StreamingModelBuilder",
]
@@ -0,0 +1,305 @@
"""Builder that constructs a BlockStreamingWrapper from safetensors checkpoints."""
from __future__ import annotations
import logging
from collections.abc import Callable
from dataclasses import dataclass, field, replace
from typing import Generic
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.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.wrapper import BlockStreamingWrapper
from ltx_core.loader.fuse_loras import apply_loras
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 (
LoraPathStrengthAndSDOps,
LoraStateDictWithStrength,
ModelBuilderProtocol,
StateDictLoader,
)
from ltx_core.loader.registry import DummyRegistry, Registry
from ltx_core.loader.sd_ops import SDOps
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader
from ltx_core.model.model_protocol import ModelConfigurator, ModelType
logger = logging.getLogger(__name__)
DISK_CPU_SLOTS = 2
_DEFAULT_GPU_SLOTS = 2
@dataclass(frozen=True)
class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType]):
"""Immutable builder for :class:`BlockStreamingWrapper`.
Reads block weights from safetensors on demand. ``cpu_slots`` and
``gpu_slots`` control the memory/speed trade-off (see :meth:`build`).
Args:
model_class_configurator: Creates the model from a config dict.
model_path: One or more ``.safetensors`` checkpoint paths.
model_sd_ops: Key remapping applied to safetensors keys.
module_ops: Module-level mutations for the meta model.
loras: LoRA adapters fused into weights at load time.
model_loader: Strategy for reading checkpoint metadata.
registry: Shared cache for loaded state dicts.
blocks_attr: Dotted path to the ``nn.ModuleList`` (e.g.
``"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."``).
model_wrapper: Optional callable wrapping the model
(e.g. ``X0Model``).
"""
model_class_configurator: type[ModelConfigurator[ModelType]]
model_path: str | tuple[str, ...]
model_sd_ops: SDOps | None = None
module_ops: tuple[ModuleOps, ...] = field(default_factory=tuple)
loras: tuple[LoraPathStrengthAndSDOps, ...] = field(default_factory=tuple)
model_loader: StateDictLoader = field(default_factory=SafetensorsModelStateDictLoader)
registry: Registry = field(default_factory=DummyRegistry)
# Streaming-specific
blocks_attr: str = ""
blocks_prefix: str = ""
state_dict_prefix: str = ""
model_wrapper: Callable[[ModelType], nn.Module] | None = None
def with_sd_ops(self, sd_ops: SDOps | None) -> StreamingModelBuilder:
return replace(self, model_sd_ops=sd_ops)
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> StreamingModelBuilder:
return replace(self, module_ops=module_ops)
def with_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> StreamingModelBuilder:
return replace(self, loras=loras)
def model_config(self) -> dict:
"""Read model configuration from the checkpoint metadata."""
return read_model_config(self.model_path, self.model_loader)
def meta_model(self, config: dict, module_ops: tuple[ModuleOps, ...]) -> ModelType:
"""Create a model on the meta device and apply module operations."""
return create_meta_model(self.model_class_configurator, config, module_ops)
def build(
self,
target_device: torch.device,
dtype: torch.dtype,
cpu_slots_count: int | None = None,
gpu_slots_count: int | None = None,
**_kwargs: object,
) -> BlockStreamingWrapper:
"""Build and return a ready-to-use :class:`BlockStreamingWrapper`.
Args:
target_device: GPU device for compute.
dtype: Weight dtype (e.g. ``torch.bfloat16``).
cpu_slots_count: Number of pinned CPU buffer slots.
``None`` = RAM streaming (all blocks pre-loaded with LoRA fusion).
gpu_slots_count: Number of GPU buffer slots.
``None`` = ``_DEFAULT_GPU_SLOTS`` (2).
"""
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:
meta_model = self.model_wrapper(meta_model)
meta_model.eval()
blocks = resolve_attr(meta_model, self.blocks_attr)
layout = build_pool_layout(blocks[0], dtype)
# 2. Determine slot counts.
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)
else:
source, lora_sources = self._build_disk_source(meta_model, layout, target_device, dtype, cpu_slots_count)
# 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)
)
provider = WeightsProvider(gpu_pool, copy_stream, target_device, source, lora_sources, self.blocks_prefix)
return BlockStreamingWrapper(
model=meta_model,
blocks=blocks,
provider=provider,
target_device=target_device,
)
def _build_pinned_source(
self,
meta_model: nn.Module,
target_device: torch.device,
dtype: torch.dtype,
cpu_slots_count: int,
) -> 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,
)
# 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 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
else:
non_block_sd[self.state_dict_prefix + key] = tensor.to(device=target_device, dtype=dtype)
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,
) -> tuple[WeightSource, list[LoraSource]]:
"""Create a DiskWeightSource backed by a DiskBlockReader for lazy loading."""
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,
non_block_keys,
meta_model,
target_device,
dtype,
sd_ops=self.model_sd_ops,
key_prefix=self.state_dict_prefix,
lora_sources=lora_sources,
matmul_device=target_device,
)
cpu_pool = WeightPool(
layout,
cpu_slots_count,
torch.device("cpu"),
reuse_barrier=lambda event: event.synchronize(),
pin_memory=True,
)
block_reader = DiskBlockReader(reader=reader, block_key_map=block_key_map, dtype=dtype)
source = DiskWeightSource(cpu_pool, block_reader)
return source, lora_sources
# ------------------------------------------------------------------
# Helpers
# ------------------------------------------------------------------
@staticmethod
def _fuse_lora_delta(
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."""
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))
return tensor
@staticmethod
@torch.inference_mode()
def _load_non_block_weights(
reader: DiskTensorReader,
non_block_keys: list[tuple[str, str]],
model: nn.Module,
device: torch.device,
dtype: torch.dtype,
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)
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)
@@ -0,0 +1,106 @@
"""Safetensors I/O and LoRA fusion for block streaming."""
from __future__ import annotations
import safetensors
import torch
from ltx_core.loader.sd_ops import SDOps
class DiskTensorReader:
"""Key-based tensor accessor over one or more safetensors files."""
def __init__(self, paths: list[str]) -> None:
self._handles: list[safetensors.safe_open] = []
self._key_to_handle_idx: dict[str, int] = {}
for path in paths:
handle = safetensors.safe_open(path, framework="pt", device="cpu")
handle_idx = len(self._handles)
self._handles.append(handle)
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)
def close(self) -> None:
self._handles.clear()
self._key_to_handle_idx.clear()
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.
"""
def __init__(
self,
reader: DiskTensorReader,
block_key_map: dict[int, list[tuple[str, str]]],
dtype: torch.dtype,
) -> None:
self._reader = reader
self._block_key_map = block_key_map
self._dtype = dtype
def read_into(self, target: dict[str, torch.Tensor], block_idx: int) -> None:
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)
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.
"""
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:
continue
if model_key.endswith(".lora_A.weight"):
a_keys[model_key[: -len(".lora_A.weight")]] = sft_key
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(),
)
def get_delta(self, param_prefix: str, device: torch.device | None = None) -> torch.Tensor | None:
"""Return ``(B * strength) @ A`` 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
def cleanup(self) -> None:
self._pinned_ab.clear()
@@ -0,0 +1,64 @@
"""Weight buffer pool for block streaming."""
from __future__ import annotations
from collections import deque
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]]
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.
Args:
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.
pin_memory: Pin buffers (for async H2D copies from CPU).
"""
def __init__(
self,
layout: BlockLayout,
capacity: int,
device: torch.device,
reuse_barrier: Callable[[torch.cuda.Event], None],
pin_memory: bool = False,
) -> None:
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))
@property
def capacity(self) -> int:
return self._capacity
def acquire(self) -> dict[str, torch.Tensor]:
"""Take a free buffer, waiting any pending event before returning."""
weights = self._free.popleft()
event = self._events.pop(id(weights), None)
if event is not None:
self._reuse_barrier(event)
return weights
def release(self, weights: dict[str, torch.Tensor], event: torch.cuda.Event | None = None) -> None:
"""Return a buffer to the free list.
If *event* is given it is waited on the next :meth:`acquire`
of this buffer, ensuring the prior operation has completed.
"""
if event is not None:
self._events[id(weights)] = event
self._free.append(weights)
@@ -0,0 +1,110 @@
"""GPU weights provider for block streaming."""
from __future__ import annotations
from collections import OrderedDict
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
class WeightsProvider:
"""Provides GPU-ready block weights via H2D copy from a pinned CPU weight source.
Args:
pool: Pre-allocated GPU weight buffer pool.
copy_stream: Dedicated CUDA stream for async H2D copies.
target_device: GPU device for compute.
source: Pinned CPU weight source.
lora_sources: LoRA adapters fused on H2D copy.
blocks_prefix: State-dict prefix for LoRA key matching.
"""
def __init__(
self,
pool: WeightPool,
copy_stream: torch.cuda.Stream,
target_device: torch.device,
source: WeightSource,
lora_sources: list[LoraSource] | None = None,
blocks_prefix: str = "",
) -> None:
self._copy_stream = copy_stream
self._pool = pool
self._cache: OrderedDict[int, dict[str, torch.Tensor]] = OrderedDict()
self._events: dict[int, torch.cuda.Event] = {}
self._target_device = target_device
self._source = source
self._lora_sources = lora_sources or []
self._blocks_prefix = blocks_prefix
def get(self, idx: int) -> dict[str, torch.Tensor]:
"""Return GPU weights for block *idx*. Does H2D copy on miss."""
if idx in self._cache:
return self._cache[idx]
# Evict oldest GPU buffer if at capacity.
if len(self._cache) >= self._pool.capacity:
evicted_idx, evicted_weights = self._cache.popitem(last=False)
self._pool.release(evicted_weights, event=self._events.pop(evicted_idx, None))
gpu_weights = self._pool.acquire()
cpu_weights = self._source.get(idx)
h2d_event = self._copy_to_gpu(idx, gpu_weights, cpu_weights)
self._source.release(idx, event=h2d_event)
self._cache[idx] = gpu_weights
return gpu_weights
def _copy_to_gpu(
self,
idx: int,
gpu_weights: dict[str, torch.Tensor],
cpu_weights: dict[str, torch.Tensor],
) -> torch.cuda.Event:
"""Enqueue H2D copy + LoRA fusion on the copy stream and wait on compute.
The wait is intentionally inside this method so callers -- and
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)
if self._lora_sources:
self._fuse_block_loras(idx, gpu_weights)
h2d_event = torch.cuda.Event()
h2d_event.record(self._copy_stream)
torch.cuda.current_stream(self._target_device).wait_event(h2d_event)
return h2d_event
def release(self, idx: int, event: torch.cuda.Event) -> None:
"""Attach a compute-done event -- waited before this buffer is recycled."""
self._events[idx] = event
def cleanup(self) -> None:
"""Synchronize streams and release all resources."""
self._copy_stream.synchronize()
torch.cuda.current_stream(self._target_device).synchronize()
self._cache.clear()
self._events.clear()
self._source.cleanup()
for lora in self._lora_sources:
lora.cleanup()
def __len__(self) -> int:
return len(self._cache)
def _fuse_block_loras(self, idx: int, weights: dict[str, torch.Tensor]) -> None:
"""Fuse LoRA deltas directly into GPU block weights."""
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))
@@ -0,0 +1,83 @@
"""Weight sources for block streaming: protocol and implementations."""
from __future__ import annotations
from collections import OrderedDict
from typing import Protocol
import torch
from ltx_core.block_streaming.disk import DiskBlockReader
from ltx_core.block_streaming.pool import WeightPool
class WeightSource(Protocol):
"""Provides pinned CPU weights for a given block index."""
def get(self, idx: int) -> dict[str, torch.Tensor]:
"""Return CPU weights for block *idx*."""
...
def release(self, idx: int, event: torch.cuda.Event) -> None:
"""Signal that an async operation using these weights is guarded by *event*."""
...
def cleanup(self) -> None:
"""Release all resources (buffers, readers, events)."""
...
class DiskWeightSource(WeightSource):
"""Reads block weights from disk into pinned CPU buffers on demand."""
def __init__(self, pool: WeightPool, reader: DiskBlockReader) -> None:
self._pool = pool
self._cache: OrderedDict[int, dict[str, torch.Tensor]] = OrderedDict()
self._events: dict[int, torch.cuda.Event] = {}
self._reader = reader
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:
return self._cache[idx]
if len(self._cache) >= self._pool.capacity:
evicted_idx, evicted_weights = self._cache.popitem(last=False)
self._pool.release(evicted_weights, event=self._events.pop(evicted_idx, None))
weights = self._pool.acquire()
self._reader.read_into(weights, idx)
self._cache[idx] = weights
return weights
def release(self, idx: int, event: torch.cuda.Event) -> None:
"""Attach an H2D event -- waited before this buffer is recycled."""
self._events[idx] = event
def cleanup(self) -> None:
"""Clear cache and close the disk reader."""
self._cache.clear()
self._events.clear()
self._reader.cleanup()
def __len__(self) -> int:
return len(self._cache)
class PinnedWeightSource(WeightSource):
"""Pre-loaded pinned CPU weights."""
def __init__(self, weights: dict[int, dict[str, torch.Tensor]]) -> None:
self._weights = weights
def get(self, idx: int) -> dict[str, torch.Tensor]:
return self._weights[idx]
def release(self, idx: int, event: torch.cuda.Event) -> None:
pass
def cleanup(self) -> None:
self._weights.clear()
def __len__(self) -> int:
return len(self._weights)
@@ -0,0 +1,60 @@
"""Shared utilities for the block_streaming package."""
from __future__ import annotations
import itertools
from typing import TYPE_CHECKING, Any
import torch
from torch import nn
if TYPE_CHECKING:
from ltx_core.block_streaming.pool import BlockLayout
def resolve_attr(module: nn.Module, dotted_path: str) -> nn.ModuleList:
"""Resolve a dotted attribute path like ``'model.language_model.layers'``."""
obj: Any = module
for part in dotted_path.split("."):
obj = getattr(obj, part)
if not isinstance(obj, nn.ModuleList):
raise TypeError(f"Expected nn.ModuleList at '{dotted_path}', got {type(obj).__name__}")
return obj
def assign_tensor_to_module(root: nn.Module, dotted_name: str, tensor: torch.Tensor) -> None:
"""Assign *tensor* to the parameter/buffer at *dotted_name* inside *root*.
Unlike ``param.data = tensor``, this works even when the existing parameter
lives on the ``meta`` device (which has an incompatible storage type).
"""
parts = dotted_name.split(".")
parent = root
for part in parts[:-1]:
parent = getattr(parent, part)
leaf = parts[-1]
if leaf in parent._parameters:
parent._parameters[leaf] = nn.Parameter(tensor, requires_grad=False)
elif leaf in parent._buffers:
parent._buffers[leaf] = tensor
else:
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.
"""
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()
}
@@ -0,0 +1,96 @@
"""Block streaming wrapper: streams transformer blocks through a WeightsProvider."""
from __future__ import annotations
import itertools
from typing import Any
import torch
from torch import nn
from ltx_core.block_streaming.provider import WeightsProvider
from ltx_core.block_streaming.utils import assign_tensor_to_module
class BlockStreamingWrapper(nn.Module):
"""Streams sequential model blocks through GPU buffer caches.
The wrapper delegates all weight management to a :class:`WeightsProvider`
which handles CPU-to-GPU copies, caching, LoRA fusion, and stream
synchronization internally.
Use :class:`StreamingModelBuilder` to construct this wrapper -- it
handles checkpoint parsing, source selection, and provider creation.
Args:
model: The wrapped model (non-block params already on GPU).
blocks: Sequential blocks to stream (``nn.ModuleList``).
provider: Provides GPU-ready weights on demand.
target_device: GPU device for compute.
"""
def __init__(
self,
model: nn.Module,
blocks: nn.ModuleList,
provider: WeightsProvider,
target_device: torch.device,
) -> None:
super().__init__()
self._model = model
self._blocks = blocks
self._target_device = target_device
self._provider = provider
self._hooks: list[torch.utils.hooks.RemovableHandle] = []
self._register_hooks()
# ------------------------------------------------------------------
# Hook registration
# ------------------------------------------------------------------
def _pre_hook(self, block_idx: int) -> None:
"""Load GPU weights for a block and inject them into its parameters."""
gpu_weights = self._provider.get(block_idx)
block = self._blocks[block_idx]
for name, _param in itertools.chain(block.named_parameters(), block.named_buffers()):
assign_tensor_to_module(block, name, gpu_weights[name])
def _post_hook(self, block_idx: int) -> None:
"""Record a compute-done event and release the block weights."""
compute_done = torch.cuda.Event()
compute_done.record(torch.cuda.current_stream(self._target_device))
self._provider.release(block_idx, event=compute_done)
def _register_hooks(self) -> None:
for idx, block in enumerate(self._blocks):
pre = block.register_forward_pre_hook(
lambda _mod, _args, *, idx=idx: self._pre_hook(idx),
)
post = block.register_forward_hook(
lambda _mod, _args, _out, *, idx=idx: self._post_hook(idx),
)
self._hooks.extend([pre, post])
# ------------------------------------------------------------------
# Teardown
# ------------------------------------------------------------------
def teardown(self) -> None:
"""Remove hooks and release all resources."""
for h in self._hooks:
h.remove()
self._hooks.clear()
self._provider.cleanup()
# ------------------------------------------------------------------
# Forward and attribute delegation
# ------------------------------------------------------------------
def forward(self, *args: Any, **kwargs: Any) -> Any: # noqa: ANN401
return self._model(*args, **kwargs)
def __getattr__(self, name: str) -> Any: # noqa: ANN401
"""Proxy attribute access to the wrapped model."""
try:
return super().__getattr__(name)
except AttributeError:
return getattr(self._model, name)
@@ -17,12 +17,20 @@ class VideoConditionByKeyframeIndex(ConditioningItem):
keyframes: Keyframe latents [B, C, F, H, W].
frame_idx: Frame index offset for positional encoding.
strength: Conditioning strength (1.0 = clean, 0.0 = fully denoised).
num_pixel_frames: Number of pixel frames the keyframe latent originally encodes.
"""
def __init__(self, keyframes: torch.Tensor, frame_idx: int, strength: float):
def __init__(
self,
keyframes: torch.Tensor,
frame_idx: int,
strength: float,
num_pixel_frames: int = 1,
):
self.keyframes = keyframes
self.frame_idx = frame_idx
self.strength = strength
self.num_pixel_frames = num_pixel_frames
def apply_to(
self,
@@ -41,6 +49,11 @@ class VideoConditionByKeyframeIndex(ConditioningItem):
)
positions[:, 0, ...] += self.frame_idx
# If the keyframe latent encodes a single pixel frame,
# narrow the temporal end to [start, start + 1) instead of the
# VAE-scaled range.
if self.num_pixel_frames == 1:
positions[:, 0, ..., 1:] = positions[:, 0, ..., :1] + 1
positions = positions.to(dtype=torch.float32)
positions[:, 0, ...] /= latent_tools.fps
+71
View File
@@ -0,0 +1,71 @@
"""HDR utilities: LogC3 compression for HDR IC-LoRA training and inference.
Provides compress/decompress and postprocess helpers for HDR video generation.
Used by ltx-pipelines for HDR IC-LoRA and by ltx-trainer for HDR validation.
"""
from __future__ import annotations
from typing import Literal
import torch
from torch import Tensor
class LogC3:
"""ARRI LogC3 (EI 800) HDR compression.
Maps linear [0, ∞) <-> LogC3 [0, 1] via the camera log curve. The log
curve allocates more precision to shadows/midtones and compresses
highlights smoothly. Callers are responsible for mapping the [0, 1]
output to the VAE's [-1, 1] input range.
"""
name = "LogC3"
A = 5.555556
B = 0.052272
C = 0.247190
D = 0.385537
E = 5.367655
F = 0.092809
CUT = 0.010591
def compress(self, hdr: Tensor) -> Tensor:
"""Compress linear HDR [0, ∞) → LogC3 [0, 1]."""
x = torch.clamp(hdr, min=0.0)
log_part = self.C * torch.log10(self.A * x + self.B) + self.D
lin_part = self.E * x + self.F
logc = torch.where(x >= self.CUT, log_part, lin_part)
return torch.clamp(logc, 0.0, 1.0)
def compress_ldr(self, ldr: Tensor) -> Tensor:
"""Compress LDR [0, 1] → [0, 1] (no log curve, just clamp)."""
return torch.clamp(ldr, 0.0, 1.0)
def decompress(self, logc: Tensor) -> Tensor:
"""Decompress LogC3 [0, 1] → linear HDR [0, ∞)."""
logc = torch.clamp(logc, 0.0, 1.0)
cut_log = self.E * self.CUT + self.F
lin_from_log = (torch.pow(10.0, (logc - self.D) / self.C) - self.B) / self.A
lin_from_lin = (logc - self.F) / self.E
return torch.where(logc >= cut_log, lin_from_log, lin_from_lin)
def decompress_ldr(self, logc: Tensor) -> Tensor:
"""Decompress [0, 1] → LDR [0, 1] (identity clamp)."""
return torch.clamp(logc, 0.0, 1.0)
def apply_hdr_decode_postprocess(
decoded_video: Tensor,
transform: Literal["logc3"] = "logc3",
) -> Tensor:
"""Apply HDR decompress to VAE decode output for HDR recovery.
Args:
decoded_video: Tensor from VAE decode in [0, 1], shape [B, C, F, H, W].
Must be float32 for sufficient color resolution.
transform: "logc3".
Returns:
HDR video tensor float32.
"""
decoded_video = decoded_video.float()
if transform == "logc3":
return LogC3().decompress(decoded_video)
raise ValueError(f"Unsupported HDR transform: {transform}")
@@ -1,306 +0,0 @@
"""Layer streaming wrapper for memory-efficient inference.
Keeps most transformer/decoder layers on CPU pinned memory and streams them
to GPU on demand, using a secondary CUDA stream to prefetch upcoming layers
so that data transfer overlaps with compute.
General-purpose: works with any ``nn.Module`` whose forward iterates over a
``nn.ModuleList`` attribute (e.g. ``transformer_blocks``, ``layers``).
Each layer is evicted back to CPU immediately after its forward completes,
and prefetch uses modular indexing so the last layer's prefetch wraps around
to prepare early layers for the next forward pass.
Example
-------
>>> model = build_my_model(device=torch.device("cpu"))
>>> model = LayerStreamingWrapper(
... model,
... layers_attr="transformer_blocks",
... target_device=torch.device("cuda:0"),
... prefetch_count=2,
... )
>>> out = model(inputs) # hooks handle layer streaming
>>> model.teardown() # move everything back to CPU
"""
from __future__ import annotations
import functools
import itertools
import logging
from typing import Any
import torch
from torch import nn
logger = logging.getLogger(__name__)
def _resolve_attr(module: nn.Module, dotted_path: str) -> nn.ModuleList:
"""Resolve a dotted attribute path like ``'model.language_model.layers'``."""
obj: Any = module
for part in dotted_path.split("."):
obj = getattr(obj, part)
if not isinstance(obj, nn.ModuleList):
raise TypeError(f"Expected nn.ModuleList at '{dotted_path}', got {type(obj).__name__}")
return obj
class _LayerStore:
"""Manages CPU-pinned copies of layer parameters/buffers.
Tracks which layers currently reside on GPU so the prefetcher and evictor
can make correct decisions.
"""
def __init__(self, layers: nn.ModuleList, target_device: torch.device) -> None:
self.target_device = target_device
self.num_layers = len(layers)
# CPU-pinned copies keyed by (layer_idx, param_name)
self._pinned: list[dict[str, torch.Tensor]] = []
self._on_gpu: set[int] = set()
for layer in layers:
pinned: dict[str, torch.Tensor] = {}
for name, tensor in itertools.chain(layer.named_parameters(), layer.named_buffers()):
pinned_tensor = tensor.data.pin_memory()
tensor.data = pinned_tensor
pinned[name] = pinned_tensor
self._pinned.append(pinned)
def _check_idx(self, idx: int) -> None:
if idx < 0 or idx >= self.num_layers:
raise IndexError(f"Layer index {idx} out of range [0, {self.num_layers})")
def is_on_gpu(self, idx: int) -> bool:
return idx in self._on_gpu
def move_to_gpu(self, idx: int, layer: nn.Module, *, non_blocking: bool = False) -> None:
"""Move layer *idx* parameters from pinned CPU to *target_device*."""
self._check_idx(idx)
if idx in self._on_gpu:
return
pinned = self._pinned[idx]
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
param.data = pinned[name].to(self.target_device, non_blocking=non_blocking)
self._on_gpu.add(idx)
def evict_to_cpu(self, idx: int, layer: nn.Module) -> None:
"""Swap layer *idx* parameters back to their pinned CPU copies."""
self._check_idx(idx)
if idx not in self._on_gpu:
return
pinned = self._pinned[idx]
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
param.data = pinned[name]
self._on_gpu.discard(idx)
def cleanup(self) -> None:
"""Release all pinned memory references.
After this call, the pinned tensors can be garbage-collected once
the layer parameters (which still reference them via ``.data``) are
also released (e.g. via ``.to("meta")``).
"""
for pinned_dict in self._pinned:
pinned_dict.clear()
self._pinned.clear()
class _AsyncPrefetcher:
"""Issues H2D transfers on a dedicated CUDA stream.
Uses per-layer CUDA events so that the compute stream only waits for the
specific layer it needs, not all pending transfers.
"""
def __init__(self, store: _LayerStore, layers: nn.ModuleList) -> None:
self._store = store
self._layers = layers
self._stream = torch.cuda.Stream(device=store.target_device)
self._events: dict[int, torch.cuda.Event] = {}
def prefetch(self, idx: int) -> None:
"""Begin async transfer of layer *idx* to GPU (no-op if already there)."""
if self._store.is_on_gpu(idx) or idx in self._events:
return
with torch.cuda.stream(self._stream):
self._store.move_to_gpu(idx, self._layers[idx], non_blocking=True)
event = torch.cuda.Event()
event.record(self._stream)
self._events[idx] = event
def wait(self, idx: int) -> None:
"""Block the compute stream until layer *idx* transfer is complete."""
event = self._events.pop(idx, None)
if event is not None:
torch.cuda.current_stream(self._store.target_device).wait_event(event)
def cleanup(self) -> None:
"""Drain pending work and release CUDA stream/event resources."""
self._events.clear()
self._stream = None
self._layers = None
self._store = None
class LayerStreamingWrapper(nn.Module):
"""Wraps a model to stream its sequential layers between CPU and GPU.
Each layer is evicted immediately after its forward completes, and
prefetch wraps around using modular indexing so the end of one forward
pass prepares early layers for the next.
Parameters
----------
model:
The model to wrap, with all parameters on **CPU**.
layers_attr:
Dotted attribute path to the ``nn.ModuleList`` of sequential layers
(e.g. ``"transformer_blocks"`` or ``"model.language_model.layers"``).
target_device:
The GPU device to use for compute.
prefetch_count:
How many layers ahead to prefetch. The maximum number of layers on
GPU at once is ``1 + prefetch_count``. Must be >= 1.
"""
def __init__(
self,
model: nn.Module,
layers_attr: str,
target_device: torch.device,
prefetch_count: int = 2,
) -> None:
if prefetch_count < 1:
raise ValueError("prefetch_count must be >= 1")
super().__init__()
# Store the wrapped model as a submodule so parameters are discoverable.
self._model = model
self._layers = _resolve_attr(model, layers_attr)
self._target_device = target_device
# Clamp: no point prefetching more than num_layers - 1 (the rest are evicted).
self._prefetch_count = min(prefetch_count, len(self._layers) - 1)
self._hooks: list[torch.utils.hooks.RemovableHandle] = []
self._setup()
# ------------------------------------------------------------------
# Setup / teardown
# ------------------------------------------------------------------
def _setup(self) -> None:
# 1. Build the pinned CPU store (copies all layer tensors to pinned memory).
self._store = _LayerStore(self._layers, self._target_device)
# 2. Move all NON-layer params/buffers to GPU.
layer_tensor_ids: set[int] = set()
for layer in self._layers:
for t in itertools.chain(layer.parameters(), layer.buffers()):
layer_tensor_ids.add(id(t))
for p in self._model.parameters():
if id(p) not in layer_tensor_ids:
p.data = p.data.to(self._target_device)
for b in self._model.buffers():
if id(b) not in layer_tensor_ids:
b.data = b.data.to(self._target_device)
# 3. Pre-load the first (1 + prefetch_count) layers synchronously.
for idx in range(min(self._prefetch_count + 1, len(self._layers))):
self._store.move_to_gpu(idx, self._layers[idx])
# 4. Create the async prefetcher and register hooks.
self._prefetcher = _AsyncPrefetcher(self._store, self._layers)
self._register_hooks()
def _register_hooks(self) -> None:
idx_map: dict[int, int] = {id(layer): idx for idx, layer in enumerate(self._layers)}
num_layers = len(self._layers)
def _pre_hook(
module: nn.Module,
_args: Any, # noqa: ANN401
*,
idx: int,
) -> None:
# Wait only for THIS layer's H2D transfer (not all pending ones).
self._prefetcher.wait(idx)
if not self._store.is_on_gpu(idx):
self._store.move_to_gpu(idx, module)
# Record that the compute stream will read these weight tensors.
# They were allocated on the prefetch stream, so without this the
# caching allocator would allow the prefetch stream to reuse their
# memory immediately after eviction — even if the compute kernel
# that reads them hasn't finished yet.
compute_stream = torch.cuda.current_stream(self._target_device)
for param in itertools.chain(module.parameters(), module.buffers()):
param.data.record_stream(compute_stream)
# Kick off prefetch for upcoming layers (wraps around for next pass).
for offset in range(1, self._prefetch_count + 1):
self._prefetcher.prefetch((idx + offset) % num_layers)
def _post_hook(
module: nn.Module,
_args: Any, # noqa: ANN401
_output: Any, # noqa: ANN401
*,
idx: int,
) -> None:
# Evict this layer immediately — its computation is done.
self._store.evict_to_cpu(idx, module)
for layer in self._layers:
idx = idx_map[id(layer)]
h1 = layer.register_forward_pre_hook(functools.partial(_pre_hook, idx=idx))
h2 = layer.register_forward_hook(functools.partial(_post_hook, idx=idx))
self._hooks.extend([h1, h2])
def teardown(self) -> None:
"""Remove hooks, release pinned memory, and move parameters back to CPU.
After this call the wrapper is inert: hooks are removed, the prefetch
stream is drained and destroyed, all parameters reside on regular
(non-pinned) CPU memory, and the ``_LayerStore`` pinned-tensor cache is
cleared. Callers should still follow up with ``.to("meta")`` to release
the CPU copies if the model is no longer needed.
"""
for h in self._hooks:
h.remove()
self._hooks.clear()
# Drain all in-flight async H2D copies, then release stream resources.
# Without the synchronize, clearing the stream/events can trigger
# use-after-free at the CUDA driver level.
torch.cuda.synchronize(device=self._target_device)
if self._prefetcher is not None:
self._prefetcher.cleanup()
self._prefetcher = None
# Move everything to CPU.
for idx, layer in enumerate(self._layers):
self._store.evict_to_cpu(idx, layer)
for p in self._model.parameters():
p.data = p.data.to("cpu")
for b in self._model.buffers():
b.data = b.data.to("cpu")
# Release pinned memory. After evict_to_cpu() the layer parameters
# still reference the pinned tensors (since .to("cpu") on a pinned
# tensor is a no-op). The caller is expected to follow up with
# .to("meta") to drop the param refs; cleanup() drops the store's refs.
self._store.cleanup()
# ------------------------------------------------------------------
# Forward and attribute delegation
# ------------------------------------------------------------------
def forward(self, *args: Any, **kwargs: Any) -> Any: # noqa: ANN401
return self._model(*args, **kwargs)
def __getattr__(self, name: str) -> Any: # noqa: ANN401
"""Proxy attribute access to the wrapped model.
This allows calling methods like ``encode()`` on a wrapped
GemmaTextEncoder without the caller needing to know about the wrapper.
``nn.Module.__getattr__`` is only called when normal attribute lookup
fails, so ``_model``, ``_store``, etc. are found first via ``__dict__``.
"""
try:
return super().__getattr__(name)
except AttributeError:
return getattr(self._model, name)
@@ -1,6 +1,11 @@
"""Loader utilities for model weights, LoRAs, and safetensor operations."""
from ltx_core.loader.fuse_loras import apply_loras
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 (
LoRAAdaptableProtocol,
@@ -45,4 +50,7 @@ __all__ = [
"StateDictLoader",
"StateDictRegistry",
"apply_loras",
"create_meta_model",
"load_state_dict",
"read_model_config",
]
@@ -0,0 +1,61 @@
"""Shared model-construction helpers used by both SingleGPUModelBuilder and StreamingModelBuilder."""
from __future__ import annotations
from typing import TypeVar
import torch
from torch import nn
from ltx_core.loader.module_ops import ModuleOps
from ltx_core.loader.primitives import StateDict, StateDictLoader
from ltx_core.loader.registry import Registry
from ltx_core.loader.sd_ops import SDOps
from ltx_core.model.model_protocol import ModelConfigurator
_M = TypeVar("_M", bound=nn.Module)
def load_state_dict(
paths: str | tuple[str, ...] | list[str],
loader: StateDictLoader,
registry: Registry,
device: torch.device | None,
sd_ops: SDOps | None = None,
) -> StateDict:
"""Load a state dict from disk, using registry caching."""
if isinstance(paths, str):
path_list = [paths]
elif isinstance(paths, tuple):
path_list = list(paths)
else:
path_list = paths
cached = registry.get(path_list, sd_ops)
if cached is not None:
return cached
result = loader.load(path_list, sd_ops=sd_ops, device=device)
registry.add(path_list, sd_ops=sd_ops, state_dict=result)
return result
def read_model_config(
model_path: str | tuple[str, ...],
loader: StateDictLoader,
) -> dict:
"""Read metadata from the first shard of a checkpoint."""
first = model_path[0] if isinstance(model_path, tuple) else model_path
return loader.metadata(first)
def create_meta_model(
configurator: type[ModelConfigurator[_M]],
config: dict,
module_ops: tuple[ModuleOps, ...] = (),
) -> _M:
"""Create a model on the meta device and apply module operations."""
with torch.device("meta"):
model = configurator.from_config(config)
for op in module_ops:
if op.matcher(model):
model = op.mutator(model)
return model
@@ -3,8 +3,10 @@ from dataclasses import dataclass, field, replace
from typing import Generic
import torch
from torch import nn
from ltx_core.loader.fuse_loras import apply_loras
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 (
LoRAAdaptableProtocol,
@@ -22,6 +24,56 @@ from ltx_core.model.model_protocol import ModelConfigurator, ModelType
logger: logging.Logger = logging.getLogger(__name__)
def _check_uninitialized(model: nn.Module) -> list[str]:
"""Return names of any parameters/buffers still on meta device."""
names = []
for name, param in model.named_parameters():
if str(param.device) == "meta":
names.append(name)
for name, buf in model.named_buffers():
if str(buf.device) == "meta":
names.append(name)
return names
def _load_model_weights(
meta_model: nn.Module,
model_path: str | tuple[str, ...],
loras: tuple[LoraPathStrengthAndSDOps, ...],
loader: StateDictLoader,
registry: Registry,
device: torch.device,
dtype: torch.dtype | None,
model_sd_ops: SDOps | None = None,
lora_load_device: torch.device | None = None,
) -> None:
"""Load base weights and fuse LoRAs into *meta_model* in-place."""
if lora_load_device is None:
lora_load_device = device
model_sd = load_state_dict(model_path, loader, registry, device, model_sd_ops)
lora_strengths = [lora.strength for lora in loras]
if not lora_strengths or (min(lora_strengths) == 0 and max(lora_strengths) == 0):
sd = model_sd.sd
if dtype is not None:
sd = {key: value.to(dtype=dtype) for key, value in model_sd.sd.items()}
meta_model.load_state_dict(sd, strict=False, assign=True)
return
lora_state_dicts = [load_state_dict([lora.path], loader, registry, lora_load_device, lora.sd_ops) for lora in loras]
lora_sd_and_strengths = [
LoraStateDictWithStrength(sd, strength) for sd, strength in zip(lora_state_dicts, lora_strengths, strict=True)
]
final_sd = apply_loras(
model_sd=model_sd,
lora_sd_and_strengths=lora_sd_and_strengths,
dtype=dtype,
destination_sd=model_sd if isinstance(registry, DummyRegistry) else None,
)
meta_model.load_state_dict(final_sd.sd, strict=False, assign=True)
@dataclass(frozen=True)
class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType], LoRAAdaptableProtocol):
"""
@@ -69,34 +121,22 @@ class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType],
return replace(self, lora_load_device=device)
def model_config(self) -> dict:
first_shard_path = self.model_path[0] if isinstance(self.model_path, tuple) else self.model_path
return self.model_loader.metadata(first_shard_path)
return read_model_config(self.model_path, self.model_loader)
def meta_model(self, config: dict, module_ops: tuple[ModuleOps, ...]) -> ModelType:
with torch.device("meta"):
model = self.model_class_configurator.from_config(config)
for module_op in module_ops:
if module_op.matcher(model):
model = module_op.mutator(model)
return model
return create_meta_model(self.model_class_configurator, config, module_ops)
def load_sd(
self, paths: list[str], registry: Registry, device: torch.device | None, sd_ops: SDOps | None = None
) -> StateDict:
state_dict = registry.get(paths, sd_ops)
if state_dict is None:
state_dict = self.model_loader.load(paths, sd_ops=sd_ops, device=device)
registry.add(paths, sd_ops=sd_ops, state_dict=state_dict)
return state_dict
return load_state_dict(paths, self.model_loader, registry, device, sd_ops)
def _return_model(self, meta_model: ModelType, device: torch.device) -> ModelType:
uninitialized_params = [name for name, param in meta_model.named_parameters() if str(param.device) == "meta"]
uninitialized_buffers = [name for name, buffer in meta_model.named_buffers() if str(buffer.device) == "meta"]
if uninitialized_params or uninitialized_buffers:
logger.warning(f"Uninitialized parameters or buffers: {uninitialized_params + uninitialized_buffers}")
uninitialized = _check_uninitialized(meta_model)
if uninitialized:
logger.warning(f"Uninitialized parameters or buffers: {uninitialized}")
return meta_model
retval = meta_model.to(device)
return retval
return meta_model.to(device)
def build(
self,
@@ -107,30 +147,16 @@ class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType],
device = torch.device("cuda") if device is None else device
config = self.model_config()
meta_model = self.meta_model(config, self.module_ops)
model_paths = list(self.model_path) if isinstance(self.model_path, tuple) else [self.model_path]
model_state_dict = self.load_sd(model_paths, sd_ops=self.model_sd_ops, registry=self.registry, device=device)
lora_strengths = [lora.strength for lora in self.loras]
if not lora_strengths or (min(lora_strengths) == 0 and max(lora_strengths) == 0):
sd = model_state_dict.sd
if dtype is not None:
sd = {key: value.to(dtype=dtype) for key, value in model_state_dict.sd.items()}
meta_model.load_state_dict(sd, strict=False, assign=True)
return self._return_model(meta_model, device)
lora_state_dicts = [
self.load_sd([lora.path], sd_ops=lora.sd_ops, registry=self.registry, device=self.lora_load_device)
for lora in self.loras
]
lora_sd_and_strengths = [
LoraStateDictWithStrength(sd, strength)
for sd, strength in zip(lora_state_dicts, lora_strengths, strict=True)
]
final_sd = apply_loras(
model_sd=model_state_dict,
lora_sd_and_strengths=lora_sd_and_strengths,
_load_model_weights(
meta_model=meta_model,
model_path=self.model_path,
loras=self.loras,
loader=self.model_loader,
registry=self.registry,
device=device,
dtype=dtype,
destination_sd=model_state_dict if isinstance(self.registry, DummyRegistry) else None,
model_sd_ops=self.model_sd_ops,
lora_load_device=self.lora_load_device,
)
meta_model.load_state_dict(final_sd.sd, strict=False, assign=True)
return self._return_model(meta_model, device)
@@ -698,6 +698,9 @@ class VideoDecoder(nn.Module):
When causal=False, allows future frame dependencies in convolutions but maintains same output shape.
"""
batch_size = sample.shape[0]
output_dtype = sample.dtype
weights_dtype = next(self.parameters()).dtype
sample = sample.to(weights_dtype)
# Add noise if timestep conditioning is enabled
if self.timestep_conditioning:
@@ -770,7 +773,7 @@ class VideoDecoder(nn.Module):
# Example: (B, 48, F, 128, 128) -> (B, 3, F, 512, 512) with patch_size=4
sample = unpatchify(sample, patch_size_hw=self.patch_size, patch_size_t=1)
return sample
return sample.to(output_dtype)
def _prepare_tiles(
self,
@@ -902,23 +905,34 @@ 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 uint8 chunks ``[f, h, w, c]``.
"""Decode a video latent tensor, yielding chunks ``[f, h, w, c]``.
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_to_uint8(frames: torch.Tensor) -> torch.Tensor:
frames = (((frames + 1.0) / 2.0).clamp(0.0, 1.0) * 255.0).to(torch.uint8)
frames = rearrange(frames[0], "c f h w -> f h w c")
return frames
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.
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)
if tiling_config is not None:
for frames in self.tiled_decode(latent, tiling_config, generator=generator):
yield convert_to_uint8(frames)
yield _convert(frames)
else:
decoded = self(latent, generator=generator)
yield convert_to_uint8(decoded)
yield _convert(decoded)
def _group_tiles_by_temporal_slice(self, tiles: List[Tile]) -> List[List[Tile]]:
"""Group tiles by their temporal output slice."""
+1 -1
View File
@@ -56,7 +56,7 @@ Inference pipelines for LTX-2 audio-video generation. Depends on `ltx-core` for
### Memory management
- **Model lifecycle**: All blocks build their model on call and free it on exit. `gpu_model()` moves params to `"meta"` device on exit, immediately releasing storage. No model persists between calls.
- **Layer streaming**: When `streaming_prefetch_count` is set, `DiffusionStage` wraps the transformer in `LayerStreamingWrapper`. Layers live on pinned CPU memory; only `1 + prefetch_count` layers are on GPU at a time, with async H2D prefetch on a separate CUDA stream.
- **Block streaming**: When offloading is enabled, `DiffusionStage` wraps the transformer in `BlockStreamingWrapper`. Blocks live on pinned CPU memory; only 2 blocks are buffered on GPU at a time (one for compute, one for async H2D copy on a separate CUDA stream).
- **Batch splitting**: `BatchSplitAdapter` wraps the transformer and splits inputs exceeding `max_batch_size` into sequential chunks. If guidance needs B=4 but `max_batch_size=1`, it runs 4 sequential B=1 passes. Higher `max_batch_size` reduces layer-streaming PCIe transfers at the cost of peak memory.
## Denoisers (`utils/denoisers.py`)
+19
View File
@@ -63,6 +63,7 @@ Available pipeline modules:
- `ltx_pipelines.keyframe_interpolation` - Keyframe interpolation.
- `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).
Use `--help` with any pipeline module to see all available options and parameters.
@@ -79,6 +80,9 @@ Do you have an existing video to modify?
Do you have an audio file to drive generation?
├─ YES → Use A2VidPipelineTwoStage (audio-to-video)
Do you need HDR output (linear float frames for EXR / tonemapping)?
├─ YES → Use HDRICLoraPipeline (video-to-video with LogC3 inverse decode)
Do you need to condition on existing images/videos?
├─ YES → Do you have reference videos for video-to-video?
│ ├─ YES → Use ICLoraPipeline
@@ -108,6 +112,7 @@ Do you need to condition on existing images/videos?
| **KeyframeInterpolationPipeline** | 2 | ✅ | ✅ | Keyframes | Animation, interpolation |
| **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) |
---
@@ -219,6 +224,20 @@ Single-stage generation that encodes the source video and audio into latents, ap
---
### 9. HDRICLoraPipeline
**Best for:** Video-to-video generation with HDR output for EXR export and offline tonemapping.
**Source**: [`src/ltx_pipelines/hdr_ic_lora.py`](src/ltx_pipelines/hdr_ic_lora.py)
Two-stage video-to-video on the distilled model with an HDR IC-LoRA. Decoded latents pass through an HDR inverse transform (ARRI LogC3, auto-detected from LoRA metadata) to produce a **linear HDR float** tensor `[f, h, w, c]`. Video-only (audio skipped). Text embeddings are pre-computed externally and loaded from a `.safetensors` file. Tonemapping and EXR saving are the caller's responsibility. LoRA and embeddings: [`Lightricks/LTX-2.3-22b-IC-LoRA-HDR`](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-HDR).
**Extra CLI arguments:** `--input` (mp4 or directory, required), `--output-dir` (required), `--hdr-lora` (required), `--text-embeddings` (pre-computed `.safetensors`, required), `--num-frames`, `--spatial-tile` (tiled VAE decode tile size; reduce on lower-VRAM GPUs), `--skip-mp4` (EXR only, no H.264 preview), `--exr-half` (float16 EXR), `--high-quality` (generates 2x frames internally for smoother output, ~2x slower), `--offload {none,cpu,disk}` (weight offloading; disables FP8 quantization when not `none`).
**Use when:** You need linear HDR float output for EXR export, color grading, or custom tonemapping workflows.
---
## 🎨 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.
+2 -2
View File
@@ -1,10 +1,10 @@
[project]
name = "ltx-pipelines"
version = "1.1.1"
version = "1.1.2"
description = "Pipelines implementation for Lightricks' LTX-2 model"
readme = "README.md"
requires-python = ">=3.10"
dependencies = ["ltx-core", "av", "tqdm", "pillow"]
dependencies = ["ltx-core", "av", "tqdm", "pillow", "openimageio"]
[build-system]
requires = ["uv_build>=0.9.8,<0.10.0"]
@@ -31,7 +31,7 @@ from ltx_pipelines.utils.helpers import (
get_device,
)
from ltx_pipelines.utils.media_io import decode_audio_from_file, encode_video
from ltx_pipelines.utils.types import ModalitySpec
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
class A2VidPipelineTwoStage:
@@ -53,12 +53,15 @@ class A2VidPipelineTwoStage:
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
offload_mode: OffloadMode = OffloadMode.NONE,
):
self.device = device or get_device()
self.dtype = torch.bfloat16
self._scheduler = LTX2Scheduler()
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
self.prompt_encoder = PromptEncoder(
checkpoint_path, gemma_root, self.dtype, self.device, registry=registry, offload_mode=offload_mode
)
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
self.audio_conditioner = AudioConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
self.stage_1 = DiffusionStage(
@@ -69,6 +72,7 @@ class A2VidPipelineTwoStage:
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
stage_2_loras = (*tuple(loras), *tuple(distilled_lora))
self.stage_2 = DiffusionStage(
@@ -79,6 +83,7 @@ class A2VidPipelineTwoStage:
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
self.upsampler = VideoUpsampler(
checkpoint_path, spatial_upsampler_path, self.dtype, self.device, registry=registry
@@ -102,7 +107,6 @@ class A2VidPipelineTwoStage:
audio_max_duration: float | None = None,
tiling_config: TilingConfig | None = None,
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
stage_1_sigmas: torch.Tensor | None = None,
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
@@ -117,7 +121,6 @@ class A2VidPipelineTwoStage:
[prompt, negative_prompt],
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_n, _ = ctx_n.video_encoding, ctx_n.audio_encoding
@@ -183,7 +186,6 @@ class A2VidPipelineTwoStage:
noise_scale=0.0,
initial_latent=encoded_audio_latent,
),
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
@@ -223,7 +225,6 @@ class A2VidPipelineTwoStage:
noise_scale=0.0,
initial_latent=encoded_audio_latent,
),
streaming_prefetch_count=streaming_prefetch_count,
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
@@ -266,6 +267,7 @@ def main() -> None:
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
offload_mode=args.offload_mode,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -294,7 +296,6 @@ def main() -> None:
audio_max_duration=args.audio_max_duration
if args.audio_max_duration is not None
else args.num_frames / args.frame_rate,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
@@ -34,7 +34,7 @@ from ltx_pipelines.utils.helpers import (
get_device,
)
from ltx_pipelines.utils.media_io import encode_video
from ltx_pipelines.utils.types import ModalitySpec
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
class DistilledPipeline:
@@ -54,12 +54,18 @@ class DistilledPipeline:
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
offload_mode: OffloadMode = OffloadMode.NONE,
):
self.device = device or get_device()
self.dtype = torch.bfloat16
self.prompt_encoder = PromptEncoder(
distilled_checkpoint_path, gemma_root, self.dtype, self.device, registry=registry
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.stage = DiffusionStage(
@@ -70,6 +76,7 @@ class DistilledPipeline:
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
@@ -88,7 +95,6 @@ class DistilledPipeline:
images: list[ImageConditioningInput],
tiling_config: TilingConfig | None = None,
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
stage_1_sigmas: torch.Tensor = DISTILLED_SIGMAS,
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
) -> tuple[Iterator[torch.Tensor], Audio]:
@@ -102,7 +108,6 @@ class DistilledPipeline:
[prompt],
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
streaming_prefetch_count=streaming_prefetch_count,
)
video_context, audio_context = ctx_p.video_encoding, ctx_p.audio_encoding
@@ -130,7 +135,6 @@ class DistilledPipeline:
fps=frame_rate,
video=ModalitySpec(context=video_context, conditionings=stage_1_conditionings),
audio=ModalitySpec(context=audio_context),
streaming_prefetch_count=streaming_prefetch_count,
)
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
@@ -167,7 +171,6 @@ class DistilledPipeline:
noise_scale=stage_2_sigmas[0].item(),
initial_latent=audio_state.latent,
),
streaming_prefetch_count=streaming_prefetch_count,
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
@@ -189,6 +192,7 @@ def main() -> None:
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
offload_mode=args.offload_mode,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -202,7 +206,6 @@ def main() -> None:
images=args.images,
tiling_config=tiling_config,
enhance_prompt=args.enhance_prompt,
streaming_prefetch_count=args.streaming_prefetch_count,
)
encode_video(
@@ -0,0 +1,886 @@
"""HDR IC-LoRA pipeline: two-stage video generation with HDR output.
Extends the standard IC-LoRA pipeline with HDR decode via LogC3 inverse
transform. ``__call__`` returns a **linear HDR float** tensor
``[f, h, w, c]``; tonemapping and EXR saving are the caller's
responsibility.
Text embeddings must be pre-computed externally (e.g. using
``PromptEncoder`` from ``ltx_pipelines.utils.blocks`` with a Gemma text
encoder) and saved as a ``.safetensors`` file with ``video_context``
and ``audio_context`` tensors (via ``safetensors.torch.save_file``).
The path is passed via ``text_embeddings_path``.
Run as a script for batch inference::
python -m ltx_pipelines.hdr_ic_lora \\
--input ./videos/ \\
--output-dir ./hdr-output \\
--distilled-checkpoint-path /models/ltx-2.3-22b-distilled.safetensors \\
--spatial-upsampler-path /models/ltx-2.3-spatial-upscaler-x2-1.0.safetensors \\
--hdr-lora /path/to/hdr_lora.safetensors \\
--text-embeddings /path/to/hdr_scene_emb.safetensors \\
--num-frames 161
Supports resolutions up to 4K (3840x2160 @ 121 frames on 80 GB,
49 frames on 48 GB). The caller is responsible for choosing a resolution
and frame count that fits in GPU memory. See ``--help`` for a reference
table, or use ``ltx_pipelines.utils.vram_budget.max_frames_for_resolution``
to query your specific configuration.
"""
import dataclasses
import logging
from dataclasses import replace
from pathlib import Path
import torch
from einops import rearrange
from safetensors import safe_open
from ltx_core.components.noisers import GaussianNoiser
from ltx_core.components.patchifiers import VideoLatentPatchifier
from ltx_core.conditioning import (
ConditioningItem,
VideoConditionByReferenceLatent,
)
from ltx_core.hdr import apply_hdr_decode_postprocess
from ltx_core.loader import LoraPathStrengthAndSDOps
from ltx_core.loader.registry import Registry
from ltx_core.loader.sd_ops import LTXV_LORA_COMFY_RENAMING_MAP
from ltx_core.modality_tiling import VideoModalityTilingHelper
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_pipelines.utils.blocks import (
DiffusionStage,
ImageConditioner,
VideoDecoder,
VideoUpsampler,
)
from ltx_pipelines.utils.constants import DISTILLED_SIGMA_VALUES, STAGE_2_DISTILLED_SIGMA_VALUES
from ltx_pipelines.utils.denoisers import SimpleDenoiser
from ltx_pipelines.utils.helpers import get_device, modality_from_latent_state
from ltx_pipelines.utils.media_io import ResizeMode, align_resolution, load_video_conditioning_hdr
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
DEFAULT_NUM_FRAMES = 161
MIN_RESOLUTION = 64
ALIGNMENT_DIVISOR = 64
# Conditioning videos whose spatial resolution (H x W) exceeds this value are
# encoded with the tiled encoder. The default (512 x 768) is suitable for
# H100-80GB. On lower-VRAM GPUs pass tiled_vae_encode_pixel_threshold=256*256
# to the pipeline constructor.
TILED_VAE_ENCODE_PIXEL_THRESHOLD = 512 * 768
_DEFAULT_QUANTIZATION = QuantizationPolicy.fp8_cast()
# Default stage-2 configuration: one refinement phase with modest 2-way tiling
# in every dimension and a short 2-step distilled sigma schedule.
_S2 = STAGE_2_DISTILLED_SIGMA_VALUES
_TILED_2F2H2W_OV8_6 = TileCountConfig(
frames=DimensionTilingConfig(2, 8),
height=DimensionTilingConfig(2, 6),
width=DimensionTilingConfig(2, 6),
)
STAGE2_TILINGS = [_TILED_2F2H2W_OV8_6]
STAGE2_SIGMAS = [[_S2[0], _S2[1], 0.0]]
STAGE2_USE_IC_LORA = [True]
def _clamp_dim_tiling(cfg: DimensionTilingConfig, dim_size: int, axis: str) -> DimensionTilingConfig:
"""Clamp a single dim's tile count and overlap to the latent's extent.
``split_by_count`` requires ``overlap < tile_size``; with
``tile_size = (dim_size + overlap*(n-1)) // n`` this reduces to
``overlap <= dim_size - n``. When the configured overlap exceeds this
bound it is clamped; if the latent is too small to hold ``n`` tiles
at all, tiling falls back to a single tile on this axis.
"""
n = cfg.num_tiles
if n <= 1:
return cfg
if dim_size < n:
logger.warning(
"%s tiling: dim_size=%d < num_tiles=%d; falling back to 1 tile on this axis.",
axis,
dim_size,
n,
)
return DimensionTilingConfig(1, 0)
max_overlap = dim_size - n
if cfg.overlap <= max_overlap:
return cfg
logger.warning(
"%s tiling: overlap=%d exceeds latent bound (%d); clamping to %d.",
axis,
cfg.overlap,
max_overlap,
max_overlap,
)
return DimensionTilingConfig(n, max_overlap)
def _clamp_tile_to_latent(tiling: TileCountConfig, latent_shape: tuple[int, int, int]) -> TileCountConfig:
"""Clamp frame, height, and width tilings to the latent's extents.
``latent_shape`` is ``(F, H, W)`` in latent units.
"""
f, h, w = latent_shape
return replace(
tiling,
frames=_clamp_dim_tiling(tiling.frames, f, "Frame"),
height=_clamp_dim_tiling(tiling.height, h, "Height"),
width=_clamp_dim_tiling(tiling.width, w, "Width"),
)
# Default tiling config (spatial tile 1280 px, overlap 256 px; temporal 32
# frames, overlap 16). On GPUs with < 80 GB VRAM you may need to shrink
# the spatial tile size (e.g. 768) to avoid OOM during VAE decode.
DEFAULT_SPATIAL_TILE = 1280
DEFAULT_SPATIAL_OVERLAP = 256
DEFAULT_TEMPORAL_TILE = 32
DEFAULT_TEMPORAL_OVERLAP = 16
# ---------------------------------------------------------------------------
# HDR LoRA config
# ---------------------------------------------------------------------------
@dataclasses.dataclass(frozen=True)
class HdrLoraConfig:
"""Explicit HDR LoRA parameters.
Read from LoRA safetensors metadata by :func:`read_hdr_lora_config`, or
constructed manually for testing.
"""
hdr_transform: str = "logc3"
reference_downscale_factor: int = 1
def read_hdr_lora_config(lora_path: str) -> HdrLoraConfig | None:
"""Read HDR config from LoRA safetensors metadata.
Returns ``None`` when the LoRA has no HDR metadata.
"""
try:
with safe_open(lora_path, framework="pt") as f:
metadata = f.metadata() or {}
except (OSError, ValueError) as e:
logger.warning("Failed to read metadata from LoRA file '%s': %s", lora_path, e)
return None
hdr_transform = metadata.get("hdr_transform", "")
has_hdr = bool(hdr_transform or metadata.get("use_hdr_transform"))
if not has_hdr:
return None
transform = hdr_transform if hdr_transform and hdr_transform != "true" else "logc3"
scale = int(metadata.get("reference_downscale_factor", 1))
return HdrLoraConfig(hdr_transform=transform, reference_downscale_factor=scale)
# ---------------------------------------------------------------------------
# Pipeline
# ---------------------------------------------------------------------------
class HDRICLoraPipeline:
"""Two-stage IC-LoRA pipeline with HDR support.
Same two-stage architecture as ICLoraPipeline (half-res generation + 2x
upscale refinement), with HDR decode via LogC3 inverse.
``__call__`` returns a **linear HDR float** tensor ``[f, h, w, c]``.
Tonemapping and EXR saving are the caller's responsibility.
"""
def __init__(
self,
distilled_checkpoint_path: str,
spatial_upsampler_path: str,
hdr_lora: str | Path,
text_embeddings_path: str | Path,
device: torch.device | None = None,
quantization: QuantizationPolicy = _DEFAULT_QUANTIZATION,
registry: Registry | None = None,
hdr_lora_config: HdrLoraConfig | None = None,
tiled_vae_encode_pixel_threshold: int = TILED_VAE_ENCODE_PIXEL_THRESHOLD,
offload_mode: OffloadMode = OffloadMode.NONE,
):
"""
Args:
distilled_checkpoint_path: Path to the distilled model checkpoint.
spatial_upsampler_path: Path to the spatial upsampler checkpoint.
hdr_lora: Path to the HDR IC-LoRA ``.safetensors`` file.
text_embeddings_path: Path to pre-computed text embeddings
(``.safetensors`` file with ``video_context`` and
``audio_context`` tensors).
device: Target device. Auto-detected when ``None``.
quantization: Quantization policy. Defaults to ``fp8_cast``.
registry: Optional model registry for caching loaded components.
hdr_lora_config: Explicit HDR LoRA config override. When ``None``,
auto-detected from LoRA safetensors metadata.
tiled_vae_encode_pixel_threshold: Conditioning videos whose spatial
area (H x W) exceeds this value are encoded with the tiled
encoder. Default ``512 * 768`` is suitable for 80 GB GPUs.
Use ``256 * 256`` on GPUs with less VRAM.
offload_mode: Weight offloading strategy for diffusion stages.
"""
self.device = device or get_device()
self._tiled_vae_encode_threshold = tiled_vae_encode_pixel_threshold
if offload_mode != OffloadMode.NONE and quantization is not None:
logger.info("Offload mode enabled — disabling quantization (not supported with layer streaming).")
quantization = None
self.dtype = torch.bfloat16
lora_path = str(Path(hdr_lora).resolve())
loras = (LoraPathStrengthAndSDOps(lora_path, 1.0, LTXV_LORA_COMFY_RENAMING_MAP),)
# Load pre-computed text embeddings from safetensors.
emb_path = Path(text_embeddings_path)
logger.info("Loading text embeddings from %s", emb_path)
with safe_open(emb_path, framework="pt", device=str(self.device)) as f:
self.text_embeddings: tuple[torch.Tensor, torch.Tensor] = (
f.get_tensor("video_context"),
f.get_tensor("audio_context"),
)
self.image_conditioner = ImageConditioner(distilled_checkpoint_path, self.dtype, self.device, registry=registry)
self.stage_1 = DiffusionStage(
distilled_checkpoint_path,
self.dtype,
self.device,
loras=loras,
quantization=quantization,
registry=registry,
offload_mode=offload_mode,
)
self.stage_2 = DiffusionStage(
distilled_checkpoint_path,
self.dtype,
self.device,
loras=loras,
quantization=quantization,
registry=registry,
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)
# HDR config: explicit override, or auto-detect from LoRA metadata.
if hdr_lora_config is not None:
self._hdr_config: HdrLoraConfig | None = hdr_lora_config
else:
self._hdr_config = read_hdr_lora_config(lora_path)
if self._hdr_config is not None:
logger.info("[HDR IC-LoRA] HDR mode enabled (%s decode)", self._hdr_config.hdr_transform)
@property
def hdr_transform(self) -> str:
"""Active HDR transform name (defaults to 'logc3')."""
return self._hdr_config.hdr_transform if self._hdr_config is not None else "logc3"
@property
def reference_downscale_factor(self) -> int:
"""Reference video downscale factor from HDR LoRA config."""
return self._hdr_config.reference_downscale_factor if self._hdr_config is not None else 1
def __call__( # noqa: PLR0913
self,
seed: int,
height: int,
width: int,
num_frames: int,
frame_rate: float,
video_conditioning: list[tuple[str, float]],
tiling_config: TilingConfig | None = None,
high_quality_hdr: bool = False,
stage2_tilings: list[TileCountConfig] | None = None,
stage2_sigmas: list[list[float]] | None = None,
stage2_use_ic_lora: list[bool] | None = None,
) -> torch.Tensor:
"""Generate video with IC-LoRA conditioning and HDR output.
Returns a linear HDR float tensor ``[f, h, w, c]``.
Args:
seed: Random seed for reproducibility.
height: Desired output video height in pixels. Aligned internally
to the nearest multiple of 64 (rounded up). Decoded output is
cropped back to this size.
width: Desired output video width in pixels. Same alignment rules
as *height*.
num_frames: Number of frames to generate.
frame_rate: Output video frame rate.
video_conditioning: List of (path, strength) tuples for IC-LoRA video conditioning.
high_quality_hdr: High-quality HDR mode. Duplicates each conditioning
frame and generates at 2x frame count, then keeps every other
output frame. Reduces temporal artifacts at the cost of ~2x
generation time.
Returns:
Linear HDR float tensor ``[f, h, w, c]``.
"""
# In high-quality HDR mode, generate 2*N - 1 frames internally
# (satisfies (n-1)%8==0 when N itself does), then keep every other frame.
if high_quality_hdr:
gen_num_frames = 2 * num_frames - 1
logger.info("[HDR IC-LoRA] High-quality HDR: %d -> %d internal frames", num_frames, gen_num_frames)
else:
gen_num_frames = num_frames
gen_w, gen_h, crop_w, crop_h = align_resolution(
width, height, ResizeMode.REFLECT_PAD, divisor=ALIGNMENT_DIVISOR
)
if gen_h < MIN_RESOLUTION or gen_w < MIN_RESOLUTION:
raise ValueError(
f"Resolution ({width}x{height}) is too small after alignment "
f"(got {gen_w}x{gen_h}, need at least {MIN_RESOLUTION}x{MIN_RESOLUTION})."
)
needs_crop = crop_w != gen_w or crop_h != gen_h
if needs_crop:
logger.info(
"[HDR IC-LoRA] Aligned %dx%d -> %dx%d, will crop to %dx%d",
width,
height,
gen_w,
gen_h,
crop_w,
crop_h,
)
generator = torch.Generator(device=self.device).manual_seed(seed)
noiser = GaussianNoiser(generator=generator)
video_context, _ = self.text_embeddings
# Stage 1: Initial low resolution video generation.
s1_w, s1_h = gen_w // 2, gen_h // 2
stage_1_conditionings = self.image_conditioner(
lambda enc: self._create_conditionings(
video_conditioning=video_conditioning,
height=s1_h,
width=s1_w,
video_encoder=enc,
num_frames=gen_num_frames,
tiling_config=tiling_config,
high_quality_hdr=high_quality_hdr,
)
)
stage_1_sigmas = torch.Tensor(DISTILLED_SIGMA_VALUES).to(self.device)
# HDR is video-only: skip the audio stream to avoid denoising 5B audio params.
video_state, _ = self.stage_1(
denoiser=SimpleDenoiser(video_context, None),
sigmas=stage_1_sigmas,
noiser=noiser,
width=s1_w,
height=s1_h,
frames=gen_num_frames,
fps=frame_rate,
video=ModalitySpec(
context=video_context,
conditionings=stage_1_conditionings,
),
)
if stage2_tilings is None:
stage2_tilings = list(STAGE2_TILINGS)
if stage2_sigmas is None:
stage2_sigmas = [list(s) for s in STAGE2_SIGMAS]
if stage2_use_ic_lora is None:
stage2_use_ic_lora = list(STAGE2_USE_IC_LORA)
if not (len(stage2_tilings) == len(stage2_sigmas) == len(stage2_use_ic_lora)):
raise ValueError("stage2_tilings, stage2_sigmas, and stage2_use_ic_lora must have equal length")
# Stage 2: Upsample and refine at full resolution.
upscaled_video_latent = self.upsampler(video_state.latent[:1])
stage_2_conditionings = self.image_conditioner(
lambda enc: self._create_conditionings(
video_conditioning=video_conditioning,
height=gen_h,
width=gen_w,
video_encoder=enc,
num_frames=gen_num_frames,
tiling_config=tiling_config,
high_quality_hdr=high_quality_hdr,
)
)
with self.stage_2.model_context() 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)
):
diffusion_tiling = _clamp_tile_to_latent(tiling, tuple(phase_latent.shape[2:5]))
conditionings = stage_2_conditionings if use_ic else []
sigma_t = torch.tensor(sigmas_list, dtype=torch.float32, device=self.device)
logger.info(
"[Stage 2 / phase %d] sigmas=%s ic_lora=%s tiling_h=%s tiling_w=%s",
phase_idx,
sigmas_list,
use_ic,
diffusion_tiling.height,
diffusion_tiling.width,
)
phase_latent = self._run_stage2_phase(
transformer=transformer,
latent=phase_latent,
conditionings=conditionings,
tiling=diffusion_tiling,
sigmas=sigma_t,
v_ctx=video_context,
frame_rate=frame_rate,
seed=seed,
)
final_video_latent = phase_latent
crop_size = (crop_w, crop_h) if needs_crop else None
return self._decode_video(
final_video_latent,
tiling_config,
generator,
crop_size,
high_quality_hdr=high_quality_hdr,
)
def _run_stage2_phase(
self,
transformer: object,
latent: torch.Tensor,
conditionings: list[ConditioningItem],
tiling: TileCountConfig,
sigmas: torch.Tensor,
v_ctx: torch.Tensor,
frame_rate: float,
seed: int,
) -> torch.Tensor:
"""Run one stage-2 denoising phase with optional IC-LoRA conditioning.
Each tile calls ``stage_2.run()`` with a tile-sized ``ModalitySpec`` for
video only (audio is omitted entirely for HDR). IC-LoRA conditionings
are sliced spatially to match each tile's extent.
"""
batch, n_channels, n_frames, n_height, n_width = latent.shape
full_shape = VideoLatentShape(batch=batch, channels=n_channels, frames=n_frames, height=n_height, width=n_width)
full_tools = VideoLatentTools(VideoLatentPatchifier(patch_size=1), full_shape, frame_rate)
helper = VideoModalityTilingHelper(tiling, full_tools)
ref_initial = full_tools.create_initial_state(device=self.device, dtype=self.dtype)
ref_modality = modality_from_latent_state(ref_initial, v_ctx, sigmas[0])
n_gen = full_tools.target_shape.token_count()
blend_output = torch.zeros(batch, n_gen, n_channels, device=self.device, dtype=self.dtype)
patchifier = VideoLatentPatchifier(patch_size=1)
df = self.reference_downscale_factor
for tile_idx, tile in enumerate(helper.tiles):
_, ctx = helper.tile_modality(ref_modality, tile, normalize_positions=True)
frame_s, height_s, width_s = tile.in_coords
tile_h = height_s.stop - height_s.start
tile_w = width_s.stop - width_s.start
tile_f = frame_s.stop - frame_s.start
tile_conditionings = [
VideoConditionByReferenceLatent(
latent=cond.latent[
:,
:,
frame_s,
slice(height_s.start // df, height_s.stop // df),
slice(width_s.start // df, width_s.stop // df),
].to(device=self.device, dtype=self.dtype),
downscale_factor=cond.downscale_factor,
strength=cond.strength,
)
for cond in conditionings
]
tile_video_state, _ = self.stage_2.run(
transformer=transformer,
denoiser=SimpleDenoiser(v_ctx, None),
sigmas=sigmas,
noiser=GaussianNoiser(generator=torch.Generator(device=self.device).manual_seed(seed + tile_idx)),
width=tile_w * 32,
height=tile_h * 32,
frames=(tile_f - 1) * 8 + 1,
fps=frame_rate,
video=ModalitySpec(
context=v_ctx,
conditionings=tile_conditionings,
noise_scale=sigmas[0].item(),
initial_latent=latent[:, :, frame_s, height_s, width_s].to(device=self.device, dtype=self.dtype),
),
)
tile_tokens = patchifier.patchify(tile_video_state.latent)
blend_output = helper.blend(tile_tokens, tile, ctx, blend_output)
return full_tools.unpatchify(replace(ref_initial, latent=blend_output)).latent
def _decode_video(
self,
latent: torch.Tensor,
tiling_config: TilingConfig | None,
generator: torch.Generator,
crop_size: tuple[int, int] | None = None,
*,
high_quality_hdr: bool = False,
) -> torch.Tensor:
"""Decode latent to HDR video, optionally cropping to target size.
Args:
crop_size: ``(width, height)`` to crop decoded frames to, or
``None`` to skip cropping.
high_quality_hdr: When True, keep only every other frame (undoes the
2x generation applied during high-quality HDR mode).
Returns:
Linear HDR float tensor ``[f, h, w, c]``.
"""
# 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.
latent = latent.float()
decoded = torch.cat(
list(self.video_decoder(latent, tiling_config, generator, output_dtype=torch.float32)),
dim=0,
)
decoded = rearrange(decoded, "f h w c -> 1 c f h w")
hdr = apply_hdr_decode_postprocess(decoded, transform=self.hdr_transform)
del decoded
out = rearrange(hdr[0], "c f h w -> f h w c")
if crop_size is not None:
out = out[:, : crop_size[1], : crop_size[0], :]
if high_quality_hdr:
out = out[::2]
return out
def _create_conditionings(
self,
video_conditioning: list[tuple[str, float]],
height: int,
width: int,
num_frames: int,
video_encoder: VideoEncoder,
tiling_config: TilingConfig | None = None,
high_quality_hdr: bool = False,
) -> list[ConditioningItem]:
"""Create conditioning items for video generation."""
conditionings: list[ConditioningItem] = []
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
# In high-quality HDR mode, load half the frames then duplicate each one.
load_frame_cap = (num_frames + 1) // 2 if high_quality_hdr else num_frames
for video_path, strength in video_conditioning:
video = torch.cat(
list(
load_video_conditioning_hdr(
video_path=video_path,
height=ref_height,
width=ref_width,
frame_cap=load_frame_cap,
dtype=self.dtype,
device=self.device,
hdr_transform=self.hdr_transform,
resize_mode=ResizeMode.REFLECT_PAD,
)
),
dim=2,
)
if high_quality_hdr:
video = video.repeat_interleave(2, dim=2)[:, :, :num_frames, :, :]
if tiling_config is not None and ref_height * ref_width > self._tiled_vae_encode_threshold:
encoded_video = video_encoder.tiled_encode(video, tiling_config)
else:
encoded_video = video_encoder(video)
cond = VideoConditionByReferenceLatent(
latent=encoded_video,
downscale_factor=scale,
strength=strength,
)
conditionings.append(cond)
if video_conditioning:
logger.info("[HDR IC-LoRA] Added %d video conditioning(s)", len(video_conditioning))
return conditionings
# ---------------------------------------------------------------------------
# CLI helpers
# ---------------------------------------------------------------------------
def _make_tiling_config(
spatial_tile: int = DEFAULT_SPATIAL_TILE,
spatial_overlap: int = DEFAULT_SPATIAL_OVERLAP,
temporal_tile: int = DEFAULT_TEMPORAL_TILE,
temporal_overlap: int = DEFAULT_TEMPORAL_OVERLAP,
) -> TilingConfig:
"""Build a TilingConfig from explicit sizes.
The defaults (1280 px spatial tile, 256 px overlap; 32 temporal frames,
16 overlap) are suitable for H100-80 GB. On GPUs with less VRAM,
reduce the spatial tile size (e.g. ``spatial_tile=768``).
"""
from ltx_core.model.video_vae.tiling import SpatialTilingConfig, TemporalTilingConfig # noqa: PLC0415
return TilingConfig(
spatial_config=SpatialTilingConfig(tile_size_in_pixels=spatial_tile, tile_overlap_in_pixels=spatial_overlap),
temporal_config=TemporalTilingConfig(
tile_size_in_frames=temporal_tile,
tile_overlap_in_frames=temporal_overlap,
),
)
_VIDEO_SUFFIXES = {".mp4", ".mov"}
def _collect_videos(input_path: Path) -> list[Path]:
"""Return a list of .mp4/.mov files from *input_path* (file or directory)."""
if input_path.is_file():
return [input_path]
if input_path.is_dir():
return sorted(p for p in input_path.iterdir() if p.is_file() and p.suffix.lower() in _VIDEO_SUFFIXES)
logger.error("Input %s is not a file or directory", input_path)
return []
def _process_single_video( # noqa: PLR0913
pipeline: HDRICLoraPipeline,
video_path: Path,
vid_w: int,
vid_h: int,
num_frames: int,
frame_rate: float,
output_dir: Path,
tiling_config: TilingConfig,
seed: int,
skip_mp4: bool,
exr_half: bool,
exr_executor: "ThreadPoolExecutor", # noqa: F821
exr_futures: list,
high_quality_hdr: bool = False,
) -> None:
"""Run inference on a single video: generate EXR frames + optional H.264 .mp4 preview."""
import gc # noqa: PLC0415
import time # noqa: PLC0415
from ltx_pipelines.utils.media_io import encode_exr_sequence_to_mp4, save_exr_tensor # noqa: PLC0415
output_mp4 = output_dir / f"{video_path.stem}.mp4"
exr_dir = output_dir / f"{video_path.stem}_exr"
t0 = time.time()
hdr_video = pipeline(
seed=seed,
height=vid_h,
width=vid_w,
num_frames=num_frames,
frame_rate=frame_rate,
video_conditioning=[(str(video_path), 1.0)],
tiling_config=tiling_config,
high_quality_hdr=high_quality_hdr,
)
exr_dir.mkdir(parents=True, exist_ok=True)
for j in range(hdr_video.shape[0]):
frame_cpu = hdr_video[j].cpu().clone()
path = exr_dir / f"frame_{j:05d}.exr"
exr_futures.append(exr_executor.submit(save_exr_tensor, frame_cpu, str(path), exr_half))
del hdr_video
gc.collect()
torch.cuda.empty_cache()
if not skip_mp4:
# Wait for EXR saves to finish before encoding.
for fut in exr_futures:
fut.result()
logger.info("Encoding H.264 sRGB preview: %s", video_path.name)
encode_exr_sequence_to_mp4(exr_dir, output_mp4, frame_rate)
elapsed = time.time() - t0
logger.info("Decode + encode: %.1fs | %s", elapsed, output_mp4)
# ---------------------------------------------------------------------------
# CLI entry point
# ---------------------------------------------------------------------------
def _build_arg_parser() -> "argparse.ArgumentParser": # noqa: F821
"""Build the argument parser for HDR IC-LoRA batch inference."""
import argparse # noqa: PLC0415
parser = argparse.ArgumentParser(
description="HDR IC-LoRA inference: EXR frames + tonemapped ProRes .mov.",
epilog="""\
Resolution & frame constraints
------------------------------
* Width and height must each be divisible by 32.
* Frame count must satisfy (frames - 1) %% 8 == 0.
Valid counts: 1, 9, 17, 25, ..., 121, 129, 137, 145, 153, 161.
Max frames by resolution (fp8_cast, bfloat16 VAE, tiled decode)
---------------------------------------------------------------
Resolution 80 GB (H100) 48 GB (A6000)
------------------------------------------------
720p 1280x720 161+ frames 161+ frames
1080p 1920x1080 161+ frames 161+ frames
2K 2048x1080 161+ frames 161+ frames
1440p 2560x1440 161+ frames 137 frames
4K 3840x2160 121 frames 49 frames
4K 4096x2160 105 frames 49 frames
Estimates from ltx_pipelines.utils.vram_budget. Run
python -c "from ltx_pipelines.utils.vram_budget import \\
max_frames_for_resolution as mf; print(mf(W, H, vram_gb=GB))"
to check your specific resolution and GPU.
* The tiled-encode threshold (%(tiled_threshold)s px) and the default
tiling config (%(stile)s px spatial tile) are tuned for 80 GB.
On lower-VRAM GPUs pass --spatial-tile 768 (or smaller).
"""
% {"tiled_threshold": TILED_VAE_ENCODE_PIXEL_THRESHOLD, "stile": DEFAULT_SPATIAL_TILE},
formatter_class=argparse.RawDescriptionHelpFormatter,
)
parser.add_argument("--input", required=True, help="Single .mp4 or directory of .mp4 videos.")
parser.add_argument("--output-dir", required=True, help="Directory for .mov and EXR folders.")
parser.add_argument("--hdr-lora", required=True, help="HDR IC-LoRA .safetensors file.")
parser.add_argument("--text-embeddings", required=True, help="Pre-computed text embeddings (.safetensors file).")
parser.add_argument("--distilled-checkpoint-path", required=True, help="Distilled model checkpoint (.safetensors).")
parser.add_argument("--spatial-upsampler-path", required=True, help="Spatial upsampler (.safetensors).")
parser.add_argument(
"--num-frames",
type=int,
default=DEFAULT_NUM_FRAMES,
help=f"Number of output frames. Must satisfy (n-1) %% 8 == 0 (default: {DEFAULT_NUM_FRAMES}).",
)
parser.add_argument(
"--spatial-tile",
type=int,
default=DEFAULT_SPATIAL_TILE,
help=f"Spatial tile size in pixels for tiled VAE decode (default: {DEFAULT_SPATIAL_TILE}). "
"Reduce on lower-VRAM GPUs (e.g. 768 for 48 GB).",
)
parser.add_argument("--skip-mp4", action="store_true", help="Skip H.264 MP4 encoding, only produce EXR.")
parser.add_argument("--exr-half", action="store_true", help="Save EXR as float16.")
parser.add_argument("--seed", type=int, default=10, help="Random seed (default: 10).")
parser.add_argument(
"--offload",
dest="offload_mode",
type=OffloadMode,
default=OffloadMode.NONE,
choices=list(OffloadMode),
help=(
"Weight offloading strategy. "
"'none' keeps all weights on GPU (default). "
"'cpu' pins weights in CPU RAM, streams to GPU per layer. "
"'disk' reads weights from disk on demand (lowest memory). "
"Example: --offload cpu"
),
)
parser.add_argument(
"--high-quality",
action="store_true",
help="High-quality HDR mode. Generates at 2x frame count internally "
"and keeps every other frame for smoother output. ~2x slower.",
)
return parser
@torch.inference_mode()
def main() -> None:
"""Batch HDR IC-LoRA inference: per-frame EXR + tonemapped ProRes .mov."""
import time # noqa: PLC0415
from concurrent.futures import ThreadPoolExecutor # noqa: PLC0415
from ltx_pipelines.utils.media_io import get_videostream_metadata # noqa: PLC0415
logging.basicConfig(level=logging.INFO)
args = _build_arg_parser().parse_args()
high_quality = args.high_quality
num_frames = args.num_frames
tiling_config = _make_tiling_config(spatial_tile=args.spatial_tile)
input_path = Path(args.input)
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
videos = _collect_videos(input_path)
if not videos:
logger.error("No valid videos to process.")
return
logger.info("Found %d video(s), generating %d frames each", len(videos), num_frames)
logger.info("Loading pipeline...")
pipeline = HDRICLoraPipeline(
distilled_checkpoint_path=args.distilled_checkpoint_path,
spatial_upsampler_path=args.spatial_upsampler_path,
hdr_lora=args.hdr_lora,
text_embeddings_path=args.text_embeddings,
offload_mode=args.offload_mode,
)
logger.info("Pipeline loaded.")
exr_executor = ThreadPoolExecutor(max_workers=4)
exr_futures: list = []
total_t0 = time.time()
successes = 0
for i, video_path in enumerate(videos, 1):
meta = get_videostream_metadata(str(video_path))
vid_w, vid_h = meta.width, meta.height
logger.info("%s", "=" * 60)
logger.info("[%d/%d] %s (%dx%d, %df)", i, len(videos), video_path.name, vid_w, vid_h, num_frames)
_process_single_video(
pipeline=pipeline,
video_path=video_path,
vid_w=vid_w,
vid_h=vid_h,
num_frames=num_frames,
frame_rate=meta.fps,
output_dir=output_dir,
tiling_config=tiling_config,
seed=args.seed,
skip_mp4=args.skip_mp4,
exr_half=args.exr_half,
exr_executor=exr_executor,
exr_futures=exr_futures,
high_quality_hdr=high_quality,
)
successes += 1
infer_elapsed = time.time() - total_t0
logger.info("%s", "=" * 60)
logger.info("All inference done in %.0fs (%d/%d OK)", infer_elapsed, successes, len(videos))
if exr_futures:
t0 = time.time()
logger.info("Waiting for %d EXR saves...", len(exr_futures))
for fut in exr_futures:
fut.result()
exr_wait = time.time() - t0
if exr_wait > 0.1:
logger.info("EXR save wait: %.1fs", exr_wait)
logger.info("Total wall time: %.0fs", time.time() - total_t0)
if __name__ == "__main__":
main()
@@ -39,7 +39,7 @@ from ltx_pipelines.utils.constants import (
from ltx_pipelines.utils.denoisers import SimpleDenoiser
from ltx_pipelines.utils.helpers import assert_resolution, combined_image_conditionings, get_device
from ltx_pipelines.utils.media_io import decode_video_by_frame, encode_video, video_preprocess
from ltx_pipelines.utils.types import ModalitySpec
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
class ICLoraPipeline:
@@ -63,12 +63,18 @@ class ICLoraPipeline:
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
offload_mode: OffloadMode = OffloadMode.NONE,
):
self.device = device or get_device()
self.dtype = torch.bfloat16
self.prompt_encoder = PromptEncoder(
distilled_checkpoint_path, gemma_root, self.dtype, self.device, registry=registry
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.stage_1 = DiffusionStage(
@@ -79,6 +85,7 @@ class ICLoraPipeline:
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
self.stage_2 = DiffusionStage(
distilled_checkpoint_path,
@@ -88,6 +95,7 @@ class ICLoraPipeline:
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
@@ -125,7 +133,6 @@ class ICLoraPipeline:
conditioning_attention_strength: float = 1.0,
skip_stage_2: bool = False,
conditioning_attention_mask: torch.Tensor | None = None,
streaming_prefetch_count: int | None = None,
stage_1_sigmas: torch.Tensor = DISTILLED_SIGMAS,
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
) -> tuple[Iterator[torch.Tensor], Audio]:
@@ -175,7 +182,6 @@ class ICLoraPipeline:
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
video_context, audio_context = ctx_p.video_encoding, ctx_p.audio_encoding
@@ -219,7 +225,6 @@ class ICLoraPipeline:
audio=ModalitySpec(
context=audio_context,
),
streaming_prefetch_count=streaming_prefetch_count,
)
if skip_stage_2:
@@ -264,7 +269,6 @@ class ICLoraPipeline:
noise_scale=stage_2_sigmas[0].item(),
initial_latent=audio_state.latent,
),
streaming_prefetch_count=streaming_prefetch_count,
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
@@ -463,6 +467,7 @@ def main() -> None:
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
offload_mode=args.offload_mode,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -479,7 +484,6 @@ def main() -> None:
conditioning_attention_strength=conditioning_attention_strength,
skip_stage_2=args.skip_stage_2,
conditioning_attention_mask=conditioning_attention_mask,
streaming_prefetch_count=args.streaming_prefetch_count,
)
encode_video(
@@ -15,7 +15,11 @@ from ltx_core.loader.registry import Registry
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
from ltx_core.quantization import QuantizationPolicy
from ltx_core.types import Audio, VideoPixelShape
from ltx_pipelines.utils.args import ImageConditioningInput, default_2_stage_arg_parser, detect_checkpoint_path
from ltx_pipelines.utils.args import (
ImageConditioningInput,
default_2_stage_arg_parser,
detect_checkpoint_path,
)
from ltx_pipelines.utils.blocks import (
AudioDecoder,
DiffusionStage,
@@ -35,7 +39,7 @@ from ltx_pipelines.utils.helpers import (
image_conditionings_by_adding_guiding_latent,
)
from ltx_pipelines.utils.media_io import encode_video
from ltx_pipelines.utils.types import ModalitySpec
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
class KeyframeInterpolationPipeline:
@@ -59,12 +63,15 @@ class KeyframeInterpolationPipeline:
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
offload_mode: OffloadMode = OffloadMode.NONE,
):
self.device = device or get_device()
self.dtype = torch.bfloat16
self._scheduler = LTX2Scheduler()
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
self.prompt_encoder = PromptEncoder(
checkpoint_path, gemma_root, self.dtype, self.device, registry=registry, offload_mode=offload_mode
)
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
self.stage_1 = DiffusionStage(
checkpoint_path,
@@ -74,6 +81,7 @@ class KeyframeInterpolationPipeline:
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
stage_2_loras = (*tuple(loras), *tuple(distilled_lora))
self.stage_2 = DiffusionStage(
@@ -84,6 +92,7 @@ class KeyframeInterpolationPipeline:
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
self.upsampler = VideoUpsampler(
checkpoint_path, spatial_upsampler_path, self.dtype, self.device, registry=registry
@@ -106,7 +115,6 @@ class KeyframeInterpolationPipeline:
images: list[ImageConditioningInput],
tiling_config: TilingConfig | None = None,
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
stage_1_sigmas: torch.Tensor | None = None,
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
@@ -122,7 +130,6 @@ class KeyframeInterpolationPipeline:
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
@@ -179,7 +186,6 @@ class KeyframeInterpolationPipeline:
audio=ModalitySpec(
context=a_context_p,
),
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
@@ -218,7 +224,6 @@ class KeyframeInterpolationPipeline:
noise_scale=stage_2_sigmas[0].item(),
initial_latent=audio_state.latent,
),
streaming_prefetch_count=streaming_prefetch_count,
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
@@ -241,6 +246,7 @@ def main() -> None:
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
offload_mode=args.offload_mode,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -271,7 +277,6 @@ def main() -> None:
),
images=args.images,
tiling_config=tiling_config,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
@@ -36,7 +36,7 @@ from ltx_pipelines.utils.media_io import (
encode_video,
get_videostream_metadata,
)
from ltx_pipelines.utils.types import ModalitySpec
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
class RetakePipeline:
@@ -74,6 +74,7 @@ class RetakePipeline:
registry: Registry | None = None,
distilled: bool = True,
torch_compile: bool = False,
offload_mode: OffloadMode = OffloadMode.NONE,
):
self.device = device or get_device()
self.dtype = torch.bfloat16
@@ -86,6 +87,7 @@ class RetakePipeline:
dtype=self.dtype,
device=self.device,
registry=registry,
offload_mode=offload_mode,
)
self.image_conditioner = ImageConditioner(
checkpoint_path=checkpoint_path,
@@ -107,6 +109,7 @@ class RetakePipeline:
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
self.video_decoder = VideoDecoder(
checkpoint_path=checkpoint_path,
@@ -141,7 +144,6 @@ class RetakePipeline:
regenerate_audio: bool = True,
enhance_prompt: bool = False,
tiling_config: TilingConfig | None = None,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
sigmas: torch.Tensor | None = None,
) -> tuple[Iterator[torch.Tensor], torch.Tensor]:
@@ -210,7 +212,6 @@ class RetakePipeline:
prompts_to_encode,
enhance_first_prompt=enhance_prompt,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = contexts[0].video_encoding, contexts[0].audio_encoding
@@ -269,7 +270,6 @@ class RetakePipeline:
fps=output_shape.fps,
video=video_modality_spec,
audio=audio_modality_spec,
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
@@ -309,6 +309,7 @@ def main() -> None:
quantization=args.quantization,
distilled=args.distilled,
torch_compile=args.compile,
offload_mode=args.offload_mode,
)
params = detect_params(args.distilled_checkpoint_path)
tiling_config = TilingConfig.default()
@@ -321,7 +322,6 @@ def main() -> None:
video_guider_params=params.video_guider_params,
audio_guider_params=params.audio_guider_params,
tiling_config=tiling_config,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
video_chunks_number = get_video_chunks_number(src.frames, tiling_config)
@@ -20,7 +20,11 @@ from ltx_pipelines.utils import (
combined_image_conditionings,
get_device,
)
from ltx_pipelines.utils.args import ImageConditioningInput, default_1_stage_arg_parser, detect_checkpoint_path
from ltx_pipelines.utils.args import (
ImageConditioningInput,
default_1_stage_arg_parser,
detect_checkpoint_path,
)
from ltx_pipelines.utils.blocks import (
AudioDecoder,
DiffusionStage,
@@ -31,7 +35,7 @@ from ltx_pipelines.utils.blocks import (
from ltx_pipelines.utils.constants import detect_params
from ltx_pipelines.utils.denoisers import FactoryGuidedDenoiser
from ltx_pipelines.utils.media_io import encode_video
from ltx_pipelines.utils.types import ModalitySpec
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
class TI2VidOneStagePipeline:
@@ -52,6 +56,7 @@ class TI2VidOneStagePipeline:
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
offload_mode: OffloadMode = OffloadMode.NONE,
):
self.dtype = torch.bfloat16
self.device = device or get_device()
@@ -62,6 +67,7 @@ class TI2VidOneStagePipeline:
dtype=self.dtype,
device=self.device,
registry=registry,
offload_mode=offload_mode,
)
self.image_conditioner = ImageConditioner(
checkpoint_path=checkpoint_path,
@@ -77,6 +83,7 @@ class TI2VidOneStagePipeline:
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
self.video_decoder = VideoDecoder(
checkpoint_path=checkpoint_path,
@@ -105,7 +112,6 @@ class TI2VidOneStagePipeline:
audio_guider_params: MultiModalGuiderParams | MultiModalGuiderFactory,
images: list[ImageConditioningInput],
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
tiling_config: TilingConfig | None = None,
max_batch_size: int = 1,
sigmas: torch.Tensor | None = None,
@@ -121,7 +127,6 @@ class TI2VidOneStagePipeline:
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
@@ -170,7 +175,6 @@ class TI2VidOneStagePipeline:
audio=ModalitySpec(
context=a_context_p,
),
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
@@ -192,6 +196,7 @@ def main() -> None:
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
offload_mode=args.offload_mode,
)
video, audio = pipeline(
prompt=args.prompt,
@@ -219,7 +224,6 @@ def main() -> None:
stg_blocks=args.audio_stg_blocks,
),
images=args.images,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
@@ -15,7 +15,11 @@ from ltx_core.loader.registry import Registry
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
from ltx_core.quantization import QuantizationPolicy
from ltx_core.types import Audio, VideoPixelShape
from ltx_pipelines.utils.args import ImageConditioningInput, default_2_stage_arg_parser, detect_checkpoint_path
from ltx_pipelines.utils.args import (
ImageConditioningInput,
default_2_stage_arg_parser,
detect_checkpoint_path,
)
from ltx_pipelines.utils.blocks import (
AudioDecoder,
DiffusionStage,
@@ -35,7 +39,7 @@ from ltx_pipelines.utils.helpers import (
get_device,
)
from ltx_pipelines.utils.media_io import encode_video
from ltx_pipelines.utils.types import ModalitySpec
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
class TI2VidTwoStagesPipeline:
@@ -58,12 +62,15 @@ class TI2VidTwoStagesPipeline:
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
offload_mode: OffloadMode = OffloadMode.NONE,
):
self.device = device or get_device()
self.dtype = torch.bfloat16
self._scheduler = LTX2Scheduler()
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
self.prompt_encoder = PromptEncoder(
checkpoint_path, gemma_root, self.dtype, self.device, registry=registry, offload_mode=offload_mode
)
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
self.upsampler = VideoUpsampler(
checkpoint_path, spatial_upsampler_path, self.dtype, self.device, registry=registry
@@ -79,6 +86,7 @@ class TI2VidTwoStagesPipeline:
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
self.stage_2 = DiffusionStage(
checkpoint_path,
@@ -88,6 +96,7 @@ class TI2VidTwoStagesPipeline:
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
def __call__( # noqa: PLR0913
@@ -105,7 +114,6 @@ class TI2VidTwoStagesPipeline:
images: list[ImageConditioningInput],
tiling_config: TilingConfig | None = None,
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
stage_1_sigmas: torch.Tensor | None = None,
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
@@ -121,7 +129,6 @@ class TI2VidTwoStagesPipeline:
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
@@ -170,7 +177,6 @@ class TI2VidTwoStagesPipeline:
fps=frame_rate,
video=ModalitySpec(context=v_context_p, conditionings=stage_1_conditionings),
audio=ModalitySpec(context=a_context_p),
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
@@ -208,7 +214,6 @@ class TI2VidTwoStagesPipeline:
noise_scale=stage_2_sigmas[0].item(),
initial_latent=audio_state.latent,
),
streaming_prefetch_count=streaming_prefetch_count,
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
@@ -231,6 +236,7 @@ def main() -> None:
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
offload_mode=args.offload_mode,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -261,7 +267,6 @@ def main() -> None:
),
images=args.images,
tiling_config=tiling_config,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
@@ -33,7 +33,7 @@ from ltx_pipelines.utils.helpers import (
)
from ltx_pipelines.utils.media_io import encode_video
from ltx_pipelines.utils.samplers import res2s_audio_video_denoising_loop
from ltx_pipelines.utils.types import ModalitySpec
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
class TI2VidTwoStagesHQPipeline:
@@ -61,6 +61,7 @@ class TI2VidTwoStagesHQPipeline:
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
offload_mode: OffloadMode = OffloadMode.NONE,
):
self.device = device or get_device()
self.dtype = torch.bfloat16
@@ -77,7 +78,9 @@ class TI2VidTwoStagesHQPipeline:
sd_ops=distilled_lora[0].sd_ops,
)
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
self.prompt_encoder = PromptEncoder(
checkpoint_path, gemma_root, self.dtype, self.device, registry=registry, offload_mode=offload_mode
)
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
self.upsampler = VideoUpsampler(
checkpoint_path, spatial_upsampler_path, self.dtype, self.device, registry=registry
@@ -93,6 +96,7 @@ class TI2VidTwoStagesHQPipeline:
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
self.stage_2 = DiffusionStage(
checkpoint_path,
@@ -102,6 +106,7 @@ class TI2VidTwoStagesHQPipeline:
quantization=quantization,
registry=registry,
torch_compile=torch_compile,
offload_mode=offload_mode,
)
@torch.inference_mode()
@@ -120,7 +125,6 @@ class TI2VidTwoStagesHQPipeline:
images: list[ImageConditioningInput],
tiling_config: TilingConfig | None = None,
enhance_prompt: bool = False,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
stage_1_sigmas: torch.Tensor | None = None,
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
@@ -136,7 +140,6 @@ class TI2VidTwoStagesHQPipeline:
enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
streaming_prefetch_count=streaming_prefetch_count,
)
v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
@@ -190,7 +193,6 @@ class TI2VidTwoStagesHQPipeline:
video=ModalitySpec(context=v_context_p, conditionings=stage_1_conditionings),
audio=ModalitySpec(context=a_context_p),
loop=res2s_audio_video_denoising_loop,
streaming_prefetch_count=streaming_prefetch_count,
max_batch_size=max_batch_size,
)
@@ -231,7 +233,6 @@ class TI2VidTwoStagesHQPipeline:
initial_latent=audio_state.latent,
),
loop=res2s_audio_video_denoising_loop,
streaming_prefetch_count=streaming_prefetch_count,
)
decoded_video = self.video_decoder(video_state.latent, tiling_config, generator)
@@ -254,6 +255,7 @@ def main() -> None:
loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization,
torch_compile=args.compile,
offload_mode=args.offload_mode,
)
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
@@ -284,7 +286,6 @@ def main() -> None:
),
images=args.images,
tiling_config=tiling_config,
streaming_prefetch_count=args.streaming_prefetch_count,
max_batch_size=args.max_batch_size,
)
@@ -12,6 +12,7 @@ from ltx_pipelines.utils.constants import (
LTX_2_3_PARAMS,
PipelineParams,
)
from ltx_pipelines.utils.types import OffloadMode
class ImageConditioningInput(NamedTuple):
@@ -231,16 +232,19 @@ def basic_arg_parser(
except ValueError as e:
raise argparse.ArgumentTypeError(f"must be an integer, got {value}") from e
# Layer streaming
# Weight offloading
parser.add_argument(
"--streaming-prefetch-count",
type=_positive_int,
default=None,
metavar="N",
"--offload",
dest="offload_mode",
type=OffloadMode,
default=OffloadMode.NONE,
choices=list(OffloadMode),
help=(
"Enable layer streaming prefetching N layers ahead. "
"At most 1 + N layers reside on GPU at once. "
"Must be >= 1. Example: --streaming-prefetch-count 2"
"Weight offloading strategy. "
"'none' keeps all weights on GPU (default). "
"'cpu' pins weights in CPU RAM, streams to GPU per layer. "
"'disk' reads weights from disk on demand (lowest memory). "
"Example: --offload cpu"
),
)
@@ -15,11 +15,11 @@ from typing import Callable, TypeVar
import torch
from ltx_core.batch_split import BatchSplitAdapter
from ltx_core.block_streaming import DISK_CPU_SLOTS, StreamingModelBuilder
from ltx_core.components.diffusion_steps import EulerDiffusionStep
from ltx_core.components.noisers import Noiser
from ltx_core.components.patchifiers import AudioPatchifier, VideoLatentPatchifier
from ltx_core.components.protocols import DiffusionStepProtocol
from ltx_core.layer_streaming import LayerStreamingWrapper
from ltx_core.loader import SDOps
from ltx_core.loader.primitives import LoraPathStrengthAndSDOps
from ltx_core.loader.registry import DummyRegistry, Registry
@@ -70,7 +70,7 @@ from ltx_pipelines.utils.helpers import (
generate_enhanced_prompt,
)
from ltx_pipelines.utils.samplers import euler_denoising_loop
from ltx_pipelines.utils.types import Denoiser, ModalitySpec
from ltx_pipelines.utils.types import Denoiser, ModalitySpec, OffloadMode
logger = logging.getLogger(__name__)
@@ -85,34 +85,24 @@ _M = TypeVar("_M", bound=torch.nn.Module)
@contextmanager
def _streaming_model(
model: _M,
layers_attr: str,
builder: StreamingModelBuilder,
offload_mode: OffloadMode,
target_device: torch.device,
prefetch_count: int,
) -> Iterator[_M]:
"""Wrap *model* with :class:`LayerStreamingWrapper`, yield it, then tear down."""
wrapped = LayerStreamingWrapper(
model,
layers_attr=layers_attr,
dtype: torch.dtype,
) -> Iterator:
"""Build a streaming wrapper, yield it, then tear down and free memory."""
cpu_slots_count = DISK_CPU_SLOTS if offload_mode == OffloadMode.DISK else None
wrapped = builder.build(
target_device=target_device,
prefetch_count=prefetch_count,
dtype=dtype,
cpu_slots_count=cpu_slots_count,
)
try:
yield wrapped # type: ignore[misc]
yield wrapped
finally:
wrapped.teardown()
wrapped.to("meta")
cleanup_memory()
# Flush the host (pinned) memory cache so that freed pinned pages are
# returned to the OS. Without this, sequential streaming models
# (e.g. text encoder then transformer) exhaust host memory because the
# CachingHostAllocator keeps freed blocks cached indefinitely.
torch.cuda.synchronize(device=target_device)
try:
if hasattr(torch._C, "_host_emptyCache"):
torch._C._host_emptyCache()
except Exception:
logger.warning("Host empty cache cleanup failed; ignoring.", exc_info=True)
def _build_state(
@@ -163,11 +153,30 @@ class DiffusionStage:
quantization: QuantizationPolicy | None = None,
registry: Registry | None = None,
torch_compile: bool = False,
offload_mode: OffloadMode = OffloadMode.NONE,
) -> None:
if offload_mode != OffloadMode.NONE:
if torch_compile:
raise ValueError("torch.compile is not supported with layer streaming")
if quantization is not None:
raise ValueError("quantization is not supported with layer streaming")
self._streaming_builder = StreamingModelBuilder(
model_class_configurator=LTXModelConfigurator,
model_path=checkpoint_path,
model_sd_ops=LTXV_MODEL_COMFY_RENAMING_MAP,
loras=tuple(loras),
registry=registry or DummyRegistry(),
blocks_attr="velocity_model.transformer_blocks",
blocks_prefix="transformer_blocks",
state_dict_prefix="velocity_model.",
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,
@@ -205,22 +214,21 @@ class DiffusionStage:
builder = self._transformer_builder.with_module_ops(module_ops).with_sd_ops(sd_ops).with_loras(loras)
return X0Model(builder.build(device=target, **kwargs)).to(target).eval()
def _transformer_ctx(
self,
streaming_prefetch_count: int | None,
**kwargs: object,
) -> AbstractContextManager:
if streaming_prefetch_count is not None:
return _streaming_model(
self._build_transformer(device=torch.device("cpu"), **kwargs),
layers_attr="velocity_model.transformer_blocks",
target_device=self._device,
prefetch_count=streaming_prefetch_count,
)
def _transformer_ctx(self, **kwargs: object) -> AbstractContextManager:
if self._offload_mode != OffloadMode.NONE:
return _streaming_model(self._streaming_builder, self._offload_mode, self._device, self._dtype)
return gpu_model(self._build_transformer(**kwargs))
def __call__( # noqa: PLR0913
def model_context(self, **kwargs: object) -> AbstractContextManager:
"""Build the transformer, yield it, then free its memory on exit.
Keyword arguments are forwarded to the underlying builder (e.g.
``video_tools`` required by ``TiledDataParallelBuilder``).
"""
return self._transformer_ctx(**kwargs)
def run( # noqa: PLR0913
self,
transformer: object,
denoiser: Denoiser,
sigmas: torch.Tensor,
noiser: Noiser,
@@ -232,27 +240,14 @@ class DiffusionStage:
audio: ModalitySpec | None = None,
stepper: DiffusionStepProtocol | None = None,
loop: Callable[..., tuple[LatentState | None, LatentState | None]] | None = None,
streaming_prefetch_count: int | None = None,
max_batch_size: int = 1,
) -> tuple[LatentState | None, LatentState | None]:
"""Build transformer → run denoising loop → free transformer.
Args:
width: Output width in pixels.
height: Output height in pixels.
frames: Number of output frames.
fps: Frame rate.
loop: Denoising loop function. Must accept
``(sigmas, video_state, audio_state, stepper, transformer, denoiser)``
as the first six positional arguments. When ``None``, resolves to
:func:`euler_denoising_loop` at call time.
streaming_prefetch_count: When set, build the transformer on CPU and
wrap with :class:`LayerStreamingWrapper` for memory-efficient
inference, prefetching this many layers ahead.
max_batch_size: Maximum batch size per transformer forward pass.
Guided denoisers make up to 4 transformer calls per step.
When set to a value > 1, the transformer batches multiple
calls together, reducing layer-streaming PCIe transfers.
Default ``1`` preserves sequential behavior.
"""Run denoising with a pre-built transformer.
Same semantics as ``__call__`` but accepts a pre-built transformer so
the model can be shared across multiple calls (e.g. tiled inference
inside a single ``model_context()`` block). Audio supports
``ModalitySpec(frozen=True)`` to keep the latent unchanged throughout
denoising while still providing cross-modal context to the transformer.
Returns ``(video_state | None, audio_state | None)`` with cleared
conditionings and unpatchified latents for present modalities.
"""
@@ -261,7 +256,6 @@ class DiffusionStage:
if loop is None:
loop = euler_denoising_loop
if stepper is None:
stepper = EulerDiffusionStep()
@@ -281,28 +275,70 @@ class DiffusionStage:
audio_tools = AudioLatentTools(AudioPatchifier(patch_size=1), a_shape)
audio_state = _build_state(audio, audio_tools, noiser, self._dtype, self._device)
with self._transformer_ctx(streaming_prefetch_count, video_tools=video_tools) as base_transformer:
transformer = BatchSplitAdapter(base_transformer, max_batch_size=max_batch_size)
video_state, audio_state = loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
transformer=transformer,
denoiser=denoiser,
)
wrapped = BatchSplitAdapter(transformer, max_batch_size=max_batch_size) # type: ignore[arg-type]
video_state, audio_state = loop(
sigmas=sigmas,
video_state=video_state,
audio_state=audio_state,
stepper=stepper,
transformer=wrapped,
denoiser=denoiser,
)
# Post-process: clear conditionings and unpatchify
if video_state is not None and video_tools is not None:
video_state = video_tools.clear_conditioning(video_state)
video_state = video_tools.unpatchify(video_state)
if audio_state is not None and audio_tools is not None:
audio_state = audio_tools.clear_conditioning(audio_state)
audio_state = audio_tools.unpatchify(audio_state)
return video_state, audio_state
def __call__( # noqa: PLR0913
self,
denoiser: Denoiser,
sigmas: torch.Tensor,
noiser: Noiser,
width: int,
height: int,
frames: int,
fps: float,
video: ModalitySpec | None = None,
audio: ModalitySpec | None = None,
stepper: DiffusionStepProtocol | None = None,
loop: Callable[..., tuple[LatentState | None, LatentState | None]] | None = None,
max_batch_size: int = 1,
) -> tuple[LatentState | None, LatentState | None]:
"""Build transformer -> run denoising loop -> free transformer.
Returns ``(video_state | None, audio_state | None)`` with cleared
conditionings and unpatchified latents for present modalities.
"""
# Build video_tools up front so it can be forwarded to the transformer
# context (required by TiledDataParallelBuilder in multi-GPU mode).
# `run()` rebuilds its own tools internally; the duplication is cheap.
video_tools: LatentTools | None = None
if video is not None:
pixel_shape = VideoPixelShape(batch=1, frames=frames, height=height, width=width, fps=fps)
v_shape = VideoLatentShape.from_pixel_shape(pixel_shape)
video_tools = VideoLatentTools(VideoLatentPatchifier(patch_size=1), v_shape, fps)
with self._transformer_ctx(video_tools=video_tools) as transformer:
return self.run(
transformer,
denoiser,
sigmas,
noiser,
width,
height,
frames,
fps,
video,
audio,
stepper,
loop,
max_batch_size,
)
# ---------------------------------------------------------------------------
# PromptEncoder
@@ -322,9 +358,11 @@ class PromptEncoder:
dtype: torch.dtype,
device: torch.device,
registry: Registry | None = None,
offload_mode: OffloadMode = OffloadMode.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
@@ -337,6 +375,15 @@ class PromptEncoder:
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,
@@ -344,17 +391,9 @@ class PromptEncoder:
registry=registry or DummyRegistry(),
)
def _text_encoder_ctx(
self,
streaming_prefetch_count: int | None,
) -> AbstractContextManager:
if streaming_prefetch_count is not None:
return _streaming_model(
self._text_encoder_builder.build(device=torch.device("cpu"), dtype=self._dtype).eval(),
layers_attr="model.model.language_model.layers",
target_device=self._device,
prefetch_count=streaming_prefetch_count,
)
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())
def __call__(
@@ -364,10 +403,9 @@ class PromptEncoder:
enhance_first_prompt: bool = False,
enhance_prompt_image: str | None = None,
enhance_prompt_seed: int = 42,
streaming_prefetch_count: int | None = None,
) -> list[EmbeddingsProcessorOutput]:
"""Encode *prompts* through Gemma embeddings processor, freeing each model after use."""
with self._text_encoder_ctx(streaming_prefetch_count) as text_encoder:
"""Encode *prompts* through Gemma -> embeddings processor, freeing each model after use."""
with self._text_encoder_ctx() as text_encoder:
if enhance_first_prompt:
prompts = list(prompts)
prompts[0] = generate_enhanced_prompt(
@@ -490,10 +528,17 @@ class VideoDecoder:
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."""
"""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.
"""
decoder = self._decoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
return _cleanup_iter(decoder.decode_video(latent, tiling_config, generator), decoder)
return _cleanup_iter(decoder.decode_video(latent, tiling_config, generator, output_dtype=output_dtype), decoder)
# ---------------------------------------------------------------------------
@@ -37,6 +37,11 @@ def cleanup_memory() -> None:
gc.collect()
torch.cuda.empty_cache()
torch.cuda.synchronize()
try:
if hasattr(torch._C, "_host_emptyCache"):
torch._C._host_emptyCache()
except Exception:
logging.warning("Host empty cache cleanup failed; ignoring.", exc_info=True)
def _conform_latent_length(latent: torch.Tensor, expected_frames_count: int) -> torch.Tensor:
@@ -1,23 +1,34 @@
import enum
import logging
import math
from collections.abc import Generator, Iterator
from fractions import Fraction
from io import BytesIO
from pathlib import Path
import av
import numpy as np
import OpenImageIO
import torch
from einops import rearrange
from PIL import Image
from torch._prims_common import DeviceLikeType
from tqdm import tqdm
from ltx_core.hdr import LogC3
from ltx_core.types import Audio, VideoPixelShape
from ltx_pipelines.utils.constants import DEFAULT_IMAGE_CRF
logger = logging.getLogger(__name__)
class ResizeMode(enum.Enum):
"""How to fit a conditioning video to the target resolution."""
CENTER_CROP = "center_crop"
REFLECT_PAD = "reflect_pad"
def resize_aspect_ratio_preserving(image: torch.Tensor, long_side: int) -> torch.Tensor:
"""
Resize image preserving aspect ratio (filling target long side).
@@ -79,6 +90,16 @@ def normalize_latent(latent: torch.Tensor, device: torch.device, dtype: torch.dt
return (latent / 127.5 - 1.0).to(device=device, dtype=dtype)
def to_vae_range(x: torch.Tensor) -> torch.Tensor:
"""Map [0, 1] to [-1, 1] (VAE input convention)."""
return torch.clamp(x, 0.0, 1.0) * 2.0 - 1.0
def from_vae_range(z: torch.Tensor) -> torch.Tensor:
"""Map [-1, 1] (VAE output convention) to [0, 1]."""
return torch.clamp((z + 1.0) / 2.0, 0.0, 1.0)
def load_image_and_preprocess(
image_path: str,
height: int,
@@ -124,6 +145,108 @@ def video_preprocess(
return result
def align_resolution(
width: int,
height: int,
resize_mode: ResizeMode,
divisor: int = 64,
) -> tuple[int, int, int, int]:
"""Compute aligned generation dimensions and crop-back size.
Args:
width: Source video width (need not be aligned).
height: Source video height (need not be aligned).
resize_mode: CENTER_CROP rounds down; REFLECT_PAD rounds up.
divisor: Alignment divisor (default 64 for two-stage pipelines).
Returns:
``(gen_width, gen_height, crop_width, crop_height)`` where
``gen_*`` are multiples of *divisor* and ``crop_*`` are the
original dimensions to trim back to after decoding. When no
cropping is needed ``crop_*`` equals ``gen_*``.
"""
if resize_mode is ResizeMode.REFLECT_PAD:
gen_w = ((width + divisor - 1) // divisor) * divisor
gen_h = ((height + divisor - 1) // divisor) * divisor
else:
gen_w = (width // divisor) * divisor
gen_h = (height // divisor) * divisor
crop_w = width if gen_w != width else gen_w
crop_h = height if gen_h != height else gen_h
return gen_w, gen_h, crop_w, crop_h
def resize_and_reflect_pad(tensor: torch.Tensor, height: int, width: int) -> torch.Tensor:
"""Resize tensor to fit within target, then reflect-pad to exact dimensions.
Unlike resize_and_center_crop which stretches and crops, this preserves the
original aspect ratio and pads the shorter dimension with reflected pixels.
When the target is already >= the source in both dimensions, interpolation
is skipped entirely to preserve original pixels.
Args:
tensor: Input with shape (H, W, C) or (F, H, W, C)
height: Target height
width: Target width
Returns:
Tensor with shape (1, C, 1, height, width) for 3D or (1, C, F, height, width) for 4D
"""
if tensor.ndim == 3:
tensor = rearrange(tensor, "h w c -> 1 c h w")
elif tensor.ndim == 4:
tensor = rearrange(tensor, "f h w c -> f c h w")
else:
raise ValueError(f"Expected input with 3 or 4 dimensions; got shape {tensor.shape}.")
_, _, src_h, src_w = tensor.shape
if height >= src_h and width >= src_w:
new_h, new_w = src_h, src_w
else:
scale = min(height / src_h, width / src_w)
new_h = round(src_h * scale)
new_w = round(src_w * scale)
tensor = torch.nn.functional.interpolate(tensor, size=(new_h, new_w), mode="bilinear", align_corners=False)
pad_bottom = height - new_h
pad_right = width - new_w
if pad_bottom > 0 or pad_right > 0:
pad_mode = "reflect" if pad_bottom < new_h and pad_right < new_w else "replicate"
tensor = torch.nn.functional.pad(tensor, (0, pad_right, 0, pad_bottom), mode=pad_mode)
tensor = rearrange(tensor, "f c h w -> 1 c f h w")
return tensor
def load_video_conditioning_hdr(
video_path: str,
height: int,
width: int,
frame_cap: int,
dtype: torch.dtype,
device: torch.device,
hdr_transform: str = "logc3",
resize_mode: ResizeMode = ResizeMode.CENTER_CROP,
) -> Iterator[torch.Tensor]:
"""Load a video and yield preprocessed frames for HDR IC-LoRA conditioning.
Decodes through the standard path and applies the LDR compression that
matches training. Callers are responsible for providing Rec.709 SDR
input the HDR IC-LoRA was trained on that color space.
Args:
hdr_transform: LDR-compression name (currently only ``logc3``).
resize_mode: How to fit the video to the target resolution.
Yields:
Per-frame tensors of shape ``(1, C, 1, height, width)``.
"""
if hdr_transform != "logc3":
raise ValueError(f"Unsupported HDR transform: {hdr_transform}")
resize_fn = resize_and_reflect_pad if resize_mode is ResizeMode.REFLECT_PAD else resize_and_center_crop
for f in decode_video_by_frame(path=video_path, frame_cap=frame_cap, device=device):
frame = resize_fn(f.to(torch.float32), height, width)
ldr = (frame / 255.0).clamp(0.0, 1.0)
compressed = LogC3().compress_ldr(ldr)
yield to_vae_range(compressed).to(device=device, dtype=dtype)
def decode_image(image_path: str) -> np.ndarray:
image = Image.open(image_path)
np_array = np.array(image)[..., :3]
@@ -481,3 +604,85 @@ def preprocess(image: np.array, crf: float = DEFAULT_IMAGE_CRF) -> np.array:
with BytesIO(video_bytes) as video_file:
image_array = decode_single_frame(video_file)
return image_array
def save_exr_tensor(tensor: torch.Tensor, file_path: str | Path, half: bool = False) -> None:
"""Save a single tensor frame as EXR with linear sRGB colorspace metadata.
Args:
tensor: ``[H, W, C]`` or ``[C, H, W]`` float tensor.
file_path: Output path (e.g. ``frame_0000.exr``).
half: Force float16 output with ZIP compression.
"""
if tensor.dim() == 3 and tensor.shape[0] == 3:
tensor = tensor.permute(1, 2, 0)
use_half = half or tensor.dtype in (torch.float16, torch.half)
img_np = np.ascontiguousarray(tensor.cpu().numpy().astype(np.float32))
file_path = str(file_path)
h, w = img_np.shape[:2]
fmt = OpenImageIO.HALF if use_half else OpenImageIO.FLOAT
spec = OpenImageIO.ImageSpec(w, h, 3, fmt)
spec.channelnames = ("R", "G", "B")
spec.attribute("compression", "zip")
spec.attribute("chromaticities", "float[8]", (0.64, 0.33, 0.30, 0.60, 0.15, 0.06, 0.3127, 0.3290))
spec.attribute("colorSpace", "sRGB")
out = OpenImageIO.ImageOutput.create(file_path)
if out is None:
raise RuntimeError(
f"Failed to create EXR writer for '{file_path}'. Ensure OpenImageIO is built with OpenEXR support."
)
try:
if not out.open(file_path, spec):
raise RuntimeError(f"Failed to open EXR file '{file_path}': {out.geterror()}")
if not out.write_image(img_np):
raise RuntimeError(f"Failed to write EXR image '{file_path}': {out.geterror()}")
finally:
out.close()
def _linear_to_srgb(x: np.ndarray) -> np.ndarray:
"""Linear -> sRGB OETF per IEC 61966-2-1. Input assumed in [0, 1]."""
x = np.clip(x, 0.0, 1.0)
return np.where(x <= 0.0031308, x * 12.92, 1.055 * np.power(x, 1.0 / 2.4) - 0.055)
def encode_exr_sequence_to_mp4(exr_dir: Path, output_mp4: Path, frame_rate: float) -> None:
"""Convert a linear EXR frame sequence to sRGB and encode to H.264 .mp4 via PyAV.
Exposure is fixed at EV=0 (no gain). Each EXR frame is clamped to [0, 1],
passed through the sRGB OETF, quantised to 8-bit BGR, and fed to a libx264
stream (crf 18, yuv420p). ``frame_rate`` is the original source video's
frame rate so playback matches the input timing.
"""
import os # noqa: PLC0415
os.environ["OPENCV_IO_ENABLE_OPENEXR"] = "1"
import cv2 # noqa: PLC0415
exr_files = sorted(exr_dir.glob("frame_*.exr"))
if not exr_files:
raise FileNotFoundError(f"No EXR frames found in {exr_dir}")
container = av.open(str(output_mp4), mode="w")
stream = container.add_stream("libx264", rate=Fraction(frame_rate).limit_denominator(1000))
stream.pix_fmt = "yuv420p"
stream.options = {"crf": "18", "movflags": "+faststart"}
try:
for i, exr_path in enumerate(exr_files):
hdr = cv2.imread(str(exr_path), cv2.IMREAD_UNCHANGED).astype(np.float32)
sdr = _linear_to_srgb(np.maximum(hdr, 0.0))
bgr8 = (sdr * 255.0 + 0.5).astype(np.uint8)
if i == 0:
stream.height = bgr8.shape[0]
stream.width = bgr8.shape[1]
frame = av.VideoFrame.from_ndarray(bgr8, format="bgr24")
for packet in stream.encode(frame):
container.mux(packet)
for packet in stream.encode():
container.mux(packet)
finally:
container.close()
@@ -1,4 +1,7 @@
from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
from typing import Protocol
import torch
@@ -74,3 +77,21 @@ class ModalitySpec:
noise_scale: float = 1.0
frozen: bool = False
initial_latent: torch.Tensor | None = None
class OffloadMode(Enum):
"""Weight offloading strategy.
Controls where model weights reside during inference:
- ``NONE``: All weights on GPU (no streaming). Fastest inference,
requires enough VRAM for the full model (~28 GB for LTX-2).
- ``CPU``: Weights pinned in CPU RAM, streamed layer-by-layer to a
small GPU buffer. First pass reads from disk; subsequent passes
reuse the CPU cache. Requires ~36 GB RAM + ~5 GB VRAM.
- ``DISK``: Weights read from disk on demand through a small CPU
buffer, then streamed to GPU. Every pass re-reads from disk.
Lowest memory: ~5 GB RAM + ~5 GB VRAM.
"""
NONE = "none"
CPU = "cpu"
DISK = "disk"
@@ -219,9 +219,6 @@ validation:
# Set to null to disable validation during training
interval: 100
# Number of videos to generate per prompt
videos_per_prompt: 1
# Classifier-free guidance scale
# Higher values = stronger adherence to prompt but may introduce artifacts
guidance_scale: 4.0
@@ -231,9 +231,6 @@ validation:
# Set to null to disable validation during training
interval: 100
# Number of videos to generate per prompt
videos_per_prompt: 1
# Classifier-free guidance scale
# Higher values = stronger adherence to prompt but may introduce artifacts
guidance_scale: 4.0
@@ -232,9 +232,6 @@ validation:
# Set to null to disable validation during training
interval: 100
# Number of videos to generate per prompt
videos_per_prompt: 1
# Classifier-free guidance scale
# Higher values = stronger adherence to prompt but may introduce artifacts
guidance_scale: 4.0
@@ -262,7 +262,6 @@ validation:
seed: 42 # Random seed for reproducibility
inference_steps: 30 # Number of inference steps
interval: 100 # Steps between validation runs
videos_per_prompt: 1 # Videos generated per prompt
guidance_scale: 4.0 # CFG guidance strength
stg_scale: 1.0 # STG guidance strength (0.0 to disable)
stg_blocks: [ 29 ] # Transformer blocks to perturb for STG
+2 -2
View File
@@ -1,6 +1,6 @@
[project]
name = "ltx-trainer"
version = "1.1.1"
version = "1.1.2"
description = "LTX-2 training, democratized."
readme = "README.md"
authors = [
@@ -48,7 +48,7 @@ build-backend = "hatchling.build"
[tool.ruff]
target-version = "1.1.1"
target-version = "1.1.2"
line-length = 120
[tool.ruff.lint]
@@ -241,8 +241,8 @@ def main(
help="Path to input video/image file or directory containing media files",
exists=True,
),
output: Path | None = typer.Option( # noqa: B008
None,
output: Path = typer.Option( # noqa: B008
...,
"--output",
"-o",
help="Path to json output file for reference video paths. "
@@ -260,12 +260,6 @@ class ValidationConfig(ConfigBaseModel):
gt=0,
)
videos_per_prompt: int = Field(
default=1,
description="Number of videos to generate per validation prompt",
gt=0,
)
guidance_scale: float = Field(
default=4.0,
description="CFG guidance scale to use during validation",
@@ -1,3 +1,4 @@
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import torch
@@ -155,40 +156,77 @@ class PrecomputedDataset(Dataset):
return source_paths
def _discover_samples(self) -> dict[str, list[Path]]:
"""Discover all valid sample files across all data sources."""
# Use first data source as the reference to discover samples
"""Discover all valid sample files across all data sources.
Uses a fast two-pass approach: first globs all sources in parallel to build
full-path sets in memory, then checks expected paths via set membership.
This avoids O(N * num_sources) stat calls on networked filesystems while
correctly handling path remapping (e.g. latent_X.pt -> condition_X.pt).
"""
if not self.data_sources:
raise ValueError("No data sources configured")
data_key = "latents" if "latents" in self.data_sources else next(iter(self.data_sources.keys()))
data_path = self.source_paths[data_key]
data_files = list(data_path.glob("**/*.pt"))
# Pass 1: Glob all sources in parallel, build full-path sets
def _glob_source(dir_name: str) -> tuple[list[Path], set[str]]:
source_path = self.source_paths[dir_name]
paths = list(source_path.glob("**/*.pt"))
path_set = {str(p) for p in paths}
return paths, path_set
with ThreadPoolExecutor(max_workers=len(self.data_sources)) as executor:
glob_results = dict(
zip(
self.data_sources.keys(),
executor.map(_glob_source, self.data_sources.keys()),
strict=True,
)
)
# Get primary source files (cached from glob, no second scan)
data_files, _ = glob_results[data_key]
if not data_files:
raise ValueError(f"No data files found in {data_path}")
data_files.sort()
# Initialize sample files dict
sample_files = {output_key: [] for output_key in self.data_sources.values()}
# Log source sizes
for dir_name, (paths, _) in glob_results.items():
logger.debug(f"Source {dir_name}: {len(paths)} files")
# Build path sets for non-primary sources
other_path_sets = {
dir_name: path_set for dir_name, (_, path_set) in glob_results.items() if dir_name != data_key
}
# Pass 2: For each primary file, check if expected paths exist in other sources' sets
sample_files: dict[str, list[Path]] = {output_key: [] for output_key in self.data_sources.values()}
valid_count = 0
# For each data file, find corresponding files in other sources
for data_file in data_files:
rel_path = data_file.relative_to(data_path)
# Check if corresponding files exist in ALL sources
if self._all_source_files_exist(data_file, rel_path):
# Check all other sources via set lookup (O(1) per source, no stat calls)
all_exist = True
for dir_name, path_set in other_path_sets.items():
expected = self._get_expected_file_path(dir_name, data_file, rel_path)
if str(expected) not in path_set:
logger.debug(f"Skipping {data_file.name}: no matching {dir_name} file at {expected}")
all_exist = False
break
if all_exist:
self._fill_sample_data_files(data_file, rel_path, sample_files)
valid_count += 1
skipped = len(data_files) - valid_count
if skipped > 0:
logger.info(f"Fast index: {valid_count} valid samples from {len(data_files)} total ({skipped} skipped)")
else:
logger.debug(f"Fast index: {valid_count} valid samples from {len(data_files)} total")
return sample_files
def _all_source_files_exist(self, data_file: Path, rel_path: Path) -> bool:
"""Check if corresponding files exist in all data sources."""
for dir_name in self.data_sources:
expected_path = self._get_expected_file_path(dir_name, data_file, rel_path)
if not expected_path.exists():
logger.warning(
f"No matching {dir_name} file found for: {data_file.name} (expected in: {expected_path})"
)
return False
return True
def _get_expected_file_path(self, dir_name: str, data_file: Path, rel_path: Path) -> Path:
"""Get the expected file path for a given data source."""
source_path = self.source_paths[dir_name]
@@ -207,11 +245,14 @@ class PrecomputedDataset(Dataset):
def _validate_setup(self) -> None:
"""Validate that the dataset setup is correct."""
if not self.sample_files:
raise ValueError("No valid samples found - all data sources must have matching files")
sample_counts = {key: len(files) for key, files in self.sample_files.items()}
if not sample_counts or all(count == 0 for count in sample_counts.values()):
raise ValueError(
f"No valid samples found in {self.data_root} - all configured data sources "
f"({list(self.data_sources)}) must have matching files (per-source counts: {sample_counts})"
)
# Verify all output keys have the same number of samples
sample_counts = {key: len(files) for key, files in self.sample_files.items()}
if len(set(sample_counts.values())) > 1:
raise ValueError(f"Mismatched sample counts across sources: {sample_counts}")
@@ -1,7 +1,7 @@
import shutil
import tempfile
from pathlib import Path
from typing import List, Union
from typing import List, Optional, Union
import imageio
from huggingface_hub import HfApi, create_repo
@@ -12,7 +12,11 @@ from ltx_trainer import logger
from ltx_trainer.config import LtxTrainerConfig
def push_to_hub(weights_path: Path, sampled_videos_paths: List[Path], config: LtxTrainerConfig) -> None:
def push_to_hub(
weights_path: Path,
sampled_videos_paths: Optional[List[Path]],
config: LtxTrainerConfig,
) -> None:
"""Push the trained LoRA weights to HuggingFace Hub."""
if not config.hub.hub_model_id:
logger.warning("⚠️ HuggingFace hub_model_id not specified, skipping push to hub")
@@ -108,7 +112,7 @@ def convert_video_to_gif(video_path: Path, output_path: Path) -> None:
def _create_model_card(
output_dir: Union[str, Path],
videos: List[Path],
videos: Optional[List[Path]],
config: LtxTrainerConfig,
) -> Path:
"""Generate and save a model card for the trained model."""
@@ -3,7 +3,6 @@ This module defines the abstract base class that all training strategies must im
along with the base configuration class.
"""
import random
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import Any, Literal
@@ -251,13 +250,17 @@ class TrainingStrategy(ABC):
device: Target device
first_frame_conditioning_p: Probability of conditioning on the first frame
Returns:
Boolean mask where True indicates first frame tokens (if conditioning is enabled)
Boolean mask where True indicates first frame tokens (if conditioning is enabled).
The conditioning decision is drawn independently per batch element so the training
signal across samples in a batch is i.i.d.
"""
conditioning_mask = torch.zeros(batch_size, sequence_length, dtype=torch.bool, device=device)
if first_frame_conditioning_p > 0 and random.random() < first_frame_conditioning_p:
if first_frame_conditioning_p > 0:
first_frame_end_idx = height * width
if first_frame_end_idx < sequence_length:
conditioning_mask[:, :first_frame_end_idx] = True
# Per-sample Bernoulli draw so each batch element is independently conditioned.
per_sample_condition = torch.rand(batch_size, device=device) < first_frame_conditioning_p
conditioning_mask[per_sample_condition, :first_frame_end_idx] = True
return conditioning_mask
@@ -5,12 +5,15 @@ with optional audio support.
from fractions import Fraction
from pathlib import Path
from typing import Literal
import av
import numpy as np
import torch
from torch import Tensor
VideoFormat = Literal["CFHW", "FCHW"]
def get_video_frame_count(video_path: str | Path) -> int:
"""Get the number of frames in a video file.
@@ -68,6 +71,7 @@ def save_video(
fps: float = 24.0,
audio: torch.Tensor | None = None,
audio_sample_rate: int | None = None,
video_format: VideoFormat | None = None,
) -> None:
"""Save a video tensor to a file using PyAV, optionally with audio.
Args:
@@ -76,12 +80,16 @@ def save_video(
fps: Frames per second for the output video
audio: Optional audio tensor of shape [C, samples] or [samples, C] in range [-1, 1]
audio_sample_rate: Sample rate for the audio (required if audio is provided)
video_format: Explicit layout of ``video_tensor``, either ``"CFHW"`` or ``"FCHW"``.
When ``None`` (default), the layout is auto-detected using a heuristic that only
works when ``shape[1] > 3`` the ambiguous ``[C=3, F=3, H, W]`` / ``[F=3, C=3, H, W]``
case requires passing this argument explicitly.
"""
output_path = Path(output_path)
output_path.parent.mkdir(parents=True, exist_ok=True)
# Normalize to [F, H, W, C] uint8 numpy array
video_np = _prepare_video_array(video_tensor)
video_np = _prepare_video_array(video_tensor, video_format=video_format)
_, height, width, _ = video_np.shape
with av.open(str(output_path), mode="w") as container:
@@ -113,11 +121,21 @@ def save_video(
_write_audio(container, audio_stream, audio, audio_sample_rate)
def _prepare_video_array(video_tensor: torch.Tensor) -> np.ndarray:
"""Convert video tensor to [F, H, W, C] uint8 numpy array."""
# Handle [C, F, H, W] vs [F, C, H, W] format
if video_tensor.shape[0] == 3 and video_tensor.shape[1] > 3:
def _prepare_video_array(
video_tensor: torch.Tensor,
video_format: VideoFormat | None = None,
) -> np.ndarray:
"""Convert video tensor to [F, H, W, C] uint8 numpy array.
If ``video_format`` is provided, it is trusted. Otherwise, the layout is auto-detected
using a heuristic that only fires when ``shape[0] == 3 and shape[1] > 3`` (CFHW). The
ambiguous ``[C=3, F=3, H, W]`` / ``[F=3, C=3, H, W]`` case cannot be disambiguated and
defaults to the FCHW interpretation callers must pass ``video_format`` explicitly for
3-frame CFHW tensors.
"""
if video_format == "CFHW":
video_tensor = video_tensor.permute(1, 0, 2, 3) # [C, F, H, W] -> [F, C, H, W]
elif video_format is None and video_tensor.shape[0] == 3 and video_tensor.shape[1] > 3:
video_tensor = video_tensor.permute(1, 0, 2, 3)
# Normalize to [0, 255] uint8
if video_tensor.max() <= 1.0:
Generated
+104 -4
View File
@@ -2063,7 +2063,7 @@ wheels = [
[[package]]
name = "ltx-core"
version = "1.1.1"
version = "1.1.2"
source = { editable = "packages/ltx-core" }
dependencies = [
{ name = "accelerate" },
@@ -2121,11 +2121,13 @@ dev = [{ name = "scikit-image", specifier = ">=0.25.2" }]
[[package]]
name = "ltx-pipelines"
version = "1.1.1"
version = "1.1.2"
source = { editable = "packages/ltx-pipelines" }
dependencies = [
{ name = "av" },
{ name = "ltx-core" },
{ name = "openimageio", version = "3.0.16.0", source = { registry = "https://pypi.org/simple" }, marker = "extra == 'extra-8-ltx-core-fp8-trtllm'" },
{ name = "openimageio", version = "3.1.11.0", source = { registry = "https://pypi.org/simple" }, marker = "extra == 'extra-8-ltx-core-xformers' or extra != 'extra-8-ltx-core-fp8-trtllm'" },
{ name = "pillow", version = "10.3.0", source = { registry = "https://pypi.org/simple" }, marker = "extra == 'extra-8-ltx-core-fp8-trtllm'" },
{ name = "pillow", version = "12.1.0", source = { registry = "https://pypi.org/simple" }, marker = "extra == 'extra-8-ltx-core-xformers' or extra != 'extra-8-ltx-core-fp8-trtllm'" },
{ name = "tqdm" },
@@ -2135,13 +2137,14 @@ dependencies = [
requires-dist = [
{ name = "av" },
{ name = "ltx-core", editable = "packages/ltx-core" },
{ name = "openimageio" },
{ name = "pillow" },
{ name = "tqdm" },
]
[[package]]
name = "ltx-trainer"
version = "1.1.1"
version = "1.1.2"
source = { editable = "packages/ltx-trainer" }
dependencies = [
{ name = "accelerate" },
@@ -4044,6 +4047,103 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/86/8a/69176a64335aed183529207ba8bc3d329c2999d852b4f3818027203f50e6/opencv_python_headless-4.11.0.86-cp37-abi3-win_amd64.whl", hash = "sha256:6c304df9caa7a6a5710b91709dd4786bf20a74d57672b3c31f7033cc638174ca", size = 39402386, upload-time = "2025-01-16T13:52:56.418Z" },
]
[[package]]
name = "openimageio"
version = "3.0.16.0"
source = { registry = "https://pypi.org/simple" }
resolution-markers = [
"python_full_version >= '3.13' and sys_platform == 'darwin'",
"python_full_version == '3.12.*' and sys_platform == 'darwin'",
"python_full_version >= '3.13' and platform_machine == 'aarch64' and sys_platform == 'linux'",
"python_full_version == '3.12.*' and platform_machine == 'aarch64' and sys_platform == 'linux'",
"python_full_version >= '3.13' and platform_machine != 'aarch64' and sys_platform == 'linux'",
"python_full_version == '3.12.*' and platform_machine != 'aarch64' and sys_platform == 'linux'",
"python_full_version >= '3.13' and sys_platform != 'darwin' and sys_platform != 'linux'",
"python_full_version == '3.12.*' and sys_platform != 'darwin' and sys_platform != 'linux'",
"python_full_version == '3.11.*' and sys_platform == 'darwin'",
"python_full_version == '3.11.*' and platform_machine == 'aarch64' and sys_platform == 'linux'",
"python_full_version == '3.11.*' and platform_machine != 'aarch64' and sys_platform == 'linux'",
"python_full_version == '3.11.*' and sys_platform != 'darwin' and sys_platform != 'linux'",
"python_full_version < '3.11' and sys_platform == 'darwin'",
"python_full_version < '3.11' and platform_machine == 'aarch64' and sys_platform == 'linux'",
"python_full_version < '3.11' and platform_machine != 'aarch64' and sys_platform == 'linux'",
"python_full_version < '3.11' and sys_platform != 'darwin' and sys_platform != 'linux'",
]
sdist = { url = "https://files.pythonhosted.org/packages/ab/45/f4fbc1f32b75493383ad6dc1820ecb789c552f2d231a7880c17a8ea7b276/openimageio-3.0.16.0.tar.gz", hash = "sha256:5308f04f30c98841012c61eba4972b4e0e94446be6ab4c2b7414904ed0c711c7", size = 6482737, upload-time = "2026-03-01T03:43:38.412Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/25/17/9160d6b4332b430a8b24bb5d0db39af2b2c3b483d36d11c5fa8bc9c9468b/openimageio-3.0.16.0-cp310-cp310-macosx_10_15_x86_64.whl", hash = "sha256:aa96e4100248ecc9bbf0b7fb06b1f6ce781592481de1ebfe2a665679c4219381", size = 6557400, upload-time = "2026-03-01T03:42:39.892Z" },
{ url = "https://files.pythonhosted.org/packages/2f/26/cf15bcbce65ed45f64fa0806889592c18b598795b7f2461429250a066d00/openimageio-3.0.16.0-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:8d4e65c80d3dca5e5daa1663efe620c6c94476be8c86181118f54aaae141cf79", size = 6194167, upload-time = "2026-03-01T03:42:41.685Z" },
{ url = "https://files.pythonhosted.org/packages/71/55/660e086ee3acbbbab2b19f8cad93c0b0cc1044a660d0189ea9bf0ed20e93/openimageio-3.0.16.0-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:84e7c8619c1a583730c73c8c8a247584f5fca980e51d9ec5ad1fd3a34a58aa4e", size = 6391703, upload-time = "2026-03-01T03:42:43.063Z" },
{ url = "https://files.pythonhosted.org/packages/0c/22/91cd80c88ec41d773e33790895a8b4732884bb5b571d760f0730ea9e184a/openimageio-3.0.16.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:b68b7406314b45eb813da12426b3d60046813b0c32c2fadc0ce54264554cfa23", size = 6595708, upload-time = "2026-03-01T03:42:44.732Z" },
{ url = "https://files.pythonhosted.org/packages/05/95/694fb96fd499c3765fb5239193bef4496e69b98a39eda7c0f1c2e3ebe056/openimageio-3.0.16.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:fe0df5af0f255f8abaa9aeb728c4cbdea4bc144a19f510fb0323959b88ddade5", size = 6492976, upload-time = "2026-03-01T03:42:46.473Z" },
{ url = "https://files.pythonhosted.org/packages/6b/b8/15a357fee630b0016b57bbcbc7e67565dca4d379b0a9916cc10ddf774ae2/openimageio-3.0.16.0-cp310-cp310-win_amd64.whl", hash = "sha256:58bd1f20d6e37759a3d6337f7d2a2929e47e1833ed1f439a1dc3ff841a335d7e", size = 7098417, upload-time = "2026-03-01T03:42:48.093Z" },
{ url = "https://files.pythonhosted.org/packages/1a/18/a6711865161c82e056ce4001b859a49b1c28830ced8c14dec99bffedb50f/openimageio-3.0.16.0-cp311-cp311-macosx_10_15_x86_64.whl", hash = "sha256:f657f4541dd469e2c9d6cda8b266b33327742142616645ecf971433955ea5eb2", size = 6558184, upload-time = "2026-03-01T03:42:50.653Z" },
{ url = "https://files.pythonhosted.org/packages/28/49/5ff1c6580ac55073ec7153d9f14b9b9fc5a0cffa3bc9188b21ed11ef6952/openimageio-3.0.16.0-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:442ac9f06c7592cdc01fc7c9ac827c3be5d2c885c2e71b55ee40e249554a878d", size = 6194989, upload-time = "2026-03-01T03:42:52.385Z" },
{ url = "https://files.pythonhosted.org/packages/e6/92/25397936a4f881cd5c9731e260009240544d6a959bfd16e14174983a4cea/openimageio-3.0.16.0-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:55ce42888c85262c3be66133be0bc4236a0f445a3568783e57a467324bcbe03c", size = 6392269, upload-time = "2026-03-01T03:42:54.679Z" },
{ url = "https://files.pythonhosted.org/packages/ca/32/590315956ec40443d6992a5e0ddb11be5615d6aaf38f6d9edf970ce28685/openimageio-3.0.16.0-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:a7f97ff4ed54d63c6241d1634d2f55004b5f83cf28cdf609302523e2eaec1a09", size = 6596973, upload-time = "2026-03-01T03:42:56.24Z" },
{ url = "https://files.pythonhosted.org/packages/07/b5/1be18ba1f2ae52598917da774820b4d0020a7ab50e1c54c33ab1820f06d0/openimageio-3.0.16.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:fa87d72410f59d969d64112c9dd2b51999f6553f392f757844bdc85060092de5", size = 6493859, upload-time = "2026-03-01T03:42:57.886Z" },
{ url = "https://files.pythonhosted.org/packages/0a/a7/a6aa117ed3fd28ed67e08a54f74c3ab10c45f1e2702dbd8f12c248b382ea/openimageio-3.0.16.0-cp311-cp311-win_amd64.whl", hash = "sha256:6af821d1371271e2072755702a7c5d073ec67cc27d1909f4e646c5a62a82879c", size = 7099505, upload-time = "2026-03-01T03:42:59.778Z" },
{ url = "https://files.pythonhosted.org/packages/9b/2a/a58a9ee65f9d60fe8be646bc1c0fdbbbaafe4c77ad36587523535e9e4c63/openimageio-3.0.16.0-cp312-cp312-macosx_10_15_x86_64.whl", hash = "sha256:c2cf2919ef287135b0cc2113dbdc85b64dc0e2fa4bd5c3887fe09c0fd1926924", size = 6643409, upload-time = "2026-03-01T03:43:01.061Z" },
{ url = "https://files.pythonhosted.org/packages/7a/c5/4aca1892f48a1bfed8b6434a02ad2762a927c0725cfd0196f3a8f46becbe/openimageio-3.0.16.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:732c7e76840b2df62975ccb91ab176e430a7dbf410acd993a30f6fd7a485e698", size = 6262459, upload-time = "2026-03-01T03:43:02.262Z" },
{ url = "https://files.pythonhosted.org/packages/b0/68/a9e1d45f347ae2ca08d5a6166af615871c323501a4a92c8825a4e604d32d/openimageio-3.0.16.0-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:7c825270d342cd071b369e69e36869271a7c811155744ebd4c9b8a6dbde24ac7", size = 6392045, upload-time = "2026-03-01T03:43:04.195Z" },
{ url = "https://files.pythonhosted.org/packages/40/a0/554e03a80c8650be3985b4aa427a80491a8cc462c3cde5890357b5976597/openimageio-3.0.16.0-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:096e0a4940395f73903a295ec423a338594ce07d405206e9d775ba7168861fd8", size = 6597227, upload-time = "2026-03-01T03:43:05.814Z" },
{ url = "https://files.pythonhosted.org/packages/cd/36/4792c357fa2b19d5ed4f5bc2b634dcb5f658c541ff33263e09002b680cd7/openimageio-3.0.16.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:331d103b35d1d402ff8c7a9ece77b7c2191f5d6e4b6d8299be08374c4c4b91f8", size = 6493718, upload-time = "2026-03-01T03:43:07.37Z" },
{ url = "https://files.pythonhosted.org/packages/94/9f/e3da32d705a704ed7f1324345dc0a0ac108667bb16a8ba50da7f4c298c0c/openimageio-3.0.16.0-cp312-cp312-win_amd64.whl", hash = "sha256:dddcd78cb35946907eb2be607ba9d620c05f502652937c9ea250078725b1745e", size = 7103136, upload-time = "2026-03-01T03:43:08.983Z" },
{ url = "https://files.pythonhosted.org/packages/67/61/0da72a9ed22127a563c12db20025255010661f7162f494741c74242d195d/openimageio-3.0.16.0-cp313-cp313-macosx_10_15_x86_64.whl", hash = "sha256:b160601825a8362eb586a760afabe1ebfd5d3b6869f012ccbe5cd637a621efd2", size = 6643458, upload-time = "2026-03-01T03:43:10.869Z" },
{ url = "https://files.pythonhosted.org/packages/65/fe/5a743e939c3ce8d822f3a826dc93a6a5914265d613bd81cfddc3c6f91c9b/openimageio-3.0.16.0-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:4259743111db8ea8daf76d44724fcf20f0e2ce641906d36202983ec6092c24d2", size = 6262560, upload-time = "2026-03-01T03:43:12.499Z" },
{ url = "https://files.pythonhosted.org/packages/de/ae/8c7069774dc198ffd7b193b49c8e3cf897a98bb8c5eb9de754be10957932/openimageio-3.0.16.0-cp313-cp313-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:5a2a0d97d7091e33b542ae0fb66d7e250ad48a4417dbffa308c48eefd96dc3c6", size = 6392031, upload-time = "2026-03-01T03:43:14.256Z" },
{ url = "https://files.pythonhosted.org/packages/29/9b/3c05dc022ae1ae9761e2a6d297b10cefaa9cbbcf2feb6a355710d1f4f045/openimageio-3.0.16.0-cp313-cp313-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:2df45d95e50dcdf4fc07df09b3cfdfb0b6dbfab1dad86a4b744b698058bce62e", size = 6596464, upload-time = "2026-03-01T03:43:16.204Z" },
{ url = "https://files.pythonhosted.org/packages/0d/82/04d1b969c56c4737783949ccd3446d63b4a0db0003684371b96877a54f53/openimageio-3.0.16.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b64ca87f11057bd60ee15f5280a6e83c719c01ed31fb6b220aa1df7e3966ad3b", size = 6493537, upload-time = "2026-03-01T03:43:17.486Z" },
{ url = "https://files.pythonhosted.org/packages/e6/97/0bbee7fd1d2e028c4be6292e5f968556a11670bec358d57ba0ff0b0d77de/openimageio-3.0.16.0-cp313-cp313-win_amd64.whl", hash = "sha256:c439c4eb80037aadd34436947eebf3fc06ac61512b0263724bafef6d0878fd76", size = 7103123, upload-time = "2026-03-01T03:43:19.06Z" },
]
[[package]]
name = "openimageio"
version = "3.1.11.0"
source = { registry = "https://pypi.org/simple" }
resolution-markers = [
"python_full_version >= '3.13' and sys_platform == 'linux'",
"python_full_version == '3.12.*' and sys_platform == 'linux'",
"python_full_version >= '3.13' and sys_platform != 'linux'",
"python_full_version == '3.12.*' and sys_platform != 'linux'",
"python_full_version == '3.11.*' and sys_platform == 'linux'",
"python_full_version == '3.11.*' and sys_platform != 'linux'",
"python_full_version < '3.11' and sys_platform == 'linux'",
"python_full_version < '3.11' and sys_platform != 'linux'",
]
dependencies = [
{ name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "(python_full_version < '3.11' and extra != 'extra-8-ltx-core-fp8-trtllm') or (extra == 'extra-8-ltx-core-fp8-trtllm' and extra == 'extra-8-ltx-core-xformers')" },
{ name = "numpy", version = "2.4.1", source = { registry = "https://pypi.org/simple" }, marker = "(python_full_version >= '3.11' and extra != 'extra-8-ltx-core-fp8-trtllm') or (extra == 'extra-8-ltx-core-fp8-trtllm' and extra == 'extra-8-ltx-core-xformers')" },
]
sdist = { url = "https://files.pythonhosted.org/packages/f6/bd/d49fe9f78b244251301da7b4caff7053b151df55ce0e04829004ac9f65ba/openimageio-3.1.11.0.tar.gz", hash = "sha256:7c096607569b0d01266da739bbf7496f1c0a8fb4449092f8f66fde2bf3a20494", size = 6555440, upload-time = "2026-03-01T03:09:00.936Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/b9/cc/eb7725c3ca6c76fd5defe92b2d558773b1a7a89a69c483bb33d791b0a723/openimageio-3.1.11.0-cp310-cp310-macosx_10_15_x86_64.whl", hash = "sha256:ef9377cdd84d67c415e6fc67f75553cacad36976cfd3f5ae1f0b2b5a6c46a7ef", size = 6851102, upload-time = "2026-03-01T03:08:18.322Z" },
{ url = "https://files.pythonhosted.org/packages/2a/54/ffc2d3b70dee7690dae808cfa8787cfa1f109fc7b4cd68632cade9fab5f3/openimageio-3.1.11.0-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:0c497762ec7993c9af51079edd723ca0a37faa085c6115b8e69f151460a60927", size = 6465995, upload-time = "2026-03-01T03:08:19.907Z" },
{ url = "https://files.pythonhosted.org/packages/7f/7c/e2f32ee394db041af8477d984aef2c04067713e916830fb113ade2dd3426/openimageio-3.1.11.0-cp310-cp310-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:1cfc57d9e790e76e76498b0e5c5a0bed83ae1be6104c635eca114265eefff34c", size = 6501598, upload-time = "2026-03-01T03:08:21.486Z" },
{ url = "https://files.pythonhosted.org/packages/e1/4d/fa9d641d424d7e16a44a502c440165687adad13ca005507dbf29d8b7febf/openimageio-3.1.11.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:0f01e9e115b1832bd27bb6a62e29bf435cb27f3f41b589ac7d2838bdda530f15", size = 6789743, upload-time = "2026-03-01T03:08:22.714Z" },
{ url = "https://files.pythonhosted.org/packages/34/0a/cbde5147af9db425a5c692147da45f849d1203304f195bcae71875b43afb/openimageio-3.1.11.0-cp310-cp310-win_amd64.whl", hash = "sha256:4b57b599294afac5f21885949014ca1469866c3b648375b834d19723c67dc52c", size = 7459665, upload-time = "2026-03-01T03:08:24.267Z" },
{ url = "https://files.pythonhosted.org/packages/0e/9d/6d67fb8358faeb7f6b195ea35483ffa58d9fd45806dbedea471c79ae8050/openimageio-3.1.11.0-cp311-cp311-macosx_10_15_x86_64.whl", hash = "sha256:d02021c7a4baf0418cde43ebddd0103886cc809c8e85c9c8cd4a1cbe69879c5a", size = 6852201, upload-time = "2026-03-01T03:08:25.864Z" },
{ url = "https://files.pythonhosted.org/packages/ec/c2/1412eb64ea7a7156ebe335a03118e8d2dc116e63e354d44fceb475b520cc/openimageio-3.1.11.0-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:daa4396f8094ed94c9b963711b2a680dc8f34c5361be5bc06804b0fc94cb4f4b", size = 6466903, upload-time = "2026-03-01T03:08:27.124Z" },
{ url = "https://files.pythonhosted.org/packages/75/20/16e4d4a310628d0960588f919a7fc85dd94be4eb37a043f5cacfb9e024fc/openimageio-3.1.11.0-cp311-cp311-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:0f457ca10c45d4e03dd4d3af13456962bee2ee6abee91bf40a8d21ffa7da1f74", size = 6502239, upload-time = "2026-03-01T03:08:28.637Z" },
{ url = "https://files.pythonhosted.org/packages/9c/6f/70f75c4fbc958e9f4f829aeb8d88558591cdcdb6dd36e97d8ee6f50c4011/openimageio-3.1.11.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:11655436841f11d268049f5494c6621dfd6370c104b4aa8c20f0724271850796", size = 6790238, upload-time = "2026-03-01T03:08:30.083Z" },
{ url = "https://files.pythonhosted.org/packages/3b/1e/f85f0393237dbbcec0559c357795462e43eb2fead40d2954a99ebf24c310/openimageio-3.1.11.0-cp311-cp311-win_amd64.whl", hash = "sha256:58eacdc68992c49257be43a30ac424fd6cf617b9de1e516debc14c0bd93af075", size = 7460293, upload-time = "2026-03-01T03:08:31.646Z" },
{ url = "https://files.pythonhosted.org/packages/cf/2c/24036c501b24f0d4136d9c4605a8861f5affc87df23fa78c279e0a9216c0/openimageio-3.1.11.0-cp312-cp312-macosx_10_15_x86_64.whl", hash = "sha256:53337dd2e11f7f0a07fbe664d21c0d578a0674695da1e5f327764a3f421aff83", size = 6942447, upload-time = "2026-03-01T03:08:33.179Z" },
{ url = "https://files.pythonhosted.org/packages/69/ca/d84f062cbceb35a8006a9a3393da1893e23ffd8147c53d5f461fa99898cc/openimageio-3.1.11.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:8d50f1d8159177561c6d5efcc0f2d7f9fd5393cc147a39b546849da8762bb4f8", size = 6540326, upload-time = "2026-03-01T03:08:34.542Z" },
{ url = "https://files.pythonhosted.org/packages/6c/f8/ca6f4ca3d5dc2b2cc1411a1083b926d5fc9eb22e57d352333027b3eb4659/openimageio-3.1.11.0-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:212dd58364f8e11c0fbd90f82bbaf617afc68b4f3910538d0ac2324ef60a4c28", size = 6503147, upload-time = "2026-03-01T03:08:35.87Z" },
{ url = "https://files.pythonhosted.org/packages/81/9c/5ea545f62790bc4f8b9c04d433963f9fe5f08548c8663e89a54f1db7002d/openimageio-3.1.11.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f845304bb17a1fd48f19c972572eb03ee3f96336f2d5599d043cb7ac1c348035", size = 6791359, upload-time = "2026-03-01T03:08:37.412Z" },
{ url = "https://files.pythonhosted.org/packages/9b/ba/c64923d962dd6c2dd0421feaa6000b081bcf1e26e07a125da94488f9cd39/openimageio-3.1.11.0-cp312-cp312-win_amd64.whl", hash = "sha256:30284fb3367e480eaa345b122ccf38df61bd3885c78b237942a585cefaf0538a", size = 7466181, upload-time = "2026-03-01T03:08:38.931Z" },
{ url = "https://files.pythonhosted.org/packages/00/81/76375647018dce289e13f8a75c42cddbd9c688baa50d1dc99f73457508f0/openimageio-3.1.11.0-cp313-cp313-macosx_10_15_x86_64.whl", hash = "sha256:35ac810ccf3aa8a4cdc1fb19289dab57e20a9d915a78041a6fa2052d3b7a8c3c", size = 6942555, upload-time = "2026-03-01T03:08:40.539Z" },
{ url = "https://files.pythonhosted.org/packages/6b/a1/b1a861bc03462bc5ac6e6a9eac3e4f9e28302f0cc14748e0cd3f1f1efadd/openimageio-3.1.11.0-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:5d294f1d631baea51cf9d940ee770f789946dcfae187cf78a5b440a3cc88c4c7", size = 6540381, upload-time = "2026-03-01T03:08:41.796Z" },
{ url = "https://files.pythonhosted.org/packages/c2/5a/79ceb6c7cf57fdce947e9b341c5de1fbf6e219a0db5ce990e69431c54f90/openimageio-3.1.11.0-cp313-cp313-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:78cffd41f1d0e575c56aa5e601a0f3c98347cb9a20d3ca6210b4ef881cef256c", size = 6503197, upload-time = "2026-03-01T03:08:43.23Z" },
{ url = "https://files.pythonhosted.org/packages/e1/49/6593704b762dfe1296229323219e216fda029c6153622a57f1e865f82160/openimageio-3.1.11.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:020e7b5a0d2898dca1c66c7c6f800aa2c307a509ffd3dbed2898562445b5f6bd", size = 6791511, upload-time = "2026-03-01T03:08:44.544Z" },
{ url = "https://files.pythonhosted.org/packages/60/c2/43c1eeb5dcc7620a212c1705dc03547d20c917f3761fd6a2bde3fa3b9fb1/openimageio-3.1.11.0-cp313-cp313-win_amd64.whl", hash = "sha256:cb331042a297bd3de87b4eed5efed4c337e1446aaed3807ac4899c129e33d61a", size = 7466220, upload-time = "2026-03-01T03:08:45.78Z" },
{ url = "https://files.pythonhosted.org/packages/2e/b9/6b09045f1a92cd38a59fd3c03dd59374f27c03270943ebd5844247633c70/openimageio-3.1.11.0-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:f7c9b14c3981cd992b33f8a47735ce099687cddc9d05ffdd03b87490c0bcdd47", size = 6940155, upload-time = "2026-03-01T03:08:47.09Z" },
{ url = "https://files.pythonhosted.org/packages/13/d8/f10343f5ddf6bc9003f8e9c1eae3b2d4b5cfed01fccde6fffe5a7cafce33/openimageio-3.1.11.0-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:4f5e106054c01c5a768c6442b712493bfe21526dd4e202e3c18245a0224e6d53", size = 6536784, upload-time = "2026-03-01T03:08:48.364Z" },
{ url = "https://files.pythonhosted.org/packages/4b/6f/4709da066a4dafef14b5bb0e7e221fa9d65e196806ce22721862627ab37c/openimageio-3.1.11.0-cp314-cp314-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:4c2affc140802fdc42b2f01aeebedfdd64ff29891fdc2e2c708d4653852eff2b", size = 6503379, upload-time = "2026-03-01T03:08:49.947Z" },
{ url = "https://files.pythonhosted.org/packages/0d/5b/de755dfb3a8837592263511fc14966901d1b8d880eeb3cde9d65b8f02e81/openimageio-3.1.11.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:ee25034ceff73b8f08ae4b562a5a62cfcdaa0e873ce913c5df72d9dc2a0d56c7", size = 6791460, upload-time = "2026-03-01T03:08:51.483Z" },
{ url = "https://files.pythonhosted.org/packages/20/fb/159531b67af2ddca76dc1bbfcc31e4b94a8d91215ae0fa168e0243327b3b/openimageio-3.1.11.0-cp314-cp314-win_amd64.whl", hash = "sha256:1945317959ac1905e97c8d9479efc716bb765dec82a647365d37e046e154a46c", size = 7767018, upload-time = "2026-03-01T03:08:52.707Z" },
]
[[package]]
name = "openmpi"
version = "5.0.9"
@@ -7005,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" },
{ 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" },
]
[[package]]