From b604d3fab3d3acb3ba35a559d81f3ce2a5206d50 Mon Sep 17 00:00:00 2001 From: "github-actions[bot]" <41898282+github-actions[bot]@users.noreply.github.com> Date: Thu, 23 Apr 2026 12:43:54 +0000 Subject: [PATCH] Automated PR - 2026-04-23 --- .gitattributes | 1 + .gitignore | 5 + README.md | 2 + packages/ltx-core/pyproject.toml | 2 +- packages/ltx-core/src/ltx_core/batch_split.py | 4 +- .../src/ltx_core/block_streaming/__init__.py | 19 + .../src/ltx_core/block_streaming/builder.py | 305 ++++++ .../src/ltx_core/block_streaming/disk.py | 106 +++ .../src/ltx_core/block_streaming/pool.py | 64 ++ .../src/ltx_core/block_streaming/provider.py | 110 +++ .../src/ltx_core/block_streaming/source.py | 83 ++ .../src/ltx_core/block_streaming/utils.py | 60 ++ .../src/ltx_core/block_streaming/wrapper.py | 96 ++ .../conditioning/types/keyframe_cond.py | 15 +- packages/ltx-core/src/ltx_core/hdr.py | 71 ++ .../ltx-core/src/ltx_core/layer_streaming.py | 306 ------ .../ltx-core/src/ltx_core/loader/__init__.py | 8 + .../ltx-core/src/ltx_core/loader/helpers.py | 61 ++ .../loader/single_gpu_model_builder.py | 110 ++- .../src/ltx_core/model/video_vae/video_vae.py | 30 +- packages/ltx-pipelines/CLAUDE.md | 2 +- packages/ltx-pipelines/README.md | 19 + packages/ltx-pipelines/pyproject.toml | 4 +- .../src/ltx_pipelines/a2vid_two_stage.py | 15 +- .../src/ltx_pipelines/distilled.py | 17 +- .../src/ltx_pipelines/hdr_ic_lora.py | 886 ++++++++++++++++++ .../src/ltx_pipelines/ic_lora.py | 18 +- .../ltx_pipelines/keyframe_interpolation.py | 21 +- .../ltx-pipelines/src/ltx_pipelines/retake.py | 10 +- .../src/ltx_pipelines/ti2vid_one_stage.py | 16 +- .../src/ltx_pipelines/ti2vid_two_stages.py | 21 +- .../src/ltx_pipelines/ti2vid_two_stages_hq.py | 15 +- .../src/ltx_pipelines/utils/args.py | 20 +- .../src/ltx_pipelines/utils/blocks.py | 211 +++-- .../src/ltx_pipelines/utils/helpers.py | 5 + .../src/ltx_pipelines/utils/media_io.py | 205 ++++ .../src/ltx_pipelines/utils/types.py | 21 + .../ltx-trainer/configs/ltx2_av_lora.yaml | 3 - .../configs/ltx2_av_lora_low_vram.yaml | 3 - .../ltx-trainer/configs/ltx2_v2v_ic_lora.yaml | 3 - .../docs/configuration-reference.md | 1 - packages/ltx-trainer/pyproject.toml | 4 +- .../ltx-trainer/scripts/compute_reference.py | 4 +- .../ltx-trainer/src/ltx_trainer/config.py | 6 - .../ltx-trainer/src/ltx_trainer/datasets.py | 87 +- .../src/ltx_trainer/hf_hub_utils.py | 10 +- .../training_strategies/base_strategy.py | 11 +- .../src/ltx_trainer/video_utils.py | 28 +- uv.lock | 108 ++- 49 files changed, 2664 insertions(+), 568 deletions(-) create mode 100644 packages/ltx-core/src/ltx_core/block_streaming/__init__.py create mode 100644 packages/ltx-core/src/ltx_core/block_streaming/builder.py create mode 100644 packages/ltx-core/src/ltx_core/block_streaming/disk.py create mode 100644 packages/ltx-core/src/ltx_core/block_streaming/pool.py create mode 100644 packages/ltx-core/src/ltx_core/block_streaming/provider.py create mode 100644 packages/ltx-core/src/ltx_core/block_streaming/source.py create mode 100644 packages/ltx-core/src/ltx_core/block_streaming/utils.py create mode 100644 packages/ltx-core/src/ltx_core/block_streaming/wrapper.py create mode 100644 packages/ltx-core/src/ltx_core/hdr.py delete mode 100644 packages/ltx-core/src/ltx_core/layer_streaming.py create mode 100644 packages/ltx-core/src/ltx_core/loader/helpers.py create mode 100644 packages/ltx-pipelines/src/ltx_pipelines/hdr_ic_lora.py diff --git a/.gitattributes b/.gitattributes index dcae1fc..cb8c63e 100644 --- a/.gitattributes +++ b/.gitattributes @@ -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 diff --git a/.gitignore b/.gitignore index 91dcb7c..82d7f42 100644 --- a/.gitignore +++ b/.gitignore @@ -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 diff --git a/README.md b/README.md index 86a28bb..0473a30 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/packages/ltx-core/pyproject.toml b/packages/ltx-core/pyproject.toml index ced26e9..c0779ae 100644 --- a/packages/ltx-core/pyproject.toml +++ b/packages/ltx-core/pyproject.toml @@ -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" diff --git a/packages/ltx-core/src/ltx_core/batch_split.py b/packages/ltx-core/src/ltx_core/batch_split.py index 014ca5a..3a77422 100644 --- a/packages/ltx-core/src/ltx_core/batch_split.py +++ b/packages/ltx-core/src/ltx_core/batch_split.py @@ -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. """ diff --git a/packages/ltx-core/src/ltx_core/block_streaming/__init__.py b/packages/ltx-core/src/ltx_core/block_streaming/__init__.py new file mode 100644 index 0000000..88a181a --- /dev/null +++ b/packages/ltx-core/src/ltx_core/block_streaming/__init__.py @@ -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", +] diff --git a/packages/ltx-core/src/ltx_core/block_streaming/builder.py b/packages/ltx-core/src/ltx_core/block_streaming/builder.py new file mode 100644 index 0000000..2492906 --- /dev/null +++ b/packages/ltx-core/src/ltx_core/block_streaming/builder.py @@ -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) diff --git a/packages/ltx-core/src/ltx_core/block_streaming/disk.py b/packages/ltx-core/src/ltx_core/block_streaming/disk.py new file mode 100644 index 0000000..90d7bc6 --- /dev/null +++ b/packages/ltx-core/src/ltx_core/block_streaming/disk.py @@ -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() diff --git a/packages/ltx-core/src/ltx_core/block_streaming/pool.py b/packages/ltx-core/src/ltx_core/block_streaming/pool.py new file mode 100644 index 0000000..4a9a163 --- /dev/null +++ b/packages/ltx-core/src/ltx_core/block_streaming/pool.py @@ -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) diff --git a/packages/ltx-core/src/ltx_core/block_streaming/provider.py b/packages/ltx-core/src/ltx_core/block_streaming/provider.py new file mode 100644 index 0000000..59646cb --- /dev/null +++ b/packages/ltx-core/src/ltx_core/block_streaming/provider.py @@ -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)) diff --git a/packages/ltx-core/src/ltx_core/block_streaming/source.py b/packages/ltx-core/src/ltx_core/block_streaming/source.py new file mode 100644 index 0000000..7cfe842 --- /dev/null +++ b/packages/ltx-core/src/ltx_core/block_streaming/source.py @@ -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) diff --git a/packages/ltx-core/src/ltx_core/block_streaming/utils.py b/packages/ltx-core/src/ltx_core/block_streaming/utils.py new file mode 100644 index 0000000..acb4a2a --- /dev/null +++ b/packages/ltx-core/src/ltx_core/block_streaming/utils.py @@ -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() + } diff --git a/packages/ltx-core/src/ltx_core/block_streaming/wrapper.py b/packages/ltx-core/src/ltx_core/block_streaming/wrapper.py new file mode 100644 index 0000000..c7fed14 --- /dev/null +++ b/packages/ltx-core/src/ltx_core/block_streaming/wrapper.py @@ -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) diff --git a/packages/ltx-core/src/ltx_core/conditioning/types/keyframe_cond.py b/packages/ltx-core/src/ltx_core/conditioning/types/keyframe_cond.py index d4af3f6..df75fd8 100644 --- a/packages/ltx-core/src/ltx_core/conditioning/types/keyframe_cond.py +++ b/packages/ltx-core/src/ltx_core/conditioning/types/keyframe_cond.py @@ -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 diff --git a/packages/ltx-core/src/ltx_core/hdr.py b/packages/ltx-core/src/ltx_core/hdr.py new file mode 100644 index 0000000..8413004 --- /dev/null +++ b/packages/ltx-core/src/ltx_core/hdr.py @@ -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}") diff --git a/packages/ltx-core/src/ltx_core/layer_streaming.py b/packages/ltx-core/src/ltx_core/layer_streaming.py deleted file mode 100644 index 92c5596..0000000 --- a/packages/ltx-core/src/ltx_core/layer_streaming.py +++ /dev/null @@ -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) diff --git a/packages/ltx-core/src/ltx_core/loader/__init__.py b/packages/ltx-core/src/ltx_core/loader/__init__.py index 3aaa01e..9bb099c 100644 --- a/packages/ltx-core/src/ltx_core/loader/__init__.py +++ b/packages/ltx-core/src/ltx_core/loader/__init__.py @@ -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", ] diff --git a/packages/ltx-core/src/ltx_core/loader/helpers.py b/packages/ltx-core/src/ltx_core/loader/helpers.py new file mode 100644 index 0000000..b95e0bd --- /dev/null +++ b/packages/ltx-core/src/ltx_core/loader/helpers.py @@ -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 diff --git a/packages/ltx-core/src/ltx_core/loader/single_gpu_model_builder.py b/packages/ltx-core/src/ltx_core/loader/single_gpu_model_builder.py index d439d3d..f98fee1 100644 --- a/packages/ltx-core/src/ltx_core/loader/single_gpu_model_builder.py +++ b/packages/ltx-core/src/ltx_core/loader/single_gpu_model_builder.py @@ -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) diff --git a/packages/ltx-core/src/ltx_core/model/video_vae/video_vae.py b/packages/ltx-core/src/ltx_core/model/video_vae/video_vae.py index 7a10409..983178c 100644 --- a/packages/ltx-core/src/ltx_core/model/video_vae/video_vae.py +++ b/packages/ltx-core/src/ltx_core/model/video_vae/video_vae.py @@ -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.""" diff --git a/packages/ltx-pipelines/CLAUDE.md b/packages/ltx-pipelines/CLAUDE.md index 995d04b..631fa01 100644 --- a/packages/ltx-pipelines/CLAUDE.md +++ b/packages/ltx-pipelines/CLAUDE.md @@ -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`) diff --git a/packages/ltx-pipelines/README.md b/packages/ltx-pipelines/README.md index 3f5a60e..14c466c 100644 --- a/packages/ltx-pipelines/README.md +++ b/packages/ltx-pipelines/README.md @@ -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. diff --git a/packages/ltx-pipelines/pyproject.toml b/packages/ltx-pipelines/pyproject.toml index 0643158..b455839 100644 --- a/packages/ltx-pipelines/pyproject.toml +++ b/packages/ltx-pipelines/pyproject.toml @@ -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"] diff --git a/packages/ltx-pipelines/src/ltx_pipelines/a2vid_two_stage.py b/packages/ltx-pipelines/src/ltx_pipelines/a2vid_two_stage.py index 0e95809..2ed84d4 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/a2vid_two_stage.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/a2vid_two_stage.py @@ -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, ) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/distilled.py b/packages/ltx-pipelines/src/ltx_pipelines/distilled.py index f1f0b64..e3cde25 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/distilled.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/distilled.py @@ -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( diff --git a/packages/ltx-pipelines/src/ltx_pipelines/hdr_ic_lora.py b/packages/ltx-pipelines/src/ltx_pipelines/hdr_ic_lora.py new file mode 100644 index 0000000..fa1637e --- /dev/null +++ b/packages/ltx-pipelines/src/ltx_pipelines/hdr_ic_lora.py @@ -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() diff --git a/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py b/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py index 016ea0e..15fe01b 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py @@ -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( diff --git a/packages/ltx-pipelines/src/ltx_pipelines/keyframe_interpolation.py b/packages/ltx-pipelines/src/ltx_pipelines/keyframe_interpolation.py index d937f80..9c27475 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/keyframe_interpolation.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/keyframe_interpolation.py @@ -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, ) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/retake.py b/packages/ltx-pipelines/src/ltx_pipelines/retake.py index 0aa58d0..ea8ee2a 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/retake.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/retake.py @@ -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) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_one_stage.py b/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_one_stage.py index 2cd9ec7..3b5389e 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_one_stage.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_one_stage.py @@ -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, ) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages.py b/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages.py index 280ab96..0f73d8e 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages.py @@ -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, ) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages_hq.py b/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages_hq.py index 502752f..e6d1c12 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages_hq.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages_hq.py @@ -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, ) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/args.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/args.py index 6dcf871..bb3236d 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/args.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/args.py @@ -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" ), ) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/blocks.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/blocks.py index 17cfb7b..abece4e 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/blocks.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/blocks.py @@ -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) # --------------------------------------------------------------------------- diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/helpers.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/helpers.py index 9f53421..9960871 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/helpers.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/helpers.py @@ -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: diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/media_io.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/media_io.py index b05dcbd..e96c65c 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/media_io.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/media_io.py @@ -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() diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/types.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/types.py index df60451..6093ead 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/types.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/types.py @@ -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" diff --git a/packages/ltx-trainer/configs/ltx2_av_lora.yaml b/packages/ltx-trainer/configs/ltx2_av_lora.yaml index 4223618..ded5347 100644 --- a/packages/ltx-trainer/configs/ltx2_av_lora.yaml +++ b/packages/ltx-trainer/configs/ltx2_av_lora.yaml @@ -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 diff --git a/packages/ltx-trainer/configs/ltx2_av_lora_low_vram.yaml b/packages/ltx-trainer/configs/ltx2_av_lora_low_vram.yaml index 811edf3..065bdb7 100644 --- a/packages/ltx-trainer/configs/ltx2_av_lora_low_vram.yaml +++ b/packages/ltx-trainer/configs/ltx2_av_lora_low_vram.yaml @@ -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 diff --git a/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml b/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml index 6190948..dbaaf0e 100644 --- a/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml +++ b/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml @@ -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 diff --git a/packages/ltx-trainer/docs/configuration-reference.md b/packages/ltx-trainer/docs/configuration-reference.md index df5c023..81ce562 100644 --- a/packages/ltx-trainer/docs/configuration-reference.md +++ b/packages/ltx-trainer/docs/configuration-reference.md @@ -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 diff --git a/packages/ltx-trainer/pyproject.toml b/packages/ltx-trainer/pyproject.toml index 81c51ff..801d73e 100644 --- a/packages/ltx-trainer/pyproject.toml +++ b/packages/ltx-trainer/pyproject.toml @@ -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] diff --git a/packages/ltx-trainer/scripts/compute_reference.py b/packages/ltx-trainer/scripts/compute_reference.py index a3ad007..4af848a 100644 --- a/packages/ltx-trainer/scripts/compute_reference.py +++ b/packages/ltx-trainer/scripts/compute_reference.py @@ -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. " diff --git a/packages/ltx-trainer/src/ltx_trainer/config.py b/packages/ltx-trainer/src/ltx_trainer/config.py index 4446c39..751f4fd 100644 --- a/packages/ltx-trainer/src/ltx_trainer/config.py +++ b/packages/ltx-trainer/src/ltx_trainer/config.py @@ -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", diff --git a/packages/ltx-trainer/src/ltx_trainer/datasets.py b/packages/ltx-trainer/src/ltx_trainer/datasets.py index 5873775..2d8f59c 100644 --- a/packages/ltx-trainer/src/ltx_trainer/datasets.py +++ b/packages/ltx-trainer/src/ltx_trainer/datasets.py @@ -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}") diff --git a/packages/ltx-trainer/src/ltx_trainer/hf_hub_utils.py b/packages/ltx-trainer/src/ltx_trainer/hf_hub_utils.py index 63ed56b..32f9a99 100644 --- a/packages/ltx-trainer/src/ltx_trainer/hf_hub_utils.py +++ b/packages/ltx-trainer/src/ltx_trainer/hf_hub_utils.py @@ -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.""" diff --git a/packages/ltx-trainer/src/ltx_trainer/training_strategies/base_strategy.py b/packages/ltx-trainer/src/ltx_trainer/training_strategies/base_strategy.py index f219d5a..c0ad765 100644 --- a/packages/ltx-trainer/src/ltx_trainer/training_strategies/base_strategy.py +++ b/packages/ltx-trainer/src/ltx_trainer/training_strategies/base_strategy.py @@ -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 diff --git a/packages/ltx-trainer/src/ltx_trainer/video_utils.py b/packages/ltx-trainer/src/ltx_trainer/video_utils.py index 3d36fc4..6e5df1b 100644 --- a/packages/ltx-trainer/src/ltx_trainer/video_utils.py +++ b/packages/ltx-trainer/src/ltx_trainer/video_utils.py @@ -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: diff --git a/uv.lock b/uv.lock index 0d106fb..3bef8ca 100644 --- a/uv.lock +++ b/uv.lock @@ -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]]