diff --git a/README.md b/README.md index 0473a30..590edc5 100644 --- a/README.md +++ b/README.md @@ -39,7 +39,7 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https:// **Temporal Upscaler** - Supported by the model and will be required for future pipeline implementations * [`ltx-2.3-temporal-upscaler-x2-1.0.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-temporal-upscaler-x2-1.0.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-temporal-upscaler-x2-1.0.safetensors) -**Distilled LoRA** - Required for current two-stage pipeline implementations in this repository (except DistilledPipeline and ICLoraPipeline) +**Distilled LoRA** - Required for current two-stage pipeline implementations in this repository (except DistilledPipeline, ICLoraPipeline, and LipDubPipeline) * [`ltx-2.3-22b-distilled-lora-384-1.1.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-22b-distilled-lora-384-1.1.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-22b-distilled-lora-384-1.1.safetensors) **Gemma Text Encoder** (download all assets from the repository) @@ -58,6 +58,7 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https:// * [`LTX-2-19b-LoRA-Camera-Control-Jib-Up`](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Jib-Up) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Jib-Up/resolve/main/ltx-2-19b-lora-camera-control-jib-up.safetensors) * [`LTX-2-19b-LoRA-Camera-Control-Static`](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Static) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Static/resolve/main/ltx-2-19b-lora-camera-control-static.safetensors) * [`LTX-2.3-22b-IC-LoRA-HDR`](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-HDR) - HDR IC-LoRA and pre-computed text embeddings for `HDRICLoraPipeline` + * [`LTX-2.3-22b-IC-LoRA-LipDub`](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-LipDub) - [Download](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-LipDub/resolve/main/ltx-2.3-22b-ic-lora-lipdub-0.9.safetensors) ### Available Pipelines @@ -70,6 +71,7 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https:// * **[A2VidPipelineTwoStage](packages/ltx-pipelines/src/ltx_pipelines/a2vid_two_stage.py)** - Audio-to-video generation conditioned on an input audio file * **[RetakePipeline](packages/ltx-pipelines/src/ltx_pipelines/retake.py)** - Regenerate a specific time region of an existing video * **[HDRICLoraPipeline](packages/ltx-pipelines/src/ltx_pipelines/hdr_ic_lora.py)** - Video-to-video with HDR output (linear float frames via LogC3 inverse decode, suitable for EXR export and tonemapping) +* **[LipDubPipeline](packages/ltx-pipelines/src/ltx_pipelines/lipdub.py)** - Lip dubbing, rephrasing, matching speaker identity (distilled model, single IC-LoRA, Two stages). ### ⚡ Optimization Tips diff --git a/packages/ltx-core/README.md b/packages/ltx-core/README.md index 2add024..dd2957c 100644 --- a/packages/ltx-core/README.md +++ b/packages/ltx-core/README.md @@ -77,13 +77,17 @@ model = builder.build(device=torch.device("cuda")) Use the `.lora()` method to attach one or more LoRA adapters before calling `.build()`: ```python +from ltx_core.loader import SDOps + +lora_sd_ops = SDOps(name="identity").with_matching() # or a model-specific key-renaming SDOps + builder = ( SingleGPUModelBuilder( model_class_configurator=MyModelConfigurator, model_path="/path/to/model.safetensors", ) - .lora("/path/to/lora_a.safetensors", strength=0.8) - .lora("/path/to/lora_b.safetensors", strength=0.5) + .lora("/path/to/lora_a.safetensors", 0.8, lora_sd_ops) + .lora("/path/to/lora_b.safetensors", 0.5, lora_sd_ops) ) model = builder.build(device=torch.device("cuda")) ``` @@ -103,7 +107,7 @@ builder = SingleGPUModelBuilder( model_class_configurator=MyModelConfigurator, model_path="/path/to/model.safetensors", lora_load_device=torch.device("cuda"), -).lora("/path/to/lora.safetensors", strength=1.0) +).lora("/path/to/lora.safetensors", 1.0, lora_sd_ops) model = builder.build(device=torch.device("cuda")) ``` diff --git a/packages/ltx-core/pyproject.toml b/packages/ltx-core/pyproject.toml index c0779ae..e47cd53 100644 --- a/packages/ltx-core/pyproject.toml +++ b/packages/ltx-core/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "ltx-core" -version = "1.1.2" +version = "1.1.3" description = "Core implementation of Lightricks' LTX-2 model" readme = "README.md" requires-python = ">=3.10" diff --git a/packages/ltx-core/src/ltx_core/block_streaming/builder.py b/packages/ltx-core/src/ltx_core/block_streaming/builder.py index 2492906..1eec905 100644 --- a/packages/ltx-core/src/ltx_core/block_streaming/builder.py +++ b/packages/ltx-core/src/ltx_core/block_streaming/builder.py @@ -7,16 +7,17 @@ from collections.abc import Callable from dataclasses import dataclass, field, replace from typing import Generic +import safetensors import torch from torch import nn from ltx_core.block_streaming.disk import DiskBlockReader, DiskTensorReader, LoraSource -from ltx_core.block_streaming.pool import BlockLayout, WeightPool +from ltx_core.block_streaming.pool import WeightPool from ltx_core.block_streaming.provider import WeightsProvider from ltx_core.block_streaming.source import DiskWeightSource, PinnedWeightSource, WeightSource -from ltx_core.block_streaming.utils import build_pool_layout, resolve_attr +from ltx_core.block_streaming.utils import allocate_layout_views, derive_layout, make_block_key, resolve_attr from ltx_core.block_streaming.wrapper import BlockStreamingWrapper -from ltx_core.loader.fuse_loras import apply_loras +from ltx_core.loader.fuse_loras import aggregate_lora_products, fuse_lora_weights from ltx_core.loader.helpers import create_meta_model, load_state_dict, read_model_config from ltx_core.loader.module_ops import ModuleOps from ltx_core.loader.primitives import ( @@ -53,8 +54,8 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType]) ``"velocity_model.transformer_blocks"``). blocks_prefix: State-dict key prefix for block weights (e.g. ``"transformer_blocks"``). - state_dict_prefix: Key prefix for non-block weights - (e.g. ``"velocity_model."``). + state_dict_prefix: Wrapper offset prepended to keys when loading into + the meta model (e.g. ``"velocity_model."`` when wrapped by ``X0Model``). model_wrapper: Optional callable wrapping the model (e.g. ``X0Model``). """ @@ -110,7 +111,6 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType]) if not self.blocks_prefix: raise ValueError("blocks_prefix must be non-empty for streaming") - # 1. Create meta model (no weights allocated). config = read_model_config(self.model_path, self.model_loader) meta_model: nn.Module = create_meta_model(self.model_class_configurator, config, self.module_ops) if self.model_wrapper is not None: @@ -118,22 +118,29 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType]) meta_model.eval() blocks = resolve_attr(meta_model, self.blocks_attr) - layout = build_pool_layout(blocks[0], dtype) - # 2. Determine slot counts. + checkpoint_paths = list(self.model_path) if isinstance(self.model_path, tuple) else [self.model_path] + block_key_map, non_block_keys = _scan_checkpoint_keys(checkpoint_paths, self.model_sd_ops, self.blocks_prefix) + cpu_slots_count = cpu_slots_count if cpu_slots_count is not None else len(blocks) gpu_slots_count = gpu_slots_count if gpu_slots_count is not None else _DEFAULT_GPU_SLOTS - # 3. Build source and load non-block weights. if cpu_slots_count >= len(blocks): - source, lora_sources = self._build_pinned_source(meta_model, target_device, dtype, cpu_slots_count) + source, lora_sources = self._build_pinned_source( + meta_model, target_device, dtype, cpu_slots_count, block_key_map, non_block_keys + ) else: - source, lora_sources = self._build_disk_source(meta_model, layout, target_device, dtype, cpu_slots_count) + reader = DiskTensorReader(checkpoint_paths) + source, lora_sources = self._build_disk_source( + meta_model, target_device, dtype, cpu_slots_count, reader, block_key_map, non_block_keys + ) - # 4. Create provider and wrapper. copy_stream = torch.cuda.Stream(device=target_device) gpu_pool = WeightPool( - layout, gpu_slots_count, target_device, reuse_barrier=lambda event: copy_stream.wait_event(event) + source.block_layout, + gpu_slots_count, + target_device, + reuse_barrier=lambda event: copy_stream.wait_event(event), ) provider = WeightsProvider(gpu_pool, copy_stream, target_device, source, lora_sources, self.blocks_prefix) return BlockStreamingWrapper( @@ -149,89 +156,90 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType]) target_device: torch.device, dtype: torch.dtype, cpu_slots_count: int, + block_key_map: dict[int, list[tuple[str, str]]], + non_block_keys: list[tuple[str, str]], ) -> tuple[WeightSource, list[LoraSource]]: """Pre-load all blocks into pinned CPU buffers with LoRA fusion.""" model_sd = load_state_dict( self.model_path, self.model_loader, self.registry, torch.device("cpu"), self.model_sd_ops ) - if self.loras: - lora_sds = [ - load_state_dict([lora.path], self.model_loader, self.registry, torch.device("cpu"), lora.sd_ops) - for lora in self.loras - ] - lora_sd_and_strengths = [ - LoraStateDictWithStrength(sd, lora.strength) for sd, lora in zip(lora_sds, self.loras, strict=True) - ] - model_sd = apply_loras( - model_sd=model_sd, - lora_sd_and_strengths=lora_sd_and_strengths, - dtype=dtype, - destination_sd=model_sd if isinstance(self.registry, DummyRegistry) else None, + lora_sd_and_strengths = [ + LoraStateDictWithStrength( + load_state_dict([lora.path], self.model_loader, self.registry, torch.device("cpu"), lora.sd_ops), + lora.strength, ) + for lora in self.loras + ] - # Partition: non-block weights go to GPU, block weights go directly - # to pinned buffers. This avoids holding the full state dict and - # pinned copies simultaneously. - non_block_sd: dict[str, torch.Tensor] = {} - block_tensors: dict[int, dict[str, torch.Tensor]] = {} - prefix_dot = self.blocks_prefix + "." + for block_idx in block_key_map: + if block_idx >= cpu_slots_count: + raise ValueError( + f"Pinned source requires one CPU slot per block; " + f"got block index {block_idx} with only {cpu_slots_count} slots." + ) - for key, tensor in model_sd.sd.items(): - if key.startswith(prefix_dot): - rest = key[len(prefix_dot) :] - idx_str, _, param_name = rest.partition(".") - try: - block_idx = int(idx_str) - except ValueError: - non_block_sd[self.state_dict_prefix + key] = tensor.to(device=target_device, dtype=dtype) - continue - block_tensors.setdefault(block_idx, {})[param_name] = tensor + blocks = resolve_attr(meta_model, self.blocks_attr) + block_tensors: dict[str, torch.Tensor] = {} + for block_idx, entries in block_key_map.items(): + block_params = dict(blocks[block_idx].named_parameters()) + for _sft_key, param_name in entries: + key = make_block_key(self.blocks_prefix, block_idx, param_name) + block_tensors[key] = block_params[param_name] + blocks_layout = derive_layout(block_tensors, dtype) + pinned_blocks = allocate_layout_views(blocks_layout, pin_memory=True) + + should_sync = False + for key, fused in fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype=None, preserve_input_device=False): + if key in pinned_blocks: + pinned_blocks[key].copy_(fused, non_blocking=True) + model_sd.sd[key] = None + should_sync = True else: - non_block_sd[self.state_dict_prefix + key] = tensor.to(device=target_device, dtype=dtype) + model_sd.sd[key] = fused + if should_sync: + torch.cuda.synchronize() + + # Fill remaining pinned keys from the source state dict. + for key in blocks_layout: + if model_sd.sd[key] is None: + continue + pinned_blocks[key].copy_(model_sd.sd[key]) + model_sd.sd[key] = None + + pinned: dict[int, dict[str, torch.Tensor]] = { + block_idx: { + param_name: pinned_blocks[make_block_key(self.blocks_prefix, block_idx, param_name)] + for _sft_key, param_name in entries + } + for block_idx, entries in block_key_map.items() + } + + non_block_sd: dict[str, torch.Tensor] = { + self.state_dict_prefix + model_key: model_sd.sd[model_key].to(device=target_device, dtype=dtype) + for _sft_key, model_key in non_block_keys + } meta_model.load_state_dict(non_block_sd, strict=False, assign=True) - del model_sd, non_block_sd - - # Pin block weights one block at a time, freeing the source tensors as we go. - pinned: dict[int, dict[str, torch.Tensor]] = {} - for idx in range(cpu_slots_count): - src = block_tensors.pop(idx) - pinned[idx] = {name: tensor.to(dtype=dtype).pin_memory() for name, tensor in src.items()} return PinnedWeightSource(pinned), [] def _build_disk_source( self, meta_model: nn.Module, - layout: BlockLayout, target_device: torch.device, dtype: torch.dtype, cpu_slots_count: int, + reader: DiskTensorReader, + block_key_map: dict[int, list[tuple[str, str]]], + non_block_keys: list[tuple[str, str]], ) -> tuple[WeightSource, list[LoraSource]]: - """Create a DiskWeightSource backed by a DiskBlockReader for lazy loading.""" + """Create a DiskWeightSource backed by a DiskBlockReader for lazy loading. + Derives the shared pool layout from the meta model's block 0 — this + relies on module_ops (e.g. fp8_cast) leaving the meta param dtype in + sync with the post-sd_ops checkpoint dtype. + """ lora_sources = [LoraSource(lora.path, lora.sd_ops, lora.strength) for lora in self.loras] - checkpoint_paths = list(self.model_path) if isinstance(self.model_path, tuple) else [self.model_path] - reader = DiskTensorReader(checkpoint_paths) - - block_key_map: dict[int, list[tuple[str, str]]] = {} - non_block_keys: list[tuple[str, str]] = [] - - for sft_key in reader.keys(): # noqa: SIM118 - model_key = self.model_sd_ops.apply_to_key(sft_key) if self.model_sd_ops else sft_key - if model_key is None: - continue - if model_key.startswith(self.blocks_prefix + "."): - rest = model_key[len(self.blocks_prefix) + 1 :] - idx_str, _, param_name = rest.partition(".") - try: - block_idx = int(idx_str) - except ValueError: - non_block_keys.append((sft_key, model_key)) - continue - block_key_map.setdefault(block_idx, []).append((sft_key, param_name)) - else: - non_block_keys.append((sft_key, model_key)) self._load_non_block_weights( reader, @@ -242,9 +250,11 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType]) sd_ops=self.model_sd_ops, key_prefix=self.state_dict_prefix, lora_sources=lora_sources, - matmul_device=target_device, ) + blocks = resolve_attr(meta_model, self.blocks_attr) + layout = derive_layout(dict(blocks[0].named_parameters()), dtype) + cpu_pool = WeightPool( layout, cpu_slots_count, @@ -252,7 +262,12 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType]) reuse_barrier=lambda event: event.synchronize(), pin_memory=True, ) - block_reader = DiskBlockReader(reader=reader, block_key_map=block_key_map, dtype=dtype) + block_reader = DiskBlockReader( + reader=reader, + block_key_map=block_key_map, + sd_ops=self.model_sd_ops, + blocks_prefix=self.blocks_prefix, + ) source = DiskWeightSource(cpu_pool, block_reader) return source, lora_sources @@ -265,17 +280,17 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType]) model_key: str, tensor: torch.Tensor, lora_sources: list[LoraSource], - matmul_device: torch.device | None = None, ) -> torch.Tensor: - """Add all matching LoRA deltas to *tensor* in-place.""" + """Add all matching LoRA deltas to *tensor* in-place via ``addmm_``.""" if not lora_sources or not model_key.endswith(".weight"): return tensor prefix = model_key[: -len(".weight")] - device = tensor.device if tensor.device.type == "cuda" else matmul_device - for source in lora_sources: - delta = source.get_delta(prefix, device=device) - if delta is not None: - tensor = tensor.add_(delta.to(device=tensor.device, dtype=tensor.dtype)) + products = ( + ab + for ab in (s.get_ab(prefix, device=tensor.device, dtype=tensor.dtype) for s in lora_sources) + if ab is not None + ) + aggregate_lora_products(products, out=tensor) return tensor @staticmethod @@ -289,17 +304,48 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType]) sd_ops: SDOps | None = None, key_prefix: str = "", lora_sources: list[LoraSource] | None = None, - matmul_device: torch.device | None = None, ) -> None: """Load non-block weights into *model* on *device*.""" state_dict: dict[str, torch.Tensor] = {} sources = lora_sources or [] for sft_key, model_key in non_block_keys: tensor = reader.get_tensor(sft_key).to(device=device, dtype=dtype) - tensor = StreamingModelBuilder._fuse_lora_delta(model_key, tensor, sources, matmul_device) + tensor = StreamingModelBuilder._fuse_lora_delta(model_key, tensor, sources) if sd_ops is not None: for kv in sd_ops.apply_to_key_value(model_key, tensor): state_dict[key_prefix + kv.new_key] = kv.new_value continue state_dict[key_prefix + model_key] = tensor model.load_state_dict(state_dict, strict=False, assign=True) + + +def _scan_checkpoint_keys( + checkpoint_paths: list[str], + sd_ops: SDOps | None, + blocks_prefix: str, +) -> tuple[dict[int, list[tuple[str, str]]], list[tuple[str, str]]]: + """Partition checkpoint keys into per-block and non-block lists. + Opens the safetensors files for header-only key enumeration; no tensor data + is read. + """ + block_key_map: dict[int, list[tuple[str, str]]] = {} + non_block_keys: list[tuple[str, str]] = [] + prefix_dot = blocks_prefix + "." + for path in checkpoint_paths: + with safetensors.safe_open(path, framework="pt", device="cpu") as handle: + for sft_key in handle.keys(): # noqa: SIM118 + model_key = sd_ops.apply_to_key(sft_key) if sd_ops else sft_key + if model_key is None: + continue + if model_key.startswith(prefix_dot): + rest = model_key[len(prefix_dot) :] + idx_str, _, param_name = rest.partition(".") + try: + block_idx = int(idx_str) + except ValueError: + non_block_keys.append((sft_key, model_key)) + continue + block_key_map.setdefault(block_idx, []).append((sft_key, param_name)) + else: + non_block_keys.append((sft_key, model_key)) + return block_key_map, non_block_keys diff --git a/packages/ltx-core/src/ltx_core/block_streaming/disk.py b/packages/ltx-core/src/ltx_core/block_streaming/disk.py index 90d7bc6..7bd1746 100644 --- a/packages/ltx-core/src/ltx_core/block_streaming/disk.py +++ b/packages/ltx-core/src/ltx_core/block_streaming/disk.py @@ -2,11 +2,22 @@ from __future__ import annotations +from collections.abc import Iterator + import safetensors import torch +from ltx_core.block_streaming.utils import allocate_layout_views, make_block_key +from ltx_core.loader.fuse_loras import LoraProduct from ltx_core.loader.sd_ops import SDOps +_SAFETENSORS_DTYPE_TO_TORCH: dict[str, torch.dtype] = { + "F64": torch.float64, + "F32": torch.float32, + "F16": torch.float16, + "BF16": torch.bfloat16, +} + class DiskTensorReader: """Key-based tensor accessor over one or more safetensors files.""" @@ -21,9 +32,6 @@ class DiskTensorReader: for sft_key in handle.keys(): # noqa: SIM118 self._key_to_handle_idx[sft_key] = handle_idx - def keys(self) -> list[str]: - return list(self._key_to_handle_idx.keys()) - def get_tensor(self, key: str) -> torch.Tensor: return self._handles[self._key_to_handle_idx[key]].get_tensor(key) @@ -31,49 +39,58 @@ class DiskTensorReader: self._handles.clear() self._key_to_handle_idx.clear() + def __contains__(self, key: str) -> bool: + return key in self._key_to_handle_idx + + def __iter__(self) -> Iterator[str]: + return iter(self._key_to_handle_idx) + class DiskBlockReader: - """Reads one block at a time from safetensors into provided buffers. - Maps block indices to safetensors keys via a pre-computed key map. - """ + """Reads one block at a time from safetensors into provided buffers.""" def __init__( self, reader: DiskTensorReader, block_key_map: dict[int, list[tuple[str, str]]], - dtype: torch.dtype, + sd_ops: SDOps | None = None, + blocks_prefix: str = "", ) -> None: self._reader = reader self._block_key_map = block_key_map - self._dtype = dtype + self._sd_ops = sd_ops + self._blocks_prefix = blocks_prefix def read_into(self, target: dict[str, torch.Tensor], block_idx: int) -> None: + block_prefix = make_block_key(self._blocks_prefix, block_idx, "") for sft_key, param_name in self._block_key_map[block_idx]: tensor = self._reader.get_tensor(sft_key) - if tensor.dtype != self._dtype: - tensor = tensor.to(self._dtype) - target[param_name].copy_(tensor) + if self._sd_ops is None: + target[param_name].copy_(tensor) + continue + full_key = make_block_key(self._blocks_prefix, block_idx, param_name) + for result in self._sd_ops.apply_to_key_value(full_key, tensor): + if not result.new_key.startswith(block_prefix): + raise ValueError( + f"SDOps output key '{result.new_key}' is outside block {block_idx} " + f"(expected prefix '{block_prefix}'); cannot route to a per-block buffer." + ) + target[result.new_key[len(block_prefix) :]].copy_(result.new_value) def cleanup(self) -> None: self._reader.close() class LoraSource: - """Pinned-memory cache of LoRA A/B matrices for on-the-fly fusion. - At init, loads all matched A/B pairs into pinned CPU memory. - :meth:`get_delta` computes ``(B * strength) @ A`` on the given device. - """ + """Pinned-memory cache of matched LoRA A/B factors backed by a single buffer.""" def __init__(self, path: str, sd_ops: SDOps | None, strength: float) -> None: self.strength = strength - - # param_prefix -> (pinned_a, pinned_b) self._pinned_ab: dict[str, tuple[torch.Tensor, torch.Tensor]] = {} a_keys: dict[str, str] = {} b_keys: dict[str, str] = {} with safetensors.safe_open(path, framework="pt", device="cpu") as handle: - # First pass: build key map. for sft_key in handle.keys(): # noqa: SIM118 model_key = sd_ops.apply_to_key(sft_key) if sd_ops is not None else sft_key if model_key is None: @@ -83,24 +100,49 @@ class LoraSource: elif model_key.endswith(".lora_B.weight"): b_keys[model_key[: -len(".lora_B.weight")]] = sft_key - # Second pass: load and pin matched A+B pairs (orphans silently skipped). - for prefix in a_keys.keys() & b_keys.keys(): - self._pinned_ab[prefix] = ( - handle.get_tensor(a_keys[prefix]).pin_memory(), - handle.get_tensor(b_keys[prefix]).pin_memory(), + matched_prefixes = list(a_keys.keys() & b_keys.keys()) + + # Build the layout from safetensors header metadata only — no tensor data is read. + layout: dict[str, tuple[torch.Size, torch.dtype]] = {} + for prefix in matched_prefixes: + a_slice_view = handle.get_slice(a_keys[prefix]) + b_slice_view = handle.get_slice(b_keys[prefix]) + layout[f"{prefix}.A"] = ( + torch.Size(a_slice_view.get_shape()), + _SAFETENSORS_DTYPE_TO_TORCH[a_slice_view.get_dtype()], + ) + layout[f"{prefix}.B"] = ( + torch.Size(b_slice_view.get_shape()), + _SAFETENSORS_DTYPE_TO_TORCH[b_slice_view.get_dtype()], ) - def get_delta(self, param_prefix: str, device: torch.device | None = None) -> torch.Tensor | None: - """Return ``(B * strength) @ A`` for *param_prefix*, or ``None``.""" + all_views = allocate_layout_views(layout, pin_memory=True) + + for prefix in matched_prefixes: + a_view = all_views[f"{prefix}.A"] + b_view = all_views[f"{prefix}.B"] + a_view.copy_(handle.get_tensor(a_keys[prefix])) + b_view.copy_(handle.get_tensor(b_keys[prefix])) + self._pinned_ab[prefix] = (a_view, b_view) + + def get_ab( + self, + param_prefix: str, + device: torch.device | None = None, + dtype: torch.dtype | None = None, + ) -> LoraProduct | None: + """Return the :class:`LoraProduct` for *param_prefix*, or ``None``.""" pair = self._pinned_ab.get(param_prefix) if pair is None: return None a, b = pair if device is not None and device.type == "cuda": - a = a.to(device=device) - b = b.to(device=device) - delta = torch.matmul(b * self.strength, a) - return delta + a = a.to(device=device, non_blocking=True) + b = b.to(device=device, non_blocking=True) + if dtype is not None: + a = a.to(dtype=dtype) + b = b.to(dtype=dtype) + return LoraProduct(a, b, self.strength) def cleanup(self) -> None: self._pinned_ab.clear() diff --git a/packages/ltx-core/src/ltx_core/block_streaming/pool.py b/packages/ltx-core/src/ltx_core/block_streaming/pool.py index 4a9a163..8dac072 100644 --- a/packages/ltx-core/src/ltx_core/block_streaming/pool.py +++ b/packages/ltx-core/src/ltx_core/block_streaming/pool.py @@ -7,20 +7,16 @@ from typing import Callable import torch -from ltx_core.block_streaming.utils import allocate_buffer - -# Type alias for the buffer layout used by slot allocation. -BlockLayout = dict[str, tuple[torch.Size, torch.dtype]] +from ltx_core.block_streaming.utils import allocate_layout_views +from ltx_core.loader.primitives import TensorLayout class WeightPool: - """Fixed pool of pre-allocated weight buffers with event-based reuse safety. - Buffers are allocated once at construction. :meth:`acquire` pops a - free buffer (waiting any pending event first). :meth:`release` - returns it, optionally attaching an event that must complete before - the buffer can be reused. + """Fixed pool of pre-allocated weight buffers with event-based reuse. + All slots share a single buffer (CPU or GPU); each slot is a + contiguous slice carved out of it via :func:`allocate_layout_views`. Args: - layout: ``{name: (shape, dtype)}`` for each buffer. + buffer_layout: ``{name: (shape, dtype)}`` for each buffer. capacity: Number of buffers to pre-allocate. device: Device for allocation. reuse_barrier: Called with the pending event before a buffer is reused. @@ -29,23 +25,34 @@ class WeightPool: def __init__( self, - layout: BlockLayout, + buffer_layout: TensorLayout, capacity: int, device: torch.device, reuse_barrier: Callable[[torch.cuda.Event], None], pin_memory: bool = False, ) -> None: + self._buffer_layout = buffer_layout self._capacity = capacity self._free: deque[dict[str, torch.Tensor]] = deque() self._events: dict[int, torch.cuda.Event] = {} self._reuse_barrier = reuse_barrier - for _ in range(capacity): - self._free.append(allocate_buffer(layout, device, pin_memory)) + memory_layout = { + _make_key(slot, name): (shape, dtype) + for slot in range(capacity) + for name, (shape, dtype) in buffer_layout.items() + } + all_views = allocate_layout_views(memory_layout, device=device, pin_memory=pin_memory) + for slot in range(capacity): + self._free.append({name: all_views[_make_key(slot, name)] for name in buffer_layout}) @property def capacity(self) -> int: return self._capacity + @property + def buffer_layout(self) -> TensorLayout: + return self._buffer_layout + def acquire(self) -> dict[str, torch.Tensor]: """Take a free buffer, waiting any pending event before returning.""" weights = self._free.popleft() @@ -62,3 +69,7 @@ class WeightPool: if event is not None: self._events[id(weights)] = event self._free.append(weights) + + +def _make_key(slot: int, name: str) -> str: + return f"{slot}/{name}" diff --git a/packages/ltx-core/src/ltx_core/block_streaming/provider.py b/packages/ltx-core/src/ltx_core/block_streaming/provider.py index 59646cb..306ab36 100644 --- a/packages/ltx-core/src/ltx_core/block_streaming/provider.py +++ b/packages/ltx-core/src/ltx_core/block_streaming/provider.py @@ -9,6 +9,29 @@ import torch from ltx_core.block_streaming.disk import LoraSource from ltx_core.block_streaming.pool import WeightPool from ltx_core.block_streaming.source import WeightSource +from ltx_core.block_streaming.utils import FP8_DTYPES +from ltx_core.loader.fuse_loras import aggregate_lora_products, fuse_cast_fp8_weight + + +def _contiguous_byte_view(weights: dict[str, torch.Tensor]) -> torch.Tensor | None: + """Return a ``uint8`` view spanning every tensor in *weights*, or ``None`` if + they don't share one contiguous storage region.""" + tensors = list(weights.values()) + if not tensors: + return None + storage = tensors[0].untyped_storage() + storage_ptr = storage.data_ptr() + start = end = tensors[0].storage_offset() * tensors[0].element_size() + for t in tensors: + if t.untyped_storage().data_ptr() != storage_ptr or not t.is_contiguous(): + return None + offset = t.storage_offset() * t.element_size() + nbytes = t.numel() * t.element_size() + start = min(start, offset) + end = max(end, offset + nbytes) + view = torch.empty(0, dtype=torch.uint8, device=tensors[0].device) + view.set_(storage, start, (end - start,), (1,)) + return view class WeightsProvider: @@ -70,8 +93,13 @@ class WeightsProvider: instrumentation regions wrapping it -- observe the full transfer time. """ with torch.cuda.stream(self._copy_stream): - for name, gpu_tensor in gpu_weights.items(): - gpu_tensor.copy_(cpu_weights[name], non_blocking=True) + gpu_view = _contiguous_byte_view(gpu_weights) + cpu_view = _contiguous_byte_view(cpu_weights) + if gpu_view is not None and cpu_view is not None and gpu_view.numel() == cpu_view.numel(): + gpu_view.copy_(cpu_view, non_blocking=True) + else: + for name, gpu_tensor in gpu_weights.items(): + gpu_tensor.copy_(cpu_weights[name], non_blocking=True) if self._lora_sources: self._fuse_block_loras(idx, gpu_weights) h2d_event = torch.cuda.Event() @@ -102,9 +130,18 @@ class WeightsProvider: for name, tensor in weights.items(): if not name.endswith(".weight"): continue - full_key = f"{self._blocks_prefix}.{idx}.{name}" - prefix = full_key[: -len(".weight")] - for source in self._lora_sources: - delta = source.get_delta(prefix, device=self._target_device) - if delta is not None: - tensor.add_(delta.to(dtype=tensor.dtype)) + prefix = f"{self._blocks_prefix}.{idx}.{name}".removesuffix(".weight") + is_fp8 = tensor.dtype in FP8_DTYPES + agg_dtype = torch.bfloat16 if is_fp8 else tensor.dtype + products = ( + ab + for ab in (s.get_ab(prefix, device=self._target_device, dtype=agg_dtype) for s in self._lora_sources) + if ab is not None + ) + aggregated = aggregate_lora_products(products, agg_dtype) + if aggregated is None: + continue + if is_fp8: + tensor.copy_(fuse_cast_fp8_weight(aggregated, tensor, tensor.dtype)) + else: + tensor.add_(aggregated) diff --git a/packages/ltx-core/src/ltx_core/block_streaming/source.py b/packages/ltx-core/src/ltx_core/block_streaming/source.py index 7cfe842..62d0d42 100644 --- a/packages/ltx-core/src/ltx_core/block_streaming/source.py +++ b/packages/ltx-core/src/ltx_core/block_streaming/source.py @@ -9,10 +9,18 @@ import torch from ltx_core.block_streaming.disk import DiskBlockReader from ltx_core.block_streaming.pool import WeightPool +from ltx_core.loader.primitives import TensorLayout class WeightSource(Protocol): - """Provides pinned CPU weights for a given block index.""" + """Provides pinned CPU weights for a given block index. + Assumes all buffers share an identical layout across all block indices. + """ + + @property + def block_layout(self) -> TensorLayout: + """Shared per-block buffer layout (shape + dtype for each param).""" + ... def get(self, idx: int) -> dict[str, torch.Tensor]: """Return CPU weights for block *idx*.""" @@ -36,6 +44,10 @@ class DiskWeightSource(WeightSource): self._events: dict[int, torch.cuda.Event] = {} self._reader = reader + @property + def block_layout(self) -> TensorLayout: + return self._pool.buffer_layout + def get(self, idx: int) -> dict[str, torch.Tensor]: """Return CPU weights for block *idx*. Reads from disk on miss.""" if idx in self._cache: @@ -68,8 +80,15 @@ class PinnedWeightSource(WeightSource): """Pre-loaded pinned CPU weights.""" def __init__(self, weights: dict[int, dict[str, torch.Tensor]]) -> None: + if not weights: + raise ValueError("PinnedWeightSource requires at least one block") self._weights = weights + @property + def block_layout(self) -> TensorLayout: + first_block = self._weights[min(self._weights)] + return {name: (t.shape, t.dtype) for name, t in first_block.items()} + def get(self, idx: int) -> dict[str, torch.Tensor]: return self._weights[idx] diff --git a/packages/ltx-core/src/ltx_core/block_streaming/utils.py b/packages/ltx-core/src/ltx_core/block_streaming/utils.py index acb4a2a..c99cc6d 100644 --- a/packages/ltx-core/src/ltx_core/block_streaming/utils.py +++ b/packages/ltx-core/src/ltx_core/block_streaming/utils.py @@ -2,14 +2,24 @@ from __future__ import annotations -import itertools -from typing import TYPE_CHECKING, Any +import math +import weakref +from dataclasses import dataclass +from typing import Any import torch from torch import nn -if TYPE_CHECKING: - from ltx_core.block_streaming.pool import BlockLayout +from ltx_core.loader.primitives import TensorLayout + +FP8_DTYPES = frozenset({torch.float8_e4m3fn, torch.float8_e5m2}) + +_BUFFER_ALIGN = 16 + + +def make_block_key(blocks_prefix: str, block_idx: int, param_name: str) -> str: + """Return the state-dict key for *param_name* under block *block_idx*.""" + return f"{blocks_prefix}.{block_idx}.{param_name}" def resolve_attr(module: nn.Module, dotted_path: str) -> nn.ModuleList: @@ -40,21 +50,85 @@ def assign_tensor_to_module(root: nn.Module, dotted_name: str, tensor: torch.Ten raise AttributeError(f"{leaf} is not a parameter or buffer of {type(parent).__name__}") -def build_pool_layout(block: nn.Module, dtype: torch.dtype) -> BlockLayout: - """Derive a buffer layout from a block's parameters and buffers. - Works on meta-device blocks (shapes are valid regardless of device). - The *dtype* argument overrides each tensor's dtype so the pool matches - the target inference precision. +def derive_layout(tensors: dict[str, torch.Tensor], dtype: torch.dtype | None = None) -> TensorLayout: + """Derive a layout from a ``{name: tensor}`` dict. + If ``dtype`` is given, non-FP8 dtypes are coerced to it (FP8 preserved). If + ``None``, the source dtype is preserved as-is. """ - layout: BlockLayout = {} - for name, tensor in itertools.chain(block.named_parameters(), block.named_buffers()): - layout[name] = (tensor.shape, dtype) - return layout - - -def allocate_buffer(layout: BlockLayout, device: torch.device, pin_memory: bool = False) -> dict[str, torch.Tensor]: - """Allocate a single buffer dict matching *layout*.""" return { - name: torch.empty(shape, dtype=dtype, device=device, pin_memory=pin_memory) - for name, (shape, dtype) in layout.items() + name: (t.shape, t.dtype if dtype is None or t.dtype in FP8_DTYPES else dtype) for name, t in tensors.items() } + + +def _align_up(offset: int, alignment: int) -> int: + return (offset + alignment - 1) & ~(alignment - 1) + + +def _alloc_pinned_exact(nbytes: int) -> torch.Tensor | None: + """Allocate exactly ``nbytes`` of pinned host memory via ``cudaHostRegister``. + Bypasses PyTorch's ``CachingHostAllocator``, which rounds every + ``pin_memory=True`` request up to ``PowerOf2Ceil(N)`` (see + ``aten/src/ATen/core/CachingHostAllocator.h``). Returns ``None`` if + registration fails. The unregister hook is bound to the storage (not the + tensor) so views of the buffer keep the registration alive until the + memory is actually freed. Caller is responsible for ensuring CUDA is + available. + """ + cudart = torch.cuda.cudart() + buf = torch.empty(nbytes, dtype=torch.uint8) + ptr = buf.data_ptr() + err = int(cudart.cudaHostRegister(ptr, nbytes, 0)) + if err != 0: + return None + weakref.finalize(buf.untyped_storage(), lambda p=ptr: cudart.cudaHostUnregister(p)) + return buf + + +def _alloc_buffer(nbytes: int, device: torch.device | None, pin_memory: bool) -> torch.Tensor: + """Allocate one ``uint8`` buffer for :func:`allocate_layout_views`. + For pinned host buffers, prefer ``cudaHostRegister`` to dodge the caching + allocator's power-of-2 rounding. Falls back to the caching allocator if + registration fails. Raises if pinning is requested without a CUDA runtime, + since pinning is fundamentally a CUDA driver operation. + """ + if pin_memory and (device is None or torch.device(device).type == "cpu"): + if not torch.cuda.is_available(): + raise RuntimeError("pin_memory=True requires CUDA, which is not available") + buf = _alloc_pinned_exact(nbytes) + if buf is not None: + return buf + return torch.empty(nbytes, dtype=torch.uint8, device=device, pin_memory=pin_memory) + + +@dataclass(frozen=True) +class _TensorSlice: + """Location of a single tensor view within the buffer.""" + + offset: int + shape: torch.Size + dtype: torch.dtype + + def size(self) -> int: + return math.prod(self.shape) * self.dtype.itemsize + + +def allocate_layout_views( + layout: TensorLayout, + device: torch.device | None = None, + pin_memory: bool = False, +) -> dict[str, torch.Tensor]: + """Allocate a single ``uint8`` buffer and return per-key tensor views into it. + All keys in *layout* live in one contiguous allocation; each returned + tensor is a non-overlapping slice of that buffer reinterpreted at the + requested shape and dtype. The views keep the underlying storage alive + via PyTorch refcounting — drop them all to release the memory. + """ + slices: dict[str, _TensorSlice] = {} + cursor = 0 + for key, (shape, dtype) in layout.items(): + cursor = _align_up(cursor, _BUFFER_ALIGN) + slices[key] = _TensorSlice(offset=cursor, shape=shape, dtype=dtype) + cursor += slices[key].size() + # Allocate at least one byte so empty layouts still produce a valid buffer. + buffer = _alloc_buffer(max(_align_up(cursor, _BUFFER_ALIGN), 1), device, pin_memory) + return {key: buffer[s.offset : s.offset + s.size()].view(s.dtype).view(s.shape) for key, s in slices.items()} diff --git a/packages/ltx-core/src/ltx_core/components/diffusion_steps.py b/packages/ltx-core/src/ltx_core/components/diffusion_steps.py index d4908cb..9d83e82 100644 --- a/packages/ltx-core/src/ltx_core/components/diffusion_steps.py +++ b/packages/ltx-core/src/ltx_core/components/diffusion_steps.py @@ -4,6 +4,24 @@ from ltx_core.components.protocols import DiffusionStepProtocol from ltx_core.utils import to_velocity +def _get_ancestral_step( + sigma_from: torch.Tensor, + sigma_to: torch.Tensor, + eta: float = 1.0, +) -> tuple[torch.Tensor, torch.Tensor]: + """Compute ``(sigma_down, sigma_up)`` for one DDIM ancestral sampling step. + Both inputs are in the rescaled parameterization ``sigma / alpha``. + Returns ``sigma_down`` (deterministic component) and ``sigma_up`` + (stochastic component) in the same rescaled space. + """ + if not eta: + return sigma_to, torch.zeros_like(sigma_to) + variance = sigma_to**2 * (sigma_from**2 - sigma_to**2).clamp(min=0) / sigma_from**2 + sigma_up = (eta * variance**0.5).clamp(max=sigma_to) + sigma_down = (sigma_to**2 - sigma_up**2).clamp(min=0) ** 0.5 + return sigma_down, sigma_up + + class EulerDiffusionStep(DiffusionStepProtocol): """ First-order Euler method for diffusion sampling. @@ -104,3 +122,65 @@ class Res2sDiffusionStep(DiffusionStepProtocol): # Mix deterministic and stochastic components x_noised = alpha_ratio * (denoised_next + sigma_down * eps_next) + sigma_up * noise return x_noised.to(output_dtype) + + +class EulerCfgPpDiffusionStep(DiffusionStepProtocol): + """Euler step using the CFG++ correction for the ODE derivative. + Instead of the standard velocity formula, the ODE derivative is computed + from the unconditioned prediction, keeping the conditioned prediction as + the target denoised state. Ancestral (DDIM) noise injection is applied + in the rescaled sigma parameterization (sigma / alpha). + All diffusion quantities (alpha, ODE derivative, ancestral coefficients) + are computed internally from ``sigmas`` and ``uncond_denoised``. + Reference: CFG++ (https://arxiv.org/abs/2406.08070). + """ + + def __init__(self, eta: float = 1.0, s_noise: float = 1.0) -> None: + self.eta = eta + self.s_noise = s_noise + + def step( + self, + sample: torch.Tensor, + denoised_sample: torch.Tensor, + sigmas: torch.Tensor, + step_index: int, + uncond_denoised: torch.Tensor, + noise: torch.Tensor | None = None, + **_kwargs, + ) -> torch.Tensor: + """Advance one CFG++ Euler step. + Args: + sample: Current noisy latent x_t. + denoised_sample: Conditioned denoised prediction x_0^cond. + sigmas: Full sigma schedule tensor. + step_index: Current step index. + uncond_denoised: Unconditioned denoised prediction x_0^uncond, + used to compute the ODE derivative direction. + noise: Noise tensor for stochastic injection; ignored when + ``eta=0`` or ``s_noise=0``. + Returns: + Updated latent x_{t-1}. + """ + sigma_s = sigmas[step_index].to(torch.float32) + sigma_t = sigmas[step_index + 1].to(torch.float32) + _eps = torch.finfo(torch.float32).eps + # Clamp to avoid division by zero when sigma == 1.0 exactly. + alpha_s = (1.0 - sigma_s).clamp(min=_eps) + alpha_t = (1.0 - sigma_t).clamp(min=_eps) + + x = sample.to(torch.float32) + denoised = denoised_sample.to(torch.float32) + uncond = uncond_denoised.to(torch.float32) + + # ODE derivative: direction toward noise using uncond prediction (CFG++ correction) + d = (x - alpha_s * uncond) / sigma_s + + # Ancestral step in rescaled sigma space (sigma / alpha) + sigma_down, sigma_up = _get_ancestral_step(sigma_s / alpha_s, sigma_t / alpha_t, eta=self.eta) + sigma_down = alpha_t * sigma_down + + x_next = alpha_t * denoised + sigma_down * d + if noise is not None and self.eta > 0 and self.s_noise > 0: + x_next = x_next + alpha_t * noise.to(torch.float32) * self.s_noise * sigma_up + return x_next.to(sample.dtype) diff --git a/packages/ltx-core/src/ltx_core/components/patchifiers.py b/packages/ltx-core/src/ltx_core/components/patchifiers.py index f9580d5..77d6521 100644 --- a/packages/ltx-core/src/ltx_core/components/patchifiers.py +++ b/packages/ltx-core/src/ltx_core/components/patchifiers.py @@ -151,17 +151,22 @@ def get_pixel_coords( that treat frame zero differently still yield non-negative timestamps. """ # Broadcast the VAE scale factors so they align with the `(batch, axis, patch, bound)` layout. + # Axis 1 of `latent_coords` is ordered (frame/time, height, width) — match that explicitly by + # pulling fields from the NamedTuple rather than relying on tuple iteration order. broadcast_shape = [1] * latent_coords.ndim broadcast_shape[1] = -1 # axis dimension corresponds to (frame/time, height, width) - scale_tensor = torch.tensor(scale_factors, device=latent_coords.device).view(*broadcast_shape) + scale_tensor = torch.tensor( + [scale_factors.time, scale_factors.height, scale_factors.width], + device=latent_coords.device, + ).view(*broadcast_shape) # Apply per-axis scaling to convert latent bounds into pixel-space coordinates. pixel_coords = latent_coords * scale_tensor if causal_fix: - # VAE temporal stride for the very first frame is 1 instead of `scale_factors[0]`. + # VAE temporal stride for the very first frame is 1 instead of `scale_factors.time`. # Shift and clamp to keep the first-frame timestamps causal and non-negative. - pixel_coords[:, 0, ...] = (pixel_coords[:, 0, ...] + 1 - scale_factors[0]).clamp(min=0) + pixel_coords[:, 0, ...] = (pixel_coords[:, 0, ...] + 1 - scale_factors.time).clamp(min=0) return pixel_coords diff --git a/packages/ltx-core/src/ltx_core/conditioning/__init__.py b/packages/ltx-core/src/ltx_core/conditioning/__init__.py index 002e91f..58b82d0 100644 --- a/packages/ltx-core/src/ltx_core/conditioning/__init__.py +++ b/packages/ltx-core/src/ltx_core/conditioning/__init__.py @@ -3,6 +3,7 @@ from ltx_core.conditioning.exceptions import ConditioningError from ltx_core.conditioning.item import ConditioningItem from ltx_core.conditioning.types import ( + AudioConditionByReferenceLatent, ConditioningItemAttentionStrengthWrapper, VideoConditionByKeyframeIndex, VideoConditionByLatentIndex, @@ -10,6 +11,7 @@ from ltx_core.conditioning.types import ( ) __all__ = [ + "AudioConditionByReferenceLatent", "ConditioningError", "ConditioningItem", "ConditioningItemAttentionStrengthWrapper", diff --git a/packages/ltx-core/src/ltx_core/conditioning/types/__init__.py b/packages/ltx-core/src/ltx_core/conditioning/types/__init__.py index 44bb920..77cfd0b 100644 --- a/packages/ltx-core/src/ltx_core/conditioning/types/__init__.py +++ b/packages/ltx-core/src/ltx_core/conditioning/types/__init__.py @@ -3,9 +3,11 @@ from ltx_core.conditioning.types.attention_strength_wrapper import ConditioningItemAttentionStrengthWrapper from ltx_core.conditioning.types.keyframe_cond import VideoConditionByKeyframeIndex from ltx_core.conditioning.types.latent_cond import VideoConditionByLatentIndex +from ltx_core.conditioning.types.reference_audio_cond import AudioConditionByReferenceLatent from ltx_core.conditioning.types.reference_video_cond import VideoConditionByReferenceLatent __all__ = [ + "AudioConditionByReferenceLatent", "ConditioningItemAttentionStrengthWrapper", "VideoConditionByKeyframeIndex", "VideoConditionByLatentIndex", diff --git a/packages/ltx-core/src/ltx_core/conditioning/types/reference_audio_cond.py b/packages/ltx-core/src/ltx_core/conditioning/types/reference_audio_cond.py new file mode 100644 index 0000000..a2d0c42 --- /dev/null +++ b/packages/ltx-core/src/ltx_core/conditioning/types/reference_audio_cond.py @@ -0,0 +1,59 @@ +"""Audio reference conditioning items.""" + +from __future__ import annotations + +import torch + +from ltx_core.conditioning.mask_utils import update_attention_mask +from ltx_core.tools import LatentTools +from ltx_core.types import LatentState + + +class AudioConditionByReferenceLatent: + """Append patchified reference audio tokens after the target audio sequence. + Mirrors :class:`ltx_core.conditioning.types.reference_video_cond.VideoConditionByReferenceLatent` + but for audio. The reference tokens are appended so the target audio tokens stay + in the first ``num_noisy_tokens`` positions and can be kept by + :meth:`ltx_core.tools.LatentTools.clear_conditioning`. + Args: + patchified: Patchified reference latent ``[B, T_ref, C]``. + positions: RoPE positions for reference tokens, ``[B, 1, T_ref, 2]``. + strength: 1.0 keeps reference clean; 0.0 would fully denoise it. + """ + + def __init__( + self, + patchified: torch.Tensor, + positions: torch.Tensor, + strength: float = 1.0, + ) -> None: + self.patchified = patchified + self.positions = positions.to(dtype=torch.float32) + self.strength = strength + + def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState: + tokens = self.patchified + denoise_mask = torch.full( + size=(*tokens.shape[:2], 1), + fill_value=1.0 - self.strength, + device=tokens.device, + dtype=tokens.dtype, + ) + + new_attention_mask = update_attention_mask( + latent_state=latent_state, + attention_mask=None, + num_noisy_tokens=latent_tools.patchifier.get_token_count(latent_tools.target_shape), + num_new_tokens=tokens.shape[1], + batch_size=tokens.shape[0], + device=tokens.device, + dtype=tokens.dtype, + ) + + return LatentState( + latent=torch.cat([latent_state.latent, tokens], dim=1), + denoise_mask=torch.cat([latent_state.denoise_mask, denoise_mask], dim=1), + positions=torch.cat([latent_state.positions, self.positions], dim=2), + clean_latent=torch.cat([latent_state.clean_latent, tokens], dim=1), + attention_mask=new_attention_mask, + ) diff --git a/packages/ltx-core/src/ltx_core/loader/fuse_loras.py b/packages/ltx-core/src/ltx_core/loader/fuse_loras.py index 00eecf6..51f0304 100644 --- a/packages/ltx-core/src/ltx_core/loader/fuse_loras.py +++ b/packages/ltx-core/src/ltx_core/loader/fuse_loras.py @@ -1,26 +1,81 @@ -from collections.abc import Iterator +from collections.abc import Iterable, Iterator +from typing import NamedTuple import torch +from ltx_core.loader.kernels import TRITON_AVAILABLE from ltx_core.loader.primitives import LoraStateDictWithStrength, StateDict -from ltx_core.quantization.fp8_cast import _fused_add_round_launch +from ltx_core.quantization.fp8_cast import fused_add_round_launch from ltx_core.quantization.fp8_scaled_mm import quantize_weight_to_fp8_per_tensor +class LoraProduct(NamedTuple): + """A LoRA's ``A``, ``B`` factors and its strength scalar.""" + + a: torch.Tensor + b: torch.Tensor + strength: float + + def _get_device() -> torch.device: if torch.cuda.is_available(): return torch.device("cuda", torch.cuda.current_device()) return torch.device("cpu") +def aggregate_lora_products( + products: Iterable[LoraProduct], + dtype: torch.dtype | None = None, + *, + out: torch.Tensor | None = None, +) -> torch.Tensor | None: + """Accumulate ``sum((B * strength) @ A)`` across :class:`LoraProduct` items. + If ``out`` is provided, ``addmm_`` accumulates directly into it — caller + ensures A/B dtypes and devices match ``out``. Otherwise the first product + materializes the ``(out, in)``-shape aggregator at ``dtype``; subsequent + products use ``addmm_`` to avoid allocating the full intermediate delta. + Returns ``out`` (or the new aggregator), or ``None`` if ``products`` was empty + and ``out`` was not given. + """ + aggregated = out + for product in products: + if aggregated is None: + aggregated = torch.matmul(product.b * product.strength, product.a).to(dtype=dtype) + else: + aggregated.addmm_(product.b, product.a, alpha=product.strength) + return aggregated + + +def fuse_cast_fp8_weight( + delta_bf16: torch.Tensor, + weight_fp8: torch.Tensor, + target_dtype: torch.dtype, +) -> torch.Tensor: + """Return ``(delta_bf16 + dequantize(weight_fp8)).to(target_dtype)``. + CUDA with Triton uses stochastic rounding; otherwise uses a deterministic bf16 add. + ``delta_bf16`` is the bf16 accumulator and is mutated in place. + """ + if delta_bf16.dtype != torch.bfloat16: + raise ValueError(f"delta_bf16 must be bfloat16, got {delta_bf16.dtype}") + if str(weight_fp8.device).startswith("cuda") and TRITON_AVAILABLE: + fused_add_round_launch(delta_bf16, weight_fp8, seed=0) + else: + delta_bf16.add_(weight_fp8.to(dtype=torch.bfloat16)) + return delta_bf16.to(dtype=target_dtype) + + def fuse_lora_weights( model_sd: StateDict, lora_sd_and_strengths: list[LoraStateDictWithStrength], dtype: torch.dtype | None = None, + preserve_input_device: bool = True, ) -> Iterator[tuple[str, torch.Tensor]]: """Yield ``(key, fused_tensor)`` for each weight modified by at least one LoRA. For scaled-FP8 weights, this includes both the updated ``.weight`` tensor and its corresponding ``.weight_scale`` tensor. + When ``preserve_input_device`` is False, fused tensors are yielded on the device + used for fusion; caller is responsible for moving them to their final + destination. """ for key, original_weight in model_sd.sd.items(): if original_weight is None or key.endswith(".weight_scale"): @@ -30,7 +85,7 @@ def fuse_lora_weights( target_dtype = dtype if dtype is not None else weight.dtype deltas_dtype = target_dtype if target_dtype not in [torch.float8_e4m3fn, torch.float8_e5m2] else torch.bfloat16 - deltas = _prepare_deltas(lora_sd_and_strengths, key, deltas_dtype, weight.device) + deltas = _aggregate_deltas(lora_sd_and_strengths, key, deltas_dtype, weight.device) if deltas is None: continue @@ -41,14 +96,15 @@ def fuse_lora_weights( if is_scaled_fp8: fused = _fuse_delta_with_scaled_fp8(deltas, weight, key, scale_key, model_sd) else: - fused = _fuse_delta_with_cast_fp8(deltas, weight, key, target_dtype) + fused = {key: fuse_cast_fp8_weight(deltas, weight, target_dtype)} elif weight.dtype == torch.bfloat16: - fused = _fuse_delta_with_bfloat16(deltas, weight, key, target_dtype) + deltas.add_(weight) + fused = {key: deltas.to(dtype=target_dtype)} else: raise ValueError(f"Unsupported dtype: {weight.dtype}") for k, v in fused.items(): - yield k, v.to(device=original_device) + yield k, v.to(device=original_device) if preserve_input_device else v def apply_loras( @@ -57,10 +113,12 @@ def apply_loras( dtype: torch.dtype | None = None, destination_sd: StateDict | None = None, ) -> StateDict: + """Fuse LoRAs into ``model_sd`` and place the results in ``destination_sd``. + When ``destination_sd`` is provided, the fused tensors are placed directly into it. + """ if destination_sd is not None: - sd = destination_sd.sd - for key, tensor in fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype): - sd[key] = tensor + for key, fused in fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype): + destination_sd.sd[key] = fused return destination_sd fused = dict(fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype)) @@ -68,26 +126,22 @@ def apply_loras( return StateDict(sd, model_sd.device, model_sd.size, model_sd.dtype) -def _prepare_deltas( +def _aggregate_deltas( lora_sd_and_strengths: list[LoraStateDictWithStrength], key: str, dtype: torch.dtype, device: torch.device ) -> torch.Tensor | None: - deltas = [] prefix = key[: -len(".weight")] key_a = f"{prefix}.lora_A.weight" key_b = f"{prefix}.lora_B.weight" - for lsd, coef in lora_sd_and_strengths: - if key_a not in lsd.sd or key_b not in lsd.sd: - continue - a = lsd.sd[key_a].to(device=device) - b = lsd.sd[key_b].to(device=device) - product = torch.matmul(b * coef, a) - del a, b - deltas.append(product.to(dtype=dtype)) - if len(deltas) == 0: - return None - elif len(deltas) == 1: - return deltas[0] - return torch.sum(torch.stack(deltas, dim=0), dim=0) + + def _ab_products() -> Iterator[LoraProduct]: + for lsd, coef in lora_sd_and_strengths: + if key_a not in lsd.sd or key_b not in lsd.sd: + continue + a = lsd.sd[key_a].to(device=device, dtype=dtype, non_blocking=True) + b = lsd.sd[key_b].to(device=device, dtype=dtype, non_blocking=True) + yield LoraProduct(a, b, coef) + + return aggregate_lora_products(_ab_products(), dtype) def _fuse_delta_with_scaled_fp8( @@ -100,34 +154,9 @@ def _fuse_delta_with_scaled_fp8( """Dequantize scaled FP8 weight, add LoRA delta, and re-quantize.""" weight_scale = model_sd.sd[scale_key] - original_weight = weight.t().to(torch.float32) * weight_scale + original_weight = weight.to(torch.float32) * weight_scale new_weight = original_weight + deltas.to(torch.float32) new_fp8_weight, new_weight_scale = quantize_weight_to_fp8_per_tensor(new_weight) return {key: new_fp8_weight, scale_key: new_weight_scale} - - -def _fuse_delta_with_cast_fp8( - deltas: torch.Tensor, - weight: torch.Tensor, - key: str, - target_dtype: torch.dtype, -) -> dict[str, torch.Tensor]: - """Fuse LoRA delta with cast-only FP8 weight (no scale factor).""" - if str(weight.device).startswith("cuda"): - _fused_add_round_launch(deltas, weight, seed=0) - else: - deltas.add_(weight.to(dtype=deltas.dtype)) - return {key: deltas.to(dtype=target_dtype)} - - -def _fuse_delta_with_bfloat16( - deltas: torch.Tensor, - weight: torch.Tensor, - key: str, - target_dtype: torch.dtype, -) -> dict[str, torch.Tensor]: - """Fuse LoRA delta with bfloat16 weight.""" - deltas.add_(weight) - return {key: deltas.to(dtype=target_dtype)} diff --git a/packages/ltx-core/src/ltx_core/loader/kernels.py b/packages/ltx-core/src/ltx_core/loader/kernels.py index ee4cefb..145cd88 100644 --- a/packages/ltx-core/src/ltx_core/loader/kernels.py +++ b/packages/ltx-core/src/ltx_core/loader/kernels.py @@ -1,72 +1,79 @@ # ruff: noqa: ANN001, ANN201, ERA001, N803, N806 -import triton -import triton.language as tl +try: + import triton + import triton.language as tl + + TRITON_AVAILABLE = True +except (ImportError, OSError): + TRITON_AVAILABLE = False -@triton.jit -def fused_add_round_kernel( - x_ptr, - output_ptr, # contents will be added to the output - seed, - n_elements, - EXPONENT_BIAS, - MANTISSA_BITS, - BLOCK_SIZE: tl.constexpr, -): - """ - A kernel to upcast 8bit quantized weights to bfloat16 with stochastic rounding - and add them to bfloat16 output weights. Might be used to upcast original model weights - and to further add them to precalculated deltas coming from LoRAs. - """ - # Get program ID and compute offsets - pid = tl.program_id(axis=0) - block_start = pid * BLOCK_SIZE - offsets = block_start + tl.arange(0, BLOCK_SIZE) - mask = offsets < n_elements +if TRITON_AVAILABLE: - # Load data - x = tl.load(x_ptr + offsets, mask=mask) - rand_vals = tl.rand(seed, offsets) - 0.5 + @triton.jit + def fused_add_round_kernel( + x_ptr, + output_ptr, # contents will be added to the output + seed, + n_elements, + EXPONENT_BIAS, + MANTISSA_BITS, + BLOCK_SIZE: tl.constexpr, + ): + """ + A kernel to upcast 8bit quantized weights to bfloat16 with stochastic rounding + and add them to bfloat16 output weights. Might be used to upcast original model weights + and to further add them to precalculated deltas coming from LoRAs. + """ + # Get program ID and compute offsets + pid = tl.program_id(axis=0) + block_start = pid * BLOCK_SIZE + offsets = block_start + tl.arange(0, BLOCK_SIZE) + mask = offsets < n_elements - x = tl.cast(x, tl.float16) - delta = tl.load(output_ptr + offsets, mask=mask) - delta = tl.cast(delta, tl.float16) - x = x + delta + # Load data + x = tl.load(x_ptr + offsets, mask=mask) + rand_vals = tl.rand(seed, offsets) - 0.5 - x_bits = tl.cast(x, tl.int16, bitcast=True) + x = tl.cast(x, tl.float16) + delta = tl.load(output_ptr + offsets, mask=mask) + delta = tl.cast(delta, tl.float16) + x = x + delta - # Calculate the exponent. Unbiased fp16 exponent is ((x_bits & 0x7C00) >> 10) - 15 for - # normal numbers and -14 for subnormals. - fp16_exponent_bits = (x_bits & 0x7C00) >> 10 - fp16_normals = fp16_exponent_bits > 0 - fp16_exponent = tl.where(fp16_normals, fp16_exponent_bits - 15, -14) + x_bits = tl.cast(x, tl.int16, bitcast=True) - # Add the target dtype's exponent bias and clamp to the target dtype's exponent range. - exponent = fp16_exponent + EXPONENT_BIAS - MAX_EXPONENT = 2 * EXPONENT_BIAS + 1 - exponent = tl.where(exponent > MAX_EXPONENT, MAX_EXPONENT, exponent) - exponent = tl.where(exponent < 0, 0, exponent) + # Calculate the exponent. Unbiased fp16 exponent is ((x_bits & 0x7C00) >> 10) - 15 for + # normal numbers and -14 for subnormals. + fp16_exponent_bits = (x_bits & 0x7C00) >> 10 + fp16_normals = fp16_exponent_bits > 0 + fp16_exponent = tl.where(fp16_normals, fp16_exponent_bits - 15, -14) - # Normal ULP exponent, expressed as an fp16 exponent field: - # (exponent - EXPONENT_BIAS - MANTISSA_BITS) + 15 - # Simplifies to: fp16_exponent - MANTISSA_BITS + 15 - # See https://en.wikipedia.org/wiki/Unit_in_the_last_place - eps_exp = tl.maximum(0, tl.minimum(31, exponent - EXPONENT_BIAS - MANTISSA_BITS + 15)) + # Add the target dtype's exponent bias and clamp to the target dtype's exponent range. + exponent = fp16_exponent + EXPONENT_BIAS + MAX_EXPONENT = 2 * EXPONENT_BIAS + 1 + exponent = tl.where(exponent > MAX_EXPONENT, MAX_EXPONENT, exponent) + exponent = tl.where(exponent < 0, 0, exponent) - # Calculate epsilon in the target dtype - eps_normal = tl.cast(tl.cast(eps_exp << 10, tl.int16), tl.float16, bitcast=True) + # Normal ULP exponent, expressed as an fp16 exponent field: + # (exponent - EXPONENT_BIAS - MANTISSA_BITS) + 15 + # Simplifies to: fp16_exponent - MANTISSA_BITS + 15 + # See https://en.wikipedia.org/wiki/Unit_in_the_last_place + eps_exp = tl.maximum(0, tl.minimum(31, exponent - EXPONENT_BIAS - MANTISSA_BITS + 15)) - # Subnormal ULP: 2^(1 - EXPONENT_BIAS - MANTISSA_BITS) -> - # fp16 exponent bits: (1 - EXPONENT_BIAS - MANTISSA_BITS) + 15 = - # 16 - EXPONENT_BIAS - MANTISSA_BITS - eps_subnormal = tl.cast(tl.cast((16 - EXPONENT_BIAS - MANTISSA_BITS) << 10, tl.int16), tl.float16, bitcast=True) - eps = tl.where(exponent > 0, eps_normal, eps_subnormal) + # Calculate epsilon in the target dtype + eps_normal = tl.cast(tl.cast(eps_exp << 10, tl.int16), tl.float16, bitcast=True) - # Apply zero mask to epsilon - eps = tl.where(x == 0, 0.0, eps) + # Subnormal ULP: 2^(1 - EXPONENT_BIAS - MANTISSA_BITS) -> + # fp16 exponent bits: (1 - EXPONENT_BIAS - MANTISSA_BITS) + 15 = + # 16 - EXPONENT_BIAS - MANTISSA_BITS + eps_subnormal = tl.cast(tl.cast((16 - EXPONENT_BIAS - MANTISSA_BITS) << 10, tl.int16), tl.float16, bitcast=True) + eps = tl.where(exponent > 0, eps_normal, eps_subnormal) - # Apply stochastic rounding - output = tl.cast(x + rand_vals * eps, tl.bfloat16) + # Apply zero mask to epsilon + eps = tl.where(x == 0, 0.0, eps) - # Store the result - tl.store(output_ptr + offsets, output, mask=mask) + # Apply stochastic rounding + output = tl.cast(x + rand_vals * eps, tl.bfloat16) + + # Store the result + tl.store(output_ptr + offsets, output, mask=mask) diff --git a/packages/ltx-core/src/ltx_core/loader/primitives.py b/packages/ltx-core/src/ltx_core/loader/primitives.py index 8a918a1..aa07635 100644 --- a/packages/ltx-core/src/ltx_core/loader/primitives.py +++ b/packages/ltx-core/src/ltx_core/loader/primitives.py @@ -13,6 +13,10 @@ if TYPE_CHECKING: from ltx_core.loader.registry import Registry +# Per-key shape and dtype description for a flat collection of tensors. +TensorLayout = dict[str, tuple[torch.Size, torch.dtype]] + + @dataclass(frozen=True) class StateDict: """ @@ -52,7 +56,15 @@ class StateDictLoader(Protocol): """ -class ModelBuilderProtocol(Protocol[ModelType]): +class BuilderProtocol(Protocol[ModelType]): + """Protocol for model builders that produce a model via ``build()``.""" + + def build( + self, device: torch.device | None = None, dtype: torch.dtype | None = None, **kwargs: object + ) -> ModelType: ... + + +class ModelBuilderProtocol(BuilderProtocol[ModelType], Protocol[ModelType]): """ Protocol for building PyTorch models from configuration dictionaries. Implementations must provide: 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 f98fee1..7c0ab21 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 @@ -102,7 +102,7 @@ class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType], registry: Registry = field(default_factory=DummyRegistry) lora_load_device: torch.device = field(default_factory=lambda: torch.device("cpu")) - def lora(self, lora_path: str, strength: float = 1.0, sd_ops: SDOps | None = None) -> "SingleGPUModelBuilder": + def lora(self, lora_path: str, strength: float, sd_ops: SDOps) -> "SingleGPUModelBuilder": return replace(self, loras=(*self.loras, LoraPathStrengthAndSDOps(lora_path, strength, sd_ops))) def with_sd_ops(self, sd_ops: SDOps | None) -> "SingleGPUModelBuilder": diff --git a/packages/ltx-core/src/ltx_core/model/transformer/attention.py b/packages/ltx-core/src/ltx_core/model/transformer/attention.py index a061344..23e0447 100644 --- a/packages/ltx-core/src/ltx_core/model/transformer/attention.py +++ b/packages/ltx-core/src/ltx_core/model/transformer/attention.py @@ -144,7 +144,7 @@ class Attention(torch.nn.Module): heads: int = 8, dim_head: int = 64, norm_eps: float = 1e-6, - rope_type: LTXRopeType = LTXRopeType.INTERLEAVED, + rope_type: LTXRopeType = LTXRopeType.SPLIT, attention_function: AttentionCallable | AttentionFunction = AttentionFunction.DEFAULT, apply_gated_attention: bool = False, ) -> None: diff --git a/packages/ltx-core/src/ltx_core/model/transformer/model.py b/packages/ltx-core/src/ltx_core/model/transformer/model.py index 4233bce..ab039f9 100644 --- a/packages/ltx-core/src/ltx_core/model/transformer/model.py +++ b/packages/ltx-core/src/ltx_core/model/transformer/model.py @@ -57,7 +57,7 @@ class LTXModel(torch.nn.Module): audio_cross_attention_dim: int = 2048, audio_positional_embedding_max_pos: list[int] | None = None, av_ca_timestep_scale_multiplier: int = 1, - rope_type: LTXRopeType = LTXRopeType.INTERLEAVED, + rope_type: LTXRopeType = LTXRopeType.SPLIT, double_precision_rope: bool = False, apply_gated_attention: bool = False, caption_projection: torch.nn.Module | None = None, diff --git a/packages/ltx-core/src/ltx_core/model/transformer/model_configurator.py b/packages/ltx-core/src/ltx_core/model/transformer/model_configurator.py index adbbce5..fb836d7 100644 --- a/packages/ltx-core/src/ltx_core/model/transformer/model_configurator.py +++ b/packages/ltx-core/src/ltx_core/model/transformer/model_configurator.py @@ -62,7 +62,7 @@ class LTXModelConfigurator(ModelConfigurator[LTXModel]): audio_cross_attention_dim=config.get("audio_cross_attention_dim", 2048), audio_positional_embedding_max_pos=config.get("audio_positional_embedding_max_pos", [20]), av_ca_timestep_scale_multiplier=config.get("av_ca_timestep_scale_multiplier", 1), - rope_type=LTXRopeType(config.get("rope_type", "interleaved")), + rope_type=LTXRopeType(config.get("rope_type", "split")), double_precision_rope=config.get("frequencies_precision", False) == "float64", apply_gated_attention=config.get("apply_gated_attention", False), caption_projection=caption_projection, @@ -114,7 +114,7 @@ class LTXVideoOnlyModelConfigurator(ModelConfigurator[LTXModel]): positional_embedding_max_pos=config.get("positional_embedding_max_pos", [20, 2048, 2048]), timestep_scale_multiplier=config.get("timestep_scale_multiplier", 1000), use_middle_indices_grid=config.get("use_middle_indices_grid", True), - rope_type=LTXRopeType(config.get("rope_type", "interleaved")), + rope_type=LTXRopeType(config.get("rope_type", "split")), double_precision_rope=config.get("frequencies_precision", False) == "float64", apply_gated_attention=config.get("apply_gated_attention", False), caption_projection=caption_projection, diff --git a/packages/ltx-core/src/ltx_core/model/transformer/rope.py b/packages/ltx-core/src/ltx_core/model/transformer/rope.py index 2ce58d9..7314b72 100644 --- a/packages/ltx-core/src/ltx_core/model/transformer/rope.py +++ b/packages/ltx-core/src/ltx_core/model/transformer/rope.py @@ -16,9 +16,10 @@ class LTXRopeType(Enum): def apply_rotary_emb( input_tensor: torch.Tensor, freqs_cis: Tuple[torch.Tensor, torch.Tensor], - rope_type: LTXRopeType = LTXRopeType.INTERLEAVED, + rope_type: LTXRopeType = LTXRopeType.SPLIT, ) -> torch.Tensor: if rope_type == LTXRopeType.INTERLEAVED: + # Note: INTERLEAVED rope is a legacy mode. Prefer SPLIT instead. return apply_interleaved_rotary_emb(input_tensor, *freqs_cis) elif rope_type == LTXRopeType.SPLIT: return apply_split_rotary_emb(input_tensor, *freqs_cis) @@ -45,6 +46,11 @@ def apply_split_rotary_emb( needs_reshape = False if input_tensor.ndim != 4 and cos_freqs.ndim == 4: b, h, t, _ = cos_freqs.shape + if input_tensor.shape[0] != b: + raise ValueError( + f"apply_split_rotary_emb: input_tensor batch ({input_tensor.shape[0]}) " + f"must equal cos_freqs batch ({b})." + ) input_tensor = input_tensor.reshape(b, t, h, -1).swapaxes(1, 2) needs_reshape = True @@ -183,7 +189,7 @@ def precompute_freqs_cis( max_pos: list[int] | None = None, use_middle_indices_grid: bool = False, num_attention_heads: int = 32, - rope_type: LTXRopeType = LTXRopeType.INTERLEAVED, + rope_type: LTXRopeType = LTXRopeType.SPLIT, freq_grid_generator: Callable[[float, int, int, torch.device], torch.Tensor] = generate_freq_grid_pytorch, ) -> tuple[torch.Tensor, torch.Tensor]: if max_pos is None: diff --git a/packages/ltx-core/src/ltx_core/model/transformer/transformer.py b/packages/ltx-core/src/ltx_core/model/transformer/transformer.py index af8b606..eb62885 100644 --- a/packages/ltx-core/src/ltx_core/model/transformer/transformer.py +++ b/packages/ltx-core/src/ltx_core/model/transformer/transformer.py @@ -27,7 +27,7 @@ class BasicAVTransformerBlock(torch.nn.Module): idx: int, video: TransformerConfig | None = None, audio: TransformerConfig | None = None, - rope_type: LTXRopeType = LTXRopeType.INTERLEAVED, + rope_type: LTXRopeType = LTXRopeType.SPLIT, norm_eps: float = 1e-6, attention_function: AttentionFunction | AttentionCallable = AttentionFunction.DEFAULT, ): diff --git a/packages/ltx-core/src/ltx_core/model/video_vae/__init__.py b/packages/ltx-core/src/ltx_core/model/video_vae/__init__.py index 122bb8f..45af025 100644 --- a/packages/ltx-core/src/ltx_core/model/video_vae/__init__.py +++ b/packages/ltx-core/src/ltx_core/model/video_vae/__init__.py @@ -1,5 +1,6 @@ """Video VAE package.""" +from ltx_core.model.video_vae.memory_efficient_decode import MEMORY_EFFICIENT_DECODE from ltx_core.model.video_vae.model_configurator import ( VAE_DECODER_COMFY_KEYS_FILTER, VAE_ENCODER_COMFY_KEYS_FILTER, @@ -10,6 +11,7 @@ from ltx_core.model.video_vae.tiling import SpatialTilingConfig, TemporalTilingC from ltx_core.model.video_vae.video_vae import VideoDecoder, VideoEncoder, get_video_chunks_number __all__ = [ + "MEMORY_EFFICIENT_DECODE", "VAE_DECODER_COMFY_KEYS_FILTER", "VAE_ENCODER_COMFY_KEYS_FILTER", "SpatialTilingConfig", diff --git a/packages/ltx-core/src/ltx_core/model/video_vae/memory_efficient_decode.py b/packages/ltx-core/src/ltx_core/model/video_vae/memory_efficient_decode.py new file mode 100644 index 0000000..4fe0846 --- /dev/null +++ b/packages/ltx-core/src/ltx_core/model/video_vae/memory_efficient_decode.py @@ -0,0 +1,612 @@ +"""Memory-efficient VAE decoder operations. +Reduces peak VRAM usage during video decoding through in-place operations +and workspace buffer reuse. The main optimizations are: +1. **Workspace buffers** -- Pre-allocated tensors with temporal padding replace + dynamic padding (``F.pad`` / ``concatenate``) in ``CausalConv3d``. A + workspace of shape ``[B, C, T+2, H, W]`` holds the data in positions + ``[1:-1]`` with replicate padding at ``[0]`` and ``[-1]``. +2. **In-place temporal-chunked Conv3d** *(non-causal only)* -- The convolution + output is written back into the workspace buffer, avoiding a separate + output allocation. Temporal chunking with boundary save/restore ensures + correct reads despite in-place writes. +3. **In-place normalization and affine transforms** -- PixelNorm, scale/shift, + and SiLU are applied in-place on workspace views. +4. **Free-before-conv** -- For ``DepthToSpaceUpsample`` blocks the input + tensor is freed before the convolution runs so that peak VRAM never holds + input *and* output simultaneously. +Both causal and non-causal modes are supported. Non-causal mode benefits +from all four optimizations. Causal mode benefits from optimizations 1, 3, +and 4; in-place conv (2) is skipped because the asymmetric causal padding +layout prevents clean in-place overwrites. +Usage via the ``ModuleOps`` pattern (preferred):: + from ltx_core.model.video_vae import MEMORY_EFFICIENT_DECODE + builder = decoder_builder.with_module_ops( + (*decoder_builder.module_ops, MEMORY_EFFICIENT_DECODE) + ) +Or applied directly to an existing decoder:: + from ltx_core.model.video_vae.memory_efficient_decode import ( + enable_memory_efficient_decode, + ) + enable_memory_efficient_decode(decoder) +""" + +from __future__ import annotations + +import math +from typing import TYPE_CHECKING + +import torch +from einops import rearrange +from torch import nn +from torch.nn import functional as F + +from ltx_core.loader.module_ops import ModuleOps +from ltx_core.model.common.normalization import PixelNorm +from ltx_core.model.video_vae.convolution import CausalConv3d +from ltx_core.model.video_vae.ops import unpatchify +from ltx_core.model.video_vae.resnet import ResnetBlock3D, UNetMidBlock3D +from ltx_core.model.video_vae.sampling import DepthToSpaceUpsample + +if TYPE_CHECKING: + from ltx_core.model.video_vae.video_vae import VideoDecoder + + +# --------------------------------------------------------------------------- +# Low-level helpers +# --------------------------------------------------------------------------- + + +def _find_temporal_split_size(num_frames: int) -> int: + """Find chunk size for in-place temporal convolution. + The chunk size ensures the last chunk has at least 3 frames + (the temporal kernel size), avoiding degenerate chunks. + """ + for s in range(16, 2, -1): + remainder = num_frames % s + if remainder == 0 or remainder >= 3: + return s + + raise ValueError( + f"Unable to find a valid temporal split size for num_frames={num_frames}. " + "Expected a split size between 3 and 16 such that the final chunk is " + "either exact or has at least 3 frames." + ) + + +def _pad_workspace_temporal(workspace: torch.Tensor) -> None: + """Apply non-causal replicate padding to temporal boundaries. + Sets ``workspace[:, :, 0]`` to a copy of ``workspace[:, :, 1]`` and + ``workspace[:, :, -1]`` to a copy of ``workspace[:, :, -2]``. + """ + workspace[:, :, 0, :, :].copy_(workspace[:, :, 1, :, :]) + workspace[:, :, -1, :, :].copy_(workspace[:, :, -2, :, :]) + + +# --------------------------------------------------------------------------- +# In-place Conv3d (non-causal only) +# --------------------------------------------------------------------------- + + +def inplace_conv3d_temporal_chunked(workspace: torch.Tensor, conv: nn.Conv3d) -> None: + """Run a 3x3x3 Conv3d in-place on a temporally-padded workspace. + The workspace has shape ``[B, C, T+2, H, W]`` where positions ``[1:-1]`` + hold the real data and positions ``[0]`` and ``[-1]`` are padding slots. + The convolution must have ``kernel_size=(3,3,3)``, ``stride=(1,1,1)``, + ``padding=(0,1,1)`` -- no temporal padding, symmetric spatial padding. + The output (T frames) overwrites positions ``[1:-1]``. Temporal chunking + with boundary save/restore ensures each chunk reads unmodified input even + though earlier chunks already wrote to the same buffer. + Only valid for **non-causal** mode (symmetric replicate padding). + Args: + workspace: Tensor ``[B, max(C_in, C_out), T+2, H, W]``. + Modified in-place; after the call ``workspace[:, :C_out, 1:-1]`` + holds the convolution result. + conv: ``nn.Conv3d`` with the constraints above. + """ + if conv.kernel_size != (3, 3, 3): + raise ValueError(f"Expected kernel_size=(3,3,3), got {conv.kernel_size}") + if conv.stride != (1, 1, 1): + raise ValueError(f"Expected stride=(1,1,1), got {conv.stride}") + if conv.padding != (0, 1, 1): + raise ValueError(f"Expected padding=(0,1,1), got {conv.padding}") + + _pad_workspace_temporal(workspace) + + total_frames = workspace.shape[2] + out_channels = conv.out_channels + in_channels = conv.in_channels + + if total_frames > 16: + split_size = _find_temporal_split_size(total_frames) + num_splits = (total_frames + split_size - 1) // split_size + else: + split_size = total_frames - 1 + num_splits = 1 + + # 1-frame buffers for saving / restoring boundary frames across chunks. + x_buf = torch.empty( + workspace.shape[0], + workspace.shape[1], + 1, + workspace.shape[3], + workspace.shape[4], + device=workspace.device, + dtype=workspace.dtype, + ) + o_buf = torch.empty_like(x_buf) + + # Helper: extract a chunk and make it contiguous. Workspace views can + # inherit strides > 2^31 from the full buffer, which makes Conv3d's + # reflect-padding path (F.pad) crash with "input tensor must fit into + # 32-bit index math". A small .clone() per chunk avoids this. + needs_clone = workspace.untyped_storage().nbytes() > (2**31 - 1) * workspace.element_size() + + def _chunk(t_start: int, t_end: int) -> torch.Tensor: + s = workspace[:, :in_channels, t_start:t_end] + return s.clone() if needs_clone else s + + # --- First chunk --- + if num_splits > 1: + # Save the boundary now so the loop below can restore it. Skipped + # when there is only one chunk: the loop never runs, and the save + # would be a wasted full HW slice copy. + x_buf[:, :, 0] = workspace[:, :, split_size - 1].clone() + workspace[:, :out_channels, 1:split_size] = conv(_chunk(0, split_size + 1)) + + # --- Remaining chunks --- + for i in range(1, num_splits): + start = i * split_size + end = min((i + 1) * split_size, total_frames - 1) + + # Save the value at start-1 (now holds previous chunk's output). + o_buf[:, :, 0] = workspace[:, :, start - 1].clone() + # Restore the original input value needed by this chunk's conv. + workspace[:, :, start - 1] = x_buf[:, :, 0] + # Save the boundary for the *next* chunk before we overwrite it. + x_buf[:, :, 0] = workspace[:, :, end - 1].clone() + + workspace[:, :out_channels, start:end] = conv(_chunk(start - 1, end + 1)) + + # Put back the previous chunk's output at the boundary. + workspace[:, :, start - 1] = o_buf[:, :, 0] + + +# --------------------------------------------------------------------------- +# Causal conv helper (free-before-conv) +# --------------------------------------------------------------------------- + + +def _causal_pad(x: torch.Tensor, pad_size: int) -> torch.Tensor: + """Build a causal-padded buffer of shape ``[B, C, T+pad_size, H, W]``. + Copies ``x`` into ``padded[:, :, pad_size:]`` and replicates the first + real frame into the leading ``pad_size`` slots. The caller still owns + ``x`` after this returns. + """ + padded = torch.empty( + x.shape[0], + x.shape[1], + x.shape[2] + pad_size, + x.shape[3], + x.shape[4], + device=x.device, + dtype=x.dtype, + ) + padded[:, :, pad_size:].copy_(x) + for i in range(pad_size): + padded[:, :, i] = padded[:, :, pad_size] + return padded + + +def _causal_pad_free_and_conv(x: torch.Tensor, causal_conv: CausalConv3d) -> torch.Tensor: + """Causal-pad *x*, free it, then run the raw ``nn.Conv3d``. + This avoids the peak where both the original and padded tensors are + live simultaneously (as happens inside ``CausalConv3d.forward``). + Args: + x: Input ``[B, C_in, T, H, W]``. **Deleted** inside this function; + the caller must not use it afterwards. + Returns: + Convolution output ``[B, C_out, T, H, W]``. + """ + padded = _causal_pad(x, causal_conv.time_kernel_size - 1) + del x + result = causal_conv.conv(padded) + del padded + return result + + +# --------------------------------------------------------------------------- +# In-place normalization +# --------------------------------------------------------------------------- + + +def _pixel_norm_inplace(x: torch.Tensor, eps: float = 1e-8) -> None: + """In-place RMS (pixel) normalization along the channel dimension.""" + rms = torch.sqrt(torch.mean(x**2, dim=1, keepdim=True) + eps) + x.div_(rms) + + +def _norm_inplace(norm: nn.Module, x: torch.Tensor) -> None: + """Apply *norm* in-place, using an optimised path for ``PixelNorm``.""" + if isinstance(norm, PixelNorm): + _pixel_norm_inplace(x, eps=norm.eps) + else: + # GroupNorm or other -- fall back to allocating a temporary. + result = norm(x) + x.copy_(result) + del result + + +# --------------------------------------------------------------------------- +# Per-block efficient forwards +# --------------------------------------------------------------------------- + + +def _resnet_block_forward_inplace( + resnet: ResnetBlock3D, + workspace: torch.Tensor, + causal: bool, + timestep: torch.Tensor | None, + generator: torch.Generator | None, +) -> None: + """Run a ``ResnetBlock3D`` in-place on a workspace buffer. + The workspace has shape ``[B, C, T+2, H, W]`` with real data in + ``[1:-1]``. After this call ``workspace[:, :, 1:-1]`` holds the + residual-branch output ``F(x)`` (without the skip connection -- + the caller adds it back to the hidden state). + Only valid when ``in_channels == out_channels`` (true for all + ``ResnetBlock3D`` instances inside a ``UNetMidBlock3D``). + """ + if resnet.in_channels != resnet.out_channels: + raise ValueError( + "In-place resnet forward requires in_channels == out_channels, " + f"got {resnet.in_channels} != {resnet.out_channels}" + ) + + interior = workspace[:, :, 1:-1] + + # --- norm1 + [ada scaling] + SiLU + conv1 --- + _norm_inplace(resnet.norm1, interior) + + if resnet.timestep_conditioning and timestep is not None: + ada = resnet.scale_shift_table[None, ..., None, None, None].to( + device=interior.device, dtype=interior.dtype + ) + timestep.reshape( + interior.shape[0], + 4, + -1, + timestep.shape[-3], + timestep.shape[-2], + timestep.shape[-1], + ) + shift1, scale1, shift2, scale2 = ada.unbind(dim=1) + interior.mul_(1 + scale1).add_(shift1) + + F.silu(interior, inplace=True) + + if causal: + result = resnet.conv1(interior, causal=True) + interior.copy_(result) + del result + else: + inplace_conv3d_temporal_chunked(workspace, resnet.conv1.conv) + + if resnet.inject_noise: + spatial_shape = interior.shape[-2:] + scale = resnet.per_channel_scale1.to(device=interior.device, dtype=interior.dtype) + noise = torch.randn(spatial_shape, device=interior.device, dtype=interior.dtype, generator=generator) + interior.add_((noise * scale)[None, :, None, ...]) + + # --- norm2 + [ada scaling] + SiLU + conv2 --- + _norm_inplace(resnet.norm2, interior) + + if resnet.timestep_conditioning and timestep is not None: + interior.mul_(1 + scale2).add_(shift2) # type: ignore[possibly-undefined] + + F.silu(interior, inplace=True) + # dropout is always 0.0 during inference -- skip. + + if causal: + result = resnet.conv2(interior, causal=True) + interior.copy_(result) + del result + else: + inplace_conv3d_temporal_chunked(workspace, resnet.conv2.conv) + + if resnet.inject_noise: + spatial_shape = interior.shape[-2:] + scale = resnet.per_channel_scale2.to(device=interior.device, dtype=interior.dtype) + noise = torch.randn(spatial_shape, device=interior.device, dtype=interior.dtype, generator=generator) + interior.add_((noise * scale)[None, :, None, ...]) + + +def _midblock_forward_efficient( + block: UNetMidBlock3D, + hidden_states: torch.Tensor, + causal: bool, + timestep: torch.Tensor | None, + generator: torch.Generator | None, +) -> torch.Tensor: + """Memory-efficient ``UNetMidBlock3D`` forward. + Allocates a single workspace buffer that is reused across all + ``ResnetBlock3D`` iterations. For each block the workspace is + populated with the current hidden state, processed in-place, and + the result is added back (residual connection). + """ + timestep_embed = None + if block.timestep_conditioning: + if timestep is None: + raise ValueError("'timestep' required when timestep_conditioning=True") + batch_size = hidden_states.shape[0] + timestep_embed = block.time_embedder( + timestep=timestep.flatten(), + hidden_dtype=hidden_states.dtype, + ) + timestep_embed = timestep_embed.view(batch_size, timestep_embed.shape[-1], 1, 1, 1) + + workspace = torch.empty( + hidden_states.shape[0], + hidden_states.shape[1], + hidden_states.shape[2] + 2, + hidden_states.shape[3], + hidden_states.shape[4], + device=hidden_states.device, + dtype=hidden_states.dtype, + ) + + for resnet in block.res_blocks: + workspace[:, :, 1:-1].copy_(hidden_states) + _resnet_block_forward_inplace(resnet, workspace, causal, timestep_embed, generator) + hidden_states.add_(workspace[:, :, 1:-1]) + + del workspace + return hidden_states + + +def _upsample_forward_efficient( + block: DepthToSpaceUpsample, + x: torch.Tensor, + causal: bool, +) -> torch.Tensor: + """Memory-efficient ``DepthToSpaceUpsample`` forward. + For non-causal mode the input is copied into a workspace and the + convolution runs in-place. For causal mode the input is manually + padded and freed before the convolution runs. Both paths avoid + the peak where input *and* output coexist. + """ + if block.residual: + x_in = rearrange( + x, + "b (c p1 p2 p3) d h w -> b c (d p1) (h p2) (w p3)", + p1=block.stride[0], + p2=block.stride[1], + p3=block.stride[2], + ) + num_repeat = math.prod(block.stride) // block.out_channels_reduction_factor + x_in = x_in.repeat(1, num_repeat, 1, 1, 1) + if block.stride[0] == 2: + x_in = x_in[:, :, 1:, :, :] + + conv = block.conv.conv # underlying nn.Conv3d inside CausalConv3d + in_channels = x.shape[1] + out_channels = conv.out_channels + + if causal: + x = _causal_pad_free_and_conv(x, block.conv) + else: + workspace = torch.empty( + x.shape[0], + max(in_channels, out_channels), + x.shape[2] + 2, + x.shape[3], + x.shape[4], + device=x.device, + dtype=x.dtype, + ) + workspace[:, :in_channels, 1:-1].copy_(x) + del x + inplace_conv3d_temporal_chunked(workspace, conv) + x = workspace[:, :out_channels, 1:-1].contiguous() + del workspace + + x = rearrange( + x, + "b (c p1 p2 p3) d h w -> b c (d p1) (h p2) (w p3)", + p1=block.stride[0], + p2=block.stride[1], + p3=block.stride[2], + ) + if block.stride[0] == 2: + x = x[:, :, 1:, :, :] + if block.residual: + x = x + x_in + del x_in + return x + + +# --------------------------------------------------------------------------- +# Final norm + conv_out +# --------------------------------------------------------------------------- + + +def _final_norm_and_conv_out( + decoder: VideoDecoder, + sample: torch.Tensor, + causal: bool, + scaled_timestep: torch.Tensor | None, + batch_size: int, +) -> torch.Tensor: + """Workspace-based final norm + [ada] + SiLU + conv_out + unpatchify.""" + conv_out_mod: CausalConv3d = decoder.conv_out # type: ignore[assignment] + conv_out = conv_out_mod.conv + feature_channels = sample.shape[1] + + workspace = torch.empty( + sample.shape[0], + max(feature_channels, conv_out.out_channels), + sample.shape[2] + 2, + sample.shape[3], + sample.shape[4], + device=sample.device, + dtype=sample.dtype, + ) + workspace[:, :feature_channels, 1:-1].copy_(sample) + del sample + + interior = workspace[:, :feature_channels, 1:-1] + _norm_inplace(decoder.conv_norm_out, interior) + + if decoder.timestep_conditioning: + embedded_timestep = decoder.last_time_embedder( + timestep=scaled_timestep.flatten(), + hidden_dtype=interior.dtype, + ) + embedded_timestep = embedded_timestep.view(batch_size, embedded_timestep.shape[-1], 1, 1, 1) + ada_values = decoder.last_scale_shift_table[None, ..., None, None, None].to( + device=interior.device, dtype=interior.dtype + ) + embedded_timestep.reshape( + batch_size, + 2, + -1, + embedded_timestep.shape[-3], + embedded_timestep.shape[-2], + embedded_timestep.shape[-1], + ) + shift, scale = ada_values.unbind(dim=1) + interior.mul_(1 + scale).add_(shift) + + F.silu(interior, inplace=True) + + if causal: + # Causal: build padded tensor directly from the interior view, + # then free the workspace before running the conv. + padded = _causal_pad(interior, conv_out_mod.time_kernel_size - 1) + del workspace, interior + result = conv_out(padded) + del padded + else: + inplace_conv3d_temporal_chunked(workspace, conv_out) + result = workspace[:, : conv_out.out_channels, 1:-1].contiguous() + del workspace, interior + + return unpatchify(result, patch_size_hw=decoder.patch_size, patch_size_t=1) + + +# --------------------------------------------------------------------------- +# Top-level efficient decoder forward +# --------------------------------------------------------------------------- + + +def _memory_efficient_forward( + decoder: VideoDecoder, + sample: torch.Tensor, + timestep: torch.Tensor | None = None, + generator: torch.Generator | None = None, +) -> torch.Tensor: + """Full memory-efficient ``VideoDecoder.forward`` replacement. + Orchestrates the entire decode through workspace-based operations: + ``UNetMidBlock3D`` and ``DepthToSpaceUpsample`` blocks use efficient + paths; standalone ``ResnetBlock3D`` blocks fall back to the standard + forward. The final norm + ada + SiLU + conv_out is also workspace-based. + """ + causal = decoder.causal + batch_size = sample.shape[0] + sample = sample.to(next(decoder.parameters()).dtype) + + # --- Noise injection and de-normalisation (identical to standard path) --- + if decoder.timestep_conditioning: + noise = ( + torch.randn(sample.size(), generator=generator, dtype=sample.dtype, device=sample.device) + * decoder.decode_noise_scale + ) + sample = noise + (1.0 - decoder.decode_noise_scale) * sample + + sample = decoder.per_channel_statistics.un_normalize(sample) + + if timestep is None and decoder.timestep_conditioning: + timestep = torch.full((batch_size,), decoder.decode_timestep, device=sample.device, dtype=sample.dtype) + + # --- conv_in (latent tensor is small -- standard path is fine) --- + sample = decoder.conv_in(sample, causal=causal) + + upscale_dtype = next(iter(decoder.up_blocks.parameters())).dtype + sample = sample.to(upscale_dtype) + + scaled_timestep = None + if decoder.timestep_conditioning: + if timestep is None: + raise ValueError("'timestep' required when timestep_conditioning=True") + scaled_timestep = timestep * decoder.timestep_scale_multiplier.to(sample) + + # --- Up blocks (dispatch to efficient path per block type) --- + for up_block in decoder.up_blocks: + if isinstance(up_block, UNetMidBlock3D): + sample = _midblock_forward_efficient( + up_block, + sample, + causal=causal, + timestep=scaled_timestep if decoder.timestep_conditioning else None, + generator=generator, + ) + elif isinstance(up_block, DepthToSpaceUpsample): + sample = _upsample_forward_efficient(up_block, sample, causal=causal) + elif isinstance(up_block, ResnetBlock3D): + sample = up_block(sample, causal=causal, generator=generator) + else: + sample = up_block(sample, causal=causal) + + return _final_norm_and_conv_out(decoder, sample, causal, scaled_timestep, batch_size) + + +# --------------------------------------------------------------------------- +# Public API +# --------------------------------------------------------------------------- + + +def enable_memory_efficient_decode(decoder: nn.Module) -> nn.Module: + """Patch a ``VideoDecoder`` to use the memory-efficient forward path. + The original ``forward`` is saved as ``decoder._original_forward`` so + that it can be restored later with :func:`disable_memory_efficient_decode`. + """ + # Import here to avoid circular dependency at module level. + from ltx_core.model.video_vae.video_vae import VideoDecoder # noqa: PLC0415 + + if not isinstance(decoder, VideoDecoder): + raise TypeError(f"Expected VideoDecoder, got {type(decoder).__name__}") + + if hasattr(decoder, "_original_forward"): + return decoder + + original_forward = decoder.forward + + def efficient_forward( + sample: torch.Tensor, + timestep: torch.Tensor | None = None, + generator: torch.Generator | None = None, + ) -> torch.Tensor: + return _memory_efficient_forward(decoder, sample, timestep, generator) + + decoder._original_forward = original_forward # type: ignore[attr-defined] + decoder.forward = efficient_forward # type: ignore[assignment] + return decoder + + +def disable_memory_efficient_decode(decoder: nn.Module) -> nn.Module: + """Restore the original ``forward`` method on a patched ``VideoDecoder``.""" + if hasattr(decoder, "_original_forward"): + decoder.forward = decoder._original_forward # type: ignore[attr-defined] + del decoder._original_forward # type: ignore[attr-defined] + return decoder + + +def _is_video_decoder(model: nn.Module) -> bool: + """Matcher for the ``MEMORY_EFFICIENT_DECODE`` module op.""" + from ltx_core.model.video_vae.video_vae import VideoDecoder # noqa: PLC0415 + + return isinstance(model, VideoDecoder) + + +MEMORY_EFFICIENT_DECODE = ModuleOps( + name="memory_efficient_vae_decode", + matcher=_is_video_decoder, + mutator=enable_memory_efficient_decode, +) diff --git a/packages/ltx-core/src/ltx_core/model/video_vae/tiling.py b/packages/ltx-core/src/ltx_core/model/video_vae/tiling.py index b32fc59..d6923f8 100644 --- a/packages/ltx-core/src/ltx_core/model/video_vae/tiling.py +++ b/packages/ltx-core/src/ltx_core/model/video_vae/tiling.py @@ -64,6 +64,6 @@ class TilingConfig: @classmethod def default(cls) -> "TilingConfig": return cls( - spatial_config=SpatialTilingConfig(tile_size_in_pixels=512, tile_overlap_in_pixels=64), - temporal_config=TemporalTilingConfig(tile_size_in_frames=64, tile_overlap_in_frames=24), + spatial_config=SpatialTilingConfig(tile_size_in_pixels=768, tile_overlap_in_pixels=64), + temporal_config=TemporalTilingConfig(tile_size_in_frames=80, tile_overlap_in_frames=24), ) 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 983178c..f75dbcc 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 @@ -259,6 +259,7 @@ class VideoEncoder(nn.Module): Args: sample: Input video (B, C, F, H, W). F should be 1 + 8*k (e.g., 1, 9, 17, 25, 33...). If not, the encoder crops the last frames to the nearest valid length. + Should be normalized to [-1, 1] range before encoding. Returns: Normalized latent means (B, 128, F', H', W') where F' = 1+(F-1)/8, H' = H/32, W' = W/32. Example: (B, 3, 33, 512, 512) -> (B, 128, 5, 16, 16). @@ -605,8 +606,8 @@ class VideoDecoder(nn.Module): # many video frames and pixels correspond to a single latent cell. self.video_downscale_factors = SpatioTemporalScaleFactors( time=8, - width=32, height=32, + width=32, ) self.patch_size = patch_size @@ -905,34 +906,22 @@ class VideoDecoder(nn.Module): latent: torch.Tensor, tiling_config: TilingConfig | None = None, generator: torch.Generator | None = None, - *, - output_dtype: torch.dtype = torch.uint8, ) -> Iterator[torch.Tensor]: - """Decode a video latent tensor, yielding chunks ``[f, h, w, c]``. + """Decode a video latent tensor, yielding float chunks ``[f, h, w, c]`` in ``[0, 1]``. Subclasses (e.g. ``DistributedVideoDecoder``) may override this to control eagerness or distribution across ranks. - Args: - output_dtype: Target dtype for output tensors. ``torch.uint8`` - (default) maps the decoder's ``[-1, 1]`` output to - ``[0, 255]``. Any floating dtype returns ``[0, 1]`` cast - to that dtype. """ - def _convert(frames: torch.Tensor) -> torch.Tensor: - # rearrange materializes a new contiguous tensor for this permutation, - # so in-place ops below do not mutate the caller's data. + def to_rgb(frames: torch.Tensor) -> torch.Tensor: video = rearrange(frames[0], "c f h w -> f h w c") - video.add_(1.0).mul_(0.5).clamp_(0.0, 1.0) - if output_dtype == torch.uint8: - return video.mul_(255.0).to(torch.uint8) - return video.to(output_dtype) + return video.add_(1.0).mul_(0.5).clamp_(0.0, 1.0) if tiling_config is not None: for frames in self.tiled_decode(latent, tiling_config, generator=generator): - yield _convert(frames) + yield to_rgb(frames) else: decoded = self(latent, generator=generator) - yield _convert(decoded) + yield to_rgb(decoded) def _group_tiles_by_temporal_slice(self, tiles: List[Tile]) -> List[List[Tile]]: """Group tiles by their temporal output slice.""" diff --git a/packages/ltx-core/src/ltx_core/quantization/__init__.py b/packages/ltx-core/src/ltx_core/quantization/__init__.py index 8910327..e23ff9f 100644 --- a/packages/ltx-core/src/ltx_core/quantization/__init__.py +++ b/packages/ltx-core/src/ltx_core/quantization/__init__.py @@ -3,12 +3,9 @@ from ltx_core.quantization.fp8_cast import ( UPCAST_DURING_INFERENCE, UpcastWithStochasticRounding, ) -from ltx_core.quantization.fp8_scaled_mm import FP8_PREPARE_MODULE_OPS, FP8_TRANSPOSE_SD_OPS from ltx_core.quantization.policy import QuantizationPolicy __all__ = [ - "FP8_PREPARE_MODULE_OPS", - "FP8_TRANSPOSE_SD_OPS", "TRANSFORMER_LINEAR_DOWNCAST_MAP", "UPCAST_DURING_INFERENCE", "QuantizationPolicy", diff --git a/packages/ltx-core/src/ltx_core/quantization/fp8_cast.py b/packages/ltx-core/src/ltx_core/quantization/fp8_cast.py index c48a8a8..8359bb0 100644 --- a/packages/ltx-core/src/ltx_core/quantization/fp8_cast.py +++ b/packages/ltx-core/src/ltx_core/quantization/fp8_cast.py @@ -1,5 +1,6 @@ import torch +from ltx_core.loader.kernels import TRITON_AVAILABLE from ltx_core.loader.module_ops import ModuleOps from ltx_core.loader.sd_ops import KeyValueOperationResult, SDOps from ltx_core.model.transformer.model import LTXModel @@ -7,8 +8,13 @@ from ltx_core.model.transformer.model import LTXModel BLOCK_SIZE = 1024 -def _fused_add_round_launch(target_weight: torch.Tensor, original_weight: torch.Tensor, seed: int) -> torch.Tensor: - # Lazy import triton - only available on CUDA platforms +def fused_add_round_launch(target_weight: torch.Tensor, original_weight: torch.Tensor, seed: int) -> torch.Tensor: + if not TRITON_AVAILABLE: + raise RuntimeError( + "fused_add_round_launch requires Triton, which is not available on this platform. " + "Callers should gate on ltx_core.loader.kernels.TRITON_AVAILABLE and use a " + "deterministic-rounding fallback instead." + ) import triton # noqa: PLC0415 from ltx_core.loader.kernels import fused_add_round_kernel # noqa: PLC0415 @@ -53,10 +59,13 @@ def _upcast_and_round( """ Upcast the weight to the given dtype and optionally apply stochastic rounding. Input weight needs to have float8_e4m3fn or float8_e5m2 dtype. + Stochastic rounding is implemented via a Triton kernel. When Triton is not + available (e.g., on Windows), this falls back to deterministic (nearest) + rounding via ``weight.to(dtype)``. """ - if not with_stochastic_rounding: + if not with_stochastic_rounding or not TRITON_AVAILABLE or weight.device.type != "cuda": return weight.to(dtype) - return _fused_add_round_launch(torch.zeros_like(weight, dtype=dtype), weight, seed) + return fused_add_round_launch(torch.zeros_like(weight, dtype=dtype), weight, seed) class Fp8CastLinear(torch.nn.Linear): @@ -82,14 +91,26 @@ class Fp8CastLinear(torch.nn.Linear): def _replace_fwd_with_upcast(layer: torch.nn.Linear, with_stochastic_rounding: bool = False, seed: int = 0) -> None: """ Intended to be applied via __class__ reassignment to existing nn.Linear - instances so that their parameter and buffer tensors are preserved in-place, - avoiding re-instantiation. Forward remains defined at the class level, which - is required for torch.compile compatibility — instance-level closure - monkey-patches cause graph breaks. + instances. Forward remains defined at the class level, which is required for + torch.compile compatibility — instance-level closure monkey-patches cause + graph breaks. + Also retypes ``weight`` and ``bias`` to fp8 so the meta param dtype matches + the post-load tensor dtype (sd_ops downcasts checkpoint bf16 -> fp8 at load). + Block streaming relies on this to derive pool buffer layout from the meta + model without an eager checkpoint read. """ layer.__class__ = Fp8CastLinear layer._with_stochastic_rounding = with_stochastic_rounding layer._seed = seed + layer.weight = torch.nn.Parameter( + torch.empty(layer.weight.shape, dtype=torch.float8_e4m3fn, device=layer.weight.device), + requires_grad=layer.weight.requires_grad, + ) + if layer.bias is not None: + layer.bias = torch.nn.Parameter( + torch.empty(layer.bias.shape, dtype=torch.float8_e4m3fn, device=layer.bias.device), + requires_grad=layer.bias.requires_grad, + ) def _amend_forward_with_upcast( diff --git a/packages/ltx-core/src/ltx_core/quantization/fp8_scaled_mm.py b/packages/ltx-core/src/ltx_core/quantization/fp8_scaled_mm.py index 5f53ac4..3453776 100644 --- a/packages/ltx-core/src/ltx_core/quantization/fp8_scaled_mm.py +++ b/packages/ltx-core/src/ltx_core/quantization/fp8_scaled_mm.py @@ -1,11 +1,21 @@ +import json +import struct from typing import Callable import torch from torch import nn from ltx_core.loader.module_ops import ModuleOps -from ltx_core.loader.sd_ops import KeyValueOperationResult, SDOps from ltx_core.model.transformer import LTXModel +from ltx_core.quantization.trtllm_scaled_usable import trtllm_scaled_mm_usable + + +def _read_safetensors_dtypes(path: str) -> dict[str, str]: + """Return ``{tensor_name: dtype_string}`` from the safetensors header.""" + with open(path, "rb") as f: + header_size = struct.unpack(" torch.Tensor: origin_shape = x.shape - # Static quantization: use pre-computed scale - qinput, cur_input_scale = torch.ops.tensorrt_llm.static_quantize_e4m3_per_tensor(x, self.input_scale) + if trtllm_scaled_mm_usable(): + qinput, cur_input_scale = torch.ops.tensorrt_llm.static_quantize_e4m3_per_tensor(x, self.input_scale) + if qinput.dim() == 3: + qinput = qinput.reshape(-1, qinput.shape[-1]) + output = torch.ops.trtllm.cublas_scaled_mm( + qinput, + self.weight.t(), + scale_a=cur_input_scale, + scale_b=self.weight_scale, + bias=None, + out_dtype=x.dtype, + ) + else: + # Clamp before cast: out-of-range values cast to NaN/saturated FP8, which + # produces black-screen output on some checkpoints (e.g. ltx-2-19b-dev-fp8). + fp8_min = torch.finfo(torch.float8_e4m3fn).min + fp8_max = torch.finfo(torch.float8_e4m3fn).max + qinput = torch.clamp(x * self.input_scale.reciprocal(), fp8_min, fp8_max).to(torch.float8_e4m3fn) + if qinput.dim() == 3: + qinput = qinput.reshape(-1, qinput.shape[-1]) + output = torch._scaled_mm( + qinput, + self.weight.t(), + scale_a=self.input_scale, + scale_b=self.weight_scale, + out_dtype=x.dtype, + use_fast_accum=True, + ) - # Flatten to 2D for matmul - if qinput.dim() == 3: - qinput = qinput.reshape(-1, qinput.shape[-1]) - - # FP8 scaled matmul - output = torch.ops.trtllm.cublas_scaled_mm( - qinput, - self.weight, - scale_a=cur_input_scale, - scale_b=self.weight_scale, - bias=None, - out_dtype=x.dtype, - ) - - # Add bias if self.bias is not None: - bias = self.bias - if bias.dtype != output.dtype: - bias = bias.to(output.dtype) - output = output + bias + output = output + self.bias.to(output.dtype) - # Restore original shape if output.dim() != len(origin_shape): output_shape = list(origin_shape) output_shape[-1] = output.shape[-1] @@ -74,15 +88,7 @@ class FP8Linear(nn.Module): def quantize_weight_to_fp8_per_tensor(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - """ - Quantize a weight tensor to FP8 (float8_e4m3fn) using per-tensor scaling. - Args: - weight: The weight tensor to quantize (any dtype, will be cast to float32) - Returns: - Tuple of (quantized_weight, weight_scale): - - quantized_weight: FP8 tensor, transposed for cublas_scaled_mm - - weight_scale: Per-tensor scale factor (reciprocal of quantization scale) - """ + """Quantize a weight tensor to ``float8_e4m3fn`` with a per-tensor scale.""" weight_fp32 = weight.to(torch.float32) fp8_min = torch.finfo(torch.float8_e4m3fn).min @@ -96,7 +102,6 @@ def quantize_weight_to_fp8_per_tensor(weight: torch.Tensor) -> tuple[torch.Tenso weight_fp32: torch.Tensor, scale: torch.Tensor, fp8_min: torch.Tensor, fp8_max: torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor]: quantized_weight = torch.clamp(weight_fp32 * scale, min=fp8_min, max=fp8_max).to(torch.float8_e4m3fn) - quantized_weight = quantized_weight.t() weight_scale = scale.reciprocal() return quantized_weight, weight_scale @@ -104,36 +109,8 @@ def quantize_weight_to_fp8_per_tensor(weight: torch.Tensor) -> tuple[torch.Tenso return quantized_weight, weight_scale -def _should_skip_layer(layer_name: str, excluded_layer_substrings: tuple[str, ...]) -> bool: - return any(substring in layer_name for substring in excluded_layer_substrings) - - -EXCLUDED_LAYER_SUBSTRINGS = ( - "patchify_proj", - "adaln_single", - "av_ca_video_scale_shift_adaln_single", - "av_ca_a2v_gate_adaln_single", - "caption_projection", - "proj_out", - "audio_patchify_proj", - "audio_adaln_single", - "av_ca_audio_scale_shift_adaln_single", - "av_ca_v2a_gate_adaln_single", - "audio_caption_projection", - "audio_proj_out", - "transformer_blocks.0.", - *[f"transformer_blocks.{i}." for i in range(43, 48)], -) - - def _linear_to_fp8linear(layer: nn.Linear) -> FP8Linear: - """ - Create an FP8Linear layer from an nn.Linear layer. - Args: - layer: The nn.Linear layer to convert (typically on meta device) - Returns: - A new FP8Linear with the same configuration - """ + """Create an ``FP8Linear`` matching the shape/bias of *layer*.""" return FP8Linear( in_features=layer.in_features, out_features=layer.out_features, @@ -142,15 +119,14 @@ def _linear_to_fp8linear(layer: nn.Linear) -> FP8Linear: ) -def _apply_fp8_prepare_to_model(model: nn.Module, excluded_layer_substrings: tuple[str, ...]) -> nn.Module: - """Replace nn.Linear layers with FP8Linear in the module tree.""" +def _swap_linears_to_fp8(model: nn.Module, should_swap: Callable[[str], bool]) -> nn.Module: + """Replace nn.Linear layers with FP8Linear where ``should_swap(name)`` returns True.""" replacements: list[tuple[nn.Module, str, nn.Linear]] = [] for name, module in model.named_modules(): if not isinstance(module, nn.Linear) or isinstance(module, FP8Linear): continue - - if _should_skip_layer(name, excluded_layer_substrings): + if not should_swap(name): continue if "." in name: @@ -168,40 +144,32 @@ def _apply_fp8_prepare_to_model(model: nn.Module, excluded_layer_substrings: tup return model -def _create_transpose_kv_operation( - excluded_layer_substrings: tuple[str, ...], -) -> Callable[[str, torch.Tensor], list[KeyValueOperationResult]]: - def transpose_if_matches(key: str, value: torch.Tensor) -> list[KeyValueOperationResult]: - # Only process .weight keys - if not key.endswith(".weight"): - return [KeyValueOperationResult(key, value)] +def get_fp8_swap_module_ops(checkpoint_path: str) -> tuple[ModuleOps, ...]: + """Return the FP8 swap ``ModuleOps`` for layers whose ``.weight`` is ``F8_E4M3`` + and which have a sibling ``.weight_scale`` tensor in the checkpoint. + Raises ``ValueError`` if no such layers are found — that combination is ambiguous + (a BF16 checkpoint with this policy would load as a no-op). + """ + dtypes = _read_safetensors_dtypes(checkpoint_path) + fp8_scale_paths = frozenset( + key.removesuffix(".weight_scale") + for key in dtypes + if key.endswith(".weight_scale") and dtypes.get(key.removesuffix(".weight_scale") + ".weight") == "F8_E4M3" + ) + if not fp8_scale_paths: + raise ValueError( + f"fp8_scaled_mm requires a pre-quantized checkpoint with F8_E4M3 .weight + .weight_scale " + f"tensors, but {checkpoint_path!r} has none. Use QuantizationPolicy.fp8_cast() for BF16 checkpoints." + ) - # Only transpose 2D FP8 tensors (Linear weights) - if value.dim() != 2 or value.dtype != torch.float8_e4m3fn: - return [KeyValueOperationResult(key, value)] + def _should_swap(name: str) -> bool: + suffix = "." + name + return any(p == name or p.endswith(suffix) for p in fp8_scale_paths) - # Check if the layer is excluded - layer_name = key.rsplit(".weight", 1)[0] - if _should_skip_layer(layer_name, excluded_layer_substrings): - return [KeyValueOperationResult(key, value)] - - # Transpose to cuBLAS layout (in, out) - transposed_weight = value.t() - - return [KeyValueOperationResult(key, transposed_weight)] - - return transpose_if_matches - - -FP8_TRANSPOSE_SD_OPS = SDOps("fp8_transpose_weights").with_kv_operation( - _create_transpose_kv_operation(EXCLUDED_LAYER_SUBSTRINGS), - key_prefix="transformer_blocks.", - key_suffix=".weight", -) - - -FP8_PREPARE_MODULE_OPS = ModuleOps( - name="fp8_prepare_for_loading", - matcher=lambda model: isinstance(model, LTXModel), - mutator=lambda model: _apply_fp8_prepare_to_model(model, EXCLUDED_LAYER_SUBSTRINGS), -) + return ( + ModuleOps( + name="fp8_swap_linears", + matcher=lambda model: isinstance(model, LTXModel), + mutator=lambda model: _swap_linears_to_fp8(model, _should_swap), + ), + ) diff --git a/packages/ltx-core/src/ltx_core/quantization/policy.py b/packages/ltx-core/src/ltx_core/quantization/policy.py index 56dad51..72e7d15 100644 --- a/packages/ltx-core/src/ltx_core/quantization/policy.py +++ b/packages/ltx-core/src/ltx_core/quantization/policy.py @@ -1,39 +1,48 @@ from dataclasses import dataclass +from enum import Enum from ltx_core.loader.module_ops import ModuleOps from ltx_core.loader.sd_ops import SDOps from ltx_core.quantization.fp8_cast import TRANSFORMER_LINEAR_DOWNCAST_MAP, UPCAST_DURING_INFERENCE -from ltx_core.quantization.fp8_scaled_mm import FP8_PREPARE_MODULE_OPS, FP8_TRANSPOSE_SD_OPS +from ltx_core.quantization.fp8_scaled_mm import get_fp8_swap_module_ops @dataclass(frozen=True) class QuantizationPolicy: """Configuration for model quantization during loading. Attributes: - sd_ops: State dict operations for weight transformation. - module_ops: Post-load module transformations. + kind: Discriminator for the policy variant. + sd_ops: State-dict operations applied to each tensor during load. + module_ops: Post-load module transformations applied to the meta model. """ + class Kind(str, Enum): + FP8_CAST = "fp8_cast" + FP8_SCALED_MM = "fp8_scaled_mm" + + kind: Kind sd_ops: SDOps | None = None module_ops: tuple[ModuleOps, ...] = () @classmethod def fp8_cast(cls) -> "QuantizationPolicy": - """Create policy using FP8 casting with upcasting during inference.""" + """FP8 casting with upcasting during inference.""" return cls( + kind=cls.Kind.FP8_CAST, sd_ops=TRANSFORMER_LINEAR_DOWNCAST_MAP, module_ops=(UPCAST_DURING_INFERENCE,), ) @classmethod - def fp8_scaled_mm(cls) -> "QuantizationPolicy": - """Create policy using FP8 scaled matrix multiplication.""" - try: - import tensorrt_llm # noqa: F401, PLC0415 - except ImportError as e: - raise ImportError("tensorrt_llm is not installed, skipping FP8 scaled MM quantization") from e - + def fp8_scaled_mm(cls, checkpoint_path: str) -> "QuantizationPolicy": + """FP8 scaled matmul for checkpoints pre-quantized with per-tensor scales. + The set of layers to swap to ``FP8Linear`` is discovered from the + checkpoint's ``.weight_scale`` tensors via suffix-matching against the + model's named modules. Requires a pre-quantized checkpoint; for BF16 + checkpoints, use :meth:`fp8_cast` instead. + """ return cls( - sd_ops=FP8_TRANSPOSE_SD_OPS, - module_ops=(FP8_PREPARE_MODULE_OPS,), + kind=cls.Kind.FP8_SCALED_MM, + sd_ops=None, + module_ops=get_fp8_swap_module_ops(checkpoint_path), ) diff --git a/packages/ltx-core/src/ltx_core/quantization/trtllm_scaled_usable.py b/packages/ltx-core/src/ltx_core/quantization/trtllm_scaled_usable.py new file mode 100644 index 0000000..deb69d0 --- /dev/null +++ b/packages/ltx-core/src/ltx_core/quantization/trtllm_scaled_usable.py @@ -0,0 +1,37 @@ +"""Runtime detection of TensorRT-LLM FP8 scaled-matmul availability. +When the TRT-LLM ops are usable on the current host (Linux + Hopper-class CUDA ++ tensorrt_llm wheel installed) we use them since they outperform the PyTorch-native +``torch._scaled_mm`` path. Otherwise we fall back to the native implementation, +which is portable across platforms (Windows, macOS, AMD GPUs). +The check runs once and is cached. +""" + +from __future__ import annotations + +import platform +from functools import cache + +import torch + + +@cache +def trtllm_scaled_mm_usable() -> bool: + if platform.system() != "Linux": + return False + + if not torch.cuda.is_available(): + return False + + major, minor = torch.cuda.get_device_capability() + sm = major * 10 + minor + + if sm < 90 or sm >= 120: + return False + + # The import is load-bearing — registers the trtllm torch ops as a side effect. + try: + import tensorrt_llm # noqa: F401, PLC0415 + except Exception: + return False + + return True diff --git a/packages/ltx-core/src/ltx_core/text_encoders/gemma/embeddings_connector.py b/packages/ltx-core/src/ltx_core/text_encoders/gemma/embeddings_connector.py index 2e614d5..c4e9d0c 100644 --- a/packages/ltx-core/src/ltx_core/text_encoders/gemma/embeddings_connector.py +++ b/packages/ltx-core/src/ltx_core/text_encoders/gemma/embeddings_connector.py @@ -18,7 +18,7 @@ class _BasicTransformerBlock1D(torch.nn.Module): dim: int, heads: int, dim_head: int, - rope_type: LTXRopeType = LTXRopeType.INTERLEAVED, + rope_type: LTXRopeType = LTXRopeType.SPLIT, apply_gated_attention: bool = False, ): super().__init__() @@ -39,7 +39,7 @@ class _BasicTransformerBlock1D(torch.nn.Module): def forward( self, hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, + additive_attention_mask: torch.Tensor | None = None, pe: torch.Tensor | None = None, ) -> torch.Tensor: # Notice that normalization is always applied before the real computation in the following blocks. @@ -49,8 +49,8 @@ class _BasicTransformerBlock1D(torch.nn.Module): norm_hidden_states = norm_hidden_states.squeeze(1) - # 2. Self-Attention - attn_output = self.attn1(norm_hidden_states, mask=attention_mask, pe=pe) + # 2. Self-Attention — `mask` is the kernel-boundary name for the additive mask. + attn_output = self.attn1(norm_hidden_states, mask=additive_attention_mask, pe=pe) hidden_states = attn_output + hidden_states if hidden_states.ndim == 4: @@ -84,7 +84,7 @@ class Embeddings1DConnector(torch.nn.Module): causal_temporal_positioning (bool): If True, uses causal attention (default=False). num_learnable_registers (int | None): Number of learnable registers to replace padded tokens. If None, disables register replacement. (default=128) - rope_type (LTXRopeType): The RoPE variant to use (default=DEFAULT_ROPE_TYPE). + rope_type (LTXRopeType): The RoPE variant to use. double_precision_rope (bool): Use double precision rope calculation (default=False). """ @@ -99,7 +99,7 @@ class Embeddings1DConnector(torch.nn.Module): positional_embedding_max_pos: list[int] | None = None, causal_temporal_positioning: bool = False, num_learnable_registers: int | None = 128, - rope_type: LTXRopeType = LTXRopeType.INTERLEAVED, + rope_type: LTXRopeType = LTXRopeType.SPLIT, double_precision_rope: bool = False, apply_gated_attention: bool = False, ): @@ -133,51 +133,40 @@ class Embeddings1DConnector(torch.nn.Module): ) def _replace_padded_with_learnable_registers( - self, hidden_states: torch.Tensor, attention_mask: torch.Tensor + self, hidden_states: torch.Tensor, additive_attention_mask: torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor]: - assert hidden_states.shape[1] % self.num_learnable_registers == 0, ( - f"Hidden states sequence length {hidden_states.shape[1]} must be divisible by num_learnable_registers " - f"{self.num_learnable_registers}." - ) + batch_size, seq_len, _ = hidden_states.shape - num_registers_duplications = hidden_states.shape[1] // self.num_learnable_registers - learnable_registers = torch.tile(self.learnable_registers, (num_registers_duplications, 1)) - attention_mask_binary = (attention_mask.squeeze(1).squeeze(1).unsqueeze(-1) >= -9000.0).int() + assert seq_len % self.num_learnable_registers == 0 - non_zero_hidden_states = hidden_states[:, attention_mask_binary.squeeze().bool(), :] - non_zero_nums = non_zero_hidden_states.shape[1] - pad_length = hidden_states.shape[1] - non_zero_nums - adjusted_hidden_states = torch.nn.functional.pad(non_zero_hidden_states, pad=(0, 0, 0, pad_length), value=0) - flipped_mask = torch.flip(attention_mask_binary, dims=[1]) - hidden_states = flipped_mask * adjusted_hidden_states + (1 - flipped_mask) * learnable_registers + registers = self.learnable_registers.repeat(seq_len // self.num_learnable_registers, 1).to(hidden_states.dtype) + registers = registers.unsqueeze(0).expand(batch_size, -1, -1) # (B, seq_len, hidden_dim) + binary_mask = additive_attention_mask[:, 0, 0, :].unsqueeze(-1) >= 0 + binary_mask = binary_mask.to(hidden_states.dtype) + hidden_states = binary_mask * hidden_states + (1 - binary_mask) * registers - attention_mask = torch.full_like( - attention_mask, - 0.0, - dtype=attention_mask.dtype, - device=attention_mask.device, - ) - - return hidden_states, attention_mask + return hidden_states, torch.zeros_like(additive_attention_mask) def forward( self, hidden_states: torch.Tensor, - attention_mask: torch.Tensor | None = None, + additive_attention_mask: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: - """ - Forward pass of Embeddings1DConnector. + """Forward pass of Embeddings1DConnector. Args: - hidden_states (torch.Tensor): Input tensor of embeddings (shape [batch, seq_len, feature_dim]). - attention_mask (torch.Tensor|None): Optional mask for valid tokens (shape compatible with hidden_states). + hidden_states: (B, S, D) input embeddings. + additive_attention_mask: optional additive mask of shape (B, 1, 1, S), where + valid = 0.0 and padding = -torch.finfo(dtype).max. Returns: - tuple[torch.Tensor, torch.Tensor]: Processed features and the corresponding (possibly modified) mask. + (hidden_states, additive_attention_mask) """ if self.num_learnable_registers: - hidden_states, attention_mask = self._replace_padded_with_learnable_registers(hidden_states, attention_mask) + hidden_states, additive_attention_mask = self._replace_padded_with_learnable_registers( + hidden_states, additive_attention_mask + ) indices_grid = torch.arange(hidden_states.shape[1], dtype=torch.float32, device=hidden_states.device) - indices_grid = indices_grid[None, None, :] + indices_grid = indices_grid[None, None, :].expand(hidden_states.shape[0], -1, -1) freq_grid_generator = generate_freq_grid_np if self.double_precision_rope else generate_freq_grid_pytorch freqs_cis = precompute_freqs_cis( indices_grid=indices_grid, @@ -191,11 +180,11 @@ class Embeddings1DConnector(torch.nn.Module): ) for block in self.transformer_1d_blocks: - hidden_states = block(hidden_states, attention_mask=attention_mask, pe=freqs_cis) + hidden_states = block(hidden_states, additive_attention_mask=additive_attention_mask, pe=freqs_cis) hidden_states = rms_norm(hidden_states) - return hidden_states, attention_mask + return hidden_states, additive_attention_mask class Embeddings1DConnectorConfigurator(ModelConfigurator[Embeddings1DConnector]): @@ -204,7 +193,7 @@ class Embeddings1DConnectorConfigurator(ModelConfigurator[Embeddings1DConnector] @classmethod def from_config(cls: type[Embeddings1DConnector], config: dict) -> Embeddings1DConnector: transformer_config = config.get("transformer", {}) - rope_type = LTXRopeType(transformer_config.get("rope_type", "interleaved")) + rope_type = LTXRopeType(transformer_config.get("rope_type", "split")) double_precision_rope = transformer_config.get("frequencies_precision", False) == "float64" pe_max_pos = transformer_config.get("connector_positional_embedding_max_pos", [1]) @@ -231,7 +220,7 @@ class AudioEmbeddings1DConnectorConfigurator(ModelConfigurator[Embeddings1DConne @classmethod def from_config(cls: type[Embeddings1DConnector], config: dict) -> Embeddings1DConnector: transformer_config = config.get("transformer", {}) - rope_type = LTXRopeType(transformer_config.get("rope_type", "interleaved")) + rope_type = LTXRopeType(transformer_config.get("rope_type", "split")) double_precision_rope = transformer_config.get("frequencies_precision", False) == "float64" pe_max_pos = transformer_config.get("connector_positional_embedding_max_pos", [1]) diff --git a/packages/ltx-core/src/ltx_core/text_encoders/gemma/embeddings_processor.py b/packages/ltx-core/src/ltx_core/text_encoders/gemma/embeddings_processor.py index d04f2d3..c12d88f 100644 --- a/packages/ltx-core/src/ltx_core/text_encoders/gemma/embeddings_processor.py +++ b/packages/ltx-core/src/ltx_core/text_encoders/gemma/embeddings_processor.py @@ -19,12 +19,32 @@ def convert_to_additive_mask(attention_mask: torch.Tensor, dtype: torch.dtype) - ) * torch.finfo(dtype).max -def _to_binary_mask(encoded: torch.Tensor, encoded_mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - """Convert connector output mask to binary mask and apply to encoded tensor.""" - binary_mask = (encoded_mask < 0.000001).to(torch.int64) - binary_mask = binary_mask.reshape([encoded.shape[0], encoded.shape[1], 1]) - encoded = encoded * binary_mask - return encoded, binary_mask +def _compute_right_pad_order(additive_mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Compute the index permutation that places valid tokens before pads in each row. + Stable sort: valid tokens keep their relative order. Idempotent for inputs already + right-padded. The sort and reordered mask depend only on the mask, so they can be + computed once and reused across multiple feature tensors that share the mask. + Args: + additive_mask: (B, 1, 1, S) additive mask, ``0.0`` for valid, ``-finfo.max`` for pad. + Returns: + ``(sort_idx, reordered_additive_mask)``: ``sort_idx`` is (B, S); the reordered mask + has the same shape as the input. + """ + binary = (additive_mask[:, 0, 0, :] >= 0).to(torch.int32) # (B, S) + sort_idx = torch.argsort(binary, dim=-1, descending=True, stable=True) # (B, S) + new_binary = torch.gather(binary, 1, sort_idx) + new_additive = (new_binary.to(additive_mask.dtype) - 1) * torch.finfo(additive_mask.dtype).max + return sort_idx, new_additive[:, None, None, :] + + +def _apply_right_pad_order(features: torch.Tensor, sort_idx: torch.Tensor) -> torch.Tensor: + """Apply a precomputed right-pad permutation (from ``_compute_right_pad_order``) to features.""" + return torch.gather(features, 1, sort_idx.unsqueeze(-1).expand_as(features)) + + +def _to_binary_mask(encoded_mask: torch.Tensor, lead_shape: tuple[int, int]) -> torch.Tensor: + """Convert connector output mask to a binary (0/1) mask shaped ``(B, S, 1)`` for broadcasting.""" + return (encoded_mask < 0.000001).to(torch.int64).reshape([lead_shape[0], lead_shape[1], 1]) class EmbeddingsProcessor(nn.Module): @@ -57,12 +77,19 @@ class EmbeddingsProcessor(nn.Module): if self.audio_connector is None and audio_features is not None: raise ValueError("Audio features were provided but no audio connector is configured.") - video_encoded, video_mask = self.video_connector(video_features, additive_attention_mask) - video_encoded, binary_mask = _to_binary_mask(video_encoded, video_mask) + # Connectors expect right-padded input ([valid, pad]). Normalize layout here so the + # upstream tokenizer can keep using either side without coupling to the connector. + # The sort index depends only on the mask, so compute it once and reuse for audio. + sort_idx, mask_for_connector = _compute_right_pad_order(additive_attention_mask) + video_features = _apply_right_pad_order(video_features, sort_idx) + video_encoded, video_mask = self.video_connector(video_features, mask_for_connector) + binary_mask = _to_binary_mask(video_mask, video_encoded.shape[:2]) + video_encoded = video_encoded * binary_mask audio_encoded = None if self.audio_connector is not None: - audio_encoded, _ = self.audio_connector(audio_features, additive_attention_mask) + audio_features = _apply_right_pad_order(audio_features, sort_idx) + audio_encoded, _ = self.audio_connector(audio_features, mask_for_connector) return video_encoded, audio_encoded, binary_mask.squeeze(-1) diff --git a/packages/ltx-core/src/ltx_core/text_encoders/gemma/feature_extractor.py b/packages/ltx-core/src/ltx_core/text_encoders/gemma/feature_extractor.py index cbd5ae6..c48bd81 100644 --- a/packages/ltx-core/src/ltx_core/text_encoders/gemma/feature_extractor.py +++ b/packages/ltx-core/src/ltx_core/text_encoders/gemma/feature_extractor.py @@ -11,37 +11,25 @@ from torch import nn def _norm_and_concat_padded_batch( encoded_text: torch.Tensor, - sequence_lengths: torch.Tensor, - padding_side: str = "right", + attention_mask: torch.Tensor, ) -> torch.Tensor: """Normalize and flatten multi-layer hidden states, respecting padding. Performs per-batch, per-layer normalization using masked mean and range, - then concatenates across the layer dimension. + then concatenates across the layer dimension. Padding-side agnostic: the + binary ``attention_mask`` already encodes which positions are valid. Args: encoded_text: Hidden states of shape [batch, seq_len, hidden_dim, num_layers]. - sequence_lengths: Number of valid (non-padded) tokens per batch item. - padding_side: Whether padding is on "left" or "right". + attention_mask: Binary mask of shape [batch, seq_len], 1 for valid tokens, 0 for padding. Returns: Normalized tensor of shape [batch, seq_len, hidden_dim * num_layers], with padded positions zeroed out. """ - b, t, d, l = encoded_text.shape # noqa: E741 - device = encoded_text.device - - token_indices = torch.arange(t, device=device)[None, :] - - if padding_side == "right": - mask = token_indices < sequence_lengths[:, None] - elif padding_side == "left": - start_indices = t - sequence_lengths[:, None] - mask = token_indices >= start_indices - else: - raise ValueError(f"padding_side must be 'left' or 'right', got {padding_side}") - - mask = rearrange(mask, "b t -> b t 1 1") - + b, _, d, l = encoded_text.shape # noqa: E741 eps = 1e-6 + sequence_lengths = attention_mask.sum(dim=-1) + mask = rearrange(attention_mask.bool(), "b t -> b t 1 1") + masked = encoded_text.masked_fill(~mask, 0.0) denom = (sequence_lengths * d).view(b, 1, 1, 1) mean = masked.sum(dim=(1, 2), keepdim=True) / (denom + eps) @@ -51,12 +39,10 @@ def _norm_and_concat_padded_batch( range_ = x_max - x_min normed = 8 * (encoded_text - mean) / (range_ + eps) - normed = normed.reshape(b, t, -1) + normed = normed.reshape(b, -1, d * l) mask_flattened = rearrange(mask, "b t 1 1 -> b t 1").expand(-1, -1, d * l) - normed = normed.masked_fill(~mask_flattened, 0.0) - - return normed + return normed.masked_fill(~mask_flattened, 0.0) def norm_and_concat_per_token_rms( @@ -97,12 +83,14 @@ class FeatureExtractorV1(nn.Module): self.is_av = is_av def forward( - self, hidden_states: torch.Tensor, attention_mask: torch.Tensor, padding_side: str = "left" + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor, + padding_side: str = "left", # noqa: ARG002 — kept for API stability; norm is layout-agnostic ) -> tuple[torch.Tensor, torch.Tensor | None]: encoded = torch.stack(hidden_states, dim=-1) if isinstance(hidden_states, (list, tuple)) else hidden_states dtype = encoded.dtype - sequence_lengths = attention_mask.sum(dim=-1) - normed = _norm_and_concat_padded_batch(encoded, sequence_lengths, padding_side) + normed = _norm_and_concat_padded_batch(encoded, attention_mask) features = self.aggregate_embed(normed.to(dtype)) if self.is_av: return features, features diff --git a/packages/ltx-core/src/ltx_core/text_encoders/gemma/tokenizer.py b/packages/ltx-core/src/ltx_core/text_encoders/gemma/tokenizer.py index f4f384c..2ac470c 100644 --- a/packages/ltx-core/src/ltx_core/text_encoders/gemma/tokenizer.py +++ b/packages/ltx-core/src/ltx_core/text_encoders/gemma/tokenizer.py @@ -1,6 +1,13 @@ +from enum import Enum + from transformers import AutoTokenizer +class PaddingSide(str, Enum): + LEFT = "left" + RIGHT = "right" + + class LTXVGemmaTokenizer: """ Tokenizer wrapper for Gemma models compatible with LTXV processes. @@ -8,18 +15,18 @@ class LTXVGemmaTokenizer: ensuring correct settings and output formatting for downstream consumption. """ - def __init__(self, tokenizer_path: str, max_length: int = 256): + def __init__(self, tokenizer_path: str, max_length: int = 256, padding_side: PaddingSide = PaddingSide.LEFT): """ Initialize the tokenizer. Args: tokenizer_path (str): Path to the pretrained tokenizer files or model directory. max_length (int, optional): Max sequence length for encoding. Defaults to 256. + padding_side (PaddingSide, optional): Side to pad on. Defaults to ``PaddingSide.LEFT``. """ self.tokenizer = AutoTokenizer.from_pretrained( tokenizer_path, local_files_only=True, model_max_length=max_length ) - # Gemma expects left padding for chat-style prompts; for plain text it doesn't matter much. - self.tokenizer.padding_side = "left" + self.tokenizer.padding_side = padding_side.value if self.tokenizer.pad_token is None: self.tokenizer.pad_token = self.tokenizer.eos_token diff --git a/packages/ltx-core/src/ltx_core/tools.py b/packages/ltx-core/src/ltx_core/tools.py index eed5f11..ec1696e 100644 --- a/packages/ltx-core/src/ltx_core/tools.py +++ b/packages/ltx-core/src/ltx_core/tools.py @@ -138,7 +138,7 @@ class VideoLatentTools(LatentTools): LatentState( latent=initial_latent, denoise_mask=denoise_mask, - positions=positions.to(dtype), + positions=positions, clean_latent=clean_latent, ) ) diff --git a/packages/ltx-core/src/ltx_core/types.py b/packages/ltx-core/src/ltx_core/types.py index 8522445..c9dac29 100644 --- a/packages/ltx-core/src/ltx_core/types.py +++ b/packages/ltx-core/src/ltx_core/types.py @@ -20,15 +20,17 @@ class SpatioTemporalScaleFactors(NamedTuple): """ Describes the spatiotemporal downscaling between decoded video space and the corresponding VAE latent grid. + Field order matches the (frame/time, height, width) axis layout used by + latent tensors and meshgrid coordinates elsewhere in the codebase. """ time: int - width: int height: int + width: int @classmethod def default(cls) -> "SpatioTemporalScaleFactors": - return cls(time=8, width=32, height=32) + return cls(time=8, height=32, width=32) VIDEO_SCALE_FACTORS = SpatioTemporalScaleFactors.default() @@ -74,9 +76,9 @@ class VideoLatentShape(NamedTuple): latent_channels: int = 128, scale_factors: SpatioTemporalScaleFactors = VIDEO_SCALE_FACTORS, ) -> "VideoLatentShape": - frames = (shape.frames - 1) // scale_factors[0] + 1 - height = shape.height // scale_factors[1] - width = shape.width // scale_factors[2] + frames = (shape.frames - 1) // scale_factors.time + 1 + height = shape.height // scale_factors.height + width = shape.width // scale_factors.width return VideoLatentShape( batch=shape.batch, diff --git a/packages/ltx-pipelines/CLAUDE.md b/packages/ltx-pipelines/CLAUDE.md index 631fa01..8b5f62e 100644 --- a/packages/ltx-pipelines/CLAUDE.md +++ b/packages/ltx-pipelines/CLAUDE.md @@ -23,6 +23,7 @@ Inference pipelines for LTX-2 audio-video generation. Depends on `ltx-core` for | `KeyframeInterpolationPipeline` | `keyframe_interpolation.py` | 2 | Full + distilled LoRA | Euler | Keyframe interpolation | | `DistilledPipeline` | `distilled.py` | 2 | Distilled only | Euler | Fastest inference | | `ICLoraPipeline` | `ic_lora.py` | 2 | Distilled only | Euler | Video-to-video with IC-LoRA control | +| `LipDubPipeline` | `lipdub.py` | 2 | Distilled only | Euler | Lip dubbing with IC-LoRA + audio ref conditioning | | `RetakePipeline` | `retake.py` | 1 | Full or distilled | Euler | Video region regeneration | ## Guidance @@ -65,6 +66,10 @@ Inference pipelines for LTX-2 audio-video generation. Depends on `ltx-core` for - `GuidedDenoiser` -- CFG/STG with static `MultiModalGuider` instances (HQ, A2Vid, Retake non-distilled). - `FactoryGuidedDenoiser` -- per-step guider creation via factory (OneStageTI2Vid, TwoStagesTI2Vid, Keyframe). +All denoisers return a `(video_result, audio_result)` tuple of `DenoisedLatentResult` (defined in `utils/types.py`), either element may be `None` for absent modalities. `DenoisedLatentResult.denoised` is the final blended tensor. Guided denoisers additionally populate per-pass fields (`.cond`, `.uncond`, `.ptb`, `.mod`) on each result; `SimpleDenoiser` leaves these `None`. + +`GuidedDenoiser` and `FactoryGuidedDenoiser` accept `force_uncond_pass=True` to run the uncond pass even when `cfg_scale=1.0` (required by CFG++ when the guidance scale is 1 but the uncond prediction is still needed for the ODE derivative). Requires `negative_context` to be set on the guider. When enabled, `DenoisedLatentResult.uncond` will be a tensor instead of `None`. + Guided denoisers batch all guidance passes into a **single transformer call**: states are repeated along the batch dimension, contexts concatenated, and a `BatchedPerturbationConfig` controls which attention ops are skipped per sample. Pass count is dynamic: B=2 for CFG-only, up to B=4 with CFG+STG+modality isolation. Results are split back and blended by the guider. ## Per-pipeline unique features @@ -72,6 +77,7 @@ Guided denoisers batch all guidance passes into a **single transformer call**: s - **HQ**: Res2s second-order sampler for **both** stages, latent-dependent sigma schedule, distilled LoRA on both stages with separate strengths. - **A2Vid**: Audio frozen in both stages (`frozen=True, noise_scale=0.0`). Returns original audio (not VAE-decoded); no `AudioDecoder`. - **IC-LoRA**: `VideoConditionByReferenceLatent`, `reference_downscale_factor` from LoRA metadata, `skip_stage_2`, attention mask downsampling. Stage 2 is LoRA-free and uses `combined_image_conditionings` (no IC-LoRA conditioning). +- **LipDub**: Standalone pipeline; IC reference **video** helpers in `iclora_utils.py`, LipDub-only **audio** patchify/negative positions in `lipdub.py`. Appends frozen audio-reference tokens via `AudioConditionByReferenceLatent` (ltx-core), matching video token order (`[target | ref]`) while keeping reference RoPE positions negative (training-compatible). Single IC-LoRA on both stages; full IC-LoRA video conditioning at stage 1 and 2; stage-2 audio is frozen with S1 latent as initial state and uses S1-derived ref. Final audio decoded from stage 1 latent. The LipDub CLI does not expose `--conditioning-attention-mask`; use `ic_lora.py` if you need spatial IC attention masking. - **Keyframe**: Uses `image_conditionings_by_adding_guiding_latent` in both stages (all frames as keyframe guidance, no replacement) -- unlike TI2Vid which uses `combined_image_conditionings` (frame_idx=0 replaces, others guide). - **Retake**: `TemporalRegionMask` for selective time-window regeneration. `regenerate_video`/`regenerate_audio` flags. Conditional distilled/full behavior. - **Distilled**: Single `self.stage` reused for both stages (not `stage_1`/`stage_2`). diff --git a/packages/ltx-pipelines/README.md b/packages/ltx-pipelines/README.md index 14c466c..3bfe0cb 100644 --- a/packages/ltx-pipelines/README.md +++ b/packages/ltx-pipelines/README.md @@ -64,6 +64,7 @@ Available pipeline modules: - `ltx_pipelines.a2vid_two_stage` - Audio-to-video generation conditioned on an input audio. - `ltx_pipelines.retake` - Regenerate a time region of an existing video. - `ltx_pipelines.hdr_ic_lora` - Video-to-video with HDR output (linear float via LogC3 inverse decode). +- `ltx_pipelines.lipdub` - Lip dubbing / re-voicing with IC-LoRA and audio reference conditioning. Use `--help` with any pipeline module to see all available options and parameters. @@ -113,6 +114,7 @@ Do you need to condition on existing images/videos? | **A2VidPipelineTwoStage** | 2 | ✅ | ✅ | Audio + Image | Audio-driven video generation | | **RetakePipeline** | 1 | ✅ | ❌ | Source Video | Regenerating a time region of a video | | **HDRICLoraPipeline** | 2 | ❌ | ✅ | Video | HDR video-to-video (linear float output for EXR) | +| **LipDubPipeline** | 2 | ✅ | ✅ | Video + Audio | Lip dubbing with audio ref conditioning | --- @@ -238,6 +240,20 @@ Two-stage video-to-video on the distilled model with an HDR IC-LoRA. Decoded lat --- +### 10. LipDubPipeline + +**Best for:** Lip dubbing, rephrasing while keeping the same speaker identity and matching lip movements to new audio. + +**Source**: [`src/ltx_pipelines/lipdub.py`](src/ltx_pipelines/lipdub.py) + +Uses IC-LoRA on a **distilled** checkpoint with a **single** lip-dub IC-LoRA applied in **both** stages. The reference clip provides video and audio reference tokens whose VAE latents are appended to the target audio sequence as frozen reference tokens. The frame count and frame rate are derived from the reference video (frame count is silently snapped to the nearest `8k+1`), so the CLI does not accept `--num-frames` or `--frame-rate`. Required: `--reference-video`. Optional: `--reference-strength`. LoRA: [`Lightricks/LTX-2.3-22b-IC-LoRA-LipDub`](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-LipDub). + +**Note:** Requires a distilled model checkpoint and one lip-dub IC-LoRA (`--lora` exactly once). + +**Use when:** Dubbing, rephrasing with matched lips and speaker identity. + +--- + ## 🎨 Conditioning Types Pipelines use different conditioning methods from [`ltx-core`](../ltx-core/) for controlling generation. See the [ltx-core conditioning documentation](../ltx-core/README.md#conditioning--control) for details. diff --git a/packages/ltx-pipelines/pyproject.toml b/packages/ltx-pipelines/pyproject.toml index b455839..f0d5fda 100644 --- a/packages/ltx-pipelines/pyproject.toml +++ b/packages/ltx-pipelines/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "ltx-pipelines" -version = "1.1.2" +version = "1.1.3" description = "Pipelines implementation for Lightricks' LTX-2 model" readme = "README.md" requires-python = ">=3.10" diff --git a/packages/ltx-pipelines/src/ltx_pipelines/__init__.py b/packages/ltx-pipelines/src/ltx_pipelines/__init__.py index 89dced7..d145f27 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/__init__.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/__init__.py @@ -5,6 +5,7 @@ This package provides ready-to-use pipelines for video generation: - TI2VidTwoStagesPipeline: Two-stage generation with upsampling - DistilledPipeline: Fast distilled two-stage generation - ICLoraPipeline: Image/video conditioning with distilled LoRA +- LipDubPipeline: Lip dubbing with IC-LoRA and audio conditioning - KeyframeInterpolationPipeline: Keyframe-based video interpolation - RetakePipeline: Regenerate a time region (retake) of an existing video For more detailed components and utilities, import from specific submodules @@ -15,6 +16,7 @@ from ltx_pipelines.a2vid_two_stage import A2VidPipelineTwoStage from ltx_pipelines.distilled import DistilledPipeline from ltx_pipelines.ic_lora import ICLoraPipeline from ltx_pipelines.keyframe_interpolation import KeyframeInterpolationPipeline +from ltx_pipelines.lipdub import LipDubPipeline from ltx_pipelines.retake import RetakePipeline from ltx_pipelines.ti2vid_one_stage import TI2VidOneStagePipeline from ltx_pipelines.ti2vid_two_stages import TI2VidTwoStagesPipeline @@ -24,6 +26,7 @@ __all__ = [ "DistilledPipeline", "ICLoraPipeline", "KeyframeInterpolationPipeline", + "LipDubPipeline", "RetakePipeline", "TI2VidOneStagePipeline", "TI2VidTwoStagesPipeline", diff --git a/packages/ltx-pipelines/src/ltx_pipelines/hdr_ic_lora.py b/packages/ltx-pipelines/src/ltx_pipelines/hdr_ic_lora.py index fa1637e..d959990 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/hdr_ic_lora.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/hdr_ic_lora.py @@ -48,7 +48,7 @@ from ltx_core.model.video_vae import TilingConfig, VideoEncoder from ltx_core.quantization import QuantizationPolicy from ltx_core.tiling import DimensionTilingConfig, TileCountConfig from ltx_core.tools import VideoLatentTools -from ltx_core.types import VideoLatentShape +from ltx_core.types import VideoLatentShape, VideoPixelShape from ltx_pipelines.utils.blocks import ( DiffusionStage, ImageConditioner, @@ -412,7 +412,22 @@ class HDRICLoraPipeline: high_quality_hdr=high_quality_hdr, ) ) - with self.stage_2.model_context() as transformer: + # video_tools is required by TiledDataParallelBuilder when stage_2 is + # wrapped for multi-GPU + stage2_video_tools = VideoLatentTools( + VideoLatentPatchifier(patch_size=1), + VideoLatentShape.from_pixel_shape( + VideoPixelShape( + batch=1, + frames=gen_num_frames, + height=gen_h, + width=gen_w, + fps=frame_rate, + ) + ), + frame_rate, + ) + with self.stage_2.model_context(video_tools=stage2_video_tools) as transformer: phase_latent = upscaled_video_latent for phase_idx, (tiling, sigmas_list, use_ic) in enumerate( zip(stage2_tilings, stage2_sigmas, stage2_use_ic_lora, strict=True) @@ -542,10 +557,10 @@ class HDRICLoraPipeline: """ # Cast to float32 so tiled-decode accumulation buffers and blending # masks run in full precision, avoiding bfloat16 seam artifacts. - # Request float32 [0, 1] output — apply_hdr_decode_postprocess expects it. + # apply_hdr_decode_postprocess expects float32 [0, 1]. latent = latent.float() decoded = torch.cat( - list(self.video_decoder(latent, tiling_config, generator, output_dtype=torch.float32)), + [chunk.float() for chunk in self.video_decoder(latent, tiling_config, generator)], dim=0, ) decoded = rearrange(decoded, "f h w c -> 1 c f h w") diff --git a/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py b/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py index 15fe01b..b8172b2 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py @@ -2,20 +2,18 @@ import logging from collections.abc import Iterator import torch -from einops import rearrange -from safetensors import safe_open from ltx_core.components.noisers import GaussianNoiser -from ltx_core.conditioning import ( - ConditioningItem, - ConditioningItemAttentionStrengthWrapper, - VideoConditionByReferenceLatent, -) +from ltx_core.conditioning import ConditioningItem from ltx_core.loader import LoraPathStrengthAndSDOps from ltx_core.loader.registry import Registry from ltx_core.model.video_vae import TilingConfig, VideoEncoder, get_video_chunks_number from ltx_core.quantization import QuantizationPolicy -from ltx_core.types import Audio, VideoLatentShape, VideoPixelShape +from ltx_core.types import Audio, VideoPixelShape +from ltx_pipelines.iclora_utils import ( + append_ic_lora_reference_video_conditionings, + read_lora_reference_downscale_factor, +) from ltx_pipelines.utils.args import ( ImageConditioningInput, VideoConditioningAction, @@ -108,7 +106,7 @@ class ICLoraPipeline: # so inference can resize reference videos to match training conditions. self.reference_downscale_factor = 1 for lora in loras: - scale = _read_lora_reference_downscale_factor(lora.path) + scale = read_lora_reference_downscale_factor(lora.path) if scale != 1: if self.reference_downscale_factor not in (1, scale): raise ValueError( @@ -309,104 +307,26 @@ class ICLoraPipeline: device=self.device, ) - # Calculate scaled dimensions for reference video conditioning. - # IC-LoRAs trained with downscaled reference videos expect the same ratio at inference. - scale = self.reference_downscale_factor - if scale != 1 and (height % scale != 0 or width % scale != 0): - raise ValueError( - f"Output dimensions ({height}x{width}) must be divisible by reference_downscale_factor ({scale})" - ) - ref_height = height // scale - ref_width = width // scale - - for video_path, strength in video_conditioning: - # Load video at scaled-down resolution (if scale > 1) - frame_gen = decode_video_by_frame(path=video_path, frame_cap=num_frames, device=self.device) - video = video_preprocess(frame_gen, ref_height, ref_width, self.dtype, self.device) - encoded_video = video_encoder(video) - reference_video_shape = VideoLatentShape.from_torch_shape(encoded_video.shape) - - # Build attention_mask for ConditioningItemAttentionStrengthWrapper - if conditioning_attention_mask is not None: - # Downsample pixel-space mask to latent space, then scale by strength - latent_mask = self._downsample_mask_to_latent( - mask=conditioning_attention_mask, - target_latent_shape=reference_video_shape, - ) - attn_mask = latent_mask * conditioning_attention_strength - elif conditioning_attention_strength < 1.0: - # Use scalar strength only - attn_mask = conditioning_attention_strength - else: - attn_mask = None - - cond = VideoConditionByReferenceLatent( - latent=encoded_video, - downscale_factor=scale, - strength=strength, - ) - if attn_mask is not None: - cond = ConditioningItemAttentionStrengthWrapper(cond, attention_mask=attn_mask) - conditionings.append(cond) + append_ic_lora_reference_video_conditionings( + conditionings, + video_conditioning, + height=height, + width=width, + num_frames=num_frames, + video_encoder=video_encoder, + dtype=self.dtype, + device=self.device, + reference_downscale_factor=self.reference_downscale_factor, + conditioning_attention_strength=conditioning_attention_strength, + conditioning_attention_mask=conditioning_attention_mask, + tiling_config=None, + ) if video_conditioning: - logging.info(f"[IC-LoRA] Added {len(video_conditioning)} video conditioning(s)") + logging.info("[IC-LoRA] Added %d video conditioning(s)", len(video_conditioning)) return conditionings - @staticmethod - def _downsample_mask_to_latent( - mask: torch.Tensor, - target_latent_shape: VideoLatentShape, - ) -> torch.Tensor: - """ - Downsample a pixel-space mask to latent space using VAE scale factors. - Handles causal temporal downsampling: the first frame is kept separately - (temporal scale factor = 1 for the first frame), while the remaining - frames are downsampled by the VAE's temporal scale factor. - Args: - mask: Pixel-space mask of shape (B, 1, F_pixel, H_pixel, W_pixel). - Values in [0, 1]. - target_latent_shape: Expected latent shape after VAE encoding. - Used to determine the target (F_latent, H_latent, W_latent). - Returns: - Flattened latent-space mask of shape (B, F_lat * H_lat * W_lat), - matching the patchifier's token ordering (f, h, w). - """ - b = mask.shape[0] - f_lat = target_latent_shape.frames - h_lat = target_latent_shape.height - w_lat = target_latent_shape.width - - # Step 1: Spatial downsampling (area interpolation per frame) - f_pix = mask.shape[2] - spatial_down = torch.nn.functional.interpolate( - rearrange(mask, "b 1 f h w -> (b f) 1 h w"), - size=(h_lat, w_lat), - mode="area", - ) - spatial_down = rearrange(spatial_down, "(b f) 1 h w -> b 1 f h w", b=b) - - # Step 2: Causal temporal downsampling - # First frame: kept as-is (causal VAE encodes first frame independently) - first_frame = spatial_down[:, :, :1, :, :] # (B, 1, 1, H_lat, W_lat) - - if f_pix > 1 and f_lat > 1: - # Remaining frames: downsample by temporal factor via group-mean - t = (f_pix - 1) // (f_lat - 1) # temporal downscale factor - assert (f_pix - 1) % (f_lat - 1) == 0, ( - f"Pixel frames ({f_pix}) not compatible with latent frames ({f_lat}): " - f"(f_pix - 1) must be divisible by (f_lat - 1)" - ) - rest = rearrange(spatial_down[:, :, 1:, :, :], "b 1 (f t) h w -> b 1 f t h w", t=t) - rest = rest.mean(dim=3) # (B, 1, F_lat-1, H_lat, W_lat) - latent_mask = torch.cat([first_frame, rest], dim=2) # (B, 1, F_lat, H_lat, W_lat) - else: - latent_mask = first_frame - - # Flatten to (B, F_lat * H_lat * W_lat) matching patchifier token order (f, h, w) - return rearrange(latent_mask, "b 1 f h w -> b (f h w)") - @torch.inference_mode() def main() -> None: @@ -523,26 +443,5 @@ def _load_mask_video( return mask.clamp(0.0, 1.0) -def _read_lora_reference_downscale_factor(lora_path: str) -> int: - """Read reference_downscale_factor from LoRA safetensors metadata. - Some IC-LoRA models are trained with reference videos at lower resolution than - the target output. This allows for more efficient training and can improve - generalization. The downscale factor indicates the ratio between target and - reference resolutions (e.g., factor=2 means reference is half the resolution). - Args: - lora_path: Path to the LoRA .safetensors file - Returns: - The reference downscale factor (1 if not specified in metadata, meaning - reference and target have the same resolution) - """ - try: - with safe_open(lora_path, framework="pt") as f: - metadata = f.metadata() or {} - return int(metadata.get("reference_downscale_factor", 1)) - except Exception as e: - logging.warning(f"Failed to read metadata from LoRA file '{lora_path}': {e}") - return 1 - - if __name__ == "__main__": main() diff --git a/packages/ltx-pipelines/src/ltx_pipelines/iclora_utils.py b/packages/ltx-pipelines/src/ltx_pipelines/iclora_utils.py new file mode 100644 index 0000000..463af81 --- /dev/null +++ b/packages/ltx-pipelines/src/ltx_pipelines/iclora_utils.py @@ -0,0 +1,120 @@ +"""Shared IC-LoRA helpers: LoRA metadata, mask downsampling, reference-video conditioning. +Used by ``ic_lora`` and ``lipdub`` (video reference path only). LipDub audio helpers live in ``lipdub.py``. +""" + +from __future__ import annotations + +import logging + +import torch +from einops import rearrange +from safetensors import safe_open + +from ltx_core.conditioning import ( + ConditioningItem, + ConditioningItemAttentionStrengthWrapper, + VideoConditionByReferenceLatent, +) +from ltx_core.model.video_vae import TilingConfig, VideoEncoder +from ltx_core.types import VideoLatentShape +from ltx_pipelines.utils.media_io import decode_video_by_frame, video_preprocess + + +def read_lora_reference_downscale_factor(lora_path: str) -> int: + """Read ``reference_downscale_factor`` from LoRA safetensors metadata (default 1).""" + try: + with safe_open(lora_path, framework="pt") as f: + metadata = f.metadata() or {} + return int(metadata.get("reference_downscale_factor", 1)) + except Exception as e: + logging.warning("Failed to read metadata from LoRA file '%s': %s", lora_path, e) + return 1 + + +def downsample_mask_video_to_latent( + mask: torch.Tensor, + target_latent_shape: VideoLatentShape, +) -> torch.Tensor: + """Downsample a pixel-space mask video to flattened latent token weights.""" + b = mask.shape[0] + f_lat = target_latent_shape.frames + h_lat = target_latent_shape.height + w_lat = target_latent_shape.width + + f_pix = mask.shape[2] + spatial_down = torch.nn.functional.interpolate( + rearrange(mask, "b 1 f h w -> (b f) 1 h w"), + size=(h_lat, w_lat), + mode="area", + ) + spatial_down = rearrange(spatial_down, "(b f) 1 h w -> b 1 f h w", b=b) + + first_frame = spatial_down[:, :, :1, :, :] + + if f_pix > 1 and f_lat > 1: + t = (f_pix - 1) // (f_lat - 1) + assert (f_pix - 1) % (f_lat - 1) == 0, ( + f"Pixel frames ({f_pix}) not compatible with latent frames ({f_lat}): " + f"(f_pix - 1) must be divisible by (f_lat - 1)" + ) + rest = rearrange(spatial_down[:, :, 1:, :, :], "b 1 (f t) h w -> b 1 f t h w", t=t) + rest = rest.mean(dim=3) + latent_mask = torch.cat([first_frame, rest], dim=2) + else: + latent_mask = first_frame + + return rearrange(latent_mask, "b 1 f h w -> b (f h w)") + + +def append_ic_lora_reference_video_conditionings( # noqa: PLR0913 + conditionings: list[ConditioningItem], + video_conditioning: list[tuple[str, float]], + *, + height: int, + width: int, + num_frames: int, + video_encoder: VideoEncoder, + dtype: torch.dtype, + device: torch.device, + reference_downscale_factor: int, + conditioning_attention_strength: float, + conditioning_attention_mask: torch.Tensor | None, + tiling_config: TilingConfig | None = None, +) -> None: + """Append :class:`VideoConditionByReferenceLatent` items for each reference path.""" + scale = reference_downscale_factor + if scale != 1 and (height % scale != 0 or width % scale != 0): + raise ValueError( + f"Output dimensions ({height}x{width}) must be divisible by reference_downscale_factor ({scale})" + ) + ref_height = height // scale + ref_width = width // scale + + for video_path, strength in video_conditioning: + frame_gen = decode_video_by_frame(path=video_path, frame_cap=num_frames, device=device) + video = video_preprocess(frame_gen, ref_height, ref_width, dtype, device) + if tiling_config is not None: + encoded_video = video_encoder.tiled_encode(video, tiling_config) + else: + encoded_video = video_encoder(video) + reference_video_shape = VideoLatentShape.from_torch_shape(encoded_video.shape) + + if conditioning_attention_mask is not None: + latent_mask = downsample_mask_video_to_latent( + mask=conditioning_attention_mask, + target_latent_shape=reference_video_shape, + ) + attn_mask = latent_mask * conditioning_attention_strength + elif conditioning_attention_strength < 1.0: + attn_mask = conditioning_attention_strength + else: + attn_mask = None + + cond = VideoConditionByReferenceLatent( + latent=encoded_video, + downscale_factor=scale, + strength=strength, + ) + if attn_mask is not None: + cond = ConditioningItemAttentionStrengthWrapper(cond, attention_mask=attn_mask) + conditionings.append(cond) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/lipdub.py b/packages/ltx-pipelines/src/ltx_pipelines/lipdub.py new file mode 100644 index 0000000..2ba8dc8 --- /dev/null +++ b/packages/ltx-pipelines/src/ltx_pipelines/lipdub.py @@ -0,0 +1,334 @@ +"""Two-stage lip-dubbing pipeline with IC-LoRA and appended audio reference conditioning.""" + +from __future__ import annotations + +import logging +from collections.abc import Iterator + +import torch + +from ltx_core.components.noisers import GaussianNoiser +from ltx_core.components.patchifiers import AudioPatchifier +from ltx_core.conditioning import AudioConditionByReferenceLatent +from ltx_core.loader import LoraPathStrengthAndSDOps +from ltx_core.loader.registry import Registry +from ltx_core.model.audio_vae import encode_audio as vae_encode_audio +from ltx_core.model.video_vae import TilingConfig, VideoEncoder, get_video_chunks_number +from ltx_core.quantization import QuantizationPolicy +from ltx_core.types import Audio, AudioLatentShape, SpatioTemporalScaleFactors, VideoPixelShape +from ltx_pipelines.iclora_utils import ( + append_ic_lora_reference_video_conditionings, + read_lora_reference_downscale_factor, +) +from ltx_pipelines.utils.args import ( + ImageConditioningInput, + detect_checkpoint_path, + lipdub_arg_parser, +) +from ltx_pipelines.utils.blocks import ( + AudioConditioner, + AudioDecoder, + DiffusionStage, + ImageConditioner, + PromptEncoder, + VideoDecoder, + VideoUpsampler, +) +from ltx_pipelines.utils.constants import DISTILLED_SIGMAS, STAGE_2_DISTILLED_SIGMAS, detect_params +from ltx_pipelines.utils.denoisers import SimpleDenoiser +from ltx_pipelines.utils.helpers import assert_resolution, combined_image_conditionings, get_device +from ltx_pipelines.utils.media_io import decode_audio_from_file, encode_video, get_videostream_metadata +from ltx_pipelines.utils.types import ModalitySpec, OffloadMode + + +def _snap_frames_to_8k1(frames: int) -> int: + """Round ``frames`` down to the nearest ``8k+1`` (the model's required frame count).""" + time_scale = SpatioTemporalScaleFactors.default().time + return ((frames - 1) // time_scale) * time_scale + 1 + + +class LipDubPipeline: + """Two-stage lip-dubbing with IC-LoRA video reference and appended audio reference tokens.""" + + def __init__( + self, + distilled_checkpoint_path: str, + spatial_upsampler_path: str, + gemma_root: str, + ic_lora: LoraPathStrengthAndSDOps, + device: torch.device | None = None, + quantization: QuantizationPolicy | None = None, + registry: Registry | None = None, + torch_compile: bool = False, + offload_mode: OffloadMode = OffloadMode.NONE, + ) -> None: + self.device = device or get_device() + self.dtype = torch.bfloat16 + self.ic_lora = ic_lora + loras = (ic_lora,) + + self.prompt_encoder = PromptEncoder( + distilled_checkpoint_path, + gemma_root, + self.dtype, + self.device, + registry=registry, + offload_mode=offload_mode, + ) + self.image_conditioner = ImageConditioner(distilled_checkpoint_path, self.dtype, self.device, registry=registry) + self.audio_conditioner = AudioConditioner( + distilled_checkpoint_path, + self.dtype, + self.device, + registry=registry, + ) + self.stage = DiffusionStage( + distilled_checkpoint_path, + self.dtype, + self.device, + loras=loras, + quantization=quantization, + registry=registry, + torch_compile=torch_compile, + offload_mode=offload_mode, + ) + self.upsampler = VideoUpsampler( + distilled_checkpoint_path, spatial_upsampler_path, self.dtype, self.device, registry=registry + ) + self.video_decoder = VideoDecoder(distilled_checkpoint_path, self.dtype, self.device, registry=registry) + self.audio_decoder = AudioDecoder(distilled_checkpoint_path, self.dtype, self.device, registry=registry) + self.reference_downscale_factor = read_lora_reference_downscale_factor(ic_lora.path) + + def _create_stage_conditionings( + self, + images: list[ImageConditioningInput], + reference_video_path: str, + reference_strength: float, + height: int, + width: int, + num_frames: int, + video_encoder: VideoEncoder, + encode_tiling: TilingConfig | None, + ) -> list: + conditionings = combined_image_conditionings( + images=images, + height=height, + width=width, + video_encoder=video_encoder, + dtype=self.dtype, + device=self.device, + ) + append_ic_lora_reference_video_conditionings( + conditionings, + [(reference_video_path, reference_strength)], + height=height, + width=width, + num_frames=num_frames, + video_encoder=video_encoder, + dtype=self.dtype, + device=self.device, + reference_downscale_factor=self.reference_downscale_factor, + conditioning_attention_strength=1.0, + conditioning_attention_mask=None, + tiling_config=encode_tiling, + ) + return conditionings + + def _encode_reference_audio_vae_latent(self, video_path: str) -> torch.Tensor: + audio = decode_audio_from_file(video_path, self.device) + if audio is None: + msg = f"No audio stream found in {video_path}" + raise ValueError(msg) + return self.audio_conditioner(lambda enc: vae_encode_audio(audio, enc, None)) + + @torch.inference_mode() + def __call__( # noqa: PLR0913 + self, + prompt: str, + seed: int, + height: int, + width: int, + images: list[ImageConditioningInput], + reference_video_path: str, + reference_strength: float = 1.0, + enhance_prompt: bool = False, + tiling_config: TilingConfig | None = None, + stage_1_sigmas: torch.Tensor = DISTILLED_SIGMAS, + stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS, + ) -> tuple[Iterator[torch.Tensor], Audio]: + assert_resolution(height=height, width=width, is_two_stage=True) + + meta = get_videostream_metadata(reference_video_path) + num_frames = _snap_frames_to_8k1(meta.frames) + frame_rate = float(meta.fps) + + generator = torch.Generator(device=self.device).manual_seed(seed) + noiser = GaussianNoiser(generator=generator) + + (ctx_p,) = self.prompt_encoder( + [prompt], + enhance_first_prompt=enhance_prompt, + enhance_prompt_image=images[0][0] if len(images) > 0 else None, + enhance_prompt_seed=seed, + ) + video_context, audio_context = ctx_p.video_encoding, ctx_p.audio_encoding + + stage_1_output_shape = VideoPixelShape( + batch=1, + frames=num_frames, + width=width // 2, + height=height // 2, + fps=frame_rate, + ) + encode_tiling = TilingConfig.default() + + def build_image_conditionings(output_shape: VideoPixelShape) -> list: + return self.image_conditioner( + lambda enc: self._create_stage_conditionings( + images=images, + reference_video_path=reference_video_path, + reference_strength=reference_strength, + height=output_shape.height, + width=output_shape.width, + num_frames=num_frames, + video_encoder=enc, + encode_tiling=encode_tiling, + ) + ) + + def build_audio_ref_conditioning(audio_latent: torch.Tensor) -> AudioConditionByReferenceLatent: + ref_patch, ref_pos = patchify_lipdub_audio_reference_latent( + audio_latent, + negative_positions=True, + device=self.device, + ) + return AudioConditionByReferenceLatent(ref_patch, ref_pos, strength=1.0) + + stage_1_conditionings = build_image_conditionings(stage_1_output_shape) + + ref_vae = self._encode_reference_audio_vae_latent(reference_video_path) + audio_conditionings = [build_audio_ref_conditioning(ref_vae)] + + stage_1_sigmas_tensor = stage_1_sigmas.to(dtype=torch.float32, device=self.device) + video_state, audio_state = self.stage( + denoiser=SimpleDenoiser(video_context, audio_context), + sigmas=stage_1_sigmas_tensor, + noiser=noiser, + width=stage_1_output_shape.width, + height=stage_1_output_shape.height, + frames=num_frames, + fps=frame_rate, + video=ModalitySpec( + context=video_context, + conditionings=stage_1_conditionings, + ), + audio=ModalitySpec( + context=audio_context, + conditionings=audio_conditionings, + ), + ) + + s1_audio_latent = audio_state.latent.clone() + + upscaled_video_latent = self.upsampler(video_state.latent[:1]) + stage_2_sigmas_tensor = stage_2_sigmas.to(dtype=torch.float32, device=self.device) + stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate) + stage_2_conditionings = build_image_conditionings(stage_2_output_shape) + + stage_2_audio_conditionings = [build_audio_ref_conditioning(s1_audio_latent)] + + video_state, _audio_unused = self.stage( + denoiser=SimpleDenoiser(video_context, audio_context), + sigmas=stage_2_sigmas_tensor, + noiser=noiser, + width=width, + height=height, + frames=num_frames, + fps=frame_rate, + video=ModalitySpec( + context=video_context, + conditionings=stage_2_conditionings, + noise_scale=stage_2_sigmas_tensor[0].item(), + initial_latent=upscaled_video_latent, + ), + audio=ModalitySpec( + context=audio_context, + conditionings=stage_2_audio_conditionings, + frozen=True, + noise_scale=0.0, + initial_latent=s1_audio_latent, + ), + ) + + decoded_video = self.video_decoder(video_state.latent, tiling_config, generator) + decoded_audio = self.audio_decoder(s1_audio_latent) + return decoded_video, decoded_audio + + +def patchify_lipdub_audio_reference_latent( + vae_latents: torch.Tensor, + *, + negative_positions: bool, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + """Patchify audio VAE latents and build RoPE positions (optional negative shift for reference).""" + patchifier = AudioPatchifier(patch_size=1) + patchified = patchifier.patchify(vae_latents) + b, c, _t, mel_bins = vae_latents.shape + seq_len = patchified.shape[1] + latent_coords = patchifier.get_patch_grid_bounds( + output_shape=AudioLatentShape(batch=b, channels=c, frames=seq_len, mel_bins=mel_bins), + device=device, + ) + positions = latent_coords.to(dtype=torch.float32) + if negative_positions: + aud_dur = positions[:, :, -1, 1].max().item() + positions = positions - aud_dur - 0.04 + return patchified, positions + + +@torch.inference_mode() +def main() -> None: + logging.getLogger().setLevel(logging.INFO) + checkpoint_path = detect_checkpoint_path(distilled=True) + params = detect_params(checkpoint_path) + parser = lipdub_arg_parser(params=params) + args = parser.parse_args() + + if not args.lora or len(args.lora) != 1: + raise ValueError("LipDub requires exactly one --lora (the lip-dub IC-LoRA).") + + pipeline = LipDubPipeline( + distilled_checkpoint_path=args.distilled_checkpoint_path, + spatial_upsampler_path=args.spatial_upsampler_path, + gemma_root=args.gemma_root, + ic_lora=args.lora[0], + quantization=args.quantization, + torch_compile=args.compile, + offload_mode=args.offload_mode, + ) + tiling_config = TilingConfig.default() + src = get_videostream_metadata(args.reference_video) + video_chunks_number = get_video_chunks_number(_snap_frames_to_8k1(src.frames), tiling_config) + video, audio = pipeline( + prompt=args.prompt, + seed=args.seed, + height=args.height, + width=args.width, + images=[], + reference_video_path=args.reference_video, + reference_strength=args.reference_strength, + tiling_config=tiling_config, + enhance_prompt=args.enhance_prompt, + ) + encode_video( + video=video, + fps=int(src.fps), + audio=audio, + output_path=args.output_path, + video_chunks_number=video_chunks_number, + ) + + +if __name__ == "__main__": + main() diff --git a/packages/ltx-pipelines/src/ltx_pipelines/retake.py b/packages/ltx-pipelines/src/ltx_pipelines/retake.py index ea8ee2a..a1e7709 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/retake.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/retake.py @@ -307,7 +307,7 @@ def main() -> None: gemma_root=args.gemma_root, loras=tuple(args.lora) if args.lora else (), quantization=args.quantization, - distilled=args.distilled, + distilled=True, torch_compile=args.compile, offload_mode=args.offload_mode, ) diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/__init__.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/__init__.py index 9d5b291..33cf327 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/__init__.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/__init__.py @@ -16,15 +16,17 @@ from ltx_pipelines.utils.helpers import ( image_conditionings_by_adding_guiding_latent, ) from ltx_pipelines.utils.samplers import ( + euler_cfg_pp_denoising_loop, euler_denoising_loop, gradient_estimating_euler_denoising_loop, res2s_audio_video_denoising_loop, ) -from ltx_pipelines.utils.types import Denoiser, ModalitySpec +from ltx_pipelines.utils.types import DenoisedLatentResult, Denoiser, ModalitySpec __all__ = [ "AudioConditioner", "AudioDecoder", + "DenoisedLatentResult", "Denoiser", "DiffusionStage", "FactoryGuidedDenoiser", @@ -38,6 +40,7 @@ __all__ = [ "assert_resolution", "cleanup_memory", "combined_image_conditionings", + "euler_cfg_pp_denoising_loop", "euler_denoising_loop", "get_device", "gradient_estimating_euler_denoising_loop", diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/args.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/args.py index bb3236d..f7e2b05 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/args.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/args.py @@ -1,4 +1,5 @@ import argparse +from collections.abc import Sequence from pathlib import Path from typing import NamedTuple @@ -115,35 +116,34 @@ def resolve_path(path: str) -> str: QUANTIZATION_POLICIES = ("fp8-cast", "fp8-scaled-mm") -class QuantizationAction(argparse.Action): - def __call__( - self, - parser: argparse.ArgumentParser, # noqa: ARG002 - namespace: argparse.Namespace, - values: list[str], - option_string: str | None = None, - ) -> None: - if len(values) > 2: - msg = ( - f"{option_string} accepts at most 2 arguments (POLICY and optional AMAX_PATH), got {len(values)} values" +def _resolve_quantization(namespace: argparse.Namespace) -> None: + # Resolution is deferred until after parse_args because fp8-scaled-mm needs the + # checkpoint path, which isn't on the namespace when the --quantization argument + # is parsed. + name = getattr(namespace, "quantization", None) + if name is None or isinstance(name, QuantizationPolicy): + return + if name == "fp8-cast": + namespace.quantization = QuantizationPolicy.fp8_cast() + return + if name == "fp8-scaled-mm": + ckpt = getattr(namespace, "checkpoint_path", None) or getattr(namespace, "distilled_checkpoint_path", None) + if ckpt is None: + raise SystemExit( + "--quantization fp8-scaled-mm requires --checkpoint-path (or --distilled-checkpoint-path)." ) - raise argparse.ArgumentError(self, msg) + namespace.quantization = QuantizationPolicy.fp8_scaled_mm(ckpt) - policy_name = values[0] - if policy_name not in QUANTIZATION_POLICIES: - msg = f"Unknown quantization policy '{policy_name}'. Choose from: {', '.join(QUANTIZATION_POLICIES)}" - raise argparse.ArgumentError(self, msg) - if policy_name == "fp8-cast": - if len(values) > 1: - msg = f"{option_string} fp8-cast does not accept additional arguments" - raise argparse.ArgumentError(self, msg) - policy = QuantizationPolicy.fp8_cast() - elif policy_name == "fp8-scaled-mm": - amax_path = resolve_path(values[1]) if len(values) > 1 else None - policy = QuantizationPolicy.fp8_scaled_mm(amax_path) - - setattr(namespace, self.dest, policy) +class _PipelineArgumentParser(argparse.ArgumentParser): + def parse_args( # type: ignore[override] + self, + args: Sequence[str] | None = None, + namespace: argparse.Namespace | None = None, + ) -> argparse.Namespace: + ns = super().parse_args(args, namespace) + _resolve_quantization(ns) + return ns def detect_checkpoint_path(distilled: bool = False) -> str: @@ -159,7 +159,7 @@ def basic_arg_parser( params: PipelineParams = LTX_2_3_PARAMS, distilled: bool = False, ) -> argparse.ArgumentParser: - parser = argparse.ArgumentParser() + parser = _PipelineArgumentParser() if distilled: parser.add_argument( "--distilled-checkpoint-path", @@ -264,16 +264,14 @@ def basic_arg_parser( parser.add_argument( "--quantization", - dest="quantization", - action=QuantizationAction, - nargs="+", - metavar=("POLICY", "AMAX_PATH"), + choices=QUANTIZATION_POLICIES, default=None, help=( f"Quantization policy: {', '.join(QUANTIZATION_POLICIES)}. " "fp8-cast uses FP8 casting with upcasting during inference. " - "fp8-scaled-mm uses FP8 scaled matrix multiplication (optionally provide amax calibration file path). " - "Example: --quantization fp8-cast or --quantization fp8-scaled-mm /path/to/amax.json" + "fp8-scaled-mm uses FP8 scaled matrix multiplication; the layer set is auto-discovered " + "from the checkpoint's .weight_scale tensors. " + "Example: --quantization fp8-cast or --quantization fp8-scaled-mm" ), ) parser.add_argument( @@ -348,6 +346,53 @@ def video_editing_arg_parser( return parser +def lipdub_arg_parser( + params: PipelineParams = LTX_2_3_PARAMS, +) -> argparse.ArgumentParser: + """Argument parser for the lip-dub pipeline. + Frame count and frame rate are derived from the reference video at runtime (the frame count + is silently snapped down to the nearest 8k+1), so this parser intentionally omits + --num-frames, --frame-rate, and --image. Distilled checkpoint only. + """ + parser = basic_arg_parser(params=params, distilled=True) + parser.add_argument( + "--height", + type=int, + default=params.stage_2_height, + help=( + f"Height of the generated video in pixels, should be divisible by 64 (default: {params.stage_2_height})." + ), + ) + parser.add_argument( + "--width", + type=int, + default=params.stage_2_width, + help=f"Width of the generated video in pixels, should be divisible by 64 (default: {params.stage_2_width}).", + ) + parser.add_argument( + "--spatial-upsampler-path", + type=resolve_path, + required=True, + help=( + "Path to the spatial upsampler model used to increase the resolution " + "of the generated video in the latent space." + ), + ) + parser.add_argument( + "--reference-video", + type=resolve_path, + required=True, + help="Reference video file (video + audio track used for IC-LoRA and audio identity).", + ) + parser.add_argument( + "--reference-strength", + type=float, + default=1.0, + help="Strength for IC-LoRA video reference conditioning (default: 1.0).", + ) + return parser + + def default_1_stage_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argparse.ArgumentParser: video_guider = params.video_guider_params audio_guider = params.audio_guider_params diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/blocks.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/blocks.py index abece4e..66c9c7b 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/blocks.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/blocks.py @@ -21,7 +21,8 @@ from ltx_core.components.noisers import Noiser from ltx_core.components.patchifiers import AudioPatchifier, VideoLatentPatchifier from ltx_core.components.protocols import DiffusionStepProtocol from ltx_core.loader import SDOps -from ltx_core.loader.primitives import LoraPathStrengthAndSDOps +from ltx_core.loader.module_ops import ModuleOps +from ltx_core.loader.primitives import BuilderProtocol, LoraPathStrengthAndSDOps, ModelBuilderProtocol from ltx_core.loader.registry import DummyRegistry, Registry from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder from ltx_core.model.audio_vae import ( @@ -37,12 +38,14 @@ from ltx_core.model.audio_vae import ( ) from ltx_core.model.transformer import ( LTXV_MODEL_COMFY_RENAMING_MAP, + LTXModel, LTXModelConfigurator, X0Model, ) from ltx_core.model.transformer.compiling import COMPILE_TRANSFORMER, modify_sd_ops_for_compilation from ltx_core.model.upsampler import LatentUpsamplerConfigurator, upsample_video from ltx_core.model.video_vae import ( + MEMORY_EFFICIENT_DECODE, VAE_DECODER_COMFY_KEYS_FILTER, VAE_ENCODER_COMFY_KEYS_FILTER, TilingConfig, @@ -59,10 +62,11 @@ from ltx_core.text_encoders.gemma import ( GemmaTextEncoderConfigurator, module_ops_from_gemma_root, ) -from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessorOutput +from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessor, EmbeddingsProcessorOutput from ltx_core.tools import AudioLatentTools, LatentTools, VideoLatentTools from ltx_core.types import Audio, AudioLatentShape, LatentState, VideoLatentShape, VideoPixelShape from ltx_core.utils import find_matching_file +from ltx_pipelines.multigpu.delegating_builder import DelegatingBuilder from ltx_pipelines.utils.gpu_model import gpu_model from ltx_pipelines.utils.helpers import ( cleanup_memory, @@ -83,6 +87,20 @@ _M = TypeVar("_M", bound=torch.nn.Module) # --------------------------------------------------------------------------- +def _chain_quantization( + sd_ops: SDOps, + module_ops: tuple[ModuleOps, ...], + quantization: QuantizationPolicy, +) -> tuple[SDOps, tuple[ModuleOps, ...]]: + chained_sd_ops = sd_ops + if quantization.sd_ops is not None: + chained_sd_ops = SDOps( + name=f"sd_ops_chain_{sd_ops.name}+{quantization.sd_ops.name}", + mapping=(*sd_ops.mapping, *quantization.sd_ops.mapping), + ) + return chained_sd_ops, (*module_ops, *quantization.module_ops) + + @contextmanager def _streaming_model( builder: StreamingModelBuilder, @@ -154,16 +172,43 @@ class DiffusionStage: registry: Registry | None = None, torch_compile: bool = False, offload_mode: OffloadMode = OffloadMode.NONE, + transformer_builder: ModelBuilderProtocol[LTXModel] | DelegatingBuilder[LTXModel] | None = None, ) -> None: + self._dtype = dtype + self._device = device + self._quantization = quantization + self._torch_compile = torch_compile + self._offload_mode = offload_mode + if transformer_builder is not None: + self._transformer_builder = transformer_builder + else: + self._transformer_builder = Builder( + model_path=checkpoint_path, + model_class_configurator=LTXModelConfigurator, + model_sd_ops=LTXV_MODEL_COMFY_RENAMING_MAP, + loras=tuple(loras), + registry=registry or DummyRegistry(), + ) + if offload_mode != OffloadMode.NONE: if torch_compile: raise ValueError("torch.compile is not supported with layer streaming") + streaming_sd_ops: SDOps = LTXV_MODEL_COMFY_RENAMING_MAP + streaming_module_ops: tuple[ModuleOps, ...] = () if quantization is not None: - raise ValueError("quantization is not supported with layer streaming") + if quantization.kind != QuantizationPolicy.Kind.FP8_CAST: + raise ValueError( + f"Layer streaming supports only QuantizationPolicy.fp8_cast(); " + f"got kind={quantization.kind!r} which produces heterogeneous block layouts." + ) + streaming_sd_ops, streaming_module_ops = _chain_quantization( + streaming_sd_ops, streaming_module_ops, quantization + ) self._streaming_builder = StreamingModelBuilder( model_class_configurator=LTXModelConfigurator, model_path=checkpoint_path, - model_sd_ops=LTXV_MODEL_COMFY_RENAMING_MAP, + model_sd_ops=streaming_sd_ops, + module_ops=streaming_module_ops, loras=tuple(loras), registry=registry or DummyRegistry(), blocks_attr="velocity_model.transformer_blocks", @@ -172,19 +217,6 @@ class DiffusionStage: model_wrapper=lambda m: X0Model(m).eval(), ) - self._dtype = dtype - self._device = device - self._quantization = quantization - self._torch_compile = torch_compile - self._offload_mode = offload_mode - self._transformer_builder = Builder( - model_path=checkpoint_path, - model_class_configurator=LTXModelConfigurator, - model_sd_ops=LTXV_MODEL_COMFY_RENAMING_MAP, - loras=tuple(loras), - registry=registry or DummyRegistry(), - ) - def _build_transformer(self, *, device: torch.device | None = None, **kwargs: object) -> X0Model: target = device or self._device sd_ops = self._transformer_builder.model_sd_ops @@ -198,18 +230,12 @@ class DiffusionStage: LoraPathStrengthAndSDOps( lora.path, lora.strength, - modify_sd_ops_for_compilation( - lora.sd_ops if lora.sd_ops is not None else SDOps(name="identity"), number_of_layers - ), + modify_sd_ops_for_compilation(lora.sd_ops, number_of_layers), ) for lora in loras ) if self._quantization is not None: - module_ops = (*module_ops, *self._quantization.module_ops) - sd_ops = SDOps( - name=f"sd_ops_chain_{sd_ops.name}+{self._quantization.sd_ops.name}", - mapping=(*sd_ops.mapping, *self._quantization.sd_ops.mapping), - ) + sd_ops, module_ops = _chain_quantization(sd_ops, module_ops, self._quantization) builder = self._transformer_builder.with_module_ops(module_ops).with_sd_ops(sd_ops).with_loras(loras) return X0Model(builder.build(device=target, **kwargs)).to(target).eval() @@ -359,31 +385,40 @@ class PromptEncoder: device: torch.device, registry: Registry | None = None, offload_mode: OffloadMode = OffloadMode.NONE, + text_encoder_builder: BuilderProtocol | None = None, ) -> None: self._dtype = dtype self._device = device self._offload_mode = offload_mode - module_ops = module_ops_from_gemma_root(gemma_root) - model_folder = find_matching_file(gemma_root, "model*.safetensors").parent - weight_paths = [str(p) for p in model_folder.rglob("*.safetensors")] - - self._text_encoder_builder = Builder( - model_path=tuple(weight_paths), - model_class_configurator=GemmaTextEncoderConfigurator, - model_sd_ops=GEMMA_LLM_KEY_OPS, - module_ops=(GEMMA_MODEL_OPS, *module_ops), - registry=registry or DummyRegistry(), - ) - self._streaming_text_encoder_builder = StreamingModelBuilder( - model_path=tuple(weight_paths), - model_class_configurator=GemmaTextEncoderConfigurator, - model_sd_ops=GEMMA_LLM_KEY_OPS, - module_ops=(GEMMA_MODEL_OPS, *module_ops), - registry=registry or DummyRegistry(), - blocks_attr="model.model.language_model.layers", - blocks_prefix="model.model.language_model.layers", - ) + if text_encoder_builder is not None: + if offload_mode != OffloadMode.NONE: + raise ValueError( + "text_encoder_builder cannot be used with offload_mode != OffloadMode.NONE " + "because no streaming text encoder builder is available." + ) + self._text_encoder_builder = text_encoder_builder + self._streaming_text_encoder_builder = None + else: + module_ops = module_ops_from_gemma_root(gemma_root) + model_folder = find_matching_file(gemma_root, "model*.safetensors").parent + weight_paths = [str(p) for p in model_folder.rglob("*.safetensors")] + self._text_encoder_builder = Builder( + model_path=tuple(weight_paths), + model_class_configurator=GemmaTextEncoderConfigurator, + model_sd_ops=GEMMA_LLM_KEY_OPS, + module_ops=(GEMMA_MODEL_OPS, *module_ops), + registry=registry or DummyRegistry(), + ) + self._streaming_text_encoder_builder = StreamingModelBuilder( + model_path=tuple(weight_paths), + model_class_configurator=GemmaTextEncoderConfigurator, + model_sd_ops=GEMMA_LLM_KEY_OPS, + module_ops=(GEMMA_MODEL_OPS, *module_ops), + registry=registry or DummyRegistry(), + blocks_attr="model.model.language_model.layers", + blocks_prefix="model.model.language_model.layers", + ) self._embeddings_processor_builder = Builder( model_path=checkpoint_path, model_class_configurator=EmbeddingsProcessorConfigurator, @@ -391,10 +426,18 @@ class PromptEncoder: registry=registry or DummyRegistry(), ) + def _build_text_encoder(self) -> torch.nn.Module: + """Build the Gemma text encoder (non-streaming path).""" + return self._text_encoder_builder.build(device=self._device, dtype=self._dtype).eval() + + def _build_embeddings_processor(self) -> EmbeddingsProcessor: + """Build the embeddings processor on the target device.""" + return self._embeddings_processor_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval() + def _text_encoder_ctx(self) -> AbstractContextManager: if self._offload_mode != OffloadMode.NONE: return _streaming_model(self._streaming_text_encoder_builder, self._offload_mode, self._device, self._dtype) - return gpu_model(self._text_encoder_builder.build(device=self._device, dtype=self._dtype).eval()) + return gpu_model(self._build_text_encoder()) def __call__( self, @@ -413,9 +456,7 @@ class PromptEncoder: ) raw_outputs = [text_encoder.encode(p) for p in prompts] - with gpu_model( - self._embeddings_processor_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval() - ) as embeddings_processor: + with gpu_model(self._build_embeddings_processor()) as embeddings_processor: return [embeddings_processor.process_hidden_states(hs, mask) for hs, mask in raw_outputs] @@ -513,32 +554,31 @@ class VideoDecoder: dtype: torch.dtype, device: torch.device, registry: Registry | None = None, + memory_efficient: bool = True, + decoder_builder: BuilderProtocol | None = None, ) -> None: self._dtype = dtype self._device = device - self._decoder_builder = Builder( - model_path=checkpoint_path, - model_class_configurator=VideoDecoderConfigurator, - model_sd_ops=VAE_DECODER_COMFY_KEYS_FILTER, - registry=registry or DummyRegistry(), - ) + if decoder_builder is not None: + self._decoder_builder = decoder_builder + else: + self._decoder_builder = Builder( + model_path=checkpoint_path, + model_class_configurator=VideoDecoderConfigurator, + model_sd_ops=VAE_DECODER_COMFY_KEYS_FILTER, + registry=registry or DummyRegistry(), + module_ops=(MEMORY_EFFICIENT_DECODE,) if memory_efficient else (), + ) def __call__( self, latent: torch.Tensor, tiling_config: TilingConfig | None = None, generator: torch.Generator | None = None, - *, - output_dtype: torch.dtype = torch.uint8, ) -> Iterator[torch.Tensor]: - """Decode *latent* to pixel-space video chunks. Decoder freed after exhaustion. - Args: - output_dtype: Target dtype for output tensors. ``torch.uint8`` - (default) maps to ``[0, 255]``. Any floating dtype returns - ``[0, 1]`` cast to that dtype. - """ + """Decode *latent* to pixel-space video chunks. Decoder freed after exhaustion.""" decoder = self._decoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval() - return _cleanup_iter(decoder.decode_video(latent, tiling_config, generator, output_dtype=output_dtype), decoder) + return _cleanup_iter(decoder.decode_video(latent, tiling_config, generator), decoder) # --------------------------------------------------------------------------- diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/color_conversion.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/color_conversion.py new file mode 100644 index 0000000..937acce --- /dev/null +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/color_conversion.py @@ -0,0 +1,224 @@ +"""Color space conversion utilities for video encoding. +Provides GPU-accelerated RGB to YUV420 conversion that runs between the +VAE decoder (which yields float RGB chunks) and ``encode_video``, bypassing +pyav's CPU-side libswscale conversion. The ``FrameConverter`` also carries +the codec metadata (pixel format, colour space, colour range) that +``encode_video`` needs to tag the output stream. +""" + +from __future__ import annotations + +import enum +from collections.abc import Callable +from dataclasses import dataclass, field + +import torch + + +class ColorSpace(enum.Enum): + """YUV color space standard.""" + + BT_709 = "bt709" + BT_2020_NCL = "bt2020ncl" + + @property + def av_colorspace(self) -> int: + """FFmpeg ``AVCOL_SPC_*`` constant for ``codec_context.colorspace``.""" + return _AV_COLORSPACE[self] + + +class ColorRange(enum.Enum): + """YUV color range.""" + + MPEG = "mpeg" + JPEG = "jpeg" + + @property + def av_color_range(self) -> int: + """FFmpeg ``AVCOL_RANGE_*`` constant for ``codec_context.color_range``.""" + return _AV_COLOR_RANGE[self] + + +class PixelFormat(enum.Enum): + """Pixel format for video frames.""" + + RGB24 = "rgb24" + YUV420P = "yuv420p" + + @property + def av_format(self) -> str: + """PyAV format string for ``VideoFrame.from_ndarray``.""" + return self.value + + +_AV_COLORSPACE = { + ColorSpace.BT_709: 1, # AVCOL_SPC_BT709 + ColorSpace.BT_2020_NCL: 9, # AVCOL_SPC_BT2020_NCL +} + +_AV_COLOR_RANGE = { + ColorRange.MPEG: 1, # AVCOL_RANGE_MPEG (limited) + ColorRange.JPEG: 2, # AVCOL_RANGE_JPEG (full) +} + +# BT.709 RGB->YUV matrix (row-major: each row produces one of Y, U, V) +_BT709_MATRIX = torch.tensor( + [ + [0.2126, 0.7152, 0.0722], + [-0.1146, -0.3854, 0.5], + [0.5, -0.4542, -0.0458], + ], + dtype=torch.float32, +) + +# BT.2020 NCL RGB->YUV matrix +_KR_2020 = 0.2627 +_KG_2020 = 0.6780 +_KB_2020 = 0.0593 +_BT2020_MATRIX = torch.tensor( + [ + [_KR_2020, _KG_2020, _KB_2020], + [-_KR_2020 / 1.8814, -_KG_2020 / 1.8814, 0.5], + [0.5, -_KG_2020 / 1.4746, -_KB_2020 / 1.4746], + ], + dtype=torch.float32, +) + +_COLOR_SPACE_MATRICES = { + ColorSpace.BT_709: _BT709_MATRIX, + ColorSpace.BT_2020_NCL: _BT2020_MATRIX, +} + + +@dataclass(frozen=True) +class FrameConverter: + """Converts ``[*, C, H, W]`` float ``[0, 1]`` frames to uint8. + Carries encoding metadata so ``encode_video`` can derive pixel format, + color space, and color range from the converter itself. + The ``fn_`` callable **may mutate its input** (PyTorch trailing-underscore + convention). Callers that need to keep the original ``frames`` afterwards + must pass ``frames.clone()``. Inside ``encode_video``'s per-chunk + generator each chunk is consumed once, so direct passthrough is safe. + """ + + pixel_format: PixelFormat + fn_: Callable[[torch.Tensor], torch.Tensor] = field(repr=False) + color_space: ColorSpace | None = None + color_range: ColorRange | None = None + + def __call__(self, frames: torch.Tensor) -> torch.Tensor: + return self.fn_(frames) + + +def rgb_to_yuv(image: torch.Tensor, color_space: ColorSpace) -> torch.Tensor: + """Convert an RGB image to YUV. + The image data is assumed to be in the range of ``[0, 1]``. + Uses a single matrix multiply for better memory locality. + Args: + image: RGB image with shape ``(*, 3, H, W)``. + color_space: Color space standard for the conversion matrix. + Returns: + YUV image with shape ``(*, 3, H, W)``. + """ + if len(image.shape) < 3 or image.shape[-3] != 3: + raise ValueError(f"Input size must have a shape of (*, 3, H, W). Got {image.shape}") + + mat = _COLOR_SPACE_MATRICES[color_space].to(device=image.device, dtype=image.dtype) + # [*, 3, H, W] -> [*, H, W, 3] @ [3, 3]^T -> [*, H, W, 3] -> [*, 3, H, W] + pixels = image.movedim(-3, -1) # [*, H, W, 3] + yuv = pixels @ mat.T # [*, H, W, 3] + return yuv.movedim(-1, -3) # [*, 3, H, W] + + +def apply_color_range_(y: torch.Tensor, uv: torch.Tensor, color_range: ColorRange) -> tuple[torch.Tensor, torch.Tensor]: + """Scale Y and UV planes to the specified color range, in-place. + Args: + y: Luma plane in ``[0, 1]``. + uv: Chroma planes centered at 0. + color_range: Target color range. + Returns: + Scaled ``(Y, UV)`` tensors (modified in-place). + """ + if color_range == ColorRange.MPEG: + y.mul_(219).add_(16) + uv.mul_(224).add_(128) + elif color_range == ColorRange.JPEG: + y.mul_(255) + uv.add_(0.5).mul_(255) + else: + raise ValueError(f"Unsupported color range: {color_range}") + return y, uv + + +def rgb_to_yuv420( + image: torch.Tensor, color_space: ColorSpace, color_range: ColorRange +) -> tuple[torch.Tensor, torch.Tensor]: + """Convert an RGB image to YUV 4:2:0 with chroma subsampling. + Chroma is subsampled by averaging 2x2 pixel blocks (chroma siting + ``(128, 128)``). + Args: + image: RGB image with shape ``(*, 3, H, W)`` in ``[0, 1]``. + H and W must be divisible by 2. + color_space: Color space standard. + color_range: Color range for the output. + Returns: + ``(Y, UV)`` where Y has shape ``(*, 1, H, W)`` and UV has shape + ``(*, 2, H//2, W//2)``. + """ + if len(image.shape) < 3 or image.shape[-3] != 3: + raise ValueError(f"Input size must have a shape of (*, 3, H, W). Got {image.shape}") + if image.shape[-2] % 2 != 0 or image.shape[-1] % 2 != 0: + raise ValueError(f"Input H and W must be divisible by 2. Got {image.shape}") + + yuv = rgb_to_yuv(image, color_space) + y = yuv[..., :1, :, :] + # Subsample chroma: average 2x2 blocks via avg_pool2d (contiguous, fused kernel) + uv_full = yuv[..., 1:3, :, :].contiguous() + # Flatten leading dims for avg_pool2d which expects [N, C, H, W] + lead = uv_full.shape[:-3] + uv_flat = uv_full.reshape(-1, 2, uv_full.shape[-2], uv_full.shape[-1]) + uv = torch.nn.functional.avg_pool2d(uv_flat, kernel_size=2, stride=2) + uv = uv.reshape(*lead, 2, uv.shape[-2], uv.shape[-1]) + + return apply_color_range_(y, uv, color_range) + + +def pack_i420(y: torch.Tensor, uv: torch.Tensor) -> torch.Tensor: + """Pack Y and UV planes into I420 layout for pyav. + I420 packs the three planes into a single 2D array of height ``H * 3 // 2`` + and width ``W``. The Y plane occupies the first ``H`` rows. The UV tensor + ``(*, 2, H//2, W//2)`` is reshaped to ``(*, H//2, W)`` -- U rows packed + two-by-two followed by V rows packed two-by-two -- and appended below. + Args: + y: Luma with shape ``(*, 1, H, W)``. + uv: Chroma with shape ``(*, 2, H//2, W//2)``. + Returns: + Packed tensor with shape ``(*, H*3//2, W)`` uint8. + """ + y_plane = y[..., 0, :, :] # [*, H, W] + uv_packed = uv.reshape(*uv.shape[:-3], uv.shape[-2], uv.shape[-1] * 2) # [*, H//2, W] + packed = torch.cat([y_plane, uv_packed], dim=-2) # [*, H*3//2, W] + return packed.clamp_(0, 255).to(torch.uint8) + + +def _rgb_uint8_fn_(frames: torch.Tensor) -> torch.Tensor: + """In-place: mutates ``frames`` via ``clamp_`` + ``mul_``, returns a uint8 view.""" + return frames.clamp_(0.0, 1.0).mul_(255.0).to(torch.uint8).movedim(-3, -1) + + +rgb_uint8_converter_ = FrameConverter(pixel_format=PixelFormat.RGB24, fn_=_rgb_uint8_fn_) +"""``(*, 3, H, W)`` float ``[0, 1]`` to ``(*, H, W, 3)`` uint8. Mutates input.""" + + +def _yuv420p_bt709_fn_(frames: torch.Tensor) -> torch.Tensor: + y, uv = rgb_to_yuv420(frames, ColorSpace.BT_709, ColorRange.MPEG) + return pack_i420(y, uv) + + +yuv420p_bt709_converter_ = FrameConverter( + pixel_format=PixelFormat.YUV420P, + fn_=_yuv420p_bt709_fn_, + color_space=ColorSpace.BT_709, + color_range=ColorRange.MPEG, +) +"""``(*, 3, H, W)`` float ``[0, 1]`` to ``(*, H*3//2, W)`` uint8 YUV420p BT.709 MPEG.""" diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/denoisers.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/denoisers.py index edceb20..316f70f 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/denoisers.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/denoisers.py @@ -20,6 +20,7 @@ from ltx_core.guidance.perturbations import ( from ltx_core.model.transformer import X0Model from ltx_core.types import LatentState from ltx_pipelines.utils.helpers import modality_from_latent_state +from ltx_pipelines.utils.types import DenoisedLatentResult _POSITIVE_ONLY_GUIDER = MultiModalGuider( params=MultiModalGuiderParams(cfg_scale=1.0, stg_scale=0.0, modality_scale=1.0), @@ -53,7 +54,7 @@ def _repeat_state(state: LatentState, n: int) -> LatentState: ) -def _guided_denoise( # noqa: PLR0913 +def _guided_denoise( # noqa: PLR0913,PLR0915 transformer: X0Model, video_state: LatentState | None, audio_state: LatentState | None, @@ -66,7 +67,8 @@ def _guided_denoise( # noqa: PLR0913 last_denoised_video: torch.Tensor | None, last_denoised_audio: torch.Tensor | None, step_index: int, -) -> tuple[torch.Tensor | None, torch.Tensor | None]: + force_uncond_pass: bool = False, +) -> tuple[DenoisedLatentResult | None, DenoisedLatentResult | None]: """Core guided denoising — batches all guidance passes into one transformer call. Collects per-pass contexts first, then builds a single batched Modality per present modality via :func:`modality_from_latent_state`. When wrapped @@ -80,7 +82,9 @@ def _guided_denoise( # noqa: PLR0913 a_skip = audio_guider.should_skip_step(step_index) if v_skip and a_skip: - return last_denoised_video, last_denoised_audio + video_result = DenoisedLatentResult.result_or_none(denoised=last_denoised_video) + audio_result = DenoisedLatentResult.result_or_none(denoised=last_denoised_audio) + return video_result, audio_result if video_state is not None and v_context is None: raise ValueError("v_context is required when video_state is provided") @@ -91,10 +95,12 @@ def _guided_denoise( # noqa: PLR0913 _pass = tuple[str, torch.Tensor | None, torch.Tensor | None, PerturbationConfig] passes: list[_pass] = [("cond", v_context, a_context, PerturbationConfig.empty())] - if video_guider.do_unconditional_generation() or audio_guider.do_unconditional_generation(): - if video_guider.do_unconditional_generation() and video_guider.negative_context is None: + v_needs_neg = video_guider.do_unconditional_generation() or (force_uncond_pass and video_state is not None) + a_needs_neg = audio_guider.do_unconditional_generation() or (force_uncond_pass and audio_state is not None) + if v_needs_neg or a_needs_neg: + if v_needs_neg and video_guider.negative_context is None: raise ValueError("Negative context is required for unconditioned denoising") - if audio_guider.do_unconditional_generation() and audio_guider.negative_context is None: + if a_needs_neg and audio_guider.negative_context is None: raise ValueError("Negative context is required for unconditioned denoising") v_neg = video_guider.negative_context if video_guider.negative_context is not None else v_context a_neg = audio_guider.negative_context if audio_guider.negative_context is not None else a_context @@ -172,7 +178,14 @@ def _guided_denoise( # noqa: PLR0913 denoised_video = last_denoised_video if v_skip else video_guider.calculate(cond_v, uncond_v, ptb_v, mod_v) denoised_audio = last_denoised_audio if a_skip else audio_guider.calculate(cond_a, uncond_a, ptb_a, mod_a) - return denoised_video, denoised_audio + return ( + DenoisedLatentResult.result_or_none( + denoised=denoised_video, uncond=uncond_v, cond=cond_v, ptb=ptb_v, mod=mod_v + ), + DenoisedLatentResult.result_or_none( + denoised=denoised_audio, uncond=uncond_a, cond=cond_a, ptb=ptb_a, mod=mod_a + ), + ) class SimpleDenoiser: @@ -195,11 +208,15 @@ class SimpleDenoiser: audio_state: LatentState | None, sigmas: torch.Tensor, step_index: int, - ) -> tuple[torch.Tensor | None, torch.Tensor | None]: + ) -> tuple[DenoisedLatentResult | None, DenoisedLatentResult | None]: sigma = sigmas[step_index] pos_video = modality_from_latent_state(video_state, self.v_context, sigma) if video_state is not None else None pos_audio = modality_from_latent_state(audio_state, self.a_context, sigma) if audio_state is not None else None - return transformer(video=pos_video, audio=pos_audio, perturbations=None) + denoised_video, denoised_audio = transformer(video=pos_video, audio=pos_audio, perturbations=None) + return ( + DenoisedLatentResult.result_or_none(denoised=denoised_video), + DenoisedLatentResult.result_or_none(denoised=denoised_audio), + ) class GuidedDenoiser: @@ -214,11 +231,13 @@ class GuidedDenoiser: a_context: torch.Tensor | None, video_guider: MultiModalGuider | None = None, audio_guider: MultiModalGuider | None = None, + force_uncond_pass: bool = False, ) -> None: self.v_context = v_context self.a_context = a_context self.video_guider = video_guider self.audio_guider = audio_guider + self.force_uncond_pass = force_uncond_pass self._last_denoised_video: torch.Tensor | None = None self._last_denoised_audio: torch.Tensor | None = None @@ -229,8 +248,8 @@ class GuidedDenoiser: audio_state: LatentState | None, sigmas: torch.Tensor, step_index: int, - ) -> tuple[torch.Tensor | None, torch.Tensor | None]: - denoised_video, denoised_audio = _guided_denoise( + ) -> tuple[DenoisedLatentResult | None, DenoisedLatentResult | None]: + guided_denoise_result_v, guided_denoise_result_a = _guided_denoise( transformer=transformer, video_state=video_state, audio_state=audio_state, @@ -242,10 +261,11 @@ class GuidedDenoiser: last_denoised_video=self._last_denoised_video, last_denoised_audio=self._last_denoised_audio, step_index=step_index, + force_uncond_pass=self.force_uncond_pass, ) - self._last_denoised_video = denoised_video - self._last_denoised_audio = denoised_audio - return denoised_video, denoised_audio + self._last_denoised_video = guided_denoise_result_v.denoised + self._last_denoised_audio = guided_denoise_result_a.denoised + return guided_denoise_result_v, guided_denoise_result_a class FactoryGuidedDenoiser: @@ -257,11 +277,13 @@ class FactoryGuidedDenoiser: a_context: torch.Tensor | None, video_guider_factory: MultiModalGuiderFactory | None = None, audio_guider_factory: MultiModalGuiderFactory | None = None, + force_uncond_pass: bool = False, ) -> None: self.v_context = v_context self.a_context = a_context self.video_guider_factory = video_guider_factory self.audio_guider_factory = audio_guider_factory + self.force_uncond_pass = force_uncond_pass self._last_denoised_video: torch.Tensor | None = None self._last_denoised_audio: torch.Tensor | None = None self._sigma_vals_cached: list[float] | None = None @@ -273,7 +295,7 @@ class FactoryGuidedDenoiser: audio_state: LatentState | None, sigmas: torch.Tensor, step_index: int, - ) -> tuple[torch.Tensor | None, torch.Tensor | None]: + ) -> tuple[DenoisedLatentResult | None, DenoisedLatentResult | None]: if self._sigma_vals_cached is None: self._sigma_vals_cached = sigmas.detach().cpu().tolist() sigma_val = self._sigma_vals_cached[step_index] @@ -287,7 +309,7 @@ class FactoryGuidedDenoiser: else None ) - denoised_video, denoised_audio = _guided_denoise( + guided_denoise_result_v, guided_denoise_result_a = _guided_denoise( transformer=transformer, video_state=video_state, audio_state=audio_state, @@ -299,7 +321,8 @@ class FactoryGuidedDenoiser: last_denoised_video=self._last_denoised_video, last_denoised_audio=self._last_denoised_audio, step_index=step_index, + force_uncond_pass=self.force_uncond_pass, ) - self._last_denoised_video = denoised_video - self._last_denoised_audio = denoised_audio - return denoised_video, denoised_audio + self._last_denoised_video = guided_denoise_result_v.denoised + self._last_denoised_audio = guided_denoise_result_a.denoised + return guided_denoise_result_v, guided_denoise_result_a 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 e96c65c..00fb506 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/media_io.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/media_io.py @@ -1,10 +1,12 @@ import enum import logging import math +import threading from collections.abc import Generator, Iterator from fractions import Fraction from io import BytesIO from pathlib import Path +from queue import Queue import av import numpy as np @@ -17,6 +19,7 @@ from tqdm import tqdm from ltx_core.hdr import LogC3 from ltx_core.types import Audio, VideoPixelShape +from ltx_pipelines.utils.color_conversion import FrameConverter, PixelFormat, yuv420p_bt709_converter_ from ltx_pipelines.utils.constants import DEFAULT_IMAGE_CRF logger = logging.getLogger(__name__) @@ -86,8 +89,8 @@ def resize_and_center_crop(tensor: torch.Tensor, height: int, width: int) -> tor return tensor -def normalize_latent(latent: torch.Tensor, device: torch.device, dtype: torch.dtype) -> torch.Tensor: - return (latent / 127.5 - 1.0).to(device=device, dtype=dtype) +def normalize_images(images: torch.Tensor, device: torch.device, dtype: torch.dtype) -> torch.Tensor: + return (images / 127.5 - 1.0).to(device=device, dtype=dtype) def to_vae_range(x: torch.Tensor) -> torch.Tensor: @@ -116,7 +119,7 @@ def load_image_and_preprocess( image = preprocess(image=image, crf=crf) image = torch.tensor(image, dtype=torch.float32, device=device) image = resize_and_center_crop(image, height, width) - image = normalize_latent(image, device, dtype) + image = normalize_images(image, device, dtype) return image @@ -137,11 +140,13 @@ def video_preprocess( Returns: Tensor of shape (1, C, F, height, width) with values in [-1, 1]. """ - result = None + result: torch.Tensor | None = None for f in frames: frame = resize_and_center_crop(f.to(torch.float32), height, width) - frame = normalize_latent(frame, device, dtype) + frame = normalize_images(frame, device, dtype) result = frame if result is None else torch.cat([result, frame], dim=2) + if result is None: + raise ValueError("video_preprocess received an empty frame generator; no frames were decoded from the source.") return result @@ -325,47 +330,120 @@ def encode_video( audio: Audio | None, output_path: str, video_chunks_number: int, + frame_converter: FrameConverter = yuv420p_bt709_converter_, + crf: int = 19, + preset: str = "veryfast", + thread_count: int = 0, ) -> None: if isinstance(video, torch.Tensor): video = iter([video]) - first_chunk = next(video) + def convert(chunk: torch.Tensor) -> torch.Tensor: + return frame_converter(chunk.movedim(-1, -3)) - _, height, width, _ = first_chunk.shape + first_chunk = convert(next(video)) + + if frame_converter.pixel_format == PixelFormat.RGB24: + height, width = first_chunk.shape[-3], first_chunk.shape[-2] + else: + height = first_chunk.shape[-2] * 2 // 3 + width = first_chunk.shape[-1] container = av.open(output_path, mode="w") - stream = container.add_stream("libx264", rate=int(fps)) - stream.width = width - stream.height = height - stream.pix_fmt = "yuv420p" + success = False + try: + stream = container.add_stream("libx264", rate=int(fps), options={"crf": str(crf), "preset": preset}) + stream.width = width + stream.height = height + stream.pix_fmt = "yuv420p" + stream.codec_context.thread_count = thread_count + stream.codec_context.thread_type = "FRAME" + if frame_converter.color_space is not None: + stream.codec_context.colorspace = frame_converter.color_space.av_colorspace + if frame_converter.color_range is not None: + stream.codec_context.color_range = frame_converter.color_range.av_color_range - if audio is not None: - audio_stream = _prepare_audio_stream(container, audio.sampling_rate) + if audio is not None: + audio_stream = _prepare_audio_stream(container, audio.sampling_rate) - def all_tiles( - first_chunk: torch.Tensor, tiles_generator: Generator[tuple[torch.Tensor, int], None, None] - ) -> Generator[tuple[torch.Tensor, int], None, None]: - yield first_chunk - yield from tiles_generator + av_format = frame_converter.pixel_format.av_format - for video_chunk in tqdm(all_tiles(first_chunk, video), total=video_chunks_number): - video_chunk_cpu = video_chunk.to("cpu").numpy() - for frame_array in video_chunk_cpu: - frame = av.VideoFrame.from_ndarray(frame_array, format="rgb24") - for packet in stream.encode(frame): - container.mux(packet) + def cpu_chunks() -> Generator[np.ndarray, None, None]: + yield first_chunk.to("cpu").numpy() + for chunk in video: + yield convert(chunk).to("cpu").numpy() - # Flush encoder - for packet in stream.encode(): - container.mux(packet) + _encode_chunks_threaded( + container=container, + stream=stream, + av_format=av_format, + chunks=cpu_chunks(), + progress_total=video_chunks_number, + ) - if audio is not None: - _write_audio(container, audio_stream, audio) - - container.close() + if audio is not None: + _write_audio(container, audio_stream, audio) + success = True + finally: + container.close() + if not success: + Path(output_path).unlink(missing_ok=True) logger.info(f"Video saved to {output_path}") +def _encode_chunks_threaded( + container: av.container.Container, + stream: av.video.stream.VideoStream, + av_format: str, + chunks: Iterator[np.ndarray], + progress_total: int, +) -> None: + """Run libx264 frame.encode + container.mux on a background thread while + the caller produces numpy chunks on the current thread. The 1-slot queue + lets the producer get one chunk ahead (so the next VAE/gather chunk + overlaps with libx264 encoding the previous chunk) without buffering more + than one chunk in CPU memory. + """ + chunk_queue: Queue[np.ndarray | None] = Queue(maxsize=1) + encoder_error: list[BaseException] = [] + + def encoder_worker() -> None: + error: BaseException | None = None + while True: + arr = chunk_queue.get() + if arr is None: + break + if error is not None: + continue + try: + for frame_array in arr: + frame = av.VideoFrame.from_ndarray(frame_array, format=av_format) + for packet in stream.encode(frame): + container.mux(packet) + except Exception as e: + error = e + if error is None: + try: + for packet in stream.encode(): + container.mux(packet) + except Exception as e: + error = e + if error is not None: + encoder_error.append(error) + + encoder_thread = threading.Thread(target=encoder_worker, name="h264-encoder") + encoder_thread.start() + try: + for arr in tqdm(chunks, total=progress_total): + chunk_queue.put(arr) + finally: + chunk_queue.put(None) + encoder_thread.join() + + if encoder_error: + raise encoder_error[0] + + _INT_FORMAT_MAX: dict[str, float] = { "u8": 128.0, "u8p": 128.0, diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/samplers.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/samplers.py index 467d1c1..130bc98 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/samplers.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/samplers.py @@ -6,7 +6,7 @@ from typing import Callable import torch from tqdm import tqdm -from ltx_core.components.diffusion_steps import Res2sDiffusionStep +from ltx_core.components.diffusion_steps import EulerCfgPpDiffusionStep, Res2sDiffusionStep from ltx_core.components.protocols import DiffusionStepProtocol from ltx_core.model.transformer import X0Model from ltx_core.utils import to_denoised, to_velocity @@ -60,13 +60,15 @@ def euler_denoising_loop( denoiser: A callable implementing :class:`Denoiser`. It is invoked as ``denoiser(transformer, video_state, audio_state, sigmas, step_index)`` - and must return ``(denoised_video, denoised_audio)``. + and must return a :class:`~ltx_pipelines.utils.types.DenoisedLatentResult`. ### Returns tuple[LatentState | None, LatentState | None] Final ``(video_state, audio_state)`` after the denoising loop. """ for step_idx, _ in enumerate(tqdm(sigmas[:-1])): - denoised_video, denoised_audio = denoiser(transformer, video_state, audio_state, sigmas, step_idx) + video_result, audio_result = denoiser(transformer, video_state, audio_state, sigmas, step_idx) + denoised_video = video_result.denoised if video_result is not None else None + denoised_audio = audio_result.denoised if audio_result is not None else None video_state = _step_state(video_state, denoised_video, stepper, sigmas, step_idx) audio_state = _step_state(audio_state, denoised_audio, stepper, sigmas, step_idx) @@ -110,7 +112,9 @@ def gradient_estimating_euler_denoising_loop( return current_velocity, denoised_sample for step_idx, _ in enumerate(tqdm(sigmas[:-1])): - denoised_video, denoised_audio = denoiser(transformer, video_state, audio_state, sigmas, step_idx) + video_result, audio_result = denoiser(transformer, video_state, audio_state, sigmas, step_idx) + denoised_video = video_result.denoised if video_result is not None else None + denoised_audio = audio_result.denoised if audio_result is not None else None if video_state is not None and denoised_video is not None: denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent) @@ -143,6 +147,11 @@ def gradient_estimating_euler_denoising_loop( return (video_state, audio_state) +def _get_plain_noise(x: torch.Tensor, generator: torch.Generator) -> torch.Tensor: + """Draw standard Gaussian noise matching the shape, dtype, and device of ``x``.""" + return torch.randn(x.shape, generator=generator, dtype=x.dtype, device=x.device) + + def _channelwise_normalize(x: torch.Tensor) -> torch.Tensor: return x.sub_(x.mean(dim=(-2, -1), keepdim=True)).div_(x.std(dim=(-2, -1), keepdim=True)) @@ -278,7 +287,9 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915,PLR0912 # ==================================================================== # STAGE 1: Evaluate at current point # ==================================================================== - denoised_video_1, denoised_audio_1 = denoiser(transformer, video_state, audio_state, sigmas, step_idx) + video_result, audio_result = denoiser(transformer, video_state, audio_state, sigmas, step_idx) + denoised_video_1 = video_result.denoised if video_result is not None else None + denoised_audio_1 = audio_result.denoised if audio_result is not None else None if video_state is not None and denoised_video_1 is not None: denoised_video_1 = post_process_latent(denoised_video_1, video_state.denoise_mask, video_state.clean_latent) if audio_state is not None and denoised_audio_1 is not None: @@ -355,13 +366,15 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915,PLR0912 else None ) - denoised_video_2, denoised_audio_2 = denoiser( + video_result_2, audio_result_2 = denoiser( transformer, video_state=mid_video_state, audio_state=mid_audio_state, sigmas=torch.stack([sub_sigma]).to(sigmas.device), step_index=0, ) + denoised_video_2 = video_result_2.denoised if video_result_2 is not None else None + denoised_audio_2 = audio_result_2.denoised if audio_result_2 is not None else None if video_state is not None and denoised_video_2 is not None: denoised_video_2 = post_process_latent(denoised_video_2, video_state.denoise_mask, video_state.clean_latent) if audio_state is not None and denoised_audio_2 is not None: @@ -410,7 +423,9 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915,PLR0912 # Final step if we need to fully remove the noise if sigmas[-1] == 0: - denoised_video_1, denoised_audio_1 = denoiser(transformer, video_state, audio_state, sigmas, n_full_steps) + video_result_final, audio_result_final = denoiser(transformer, video_state, audio_state, sigmas, n_full_steps) + denoised_video_1 = video_result_final.denoised if video_result_final is not None else None + denoised_audio_1 = audio_result_final.denoised if audio_result_final is not None else None if video_state is not None and denoised_video_1 is not None: denoised_video_1 = post_process_latent(denoised_video_1, video_state.denoise_mask, video_state.clean_latent) video_state = replace(video_state, latent=denoised_video_1.to(model_dtype)) @@ -419,3 +434,121 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915,PLR0912 audio_state = replace(audio_state, latent=denoised_audio_1.to(model_dtype)) return video_state, audio_state + + +def euler_cfg_pp_denoising_loop( + sigmas: torch.Tensor, + video_state: LatentState | None, + audio_state: LatentState | None, + stepper: EulerCfgPpDiffusionStep, + transformer: X0Model, + denoiser: Denoiser, + noise_seed: int = -1, + new_noise_fn: Callable[[torch.Tensor, torch.Generator], torch.Tensor] = _get_plain_noise, + model_dtype: torch.dtype = torch.bfloat16, +) -> tuple[LatentState | None, LatentState | None]: + """ + Joint audio-video denoising loop using the CFG++ corrected Euler sampler. + Applies the CFG++ update rule at each step: the ODE derivative is computed + from the unconditioned denoised prediction rather than the standard velocity, + and an ancestral DDIM noise injection is applied in the rescaled sigma space. + Requires a guided denoiser whose :class:`~ltx_pipelines.utils.types.DenoisedLatentResult` + carries ``uncond`` tensors (i.e. CFG must be enabled). + Either ``video_state`` or ``audio_state`` may be ``None`` for absent modalities. + When both are present, noise is drawn from the same seeded generator (video + first, audio second) to produce a consistent random sequence. + ### Parameters + sigmas: + 1-D tensor of noise levels defining the sampling schedule. + video_state: + Current video :class:`~ltx_core.types.LatentState`, or ``None``. + audio_state: + Current audio :class:`~ltx_core.types.LatentState`, or ``None``. + stepper: + :class:`~ltx_core.components.diffusion_steps.EulerCfgPpDiffusionStep` + instance carrying ``eta`` and ``s_noise`` parameters. + transformer: + The diffusion model passed to the denoiser at each step. + denoiser: + Callable implementing :class:`~ltx_pipelines.utils.types.Denoiser`. + noise_seed: + Integer seed for the noise generator. Default ``-1``. + new_noise_fn: + ``(latent, generator) -> noise`` callable. Defaults to plain + ``torch.randn`` (no channel-wise normalization). Pass + :func:`_get_new_noise` for the normalized variant used in res2s. + model_dtype: + Dtype for latent state updates. Default ``bfloat16``. + ### Returns + tuple[LatentState | None, LatentState | None] + Final ``(video_state, audio_state)`` after the denoising loop. + """ + if not isinstance(stepper, EulerCfgPpDiffusionStep): + raise ValueError(f"stepper must be an instance of EulerCfgPpDiffusionStep, got {type(stepper).__name__}") + + present_state = video_state or audio_state + if present_state is None: + raise ValueError("At least one of video_state or audio_state must be provided") + + generator = torch.Generator(device=present_state.latent.device).manual_seed(noise_seed) + draw_noise = stepper.eta > 0 and stepper.s_noise > 0 + + for step_idx, _ in enumerate(tqdm(sigmas[:-1])): + video_result, audio_result = denoiser(transformer, video_state, audio_state, sigmas, step_idx) + denoised_video = video_result.denoised if video_result is not None else None + denoised_audio = audio_result.denoised if audio_result is not None else None + uncond_video = video_result.uncond if video_result is not None else None + uncond_audio = audio_result.uncond if audio_result is not None else None + + if video_state is not None and not isinstance(uncond_video, torch.Tensor): + raise ValueError( + "euler_cfg_pp_denoising_loop requires video DenoisedLatentResult.uncond to be a tensor. " + "Use GuidedDenoiser or FactoryGuidedDenoiser with cfg_scale != 1 " + "or force_uncond_pass=True and a negative_context." + ) + if audio_state is not None and not isinstance(uncond_audio, torch.Tensor): + raise ValueError( + "euler_cfg_pp_denoising_loop requires audio DenoisedLatentResult.uncond to be a tensor. " + "Use GuidedDenoiser or FactoryGuidedDenoiser with cfg_scale != 1 " + "or force_uncond_pass=True and a negative_context." + ) + + if video_state is not None and denoised_video is not None: + denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent) + if audio_state is not None and denoised_audio is not None: + denoised_audio = post_process_latent(denoised_audio, audio_state.denoise_mask, audio_state.clean_latent) + + if sigmas[step_idx + 1] == 0: + if video_state is not None and denoised_video is not None: + video_state = replace(video_state, latent=denoised_video.to(model_dtype)) + if audio_state is not None and denoised_audio is not None: + audio_state = replace(audio_state, latent=denoised_audio.to(model_dtype)) + return video_state, audio_state + + # Draw noise consecutively from the same generator: video first, audio second. + noise_video = new_noise_fn(video_state.latent, generator) if (video_state is not None and draw_noise) else None + noise_audio = new_noise_fn(audio_state.latent, generator) if (audio_state is not None and draw_noise) else None + + if video_state is not None and denoised_video is not None: + x_next = stepper.step( + sample=video_state.latent, + denoised_sample=denoised_video, + sigmas=sigmas, + step_index=step_idx, + uncond_denoised=uncond_video, + noise=noise_video, + ) + video_state = replace(video_state, latent=x_next.to(model_dtype)) + + if audio_state is not None and denoised_audio is not None: + x_next = stepper.step( + sample=audio_state.latent, + denoised_sample=denoised_audio, + sigmas=sigmas, + step_index=step_idx, + uncond_denoised=uncond_audio, + noise=noise_audio, + ) + audio_state = replace(audio_state, latent=x_next.to(model_dtype)) + + return video_state, audio_state diff --git a/packages/ltx-pipelines/src/ltx_pipelines/utils/types.py b/packages/ltx-pipelines/src/ltx_pipelines/utils/types.py index 6093ead..8ea3727 100644 --- a/packages/ltx-pipelines/src/ltx_pipelines/utils/types.py +++ b/packages/ltx-pipelines/src/ltx_pipelines/utils/types.py @@ -40,6 +40,36 @@ class PipelineComponents: self.audio_patchifier = AudioPatchifier(patch_size=1) +@dataclass(frozen=True) +class DenoisedLatentResult: + """Output of one denoiser call for a single modality. + ``denoised`` is the final blended prediction for this modality. + The remaining fields carry the per-pass raw outputs from ``_guided_denoise`` + (all ``None`` for ``SimpleDenoiser``). Denoisers return a + ``(video_result, audio_result)`` tuple; either element may be ``None`` + for absent modalities. + """ + + denoised: torch.Tensor + uncond: torch.Tensor | None = None + cond: torch.Tensor | None = None + ptb: torch.Tensor | None = None + mod: torch.Tensor | None = None + + @classmethod + def result_or_none( + cls, + denoised: torch.Tensor | None, + uncond: torch.Tensor | None = None, + cond: torch.Tensor | None = None, + ptb: torch.Tensor | None = None, + mod: torch.Tensor | None = None, + ) -> DenoisedLatentResult | None: + if denoised is None: + return None + return cls(denoised=denoised, uncond=uncond, cond=cond, ptb=ptb, mod=mod) + + class Denoiser(Protocol): """Protocol for a denoiser that receives the transformer at call time. The transformer is not stored — it is passed as the first argument so the @@ -51,7 +81,8 @@ class Denoiser(Protocol): sigmas: 1-D tensor of sigma values for each diffusion step. step_index: Index of the current denoising step. Returns: - ``(denoised_video, denoised_audio)`` tensors (either may be ``None``). + A ``(video_result, audio_result)`` tuple of :class:`DenoisedLatentResult`, + either may be ``None`` for absent modalities. """ def __call__( @@ -61,7 +92,7 @@ class Denoiser(Protocol): audio_state: LatentState | None, sigmas: torch.Tensor, step_index: int, - ) -> tuple[torch.Tensor | None, torch.Tensor | None]: ... + ) -> tuple[DenoisedLatentResult | None, DenoisedLatentResult | None]: ... @dataclass(frozen=True) diff --git a/packages/ltx-trainer/configs/ltx2_av_lora.yaml b/packages/ltx-trainer/configs/ltx2_av_lora.yaml index ded5347..45e3a61 100644 --- a/packages/ltx-trainer/configs/ltx2_av_lora.yaml +++ b/packages/ltx-trainer/configs/ltx2_av_lora.yaml @@ -166,6 +166,11 @@ acceleration: # Useful when GPU memory is limited load_text_encoder_in_8bit: false + # Offload optimizer state to CPU during validation video sampling. + # Helps avoid OOM when VAE decoder + transformer + optimizer state can't coexist + # on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP. + offload_optimizer_during_validation: false + # ----------------------------------------------------------------------------- # Data Configuration # ----------------------------------------------------------------------------- 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 065bdb7..1868171 100644 --- a/packages/ltx-trainer/configs/ltx2_av_lora_low_vram.yaml +++ b/packages/ltx-trainer/configs/ltx2_av_lora_low_vram.yaml @@ -178,6 +178,11 @@ acceleration: # Useful when GPU memory is limited load_text_encoder_in_8bit: true + # Offload optimizer state to CPU during validation video sampling. + # Helps avoid OOM when VAE decoder + transformer + optimizer state can't coexist + # on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP. + offload_optimizer_during_validation: true + # ----------------------------------------------------------------------------- # Data Configuration # ----------------------------------------------------------------------------- diff --git a/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml b/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml index dbaaf0e..4470ad9 100644 --- a/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml +++ b/packages/ltx-trainer/configs/ltx2_v2v_ic_lora.yaml @@ -166,6 +166,11 @@ acceleration: # Useful when GPU memory is limited load_text_encoder_in_8bit: false + # Offload optimizer state to CPU during validation video sampling. + # Helps avoid OOM when VAE decoder + transformer + optimizer state can't coexist + # on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP. + offload_optimizer_during_validation: false + # ----------------------------------------------------------------------------- # Data Configuration # ----------------------------------------------------------------------------- diff --git a/packages/ltx-trainer/docs/configuration-reference.md b/packages/ltx-trainer/docs/configuration-reference.md index 81ce562..d40961f 100644 --- a/packages/ltx-trainer/docs/configuration-reference.md +++ b/packages/ltx-trainer/docs/configuration-reference.md @@ -215,18 +215,20 @@ Hardware acceleration and compute optimization settings. ```yaml acceleration: - mixed_precision_mode: "bf16" # "no", "fp16", or "bf16" - quantization: null # Quantization options - load_text_encoder_in_8bit: false # Load text encoder in 8-bit + mixed_precision_mode: "bf16" # "no", "fp16", or "bf16" + quantization: null # Quantization options + load_text_encoder_in_8bit: false # Load text encoder in 8-bit + offload_optimizer_during_validation: false # Offload optimizer state to CPU during validation ``` **Key parameters:** -| Parameter | Description | -|-----------------------------|------------------------------------------------------------------------------------| -| `mixed_precision_mode` | Precision mode - `"bf16"` recommended for modern GPUs | -| `quantization` | Model quantization: `null`, `"int8-quanto"`, `"int4-quanto"`, `"fp8-quanto"`, etc. | -| `load_text_encoder_in_8bit` | Load the Gemma text encoder in 8-bit to save GPU memory | +| Parameter | Description | +|---------------------------------------|------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| +| `mixed_precision_mode` | Precision mode - `"bf16"` recommended for modern GPUs | +| `quantization` | Model quantization: `null`, `"int8-quanto"`, `"int4-quanto"`, `"fp8-quanto"`, etc. | +| `load_text_encoder_in_8bit` | Load the Gemma text encoder in 8-bit to save GPU memory | +| `offload_optimizer_during_validation` | Move optimizer state to CPU before validation video sampling and back afterwards. Useful when validation OOMs because VAE decoder + transformer + optimizer state can't coexist on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP. | ### DataConfig diff --git a/packages/ltx-trainer/docs/dataset-preparation.md b/packages/ltx-trainer/docs/dataset-preparation.md index 53528ce..6b8b151 100644 --- a/packages/ltx-trainer/docs/dataset-preparation.md +++ b/packages/ltx-trainer/docs/dataset-preparation.md @@ -50,17 +50,20 @@ This will create a `dataset.json` file containing video paths and their captions **Captioning options:** -| Option | Description | -|--------|-------------| -| `--captioner-type` | `qwen_omni` (default, local) or `gemini_flash` (API) | -| `--use-8bit` | Enable 8-bit quantization for lower VRAM usage | -| `--no-audio` | Disable audio processing (video-only captions) | -| `--override` | Re-caption files that already have captions | -| `--api-key` | API key for Gemini Flash (or set `GOOGLE_API_KEY` env var) | + +| Option | Description | +| ------------------ | ---------------------------------------------------------- | +| `--captioner-type` | `qwen_omni` (default, local) or `gemini_flash` (API) | +| `--use-8bit` | Enable 8-bit quantization for lower VRAM usage | +| `--no-audio` | Disable audio processing (video-only captions) | +| `--override` | Re-caption files that already have captions | +| `--api-key` | API key for Gemini Flash (or set `GOOGLE_API_KEY` env var) | + **Caption format:** The captioner produces structured captions with sections for: + - **Visual content**: People, objects, actions, settings, colors, movements - **Speech transcription**: Word-for-word transcription of spoken content - **Sounds**: Music, ambient sounds, sound effects @@ -106,15 +109,57 @@ uv run python scripts/process_dataset.py dataset.json \ --with-audio ``` +### 🚀 Multi-GPU Preprocessing + +Preprocessing large datasets can take a while. To run it across multiple GPUs in parallel, wrap the command with +`accelerate launch` (for example `--num_processes 4`). Each process handles an interleaved slice of the dataset. +The same approach applies to `process_videos.py` and `process_captions.py` when you run them standalone. + +```bash +uv run accelerate launch --num_processes 4 scripts/process_dataset.py dataset.json \ + --resolution-buckets "960x544x49" \ + --model-path /path/to/ltx-2-model.safetensors \ + --text-encoder-path /path/to/gemma-model +``` + +Outputs are written atomically (via a per-process temporary file, then renamed), so an interrupted run leaves no +corrupt files. By default a rerun **resumes** — items whose output `.pt` already exists are skipped. + +> [!IMPORTANT] +> Pass `**--overwrite`** when rerunning with changed parameters (different model checkpoint, resolution buckets, +> text encoder, `--lora-trigger`, etc.). Without it the script keeps the stale outputs from the previous run. +> +> ```bash +> uv run accelerate launch --num_processes 4 scripts/process_dataset.py dataset.json \ +> --resolution-buckets "960x544x49" \ +> --model-path /path/to/ltx-2.3-model.safetensors \ +> --text-encoder-path /path/to/gemma-model \ +> --overwrite +> ``` + ### 📊 Dataset Format -The trainer supports either videos or single images. -Note that your dataset must be homogeneous - either all videos or all images, mixing is not supported. +The trainer supports videos, single images, or a mix of both in the same dataset. > [!TIP] > **Image Datasets:** When using images, follow the same preprocessing steps and format requirements as with videos, > but use `1` for the frame count in the resolution bucket (e.g., `960x544x1`). +> [!NOTE] +> **Mixed image + video datasets:** Mixing stills and videos in a single dataset is supported, but requires some care: +> +> - Preprocess with **multiple resolution buckets** covering both frame counts — e.g. +> `--resolution-buckets "960x544x1;960x544x49"`. Images are automatically assigned to the `F=1` bucket and +> videos to an `F>1` bucket. +> - You **must** set `optimization.batch_size: 1` in your training config (see the warning under +> [Resolution Buckets](#-resolution-buckets)), since samples with different shapes cannot be collated into a +> single batch. Use `gradient_accumulation_steps` if you need a larger effective batch. +> - Per-step cost differs substantially between a single-frame sample and a many-frame sample, which can lead to +> uneven gradient magnitudes across steps. Consider weighting the two subsets or tuning the learning rate if +> you observe instability. +> - If you prefer a fully officially-supported path, train two separate LoRAs (one on stills, one on video) and +> stack them at inference. + The dataset must be a CSV, JSON, or JSONL metadata file with columns for captions and video paths: **JSON format example:** @@ -197,6 +242,7 @@ uv run python scripts/process_dataset.py dataset.json \ > ``` > > Where: +> > - H = Height of video > - W = Width of video > - F = Number of frames @@ -204,6 +250,7 @@ uv run python scripts/process_dataset.py dataset.json \ > - 8 = VAE's temporal downsampling factor > > For example, a 768×448×89 video would have sequence length: +> > ``` > (768/32) * (448/32) * ((89-1)/8 + 1) = 24 * 14 * 12 = 4,032 > ``` @@ -268,7 +315,6 @@ uv run python scripts/process_dataset.py dataset.json \ This will create an additional `reference_latents/` directory containing the preprocessed reference video latents. - ### Generating Reference Videos **Dataset Requirements for IC-LoRA:** @@ -277,7 +323,7 @@ This will create an additional `reference_latents/` directory containing the pre - Reference and target videos must have *identical* resolution and length - Both reference and target videos should be preprocessed together using the same resolution buckets -We provide an example script, [`scripts/compute_reference.py`](../scripts/compute_reference.py), to generate reference +We provide an example script, `[scripts/compute_reference.py](../scripts/compute_reference.py)`, to generate reference videos for a given dataset. The default implementation generates Canny edge reference videos. ```bash @@ -293,7 +339,6 @@ If you want to generate a different type of condition (depth maps, pose skeleton For reference, see our **[Canny Control Dataset](https://huggingface.co/datasets/Lightricks/Canny-Control-Dataset)** which demonstrates proper IC-LoRA dataset structure with paired videos and Canny edge maps. - ## 🎯 LoRA Trigger Words When training a LoRA, you can specify a trigger token that will be prepended to all captions: diff --git a/packages/ltx-trainer/docs/troubleshooting.md b/packages/ltx-trainer/docs/troubleshooting.md index aeb66c6..068b268 100644 --- a/packages/ltx-trainer/docs/troubleshooting.md +++ b/packages/ltx-trainer/docs/troubleshooting.md @@ -84,6 +84,20 @@ optimization: optimizer_type: "adamw8bit" ``` +#### 7. Offload Optimizer State During Validation + +If you OOM specifically during validation video sampling — typically in +full fine-tunes or high-rank LoRA runs where AdamW state and the VAE decoder +can't coexist on the GPU — offload optimizer state to CPU during sampling: + +```yaml +acceleration: + offload_optimizer_during_validation: true +``` + +The offload + reload happens once per validation interval, not per step. +No effect for FSDP (sharded state). + --- ## ⚠️ Common Usage Issues diff --git a/packages/ltx-trainer/docs/utility-scripts.md b/packages/ltx-trainer/docs/utility-scripts.md index e25920c..1838268 100644 --- a/packages/ltx-trainer/docs/utility-scripts.md +++ b/packages/ltx-trainer/docs/utility-scripts.md @@ -143,6 +143,27 @@ uv run python scripts/process_dataset.py dataset.json \ > [!NOTE] > When training with multiple resolution buckets, set `optimization.batch_size: 1`. +**Multi-GPU preprocessing.** Launch with `accelerate launch` to shard the dataset across processes. Reruns resume +by default (existing `.pt` outputs are skipped); writes are atomic so interrupted runs are safe. Pass `--overwrite` +when rerunning with changed parameters (different model, resolution buckets, text encoder, `--lora-trigger`, etc.) +so stale outputs are replaced. Use the same `accelerate launch` pattern (and `--overwrite` when needed) with +`process_videos.py` or `process_captions.py` when you run those scripts standalone. + +```bash +# Multi-GPU preprocessing +uv run accelerate launch --num_processes 4 scripts/process_dataset.py dataset.json \ + --resolution-buckets "960x544x49" \ + --model-path /path/to/ltx-2-model.safetensors \ + --text-encoder-path /path/to/gemma-model + +# Force re-encoding of all items (e.g. after switching model or resolution) +uv run accelerate launch --num_processes 4 scripts/process_dataset.py dataset.json \ + --resolution-buckets "960x544x49" \ + --model-path /path/to/ltx-2.3-model.safetensors \ + --text-encoder-path /path/to/gemma-model \ + --overwrite +``` + For detailed usage, see the [Dataset Preparation Guide](dataset-preparation.md). ### Reference Video Generation diff --git a/packages/ltx-trainer/pyproject.toml b/packages/ltx-trainer/pyproject.toml index 801d73e..986e27e 100644 --- a/packages/ltx-trainer/pyproject.toml +++ b/packages/ltx-trainer/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "ltx-trainer" -version = "1.1.2" +version = "1.1.3" description = "LTX-2 training, democratized." readme = "README.md" authors = [ @@ -48,7 +48,7 @@ build-backend = "hatchling.build" [tool.ruff] -target-version = "1.1.2" +target-version = "1.1.3" line-length = 120 [tool.ruff.lint] diff --git a/packages/ltx-trainer/scripts/process_captions.py b/packages/ltx-trainer/scripts/process_captions.py index bc9759a..9523209 100755 --- a/packages/ltx-trainer/scripts/process_captions.py +++ b/packages/ltx-trainer/scripts/process_captions.py @@ -13,12 +13,14 @@ Can be used as a standalone script: import json import os +from collections.abc import Callable from pathlib import Path from typing import Any import pandas as pd import torch import typer +from accelerate import PartialState from rich.console import Console from rich.progress import ( BarColumn, @@ -30,7 +32,7 @@ from rich.progress import ( TimeElapsedColumn, TimeRemainingColumn, ) -from torch.utils.data import DataLoader, Dataset +from torch.utils.data import DataLoader, Dataset, Subset from transformers.utils.logging import disable_progress_bar from ltx_trainer import logger @@ -232,9 +234,14 @@ def compute_captions_embeddings( # noqa: PLR0913 batch_size: int = 8, device: str = "cuda", load_in_8bit: bool = False, + overwrite: bool = False, ) -> None: """ Process captions and save text embeddings. + Under ``accelerate launch``, each process handles an interleaved shard of + the dataset (rank/world read from ``accelerate.PartialState``). Already- + computed ``.pt`` outputs are skipped unless ``overwrite=True``; writes are + atomic so an interrupted run is safe to resume. Args: dataset_file: Path to metadata file (CSV/JSON/JSONL) containing captions and media paths output_dir: Directory to save embeddings @@ -247,11 +254,12 @@ def compute_captions_embeddings( # noqa: PLR0913 batch_size: Batch size for processing device: Device to use for computation load_in_8bit: Whether to load the Gemma text encoder in 8-bit precision + overwrite: Re-encode every item even if its output exists. Use when rerunning with + changed parameters (different text encoder, lora_trigger, etc.) so stale + outputs are replaced. """ - console = Console() - # Create dataset dataset = CaptionsDataset( dataset_file=dataset_file, caption_column=caption_column, @@ -264,6 +272,24 @@ def compute_captions_embeddings( # noqa: PLR0913 output_path = Path(output_dir) output_path.mkdir(parents=True, exist_ok=True) + # TODO(batch-tokenization): The current Gemma tokenizer doesn't support batched tokenization. + if batch_size > 1: + logger.warning( + "Batch size greater than 1 is not currently supported with the Gemma tokenizer. " + "Overriding batch_size to 1. This will be fixed in a future update." + ) + batch_size = 1 + + dataloader = _build_sharded_dataloader( + dataset, + batch_size=batch_size, + num_workers=2, + is_done=lambda idx: (output_path / dataset.output_paths[idx]).is_file(), + overwrite=overwrite, + ) + if dataloader is None: + return + # Load text encoder and embeddings processor with console.status("[bold]Loading Gemma text encoder...", spinner="dots"): text_encoder = load_text_encoder( @@ -279,21 +305,7 @@ def compute_captions_embeddings( # noqa: PLR0913 ) logger.info("Text encoder and embeddings processor loaded successfully") - - # TODO(batch-tokenization): The current Gemma tokenizer doesn't support batched tokenization. - if batch_size > 1: - logger.warning( - "Batch size greater than 1 is not currently supported with the Gemma tokenizer. " - "Overriding batch_size to 1. This will be fixed in a future update." - ) - batch_size = 1 - - # Create dataloader - dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=False, num_workers=2) - - # Process batches - total_batches = len(dataloader) - logger.info(f"Processing captions in {total_batches:,} batches...") + logger.info(f"Processing captions in {len(dataloader):,} batches...") with Progress( SpinnerColumn(), @@ -333,11 +345,44 @@ def compute_captions_embeddings( # noqa: PLR0913 embedding_data["audio_prompt_embeds"] = audio_prompt_embeds[0].cpu().contiguous() output_file = output_path / output_rel_path - torch.save(embedding_data, output_file) + _atomic_save(embedding_data, output_file) progress.advance(task) - logger.info(f"Processed {len(dataset):,} captions. Embeddings saved to {output_path}") + logger.info(f"Processed {len(dataloader.dataset):,} captions -> {output_path}") # type: ignore[arg-type] + + +def _atomic_save(data: Any, out: Path) -> None: # noqa: ANN401 + """Save to ``out`` atomically via per-PID temp file + replace. + Crash mid-write leaves an orphan ``.tmp.`` file that the skip logic + ignores. The per-PID suffix makes concurrent writes from multiple ranks + collision-free. + """ + tmp = out.with_suffix(f"{out.suffix}.tmp.{os.getpid()}") + torch.save(data, tmp) + tmp.replace(out) + + +def _build_sharded_dataloader( + dataset: Dataset, + *, + batch_size: int, + num_workers: int, + is_done: Callable[[int], bool], + overwrite: bool, +) -> DataLoader | None: + """Return a DataLoader over this rank's interleaved shard of ``dataset``. + When ``overwrite`` is False, items whose outputs already exist (per + ``is_done``) are filtered out. Returns ``None`` if this rank has nothing + to do, so the caller can early-return without loading any models. + """ + state = PartialState() + todo = [i for i in range(state.process_index, len(dataset), state.num_processes) if overwrite or not is_done(i)] + if not todo: + logger.info(f"Rank {state.process_index}/{state.num_processes}: nothing to do") + return None + logger.info(f"Rank {state.process_index}/{state.num_processes}: processing {len(todo):,} of {len(dataset):,} items") + return DataLoader(Subset(dataset, todo), batch_size=batch_size, shuffle=False, num_workers=num_workers) @app.command() @@ -387,8 +432,15 @@ def main( # noqa: PLR0913 default=False, help="Load the Gemma text encoder in 8-bit precision to save GPU memory (requires bitsandbytes)", ), + overwrite: bool = typer.Option( + default=False, + help="Re-encode every caption even if its output exists. Use when rerunning with " + "changed parameters (different text encoder, lora_trigger, etc.) so stale outputs are replaced.", + ), ) -> None: """Process text captions and save embeddings for video generation training. + For multi-GPU preprocessing, invoke under ``accelerate launch`` - each process + will handle an interleaved shard of the dataset. This script processes captions from metadata files and saves text embeddings that can be used for training video generation models. The output embeddings will maintain the same folder structure and naming as the corresponding media files. @@ -428,6 +480,7 @@ def main( # noqa: PLR0913 batch_size=batch_size, device=device, load_in_8bit=load_text_encoder_in_8bit, + overwrite=overwrite, ) diff --git a/packages/ltx-trainer/scripts/process_dataset.py b/packages/ltx-trainer/scripts/process_dataset.py index 22c257b..fa827b3 100755 --- a/packages/ltx-trainer/scripts/process_dataset.py +++ b/packages/ltx-trainer/scripts/process_dataset.py @@ -50,6 +50,7 @@ def preprocess_dataset( # noqa: PLR0913 reference_downscale_factor: int = 1, with_audio: bool = False, load_text_encoder_in_8bit: bool = False, + overwrite: bool = False, ) -> None: """Run the preprocessing pipeline with the given arguments.""" # Validate dataset file @@ -77,6 +78,7 @@ def preprocess_dataset( # noqa: PLR0913 batch_size=batch_size, device=device, load_in_8bit=load_text_encoder_in_8bit, + overwrite=overwrite, ) # Process videos using the dedicated function @@ -97,6 +99,7 @@ def preprocess_dataset( # noqa: PLR0913 vae_tiling=vae_tiling, with_audio=with_audio, audio_output_dir=str(audio_latents_dir) if audio_latents_dir else None, + overwrite=overwrite, ) # Process reference videos if reference_column is provided @@ -133,6 +136,7 @@ def preprocess_dataset( # noqa: PLR0913 batch_size=batch_size, device=device, vae_tiling=vae_tiling, + overwrite=overwrite, ) # Handle decoding if requested (for verification) @@ -252,8 +256,15 @@ def main( # noqa: PLR0913 help="Downscale factor for reference video resolution. When > 1, reference videos are processed at " "1/n resolution (e.g., 2 means half resolution). Used for efficient IC-LoRA training.", ), + overwrite: bool = typer.Option( + default=False, + help="Re-compute every item even if its output exists. Use when rerunning with " + "changed parameters (different model, resolution, etc.) so stale outputs are replaced.", + ), ) -> None: """Preprocess a video dataset by computing and saving latents and text embeddings. + For multi-GPU preprocessing, invoke under ``accelerate launch`` - each process + will handle an interleaved shard of the dataset. The dataset must be a CSV, JSON, or JSONL file with columns for captions and video paths. This script is designed for LTX-2 models which use the Gemma text encoder. Examples: @@ -310,6 +321,7 @@ def main( # noqa: PLR0913 reference_downscale_factor=reference_downscale_factor, with_audio=with_audio, load_text_encoder_in_8bit=load_text_encoder_in_8bit, + overwrite=overwrite, ) diff --git a/packages/ltx-trainer/scripts/process_videos.py b/packages/ltx-trainer/scripts/process_videos.py index 40c0c08..c33a815 100755 --- a/packages/ltx-trainer/scripts/process_videos.py +++ b/packages/ltx-trainer/scripts/process_videos.py @@ -14,6 +14,8 @@ Can be used as a standalone script: import json import math +import os +from collections.abc import Callable from pathlib import Path from typing import Any @@ -22,6 +24,7 @@ import pandas as pd import torch import torchaudio import typer +from accelerate import PartialState from pillow_heif import register_heif_opener from rich.console import Console from rich.progress import ( @@ -34,7 +37,7 @@ from rich.progress import ( TimeElapsedColumn, TimeRemainingColumn, ) -from torch.utils.data import DataLoader, Dataset +from torch.utils.data import DataLoader, Dataset, Subset from torchvision import transforms from torchvision.transforms import InterpolationMode from torchvision.transforms.functional import crop, resize, to_tensor @@ -444,9 +447,14 @@ def compute_latents( # noqa: PLR0913, PLR0915 vae_tiling: bool = False, with_audio: bool = False, audio_output_dir: str | None = None, + overwrite: bool = False, ) -> None: """ Process videos and save latent representations. + Under ``accelerate launch``, each process handles an interleaved shard of + the dataset (rank/world read from ``accelerate.PartialState``). Already- + computed ``.pt`` outputs are skipped unless ``overwrite=True``; writes are + atomic so an interrupted run is safe to resume. Args: dataset_file: Path to metadata file (CSV/JSON/JSONL) containing video paths video_column: Column name for video paths in the metadata file @@ -460,15 +468,15 @@ def compute_latents( # noqa: PLR0913, PLR0915 vae_tiling: Whether to enable VAE tiling with_audio: Whether to extract and encode audio from videos audio_output_dir: Directory to save audio latents (required if with_audio=True) + overwrite: Re-process every item even if its output exists. Use when rerunning with + changed parameters (different model, resolution, etc.) so stale outputs are replaced. """ - # Validate audio parameters if with_audio and audio_output_dir is None: raise ValueError("audio_output_dir must be provided when with_audio=True") console = Console() torch_device = torch.device(device) - # Create dataset dataset = MediaDataset( dataset_file=dataset_file, main_media_column=main_media_column or video_column, @@ -481,13 +489,34 @@ def compute_latents( # noqa: PLR0913, PLR0915 output_path = Path(output_dir) output_path.mkdir(parents=True, exist_ok=True) - - # Set up audio output directory if needed - audio_output_path = None + audio_output_path: Path | None = None if with_audio: audio_output_path = Path(audio_output_dir) audio_output_path.mkdir(parents=True, exist_ok=True) + # Audio processing requires batch_size=1; must be applied before the dataloader is built. + if with_audio and batch_size > 1: + logger.warning("Audio processing requires batch_size=1. Overriding batch_size to 1.") + batch_size = 1 + + data_root = dataset.dataset_file.parent + + def _is_done(idx: int) -> bool: + rel = dataset.main_media_paths[idx].relative_to(data_root).with_suffix(".pt") + if not (output_path / rel).is_file(): + return False + return audio_output_path is None or (audio_output_path / rel).is_file() + + dataloader = _build_sharded_dataloader( + dataset, + batch_size=batch_size, + num_workers=4, + is_done=_is_done, + overwrite=overwrite, + ) + if dataloader is None: + return + # Load video VAE encoder with console.status(f"[bold]Loading video VAE encoder from [cyan]{model_path}[/]...", spinner="dots"): vae = load_video_vae_encoder(model_path, device=torch_device, dtype=torch.bfloat16) @@ -510,14 +539,6 @@ def compute_latents( # noqa: PLR0913, PLR0915 n_fft=audio_vae_encoder.n_fft, ).to(torch_device) - # Create dataloader - # Note: batch_size=1 required when with_audio because audio extraction can fail for some videos, - # and the default collate function can't handle mixed None/dict values across a batch. - if with_audio and batch_size > 1: - logger.warning("Audio processing requires batch_size=1. Overriding batch_size to 1.") - batch_size = 1 - dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=False, num_workers=4) - # Track audio statistics audio_success_count = 0 audio_skip_count = 0 @@ -560,7 +581,7 @@ def compute_latents( # noqa: PLR0913, PLR0915 "fps": batch["video_metadata"]["fps"][i].item(), } - torch.save(latent_data, output_file) + _atomic_save(latent_data, output_file) # Process audio if enabled (audio is already extracted by the dataset) if with_audio: @@ -588,7 +609,7 @@ def compute_latents( # noqa: PLR0913, PLR0915 "duration": audio_latents["duration"], } - torch.save(audio_save_data, audio_output_file) + _atomic_save(audio_save_data, audio_output_file) audio_success_count += 1 else: # Video has no audio track @@ -596,8 +617,7 @@ def compute_latents( # noqa: PLR0913, PLR0915 progress.advance(task) - # Log summary - logger.info(f"Processed {len(dataset)} videos. Latents saved to {output_path}") + logger.info(f"Processed {len(dataloader.dataset)} videos -> {output_path}") # type: ignore[arg-type] if with_audio: logger.info( f"Audio processing: {audio_success_count} videos with audio, " @@ -935,6 +955,39 @@ def compute_scaled_resolution_buckets( return scaled_buckets +def _atomic_save(data: Any, out: Path) -> None: # noqa: ANN401 + """Save to ``out`` atomically via per-PID temp file + replace. + Crash mid-write leaves an orphan ``.tmp.`` file that the skip logic + ignores. The per-PID suffix makes concurrent writes from multiple ranks + collision-free. + """ + tmp = out.with_suffix(f"{out.suffix}.tmp.{os.getpid()}") + torch.save(data, tmp) + tmp.replace(out) + + +def _build_sharded_dataloader( + dataset: Dataset, + *, + batch_size: int, + num_workers: int, + is_done: Callable[[int], bool], + overwrite: bool, +) -> DataLoader | None: + """Return a DataLoader over this rank's interleaved shard of ``dataset``. + When ``overwrite`` is False, items whose outputs already exist (per + ``is_done``) are filtered out. Returns ``None`` if this rank has nothing + to do, so the caller can early-return without loading any models. + """ + state = PartialState() + todo = [i for i in range(state.process_index, len(dataset), state.num_processes) if overwrite or not is_done(i)] + if not todo: + logger.info(f"Rank {state.process_index}/{state.num_processes}: nothing to do") + return None + logger.info(f"Rank {state.process_index}/{state.num_processes}: processing {len(todo):,} of {len(dataset):,} items") + return DataLoader(Subset(dataset, todo), batch_size=batch_size, shuffle=False, num_workers=num_workers) + + @app.command() def main( # noqa: PLR0913 dataset_file: str = typer.Argument( @@ -981,8 +1034,15 @@ def main( # noqa: PLR0913 default=None, help="Output directory for audio latents (required if --with-audio is set)", ), + overwrite: bool = typer.Option( + default=False, + help="Re-encode every item even if its output exists. Use when rerunning with " + "changed parameters (different model, resolution, etc.) so stale outputs are replaced.", + ), ) -> None: """Process videos/images and save latent representations for video generation training. + For multi-GPU preprocessing, invoke under ``accelerate launch`` - each process + will handle an interleaved shard of the dataset. This script processes videos and images from metadata files and saves latent representations that can be used for training video generation models. The output latents will maintain the same folder structure and naming as the corresponding media files. @@ -1032,6 +1092,7 @@ def main( # noqa: PLR0913 vae_tiling=vae_tiling, with_audio=with_audio, audio_output_dir=audio_output_dir, + overwrite=overwrite, ) diff --git a/packages/ltx-trainer/src/ltx_trainer/config.py b/packages/ltx-trainer/src/ltx_trainer/config.py index 751f4fd..3220b4d 100644 --- a/packages/ltx-trainer/src/ltx_trainer/config.py +++ b/packages/ltx-trainer/src/ltx_trainer/config.py @@ -168,6 +168,15 @@ class AccelerationConfig(ConfigBaseModel): description="Whether to load the text encoder in 8-bit precision to save memory", ) + offload_optimizer_during_validation: bool = Field( + default=False, + description="Offload optimizer state to CPU before validation video sampling and reload " + "it afterwards, to free VRAM for inference. Useful when optimizer state is large " + "(e.g. AdamW for full fine-tuning or high-rank LoRA) and validation OOMs because the " + "VAE decoder + transformer + optimizer state cannot coexist on the GPU. Has no effect " + "for FSDP (sharded state). Disabled by default.", + ) + class DataConfig(ConfigBaseModel): """Configuration for data loading and processing""" diff --git a/packages/ltx-trainer/src/ltx_trainer/config_display.py b/packages/ltx-trainer/src/ltx_trainer/config_display.py index a80b1eb..c9ddb20 100644 --- a/packages/ltx-trainer/src/ltx_trainer/config_display.py +++ b/packages/ltx-trainer/src/ltx_trainer/config_display.py @@ -85,6 +85,7 @@ def print_config(config: LtxTrainerConfig) -> None: ("Mixed Precision", accel.mixed_precision_mode or "[dim]—[/]"), ("Quantization", str(accel.quantization) if accel.quantization else "[dim]—[/]"), ("Text Encoder 8bit", fmt(accel.load_text_encoder_in_8bit)), + ("Optimizer CPU Offload", fmt(accel.offload_optimizer_during_validation)), ], ), ( diff --git a/packages/ltx-trainer/src/ltx_trainer/gemma_8bit.py b/packages/ltx-trainer/src/ltx_trainer/gemma_8bit.py index 813454e..8372790 100644 --- a/packages/ltx-trainer/src/ltx_trainer/gemma_8bit.py +++ b/packages/ltx-trainer/src/ltx_trainer/gemma_8bit.py @@ -12,6 +12,7 @@ Example usage: from __future__ import annotations import logging +import os from collections.abc import Generator from contextlib import contextmanager from pathlib import Path @@ -22,7 +23,11 @@ from ltx_core.text_encoders.gemma.encoders.base_encoder import GemmaTextEncoder from ltx_core.text_encoders.gemma.tokenizer import LTXVGemmaTokenizer -def load_8bit_gemma(gemma_model_path: str | Path, dtype: torch.dtype = torch.bfloat16) -> GemmaTextEncoder: +def load_8bit_gemma( + gemma_model_path: str | Path, + dtype: torch.dtype = torch.bfloat16, + device: torch.device | str | int | None = None, +) -> GemmaTextEncoder: """Load the Gemma text encoder in 8-bit precision using bitsandbytes. Only the Gemma LLM backbone is loaded here. The embeddings processor (feature extractor + connectors) should be loaded separately via @@ -30,6 +35,10 @@ def load_8bit_gemma(gemma_model_path: str | Path, dtype: torch.dtype = torch.bfl Args: gemma_model_path: Path to Gemma model directory dtype: Data type for non-quantized model weights + device: Device to place the quantized model on. When ``None`` (default), + the device is inferred from ``LOCAL_RANK`` if CUDA is available, so + multi-process launches put each rank's encoder on its own GPU + instead of all colliding on ``cuda:0``. Returns: GemmaTextEncoder with 8-bit quantized Gemma backbone Raises: @@ -46,13 +55,23 @@ def load_8bit_gemma(gemma_model_path: str | Path, dtype: torch.dtype = torch.bfl gemma_path = _find_gemma_subpath(gemma_model_path, "model*.safetensors") tokenizer_path = _find_gemma_subpath(gemma_model_path, "tokenizer.model") + # Pin the entire model to a single device. `device_map="auto"` collides on cuda:0 + # in multi-process launches because every rank picks the same default device. + device_map: str | dict[str, int | str | torch.device] + if device is not None: + device_map = {"": device} + elif torch.cuda.is_available(): + device_map = {"": int(os.environ.get("LOCAL_RANK", "0"))} + else: + device_map = "auto" + quantization_config = BitsAndBytesConfig(load_in_8bit=True) with _suppress_accelerate_memory_warnings(): gemma_model = Gemma3ForConditionalGeneration.from_pretrained( gemma_path, quantization_config=quantization_config, torch_dtype=torch.bfloat16, - device_map="auto", + device_map=device_map, local_files_only=True, ) diff --git a/packages/ltx-trainer/src/ltx_trainer/model_loader.py b/packages/ltx-trainer/src/ltx_trainer/model_loader.py index c2bd52a..f6aeba6 100644 --- a/packages/ltx-trainer/src/ltx_trainer/model_loader.py +++ b/packages/ltx-trainer/src/ltx_trainer/model_loader.py @@ -199,8 +199,6 @@ def load_text_encoder( device: Device to load model on dtype: Data type for model weights load_in_8bit: Whether to load the Gemma model in 8-bit precision using bitsandbytes. - When True, the model is loaded with device_map="auto" and the device argument - is ignored for the Gemma backbone. Returns: Loaded GemmaTextEncoder """ @@ -211,7 +209,7 @@ def load_text_encoder( if load_in_8bit: from ltx_trainer.gemma_8bit import load_8bit_gemma - return load_8bit_gemma(gemma_model_path, dtype) + return load_8bit_gemma(gemma_model_path, dtype, device=device) # Standard loading path from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder diff --git a/packages/ltx-trainer/src/ltx_trainer/trainer.py b/packages/ltx-trainer/src/ltx_trainer/trainer.py index dda3d60..92cfec4 100644 --- a/packages/ltx-trainer/src/ltx_trainer/trainer.py +++ b/packages/ltx-trainer/src/ltx_trainer/trainer.py @@ -1,7 +1,10 @@ +import contextlib +import math import os import re import time import warnings +from collections.abc import Iterator from dataclasses import dataclass from pathlib import Path from typing import Any, Callable @@ -9,8 +12,8 @@ from typing import Any, Callable import torch import wandb import yaml -from accelerate import Accelerator, DistributedType -from accelerate.utils import set_seed +from accelerate import Accelerator, DistributedDataParallelKwargs, DistributedType +from accelerate.utils import gather_object, set_seed from peft import LoraConfig, get_peft_model, get_peft_model_state_dict, set_peft_model_state_dict from peft.tuners.tuners_utils import BaseTunerLayer from peft.utils import ModulesToSaveWrapper @@ -63,7 +66,7 @@ if not IS_MAIN_PROCESS: disable_progress_bar() -StepCallback = Callable[[int, int, list[Path]], None] # (step, total, list[sampled_video_path]) -> None +StepCallback = Callable[[int, int, list[Path] | None], None] # (step, total, sampled paths or None) -> None MEMORY_CHECK_INTERVAL = 200 @@ -186,9 +189,8 @@ class LtxvTrainer: with progress: if cfg.validation.interval and not cfg.validation.skip_initial_validation: - sampled_videos_paths = self._sample_videos(progress) - if IS_MAIN_PROCESS and sampled_videos_paths and self._config.wandb.log_validation_videos: - self._log_validation_samples(sampled_videos_paths, cfg.validation.prompts) + with self._offloaded_optimizer_state(): + sampled_videos_paths = self._run_distributed_validation(progress) self._accelerator.wait_for_everyone() @@ -228,16 +230,8 @@ class LtxvTrainer: and self._global_step % cfg.validation.interval == 0 and is_optimization_step ): - if self._accelerator.distributed_type == DistributedType.FSDP: - # FSDP: All processes must participate in validation - sampled_videos_paths = self._sample_videos(progress) - if IS_MAIN_PROCESS and sampled_videos_paths and self._config.wandb.log_validation_videos: - self._log_validation_samples(sampled_videos_paths, cfg.validation.prompts) - # DDP: Only main process runs validation - elif IS_MAIN_PROCESS: - sampled_videos_paths = self._sample_videos(progress) - if sampled_videos_paths and self._config.wandb.log_validation_videos: - self._log_validation_samples(sampled_videos_paths, cfg.validation.prompts) + with self._offloaded_optimizer_state(): + sampled_videos_paths = self._run_distributed_validation(progress) # Save checkpoint if needed if ( @@ -398,11 +392,14 @@ class LtxvTrainer: # 3. If validation prompts are configured, computes and caches their embeddings # 4. Unloads the Gemma model entirely, keeps the embeddings processor for training - # Load text encoder (pure Gemma LLM) on GPU + # Load text encoder (pure Gemma LLM) on GPU — LOCAL_RANK before Accelerator exists + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + init_device = torch.device(f"cuda:{local_rank}" if torch.cuda.is_available() else "cpu") + logger.debug("Loading text encoder...") text_encoder = load_text_encoder( gemma_model_path=self._config.model.text_encoder_path, - device="cuda", + device=init_device, dtype=torch.bfloat16, load_in_8bit=self._config.acceleration.load_text_encoder_in_8bit, ) @@ -411,7 +408,7 @@ class LtxvTrainer: logger.debug("Loading embeddings processor...") self._embeddings_processor = load_embeddings_processor( checkpoint_path=self._config.model.model_path, - device="cuda", + device=init_device, dtype=torch.bfloat16, ) @@ -788,6 +785,41 @@ class LtxvTrainer: # noinspection PyTypeChecker self._optimizer, self._lr_scheduler = self._accelerator.prepare(optimizer, lr_scheduler) + @contextlib.contextmanager + def _offloaded_optimizer_state(self) -> Iterator[None]: + """Context manager that offloads optimizer state to CPU during validation. + Opt-in via `acceleration.offload_optimizer_during_validation`. Frees VRAM for + validation video generation when optimizer state is large (e.g. full fine-tune + AdamW, high-rank LoRA). No-op for FSDP (sharded state -- manual `.cpu()` breaks + metadata). + """ + enabled = ( + self._config.acceleration.offload_optimizer_during_validation + and self._accelerator.distributed_type != DistributedType.FSDP + ) + + # Track exactly which tensors we move so we don't promote ones that were + # intentionally on CPU (e.g. AdamW's `step` scalar on recent PyTorch). + offloaded: list[tuple[dict, str]] = [] + if enabled: + offloaded_bytes = 0 + for state in self._optimizer.state.values(): + for k, v in state.items(): + if isinstance(v, torch.Tensor) and v.is_cuda: + offloaded.append((state, k)) + offloaded_bytes += v.nbytes + if offloaded: + logger.info(f"Offloading optimizer state to CPU ({offloaded_bytes / 1e9:.1f} GB)") + for state, k in offloaded: + state[k] = state[k].cpu() + + try: + yield + finally: + device = self._accelerator.device + for state, k in offloaded: + state[k] = state[k].to(device) + def _create_scheduler(self, optimizer: torch.optim.Optimizer) -> LRScheduler | None: """Create learning rate scheduler based on config.""" scheduler_type = self._config.optimization.scheduler_type @@ -844,11 +876,18 @@ class LtxvTrainer: def _setup_accelerator(self) -> None: """Initialize the Accelerator with the appropriate settings.""" + # find_unused_parameters=True keeps DDP happy when LoRA targets a branch the forward + # pass skips (e.g. audio LoRA with `with_audio: false`, or short module patterns like + # "to_k" that match the audio branch unintentionally). It's a no-op for FSDP and + # single-GPU runs. The probing cost is paid only on the first step. + ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True) + # All distributed setup (DDP/FSDP, number of processes, etc.) is controlled by # the user's Accelerate configuration (accelerate config / accelerate launch). self._accelerator = Accelerator( mixed_precision=self._config.acceleration.mixed_precision_mode, gradient_accumulation_steps=self._config.optimization.gradient_accumulation_steps, + kwargs_handlers=[ddp_kwargs], ) if self._accelerator.num_processes > 1: @@ -881,11 +920,42 @@ class LtxvTrainer: "Monitor training stability and consider disabling quantization if issues arise." ) + def _run_distributed_validation(self, progress: TrainingProgress) -> list[Path]: + """Run validation across all ranks and log gathered results on rank 0. + Each rank generates only its assigned subset of prompts (see `_sample_videos`), + so all GPUs stay busy and no rank idles long enough to trigger NCCL timeouts. + Paths are gathered across ranks so rank 0 has the full list for W&B logging. + Note: Multi-node training requires a shared filesystem so rank 0 can read + videos written by other ranks. + """ + sampled = self._sample_videos(progress) + + if self._accelerator.num_processes > 1: + # gather_object returns a flat list from all ranks + sampled = sorted(gather_object(sampled), key=lambda x: x[0]) + + paths = [p for _, p in sampled] + + if self._accelerator.is_main_process and paths: + self._log_validation_samples(paths, self._config.validation.prompts) + + # Non-main ranks must not reach checkpoint collectives while main is still logging to W&B. + self._accelerator.wait_for_everyone() + + return paths + # Note: Use @torch.no_grad() instead of @torch.inference_mode() to avoid FSDP inplace update errors after validation @torch.no_grad() @free_gpu_memory_context(after=True) - def _sample_videos(self, progress: TrainingProgress) -> list[Path] | None: - """Run validation by generating videos from validation prompts.""" + def _sample_videos(self, progress: TrainingProgress) -> list[tuple[int, Path]]: + """Run validation by generating videos from this rank's share of the validation prompts. + Prompts are split round-robin across ranks via `process_index` / `num_processes`, + which collapses to "all prompts" when running on a single GPU. Returns + (prompt_idx, path) tuples so the caller can reconstruct global order without + relying on filename conventions. + Under FSDP with multiple processes, ranks pad with extra generate passes (same prompt, + no disk write) so every rank runs the same number of forwards — avoids collective mismatch. + """ use_images = self._config.validation.images is not None use_reference_videos = self._config.validation.reference_videos is not None generate_audio = self._config.validation.generate_audio @@ -895,13 +965,24 @@ class LtxvTrainer: self._optimizer.zero_grad(set_to_none=True) free_gpu_memory() - # Start sampling progress tracking + prompts = self._config.validation.prompts + rank = self._accelerator.process_index + world_size = self._accelerator.num_processes + rank_indices = list(range(rank, len(prompts), world_size)) + + # FSDP: every rank must run the same number of forwards; pad with duplicate generates (no save). + work: list[tuple[int, bool]] = [(i, True) for i in rank_indices] + if self._accelerator.distributed_type == DistributedType.FSDP and world_size > 1: + max_per_rank = math.ceil(len(prompts) / world_size) + pad_seed = rank_indices[-1] if rank_indices else 0 + work += [(pad_seed, False)] * (max_per_rank - len(work)) + sampling_ctx = progress.start_sampling( - num_prompts=len(self._config.validation.prompts), + num_prompts=len(work), num_steps=inference_steps, ) - # Create validation sampler with loaded models and progress tracking + # Create a validation sampler with loaded models and progress tracking sampler = ValidationSampler( transformer=self._transformer, vae_decoder=self._vae_decoder, @@ -915,12 +996,12 @@ class LtxvTrainer: output_dir = Path(self._config.output_dir) / "samples" output_dir.mkdir(exist_ok=True, parents=True) - video_paths = [] + results: list[tuple[int, Path]] = [] width, height, num_frames = self._config.validation.video_dims - for prompt_idx, prompt in enumerate(self._config.validation.prompts): - # Update progress to show current video - sampling_ctx.start_video(prompt_idx) + for local_i, (prompt_idx, save_output) in enumerate(work): + prompt = prompts[prompt_idx] + sampling_ctx.start_video(local_i) # Load conditioning image if provided condition_image = None @@ -972,28 +1053,30 @@ class LtxvTrainer: device=self._accelerator.device, ) + if not save_output: + continue + # Save output (image for single frame, video otherwise) - if IS_MAIN_PROCESS: - ext = "png" if num_frames == 1 else "mp4" - output_path = output_dir / f"step_{self._global_step:06d}_{prompt_idx + 1}.{ext}" - if num_frames == 1: - save_image(video, output_path) - else: - save_video( - video_tensor=video, - output_path=output_path, - fps=self._config.validation.frame_rate, - audio=audio, - audio_sample_rate=self._vocoder.output_sampling_rate if audio is not None else None, - ) - video_paths.append(output_path) + ext = "png" if num_frames == 1 else "mp4" + output_path = output_dir / f"step_{self._global_step:06d}_{prompt_idx + 1:02d}.{ext}" + if num_frames == 1: + save_image(video, output_path) + else: + save_video( + video_tensor=video, + output_path=output_path, + fps=self._config.validation.frame_rate, + audio=audio, + audio_sample_rate=self._vocoder.output_sampling_rate if audio is not None else None, + ) + results.append((prompt_idx, output_path)) # Clean up progress tasks sampling_ctx.cleanup() rel_outputs_path = output_dir.relative_to(self._config.output_dir) logger.info(f"🎥 Validation samples for step {self._global_step} saved in {rel_outputs_path}") - return video_paths + return results @staticmethod def _log_training_stats(stats: TrainingStats) -> None: diff --git a/uv.lock b/uv.lock index 3bef8ca..af48c4f 100644 --- a/uv.lock +++ b/uv.lock @@ -2063,7 +2063,7 @@ wheels = [ [[package]] name = "ltx-core" -version = "1.1.2" +version = "1.1.3" source = { editable = "packages/ltx-core" } dependencies = [ { name = "accelerate" }, @@ -2121,7 +2121,7 @@ dev = [{ name = "scikit-image", specifier = ">=0.25.2" }] [[package]] name = "ltx-pipelines" -version = "1.1.2" +version = "1.1.3" source = { editable = "packages/ltx-pipelines" } dependencies = [ { name = "av" }, @@ -2144,7 +2144,7 @@ requires-dist = [ [[package]] name = "ltx-trainer" -version = "1.1.2" +version = "1.1.3" source = { editable = "packages/ltx-trainer" } dependencies = [ { name = "accelerate" }, @@ -7105,7 +7105,7 @@ dependencies = [ { name = "torch", version = "2.9.1", source = { registry = "https://pypi.org/simple" } }, ] wheels = [ - { url = "https://download.pytorch.org/whl/cu129/xformers-0.0.33%2B5d4b92a5.d20251029-cp39-abi3-linux_x86_64.whl", upload-time = "2025-10-30T00:15:46Z" }, + { url = "https://download.pytorch.org/whl/cu129/xformers-0.0.33%2B5d4b92a5.d20251029-cp39-abi3-linux_x86_64.whl", hash = "sha256:4e4f2dea153b60ca4f21cc44b82c072b97b096a16f1920c7539c91ec8ffd7ba4", upload-time = "2026-04-27T19:01:32Z" }, ] [[package]]