Automated PR - 2026-04-13
This commit is contained in:
@@ -0,0 +1,3 @@
|
|||||||
|
BasedOnStyle: LLVM
|
||||||
|
IndentWidth: 2
|
||||||
|
ColumnLimit: 120
|
||||||
@@ -39,3 +39,6 @@ tmp
|
|||||||
*.png
|
*.png
|
||||||
*.wav
|
*.wav
|
||||||
*.webp
|
*.webp
|
||||||
|
|
||||||
|
# Binary files
|
||||||
|
*.so
|
||||||
|
|||||||
@@ -30,17 +30,17 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https://
|
|||||||
|
|
||||||
**LTX-2.3 Model Checkpoint** (choose and download one of the following)
|
**LTX-2.3 Model Checkpoint** (choose and download one of the following)
|
||||||
* [`ltx-2.3-22b-dev.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-22b-dev.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-22b-dev.safetensors)
|
* [`ltx-2.3-22b-dev.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-22b-dev.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-22b-dev.safetensors)
|
||||||
* [`ltx-2.3-22b-distilled.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-22b-distilled.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-22b-distilled.safetensors)
|
* [`ltx-2.3-22b-distilled-1.1.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-22b-distilled-1.1.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-22b-distilled-1.1.safetensors)
|
||||||
|
|
||||||
**Spatial Upscaler** - Required for current two-stage pipeline implementations in this repository
|
**Spatial Upscaler** - Required for current two-stage pipeline implementations in this repository
|
||||||
* [`ltx-2.3-spatial-upscaler-x2-1.0.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-spatial-upscaler-x2-1.0.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-spatial-upscaler-x2-1.0.safetensors)
|
* [`ltx-2.3-spatial-upscaler-x2-1.1.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-spatial-upscaler-x2-1.1.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-spatial-upscaler-x2-1.1.safetensors)
|
||||||
* [`ltx-2.3-spatial-upscaler-x1.5-1.0.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-spatial-upscaler-x1.5-1.0.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-spatial-upscaler-x1.5-1.0.safetensors)
|
* [`ltx-2.3-spatial-upscaler-x1.5-1.0.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-spatial-upscaler-x1.5-1.0.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-spatial-upscaler-x1.5-1.0.safetensors)
|
||||||
|
|
||||||
**Temporal Upscaler** - Supported by the model and will be required for future pipeline implementations
|
**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)
|
* [`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 and ICLoraPipeline)
|
||||||
* [`ltx-2.3-22b-distilled-lora-384.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-22b-distilled-lora-384.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-22b-distilled-lora-384.safetensors)
|
* [`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)
|
**Gemma Text Encoder** (download all assets from the repository)
|
||||||
* [`Gemma 3`](https://huggingface.co/google/gemma-3-12b-it-qat-q4_0-unquantized/tree/main)
|
* [`Gemma 3`](https://huggingface.co/google/gemma-3-12b-it-qat-q4_0-unquantized/tree/main)
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "ltx-core"
|
name = "ltx-core"
|
||||||
version = "1.1.0"
|
version = "1.1.1"
|
||||||
description = "Core implementation of Lightricks' LTX-2 model"
|
description = "Core implementation of Lightricks' LTX-2 model"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
requires-python = ">=3.10"
|
requires-python = ">=3.10"
|
||||||
|
|||||||
@@ -44,32 +44,26 @@ def _resolve_attr(module: nn.Module, dotted_path: str) -> nn.ModuleList:
|
|||||||
|
|
||||||
|
|
||||||
class _LayerStore:
|
class _LayerStore:
|
||||||
"""Manages on-demand pinning of layer parameters for GPU streaming.
|
"""Manages CPU-pinned copies of layer parameters/buffers.
|
||||||
Stores references to each layer's source data (which may be file-backed
|
Tracks which layers currently reside on GPU so the prefetcher and evictor
|
||||||
mmap views or in-memory tensors). When a layer needs to be transferred
|
can make correct decisions.
|
||||||
to GPU, its source data is pinned on demand and copied; on eviction the
|
|
||||||
pinned copy is freed and the source data is restored.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, layers: nn.ModuleList, target_device: torch.device) -> None:
|
def __init__(self, layers: nn.ModuleList, target_device: torch.device) -> None:
|
||||||
self.target_device = target_device
|
self.target_device = target_device
|
||||||
self.num_layers = len(layers)
|
self.num_layers = len(layers)
|
||||||
|
|
||||||
|
# CPU-pinned copies keyed by (layer_idx, param_name)
|
||||||
|
self._pinned: list[dict[str, torch.Tensor]] = []
|
||||||
self._on_gpu: set[int] = set()
|
self._on_gpu: set[int] = set()
|
||||||
|
|
||||||
# Keep a reference to the source data for each layer so we can pin it
|
|
||||||
# on demand and restore it after eviction.
|
|
||||||
self._source_data: list[dict[str, torch.Tensor]] = []
|
|
||||||
for layer in layers:
|
for layer in layers:
|
||||||
source: dict[str, torch.Tensor] = {}
|
pinned: dict[str, torch.Tensor] = {}
|
||||||
for name, tensor in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
for name, tensor in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
||||||
source[name] = tensor.data
|
pinned_tensor = tensor.data.pin_memory()
|
||||||
self._source_data.append(source)
|
tensor.data = pinned_tensor
|
||||||
|
pinned[name] = pinned_tensor
|
||||||
# Hold pinned tensors alive until the H2D transfer completes.
|
self._pinned.append(pinned)
|
||||||
# Without this, the CachingHostAllocator can reclaim a pinned tensor
|
|
||||||
# as soon as its Python reference is dropped, even if an async H2D
|
|
||||||
# transfer is still reading from it.
|
|
||||||
self._pinned_in_flight: dict[int, list[torch.Tensor]] = {}
|
|
||||||
|
|
||||||
def _check_idx(self, idx: int) -> None:
|
def _check_idx(self, idx: int) -> None:
|
||||||
if idx < 0 or idx >= self.num_layers:
|
if idx < 0 or idx >= self.num_layers:
|
||||||
@@ -79,45 +73,34 @@ class _LayerStore:
|
|||||||
return idx in self._on_gpu
|
return idx in self._on_gpu
|
||||||
|
|
||||||
def move_to_gpu(self, idx: int, layer: nn.Module, *, non_blocking: bool = False) -> None:
|
def move_to_gpu(self, idx: int, layer: nn.Module, *, non_blocking: bool = False) -> None:
|
||||||
"""Pin layer *idx* on demand, then transfer to GPU."""
|
"""Move layer *idx* parameters from pinned CPU to *target_device*."""
|
||||||
self._check_idx(idx)
|
self._check_idx(idx)
|
||||||
if idx in self._on_gpu:
|
if idx in self._on_gpu:
|
||||||
return
|
return
|
||||||
source = self._source_data[idx]
|
pinned = self._pinned[idx]
|
||||||
pinned_refs: list[torch.Tensor] = []
|
|
||||||
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
||||||
pinned = source[name].pin_memory()
|
param.data = pinned[name].to(self.target_device, non_blocking=non_blocking)
|
||||||
param.data = pinned.to(self.target_device, non_blocking=non_blocking)
|
|
||||||
pinned_refs.append(pinned)
|
|
||||||
# Keep pinned tensors alive until eviction — the async H2D transfer
|
|
||||||
# may still be reading from them.
|
|
||||||
self._pinned_in_flight[idx] = pinned_refs
|
|
||||||
self._on_gpu.add(idx)
|
self._on_gpu.add(idx)
|
||||||
|
|
||||||
def evict_to_cpu(self, idx: int, layer: nn.Module) -> None:
|
def evict_to_cpu(self, idx: int, layer: nn.Module) -> None:
|
||||||
"""Restore source data, freeing the GPU and pinned copies."""
|
"""Swap layer *idx* parameters back to their pinned CPU copies."""
|
||||||
self._check_idx(idx)
|
self._check_idx(idx)
|
||||||
if idx not in self._on_gpu:
|
if idx not in self._on_gpu:
|
||||||
return
|
return
|
||||||
source = self._source_data[idx]
|
pinned = self._pinned[idx]
|
||||||
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
for name, param in itertools.chain(layer.named_parameters(), layer.named_buffers()):
|
||||||
param.data = source[name]
|
param.data = pinned[name]
|
||||||
# Release pinned tensors — the H2D transfer is complete by now
|
|
||||||
# (the compute stream waited on the prefetch event before using
|
|
||||||
# the layer, and we only evict after compute finishes).
|
|
||||||
self._pinned_in_flight.pop(idx, None)
|
|
||||||
self._on_gpu.discard(idx)
|
self._on_gpu.discard(idx)
|
||||||
|
|
||||||
def cleanup(self) -> None:
|
def cleanup(self) -> None:
|
||||||
"""Release all source data and in-flight pinned references.
|
"""Release all pinned memory references.
|
||||||
After this call, the source tensors can be garbage-collected once
|
After this call, the pinned tensors can be garbage-collected once
|
||||||
the layer parameters (which still reference them via ``.data``) are
|
the layer parameters (which still reference them via ``.data``) are
|
||||||
also released (e.g. via ``.to("meta")``).
|
also released (e.g. via ``.to("meta")``).
|
||||||
"""
|
"""
|
||||||
for source_dict in self._source_data:
|
for pinned_dict in self._pinned:
|
||||||
source_dict.clear()
|
pinned_dict.clear()
|
||||||
self._source_data.clear()
|
self._pinned.clear()
|
||||||
self._pinned_in_flight.clear()
|
|
||||||
|
|
||||||
|
|
||||||
class _AsyncPrefetcher:
|
class _AsyncPrefetcher:
|
||||||
@@ -228,8 +211,6 @@ class LayerStreamingWrapper(nn.Module):
|
|||||||
idx_map: dict[int, int] = {id(layer): idx for idx, layer in enumerate(self._layers)}
|
idx_map: dict[int, int] = {id(layer): idx for idx, layer in enumerate(self._layers)}
|
||||||
num_layers = len(self._layers)
|
num_layers = len(self._layers)
|
||||||
|
|
||||||
compute_stream = torch.cuda.current_stream(self._target_device)
|
|
||||||
|
|
||||||
def _pre_hook(
|
def _pre_hook(
|
||||||
module: nn.Module,
|
module: nn.Module,
|
||||||
_args: Any, # noqa: ANN401
|
_args: Any, # noqa: ANN401
|
||||||
@@ -246,6 +227,7 @@ class LayerStreamingWrapper(nn.Module):
|
|||||||
# caching allocator would allow the prefetch stream to reuse their
|
# caching allocator would allow the prefetch stream to reuse their
|
||||||
# memory immediately after eviction — even if the compute kernel
|
# memory immediately after eviction — even if the compute kernel
|
||||||
# that reads them hasn't finished yet.
|
# that reads them hasn't finished yet.
|
||||||
|
compute_stream = torch.cuda.current_stream(self._target_device)
|
||||||
for param in itertools.chain(module.parameters(), module.buffers()):
|
for param in itertools.chain(module.parameters(), module.buffers()):
|
||||||
param.data.record_stream(compute_stream)
|
param.data.record_stream(compute_stream)
|
||||||
|
|
||||||
@@ -270,12 +252,12 @@ class LayerStreamingWrapper(nn.Module):
|
|||||||
self._hooks.extend([h1, h2])
|
self._hooks.extend([h1, h2])
|
||||||
|
|
||||||
def teardown(self) -> None:
|
def teardown(self) -> None:
|
||||||
"""Remove hooks, release resources, and move parameters back to CPU.
|
"""Remove hooks, release pinned memory, and move parameters back to CPU.
|
||||||
After this call the wrapper is inert: hooks are removed, the prefetch
|
After this call the wrapper is inert: hooks are removed, the prefetch
|
||||||
stream is drained and destroyed, all parameters reside on CPU, and the
|
stream is drained and destroyed, all parameters reside on regular
|
||||||
``_LayerStore`` source data references are cleared. Callers should
|
(non-pinned) CPU memory, and the ``_LayerStore`` pinned-tensor cache is
|
||||||
still follow up with ``.to("meta")`` to release the CPU copies if the
|
cleared. Callers should still follow up with ``.to("meta")`` to release
|
||||||
model is no longer needed.
|
the CPU copies if the model is no longer needed.
|
||||||
"""
|
"""
|
||||||
for h in self._hooks:
|
for h in self._hooks:
|
||||||
h.remove()
|
h.remove()
|
||||||
@@ -298,10 +280,10 @@ class LayerStreamingWrapper(nn.Module):
|
|||||||
for b in self._model.buffers():
|
for b in self._model.buffers():
|
||||||
b.data = b.data.to("cpu")
|
b.data = b.data.to("cpu")
|
||||||
|
|
||||||
# Release source data references. After evict_to_cpu() the layer
|
# Release pinned memory. After evict_to_cpu() the layer parameters
|
||||||
# params point to the source data. The caller is expected to follow
|
# still reference the pinned tensors (since .to("cpu") on a pinned
|
||||||
# up with .to("meta") to drop the param refs; cleanup() drops the
|
# tensor is a no-op). The caller is expected to follow up with
|
||||||
# store's refs.
|
# .to("meta") to drop the param refs; cleanup() drops the store's refs.
|
||||||
self._store.cleanup()
|
self._store.cleanup()
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|||||||
@@ -67,11 +67,16 @@ class VideoModalityTilingHelper:
|
|||||||
|
|
||||||
# -- tile modality -----------------------------------------------------
|
# -- tile modality -----------------------------------------------------
|
||||||
|
|
||||||
def tile_modality(self, modality: Modality, tile: Tile) -> tuple[Modality, TilingContext]:
|
def tile_modality(
|
||||||
|
self, modality: Modality, tile: Tile, *, normalize_positions: bool = True
|
||||||
|
) -> tuple[Modality, TilingContext]:
|
||||||
"""Slice *modality* to the tokens covered by *tile*.
|
"""Slice *modality* to the tokens covered by *tile*.
|
||||||
Selects generated tokens belonging to the tile's spatial region
|
Selects generated tokens belonging to the tile's spatial region
|
||||||
and conditioning tokens that overlap with the tile (or have
|
and conditioning tokens that overlap with the tile (or have
|
||||||
negative time coordinates).
|
negative time coordinates).
|
||||||
|
Args:
|
||||||
|
normalize_positions: When True, shift all positions so the
|
||||||
|
tile's generated tokens start at zero in every dimension.
|
||||||
Returns:
|
Returns:
|
||||||
A ``(tiled_modality, context)`` tuple. Pass *context* to
|
A ``(tiled_modality, context)`` tuple. Pass *context* to
|
||||||
:meth:`blend` together with the model output.
|
:meth:`blend` together with the model output.
|
||||||
@@ -83,11 +88,18 @@ class VideoModalityTilingHelper:
|
|||||||
keep_indices = keep_mask.nonzero(as_tuple=False).squeeze(1)
|
keep_indices = keep_mask.nonzero(as_tuple=False).squeeze(1)
|
||||||
tile_attention_mask = modality.attention_mask[:, keep_indices, :][:, :, keep_indices]
|
tile_attention_mask = modality.attention_mask[:, keep_indices, :][:, :, keep_indices]
|
||||||
|
|
||||||
|
positions = modality.positions[:, :, keep_mask, :]
|
||||||
|
if normalize_positions:
|
||||||
|
num_tile_gen = self._tile_generated_token_count(tile)
|
||||||
|
gen_pos = positions[:, :, :num_tile_gen, :] # (B, 3, num_tile_gen, 2)
|
||||||
|
offset = gen_pos[..., 0].amin(dim=2, keepdim=True).unsqueeze(-1) # (B, 3, 1, 1)
|
||||||
|
positions = positions - offset
|
||||||
|
|
||||||
tiled = replace(
|
tiled = replace(
|
||||||
modality,
|
modality,
|
||||||
latent=modality.latent[:, keep_mask, :],
|
latent=modality.latent[:, keep_mask, :],
|
||||||
timesteps=modality.timesteps[:, keep_mask],
|
timesteps=modality.timesteps[:, keep_mask],
|
||||||
positions=modality.positions[:, :, keep_mask, :],
|
positions=positions,
|
||||||
attention_mask=tile_attention_mask,
|
attention_mask=tile_attention_mask,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,82 @@
|
|||||||
|
<!--
|
||||||
|
MAINTENANCE: When modifying any pipeline class in src/ltx_pipelines/,
|
||||||
|
update this document to reflect changes to:
|
||||||
|
- __init__ / __call__ signatures
|
||||||
|
- sigma handling or step counts
|
||||||
|
- denoiser types or guidance
|
||||||
|
- new or removed pipelines
|
||||||
|
Run: ls src/ltx_pipelines/*.py to check for new pipeline files.
|
||||||
|
-->
|
||||||
|
|
||||||
|
# ltx-pipelines
|
||||||
|
|
||||||
|
Inference pipelines for LTX-2 audio-video generation. Depends on `ltx-core` for model definitions, diffusion components, and loading. All pipelines live in `packages/ltx-pipelines/src/ltx_pipelines/`.
|
||||||
|
|
||||||
|
## Pipeline selection
|
||||||
|
|
||||||
|
| Pipeline | File | Stages | Model | Sampler | Use case |
|
||||||
|
|----------|------|--------|-------|---------|----------|
|
||||||
|
| `TI2VidOneStagePipeline` | `ti2vid_one_stage.py` | 1 | Full | Euler | Simple text/image-to-video |
|
||||||
|
| `TI2VidTwoStagesPipeline` | `ti2vid_two_stages.py` | 2 | Full + distilled LoRA | Euler | Production quality |
|
||||||
|
| `TI2VidTwoStagesHQPipeline` | `ti2vid_two_stages_hq.py` | 2 | Full + distilled LoRA (both stages) | Res2s | Highest quality, fewer steps |
|
||||||
|
| `A2VidPipelineTwoStage` | `a2vid_two_stage.py` | 2 | Full + distilled LoRA | Euler | Audio-conditioned video |
|
||||||
|
| `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 |
|
||||||
|
| `RetakePipeline` | `retake.py` | 1 | Full or distilled | Euler | Video region regeneration |
|
||||||
|
|
||||||
|
## Guidance
|
||||||
|
|
||||||
|
- **CFG**: Blends conditioned/unconditioned predictions. Defaults: `cfg_scale=3.0` (video), `7.0` (audio).
|
||||||
|
- **STG**: Perturbs self-attention in transformer blocks. Default `stg_scale=1.0`, `stg_blocks=[28]` (LTX-2.3) / `[29]` (LTX-2). HQ disables STG (`stg_scale=0.0`).
|
||||||
|
- **Modality guidance**: Cross-modal attention scaling (`modality_scale=3.0`).
|
||||||
|
- All guidance is stage 1 only. Stage 2 always uses `SimpleDenoiser`.
|
||||||
|
|
||||||
|
## Sigma schedules and step counts
|
||||||
|
|
||||||
|
- **Scheduler-based** (full model): `self._scheduler = LTX2Scheduler()` with `execute(steps=N)` (HQ also passes `latent=` for token-count-dependent shift). Defaults: 30 steps (LTX-2.3), 40 (LTX-2), 15 (HQ).
|
||||||
|
- **Distilled**: Fixed 8-step `DISTILLED_SIGMA_VALUES` (9 values). Stage 2 uses 3-step `STAGE_2_DISTILLED_SIGMA_VALUES` (4 values). No `num_inference_steps` param.
|
||||||
|
- **Retake**: `num_inference_steps=40` default; ignored when `distilled=True` (fixed 8-step).
|
||||||
|
- **Overrides**: All pipelines accept optional sigma tensors in `__call__`: `sigmas` (one-stage), `stage_1_sigmas` + `stage_2_sigmas` (two-stage).
|
||||||
|
|
||||||
|
## LoRA conventions
|
||||||
|
|
||||||
|
- No default LoRAs. `loras` param defaults to empty list/tuple. `DEFAULT_LORA_STRENGTH = 1.0`.
|
||||||
|
- Two-stage non-distilled pipelines require `distilled_lora` (applied to stage 2 only in TI2Vid/A2Vid/Keyframe).
|
||||||
|
- HQ is unique: applies distilled LoRA to **both** stages with separate `distilled_lora_strength_stage_1` / `_stage_2` params.
|
||||||
|
|
||||||
|
## Shared building blocks (`utils/blocks.py`)
|
||||||
|
|
||||||
|
- `DiffusionStage` -- owns transformer lifecycle; builds model on call, frees on exit via `gpu_model()` context manager (moves params to meta device to release GPU/CPU memory). Accepts optional `stepper` and `loop` overrides.
|
||||||
|
- `PromptEncoder` -- Gemma text encoder + embeddings processor (video 4096-dim, audio 2048-dim).
|
||||||
|
- `ImageConditioner` / `AudioConditioner` -- temporary encoder scope; builds encoder, passes to callable, frees.
|
||||||
|
- `VideoUpsampler` -- 2x spatial upsampling via encoder + upsampler.
|
||||||
|
- `VideoDecoder` / `AudioDecoder` -- latent-to-pixel decoding (iterator for video, `Audio` for audio).
|
||||||
|
|
||||||
|
### Memory management
|
||||||
|
|
||||||
|
- **Model lifecycle**: All blocks build their model on call and free it on exit. `gpu_model()` moves params to `"meta"` device on exit, immediately releasing storage. No model persists between calls.
|
||||||
|
- **Layer streaming**: When `streaming_prefetch_count` is set, `DiffusionStage` wraps the transformer in `LayerStreamingWrapper`. Layers live on pinned CPU memory; only `1 + prefetch_count` layers are on GPU at a time, with async H2D prefetch on a separate CUDA stream.
|
||||||
|
- **Batch splitting**: `BatchSplitAdapter` wraps the transformer and splits inputs exceeding `max_batch_size` into sequential chunks. If guidance needs B=4 but `max_batch_size=1`, it runs 4 sequential B=1 passes. Higher `max_batch_size` reduces layer-streaming PCIe transfers at the cost of peak memory.
|
||||||
|
|
||||||
|
## Denoisers (`utils/denoisers.py`)
|
||||||
|
|
||||||
|
- `SimpleDenoiser` -- single forward pass (B=1), no guidance. Used by distilled pipelines and all stage 2.
|
||||||
|
- `GuidedDenoiser` -- CFG/STG with static `MultiModalGuider` instances (HQ, A2Vid, Retake non-distilled).
|
||||||
|
- `FactoryGuidedDenoiser` -- per-step guider creation via factory (OneStageTI2Vid, TwoStagesTI2Vid, Keyframe).
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
- **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).
|
||||||
|
- **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`).
|
||||||
|
|
||||||
|
## Image conditioning helpers (`utils/helpers.py`)
|
||||||
|
|
||||||
|
- `combined_image_conditionings()` -- images with `frame_idx==0` replace latent (`VideoConditionByLatentIndex`), others guide (`VideoConditionByKeyframeIndex`).
|
||||||
|
- `image_conditionings_by_adding_guiding_latent()` -- all images become keyframe guidance regardless of `frame_idx`.
|
||||||
@@ -1,6 +1,6 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "ltx-pipelines"
|
name = "ltx-pipelines"
|
||||||
version = "1.1.0"
|
version = "1.1.1"
|
||||||
description = "Pipelines implementation for Lightricks' LTX-2 model"
|
description = "Pipelines implementation for Lightricks' LTX-2 model"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
requires-python = ">=3.10"
|
requires-python = ">=3.10"
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ from ltx_pipelines.utils.blocks import (
|
|||||||
VideoUpsampler,
|
VideoUpsampler,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.constants import (
|
from ltx_pipelines.utils.constants import (
|
||||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
STAGE_2_DISTILLED_SIGMAS,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
|
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
|
||||||
from ltx_pipelines.utils.helpers import (
|
from ltx_pipelines.utils.helpers import (
|
||||||
@@ -56,6 +56,7 @@ class A2VidPipelineTwoStage:
|
|||||||
):
|
):
|
||||||
self.device = device or get_device()
|
self.device = device or get_device()
|
||||||
self.dtype = torch.bfloat16
|
self.dtype = torch.bfloat16
|
||||||
|
self._scheduler = LTX2Scheduler()
|
||||||
|
|
||||||
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
|
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
|
||||||
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
|
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
|
||||||
@@ -103,6 +104,8 @@ class A2VidPipelineTwoStage:
|
|||||||
enhance_prompt: bool = False,
|
enhance_prompt: bool = False,
|
||||||
streaming_prefetch_count: int | None = None,
|
streaming_prefetch_count: int | None = None,
|
||||||
max_batch_size: int = 1,
|
max_batch_size: int = 1,
|
||||||
|
stage_1_sigmas: torch.Tensor | None = None,
|
||||||
|
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
|
||||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||||
assert_resolution(height=height, width=width, is_two_stage=True)
|
assert_resolution(height=height, width=width, is_two_stage=True)
|
||||||
|
|
||||||
@@ -148,7 +151,9 @@ class A2VidPipelineTwoStage:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
|
sigmas = (
|
||||||
|
stage_1_sigmas if stage_1_sigmas is not None else self._scheduler.execute(steps=num_inference_steps)
|
||||||
|
).to(dtype=torch.float32, device=self.device)
|
||||||
|
|
||||||
video_state, _ = self.stage_1(
|
video_state, _ = self.stage_1(
|
||||||
denoiser=GuidedDenoiser(
|
denoiser=GuidedDenoiser(
|
||||||
@@ -185,7 +190,7 @@ class A2VidPipelineTwoStage:
|
|||||||
# Stage 2: Upsample and refine the video at higher resolution with distilled LoRA.
|
# Stage 2: Upsample and refine the video at higher resolution with distilled LoRA.
|
||||||
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
||||||
|
|
||||||
distilled_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
|
stage_2_sigmas = 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_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
|
||||||
stage_2_conditionings = self.image_conditioner(
|
stage_2_conditionings = self.image_conditioner(
|
||||||
lambda enc: combined_image_conditionings(
|
lambda enc: combined_image_conditionings(
|
||||||
@@ -200,7 +205,7 @@ class A2VidPipelineTwoStage:
|
|||||||
|
|
||||||
video_state, _ = self.stage_2(
|
video_state, _ = self.stage_2(
|
||||||
denoiser=SimpleDenoiser(v_context_p, a_context_p),
|
denoiser=SimpleDenoiser(v_context_p, a_context_p),
|
||||||
sigmas=distilled_sigmas,
|
sigmas=stage_2_sigmas,
|
||||||
noiser=noiser,
|
noiser=noiser,
|
||||||
width=width,
|
width=width,
|
||||||
height=height,
|
height=height,
|
||||||
@@ -209,7 +214,7 @@ class A2VidPipelineTwoStage:
|
|||||||
video=ModalitySpec(
|
video=ModalitySpec(
|
||||||
context=v_context_p,
|
context=v_context_p,
|
||||||
conditionings=stage_2_conditionings,
|
conditionings=stage_2_conditionings,
|
||||||
noise_scale=distilled_sigmas[0].item(),
|
noise_scale=stage_2_sigmas[0].item(),
|
||||||
initial_latent=upscaled_video_latent,
|
initial_latent=upscaled_video_latent,
|
||||||
),
|
),
|
||||||
audio=ModalitySpec(
|
audio=ModalitySpec(
|
||||||
|
|||||||
@@ -23,8 +23,8 @@ from ltx_pipelines.utils.blocks import (
|
|||||||
VideoUpsampler,
|
VideoUpsampler,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.constants import (
|
from ltx_pipelines.utils.constants import (
|
||||||
DISTILLED_SIGMA_VALUES,
|
DISTILLED_SIGMAS,
|
||||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
STAGE_2_DISTILLED_SIGMAS,
|
||||||
detect_params,
|
detect_params,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.denoisers import SimpleDenoiser
|
from ltx_pipelines.utils.denoisers import SimpleDenoiser
|
||||||
@@ -77,7 +77,7 @@ class DistilledPipeline:
|
|||||||
self.video_decoder = VideoDecoder(distilled_checkpoint_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.audio_decoder = AudioDecoder(distilled_checkpoint_path, self.dtype, self.device, registry=registry)
|
||||||
|
|
||||||
def __call__(
|
def __call__( # noqa: PLR0913
|
||||||
self,
|
self,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
seed: int,
|
seed: int,
|
||||||
@@ -89,6 +89,8 @@ class DistilledPipeline:
|
|||||||
tiling_config: TilingConfig | None = None,
|
tiling_config: TilingConfig | None = None,
|
||||||
enhance_prompt: bool = False,
|
enhance_prompt: bool = False,
|
||||||
streaming_prefetch_count: int | None = None,
|
streaming_prefetch_count: int | None = None,
|
||||||
|
stage_1_sigmas: torch.Tensor = DISTILLED_SIGMAS,
|
||||||
|
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
|
||||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||||
assert_resolution(height=height, width=width, is_two_stage=True)
|
assert_resolution(height=height, width=width, is_two_stage=True)
|
||||||
|
|
||||||
@@ -105,7 +107,7 @@ class DistilledPipeline:
|
|||||||
video_context, audio_context = ctx_p.video_encoding, ctx_p.audio_encoding
|
video_context, audio_context = ctx_p.video_encoding, ctx_p.audio_encoding
|
||||||
|
|
||||||
# Stage 1: Initial low resolution video generation.
|
# Stage 1: Initial low resolution video generation.
|
||||||
stage_1_sigmas = torch.Tensor(DISTILLED_SIGMA_VALUES).to(self.device)
|
stage_1_sigmas = stage_1_sigmas.to(dtype=torch.float32, device=self.device)
|
||||||
stage_1_w, stage_1_h = width // 2, height // 2
|
stage_1_w, stage_1_h = width // 2, height // 2
|
||||||
stage_1_conditionings = self.image_conditioner(
|
stage_1_conditionings = self.image_conditioner(
|
||||||
lambda enc: combined_image_conditionings(
|
lambda enc: combined_image_conditionings(
|
||||||
@@ -134,7 +136,7 @@ class DistilledPipeline:
|
|||||||
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
|
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
|
||||||
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
||||||
|
|
||||||
stage_2_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
|
stage_2_sigmas = stage_2_sigmas.to(dtype=torch.float32, device=self.device)
|
||||||
stage_2_conditionings = self.image_conditioner(
|
stage_2_conditionings = self.image_conditioner(
|
||||||
lambda enc: combined_image_conditionings(
|
lambda enc: combined_image_conditionings(
|
||||||
images=images,
|
images=images,
|
||||||
|
|||||||
@@ -32,8 +32,8 @@ from ltx_pipelines.utils.blocks import (
|
|||||||
VideoUpsampler,
|
VideoUpsampler,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.constants import (
|
from ltx_pipelines.utils.constants import (
|
||||||
DISTILLED_SIGMA_VALUES,
|
DISTILLED_SIGMAS,
|
||||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
STAGE_2_DISTILLED_SIGMAS,
|
||||||
detect_params,
|
detect_params,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.denoisers import SimpleDenoiser
|
from ltx_pipelines.utils.denoisers import SimpleDenoiser
|
||||||
@@ -126,6 +126,8 @@ class ICLoraPipeline:
|
|||||||
skip_stage_2: bool = False,
|
skip_stage_2: bool = False,
|
||||||
conditioning_attention_mask: torch.Tensor | None = None,
|
conditioning_attention_mask: torch.Tensor | None = None,
|
||||||
streaming_prefetch_count: int | None = None,
|
streaming_prefetch_count: int | None = None,
|
||||||
|
stage_1_sigmas: torch.Tensor = DISTILLED_SIGMAS,
|
||||||
|
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
|
||||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||||
"""
|
"""
|
||||||
Generate video with IC-LoRA conditioning.
|
Generate video with IC-LoRA conditioning.
|
||||||
@@ -200,7 +202,7 @@ class ICLoraPipeline:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
stage_1_sigmas = torch.Tensor(DISTILLED_SIGMA_VALUES).to(self.device)
|
stage_1_sigmas = stage_1_sigmas.to(dtype=torch.float32, device=self.device)
|
||||||
|
|
||||||
video_state, audio_state = self.stage_1(
|
video_state, audio_state = self.stage_1(
|
||||||
denoiser=SimpleDenoiser(video_context, audio_context),
|
denoiser=SimpleDenoiser(video_context, audio_context),
|
||||||
@@ -230,7 +232,7 @@ class ICLoraPipeline:
|
|||||||
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
|
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
|
||||||
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
||||||
|
|
||||||
distilled_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
|
stage_2_sigmas = 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_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
|
||||||
stage_2_conditionings = self.image_conditioner(
|
stage_2_conditionings = self.image_conditioner(
|
||||||
lambda enc: combined_image_conditionings(
|
lambda enc: combined_image_conditionings(
|
||||||
@@ -245,7 +247,7 @@ class ICLoraPipeline:
|
|||||||
|
|
||||||
video_state, audio_state = self.stage_2(
|
video_state, audio_state = self.stage_2(
|
||||||
denoiser=SimpleDenoiser(video_context, audio_context),
|
denoiser=SimpleDenoiser(video_context, audio_context),
|
||||||
sigmas=distilled_sigmas,
|
sigmas=stage_2_sigmas,
|
||||||
noiser=noiser,
|
noiser=noiser,
|
||||||
width=width,
|
width=width,
|
||||||
height=height,
|
height=height,
|
||||||
@@ -254,12 +256,12 @@ class ICLoraPipeline:
|
|||||||
video=ModalitySpec(
|
video=ModalitySpec(
|
||||||
context=video_context,
|
context=video_context,
|
||||||
conditionings=stage_2_conditionings,
|
conditionings=stage_2_conditionings,
|
||||||
noise_scale=distilled_sigmas[0].item(),
|
noise_scale=stage_2_sigmas[0].item(),
|
||||||
initial_latent=upscaled_video_latent,
|
initial_latent=upscaled_video_latent,
|
||||||
),
|
),
|
||||||
audio=ModalitySpec(
|
audio=ModalitySpec(
|
||||||
context=audio_context,
|
context=audio_context,
|
||||||
noise_scale=distilled_sigmas[0].item(),
|
noise_scale=stage_2_sigmas[0].item(),
|
||||||
initial_latent=audio_state.latent,
|
initial_latent=audio_state.latent,
|
||||||
),
|
),
|
||||||
streaming_prefetch_count=streaming_prefetch_count,
|
streaming_prefetch_count=streaming_prefetch_count,
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ from ltx_pipelines.utils.blocks import (
|
|||||||
VideoUpsampler,
|
VideoUpsampler,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.constants import (
|
from ltx_pipelines.utils.constants import (
|
||||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
STAGE_2_DISTILLED_SIGMAS,
|
||||||
detect_params,
|
detect_params,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.denoisers import FactoryGuidedDenoiser, SimpleDenoiser
|
from ltx_pipelines.utils.denoisers import FactoryGuidedDenoiser, SimpleDenoiser
|
||||||
@@ -62,6 +62,7 @@ class KeyframeInterpolationPipeline:
|
|||||||
):
|
):
|
||||||
self.device = device or get_device()
|
self.device = device or get_device()
|
||||||
self.dtype = torch.bfloat16
|
self.dtype = torch.bfloat16
|
||||||
|
self._scheduler = LTX2Scheduler()
|
||||||
|
|
||||||
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
|
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
|
||||||
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
|
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
|
||||||
@@ -107,6 +108,8 @@ class KeyframeInterpolationPipeline:
|
|||||||
enhance_prompt: bool = False,
|
enhance_prompt: bool = False,
|
||||||
streaming_prefetch_count: int | None = None,
|
streaming_prefetch_count: int | None = None,
|
||||||
max_batch_size: int = 1,
|
max_batch_size: int = 1,
|
||||||
|
stage_1_sigmas: torch.Tensor | None = None,
|
||||||
|
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
|
||||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||||
assert_resolution(height=height, width=width, is_two_stage=True)
|
assert_resolution(height=height, width=width, is_two_stage=True)
|
||||||
|
|
||||||
@@ -125,7 +128,9 @@ class KeyframeInterpolationPipeline:
|
|||||||
v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
|
v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
|
||||||
|
|
||||||
# Stage 1: Initial low resolution video generation.
|
# Stage 1: Initial low resolution video generation.
|
||||||
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
|
sigmas = (
|
||||||
|
stage_1_sigmas if stage_1_sigmas is not None else self._scheduler.execute(steps=num_inference_steps)
|
||||||
|
).to(dtype=torch.float32, device=self.device)
|
||||||
|
|
||||||
stage_1_output_shape = VideoPixelShape(
|
stage_1_output_shape = VideoPixelShape(
|
||||||
batch=1,
|
batch=1,
|
||||||
@@ -181,7 +186,7 @@ class KeyframeInterpolationPipeline:
|
|||||||
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
|
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
|
||||||
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
||||||
|
|
||||||
distilled_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
|
stage_2_sigmas = 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_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
|
||||||
stage_2_conditionings = self.image_conditioner(
|
stage_2_conditionings = self.image_conditioner(
|
||||||
lambda enc: image_conditionings_by_adding_guiding_latent(
|
lambda enc: image_conditionings_by_adding_guiding_latent(
|
||||||
@@ -196,7 +201,7 @@ class KeyframeInterpolationPipeline:
|
|||||||
|
|
||||||
video_state, audio_state = self.stage_2(
|
video_state, audio_state = self.stage_2(
|
||||||
denoiser=SimpleDenoiser(v_context_p, a_context_p),
|
denoiser=SimpleDenoiser(v_context_p, a_context_p),
|
||||||
sigmas=distilled_sigmas,
|
sigmas=stage_2_sigmas,
|
||||||
noiser=noiser,
|
noiser=noiser,
|
||||||
width=width,
|
width=width,
|
||||||
height=height,
|
height=height,
|
||||||
@@ -205,12 +210,12 @@ class KeyframeInterpolationPipeline:
|
|||||||
video=ModalitySpec(
|
video=ModalitySpec(
|
||||||
context=v_context_p,
|
context=v_context_p,
|
||||||
conditionings=stage_2_conditionings,
|
conditionings=stage_2_conditionings,
|
||||||
noise_scale=distilled_sigmas[0].item(),
|
noise_scale=stage_2_sigmas[0].item(),
|
||||||
initial_latent=upscaled_video_latent,
|
initial_latent=upscaled_video_latent,
|
||||||
),
|
),
|
||||||
audio=ModalitySpec(
|
audio=ModalitySpec(
|
||||||
context=a_context_p,
|
context=a_context_p,
|
||||||
noise_scale=distilled_sigmas[0].item(),
|
noise_scale=stage_2_sigmas[0].item(),
|
||||||
initial_latent=audio_state.latent,
|
initial_latent=audio_state.latent,
|
||||||
),
|
),
|
||||||
streaming_prefetch_count=streaming_prefetch_count,
|
streaming_prefetch_count=streaming_prefetch_count,
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ from ltx_pipelines.utils.blocks import (
|
|||||||
PromptEncoder,
|
PromptEncoder,
|
||||||
VideoDecoder,
|
VideoDecoder,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.constants import DISTILLED_SIGMA_VALUES, detect_params
|
from ltx_pipelines.utils.constants import DISTILLED_SIGMAS, detect_params
|
||||||
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
|
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
|
||||||
from ltx_pipelines.utils.helpers import (
|
from ltx_pipelines.utils.helpers import (
|
||||||
audio_latent_from_file,
|
audio_latent_from_file,
|
||||||
@@ -78,6 +78,8 @@ class RetakePipeline:
|
|||||||
self.device = device or get_device()
|
self.device = device or get_device()
|
||||||
self.dtype = torch.bfloat16
|
self.dtype = torch.bfloat16
|
||||||
self.distilled = distilled
|
self.distilled = distilled
|
||||||
|
if not distilled:
|
||||||
|
self._scheduler = LTX2Scheduler()
|
||||||
self.prompt_encoder = PromptEncoder(
|
self.prompt_encoder = PromptEncoder(
|
||||||
checkpoint_path=checkpoint_path,
|
checkpoint_path=checkpoint_path,
|
||||||
gemma_root=gemma_root,
|
gemma_root=gemma_root,
|
||||||
@@ -141,6 +143,7 @@ class RetakePipeline:
|
|||||||
tiling_config: TilingConfig | None = None,
|
tiling_config: TilingConfig | None = None,
|
||||||
streaming_prefetch_count: int | None = None,
|
streaming_prefetch_count: int | None = None,
|
||||||
max_batch_size: int = 1,
|
max_batch_size: int = 1,
|
||||||
|
sigmas: torch.Tensor | None = None,
|
||||||
) -> tuple[Iterator[torch.Tensor], torch.Tensor]:
|
) -> tuple[Iterator[torch.Tensor], torch.Tensor]:
|
||||||
"""Regenerate ``[start_time, end_time]`` of the source video (retake).
|
"""Regenerate ``[start_time, end_time]`` of the source video (retake).
|
||||||
Parameters
|
Parameters
|
||||||
@@ -227,15 +230,18 @@ class RetakePipeline:
|
|||||||
initial_latent=initial_audio_latent,
|
initial_latent=initial_audio_latent,
|
||||||
frozen=initial_audio_latent is not None and not regenerate_audio,
|
frozen=initial_audio_latent is not None and not regenerate_audio,
|
||||||
)
|
)
|
||||||
# Build denoiser
|
|
||||||
|
# Build denoiser and resolve sigma schedule.
|
||||||
|
if sigmas is None:
|
||||||
|
sigmas = DISTILLED_SIGMAS if self.distilled else self._scheduler.execute(steps=num_inference_steps)
|
||||||
|
sigmas = sigmas.to(dtype=torch.float32, device=self.device)
|
||||||
|
|
||||||
if self.distilled:
|
if self.distilled:
|
||||||
sigmas = torch.tensor(DISTILLED_SIGMA_VALUES).to(dtype=torch.float32, device=self.device)
|
|
||||||
denoiser = SimpleDenoiser(
|
denoiser = SimpleDenoiser(
|
||||||
v_context=v_context_p,
|
v_context=v_context_p,
|
||||||
a_context=a_context_p,
|
a_context=a_context_p,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
|
|
||||||
v_context_n, a_context_n = contexts[1].video_encoding, contexts[1].audio_encoding
|
v_context_n, a_context_n = contexts[1].video_encoding, contexts[1].audio_encoding
|
||||||
video_guider = MultiModalGuider(
|
video_guider = MultiModalGuider(
|
||||||
params=video_guider_params,
|
params=video_guider_params,
|
||||||
|
|||||||
@@ -55,6 +55,7 @@ class TI2VidOneStagePipeline:
|
|||||||
):
|
):
|
||||||
self.dtype = torch.bfloat16
|
self.dtype = torch.bfloat16
|
||||||
self.device = device or get_device()
|
self.device = device or get_device()
|
||||||
|
self._scheduler = LTX2Scheduler()
|
||||||
self.prompt_encoder = PromptEncoder(
|
self.prompt_encoder = PromptEncoder(
|
||||||
checkpoint_path=checkpoint_path,
|
checkpoint_path=checkpoint_path,
|
||||||
gemma_root=gemma_root,
|
gemma_root=gemma_root,
|
||||||
@@ -107,6 +108,7 @@ class TI2VidOneStagePipeline:
|
|||||||
streaming_prefetch_count: int | None = None,
|
streaming_prefetch_count: int | None = None,
|
||||||
tiling_config: TilingConfig | None = None,
|
tiling_config: TilingConfig | None = None,
|
||||||
max_batch_size: int = 1,
|
max_batch_size: int = 1,
|
||||||
|
sigmas: torch.Tensor | None = None,
|
||||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||||
assert_resolution(height=height, width=width, is_two_stage=False)
|
assert_resolution(height=height, width=width, is_two_stage=False)
|
||||||
|
|
||||||
@@ -135,7 +137,9 @@ class TI2VidOneStagePipeline:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
|
sigmas = (sigmas if sigmas is not None else self._scheduler.execute(steps=num_inference_steps)).to(
|
||||||
|
dtype=torch.float32, device=self.device
|
||||||
|
)
|
||||||
|
|
||||||
video_guider_factory = create_multimodal_guider_factory(
|
video_guider_factory = create_multimodal_guider_factory(
|
||||||
params=video_guider_params,
|
params=video_guider_params,
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ from ltx_pipelines.utils.blocks import (
|
|||||||
VideoUpsampler,
|
VideoUpsampler,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.constants import (
|
from ltx_pipelines.utils.constants import (
|
||||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
STAGE_2_DISTILLED_SIGMAS,
|
||||||
detect_params,
|
detect_params,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.denoisers import FactoryGuidedDenoiser, SimpleDenoiser
|
from ltx_pipelines.utils.denoisers import FactoryGuidedDenoiser, SimpleDenoiser
|
||||||
@@ -61,6 +61,7 @@ class TI2VidTwoStagesPipeline:
|
|||||||
):
|
):
|
||||||
self.device = device or get_device()
|
self.device = device or get_device()
|
||||||
self.dtype = torch.bfloat16
|
self.dtype = torch.bfloat16
|
||||||
|
self._scheduler = LTX2Scheduler()
|
||||||
|
|
||||||
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
|
self.prompt_encoder = PromptEncoder(checkpoint_path, gemma_root, self.dtype, self.device, registry=registry)
|
||||||
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
|
self.image_conditioner = ImageConditioner(checkpoint_path, self.dtype, self.device, registry=registry)
|
||||||
@@ -106,6 +107,8 @@ class TI2VidTwoStagesPipeline:
|
|||||||
enhance_prompt: bool = False,
|
enhance_prompt: bool = False,
|
||||||
streaming_prefetch_count: int | None = None,
|
streaming_prefetch_count: int | None = None,
|
||||||
max_batch_size: int = 1,
|
max_batch_size: int = 1,
|
||||||
|
stage_1_sigmas: torch.Tensor | None = None,
|
||||||
|
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
|
||||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||||
assert_resolution(height=height, width=width, is_two_stage=True)
|
assert_resolution(height=height, width=width, is_two_stage=True)
|
||||||
|
|
||||||
@@ -142,7 +145,9 @@ class TI2VidTwoStagesPipeline:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
|
sigmas = (
|
||||||
|
stage_1_sigmas if stage_1_sigmas is not None else self._scheduler.execute(steps=num_inference_steps)
|
||||||
|
).to(dtype=torch.float32, device=self.device)
|
||||||
|
|
||||||
video_state, audio_state = self.stage_1(
|
video_state, audio_state = self.stage_1(
|
||||||
denoiser=FactoryGuidedDenoiser(
|
denoiser=FactoryGuidedDenoiser(
|
||||||
@@ -172,7 +177,7 @@ class TI2VidTwoStagesPipeline:
|
|||||||
# Stage 2: Upsample and refine the video at higher resolution with distilled LoRA.
|
# Stage 2: Upsample and refine the video at higher resolution with distilled LoRA.
|
||||||
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
||||||
|
|
||||||
distilled_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
|
stage_2_sigmas = stage_2_sigmas.to(dtype=torch.float32, device=self.device)
|
||||||
stage_2_conditionings = self.image_conditioner(
|
stage_2_conditionings = self.image_conditioner(
|
||||||
lambda enc: combined_image_conditionings(
|
lambda enc: combined_image_conditionings(
|
||||||
images=images,
|
images=images,
|
||||||
@@ -186,7 +191,7 @@ class TI2VidTwoStagesPipeline:
|
|||||||
|
|
||||||
video_state, audio_state = self.stage_2(
|
video_state, audio_state = self.stage_2(
|
||||||
denoiser=SimpleDenoiser(v_context=v_context_p, a_context=a_context_p),
|
denoiser=SimpleDenoiser(v_context=v_context_p, a_context=a_context_p),
|
||||||
sigmas=distilled_sigmas,
|
sigmas=stage_2_sigmas,
|
||||||
noiser=noiser,
|
noiser=noiser,
|
||||||
width=width,
|
width=width,
|
||||||
height=height,
|
height=height,
|
||||||
@@ -195,12 +200,12 @@ class TI2VidTwoStagesPipeline:
|
|||||||
video=ModalitySpec(
|
video=ModalitySpec(
|
||||||
context=v_context_p,
|
context=v_context_p,
|
||||||
conditionings=stage_2_conditionings,
|
conditionings=stage_2_conditionings,
|
||||||
noise_scale=distilled_sigmas[0].item(),
|
noise_scale=stage_2_sigmas[0].item(),
|
||||||
initial_latent=upscaled_video_latent,
|
initial_latent=upscaled_video_latent,
|
||||||
),
|
),
|
||||||
audio=ModalitySpec(
|
audio=ModalitySpec(
|
||||||
context=a_context_p,
|
context=a_context_p,
|
||||||
noise_scale=distilled_sigmas[0].item(),
|
noise_scale=stage_2_sigmas[0].item(),
|
||||||
initial_latent=audio_state.latent,
|
initial_latent=audio_state.latent,
|
||||||
),
|
),
|
||||||
streaming_prefetch_count=streaming_prefetch_count,
|
streaming_prefetch_count=streaming_prefetch_count,
|
||||||
|
|||||||
@@ -23,7 +23,7 @@ from ltx_pipelines.utils.blocks import (
|
|||||||
)
|
)
|
||||||
from ltx_pipelines.utils.constants import (
|
from ltx_pipelines.utils.constants import (
|
||||||
LTX_2_3_HQ_PARAMS,
|
LTX_2_3_HQ_PARAMS,
|
||||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
STAGE_2_DISTILLED_SIGMAS,
|
||||||
)
|
)
|
||||||
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
|
from ltx_pipelines.utils.denoisers import GuidedDenoiser, SimpleDenoiser
|
||||||
from ltx_pipelines.utils.helpers import (
|
from ltx_pipelines.utils.helpers import (
|
||||||
@@ -64,6 +64,7 @@ class TI2VidTwoStagesHQPipeline:
|
|||||||
):
|
):
|
||||||
self.device = device or get_device()
|
self.device = device or get_device()
|
||||||
self.dtype = torch.bfloat16
|
self.dtype = torch.bfloat16
|
||||||
|
self._scheduler = LTX2Scheduler()
|
||||||
|
|
||||||
distilled_lora_stage_1 = LoraPathStrengthAndSDOps(
|
distilled_lora_stage_1 = LoraPathStrengthAndSDOps(
|
||||||
path=distilled_lora[0].path,
|
path=distilled_lora[0].path,
|
||||||
@@ -121,6 +122,8 @@ class TI2VidTwoStagesHQPipeline:
|
|||||||
enhance_prompt: bool = False,
|
enhance_prompt: bool = False,
|
||||||
streaming_prefetch_count: int | None = None,
|
streaming_prefetch_count: int | None = None,
|
||||||
max_batch_size: int = 1,
|
max_batch_size: int = 1,
|
||||||
|
stage_1_sigmas: torch.Tensor | None = None,
|
||||||
|
stage_2_sigmas: torch.Tensor = STAGE_2_DISTILLED_SIGMAS,
|
||||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||||
assert_resolution(height=height, width=width, is_two_stage=True)
|
assert_resolution(height=height, width=width, is_two_stage=True)
|
||||||
|
|
||||||
@@ -157,13 +160,12 @@ class TI2VidTwoStagesHQPipeline:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
empty_latent = torch.empty(VideoLatentShape.from_pixel_shape(stage_1_output_shape).to_torch_shape())
|
|
||||||
stepper = Res2sDiffusionStep()
|
stepper = Res2sDiffusionStep()
|
||||||
sigmas = (
|
|
||||||
LTX2Scheduler()
|
if stage_1_sigmas is None:
|
||||||
.execute(latent=empty_latent, steps=num_inference_steps)
|
empty_latent = torch.empty(VideoLatentShape.from_pixel_shape(stage_1_output_shape).to_torch_shape())
|
||||||
.to(dtype=torch.float32, device=self.device)
|
stage_1_sigmas = self._scheduler.execute(latent=empty_latent, steps=num_inference_steps)
|
||||||
)
|
sigmas = stage_1_sigmas.to(dtype=torch.float32, device=self.device)
|
||||||
|
|
||||||
video_state, audio_state = self.stage_1(
|
video_state, audio_state = self.stage_1(
|
||||||
denoiser=GuidedDenoiser(
|
denoiser=GuidedDenoiser(
|
||||||
@@ -195,7 +197,7 @@ class TI2VidTwoStagesHQPipeline:
|
|||||||
# Stage 2: Upsample and refine the video at higher resolution with distilled LoRA.
|
# Stage 2: Upsample and refine the video at higher resolution with distilled LoRA.
|
||||||
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
upscaled_video_latent = self.upsampler(video_state.latent[:1])
|
||||||
|
|
||||||
distilled_sigmas = torch.tensor(STAGE_2_DISTILLED_SIGMA_VALUES, device=self.device)
|
stage_2_sigmas = 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_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
|
||||||
stage_2_conditionings = self.image_conditioner(
|
stage_2_conditionings = self.image_conditioner(
|
||||||
lambda enc: combined_image_conditionings(
|
lambda enc: combined_image_conditionings(
|
||||||
@@ -210,7 +212,7 @@ class TI2VidTwoStagesHQPipeline:
|
|||||||
|
|
||||||
video_state, audio_state = self.stage_2(
|
video_state, audio_state = self.stage_2(
|
||||||
denoiser=SimpleDenoiser(v_context=v_context_p, a_context=a_context_p),
|
denoiser=SimpleDenoiser(v_context=v_context_p, a_context=a_context_p),
|
||||||
sigmas=distilled_sigmas,
|
sigmas=stage_2_sigmas,
|
||||||
noiser=noiser,
|
noiser=noiser,
|
||||||
stepper=stepper,
|
stepper=stepper,
|
||||||
width=width,
|
width=width,
|
||||||
@@ -220,12 +222,12 @@ class TI2VidTwoStagesHQPipeline:
|
|||||||
video=ModalitySpec(
|
video=ModalitySpec(
|
||||||
context=v_context_p,
|
context=v_context_p,
|
||||||
conditionings=stage_2_conditionings,
|
conditionings=stage_2_conditionings,
|
||||||
noise_scale=distilled_sigmas[0].item(),
|
noise_scale=stage_2_sigmas[0].item(),
|
||||||
initial_latent=upscaled_video_latent,
|
initial_latent=upscaled_video_latent,
|
||||||
),
|
),
|
||||||
audio=ModalitySpec(
|
audio=ModalitySpec(
|
||||||
context=a_context_p,
|
context=a_context_p,
|
||||||
noise_scale=distilled_sigmas[0].item(),
|
noise_scale=stage_2_sigmas[0].item(),
|
||||||
initial_latent=audio_state.latent,
|
initial_latent=audio_state.latent,
|
||||||
),
|
),
|
||||||
loop=res2s_audio_video_denoising_loop,
|
loop=res2s_audio_video_denoising_loop,
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
import logging
|
import logging
|
||||||
from dataclasses import dataclass, field, replace
|
from dataclasses import dataclass, field, replace
|
||||||
|
|
||||||
|
import torch
|
||||||
from safetensors import safe_open
|
from safetensors import safe_open
|
||||||
|
|
||||||
from ltx_core.components.guiders import MultiModalGuiderParams
|
from ltx_core.components.guiders import MultiModalGuiderParams
|
||||||
@@ -17,6 +18,9 @@ DISTILLED_SIGMA_VALUES = [1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725,
|
|||||||
# Reduced schedule for super-resolution stage 2 (subset of distilled values)
|
# Reduced schedule for super-resolution stage 2 (subset of distilled values)
|
||||||
STAGE_2_DISTILLED_SIGMA_VALUES = [0.909375, 0.725, 0.421875, 0.0]
|
STAGE_2_DISTILLED_SIGMA_VALUES = [0.909375, 0.725, 0.421875, 0.0]
|
||||||
|
|
||||||
|
DISTILLED_SIGMAS = torch.tensor(DISTILLED_SIGMA_VALUES)
|
||||||
|
STAGE_2_DISTILLED_SIGMAS = torch.tensor(STAGE_2_DISTILLED_SIGMA_VALUES)
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
# Pipeline Parameters
|
# Pipeline Parameters
|
||||||
|
|||||||
@@ -307,15 +307,14 @@ class InpaintingStrategy(TrainingStrategy):
|
|||||||
audio_pred: Tensor | None,
|
audio_pred: Tensor | None,
|
||||||
inputs: ModelInputs,
|
inputs: ModelInputs,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
"""Compute training loss on inpaint regions only."""
|
"""Compute training loss on inpaint regions only. Returns [B,]."""
|
||||||
# MSE loss
|
# MSE loss
|
||||||
loss = (video_pred - inputs.video_targets).pow(2)
|
loss = (video_pred - inputs.video_targets).pow(2)
|
||||||
|
|
||||||
# Apply loss mask
|
# Apply loss mask and reduce to per-element [B,]
|
||||||
loss_mask = inputs.video_loss_mask.unsqueeze(-1).float()
|
loss_mask = inputs.video_loss_mask.unsqueeze(-1).float()
|
||||||
loss = loss.mul(loss_mask).div(loss_mask.mean() + 1e-8)
|
masked = loss.mul(loss_mask)
|
||||||
|
return masked.mean(dim=[-2, -1]) / loss_mask.mean(dim=[-2, -1]).clamp(min=1e-8)
|
||||||
return loss.mean()
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### Step 5: Register the Strategy
|
### Step 5: Register the Strategy
|
||||||
|
|||||||
@@ -48,6 +48,10 @@ uv run python scripts/caption_videos.py videos_dir/ --output dataset.json --use-
|
|||||||
uv run python scripts/caption_videos.py videos_dir/ --output dataset.json \
|
uv run python scripts/caption_videos.py videos_dir/ --output dataset.json \
|
||||||
--captioner-type gemini_flash --api-key YOUR_API_KEY
|
--captioner-type gemini_flash --api-key YOUR_API_KEY
|
||||||
|
|
||||||
|
# Use Gemini Flash with parallel workers for faster throughput
|
||||||
|
uv run python scripts/caption_videos.py videos_dir/ --output dataset.json \
|
||||||
|
--captioner-type gemini_flash --num-workers 5
|
||||||
|
|
||||||
# Caption without audio processing (video-only)
|
# Caption without audio processing (video-only)
|
||||||
uv run python scripts/caption_videos.py videos_dir/ --output dataset.json --no-audio
|
uv run python scripts/caption_videos.py videos_dir/ --output dataset.json --no-audio
|
||||||
|
|
||||||
@@ -61,9 +65,10 @@ uv run python scripts/caption_videos.py videos_dir/ --output dataset.json --over
|
|||||||
- **Multiple backends**:
|
- **Multiple backends**:
|
||||||
- `qwen_omni` (default): Local Qwen2.5-Omni model - processes video + audio locally
|
- `qwen_omni` (default): Local Qwen2.5-Omni model - processes video + audio locally
|
||||||
- `gemini_flash`: Google Gemini Flash API - cloud-based, requires API key
|
- `gemini_flash`: Google Gemini Flash API - cloud-based, requires API key
|
||||||
|
- **Parallel captioning** (Gemini Flash only): Use `--num-workers` to run multiple API calls concurrently for faster throughput on large datasets
|
||||||
- **Structured output**: Captions include visual description, speech transcription, sounds, and on-screen text
|
- **Structured output**: Captions include visual description, speech transcription, sounds, and on-screen text
|
||||||
- **Memory optimization**: 8-bit quantization option for limited VRAM
|
- **Memory optimization**: 8-bit quantization option for limited VRAM
|
||||||
- **Incremental processing**: Skips already-captioned files by default
|
- **Incremental processing**: Skips already-captioned files by default; progress is saved every 5 videos
|
||||||
- **Multiple output formats**: JSON, JSONL, CSV, or TXT
|
- **Multiple output formats**: JSON, JSONL, CSV, or TXT
|
||||||
|
|
||||||
**Caption format:**
|
**Caption format:**
|
||||||
@@ -74,6 +79,26 @@ The captioner produces structured captions with four sections:
|
|||||||
- `[SOUNDS]`: Description of music, ambient sounds, sound effects
|
- `[SOUNDS]`: Description of music, ambient sounds, sound effects
|
||||||
- `[TEXT]`: Any on-screen text visible in the video
|
- `[TEXT]`: Any on-screen text visible in the video
|
||||||
|
|
||||||
|
**Parallel captioning with Gemini Flash:**
|
||||||
|
|
||||||
|
When using `--captioner-type gemini_flash`, you can speed up large dataset captioning by running multiple API calls at the same time using `--num-workers` (accepts 1–10, default is 1):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export GEMINI_API_KEY="your-key-here"
|
||||||
|
|
||||||
|
# Caption a large dataset with 5 workers running concurrently
|
||||||
|
uv run python scripts/caption_videos.py videos_dir/ \
|
||||||
|
--output dataset.json \
|
||||||
|
--captioner-type gemini_flash \
|
||||||
|
--num-workers 5
|
||||||
|
```
|
||||||
|
|
||||||
|
> [!NOTE]
|
||||||
|
> `--num-workers` is only supported with `gemini_flash`. Using it with `qwen_omni` or any other local model will raise an error, because local GPU models are not thread-safe.
|
||||||
|
|
||||||
|
> [!TIP]
|
||||||
|
> Keep `--num-workers` between 3–5 for most use cases. Very high values (8–10) may hit Gemini API rate limits depending on your quota tier.
|
||||||
|
|
||||||
**Environment variables (for Gemini Flash):**
|
**Environment variables (for Gemini Flash):**
|
||||||
|
|
||||||
Set one of these to use Gemini Flash without passing `--api-key`:
|
Set one of these to use Gemini Flash without passing `--api-key`:
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "ltx-trainer"
|
name = "ltx-trainer"
|
||||||
version = "1.1.0"
|
version = "1.1.1"
|
||||||
description = "LTX-2 training, democratized."
|
description = "LTX-2 training, democratized."
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
authors = [
|
authors = [
|
||||||
@@ -48,7 +48,7 @@ build-backend = "hatchling.build"
|
|||||||
|
|
||||||
|
|
||||||
[tool.ruff]
|
[tool.ruff]
|
||||||
target-version = "1.1.0"
|
target-version = "1.1.1"
|
||||||
line-length = 120
|
line-length = 120
|
||||||
|
|
||||||
[tool.ruff.lint]
|
[tool.ruff.lint]
|
||||||
|
|||||||
@@ -19,6 +19,8 @@ Basic usage:
|
|||||||
Advanced usage:
|
Advanced usage:
|
||||||
# Use Gemini Flash API (requires GEMINI_API_KEY or GOOGLE_API_KEY env var)
|
# Use Gemini Flash API (requires GEMINI_API_KEY or GOOGLE_API_KEY env var)
|
||||||
caption_videos.py videos_dir/ --captioner-type gemini_flash
|
caption_videos.py videos_dir/ --captioner-type gemini_flash
|
||||||
|
# Use Gemini Flash with parallel workers (2-10 workers, cloud API only)
|
||||||
|
caption_videos.py videos_dir/ --captioner-type gemini_flash --num-workers 5
|
||||||
# Disable audio processing (video-only captions)
|
# Disable audio processing (video-only captions)
|
||||||
caption_videos.py videos_dir/ --no-audio
|
caption_videos.py videos_dir/ --no-audio
|
||||||
# Process videos with specific extensions and save as JSON
|
# Process videos with specific extensions and save as JSON
|
||||||
@@ -27,6 +29,7 @@ Advanced usage:
|
|||||||
|
|
||||||
import csv
|
import csv
|
||||||
import json
|
import json
|
||||||
|
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
@@ -70,7 +73,7 @@ class OutputFormat(str, Enum):
|
|||||||
JSONL = "jsonl" # JSON Lines file with one JSON object per line
|
JSONL = "jsonl" # JSON Lines file with one JSON object per line
|
||||||
|
|
||||||
|
|
||||||
def caption_media(
|
def caption_media( # noqa: PLR0913
|
||||||
input_path: Path,
|
input_path: Path,
|
||||||
output_path: Path,
|
output_path: Path,
|
||||||
captioner: MediaCaptioningModel,
|
captioner: MediaCaptioningModel,
|
||||||
@@ -81,6 +84,7 @@ def caption_media(
|
|||||||
clean_caption: bool,
|
clean_caption: bool,
|
||||||
output_format: OutputFormat,
|
output_format: OutputFormat,
|
||||||
override: bool,
|
override: bool,
|
||||||
|
num_workers: int = 1,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Caption videos and images using the provided captioning model.
|
"""Caption videos and images using the provided captioning model.
|
||||||
Args:
|
Args:
|
||||||
@@ -94,6 +98,7 @@ def caption_media(
|
|||||||
clean_caption: Whether to clean up captions
|
clean_caption: Whether to clean up captions
|
||||||
output_format: Format to save the captions in
|
output_format: Format to save the captions in
|
||||||
override: Whether to override existing captions
|
override: Whether to override existing captions
|
||||||
|
num_workers: Number of parallel workers (only for cloud-based captioners like Gemini)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Get list of media files to process
|
# Get list of media files to process
|
||||||
@@ -121,9 +126,13 @@ def caption_media(
|
|||||||
console.print("[bold yellow]All media already have captions. Use --override to recaption.[/]")
|
console.print("[bold yellow]All media already have captions. Use --override to recaption.[/]")
|
||||||
return
|
return
|
||||||
|
|
||||||
# Process media files
|
if num_workers > 1:
|
||||||
|
console.print(f"Running with [bold cyan]{num_workers}[/] parallel workers.")
|
||||||
|
|
||||||
captions = existing_captions.copy()
|
captions = existing_captions.copy()
|
||||||
successfully_captioned = 0
|
successfully_captioned = 0
|
||||||
|
completed_since_save = 0
|
||||||
|
|
||||||
progress = Progress(
|
progress = Progress(
|
||||||
SpinnerColumn(),
|
SpinnerColumn(),
|
||||||
TextColumn("{task.description}"),
|
TextColumn("{task.description}"),
|
||||||
@@ -135,36 +144,47 @@ def caption_media(
|
|||||||
console=console,
|
console=console,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def process_one(media_file: Path) -> tuple[str, str]:
|
||||||
|
"""Caption a single media file and return (relative_path, caption)."""
|
||||||
|
caption = captioner.caption(
|
||||||
|
path=media_file,
|
||||||
|
fps=fps,
|
||||||
|
include_audio=include_audio,
|
||||||
|
clean_caption=clean_caption,
|
||||||
|
)
|
||||||
|
rel_path = str(media_file.resolve().relative_to(base_dir))
|
||||||
|
return rel_path, caption
|
||||||
|
|
||||||
with progress:
|
with progress:
|
||||||
task = progress.add_task("Captioning", total=len(media_to_process))
|
task = progress.add_task(
|
||||||
|
f"Captioning (workers: {num_workers})" if num_workers > 1 else "Captioning",
|
||||||
|
total=len(media_to_process),
|
||||||
|
)
|
||||||
|
|
||||||
for i, media_file in enumerate(media_to_process):
|
with ThreadPoolExecutor(max_workers=num_workers) as executor:
|
||||||
progress.update(task, description=f"Captioning [bold blue]{media_file.name}[/]")
|
futures = {executor.submit(process_one, f): f for f in media_to_process}
|
||||||
|
|
||||||
try:
|
for future in as_completed(futures):
|
||||||
# Generate caption for the media
|
media_file = futures[future]
|
||||||
caption = captioner.caption(
|
progress.update(task, description=f"Captioning [bold blue]{media_file.name}[/]")
|
||||||
path=media_file,
|
|
||||||
fps=fps,
|
|
||||||
include_audio=include_audio,
|
|
||||||
clean_caption=clean_caption,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Convert absolute path to relative path (relative to the output file's directory)
|
try:
|
||||||
rel_path = str(media_file.resolve().relative_to(base_dir))
|
rel_path, caption = future.result()
|
||||||
# Store the caption with the relative path as key
|
|
||||||
captions[rel_path] = caption
|
|
||||||
successfully_captioned += 1
|
|
||||||
except Exception as e:
|
|
||||||
console.print(f"[bold red]Error captioning {media_file}: {e}[/]")
|
|
||||||
|
|
||||||
if i % SAVE_INTERVAL == 0:
|
captions[rel_path] = caption
|
||||||
_save_captions(captions, output_path, output_format)
|
successfully_captioned += 1
|
||||||
|
completed_since_save += 1
|
||||||
|
|
||||||
# Advance progress bar
|
if completed_since_save >= SAVE_INTERVAL:
|
||||||
progress.advance(task)
|
_save_captions(captions, output_path, output_format)
|
||||||
|
completed_since_save = 0
|
||||||
|
|
||||||
# Save captions to file
|
except Exception as e:
|
||||||
|
console.print(f"[bold red]Error captioning {media_file.name}: {e}[/]")
|
||||||
|
|
||||||
|
progress.advance(task)
|
||||||
|
|
||||||
|
# Final save with everything accumulated
|
||||||
_save_captions(captions, output_path, output_format)
|
_save_captions(captions, output_path, output_format)
|
||||||
|
|
||||||
# Print summary
|
# Print summary
|
||||||
@@ -407,6 +427,18 @@ def main( # noqa: PLR0913
|
|||||||
envvar=["GOOGLE_API_KEY", "GEMINI_API_KEY"],
|
envvar=["GOOGLE_API_KEY", "GEMINI_API_KEY"],
|
||||||
help="API key for Gemini Flash (can also use GOOGLE_API_KEY or GEMINI_API_KEY env var)",
|
help="API key for Gemini Flash (can also use GOOGLE_API_KEY or GEMINI_API_KEY env var)",
|
||||||
),
|
),
|
||||||
|
num_workers: int = typer.Option(
|
||||||
|
1,
|
||||||
|
"--num-workers",
|
||||||
|
"-w",
|
||||||
|
min=1,
|
||||||
|
max=10,
|
||||||
|
help=(
|
||||||
|
"Number of parallel workers for captioning (1-10). "
|
||||||
|
"Values above 1 are only supported for cloud-based captioners (gemini_flash). "
|
||||||
|
"Using multiple workers with a local model will raise an error."
|
||||||
|
),
|
||||||
|
),
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Auto-caption videos with audio using multimodal models.
|
"""Auto-caption videos with audio using multimodal models.
|
||||||
This script supports audio-visual captioning using:
|
This script supports audio-visual captioning using:
|
||||||
@@ -424,6 +456,17 @@ def main( # noqa: PLR0913
|
|||||||
caption_videos.py video.mp4 -o captions.json -i "Describe this video in detail"
|
caption_videos.py video.mp4 -o captions.json -i "Describe this video in detail"
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
# Parallel workers are only safe for cloud-based (stateless) captioners.
|
||||||
|
# Local models like Qwen-Omni hold GPU state and are not thread-safe.
|
||||||
|
if num_workers > 1 and captioner_type != CaptionerType.GEMINI_FLASH:
|
||||||
|
console.print(
|
||||||
|
"[bold red]Error:[/] --num-workers > 1 is only supported with [bold]--captioner-type gemini_flash[/].\n"
|
||||||
|
"Local models (e.g. qwen_omni) run on GPU and are not thread-safe — "
|
||||||
|
"parallel calls would cause memory corruption or incorrect results.\n"
|
||||||
|
"Either set [bold]--num-workers 1[/] (default) or switch to [bold]--captioner-type gemini_flash[/]."
|
||||||
|
)
|
||||||
|
raise typer.Exit(code=1)
|
||||||
|
|
||||||
# Determine device for local models
|
# Determine device for local models
|
||||||
device_str = device or ("cuda" if torch.cuda.is_available() else "cpu")
|
device_str = device or ("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
|
||||||
@@ -479,6 +522,7 @@ def main( # noqa: PLR0913
|
|||||||
clean_caption=clean_caption,
|
clean_caption=clean_caption,
|
||||||
output_format=output_format,
|
output_format=output_format,
|
||||||
override=override,
|
override=override,
|
||||||
|
num_workers=num_workers,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,59 @@
|
|||||||
|
"""Sigma-bucketed loss tracking.
|
||||||
|
Maps each training step's per-element sigmas and losses to buckets.
|
||||||
|
Smoothing is left to wandb's UI.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import bisect
|
||||||
|
from collections import defaultdict
|
||||||
|
|
||||||
|
|
||||||
|
class SigmaBucketTracker:
|
||||||
|
"""Map per-element sigma values to named buckets for per-bucket loss logging.
|
||||||
|
By default, partitions [0, 1] into four equal-width buckets.
|
||||||
|
Custom boundaries can be provided for non-uniform bucketing.
|
||||||
|
Each call to update() receives per-element sigmas and losses (both [B,]),
|
||||||
|
buckets each element, and computes the mean loss per bucket. This gives
|
||||||
|
accurate per-sigma loss tracking even for batch_size > 1.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
bucket_boundaries: list[float] | None = None,
|
||||||
|
) -> None:
|
||||||
|
if bucket_boundaries is None:
|
||||||
|
bucket_boundaries = [0.0, 0.25, 0.5, 0.75, 1.0]
|
||||||
|
if len(bucket_boundaries) < 2:
|
||||||
|
raise ValueError("bucket_boundaries must have at least 2 elements")
|
||||||
|
if any(bucket_boundaries[i] >= bucket_boundaries[i + 1] for i in range(len(bucket_boundaries) - 1)):
|
||||||
|
raise ValueError("bucket_boundaries must be strictly increasing")
|
||||||
|
self._boundaries = list(bucket_boundaries)
|
||||||
|
self._num_buckets = len(bucket_boundaries) - 1
|
||||||
|
self._bucket_labels = [
|
||||||
|
f"{bucket_boundaries[i]:.2f}-{bucket_boundaries[i + 1]:.2f}" for i in range(self._num_buckets)
|
||||||
|
]
|
||||||
|
self._last_metrics: dict[str, float] = {}
|
||||||
|
|
||||||
|
def _get_bucket_index(self, sigma: float) -> int:
|
||||||
|
"""Map sigma value to bucket index."""
|
||||||
|
idx = bisect.bisect_right(self._boundaries, sigma) - 1
|
||||||
|
return max(0, min(idx, self._num_buckets - 1))
|
||||||
|
|
||||||
|
def update(self, sigmas: list[float], losses: list[float]) -> None:
|
||||||
|
"""Record per-element losses into their sigma buckets.
|
||||||
|
Args:
|
||||||
|
sigmas: Per-element sigma values, one per batch element.
|
||||||
|
losses: Per-element losses, one per batch element.
|
||||||
|
"""
|
||||||
|
if not sigmas:
|
||||||
|
self._last_metrics = {}
|
||||||
|
return
|
||||||
|
bucket_losses: dict[int, list[float]] = defaultdict(list)
|
||||||
|
for sigma, loss in zip(sigmas, losses, strict=True):
|
||||||
|
bucket_losses[self._get_bucket_index(sigma)].append(loss)
|
||||||
|
self._last_metrics = {self._bucket_labels[b]: sum(vals) / len(vals) for b, vals in bucket_losses.items()}
|
||||||
|
|
||||||
|
def get_metrics(self, prefix: str = "train") -> dict[str, float]:
|
||||||
|
"""Return the mean loss for each bucket hit on the last update.
|
||||||
|
Wandb handles smoothing in the UI.
|
||||||
|
"""
|
||||||
|
return {f"{prefix}/loss_sigma_{label}": loss for label, loss in self._last_metrics.items()}
|
||||||
@@ -2,8 +2,9 @@ import os
|
|||||||
import re
|
import re
|
||||||
import time
|
import time
|
||||||
import warnings
|
import warnings
|
||||||
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Callable
|
from typing import Any, Callable
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import wandb
|
import wandb
|
||||||
@@ -39,6 +40,7 @@ from ltx_trainer.model_loader import load_embeddings_processor, load_text_encode
|
|||||||
from ltx_trainer.model_loader import load_model as load_ltx_model
|
from ltx_trainer.model_loader import load_model as load_ltx_model
|
||||||
from ltx_trainer.progress import TrainingProgress
|
from ltx_trainer.progress import TrainingProgress
|
||||||
from ltx_trainer.quantization import quantize_model
|
from ltx_trainer.quantization import quantize_model
|
||||||
|
from ltx_trainer.sigma_tracker import SigmaBucketTracker
|
||||||
from ltx_trainer.timestep_samplers import SAMPLERS
|
from ltx_trainer.timestep_samplers import SAMPLERS
|
||||||
from ltx_trainer.training_state import ConfigFingerprint, RngStates, TrainingState
|
from ltx_trainer.training_state import ConfigFingerprint, RngStates, TrainingState
|
||||||
from ltx_trainer.training_strategies import get_training_strategy
|
from ltx_trainer.training_strategies import get_training_strategy
|
||||||
@@ -77,6 +79,14 @@ class TrainingStats(BaseModel):
|
|||||||
num_processes: int
|
num_processes: int
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class TrainingStepOutput:
|
||||||
|
"""Output from a single training step."""
|
||||||
|
|
||||||
|
loss: Tensor # [B,] per-element loss (unreduced)
|
||||||
|
sigma: Tensor # [B,] sampled sigma, detached from computational graph
|
||||||
|
|
||||||
|
|
||||||
class LtxvTrainer:
|
class LtxvTrainer:
|
||||||
def __init__(self, trainer_config: LtxTrainerConfig) -> None:
|
def __init__(self, trainer_config: LtxTrainerConfig) -> None:
|
||||||
self._config = trainer_config
|
self._config = trainer_config
|
||||||
@@ -95,7 +105,8 @@ class LtxvTrainer:
|
|||||||
self._checkpoint_paths: list[Path] = []
|
self._checkpoint_paths: list[Path] = []
|
||||||
self._training_state_paths: list[Path] = []
|
self._training_state_paths: list[Path] = []
|
||||||
self._training_state_size_warned = False
|
self._training_state_size_warned = False
|
||||||
self._init_wandb()
|
self._wandb_run = None
|
||||||
|
self._sigma_tracker = SigmaBucketTracker()
|
||||||
|
|
||||||
def train( # noqa: PLR0912, PLR0915
|
def train( # noqa: PLR0912, PLR0915
|
||||||
self,
|
self,
|
||||||
@@ -128,6 +139,10 @@ class LtxvTrainer:
|
|||||||
initial_step = 0
|
initial_step = 0
|
||||||
resuming = False
|
resuming = False
|
||||||
|
|
||||||
|
# Initialize W&B after restore so we only resume the run when state restore succeeds.
|
||||||
|
resume_run_id = training_state.wandb_run_id if resuming and training_state is not None else None
|
||||||
|
self._init_wandb(resume_run_id=resume_run_id)
|
||||||
|
|
||||||
self._init_dataloader()
|
self._init_dataloader()
|
||||||
data_iter = iter(self._dataloader)
|
data_iter = iter(self._dataloader)
|
||||||
self._init_timestep_sampler()
|
self._init_timestep_sampler()
|
||||||
@@ -191,8 +206,8 @@ class LtxvTrainer:
|
|||||||
if is_optimization_step:
|
if is_optimization_step:
|
||||||
self._global_step += 1
|
self._global_step += 1
|
||||||
|
|
||||||
loss = self._training_step(batch)
|
output = self._training_step(batch)
|
||||||
self._accelerator.backward(loss)
|
self._accelerator.backward(output.loss.mean())
|
||||||
|
|
||||||
if self._accelerator.sync_gradients and cfg.optimization.max_grad_norm > 0:
|
if self._accelerator.sync_gradients and cfg.optimization.max_grad_norm > 0:
|
||||||
self._accelerator.clip_grad_norm_(
|
self._accelerator.clip_grad_norm_(
|
||||||
@@ -244,9 +259,10 @@ class LtxvTrainer:
|
|||||||
# Update progress and log metrics
|
# Update progress and log metrics
|
||||||
current_lr = self._optimizer.param_groups[0]["lr"]
|
current_lr = self._optimizer.param_groups[0]["lr"]
|
||||||
step_time = (time.time() - step_start_time) * cfg.optimization.gradient_accumulation_steps
|
step_time = (time.time() - step_start_time) * cfg.optimization.gradient_accumulation_steps
|
||||||
|
step_loss = output.loss.detach().mean().item()
|
||||||
|
|
||||||
progress.update_training(
|
progress.update_training(
|
||||||
loss=loss.item(),
|
loss=step_loss,
|
||||||
lr=current_lr,
|
lr=current_lr,
|
||||||
step_time=step_time,
|
step_time=step_time,
|
||||||
advance=is_optimization_step,
|
advance=is_optimization_step,
|
||||||
@@ -254,14 +270,16 @@ class LtxvTrainer:
|
|||||||
|
|
||||||
# Log metrics to W&B (only on main process and optimization steps)
|
# Log metrics to W&B (only on main process and optimization steps)
|
||||||
if IS_MAIN_PROCESS and is_optimization_step:
|
if IS_MAIN_PROCESS and is_optimization_step:
|
||||||
self._log_metrics(
|
# Track per-element loss by sigma bucket
|
||||||
{
|
self._sigma_tracker.update(output.sigma.cpu().tolist(), output.loss.detach().cpu().tolist())
|
||||||
"train/loss": loss.item(),
|
metrics = {
|
||||||
"train/learning_rate": current_lr,
|
"train/loss": step_loss,
|
||||||
"train/step_time": step_time,
|
"train/learning_rate": current_lr,
|
||||||
"train/global_step": self._global_step,
|
"train/step_time": step_time,
|
||||||
}
|
"train/global_step": self._global_step,
|
||||||
)
|
}
|
||||||
|
metrics.update(self._sigma_tracker.get_metrics())
|
||||||
|
self._log_metrics(metrics)
|
||||||
|
|
||||||
# Fallback logging when progress bars are disabled
|
# Fallback logging when progress bars are disabled
|
||||||
if disable_progress_bars and IS_MAIN_PROCESS and self._global_step % 20 == 0:
|
if disable_progress_bars and IS_MAIN_PROCESS and self._global_step % 20 == 0:
|
||||||
@@ -274,7 +292,7 @@ class LtxvTrainer:
|
|||||||
total_time = "calculating..."
|
total_time = "calculating..."
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Step {self._global_step}/{cfg.optimization.steps} - "
|
f"Step {self._global_step}/{cfg.optimization.steps} - "
|
||||||
f"Loss: {loss.item():.4f}, LR: {current_lr:.2e}, "
|
f"Loss: {step_loss:.4f}, LR: {current_lr:.2e}, "
|
||||||
f"Time/Step: {step_time:.2f}s, Total Time: {total_time}",
|
f"Time/Step: {step_time:.2f}s, Total Time: {total_time}",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -330,7 +348,7 @@ class LtxvTrainer:
|
|||||||
|
|
||||||
return saved_path, stats
|
return saved_path, stats
|
||||||
|
|
||||||
def _training_step(self, batch: dict[str, dict[str, Tensor]]) -> Tensor:
|
def _training_step(self, batch: dict[str, dict[str, Tensor]]) -> TrainingStepOutput:
|
||||||
"""Perform a single training step using the configured strategy."""
|
"""Perform a single training step using the configured strategy."""
|
||||||
# Apply embedding connectors to transform pre-computed text embeddings
|
# Apply embedding connectors to transform pre-computed text embeddings
|
||||||
conditions = batch["conditions"]
|
conditions = batch["conditions"]
|
||||||
@@ -366,8 +384,9 @@ class LtxvTrainer:
|
|||||||
|
|
||||||
# Use strategy to compute loss
|
# Use strategy to compute loss
|
||||||
loss = self._training_strategy.compute_loss(video_pred, audio_pred, model_inputs)
|
loss = self._training_strategy.compute_loss(video_pred, audio_pred, model_inputs)
|
||||||
|
sigma = model_inputs.video.sigma.detach() if model_inputs.video.enabled else model_inputs.audio.sigma.detach()
|
||||||
|
|
||||||
return loss
|
return TrainingStepOutput(loss=loss, sigma=sigma)
|
||||||
|
|
||||||
@free_gpu_memory_context(after=True)
|
@free_gpu_memory_context(after=True)
|
||||||
def _load_text_encoder_and_cache_embeddings(self) -> list[CachedPromptEmbeddings] | None:
|
def _load_text_encoder_and_cache_embeddings(self) -> list[CachedPromptEmbeddings] | None:
|
||||||
@@ -1064,8 +1083,8 @@ class LtxvTrainer:
|
|||||||
def _save_training_state(self, save_dir: Path) -> None:
|
def _save_training_state(self, save_dir: Path) -> None:
|
||||||
"""Save training state alongside checkpoint for resume.
|
"""Save training state alongside checkpoint for resume.
|
||||||
Respects checkpoints.save_training_state config:
|
Respects checkpoints.save_training_state config:
|
||||||
- "full": optimizer + scheduler + RNG + step
|
- "full": optimizer + scheduler + RNG + step + wandb_run_id
|
||||||
- "minimal": scheduler + RNG + step only
|
- "minimal": scheduler + RNG + step + wandb_run_id
|
||||||
- "off": skip entirely
|
- "off": skip entirely
|
||||||
"""
|
"""
|
||||||
if not IS_MAIN_PROCESS:
|
if not IS_MAIN_PROCESS:
|
||||||
@@ -1101,6 +1120,7 @@ class LtxvTrainer:
|
|||||||
),
|
),
|
||||||
lr_scheduler_state_dict=self._lr_scheduler.state_dict() if self._lr_scheduler is not None else None,
|
lr_scheduler_state_dict=self._lr_scheduler.state_dict() if self._lr_scheduler is not None else None,
|
||||||
optimizer_state_dict=optimizer_state,
|
optimizer_state_dict=optimizer_state,
|
||||||
|
wandb_run_id=self._wandb_run.id if self._wandb_run is not None else None,
|
||||||
)
|
)
|
||||||
|
|
||||||
state_path = save_dir / f"training_state_step_{self._global_step:05d}.pt"
|
state_path = save_dir / f"training_state_step_{self._global_step:05d}.pt"
|
||||||
@@ -1166,20 +1186,24 @@ class LtxvTrainer:
|
|||||||
|
|
||||||
logger.info(f"💾 Training configuration saved to: {config_path.relative_to(self._config.output_dir)}")
|
logger.info(f"💾 Training configuration saved to: {config_path.relative_to(self._config.output_dir)}")
|
||||||
|
|
||||||
def _init_wandb(self) -> None:
|
def _init_wandb(self, resume_run_id: str | None = None) -> None:
|
||||||
"""Initialize Weights & Biases run."""
|
"""Initialize Weights & Biases run."""
|
||||||
if not self._config.wandb.enabled or not IS_MAIN_PROCESS:
|
if not self._config.wandb.enabled or not IS_MAIN_PROCESS:
|
||||||
self._wandb_run = None
|
self._wandb_run = None
|
||||||
return
|
return
|
||||||
|
|
||||||
wandb_config = self._config.wandb
|
wandb_config = self._config.wandb
|
||||||
run = wandb.init(
|
init_kwargs: dict[str, Any] = {
|
||||||
project=wandb_config.project,
|
"project": wandb_config.project,
|
||||||
entity=wandb_config.entity,
|
"entity": wandb_config.entity,
|
||||||
name=Path(self._config.output_dir).name,
|
"name": Path(self._config.output_dir).name,
|
||||||
tags=wandb_config.tags,
|
"tags": wandb_config.tags,
|
||||||
config=self._config.model_dump(),
|
"config": self._config.model_dump(),
|
||||||
)
|
}
|
||||||
|
if resume_run_id is not None:
|
||||||
|
init_kwargs["id"] = resume_run_id
|
||||||
|
init_kwargs["resume"] = "allow"
|
||||||
|
run = wandb.init(**init_kwargs)
|
||||||
self._wandb_run = run
|
self._wandb_run = run
|
||||||
|
|
||||||
def _log_metrics(self, metrics: dict[str, float]) -> None:
|
def _log_metrics(self, metrics: dict[str, float]) -> None:
|
||||||
|
|||||||
@@ -28,6 +28,7 @@ class TrainingState(BaseModel):
|
|||||||
rng_states: RngStates
|
rng_states: RngStates
|
||||||
lr_scheduler_state_dict: dict[str, Any] | None = None
|
lr_scheduler_state_dict: dict[str, Any] | None = None
|
||||||
optimizer_state_dict: dict[str, Any] | None = None
|
optimizer_state_dict: dict[str, Any] | None = None
|
||||||
|
wandb_run_id: str | None = None
|
||||||
|
|
||||||
def to_save_dict(self) -> dict[str, Any]:
|
def to_save_dict(self) -> dict[str, Any]:
|
||||||
"""Build dict suitable for torch.save -- recurses BaseModel sub-models, passes tensors/dicts through."""
|
"""Build dict suitable for torch.save -- recurses BaseModel sub-models, passes tensors/dicts through."""
|
||||||
@@ -48,4 +49,5 @@ class TrainingState(BaseModel):
|
|||||||
rng_states=RngStates(**data["rng_states"]),
|
rng_states=RngStates(**data["rng_states"]),
|
||||||
lr_scheduler_state_dict=data.get("lr_scheduler_state_dict"),
|
lr_scheduler_state_dict=data.get("lr_scheduler_state_dict"),
|
||||||
optimizer_state_dict=data.get("optimizer_state_dict"),
|
optimizer_state_dict=data.get("optimizer_state_dict"),
|
||||||
|
wandb_run_id=data.get("wandb_run_id"),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -125,7 +125,8 @@ class TrainingStrategy(ABC):
|
|||||||
audio_pred: Audio prediction from the transformer model (None for video-only)
|
audio_pred: Audio prediction from the transformer model (None for video-only)
|
||||||
inputs: The prepared model inputs containing targets and masks
|
inputs: The prepared model inputs containing targets and masks
|
||||||
Returns:
|
Returns:
|
||||||
Scalar loss tensor
|
Per-element loss tensor of shape [B,]. The trainer reduces to a scalar
|
||||||
|
before backward(). Returning unreduced loss enables per-sigma-bucket tracking.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def get_checkpoint_metadata(self) -> dict[str, Any]:
|
def get_checkpoint_metadata(self) -> dict[str, Any]:
|
||||||
|
|||||||
@@ -273,19 +273,19 @@ class TextToVideoStrategy(TrainingStrategy):
|
|||||||
audio_pred: Tensor | None,
|
audio_pred: Tensor | None,
|
||||||
inputs: ModelInputs,
|
inputs: ModelInputs,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
"""Compute masked MSE loss for video and optionally audio."""
|
"""Compute masked MSE loss for video and optionally audio. Returns [B,]."""
|
||||||
# Video loss
|
# Video loss: per-element mean over (seq, channels), [B,]
|
||||||
video_loss = (video_pred - inputs.video_targets).pow(2)
|
video_loss = (video_pred - inputs.video_targets).pow(2)
|
||||||
video_loss_mask = inputs.video_loss_mask.unsqueeze(-1).float()
|
video_loss_mask = inputs.video_loss_mask.unsqueeze(-1).float()
|
||||||
video_loss = video_loss.mul(video_loss_mask).div(video_loss_mask.mean())
|
masked = video_loss.mul(video_loss_mask)
|
||||||
video_loss = video_loss.mean()
|
video_loss = masked.mean(dim=[-2, -1]) / video_loss_mask.mean(dim=[-2, -1]).clamp(min=1e-8)
|
||||||
|
|
||||||
# If no audio, return video loss only
|
# If no audio, return video loss only
|
||||||
if not self.config.with_audio or audio_pred is None or inputs.audio_targets is None:
|
if not self.config.with_audio or audio_pred is None or inputs.audio_targets is None:
|
||||||
return video_loss
|
return video_loss
|
||||||
|
|
||||||
# Audio loss (no conditioning mask)
|
# Audio loss: per-element mean over (seq, channels), [B,]
|
||||||
audio_loss = (audio_pred - inputs.audio_targets).pow(2).mean()
|
audio_loss = (audio_pred - inputs.audio_targets).pow(2).mean(dim=[-2, -1])
|
||||||
|
|
||||||
# Combined loss
|
# Combined loss [B,]
|
||||||
return video_loss + audio_loss
|
return video_loss + audio_loss
|
||||||
|
|||||||
@@ -240,7 +240,7 @@ class VideoToVideoStrategy(TrainingStrategy):
|
|||||||
_audio_pred: Tensor | None,
|
_audio_pred: Tensor | None,
|
||||||
inputs: ModelInputs,
|
inputs: ModelInputs,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
"""Compute masked loss only on target portion."""
|
"""Compute masked loss only on target portion. Returns [B,]."""
|
||||||
# Extract target portion of prediction
|
# Extract target portion of prediction
|
||||||
ref_seq_len = inputs.ref_seq_len
|
ref_seq_len = inputs.ref_seq_len
|
||||||
target_pred = video_pred[:, ref_seq_len:, :]
|
target_pred = video_pred[:, ref_seq_len:, :]
|
||||||
@@ -248,14 +248,11 @@ class VideoToVideoStrategy(TrainingStrategy):
|
|||||||
# Get target portion of loss mask
|
# Get target portion of loss mask
|
||||||
target_loss_mask = inputs.video_loss_mask[:, ref_seq_len:]
|
target_loss_mask = inputs.video_loss_mask[:, ref_seq_len:]
|
||||||
|
|
||||||
# Compute loss
|
# Compute per-element loss [B,]
|
||||||
loss = (target_pred - inputs.video_targets).pow(2)
|
loss = (target_pred - inputs.video_targets).pow(2)
|
||||||
|
|
||||||
# Apply loss mask
|
|
||||||
loss_mask = target_loss_mask.unsqueeze(-1).float()
|
loss_mask = target_loss_mask.unsqueeze(-1).float()
|
||||||
loss = loss.mul(loss_mask).div(loss_mask.mean())
|
masked = loss.mul(loss_mask)
|
||||||
|
return masked.mean(dim=[-2, -1]) / loss_mask.mean(dim=[-2, -1]).clamp(min=1e-8)
|
||||||
return loss.mean()
|
|
||||||
|
|
||||||
def get_checkpoint_metadata(self) -> dict[str, Any]:
|
def get_checkpoint_metadata(self) -> dict[str, Any]:
|
||||||
"""Get metadata for checkpoint files."""
|
"""Get metadata for checkpoint files."""
|
||||||
|
|||||||
@@ -14,6 +14,9 @@ from torch import Tensor
|
|||||||
|
|
||||||
def get_video_frame_count(video_path: str | Path) -> int:
|
def get_video_frame_count(video_path: str | Path) -> int:
|
||||||
"""Get the number of frames in a video file.
|
"""Get the number of frames in a video file.
|
||||||
|
Tries three approaches in order: stream metadata, duration*fps estimate,
|
||||||
|
full decode. The estimate may be off by a few frames for VFR videos or
|
||||||
|
containers with edit lists — exact for the min_frames filtering use case.
|
||||||
Args:
|
Args:
|
||||||
video_path: Path to the video file
|
video_path: Path to the video file
|
||||||
Returns:
|
Returns:
|
||||||
@@ -21,11 +24,19 @@ def get_video_frame_count(video_path: str | Path) -> int:
|
|||||||
"""
|
"""
|
||||||
with av.open(str(video_path)) as container:
|
with av.open(str(video_path)) as container:
|
||||||
video_stream = container.streams.video[0]
|
video_stream = container.streams.video[0]
|
||||||
frame_count = video_stream.frames
|
|
||||||
if frame_count == 0:
|
if video_stream.frames > 0:
|
||||||
# Fallback: count frames by decoding
|
return video_stream.frames
|
||||||
frame_count = sum(1 for _ in container.decode(video=0))
|
|
||||||
return frame_count
|
# Fast estimate from container metadata (avoids full decode).
|
||||||
|
# Uses Fraction arithmetic to prevent float precision loss.
|
||||||
|
rate = video_stream.average_rate or video_stream.base_rate
|
||||||
|
if video_stream.duration and video_stream.time_base and rate:
|
||||||
|
duration = Fraction(video_stream.duration) * Fraction(video_stream.time_base)
|
||||||
|
return round(duration * Fraction(rate))
|
||||||
|
|
||||||
|
# Last resort: full decode (very slow for 4K)
|
||||||
|
return sum(1 for _ in container.decode(video=0))
|
||||||
|
|
||||||
|
|
||||||
def read_video(video_path: str | Path, max_frames: int | None = None) -> tuple[Tensor, float]:
|
def read_video(video_path: str | Path, max_frames: int | None = None) -> tuple[Tensor, float]:
|
||||||
|
|||||||
@@ -2063,7 +2063,7 @@ wheels = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "ltx-core"
|
name = "ltx-core"
|
||||||
version = "1.1.0"
|
version = "1.1.1"
|
||||||
source = { editable = "packages/ltx-core" }
|
source = { editable = "packages/ltx-core" }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
{ name = "accelerate" },
|
{ name = "accelerate" },
|
||||||
@@ -2121,7 +2121,7 @@ dev = [{ name = "scikit-image", specifier = ">=0.25.2" }]
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "ltx-pipelines"
|
name = "ltx-pipelines"
|
||||||
version = "1.1.0"
|
version = "1.1.1"
|
||||||
source = { editable = "packages/ltx-pipelines" }
|
source = { editable = "packages/ltx-pipelines" }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
{ name = "av" },
|
{ name = "av" },
|
||||||
@@ -2141,7 +2141,7 @@ requires-dist = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "ltx-trainer"
|
name = "ltx-trainer"
|
||||||
version = "1.1.0"
|
version = "1.1.1"
|
||||||
source = { editable = "packages/ltx-trainer" }
|
source = { editable = "packages/ltx-trainer" }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
{ name = "accelerate" },
|
{ name = "accelerate" },
|
||||||
|
|||||||
Reference in New Issue
Block a user