Automated PR - 2026-03-04
This commit is contained in:
@@ -12,7 +12,7 @@ LTX-2 Pipelines provides production-ready implementations that abstract away the
|
||||
|
||||
**Key Features:**
|
||||
|
||||
- 🎬 **Multiple Pipeline Types**: Text-to-video, image-to-video, video-to-video, and keyframe interpolation
|
||||
- 🎬 **Multiple Pipeline Types**: Text-to-video, image-to-video, video-to-video, audio-to-video, keyframe interpolation, and retake
|
||||
- ⚡ **Optimized Performance**: Support for FP8 transformers, gradient estimation, and memory optimization
|
||||
- 🎯 **Production Ready**: Two-stage pipelines for best quality output
|
||||
- 🔧 **LoRA Support**: Easy integration with trained LoRA adapters
|
||||
@@ -23,7 +23,7 @@ LTX-2 Pipelines provides production-ready implementations that abstract away the
|
||||
|
||||
## 🚀 Quick Start
|
||||
|
||||
`ltx-pipelines` provides ready-made inference pipelines for text-to-video, image-to-video, video-to-video, and keyframe interpolation. Built using building blocks from [`ltx-core`](../ltx-core/), these pipelines handle the complete inference flow including model loading, encoding, decoding, and file I/O.
|
||||
`ltx-pipelines` provides ready-made inference pipelines for text-to-video, image-to-video, video-to-video, audio-to-video, keyframe interpolation, and retake. Built using building blocks from [`ltx-core`](../ltx-core/), these pipelines handle the complete inference flow including model loading, encoding, decoding, and file I/O.
|
||||
|
||||
## 🔧 Installation
|
||||
|
||||
@@ -56,10 +56,13 @@ python -m ltx_pipelines.ti2vid_two_stages --help
|
||||
Available pipeline modules:
|
||||
|
||||
- `ltx_pipelines.ti2vid_two_stages` - Two-stage text/image-to-video (recommended).
|
||||
- `ltx_pipelines.ti2vid_two_stages_res2s` - Two-stage text/image-to-video (use 2 times less steps).
|
||||
- `ltx_pipelines.ti2vid_one_stage` - Single-stage text/image-to-video.
|
||||
- `ltx_pipelines.distilled` - Fast text/image-to-video pipeline using only the distilled model.
|
||||
- `ltx_pipelines.ic_lora` - Video-to-video with IC-LoRA.
|
||||
- `ltx_pipelines.keyframe_interpolation` - Keyframe interpolation.
|
||||
- `ltx_pipelines.a2vid_two_stage` - Audio-to-video generation conditioned on an input audio.
|
||||
- `ltx_pipelines.retake` - Regenerate a time region of an existing video.
|
||||
|
||||
Use `--help` with any pipeline module to see all available options and parameters.
|
||||
|
||||
@@ -70,6 +73,12 @@ Use `--help` with any pipeline module to see all available options and parameter
|
||||
### Quick Decision Tree
|
||||
|
||||
```text
|
||||
Do you have an existing video to modify?
|
||||
├─ YES → Use RetakePipeline (regenerate a specific time region)
|
||||
│
|
||||
Do you have an audio file to drive generation?
|
||||
├─ YES → Use A2VidPipelineTwoStage (audio-to-video)
|
||||
│
|
||||
Do you need to condition on existing images/videos?
|
||||
├─ YES → Do you have reference videos for video-to-video?
|
||||
│ ├─ YES → Use ICLoraPipeline
|
||||
@@ -85,17 +94,20 @@ Do you need to condition on existing images/videos?
|
||||
└─ YES → Use DistilledPipeline (with 8 predefined sigmas)
|
||||
```
|
||||
|
||||
> **Note:** [`TI2VidOneStagePipeline`](src/ltx_pipelines/ti2vid_one_stage.py) is primarily for educational purposes. For best quality, use two-stage pipelines ([`TI2VidTwoStagesPipeline`](src/ltx_pipelines/ti2vid_two_stages.py), [`ICLoraPipeline`](src/ltx_pipelines/ic_lora.py), [`KeyframeInterpolationPipeline`](src/ltx_pipelines/keyframe_interpolation.py), or [`DistilledPipeline`](src/ltx_pipelines/distilled.py)).
|
||||
> **Note:** [`TI2VidOneStagePipeline`](src/ltx_pipelines/ti2vid_one_stage.py) is primarily for educational purposes. For best quality, use two-stage pipelines ([`TI2VidTwoStagesPipeline`](src/ltx_pipelines/ti2vid_two_stages.py), [`TI2VidTwoStagesRes2sPipeline`](src/ltx_pipelines/ti2vid_two_stages_res2s.py), [`ICLoraPipeline`](src/ltx_pipelines/ic_lora.py), [`KeyframeInterpolationPipeline`](src/ltx_pipelines/keyframe_interpolation.py), [`A2VidPipelineTwoStage`](src/ltx_pipelines/a2vid_two_stage.py), or [`DistilledPipeline`](src/ltx_pipelines/distilled.py)). For editing existing videos, use [`RetakePipeline`](src/ltx_pipelines/retake.py).
|
||||
|
||||
### Features Comparison
|
||||
|
||||
| Pipeline | Stages | [Multimodal Guidance](#%EF%B8%8F-multimodal-guidance) | Upsampling | Conditioning | Best For |
|
||||
| -------- | ------ | --- | ---------- | ------------- | -------- |
|
||||
| **TI2VidTwoStagesPipeline** | 2 | ✅ | ✅ | Image | **Production quality** (recommended) |
|
||||
| **TI2VidTwoStagesRes2sPipeline** | 2 | ✅ | ✅ | Image | Same as above, res_2s sampler (fewer steps) |
|
||||
| **TI2VidOneStagePipeline** | 1 | ✅ | ❌ | Image | Educational, prototyping |
|
||||
| **DistilledPipeline** | 2 | ❌ | ✅ | Image | Fastest inference (8 sigmas) |
|
||||
| **ICLoraPipeline** | 2 | ✅ | ✅ | Image + Video | Video-to-video transformations |
|
||||
| **KeyframeInterpolationPipeline** | 2 | ✅ | ✅ | Keyframes | Animation, interpolation |
|
||||
| **A2VidPipelineTwoStage** | 2 | ✅ | ✅ | Audio + Image | Audio-driven video generation |
|
||||
| **RetakePipeline** | 1 | ✅ | ❌ | Source Video | Regenerating a time region of a video |
|
||||
|
||||
---
|
||||
|
||||
@@ -113,7 +125,19 @@ Two-stage generation: Stage 1 generates low-resolution video with [multimodal gu
|
||||
|
||||
---
|
||||
|
||||
### 2. TI2VidOneStagePipeline
|
||||
### 2. TI2VidTwoStagesRes2sPipeline
|
||||
|
||||
**Best for:** Same two-stage text/image-to-video as TI2VidTwoStagesPipeline but with a different sampler and step count.
|
||||
|
||||
**Source**: [`src/ltx_pipelines/ti2vid_two_stages_res2s.py`](src/ltx_pipelines/ti2vid_two_stages_res2s.py)
|
||||
|
||||
Uses the **res_2s** second-order sampler instead of Euler. Same stage structure (stage 1 at target resolution with CFG, stage 2 upsampling with distilled LoRA) and image conditioning support. Typically allows fewer steps for comparable quality; trade-offs differ from the default Euler-based pipeline.
|
||||
|
||||
**Use when:** You want the same two-stage workflow with fewer steps or prefer the res_2s sampling behavior.
|
||||
|
||||
---
|
||||
|
||||
### 3. TI2VidOneStagePipeline
|
||||
|
||||
**Best for:** Educational purposes and quick prototyping.
|
||||
|
||||
@@ -127,7 +151,7 @@ Single-stage generation (no upsampling) with [multimodal guidance](#%EF%B8%8F-mu
|
||||
|
||||
---
|
||||
|
||||
### 3. DistilledPipeline
|
||||
### 4. DistilledPipeline
|
||||
|
||||
**Best for:** Fastest inference with good quality using a distilled model with predefined sigma schedule.
|
||||
|
||||
@@ -139,7 +163,7 @@ Two-stage generation with 8 predefined sigmas (8 steps in stage 1, 4 steps in st
|
||||
|
||||
---
|
||||
|
||||
### 4. ICLoraPipeline
|
||||
### 5. ICLoraPipeline
|
||||
|
||||
**Best for:** Video-to-video and image-to-video transformations using IC-LoRA.
|
||||
|
||||
@@ -147,11 +171,13 @@ Two-stage generation with 8 predefined sigmas (8 steps in stage 1, 4 steps in st
|
||||
|
||||
Two-stage generation with IC-LoRA support. Can condition on reference videos (video-to-video) or images at specific frames. CFG guidance in stage 1, upsampling in stage 2. Requires IC-LoRA trained model.
|
||||
|
||||
**Note:** ICLoraPipeline can only be used with a distilled model.
|
||||
|
||||
**Use when:** Video-to-video transformations, image-to-video with strong control, or when you have reference videos to guide generation.
|
||||
|
||||
---
|
||||
|
||||
### 5. KeyframeInterpolationPipeline
|
||||
### 6. KeyframeInterpolationPipeline
|
||||
|
||||
**Best for:** Generating videos by interpolating between keyframe images.
|
||||
|
||||
@@ -163,6 +189,36 @@ Two-stage generation with keyframe interpolation. Uses guiding latents (additive
|
||||
|
||||
---
|
||||
|
||||
### 7. A2VidPipelineTwoStage
|
||||
|
||||
**Best for:** Generating video driven by an input audio.
|
||||
|
||||
**Source**: [`src/ltx_pipelines/a2vid_two_stage.py`](src/ltx_pipelines/a2vid_two_stage.py)
|
||||
|
||||
Two-stage audio-to-video generation. Stage 1 generates video at half resolution with audio conditioning (video-only denoising with the audio frozen), then Stage 2 upsamples by 2x and refines the video while keeping the audio fixed, using a distilled LoRA. The input audio is encoded via the audio VAE and used as the initial audio latent, but the original audio waveform is passed through and returned in the output to preserve fidelity. Supports image conditioning and prompt enhancement.
|
||||
|
||||
**Extra CLI arguments:** `--audio-path` (required), `--audio-start-time`, `--audio-max-duration`.
|
||||
|
||||
**Use when:** You have an audio clip and want to generate a matching video, audio-reactive video generation, or music visualization.
|
||||
|
||||
---
|
||||
|
||||
### 8. RetakePipeline
|
||||
|
||||
**Best for:** Regenerating a specific time region of an existing video while keeping the rest unchanged.
|
||||
|
||||
**Source**: [`src/ltx_pipelines/retake.py`](src/ltx_pipelines/retake.py)
|
||||
|
||||
Single-stage generation that encodes the source video and audio into latents, applies a temporal region mask to mark `[start_time, end_time]` for regeneration, and denoises only the masked region from a text prompt. Content outside the time window is preserved. Supports independent control over video and audio regeneration (`regenerate_video`, `regenerate_audio` flags), and can use either the full model with CFG guidance or the distilled model with a fixed sigma schedule.
|
||||
|
||||
**Extra CLI arguments:** `--video-path` (required), `--start-time` (required), `--end-time` (required).
|
||||
|
||||
**Constraints:** Source video frame count must satisfy the 8k+1 format (e.g. 97, 193) and resolution must be multiples of 32.
|
||||
|
||||
**Use when:** You want to re-do a specific section of a generated video (e.g. fix a bad segment), selectively regenerate audio or video in a time window, or iterate on part of a result without re-generating the entire clip.
|
||||
|
||||
---
|
||||
|
||||
## 🎨 Conditioning Types
|
||||
|
||||
Pipelines use different conditioning methods from [`ltx-core`](../ltx-core/) for controlling generation. See the [ltx-core conditioning documentation](../ltx-core/README.md#conditioning--control) for details.
|
||||
@@ -250,28 +306,43 @@ audio_guider_params = MultiModalGuiderParams(
|
||||
|
||||
### Memory Optimization
|
||||
|
||||
**FP8 Transformer (Lower Memory Footprint):**
|
||||
**FP8 Quantization (Lower Memory Footprint):**
|
||||
|
||||
For smaller GPU memory footprint, use the `enable-fp8` flag and use the `PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True` environment variable.
|
||||
For smaller GPU memory footprint, use the `--quantization` flag and set `PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True`.
|
||||
|
||||
Two quantization policies are available:
|
||||
|
||||
| Policy | CLI Flag | Description |
|
||||
| ------ | -------- | ----------- |
|
||||
| **FP8 Cast** | `--quantization fp8-cast` | Downcasts transformer linear weights to FP8 during loading; upcasts on the fly during inference. No extra dependencies. |
|
||||
| **FP8 Scaled MM** | `--quantization fp8-scaled-mm` | Uses FP8 scaled matrix multiplication via TensorRT-LLM (`tensorrt_llm` must be installed). Best performance on Hopper GPUs. |
|
||||
|
||||
**CLI:**
|
||||
|
||||
```bash
|
||||
PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True python -m ltx_pipelines.ti2vid_one_stage --enable-fp8 --checkpoint-path=...
|
||||
# FP8 Cast (works on any GPU with FP8 support)
|
||||
PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True python -m ltx_pipelines.ti2vid_two_stages \
|
||||
--quantization fp8-cast --checkpoint-path=...
|
||||
|
||||
# FP8 Scaled MM (requires tensorrt_llm, best on Hopper GPUs)
|
||||
PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True python -m ltx_pipelines.ti2vid_two_stages \
|
||||
--quantization fp8-scaled-mm --checkpoint-path=...
|
||||
```
|
||||
|
||||
**Programmatically:**
|
||||
|
||||
When authoring custom scripts, pass the `fp8transformer` flag to pipeline classes or construct your own by analogy:
|
||||
When authoring custom scripts, pass a `QuantizationPolicy` to pipeline classes:
|
||||
|
||||
```python
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
|
||||
pipeline = TI2VidTwoStagesPipeline(
|
||||
checkpoint_path=ltx_model_path,
|
||||
distilled_lora=distilled_lora,
|
||||
spatial_upsampler_path=upsampler_path,
|
||||
gemma_root=gemma_root_path,
|
||||
loras=[],
|
||||
fp8transformer=True,
|
||||
quantization=QuantizationPolicy.fp8_cast(), # or QuantizationPolicy.fp8_scaled_mm()
|
||||
)
|
||||
pipeline(...)
|
||||
```
|
||||
@@ -299,7 +370,7 @@ By default, pipelines clean GPU memory (especially transformer weights) between
|
||||
Instead of the standard Euler denoising loop, you can use gradient estimation for fewer steps (~20-30 instead of 40):
|
||||
|
||||
```python
|
||||
from ltx_pipelines.utils.helpers import gradient_estimating_euler_denoising_loop
|
||||
from ltx_pipelines.utils import gradient_estimating_euler_denoising_loop
|
||||
|
||||
# Use gradient estimation denoising loop
|
||||
def denoising_loop(sigmas, video_state, audio_state, stepper):
|
||||
@@ -379,7 +450,7 @@ pipeline(
|
||||
num_inference_steps=40,
|
||||
video_guider_params=video_guider_params,
|
||||
audio_guider_params=audio_guider_params,
|
||||
images=[("input_image.jpg", 0, 1.0)], # Image at frame 0, strength 1.0
|
||||
images=[ImageConditioningInput("input_image.jpg", 0, 1.0, 33)], # Image at frame 0, strength 1.0, CRF 33
|
||||
)
|
||||
```
|
||||
|
||||
|
||||
@@ -6,21 +6,26 @@ This package provides ready-to-use pipelines for video generation:
|
||||
- DistilledPipeline: Fast distilled two-stage generation
|
||||
- ICLoraPipeline: Image/video conditioning with distilled LoRA
|
||||
- KeyframeInterpolationPipeline: Keyframe-based video interpolation
|
||||
- RetakePipeline: Regenerate a time region (retake) of an existing video
|
||||
- ModelLedger: Central coordinator for loading and building models
|
||||
For more detailed components and utilities, import from specific submodules
|
||||
like `ltx_pipelines.utils.media_io` or `ltx_pipelines.utils.constants`.
|
||||
"""
|
||||
|
||||
from ltx_pipelines.a2vid_two_stage import A2VidPipelineTwoStage
|
||||
from ltx_pipelines.distilled import DistilledPipeline
|
||||
from ltx_pipelines.ic_lora import ICLoraPipeline
|
||||
from ltx_pipelines.keyframe_interpolation import KeyframeInterpolationPipeline
|
||||
from ltx_pipelines.retake import RetakePipeline
|
||||
from ltx_pipelines.ti2vid_one_stage import TI2VidOneStagePipeline
|
||||
from ltx_pipelines.ti2vid_two_stages import TI2VidTwoStagesPipeline
|
||||
|
||||
__all__ = [
|
||||
"A2VidPipelineTwoStage",
|
||||
"DistilledPipeline",
|
||||
"ICLoraPipeline",
|
||||
"KeyframeInterpolationPipeline",
|
||||
"RetakePipeline",
|
||||
"TI2VidOneStagePipeline",
|
||||
"TI2VidTwoStagesPipeline",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,323 @@
|
||||
import logging
|
||||
from collections.abc import Iterator
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.components.diffusion_steps import EulerDiffusionStep
|
||||
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
|
||||
from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
|
||||
from ltx_core.model.upsampler import upsample_video
|
||||
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.model.video_vae import decode_video as vae_decode_video
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.text_encoders.gemma import encode_text
|
||||
from ltx_core.types import Audio, AudioLatentShape, LatentState, VideoPixelShape
|
||||
from ltx_pipelines.utils import ModelLedger
|
||||
from ltx_pipelines.utils.args import default_2_stage_arg_parser
|
||||
from ltx_pipelines.utils.constants import (
|
||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
||||
)
|
||||
from ltx_pipelines.utils.helpers import (
|
||||
assert_resolution,
|
||||
cleanup_memory,
|
||||
denoise_video_only,
|
||||
generate_enhanced_prompt,
|
||||
get_device,
|
||||
image_conditionings_by_replacing_latent,
|
||||
multi_modal_guider_denoising_func,
|
||||
simple_denoising_func,
|
||||
)
|
||||
from ltx_pipelines.utils.media_io import decode_audio_from_file, encode_video
|
||||
from ltx_pipelines.utils.samplers import euler_denoising_loop
|
||||
from ltx_pipelines.utils.types import PipelineComponents
|
||||
|
||||
device = get_device()
|
||||
|
||||
|
||||
class A2VidPipelineTwoStage:
|
||||
"""
|
||||
Two-stage audio to video generation pipeline.
|
||||
Stage 1 generates video at half the target resolution with audio conditioning
|
||||
(video-only denoising, audio frozen), then Stage 2 upsamples by 2x and refines
|
||||
both video and audio using a distilled LoRA for higher quality output.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
checkpoint_path: str,
|
||||
distilled_lora: list[LoraPathStrengthAndSDOps],
|
||||
spatial_upsampler_path: str,
|
||||
gemma_root: str,
|
||||
loras: list[LoraPathStrengthAndSDOps],
|
||||
device: torch.device = device,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
):
|
||||
self.device = device
|
||||
self.dtype = torch.bfloat16
|
||||
|
||||
self.stage_1_model_ledger = ModelLedger(
|
||||
dtype=self.dtype,
|
||||
device=device,
|
||||
checkpoint_path=checkpoint_path,
|
||||
gemma_root_path=gemma_root,
|
||||
spatial_upsampler_path=spatial_upsampler_path,
|
||||
loras=loras,
|
||||
quantization=quantization,
|
||||
)
|
||||
|
||||
self.stage_2_model_ledger = self.stage_1_model_ledger.with_loras(
|
||||
loras=distilled_lora,
|
||||
)
|
||||
|
||||
self.pipeline_components = PipelineComponents(
|
||||
dtype=self.dtype,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def __call__( # noqa: PLR0913
|
||||
self,
|
||||
prompt: str,
|
||||
negative_prompt: str,
|
||||
seed: int,
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
frame_rate: float,
|
||||
num_inference_steps: int,
|
||||
video_guider_params: MultiModalGuiderParams,
|
||||
images: list[tuple[str, int, float]],
|
||||
audio_path: str,
|
||||
audio_start_time: float = 0.0,
|
||||
audio_max_duration: float | None = None,
|
||||
tiling_config: TilingConfig | None = None,
|
||||
enhance_prompt: bool = False,
|
||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||
assert_resolution(height=height, width=width, is_two_stage=True)
|
||||
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
noiser = GaussianNoiser(generator=generator)
|
||||
stepper = EulerDiffusionStep()
|
||||
dtype = torch.bfloat16
|
||||
|
||||
text_encoder = self.stage_1_model_ledger.text_encoder()
|
||||
if enhance_prompt:
|
||||
prompt = generate_enhanced_prompt(text_encoder, prompt, images[0][0] if len(images) > 0 else None)
|
||||
context_p, context_n = encode_text(text_encoder, prompts=[prompt, negative_prompt])
|
||||
v_context_p, a_context_p = context_p
|
||||
v_context_n, _ = context_n
|
||||
|
||||
torch.cuda.synchronize()
|
||||
del text_encoder
|
||||
cleanup_memory()
|
||||
|
||||
# Encode audio.
|
||||
decoded_audio = decode_audio_from_file(audio_path, self.device, audio_start_time, audio_max_duration)
|
||||
encoded_audio_latent = vae_encode_audio(decoded_audio, self.stage_1_model_ledger.audio_encoder())
|
||||
audio_shape = AudioLatentShape.from_duration(batch=1, duration=num_frames / frame_rate, channels=8, mel_bins=16)
|
||||
encoded_audio_latent = encoded_audio_latent[:, :, : audio_shape.frames]
|
||||
|
||||
cleanup_memory()
|
||||
# Stage 1: Initial low resolution video generation with audio conditioning.
|
||||
transformer = self.stage_1_model_ledger.transformer()
|
||||
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
|
||||
|
||||
def first_stage_denoising_loop(
|
||||
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
|
||||
) -> tuple[LatentState, LatentState]:
|
||||
return euler_denoising_loop(
|
||||
sigmas=sigmas,
|
||||
video_state=video_state,
|
||||
audio_state=audio_state,
|
||||
stepper=stepper,
|
||||
denoise_fn=multi_modal_guider_denoising_func(
|
||||
video_guider=MultiModalGuider(
|
||||
params=video_guider_params,
|
||||
negative_context=v_context_n,
|
||||
),
|
||||
audio_guider=MultiModalGuider(
|
||||
params=MultiModalGuiderParams(),
|
||||
),
|
||||
v_context=v_context_p,
|
||||
a_context=a_context_p,
|
||||
transformer=transformer, # noqa: F821
|
||||
),
|
||||
)
|
||||
|
||||
stage_1_output_shape = VideoPixelShape(
|
||||
batch=1,
|
||||
frames=num_frames,
|
||||
width=width // 2,
|
||||
height=height // 2,
|
||||
fps=frame_rate,
|
||||
)
|
||||
|
||||
video_encoder = self.stage_1_model_ledger.video_encoder()
|
||||
stage_1_conditionings = image_conditionings_by_replacing_latent(
|
||||
images=images,
|
||||
height=stage_1_output_shape.height,
|
||||
width=stage_1_output_shape.width,
|
||||
video_encoder=video_encoder,
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
video_state = denoise_video_only(
|
||||
output_shape=stage_1_output_shape,
|
||||
conditionings=stage_1_conditionings,
|
||||
noiser=noiser,
|
||||
sigmas=sigmas,
|
||||
stepper=stepper,
|
||||
denoising_loop_fn=first_stage_denoising_loop,
|
||||
components=self.pipeline_components,
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
initial_audio_latent=encoded_audio_latent,
|
||||
)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
del transformer
|
||||
cleanup_memory()
|
||||
|
||||
# Stage 2: Upsample and refine the video at higher resolution with distilled LoRA.
|
||||
upscaled_video_latent = upsample_video(
|
||||
latent=video_state.latent[:1],
|
||||
video_encoder=video_encoder,
|
||||
upsampler=self.stage_2_model_ledger.spatial_upsampler(),
|
||||
)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
cleanup_memory()
|
||||
|
||||
transformer = self.stage_2_model_ledger.transformer()
|
||||
distilled_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
|
||||
|
||||
def second_stage_denoising_loop(
|
||||
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
|
||||
) -> tuple[LatentState, LatentState]:
|
||||
return euler_denoising_loop(
|
||||
sigmas=sigmas,
|
||||
video_state=video_state,
|
||||
audio_state=audio_state,
|
||||
stepper=stepper,
|
||||
denoise_fn=simple_denoising_func(
|
||||
video_context=v_context_p,
|
||||
audio_context=a_context_p,
|
||||
transformer=transformer, # noqa: F821
|
||||
),
|
||||
)
|
||||
|
||||
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
|
||||
stage_2_conditionings = image_conditionings_by_replacing_latent(
|
||||
images=images,
|
||||
height=stage_2_output_shape.height,
|
||||
width=stage_2_output_shape.width,
|
||||
video_encoder=video_encoder,
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
)
|
||||
video_state = denoise_video_only(
|
||||
output_shape=stage_2_output_shape,
|
||||
conditionings=stage_2_conditionings,
|
||||
noiser=noiser,
|
||||
sigmas=distilled_sigmas,
|
||||
stepper=stepper,
|
||||
denoising_loop_fn=second_stage_denoising_loop,
|
||||
components=self.pipeline_components,
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
noise_scale=distilled_sigmas[0],
|
||||
initial_video_latent=upscaled_video_latent,
|
||||
initial_audio_latent=encoded_audio_latent,
|
||||
)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
del transformer
|
||||
del video_encoder
|
||||
cleanup_memory()
|
||||
|
||||
decoded_video = vae_decode_video(
|
||||
video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config, generator
|
||||
)
|
||||
|
||||
# Return the original input audio instead of VAE-decoded audio to preserve fidelity.
|
||||
# decode_audio_from_file already returns normalised [-1, 1] float values.
|
||||
original_audio = Audio(waveform=decoded_audio.waveform.squeeze(0), sampling_rate=decoded_audio.sampling_rate)
|
||||
|
||||
return decoded_video, original_audio
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
parser = default_2_stage_arg_parser()
|
||||
parser.add_argument(
|
||||
"--audio-path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to the audio file to condition the video generation.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-start-time",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Start time in seconds to read audio from (default: 0.0).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-max-duration",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Maximum audio duration in seconds. Defaults to video duration (num_frames / frame_rate).",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
pipeline = A2VidPipelineTwoStage(
|
||||
checkpoint_path=args.checkpoint_path,
|
||||
distilled_lora=args.distilled_lora,
|
||||
spatial_upsampler_path=args.spatial_upsampler_path,
|
||||
gemma_root=args.gemma_root,
|
||||
loras=args.lora,
|
||||
quantization=args.quantization,
|
||||
)
|
||||
tiling_config = TilingConfig.default()
|
||||
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
|
||||
video, audio = pipeline(
|
||||
prompt=args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
seed=args.seed,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
frame_rate=args.frame_rate,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
video_guider_params=MultiModalGuiderParams(
|
||||
cfg_scale=args.video_cfg_guidance_scale,
|
||||
stg_scale=args.video_stg_guidance_scale,
|
||||
rescale_scale=args.video_rescale_scale,
|
||||
modality_scale=args.a2v_guidance_scale,
|
||||
skip_step=args.video_skip_step,
|
||||
stg_blocks=args.video_stg_blocks,
|
||||
),
|
||||
images=args.images,
|
||||
tiling_config=tiling_config,
|
||||
enhance_prompt=args.enhance_prompt,
|
||||
audio_path=args.audio_path,
|
||||
audio_start_time=args.audio_start_time,
|
||||
audio_max_duration=args.audio_max_duration
|
||||
if args.audio_max_duration is not None
|
||||
else args.num_frames / args.frame_rate,
|
||||
)
|
||||
|
||||
encode_video(
|
||||
video=video,
|
||||
fps=args.frame_rate,
|
||||
audio=audio,
|
||||
output_path=args.output_path,
|
||||
video_chunks_number=video_chunks_number,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -13,19 +13,22 @@ from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.model.video_vae import decode_video as vae_decode_video
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.text_encoders.gemma import encode_text
|
||||
from ltx_core.types import LatentState, VideoPixelShape
|
||||
from ltx_pipelines.utils import ModelLedger
|
||||
from ltx_pipelines.utils.args import default_2_stage_distilled_arg_parser
|
||||
from ltx_core.types import Audio, LatentState, VideoPixelShape
|
||||
from ltx_pipelines.utils import ModelLedger, euler_denoising_loop
|
||||
from ltx_pipelines.utils.args import (
|
||||
ImageConditioningInput,
|
||||
default_2_stage_distilled_arg_parser,
|
||||
detect_checkpoint_path,
|
||||
)
|
||||
from ltx_pipelines.utils.constants import (
|
||||
AUDIO_SAMPLE_RATE,
|
||||
DISTILLED_SIGMA_VALUES,
|
||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
||||
detect_params,
|
||||
)
|
||||
from ltx_pipelines.utils.helpers import (
|
||||
assert_resolution,
|
||||
cleanup_memory,
|
||||
denoise_audio_video,
|
||||
euler_denoising_loop,
|
||||
generate_enhanced_prompt,
|
||||
get_device,
|
||||
image_conditionings_by_replacing_latent,
|
||||
@@ -40,13 +43,13 @@ device = get_device()
|
||||
class DistilledPipeline:
|
||||
"""
|
||||
Two-stage distilled video generation pipeline.
|
||||
Stage 1 generates video at the target resolution, then Stage 2 upsamples
|
||||
Stage 1 generates video at half of the target resolution, then Stage 2 upsamples
|
||||
by 2x and refines with additional denoising steps for higher quality output.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
checkpoint_path: str,
|
||||
distilled_checkpoint_path: str,
|
||||
gemma_root: str,
|
||||
spatial_upsampler_path: str,
|
||||
loras: list[LoraPathStrengthAndSDOps],
|
||||
@@ -59,7 +62,7 @@ class DistilledPipeline:
|
||||
self.model_ledger = ModelLedger(
|
||||
dtype=self.dtype,
|
||||
device=device,
|
||||
checkpoint_path=checkpoint_path,
|
||||
checkpoint_path=distilled_checkpoint_path,
|
||||
spatial_upsampler_path=spatial_upsampler_path,
|
||||
gemma_root_path=gemma_root,
|
||||
loras=loras,
|
||||
@@ -79,10 +82,10 @@ class DistilledPipeline:
|
||||
width: int,
|
||||
num_frames: int,
|
||||
frame_rate: float,
|
||||
images: list[tuple[str, int, float]],
|
||||
images: list[ImageConditioningInput],
|
||||
tiling_config: TilingConfig | None = None,
|
||||
enhance_prompt: bool = False,
|
||||
) -> tuple[Iterator[torch.Tensor], torch.Tensor]:
|
||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||
assert_resolution(height=height, width=width, is_two_stage=True)
|
||||
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
@@ -198,10 +201,12 @@ class DistilledPipeline:
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
parser = default_2_stage_distilled_arg_parser()
|
||||
checkpoint_path = detect_checkpoint_path(distilled=True)
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_2_stage_distilled_arg_parser(params=params)
|
||||
args = parser.parse_args()
|
||||
pipeline = DistilledPipeline(
|
||||
checkpoint_path=args.checkpoint_path,
|
||||
distilled_checkpoint_path=args.distilled_checkpoint_path,
|
||||
spatial_upsampler_path=args.spatial_upsampler_path,
|
||||
gemma_root=args.gemma_root,
|
||||
loras=args.lora,
|
||||
@@ -225,7 +230,6 @@ def main() -> None:
|
||||
video=video,
|
||||
fps=args.frame_rate,
|
||||
audio=audio,
|
||||
audio_sample_rate=AUDIO_SAMPLE_RATE,
|
||||
output_path=args.output_path,
|
||||
video_chunks_number=video_chunks_number,
|
||||
)
|
||||
|
||||
@@ -2,12 +2,17 @@ import logging
|
||||
from collections.abc import Iterator
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from safetensors import safe_open
|
||||
|
||||
from ltx_core.components.diffusion_steps import EulerDiffusionStep
|
||||
from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.conditioning import ConditioningItem, VideoConditionByReferenceLatent
|
||||
from ltx_core.conditioning import (
|
||||
ConditioningItem,
|
||||
ConditioningItemAttentionStrengthWrapper,
|
||||
VideoConditionByReferenceLatent,
|
||||
)
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.model.audio_vae import decode_audio as vae_decode_audio
|
||||
from ltx_core.model.upsampler import upsample_video
|
||||
@@ -15,15 +20,9 @@ from ltx_core.model.video_vae import TilingConfig, VideoEncoder, get_video_chunk
|
||||
from ltx_core.model.video_vae import decode_video as vae_decode_video
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.text_encoders.gemma import encode_text
|
||||
from ltx_core.types import LatentState, VideoPixelShape
|
||||
from ltx_pipelines.utils import ModelLedger
|
||||
from ltx_pipelines.utils.args import VideoConditioningAction, default_2_stage_distilled_arg_parser
|
||||
from ltx_pipelines.utils.constants import (
|
||||
AUDIO_SAMPLE_RATE,
|
||||
DISTILLED_SIGMA_VALUES,
|
||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
||||
)
|
||||
from ltx_pipelines.utils.helpers import (
|
||||
from ltx_core.types import Audio, LatentState, VideoLatentShape, VideoPixelShape
|
||||
from ltx_pipelines.utils import (
|
||||
ModelLedger,
|
||||
assert_resolution,
|
||||
cleanup_memory,
|
||||
denoise_audio_video,
|
||||
@@ -33,6 +32,18 @@ from ltx_pipelines.utils.helpers import (
|
||||
image_conditionings_by_replacing_latent,
|
||||
simple_denoising_func,
|
||||
)
|
||||
from ltx_pipelines.utils.args import (
|
||||
ImageConditioningInput,
|
||||
VideoConditioningAction,
|
||||
VideoMaskConditioningAction,
|
||||
default_2_stage_distilled_arg_parser,
|
||||
detect_checkpoint_path,
|
||||
)
|
||||
from ltx_pipelines.utils.constants import (
|
||||
DISTILLED_SIGMA_VALUES,
|
||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
||||
detect_params,
|
||||
)
|
||||
from ltx_pipelines.utils.media_io import encode_video, load_video_conditioning
|
||||
from ltx_pipelines.utils.types import PipelineComponents
|
||||
|
||||
@@ -45,13 +56,14 @@ class ICLoraPipeline:
|
||||
Allows conditioning the generated video on control signals such as depth maps,
|
||||
human pose, or image edges via the video_conditioning parameter.
|
||||
The specific IC-LoRA model should be provided via the loras parameter.
|
||||
Stage 1 generates video at the target resolution, then Stage 2 upsamples
|
||||
Stage 1 generates video at half of the target resolution, then Stage 2 upsamples
|
||||
by 2x and refines with additional denoising steps for higher quality output.
|
||||
Both stages use distilled models for efficiency.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
checkpoint_path: str,
|
||||
distilled_checkpoint_path: str,
|
||||
spatial_upsampler_path: str,
|
||||
gemma_root: str,
|
||||
loras: list[LoraPathStrengthAndSDOps],
|
||||
@@ -62,7 +74,7 @@ class ICLoraPipeline:
|
||||
self.stage_1_model_ledger = ModelLedger(
|
||||
dtype=self.dtype,
|
||||
device=device,
|
||||
checkpoint_path=checkpoint_path,
|
||||
checkpoint_path=distilled_checkpoint_path,
|
||||
spatial_upsampler_path=spatial_upsampler_path,
|
||||
gemma_root_path=gemma_root,
|
||||
loras=loras,
|
||||
@@ -71,7 +83,7 @@ class ICLoraPipeline:
|
||||
self.stage_2_model_ledger = ModelLedger(
|
||||
dtype=self.dtype,
|
||||
device=device,
|
||||
checkpoint_path=checkpoint_path,
|
||||
checkpoint_path=distilled_checkpoint_path,
|
||||
spatial_upsampler_path=spatial_upsampler_path,
|
||||
gemma_root_path=gemma_root,
|
||||
loras=[],
|
||||
@@ -98,8 +110,7 @@ class ICLoraPipeline:
|
||||
)
|
||||
self.reference_downscale_factor = scale
|
||||
|
||||
@torch.inference_mode()
|
||||
def __call__(
|
||||
def __call__( # noqa: PLR0913
|
||||
self,
|
||||
prompt: str,
|
||||
seed: int,
|
||||
@@ -107,12 +118,51 @@ class ICLoraPipeline:
|
||||
width: int,
|
||||
num_frames: int,
|
||||
frame_rate: float,
|
||||
images: list[tuple[str, int, float]],
|
||||
images: list[ImageConditioningInput],
|
||||
video_conditioning: list[tuple[str, float]],
|
||||
enhance_prompt: bool = False,
|
||||
tiling_config: TilingConfig | None = None,
|
||||
) -> tuple[Iterator[torch.Tensor], torch.Tensor]:
|
||||
conditioning_attention_strength: float = 1.0,
|
||||
skip_stage_2: bool = False,
|
||||
conditioning_attention_mask: torch.Tensor | None = None,
|
||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||
"""
|
||||
Generate video with IC-LoRA conditioning.
|
||||
Args:
|
||||
prompt: Text prompt for video generation.
|
||||
seed: Random seed for reproducibility.
|
||||
height: Output video height in pixels (must be divisible by 64).
|
||||
width: Output video width in pixels (must be divisible by 64).
|
||||
num_frames: Number of frames to generate.
|
||||
frame_rate: Output video frame rate.
|
||||
images: List of (path, frame_idx, strength) tuples for image conditioning.
|
||||
video_conditioning: List of (path, strength) tuples for IC-LoRA video conditioning.
|
||||
enhance_prompt: Whether to enhance the prompt using the text encoder.
|
||||
tiling_config: Optional tiling configuration for VAE decoding.
|
||||
conditioning_attention_strength: Scale factor for IC-LoRA conditioning attention.
|
||||
Controls how strongly the conditioning video influences the output.
|
||||
0.0 = ignore conditioning, 1.0 = full conditioning influence. Default 1.0.
|
||||
When conditioning_attention_mask is provided, the mask is multiplied by
|
||||
this strength before being passed to the conditioning items.
|
||||
skip_stage_2: If True, skip Stage 2 upsampling and refinement. Output will be
|
||||
at half resolution (height//2, width//2). Default is False.
|
||||
conditioning_attention_mask: Optional pixel-space attention mask with the same
|
||||
spatial-temporal dimensions as the input reference video. Shape should be
|
||||
(B, 1, F, H, W) or (1, 1, F, H, W) where F, H, W match the reference
|
||||
video's pixel dimensions. Values in [0, 1].
|
||||
The mask is downsampled to latent space using VAE scale factors (with
|
||||
causal temporal handling for the first frame), then multiplied by
|
||||
conditioning_attention_strength.
|
||||
When None (default): scalar conditioning_attention_strength is used
|
||||
directly.
|
||||
Returns:
|
||||
Tuple of (video_iterator, audio_tensor).
|
||||
"""
|
||||
assert_resolution(height=height, width=width, is_two_stage=True)
|
||||
if not (0.0 <= conditioning_attention_strength <= 1.0):
|
||||
raise ValueError(
|
||||
f"conditioning_attention_strength must be in [0.0, 1.0], got {conditioning_attention_strength}"
|
||||
)
|
||||
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
noiser = GaussianNoiser(generator=generator)
|
||||
@@ -158,6 +208,7 @@ class ICLoraPipeline:
|
||||
height=height // 2,
|
||||
fps=frame_rate,
|
||||
)
|
||||
|
||||
stage_1_conditionings = self._create_conditionings(
|
||||
images=images,
|
||||
video_conditioning=video_conditioning,
|
||||
@@ -165,7 +216,10 @@ class ICLoraPipeline:
|
||||
width=stage_1_output_shape.width,
|
||||
video_encoder=video_encoder,
|
||||
num_frames=num_frames,
|
||||
conditioning_attention_strength=conditioning_attention_strength,
|
||||
conditioning_attention_mask=conditioning_attention_mask,
|
||||
)
|
||||
|
||||
video_state, audio_state = denoise_audio_video(
|
||||
output_shape=stage_1_output_shape,
|
||||
conditionings=stage_1_conditionings,
|
||||
@@ -182,6 +236,19 @@ class ICLoraPipeline:
|
||||
del transformer
|
||||
cleanup_memory()
|
||||
|
||||
if skip_stage_2:
|
||||
# Skip Stage 2: Decode directly from Stage 1 output at half resolution
|
||||
logging.info("[IC-LoRA] Skipping Stage 2 (--skip-stage-2 enabled)")
|
||||
decoded_video = vae_decode_video(
|
||||
video_state.latent, self.stage_1_model_ledger.video_decoder(), tiling_config, generator
|
||||
)
|
||||
decoded_audio = vae_decode_audio(
|
||||
audio_state.latent, self.stage_1_model_ledger.audio_decoder(), self.stage_1_model_ledger.vocoder()
|
||||
)
|
||||
del video_encoder
|
||||
cleanup_memory()
|
||||
return decoded_video, decoded_audio
|
||||
|
||||
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
|
||||
upscaled_video_latent = upsample_video(
|
||||
latent=video_state.latent[:1],
|
||||
@@ -250,13 +317,29 @@ class ICLoraPipeline:
|
||||
|
||||
def _create_conditionings(
|
||||
self,
|
||||
images: list[tuple[str, int, float]],
|
||||
images: list[ImageConditioningInput],
|
||||
video_conditioning: list[tuple[str, float]],
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
video_encoder: VideoEncoder,
|
||||
conditioning_attention_strength: float = 1.0,
|
||||
conditioning_attention_mask: torch.Tensor | None = None,
|
||||
) -> list[ConditioningItem]:
|
||||
"""
|
||||
Create conditioning items for video generation.
|
||||
Args:
|
||||
conditioning_attention_strength: Scalar attention weight in [0, 1].
|
||||
If conditioning_attention_mask is also provided, the downsampled mask
|
||||
is multiplied by this strength. Otherwise this scalar is passed
|
||||
directly as the attention mask.
|
||||
conditioning_attention_mask: Optional pixel-space attention mask with shape
|
||||
(B, 1, F_pixel, H_pixel, W_pixel) matching the reference video's
|
||||
pixel dimensions. Downsampled to latent space with causal temporal
|
||||
handling, then multiplied by conditioning_attention_strength.
|
||||
Returns:
|
||||
List of conditioning items. IC-LoRA conditionings are appended last.
|
||||
"""
|
||||
conditionings = image_conditionings_by_replacing_latent(
|
||||
images=images,
|
||||
height=height,
|
||||
@@ -287,21 +370,96 @@ class ICLoraPipeline:
|
||||
device=self.device,
|
||||
)
|
||||
encoded_video = video_encoder(video)
|
||||
conditionings.append(
|
||||
VideoConditionByReferenceLatent(
|
||||
latent=encoded_video,
|
||||
downscale_factor=scale,
|
||||
strength=strength,
|
||||
reference_video_shape = VideoLatentShape.from_torch_shape(encoded_video.shape)
|
||||
|
||||
# Build attention_mask for ConditioningItemAttentionStrengthWrapper
|
||||
if conditioning_attention_mask is not None:
|
||||
# Downsample pixel-space mask to latent space, then scale by strength
|
||||
latent_mask = self._downsample_mask_to_latent(
|
||||
mask=conditioning_attention_mask,
|
||||
target_latent_shape=reference_video_shape,
|
||||
)
|
||||
attn_mask = latent_mask * conditioning_attention_strength
|
||||
elif conditioning_attention_strength < 1.0:
|
||||
# Use scalar strength only
|
||||
attn_mask = conditioning_attention_strength
|
||||
else:
|
||||
attn_mask = None
|
||||
|
||||
cond = VideoConditionByReferenceLatent(
|
||||
latent=encoded_video,
|
||||
downscale_factor=scale,
|
||||
strength=strength,
|
||||
)
|
||||
if attn_mask is not None:
|
||||
cond = ConditioningItemAttentionStrengthWrapper(cond, attention_mask=attn_mask)
|
||||
conditionings.append(cond)
|
||||
|
||||
if video_conditioning:
|
||||
logging.info(f"[IC-LoRA] Added {len(video_conditioning)} video conditioning(s)")
|
||||
|
||||
return conditionings
|
||||
|
||||
@staticmethod
|
||||
def _downsample_mask_to_latent(
|
||||
mask: torch.Tensor,
|
||||
target_latent_shape: VideoLatentShape,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Downsample a pixel-space mask to latent space using VAE scale factors.
|
||||
Handles causal temporal downsampling: the first frame is kept separately
|
||||
(temporal scale factor = 1 for the first frame), while the remaining
|
||||
frames are downsampled by the VAE's temporal scale factor.
|
||||
Args:
|
||||
mask: Pixel-space mask of shape (B, 1, F_pixel, H_pixel, W_pixel).
|
||||
Values in [0, 1].
|
||||
target_latent_shape: Expected latent shape after VAE encoding.
|
||||
Used to determine the target (F_latent, H_latent, W_latent).
|
||||
Returns:
|
||||
Flattened latent-space mask of shape (B, F_lat * H_lat * W_lat),
|
||||
matching the patchifier's token ordering (f, h, w).
|
||||
"""
|
||||
b = mask.shape[0]
|
||||
f_lat = target_latent_shape.frames
|
||||
h_lat = target_latent_shape.height
|
||||
w_lat = target_latent_shape.width
|
||||
|
||||
# Step 1: Spatial downsampling (area interpolation per frame)
|
||||
f_pix = mask.shape[2]
|
||||
spatial_down = torch.nn.functional.interpolate(
|
||||
rearrange(mask, "b 1 f h w -> (b f) 1 h w"),
|
||||
size=(h_lat, w_lat),
|
||||
mode="area",
|
||||
)
|
||||
spatial_down = rearrange(spatial_down, "(b f) 1 h w -> b 1 f h w", b=b)
|
||||
|
||||
# Step 2: Causal temporal downsampling
|
||||
# First frame: kept as-is (causal VAE encodes first frame independently)
|
||||
first_frame = spatial_down[:, :, :1, :, :] # (B, 1, 1, H_lat, W_lat)
|
||||
|
||||
if f_pix > 1 and f_lat > 1:
|
||||
# Remaining frames: downsample by temporal factor via group-mean
|
||||
t = (f_pix - 1) // (f_lat - 1) # temporal downscale factor
|
||||
assert (f_pix - 1) % (f_lat - 1) == 0, (
|
||||
f"Pixel frames ({f_pix}) not compatible with latent frames ({f_lat}): "
|
||||
f"(f_pix - 1) must be divisible by (f_lat - 1)"
|
||||
)
|
||||
rest = rearrange(spatial_down[:, :, 1:, :, :], "b 1 (f t) h w -> b 1 f t h w", t=t)
|
||||
rest = rest.mean(dim=3) # (B, 1, F_lat-1, H_lat, W_lat)
|
||||
latent_mask = torch.cat([first_frame, rest], dim=2) # (B, 1, F_lat, H_lat, W_lat)
|
||||
else:
|
||||
latent_mask = first_frame
|
||||
|
||||
# Flatten to (B, F_lat * H_lat * W_lat) matching patchifier token order (f, h, w)
|
||||
return rearrange(latent_mask, "b 1 f h w -> b (f h w)")
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
parser = default_2_stage_distilled_arg_parser()
|
||||
checkpoint_path = detect_checkpoint_path(distilled=True)
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_2_stage_distilled_arg_parser(params=params)
|
||||
parser.add_argument(
|
||||
"--video-conditioning",
|
||||
action=VideoConditioningAction,
|
||||
@@ -309,9 +467,47 @@ def main() -> None:
|
||||
metavar=("PATH", "STRENGTH"),
|
||||
required=True,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--conditioning-attention-mask",
|
||||
action=VideoMaskConditioningAction,
|
||||
nargs=2,
|
||||
metavar=("MASK_PATH", "STRENGTH"),
|
||||
default=None,
|
||||
help=(
|
||||
"Optional spatial attention mask: path to a grayscale mask video and "
|
||||
"attention strength. The mask video pixel values in [0,1] control "
|
||||
"per-region conditioning attention strength. The strength scalar is "
|
||||
"multiplied with the spatial mask. "
|
||||
"0.0 = ignore IC-LoRA conditioning, 1.0 = full conditioning influence. "
|
||||
"When not provided, full conditioning strength (1.0) is used. "
|
||||
"Example: --conditioning-attention-mask path/to/mask.mp4 0.5"
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip-stage-2",
|
||||
action="store_true",
|
||||
help=(
|
||||
"Skip Stage 2 upsampling and refinement. Output will be at half resolution "
|
||||
"(height//2, width//2). Useful for faster iteration or when GPU memory is limited."
|
||||
),
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
# Load mask video if provided via --conditioning-attention-mask
|
||||
conditioning_attention_mask = None
|
||||
conditioning_attention_strength = 1.0
|
||||
if args.conditioning_attention_mask is not None:
|
||||
mask_path, mask_strength = args.conditioning_attention_mask
|
||||
conditioning_attention_strength = mask_strength
|
||||
conditioning_attention_mask = _load_mask_video(
|
||||
mask_path=mask_path,
|
||||
height=args.height // 2, # Stage 1 operates at half resolution
|
||||
width=args.width // 2,
|
||||
num_frames=args.num_frames,
|
||||
)
|
||||
|
||||
pipeline = ICLoraPipeline(
|
||||
checkpoint_path=args.checkpoint_path,
|
||||
distilled_checkpoint_path=args.distilled_checkpoint_path,
|
||||
spatial_upsampler_path=args.spatial_upsampler_path,
|
||||
gemma_root=args.gemma_root,
|
||||
loras=args.lora,
|
||||
@@ -329,18 +525,53 @@ def main() -> None:
|
||||
images=args.images,
|
||||
video_conditioning=args.video_conditioning,
|
||||
tiling_config=tiling_config,
|
||||
conditioning_attention_strength=conditioning_attention_strength,
|
||||
skip_stage_2=args.skip_stage_2,
|
||||
conditioning_attention_mask=conditioning_attention_mask,
|
||||
)
|
||||
|
||||
encode_video(
|
||||
video=video,
|
||||
fps=args.frame_rate,
|
||||
audio=audio,
|
||||
audio_sample_rate=AUDIO_SAMPLE_RATE,
|
||||
output_path=args.output_path,
|
||||
video_chunks_number=video_chunks_number,
|
||||
)
|
||||
|
||||
|
||||
def _load_mask_video(
|
||||
mask_path: str,
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
) -> torch.Tensor:
|
||||
"""Load a mask video and return a pixel-space tensor of shape (1, 1, F, H, W).
|
||||
The mask video is loaded, resized to (height, width), converted to
|
||||
grayscale, and normalised to [0, 1].
|
||||
Args:
|
||||
mask_path: Path to the mask video file.
|
||||
height: Target height in pixels.
|
||||
width: Target width in pixels.
|
||||
num_frames: Maximum number of frames to load.
|
||||
Returns:
|
||||
Tensor of shape ``(1, 1, F, H, W)`` with values in ``[0, 1]``.
|
||||
"""
|
||||
mask_video = load_video_conditioning(
|
||||
video_path=mask_path,
|
||||
height=height,
|
||||
width=width,
|
||||
frame_cap=num_frames,
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
# mask_video shape: (1, C, F, H, W) — take mean over channels for grayscale
|
||||
mask = mask_video.mean(dim=1, keepdim=True) # (1, 1, F, H, W)
|
||||
# Normalise to [0, 1] — load_video_conditioning applies normalize_latent,
|
||||
# so undo that: values are in [-1, 1], remap to [0, 1]
|
||||
mask = (mask + 1.0) / 2.0
|
||||
return mask.clamp(0.0, 1.0)
|
||||
|
||||
|
||||
def _read_lora_reference_downscale_factor(lora_path: str) -> int:
|
||||
"""Read reference_downscale_factor from LoRA safetensors metadata.
|
||||
Some IC-LoRA models are trained with reference videos at lower resolution than
|
||||
|
||||
@@ -4,7 +4,11 @@ from collections.abc import Iterator
|
||||
import torch
|
||||
|
||||
from ltx_core.components.diffusion_steps import EulerDiffusionStep
|
||||
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
|
||||
from ltx_core.components.guiders import (
|
||||
MultiModalGuiderFactory,
|
||||
MultiModalGuiderParams,
|
||||
create_multimodal_guider_factory,
|
||||
)
|
||||
from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
@@ -15,25 +19,22 @@ from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.model.video_vae import decode_video as vae_decode_video
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.text_encoders.gemma import encode_text
|
||||
from ltx_core.types import LatentState, VideoPixelShape
|
||||
from ltx_core.types import Audio, LatentState, VideoPixelShape
|
||||
from ltx_pipelines.utils import ModelLedger
|
||||
from ltx_pipelines.utils.args import default_2_stage_arg_parser
|
||||
from ltx_pipelines.utils.constants import (
|
||||
AUDIO_SAMPLE_RATE,
|
||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
||||
)
|
||||
from ltx_pipelines.utils.args import ImageConditioningInput, default_2_stage_arg_parser, detect_checkpoint_path
|
||||
from ltx_pipelines.utils.constants import STAGE_2_DISTILLED_SIGMA_VALUES, detect_params
|
||||
from ltx_pipelines.utils.helpers import (
|
||||
assert_resolution,
|
||||
cleanup_memory,
|
||||
denoise_audio_video,
|
||||
euler_denoising_loop,
|
||||
generate_enhanced_prompt,
|
||||
get_device,
|
||||
image_conditionings_by_adding_guiding_latent,
|
||||
multi_modal_guider_denoising_func,
|
||||
multi_modal_guider_factory_denoising_func,
|
||||
simple_denoising_func,
|
||||
)
|
||||
from ltx_pipelines.utils.media_io import encode_video
|
||||
from ltx_pipelines.utils.samplers import euler_denoising_loop
|
||||
from ltx_pipelines.utils.types import PipelineComponents
|
||||
|
||||
device = get_device()
|
||||
@@ -43,8 +44,10 @@ class KeyframeInterpolationPipeline:
|
||||
"""
|
||||
Keyframe-based Two-stage video interpolation pipeline.
|
||||
Interpolates between keyframes to generate a video with smoother transitions.
|
||||
Stage 1 generates video at the target resolution, then Stage 2 upsamples
|
||||
Stage 1 generates video at half of the target resolution, then Stage 2 upsamples
|
||||
by 2x and refines with additional denoising steps for higher quality output.
|
||||
Stage 1 uses full model while Stage 2 uses distilled LORA for efficiency,
|
||||
as the upsampled video already has good quality and just needs refinement.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -76,7 +79,6 @@ class KeyframeInterpolationPipeline:
|
||||
device=device,
|
||||
)
|
||||
|
||||
@torch.inference_mode()
|
||||
def __call__( # noqa: PLR0913
|
||||
self,
|
||||
prompt: str,
|
||||
@@ -87,12 +89,12 @@ class KeyframeInterpolationPipeline:
|
||||
num_frames: int,
|
||||
frame_rate: float,
|
||||
num_inference_steps: int,
|
||||
video_guider_params: MultiModalGuiderParams,
|
||||
audio_guider_params: MultiModalGuiderParams,
|
||||
images: list[tuple[str, int, float]],
|
||||
video_guider_params: MultiModalGuiderParams | MultiModalGuiderFactory,
|
||||
audio_guider_params: MultiModalGuiderParams | MultiModalGuiderFactory,
|
||||
images: list[ImageConditioningInput],
|
||||
tiling_config: TilingConfig | None = None,
|
||||
enhance_prompt: bool = False,
|
||||
) -> tuple[Iterator[torch.Tensor], torch.Tensor]:
|
||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||
assert_resolution(height=height, width=width, is_two_stage=True)
|
||||
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
@@ -126,12 +128,12 @@ class KeyframeInterpolationPipeline:
|
||||
video_state=video_state,
|
||||
audio_state=audio_state,
|
||||
stepper=stepper,
|
||||
denoise_fn=multi_modal_guider_denoising_func(
|
||||
video_guider=MultiModalGuider(
|
||||
denoise_fn=multi_modal_guider_factory_denoising_func(
|
||||
video_guider_factory=create_multimodal_guider_factory(
|
||||
params=video_guider_params,
|
||||
negative_context=v_context_n,
|
||||
),
|
||||
audio_guider=MultiModalGuider(
|
||||
audio_guider_factory=create_multimodal_guider_factory(
|
||||
params=audio_guider_params,
|
||||
negative_context=a_context_n,
|
||||
),
|
||||
@@ -241,7 +243,9 @@ class KeyframeInterpolationPipeline:
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
parser = default_2_stage_arg_parser()
|
||||
checkpoint_path = detect_checkpoint_path()
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_2_stage_arg_parser(params=params)
|
||||
args = parser.parse_args()
|
||||
pipeline = KeyframeInterpolationPipeline(
|
||||
checkpoint_path=args.checkpoint_path,
|
||||
@@ -286,7 +290,6 @@ def main() -> None:
|
||||
video=video,
|
||||
fps=args.frame_rate,
|
||||
audio=audio,
|
||||
audio_sample_rate=AUDIO_SAMPLE_RATE,
|
||||
output_path=args.output_path,
|
||||
video_chunks_number=video_chunks_number,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,476 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.components.diffusion_steps import EulerDiffusionStep
|
||||
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
|
||||
from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.patchifiers import get_pixel_coords
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
from ltx_core.conditioning import ConditioningItem
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.model.audio_vae import decode_audio as vae_decode_audio
|
||||
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
|
||||
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.model.video_vae import decode_video as vae_decode_video
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.text_encoders.gemma import encode_text
|
||||
from ltx_core.tools import LatentTools
|
||||
from ltx_core.types import (
|
||||
Audio,
|
||||
AudioLatentShape,
|
||||
LatentState,
|
||||
SpatioTemporalScaleFactors,
|
||||
VideoPixelShape,
|
||||
)
|
||||
from ltx_pipelines.utils import ModelLedger
|
||||
from ltx_pipelines.utils.constants import DISTILLED_SIGMA_VALUES
|
||||
from ltx_pipelines.utils.helpers import (
|
||||
cleanup_memory,
|
||||
generate_enhanced_prompt,
|
||||
get_device,
|
||||
multi_modal_guider_denoising_func,
|
||||
noise_audio_state,
|
||||
noise_video_state,
|
||||
simple_denoising_func,
|
||||
)
|
||||
from ltx_pipelines.utils.media_io import (
|
||||
decode_audio_from_file,
|
||||
encode_video,
|
||||
get_videostream_metadata,
|
||||
load_video_conditioning,
|
||||
)
|
||||
from ltx_pipelines.utils.samplers import euler_denoising_loop
|
||||
from ltx_pipelines.utils.types import PipelineComponents
|
||||
|
||||
device = get_device()
|
||||
|
||||
|
||||
def _encode_video_for_retake(
|
||||
video_encoder: torch.nn.Module,
|
||||
video_path: str,
|
||||
output_shape: VideoPixelShape,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> torch.Tensor:
|
||||
"""Load video and encode to latents."""
|
||||
pixel_video = load_video_conditioning(
|
||||
video_path=video_path,
|
||||
height=output_shape.height,
|
||||
width=output_shape.width,
|
||||
frame_cap=output_shape.frames,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
) # (1, C, F, H, W)
|
||||
return video_encoder(pixel_video)
|
||||
|
||||
|
||||
def _encode_audio_for_retake(
|
||||
audio_encoder: torch.nn.Module,
|
||||
waveform: torch.Tensor,
|
||||
waveform_sr: int,
|
||||
output_shape: VideoPixelShape,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
"""Encode audio to latents and trim/pad to match output_shape."""
|
||||
waveform_batch = waveform.unsqueeze(0) if waveform.dim() == 2 else waveform
|
||||
initial_audio_latent = vae_encode_audio(
|
||||
Audio(waveform=waveform_batch.to(dtype), sampling_rate=waveform_sr), audio_encoder, None
|
||||
)
|
||||
expected_audio_shape = AudioLatentShape.from_video_pixel_shape(output_shape)
|
||||
expected_frames = expected_audio_shape.frames
|
||||
actual_frames = initial_audio_latent.shape[2]
|
||||
if actual_frames > expected_frames:
|
||||
initial_audio_latent = initial_audio_latent[:, :, :expected_frames, :]
|
||||
elif actual_frames < expected_frames:
|
||||
pad = torch.zeros(
|
||||
initial_audio_latent.shape[0],
|
||||
initial_audio_latent.shape[1],
|
||||
expected_frames - actual_frames,
|
||||
initial_audio_latent.shape[3],
|
||||
device=initial_audio_latent.device,
|
||||
dtype=initial_audio_latent.dtype,
|
||||
)
|
||||
initial_audio_latent = torch.cat([initial_audio_latent, pad], dim=2)
|
||||
return initial_audio_latent
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Custom conditioning item: temporal region mask
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TemporalRegionMask:
|
||||
"""Conditioning item that sets ``denoise_mask = 0`` outside a time range
|
||||
and ``1`` inside, so only the specified temporal region is regenerated.
|
||||
Uses ``start_time`` and ``end_time`` in seconds. Works in *patchified*
|
||||
(token) space using the patchifier's ``get_patch_grid_bounds``: for video
|
||||
coords are latent frame indices (converted from seconds via ``fps``), for
|
||||
audio coords are already in seconds.
|
||||
"""
|
||||
|
||||
start_time: float # seconds, inclusive
|
||||
end_time: float # seconds, exclusive
|
||||
fps: float
|
||||
|
||||
def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState:
|
||||
coords = latent_tools.patchifier.get_patch_grid_bounds(
|
||||
latent_tools.target_shape, device=latent_state.denoise_mask.device
|
||||
)
|
||||
# coords: [B, 3, N, 2] (video) or [B, 1, N, 2] (audio); temporal dim is index 0
|
||||
if coords.shape[1] == 1:
|
||||
# Audio: patchifier returns seconds
|
||||
t_start = coords[:, 0, :, 0] # [B, N]
|
||||
t_end = coords[:, 0, :, 1] # [B, N]
|
||||
in_region = (t_end > self.start_time) & (t_start < self.end_time)
|
||||
else:
|
||||
# Video: get pixel bounds per patch, find patches for start/end frame, read latent from coords.
|
||||
scale_factors = getattr(latent_tools, "scale_factors", SpatioTemporalScaleFactors.default())
|
||||
pixel_bounds = get_pixel_coords(coords, scale_factors, causal_fix=getattr(latent_tools, "causal_fix", True))
|
||||
timestamp_bounds = pixel_bounds[0, 0] / self.fps
|
||||
t_start, t_end = timestamp_bounds.unbind(dim=-1)
|
||||
in_region = (t_end > self.start_time) & (t_start < self.end_time)
|
||||
state = latent_state.clone()
|
||||
mask_val = in_region.to(state.denoise_mask.dtype)
|
||||
if state.denoise_mask.dim() == 3:
|
||||
mask_val = mask_val.unsqueeze(-1)
|
||||
state.denoise_mask.copy_(mask_val)
|
||||
return state
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pipeline
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class RetakePipeline:
|
||||
"""Regenerate a time region (retake) of an existing video.
|
||||
Given a source video file and a time window ``[start_time, end_time]``
|
||||
(in seconds), this pipeline keeps the video/audio outside that window
|
||||
unchanged and *regenerates* the content inside the window from a text
|
||||
prompt using the LTX-2 diffusion model.
|
||||
Parameters
|
||||
----------
|
||||
checkpoint_path : str
|
||||
Path to the LTX-2 model checkpoint.
|
||||
gemma_root : str
|
||||
Root directory containing Gemma text-encoder weights.
|
||||
loras : list[LoraPathStrengthAndSDOps]
|
||||
Optional LoRA configs applied to the transformer.
|
||||
device : torch.device
|
||||
Target device (default: CUDA if available).
|
||||
quantization : QuantizationPolicy | None
|
||||
Optional quantization policy for the transformer.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
checkpoint_path: str,
|
||||
gemma_root: str,
|
||||
loras: list[LoraPathStrengthAndSDOps],
|
||||
device: torch.device = device,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
):
|
||||
self.device = device
|
||||
self.dtype = torch.bfloat16
|
||||
self.model_ledger = ModelLedger(
|
||||
dtype=self.dtype,
|
||||
device=device,
|
||||
checkpoint_path=checkpoint_path,
|
||||
gemma_root_path=gemma_root,
|
||||
loras=loras,
|
||||
quantization=quantization,
|
||||
)
|
||||
self.pipeline_components = PipelineComponents(
|
||||
dtype=self.dtype,
|
||||
device=device,
|
||||
)
|
||||
|
||||
# --------------------------------------------------------------------- #
|
||||
# Public entry point #
|
||||
# --------------------------------------------------------------------- #
|
||||
|
||||
@torch.inference_mode()
|
||||
def __call__( # noqa: PLR0913, PLR0915
|
||||
self,
|
||||
video_path: str,
|
||||
prompt: str,
|
||||
start_time: float,
|
||||
end_time: float,
|
||||
seed: int,
|
||||
*,
|
||||
negative_prompt: str = "",
|
||||
num_inference_steps: int = 40,
|
||||
video_guider_params: MultiModalGuiderParams | None = None,
|
||||
audio_guider_params: MultiModalGuiderParams | None = None,
|
||||
regenerate_video: bool = True,
|
||||
regenerate_audio: bool = True,
|
||||
enhance_prompt: bool = False,
|
||||
distilled: bool = False,
|
||||
) -> tuple[Iterator[torch.Tensor], torch.Tensor]:
|
||||
"""Regenerate ``[start_time, end_time]`` of the source video (retake).
|
||||
Parameters
|
||||
----------
|
||||
video_path : str
|
||||
Path to the source video file (must contain video; audio is optional).
|
||||
prompt : str
|
||||
Text prompt describing the *regenerated* section.
|
||||
start_time, end_time : float
|
||||
Time window (in seconds) of the section to regenerate.
|
||||
seed : int
|
||||
Random seed for reproducibility.
|
||||
negative_prompt : str
|
||||
Negative prompt for CFG guidance (ignored in distilled mode).
|
||||
num_inference_steps : int
|
||||
Number of Euler denoising steps (ignored in distilled mode which
|
||||
uses a fixed 8-step schedule).
|
||||
video_guider_params, audio_guider_params : MultiModalGuiderParams | None
|
||||
Guidance parameters for video and audio modalities. Ignored in
|
||||
distilled mode.
|
||||
regenerate_video : bool
|
||||
If ``True`` (default), preserve video outside ``[start_time, end_time]``
|
||||
and only regenerate the masked region. If ``False``, fully regenerate
|
||||
all video frames (the encoded video is still used as the initial latent
|
||||
but with ``denoise_mask = 1`` everywhere).
|
||||
regenerate_audio : bool
|
||||
If True, regenerate audio in the [start_time, end_time] window; if False,
|
||||
audio is preserved as-is (no regeneration).
|
||||
enhance_prompt : bool
|
||||
Whether to enhance the prompt via the text encoder.
|
||||
distilled : bool
|
||||
If ``True``, use the distilled sigma schedule
|
||||
(``DISTILLED_SIGMA_VALUES``) and a simple (non-guided) denoising
|
||||
function. The model checkpoint must be the distilled variant.
|
||||
Returns
|
||||
-------
|
||||
tuple[Iterator[torch.Tensor], torch.Tensor]
|
||||
``(video_frames_iterator, audio_waveform)``
|
||||
"""
|
||||
if start_time >= end_time:
|
||||
raise ValueError(f"start_time ({start_time}) must be less than end_time ({end_time})")
|
||||
|
||||
effective_seed = torch.randint(0, 2**31, (1,), device=self.device).item() if seed < 0 else seed
|
||||
generator = torch.Generator(device=self.device).manual_seed(effective_seed)
|
||||
noiser = GaussianNoiser(generator=generator)
|
||||
stepper = EulerDiffusionStep()
|
||||
dtype = self.dtype
|
||||
|
||||
video_encoder = self.model_ledger.video_encoder()
|
||||
|
||||
# Use av to get metadata
|
||||
fps, num_pixel_frames, src_width, src_height = get_videostream_metadata(video_path)
|
||||
|
||||
output_shape = VideoPixelShape(
|
||||
batch=1,
|
||||
frames=num_pixel_frames,
|
||||
width=src_width,
|
||||
height=src_height,
|
||||
fps=fps,
|
||||
)
|
||||
initial_video_latent = _encode_video_for_retake(
|
||||
video_encoder=video_encoder,
|
||||
video_path=video_path,
|
||||
output_shape=output_shape,
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
)
|
||||
video_conditionings: list[ConditioningItem] = [
|
||||
TemporalRegionMask(
|
||||
start_time=start_time if regenerate_video else 0.0,
|
||||
end_time=end_time if regenerate_video else 0.0,
|
||||
fps=fps,
|
||||
)
|
||||
]
|
||||
del video_encoder
|
||||
cleanup_memory()
|
||||
|
||||
initial_audio_latent: torch.Tensor | None = None
|
||||
audio_conditionings: list[ConditioningItem] = []
|
||||
|
||||
audio_in = decode_audio_from_file(video_path, self.device)
|
||||
audio_encoder = self.model_ledger.audio_encoder()
|
||||
|
||||
if audio_in is not None:
|
||||
waveform = audio_in.waveform.squeeze(0)
|
||||
waveform_sr = audio_in.sampling_rate
|
||||
else:
|
||||
waveform, waveform_sr = None, None
|
||||
if waveform is not None:
|
||||
initial_audio_latent = _encode_audio_for_retake(
|
||||
audio_encoder=audio_encoder,
|
||||
waveform=waveform,
|
||||
waveform_sr=waveform_sr,
|
||||
output_shape=output_shape,
|
||||
dtype=dtype,
|
||||
)
|
||||
audio_conditionings = [
|
||||
TemporalRegionMask(
|
||||
start_time=start_time if regenerate_audio else 0.0,
|
||||
end_time=end_time if regenerate_audio else 0.0,
|
||||
fps=fps,
|
||||
)
|
||||
]
|
||||
|
||||
del audio_encoder
|
||||
cleanup_memory()
|
||||
|
||||
text_encoder = self.model_ledger.text_encoder()
|
||||
if enhance_prompt:
|
||||
prompt = generate_enhanced_prompt(text_encoder, prompt, None, seed=effective_seed)
|
||||
|
||||
if distilled:
|
||||
# Distilled mode: single prompt, no negative
|
||||
context_p = encode_text(text_encoder, prompts=[prompt])[0]
|
||||
v_context_p, a_context_p = context_p
|
||||
else:
|
||||
context_p, context_n = encode_text(text_encoder, prompts=[prompt, negative_prompt])
|
||||
v_context_p, a_context_p = context_p
|
||||
v_context_n, a_context_n = context_n
|
||||
|
||||
torch.cuda.synchronize()
|
||||
del text_encoder
|
||||
cleanup_memory()
|
||||
|
||||
transformer = self.model_ledger.transformer()
|
||||
|
||||
sigmas = (
|
||||
torch.tensor(DISTILLED_SIGMA_VALUES) if distilled else LTX2Scheduler().execute(steps=num_inference_steps)
|
||||
).to(dtype=torch.float32, device=self.device)
|
||||
if distilled:
|
||||
denoise_fn = simple_denoising_func(
|
||||
video_context=v_context_p,
|
||||
audio_context=a_context_p,
|
||||
transformer=transformer,
|
||||
)
|
||||
else:
|
||||
video_guider = MultiModalGuider(
|
||||
params=video_guider_params,
|
||||
negative_context=v_context_n,
|
||||
)
|
||||
audio_guider = MultiModalGuider(
|
||||
params=audio_guider_params,
|
||||
negative_context=a_context_n,
|
||||
)
|
||||
denoise_fn = multi_modal_guider_denoising_func(
|
||||
video_guider,
|
||||
audio_guider,
|
||||
v_context=v_context_p,
|
||||
a_context=a_context_p,
|
||||
transformer=transformer,
|
||||
)
|
||||
|
||||
def denoising_loop(
|
||||
sigmas: torch.Tensor,
|
||||
video_state: LatentState,
|
||||
audio_state: LatentState,
|
||||
stepper: DiffusionStepProtocol,
|
||||
) -> tuple[LatentState, LatentState]:
|
||||
return euler_denoising_loop(
|
||||
sigmas=sigmas,
|
||||
video_state=video_state,
|
||||
audio_state=audio_state,
|
||||
stepper=stepper,
|
||||
denoise_fn=denoise_fn,
|
||||
)
|
||||
|
||||
# Build noised states with the encoded latents as initial values and
|
||||
# the temporal masks applied via conditionings.
|
||||
video_state, video_tools = noise_video_state(
|
||||
output_shape=output_shape,
|
||||
noiser=noiser,
|
||||
conditionings=video_conditionings,
|
||||
components=self.pipeline_components,
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
initial_latent=initial_video_latent,
|
||||
)
|
||||
audio_state, audio_tools = noise_audio_state(
|
||||
output_shape=output_shape,
|
||||
noiser=noiser,
|
||||
conditionings=audio_conditionings,
|
||||
components=self.pipeline_components,
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
initial_latent=initial_audio_latent,
|
||||
)
|
||||
|
||||
video_state, audio_state = denoising_loop(sigmas, video_state, audio_state, stepper)
|
||||
|
||||
video_state = video_tools.clear_conditioning(video_state)
|
||||
video_state = video_tools.unpatchify(video_state)
|
||||
audio_state = audio_tools.clear_conditioning(audio_state)
|
||||
audio_state = audio_tools.unpatchify(audio_state)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
del transformer
|
||||
cleanup_memory()
|
||||
|
||||
decoded_video = vae_decode_video(video_state.latent, self.model_ledger.video_decoder(), generator=generator)
|
||||
decoded_audio = vae_decode_audio(
|
||||
audio_state.latent, self.model_ledger.audio_decoder(), self.model_ledger.vocoder()
|
||||
)
|
||||
|
||||
return decoded_video, decoded_audio
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""CLI entry point for retake (regenerate a time region)."""
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
parser = argparse.ArgumentParser(description="Retake: regenerate a time region of a video with LTX-2.")
|
||||
parser.add_argument("--video-path", type=str, required=True, help="Path to the source video.")
|
||||
parser.add_argument("--prompt", type=str, required=True, help="Text prompt for the regenerated region.")
|
||||
parser.add_argument("--start-time", type=float, required=True, help="Start time of the region to regenerate (s).")
|
||||
parser.add_argument("--end-time", type=float, required=True, help="End time of the region to regenerate (s).")
|
||||
parser.add_argument("--output-path", type=str, required=True, help="Path for the output video.")
|
||||
parser.add_argument("--checkpoint-path", type=str, required=True, help="Path to the LTX-2 checkpoint.")
|
||||
parser.add_argument("--gemma-root", type=str, required=True, help="Path to Gemma text encoder weights.")
|
||||
parser.add_argument("--seed", type=int, default=42, help="Random seed. Use -1 for a random seed.")
|
||||
parser.add_argument("--loras", nargs="*", default=[], help="LoRA paths (optional).")
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.start_time >= args.end_time:
|
||||
raise ValueError("start_time must be less than end_time")
|
||||
|
||||
# Validate frame count (8k+1) and resolution (multiples of 32) at CLI stage
|
||||
video_scale = SpatioTemporalScaleFactors.default()
|
||||
fps, num_frames, width, height = get_videostream_metadata(args.video_path)
|
||||
if (num_frames - 1) % video_scale.time != 0:
|
||||
snapped = ((num_frames - 1) // video_scale.time) * video_scale.time + 1
|
||||
raise ValueError(
|
||||
f"Video frame count must satisfy 8k+1 (e.g. 97, 193). Got {num_frames}; use a video with {snapped} frames."
|
||||
)
|
||||
if width % 32 != 0 or height % 32 != 0:
|
||||
raise ValueError(f"Video width and height must be multiples of 32. Got {width}x{height}.")
|
||||
|
||||
pipeline = RetakePipeline(
|
||||
checkpoint_path=args.checkpoint_path,
|
||||
gemma_root=args.gemma_root,
|
||||
loras=args.loras or [],
|
||||
)
|
||||
video_iter, audio = pipeline(
|
||||
video_path=args.video_path,
|
||||
prompt=args.prompt,
|
||||
start_time=args.start_time,
|
||||
end_time=args.end_time,
|
||||
seed=args.seed,
|
||||
)
|
||||
tiling_config = TilingConfig.default()
|
||||
video_chunks_number = get_video_chunks_number(num_frames, tiling_config)
|
||||
encode_video(
|
||||
video=video_iter,
|
||||
fps=int(fps),
|
||||
audio=audio,
|
||||
output_path=args.output_path,
|
||||
video_chunks_number=video_chunks_number,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -4,7 +4,11 @@ from collections.abc import Iterator
|
||||
import torch
|
||||
|
||||
from ltx_core.components.diffusion_steps import EulerDiffusionStep
|
||||
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
|
||||
from ltx_core.components.guiders import (
|
||||
MultiModalGuiderFactory,
|
||||
MultiModalGuiderParams,
|
||||
create_multimodal_guider_factory,
|
||||
)
|
||||
from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
@@ -13,11 +17,9 @@ from ltx_core.model.audio_vae import decode_audio as vae_decode_audio
|
||||
from ltx_core.model.video_vae import decode_video as vae_decode_video
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.text_encoders.gemma import encode_text
|
||||
from ltx_core.types import LatentState, VideoPixelShape
|
||||
from ltx_pipelines.utils import ModelLedger
|
||||
from ltx_pipelines.utils.args import default_1_stage_arg_parser
|
||||
from ltx_pipelines.utils.constants import AUDIO_SAMPLE_RATE
|
||||
from ltx_pipelines.utils.helpers import (
|
||||
from ltx_core.types import Audio, LatentState, VideoPixelShape
|
||||
from ltx_pipelines.utils import (
|
||||
ModelLedger,
|
||||
assert_resolution,
|
||||
cleanup_memory,
|
||||
denoise_audio_video,
|
||||
@@ -25,8 +27,10 @@ from ltx_pipelines.utils.helpers import (
|
||||
generate_enhanced_prompt,
|
||||
get_device,
|
||||
image_conditionings_by_replacing_latent,
|
||||
multi_modal_guider_denoising_func,
|
||||
multi_modal_guider_factory_denoising_func,
|
||||
)
|
||||
from ltx_pipelines.utils.args import ImageConditioningInput, default_1_stage_arg_parser, detect_checkpoint_path
|
||||
from ltx_pipelines.utils.constants import detect_params
|
||||
from ltx_pipelines.utils.media_io import encode_video
|
||||
from ltx_pipelines.utils.types import PipelineComponents
|
||||
|
||||
@@ -39,6 +43,7 @@ class TI2VidOneStagePipeline:
|
||||
Generates video at the target resolution in a single diffusion pass with
|
||||
classifier-free guidance (CFG). Supports optional image conditioning via
|
||||
the images parameter.
|
||||
Assumes full non distilled model is provided in the checkpoint_path.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -74,11 +79,11 @@ class TI2VidOneStagePipeline:
|
||||
num_frames: int,
|
||||
frame_rate: float,
|
||||
num_inference_steps: int,
|
||||
video_guider_params: MultiModalGuiderParams,
|
||||
audio_guider_params: MultiModalGuiderParams,
|
||||
images: list[tuple[str, int, float]],
|
||||
video_guider_params: MultiModalGuiderParams | MultiModalGuiderFactory,
|
||||
audio_guider_params: MultiModalGuiderParams | MultiModalGuiderFactory,
|
||||
images: list[ImageConditioningInput],
|
||||
enhance_prompt: bool = False,
|
||||
) -> tuple[Iterator[torch.Tensor], torch.Tensor]:
|
||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||
assert_resolution(height=height, width=width, is_two_stage=False)
|
||||
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
@@ -104,6 +109,15 @@ class TI2VidOneStagePipeline:
|
||||
transformer = self.model_ledger.transformer()
|
||||
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
|
||||
|
||||
video_guider_factory = create_multimodal_guider_factory(
|
||||
params=video_guider_params,
|
||||
negative_context=v_context_n,
|
||||
)
|
||||
audio_guider_factory = create_multimodal_guider_factory(
|
||||
params=audio_guider_params,
|
||||
negative_context=a_context_n,
|
||||
)
|
||||
|
||||
def first_stage_denoising_loop(
|
||||
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
|
||||
) -> tuple[LatentState, LatentState]:
|
||||
@@ -112,15 +126,9 @@ class TI2VidOneStagePipeline:
|
||||
video_state=video_state,
|
||||
audio_state=audio_state,
|
||||
stepper=stepper,
|
||||
denoise_fn=multi_modal_guider_denoising_func(
|
||||
video_guider=MultiModalGuider(
|
||||
params=video_guider_params,
|
||||
negative_context=v_context_n,
|
||||
),
|
||||
audio_guider=MultiModalGuider(
|
||||
params=audio_guider_params,
|
||||
negative_context=a_context_n,
|
||||
),
|
||||
denoise_fn=multi_modal_guider_factory_denoising_func(
|
||||
video_guider_factory=video_guider_factory,
|
||||
audio_guider_factory=audio_guider_factory,
|
||||
v_context=v_context_p,
|
||||
a_context=a_context_p,
|
||||
transformer=transformer, # noqa: F821
|
||||
@@ -157,14 +165,15 @@ class TI2VidOneStagePipeline:
|
||||
decoded_audio = vae_decode_audio(
|
||||
audio_state.latent, self.model_ledger.audio_decoder(), self.model_ledger.vocoder()
|
||||
)
|
||||
|
||||
return decoded_video, decoded_audio
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
parser = default_1_stage_arg_parser()
|
||||
checkpoint_path = detect_checkpoint_path()
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_1_stage_arg_parser(params=params)
|
||||
args = parser.parse_args()
|
||||
pipeline = TI2VidOneStagePipeline(
|
||||
checkpoint_path=args.checkpoint_path,
|
||||
@@ -204,7 +213,6 @@ def main() -> None:
|
||||
video=video,
|
||||
fps=args.frame_rate,
|
||||
audio=audio,
|
||||
audio_sample_rate=AUDIO_SAMPLE_RATE,
|
||||
output_path=args.output_path,
|
||||
video_chunks_number=1,
|
||||
)
|
||||
|
||||
@@ -4,7 +4,11 @@ from collections.abc import Iterator
|
||||
import torch
|
||||
|
||||
from ltx_core.components.diffusion_steps import EulerDiffusionStep
|
||||
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
|
||||
from ltx_core.components.guiders import (
|
||||
MultiModalGuiderFactory,
|
||||
MultiModalGuiderParams,
|
||||
create_multimodal_guider_factory,
|
||||
)
|
||||
from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
@@ -15,14 +19,9 @@ from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.model.video_vae import decode_video as vae_decode_video
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.text_encoders.gemma import encode_text
|
||||
from ltx_core.types import LatentState, VideoPixelShape
|
||||
from ltx_pipelines.utils import ModelLedger
|
||||
from ltx_pipelines.utils.args import default_2_stage_arg_parser
|
||||
from ltx_pipelines.utils.constants import (
|
||||
AUDIO_SAMPLE_RATE,
|
||||
STAGE_2_DISTILLED_SIGMA_VALUES,
|
||||
)
|
||||
from ltx_pipelines.utils.helpers import (
|
||||
from ltx_core.types import Audio, LatentState, VideoPixelShape
|
||||
from ltx_pipelines.utils import (
|
||||
ModelLedger,
|
||||
assert_resolution,
|
||||
cleanup_memory,
|
||||
denoise_audio_video,
|
||||
@@ -30,9 +29,11 @@ from ltx_pipelines.utils.helpers import (
|
||||
generate_enhanced_prompt,
|
||||
get_device,
|
||||
image_conditionings_by_replacing_latent,
|
||||
multi_modal_guider_denoising_func,
|
||||
multi_modal_guider_factory_denoising_func,
|
||||
simple_denoising_func,
|
||||
)
|
||||
from ltx_pipelines.utils.args import ImageConditioningInput, default_2_stage_arg_parser, detect_checkpoint_path
|
||||
from ltx_pipelines.utils.constants import STAGE_2_DISTILLED_SIGMA_VALUES, detect_params
|
||||
from ltx_pipelines.utils.media_io import encode_video
|
||||
from ltx_pipelines.utils.types import PipelineComponents
|
||||
|
||||
@@ -42,9 +43,10 @@ device = get_device()
|
||||
class TI2VidTwoStagesPipeline:
|
||||
"""
|
||||
Two-stage text/image-to-video generation pipeline.
|
||||
Stage 1 generates video at the target resolution with CFG guidance, then
|
||||
Stage 2 upsamples by 2x and refines using a distilled LoRA for higher
|
||||
quality output. Supports optional image conditioning via the images parameter.
|
||||
Stage 1 generates video at half of the target resolution with CFG guidance (assuming
|
||||
full model is used), then Stage 2 upsamples by 2x and refines using a distilled
|
||||
LoRA for higher quality output. Supports optional image conditioning via the
|
||||
images parameter.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -54,7 +56,7 @@ class TI2VidTwoStagesPipeline:
|
||||
spatial_upsampler_path: str,
|
||||
gemma_root: str,
|
||||
loras: list[LoraPathStrengthAndSDOps],
|
||||
device: str = device,
|
||||
device: torch.device = device,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
):
|
||||
self.device = device
|
||||
@@ -78,7 +80,6 @@ class TI2VidTwoStagesPipeline:
|
||||
device=device,
|
||||
)
|
||||
|
||||
@torch.inference_mode()
|
||||
def __call__( # noqa: PLR0913
|
||||
self,
|
||||
prompt: str,
|
||||
@@ -89,12 +90,12 @@ class TI2VidTwoStagesPipeline:
|
||||
num_frames: int,
|
||||
frame_rate: float,
|
||||
num_inference_steps: int,
|
||||
video_guider_params: MultiModalGuiderParams,
|
||||
audio_guider_params: MultiModalGuiderParams,
|
||||
images: list[tuple[str, int, float]],
|
||||
video_guider_params: MultiModalGuiderParams | MultiModalGuiderFactory,
|
||||
audio_guider_params: MultiModalGuiderParams | MultiModalGuiderFactory,
|
||||
images: list[ImageConditioningInput],
|
||||
tiling_config: TilingConfig | None = None,
|
||||
enhance_prompt: bool = False,
|
||||
) -> tuple[Iterator[torch.Tensor], torch.Tensor]:
|
||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||
assert_resolution(height=height, width=width, is_two_stage=True)
|
||||
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
@@ -128,12 +129,12 @@ class TI2VidTwoStagesPipeline:
|
||||
video_state=video_state,
|
||||
audio_state=audio_state,
|
||||
stepper=stepper,
|
||||
denoise_fn=multi_modal_guider_denoising_func(
|
||||
video_guider=MultiModalGuider(
|
||||
denoise_fn=multi_modal_guider_factory_denoising_func(
|
||||
video_guider_factory=create_multimodal_guider_factory(
|
||||
params=video_guider_params,
|
||||
negative_context=v_context_n,
|
||||
),
|
||||
audio_guider=MultiModalGuider(
|
||||
audio_guider_factory=create_multimodal_guider_factory(
|
||||
params=audio_guider_params,
|
||||
negative_context=a_context_n,
|
||||
),
|
||||
@@ -237,14 +238,15 @@ class TI2VidTwoStagesPipeline:
|
||||
decoded_audio = vae_decode_audio(
|
||||
audio_state.latent, self.stage_2_model_ledger.audio_decoder(), self.stage_2_model_ledger.vocoder()
|
||||
)
|
||||
|
||||
return decoded_video, decoded_audio
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
parser = default_2_stage_arg_parser()
|
||||
checkpoint_path = detect_checkpoint_path()
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_2_stage_arg_parser(params=params)
|
||||
args = parser.parse_args()
|
||||
pipeline = TI2VidTwoStagesPipeline(
|
||||
checkpoint_path=args.checkpoint_path,
|
||||
@@ -289,7 +291,6 @@ def main() -> None:
|
||||
video=video,
|
||||
fps=args.frame_rate,
|
||||
audio=audio,
|
||||
audio_sample_rate=AUDIO_SAMPLE_RATE,
|
||||
output_path=args.output_path,
|
||||
video_chunks_number=video_chunks_number,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,307 @@
|
||||
import logging
|
||||
from collections.abc import Iterator
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.components.diffusion_steps import Res2sDiffusionStep
|
||||
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderParams
|
||||
from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.model.audio_vae import decode_audio as vae_decode_audio
|
||||
from ltx_core.model.upsampler import upsample_video
|
||||
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.model.video_vae import decode_video as vae_decode_video
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.text_encoders.gemma import encode_text
|
||||
from ltx_core.tools import VideoLatentShape
|
||||
from ltx_core.types import Audio, LatentState, VideoPixelShape
|
||||
from ltx_pipelines.utils import (
|
||||
ModelLedger,
|
||||
assert_resolution,
|
||||
cleanup_memory,
|
||||
denoise_audio_video,
|
||||
generate_enhanced_prompt,
|
||||
get_device,
|
||||
image_conditionings_by_replacing_latent,
|
||||
multi_modal_guider_denoising_func,
|
||||
res2s_audio_video_denoising_loop,
|
||||
simple_denoising_func,
|
||||
)
|
||||
from ltx_pipelines.utils.args import ImageConditioningInput, default_2_stage_arg_parser, detect_checkpoint_path
|
||||
from ltx_pipelines.utils.constants import STAGE_2_DISTILLED_SIGMA_VALUES, detect_params
|
||||
from ltx_pipelines.utils.media_io import encode_video
|
||||
from ltx_pipelines.utils.types import PipelineComponents
|
||||
|
||||
device = get_device()
|
||||
|
||||
|
||||
class TI2VidTwoStagesRes2sPipeline:
|
||||
"""
|
||||
Two-stage text/image-to-video generation pipeline using the res_2s sampler.
|
||||
Same structure as :class:`TI2VidTwoStagesPipeline`: stage 1 generates video at
|
||||
half of the target resolution with CFG guidance (assuming full model is used),
|
||||
then Stage 2 upsamples by 2x and refines using a distilled LoRA for higher
|
||||
quality output.
|
||||
Uses the res_2s second-order sampler instead of Euler, allowing fewer
|
||||
steps for comparable quality. Supports optional image conditioning via
|
||||
the images parameter.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
checkpoint_path: str,
|
||||
distilled_lora: list[LoraPathStrengthAndSDOps],
|
||||
spatial_upsampler_path: str,
|
||||
gemma_root: str,
|
||||
loras: list[LoraPathStrengthAndSDOps],
|
||||
device: str = device,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
):
|
||||
self.device = device
|
||||
self.dtype = torch.bfloat16
|
||||
self.stage_1_model_ledger = ModelLedger(
|
||||
dtype=self.dtype,
|
||||
device=device,
|
||||
checkpoint_path=checkpoint_path,
|
||||
gemma_root_path=gemma_root,
|
||||
spatial_upsampler_path=spatial_upsampler_path,
|
||||
loras=loras,
|
||||
quantization=quantization,
|
||||
)
|
||||
|
||||
self.stage_2_model_ledger = self.stage_1_model_ledger.with_loras(
|
||||
loras=distilled_lora,
|
||||
)
|
||||
|
||||
self.pipeline_components = PipelineComponents(
|
||||
dtype=self.dtype,
|
||||
device=device,
|
||||
)
|
||||
|
||||
@torch.inference_mode()
|
||||
def __call__( # noqa: PLR0913
|
||||
self,
|
||||
prompt: str,
|
||||
negative_prompt: str,
|
||||
seed: int,
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
frame_rate: float,
|
||||
num_inference_steps: int,
|
||||
video_guider_params: MultiModalGuiderParams,
|
||||
audio_guider_params: MultiModalGuiderParams,
|
||||
images: list[ImageConditioningInput],
|
||||
tiling_config: TilingConfig | None = None,
|
||||
enhance_prompt: bool = False,
|
||||
) -> tuple[Iterator[torch.Tensor], Audio]:
|
||||
assert_resolution(height=height, width=width, is_two_stage=True)
|
||||
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
noiser = GaussianNoiser(generator=generator)
|
||||
dtype = torch.bfloat16
|
||||
|
||||
text_encoder = self.stage_1_model_ledger.text_encoder()
|
||||
if enhance_prompt:
|
||||
prompt = generate_enhanced_prompt(
|
||||
text_encoder, prompt, images[0][0] if len(images) > 0 else None, seed=seed
|
||||
)
|
||||
context_p, context_n = encode_text(text_encoder, prompts=[prompt, negative_prompt])
|
||||
v_context_p, a_context_p = context_p
|
||||
v_context_n, a_context_n = context_n
|
||||
|
||||
torch.cuda.synchronize()
|
||||
del text_encoder
|
||||
cleanup_memory()
|
||||
|
||||
# Stage 1: Initial low resolution video generation.
|
||||
video_encoder = self.stage_1_model_ledger.video_encoder()
|
||||
transformer = self.stage_1_model_ledger.transformer()
|
||||
|
||||
stage_1_output_shape = VideoPixelShape(
|
||||
batch=1,
|
||||
frames=num_frames,
|
||||
width=width // 2,
|
||||
height=height // 2,
|
||||
fps=frame_rate,
|
||||
)
|
||||
empty_latent = torch.empty(VideoLatentShape.from_pixel_shape(stage_1_output_shape).to_torch_shape())
|
||||
stepper = Res2sDiffusionStep()
|
||||
sigmas = (
|
||||
LTX2Scheduler()
|
||||
.execute(latent=empty_latent, steps=num_inference_steps)
|
||||
.to(dtype=torch.float32, device=self.device)
|
||||
)
|
||||
|
||||
def first_stage_denoising_loop(
|
||||
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
|
||||
) -> tuple[LatentState, LatentState]:
|
||||
return res2s_audio_video_denoising_loop(
|
||||
sigmas=sigmas,
|
||||
video_state=video_state,
|
||||
audio_state=audio_state,
|
||||
stepper=stepper,
|
||||
denoise_fn=multi_modal_guider_denoising_func(
|
||||
video_guider=MultiModalGuider(
|
||||
params=video_guider_params,
|
||||
negative_context=v_context_n,
|
||||
),
|
||||
audio_guider=MultiModalGuider(
|
||||
params=audio_guider_params,
|
||||
negative_context=a_context_n,
|
||||
),
|
||||
v_context=v_context_p,
|
||||
a_context=a_context_p,
|
||||
transformer=transformer, # noqa: F821
|
||||
),
|
||||
)
|
||||
|
||||
stage_1_conditionings = image_conditionings_by_replacing_latent(
|
||||
images=images,
|
||||
height=stage_1_output_shape.height,
|
||||
width=stage_1_output_shape.width,
|
||||
video_encoder=video_encoder,
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
)
|
||||
video_state, audio_state = denoise_audio_video(
|
||||
output_shape=stage_1_output_shape,
|
||||
conditionings=stage_1_conditionings,
|
||||
noiser=noiser,
|
||||
sigmas=sigmas,
|
||||
stepper=stepper,
|
||||
denoising_loop_fn=first_stage_denoising_loop,
|
||||
components=self.pipeline_components,
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
del transformer
|
||||
cleanup_memory()
|
||||
|
||||
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
|
||||
upscaled_video_latent = upsample_video(
|
||||
latent=video_state.latent[:1],
|
||||
video_encoder=video_encoder,
|
||||
upsampler=self.stage_2_model_ledger.spatial_upsampler(),
|
||||
)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
cleanup_memory()
|
||||
|
||||
transformer = self.stage_2_model_ledger.transformer()
|
||||
distilled_sigmas = torch.tensor(STAGE_2_DISTILLED_SIGMA_VALUES, device=self.device)
|
||||
|
||||
def second_stage_denoising_loop(
|
||||
sigmas: torch.Tensor, video_state: LatentState, audio_state: LatentState, stepper: DiffusionStepProtocol
|
||||
) -> tuple[LatentState, LatentState]:
|
||||
return res2s_audio_video_denoising_loop(
|
||||
sigmas=sigmas,
|
||||
video_state=video_state,
|
||||
audio_state=audio_state,
|
||||
stepper=stepper,
|
||||
denoise_fn=simple_denoising_func(
|
||||
video_context=v_context_p,
|
||||
audio_context=a_context_p,
|
||||
transformer=transformer, # noqa: F821
|
||||
),
|
||||
)
|
||||
|
||||
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
|
||||
stage_2_conditionings = image_conditionings_by_replacing_latent(
|
||||
images=images,
|
||||
height=stage_2_output_shape.height,
|
||||
width=stage_2_output_shape.width,
|
||||
video_encoder=video_encoder,
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
)
|
||||
video_state, audio_state = denoise_audio_video(
|
||||
output_shape=stage_2_output_shape,
|
||||
conditionings=stage_2_conditionings,
|
||||
noiser=noiser,
|
||||
sigmas=distilled_sigmas,
|
||||
stepper=stepper,
|
||||
denoising_loop_fn=second_stage_denoising_loop,
|
||||
components=self.pipeline_components,
|
||||
dtype=dtype,
|
||||
device=self.device,
|
||||
noise_scale=distilled_sigmas[0],
|
||||
initial_video_latent=upscaled_video_latent,
|
||||
initial_audio_latent=audio_state.latent,
|
||||
)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
del transformer
|
||||
del video_encoder
|
||||
cleanup_memory()
|
||||
|
||||
decoded_video = vae_decode_video(
|
||||
video_state.latent, self.stage_2_model_ledger.video_decoder(), tiling_config, generator
|
||||
)
|
||||
decoded_audio = vae_decode_audio(
|
||||
audio_state.latent, self.stage_2_model_ledger.audio_decoder(), self.stage_2_model_ledger.vocoder()
|
||||
)
|
||||
return decoded_video, decoded_audio
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
checkpoint_path = detect_checkpoint_path()
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_2_stage_arg_parser(params=params)
|
||||
args = parser.parse_args()
|
||||
pipeline = TI2VidTwoStagesRes2sPipeline(
|
||||
checkpoint_path=args.checkpoint_path,
|
||||
distilled_lora=args.distilled_lora,
|
||||
spatial_upsampler_path=args.spatial_upsampler_path,
|
||||
gemma_root=args.gemma_root,
|
||||
loras=args.lora,
|
||||
quantization=args.quantization,
|
||||
)
|
||||
tiling_config = TilingConfig.default()
|
||||
video_chunks_number = get_video_chunks_number(args.num_frames, tiling_config)
|
||||
video, audio = pipeline(
|
||||
prompt=args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
seed=args.seed,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
frame_rate=args.frame_rate,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
video_guider_params=MultiModalGuiderParams(
|
||||
cfg_scale=args.video_cfg_guidance_scale,
|
||||
stg_scale=args.video_stg_guidance_scale,
|
||||
rescale_scale=args.video_rescale_scale,
|
||||
modality_scale=args.a2v_guidance_scale,
|
||||
skip_step=args.video_skip_step,
|
||||
stg_blocks=args.video_stg_blocks,
|
||||
),
|
||||
audio_guider_params=MultiModalGuiderParams(
|
||||
cfg_scale=args.audio_cfg_guidance_scale,
|
||||
stg_scale=args.audio_stg_guidance_scale,
|
||||
rescale_scale=args.audio_rescale_scale,
|
||||
modality_scale=args.v2a_guidance_scale,
|
||||
skip_step=args.audio_skip_step,
|
||||
stg_blocks=args.audio_stg_blocks,
|
||||
),
|
||||
images=args.images,
|
||||
tiling_config=tiling_config,
|
||||
)
|
||||
|
||||
encode_video(
|
||||
video=video,
|
||||
fps=args.frame_rate,
|
||||
audio=audio,
|
||||
output_path=args.output_path,
|
||||
video_chunks_number=video_chunks_number,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,5 +1,33 @@
|
||||
from ltx_pipelines.utils.helpers import (
|
||||
assert_resolution,
|
||||
cleanup_memory,
|
||||
denoise_audio_video,
|
||||
generate_enhanced_prompt,
|
||||
get_device,
|
||||
image_conditionings_by_replacing_latent,
|
||||
multi_modal_guider_denoising_func,
|
||||
multi_modal_guider_factory_denoising_func,
|
||||
simple_denoising_func,
|
||||
)
|
||||
from ltx_pipelines.utils.model_ledger import ModelLedger
|
||||
from ltx_pipelines.utils.samplers import (
|
||||
euler_denoising_loop,
|
||||
gradient_estimating_euler_denoising_loop,
|
||||
res2s_audio_video_denoising_loop,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ModelLedger",
|
||||
"assert_resolution",
|
||||
"cleanup_memory",
|
||||
"denoise_audio_video",
|
||||
"euler_denoising_loop",
|
||||
"generate_enhanced_prompt",
|
||||
"get_device",
|
||||
"gradient_estimating_euler_denoising_loop",
|
||||
"image_conditionings_by_replacing_latent",
|
||||
"multi_modal_guider_denoising_func",
|
||||
"multi_modal_guider_factory_denoising_func",
|
||||
"res2s_audio_video_denoising_loop",
|
||||
"simple_denoising_func",
|
||||
]
|
||||
|
||||
@@ -1,24 +1,25 @@
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
from typing import NamedTuple
|
||||
|
||||
from ltx_core.loader import LTXV_LORA_COMFY_RENAMING_MAP, LoraPathStrengthAndSDOps
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_pipelines.utils.constants import (
|
||||
DEFAULT_1_STAGE_HEIGHT,
|
||||
DEFAULT_1_STAGE_WIDTH,
|
||||
DEFAULT_2_STAGE_HEIGHT,
|
||||
DEFAULT_2_STAGE_WIDTH,
|
||||
DEFAULT_AUDIO_GUIDER_PARAMS,
|
||||
DEFAULT_FRAME_RATE,
|
||||
DEFAULT_IMAGE_CRF,
|
||||
DEFAULT_LORA_STRENGTH,
|
||||
DEFAULT_NEGATIVE_PROMPT,
|
||||
DEFAULT_NUM_FRAMES,
|
||||
DEFAULT_NUM_INFERENCE_STEPS,
|
||||
DEFAULT_SEED,
|
||||
DEFAULT_VIDEO_GUIDER_PARAMS,
|
||||
LTX_2_3_PARAMS,
|
||||
PipelineParams,
|
||||
)
|
||||
|
||||
|
||||
class ImageConditioningInput(NamedTuple):
|
||||
path: str
|
||||
frame_idx: int
|
||||
strength: float
|
||||
crf: int = DEFAULT_IMAGE_CRF
|
||||
|
||||
|
||||
class VideoConditioningAction(argparse.Action):
|
||||
def __call__(
|
||||
self,
|
||||
@@ -35,20 +36,50 @@ class VideoConditioningAction(argparse.Action):
|
||||
setattr(namespace, self.dest, current)
|
||||
|
||||
|
||||
class VideoMaskConditioningAction(argparse.Action):
|
||||
"""Parse ``--conditioning-attention-mask PATH STRENGTH``.
|
||||
Stores a ``(mask_path, strength)`` tuple on the namespace. The mask video
|
||||
should be grayscale with pixel values in [0, 1] controlling per-region
|
||||
conditioning attention strength. The scalar *STRENGTH* is multiplied with
|
||||
the spatial mask before it is applied.
|
||||
"""
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
parser: argparse.ArgumentParser, # noqa: ARG002
|
||||
namespace: argparse.Namespace,
|
||||
values: list[str],
|
||||
option_string: str | None = None,
|
||||
) -> None:
|
||||
if len(values) != 2:
|
||||
msg = f"{option_string} requires exactly 2 arguments (MASK_PATH STRENGTH), got {len(values)}"
|
||||
raise argparse.ArgumentError(self, msg)
|
||||
|
||||
mask_path = resolve_path(values[0])
|
||||
strength = float(values[1])
|
||||
setattr(namespace, self.dest, (mask_path, strength))
|
||||
|
||||
|
||||
class ImageAction(argparse.Action):
|
||||
def __call__(
|
||||
self,
|
||||
parser: argparse.ArgumentParser, # noqa: ARG002
|
||||
namespace: argparse.Namespace,
|
||||
values: list[str],
|
||||
option_string: str | None = None, # noqa: ARG002
|
||||
option_string: str | None = None,
|
||||
) -> None:
|
||||
path, frame_idx, strength_str = values
|
||||
resolved_path = resolve_path(path)
|
||||
frame_idx = int(frame_idx)
|
||||
strength = float(strength_str)
|
||||
if len(values) not in (3, 4):
|
||||
msg = f"{option_string} requires 3 or 4 arguments (PATH FRAME_IDX STRENGTH [CRF]), got {len(values)}"
|
||||
raise argparse.ArgumentError(self, msg)
|
||||
|
||||
conditioning = ImageConditioningInput(
|
||||
path=resolve_path(values[0]),
|
||||
frame_idx=int(values[1]),
|
||||
strength=float(values[2]),
|
||||
crf=int(values[3]) if len(values) > 3 else DEFAULT_IMAGE_CRF,
|
||||
)
|
||||
current = getattr(namespace, self.dest) or []
|
||||
current.append((resolved_path, frame_idx, strength))
|
||||
current.append(conditioning)
|
||||
setattr(namespace, self.dest, current)
|
||||
|
||||
|
||||
@@ -113,14 +144,34 @@ class QuantizationAction(argparse.Action):
|
||||
setattr(namespace, self.dest, policy)
|
||||
|
||||
|
||||
def basic_arg_parser() -> argparse.ArgumentParser:
|
||||
def detect_checkpoint_path(distilled: bool = False) -> str:
|
||||
"""Pre-parse argv to extract the checkpoint path before building the full parser."""
|
||||
pre = argparse.ArgumentParser(add_help=False)
|
||||
flag = "--distilled-checkpoint-path" if distilled else "--checkpoint-path"
|
||||
pre.add_argument(flag, type=resolve_path, required=True)
|
||||
known, _ = pre.parse_known_args()
|
||||
return known.distilled_checkpoint_path if distilled else known.checkpoint_path
|
||||
|
||||
|
||||
def basic_arg_parser(
|
||||
params: PipelineParams = LTX_2_3_PARAMS,
|
||||
distilled: bool = False,
|
||||
) -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--checkpoint-path",
|
||||
type=resolve_path,
|
||||
required=True,
|
||||
help="Path to LTX-2 model checkpoint (.safetensors file).",
|
||||
)
|
||||
if distilled:
|
||||
parser.add_argument(
|
||||
"--distilled-checkpoint-path",
|
||||
type=resolve_path,
|
||||
required=True,
|
||||
help="Path to LTX-2 distilled model checkpoint (.safetensors file).",
|
||||
)
|
||||
else:
|
||||
parser.add_argument(
|
||||
"--checkpoint-path",
|
||||
type=resolve_path,
|
||||
required=True,
|
||||
help="Path to LTX-2 model checkpoint (.safetensors file).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--gemma-root",
|
||||
type=resolve_path,
|
||||
@@ -142,58 +193,57 @@ def basic_arg_parser() -> argparse.ArgumentParser:
|
||||
parser.add_argument(
|
||||
"--seed",
|
||||
type=int,
|
||||
default=DEFAULT_SEED,
|
||||
help=(
|
||||
f"Random seed value used to initialize the noise tensor for "
|
||||
f"reproducible generation (default: {DEFAULT_SEED})."
|
||||
),
|
||||
default=params.seed,
|
||||
help=f"Random seed for reproducible generation (default: {params.seed}).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--height",
|
||||
type=int,
|
||||
default=DEFAULT_1_STAGE_HEIGHT,
|
||||
help=f"Height of the generated video in pixels, should be divisible by 32 (default: {DEFAULT_1_STAGE_HEIGHT}).",
|
||||
default=params.stage_1_height,
|
||||
help=f"Video height in pixels, divisible by 32 (default: {params.stage_1_height}).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--width",
|
||||
type=int,
|
||||
default=DEFAULT_1_STAGE_WIDTH,
|
||||
help=f"Width of the generated video in pixels, should be divisible by 32 (default: {DEFAULT_1_STAGE_WIDTH}).",
|
||||
default=params.stage_1_width,
|
||||
help=f"Width of the generated video in pixels, should be divisible by 32 (default: {params.stage_1_width}).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-frames",
|
||||
type=int,
|
||||
default=DEFAULT_NUM_FRAMES,
|
||||
default=params.num_frames,
|
||||
help=f"Number of frames to generate in the output video sequence, num-frames = (8 x K) + 1, "
|
||||
f"where k is a non-negative integer (default: {DEFAULT_NUM_FRAMES}).",
|
||||
f"where k is a non-negative integer (default: {params.num_frames}).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--frame-rate",
|
||||
type=float,
|
||||
default=DEFAULT_FRAME_RATE,
|
||||
help=f"Frame rate of the generated video (fps) (default: {DEFAULT_FRAME_RATE}).",
|
||||
default=params.frame_rate,
|
||||
help=f"Frame rate of the generated video (fps) (default: {params.frame_rate}).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-inference-steps",
|
||||
type=int,
|
||||
default=DEFAULT_NUM_INFERENCE_STEPS,
|
||||
default=params.num_inference_steps,
|
||||
help=(
|
||||
f"Number of denoising steps in the diffusion sampling process. "
|
||||
f"Higher values improve quality but increase generation time (default: {DEFAULT_NUM_INFERENCE_STEPS})."
|
||||
f"Higher values improve quality but increase generation time (default: {params.num_inference_steps})."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image",
|
||||
dest="images",
|
||||
action=ImageAction,
|
||||
nargs=3,
|
||||
metavar=("PATH", "FRAME_IDX", "STRENGTH"),
|
||||
nargs="+",
|
||||
metavar="ARG",
|
||||
default=[],
|
||||
help=(
|
||||
"Image conditioning input: path to image file, target frame index, "
|
||||
"and conditioning strength (all three required). Default: empty list [] (no image conditioning). "
|
||||
"Image conditioning input: PATH FRAME_IDX STRENGTH [CRF]. "
|
||||
"PATH is the image file, FRAME_IDX is the target frame index, "
|
||||
"STRENGTH is the conditioning strength (all three required). "
|
||||
f"CRF is the optional H.264 compression quality (0=lossless, default: {DEFAULT_IMAGE_CRF}). "
|
||||
"Can be specified multiple times. Example: --image path/to/image1.jpg 0 0.8 "
|
||||
"--image path/to/image2.jpg 160 0.9"
|
||||
"--image path/to/image2.jpg 160 0.9 0"
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -228,8 +278,10 @@ def basic_arg_parser() -> argparse.ArgumentParser:
|
||||
return parser
|
||||
|
||||
|
||||
def default_1_stage_arg_parser() -> argparse.ArgumentParser:
|
||||
parser = basic_arg_parser()
|
||||
def default_1_stage_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argparse.ArgumentParser:
|
||||
video_guider = params.video_guider_params
|
||||
audio_guider = params.audio_guider_params
|
||||
parser = basic_arg_parser(params=params)
|
||||
parser.add_argument(
|
||||
"--negative-prompt",
|
||||
type=str,
|
||||
@@ -243,139 +295,139 @@ def default_1_stage_arg_parser() -> argparse.ArgumentParser:
|
||||
parser.add_argument(
|
||||
"--video-cfg-guidance-scale",
|
||||
type=float,
|
||||
default=DEFAULT_VIDEO_GUIDER_PARAMS.cfg_scale,
|
||||
default=video_guider.cfg_scale,
|
||||
help=(
|
||||
f"Classifier-free guidance (CFG) scale controlling how strongly "
|
||||
f"the model adheres to the video prompt. Higher values increase prompt "
|
||||
"adherence but may reduce diversity. 1.0 means no effect "
|
||||
f"(default: {DEFAULT_VIDEO_GUIDER_PARAMS.cfg_scale})."
|
||||
f"adherence but may reduce diversity. 1.0 means no effect "
|
||||
f"(default: {video_guider.cfg_scale})."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video-stg-guidance-scale",
|
||||
type=float,
|
||||
default=DEFAULT_VIDEO_GUIDER_PARAMS.stg_scale,
|
||||
default=video_guider.stg_scale,
|
||||
help=(
|
||||
f"STG (Spatio-Temporal Guidance) scale controlling how strongly "
|
||||
f"the model reacts to the perturbation of the video modality. Higher values increase "
|
||||
f"the effect but may reduce quality. 0.0 means no effect "
|
||||
f"(default: {DEFAULT_VIDEO_GUIDER_PARAMS.stg_scale})."
|
||||
f"(default: {video_guider.stg_scale})."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video-rescale-scale",
|
||||
type=float,
|
||||
default=DEFAULT_VIDEO_GUIDER_PARAMS.rescale_scale,
|
||||
default=video_guider.rescale_scale,
|
||||
help=(
|
||||
f"Rescale scale controlling how strongly "
|
||||
f"the model rescales the video modality after applying other guidance. Higher values tend to decrease "
|
||||
f"oversaturation effects. 0.0 means no effect (default: {DEFAULT_VIDEO_GUIDER_PARAMS.rescale_scale})."
|
||||
f"oversaturation effects. 0.0 means no effect (default: {video_guider.rescale_scale})."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video-stg-blocks",
|
||||
type=int,
|
||||
nargs="*",
|
||||
default=DEFAULT_VIDEO_GUIDER_PARAMS.stg_blocks,
|
||||
help=(f"Which transformer blocks to perturb for STG. Default: {DEFAULT_VIDEO_GUIDER_PARAMS.stg_blocks}."),
|
||||
default=video_guider.stg_blocks,
|
||||
help=(f"Which transformer blocks to perturb for STG. Default: {video_guider.stg_blocks}."),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--a2v-guidance-scale",
|
||||
type=float,
|
||||
default=DEFAULT_VIDEO_GUIDER_PARAMS.modality_scale,
|
||||
default=video_guider.modality_scale,
|
||||
help=(
|
||||
f"A2V (Audio-to-Video) guidance scale controlling how strongly "
|
||||
f"the model reacts to the perturbation of the audio-to-video cross-attention. Higher values may increase "
|
||||
f"lipsync quality. 1.0 means no effect (default: {DEFAULT_VIDEO_GUIDER_PARAMS.modality_scale})."
|
||||
f"lipsync quality. 1.0 means no effect (default: {video_guider.modality_scale})."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video-skip-step",
|
||||
type=int,
|
||||
default=DEFAULT_VIDEO_GUIDER_PARAMS.skip_step,
|
||||
default=video_guider.skip_step,
|
||||
help=(
|
||||
"Video skip step N controls periodic skipping during the video diffusion process: "
|
||||
"only steps where step_index % (N + 1) == 0 are processed, all others are skipped "
|
||||
f"(e.g., 0 = no skipping; 1 = skip every other step; 2 = skip 2 of every 3 steps; "
|
||||
f"default: {DEFAULT_VIDEO_GUIDER_PARAMS.skip_step})."
|
||||
f"default: {video_guider.skip_step})."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-cfg-guidance-scale",
|
||||
type=float,
|
||||
default=DEFAULT_AUDIO_GUIDER_PARAMS.cfg_scale,
|
||||
default=audio_guider.cfg_scale,
|
||||
help=(
|
||||
f"Audio CFG (Classifier-free guidance) scale controlling how strongly "
|
||||
f"the model adheres to the audio prompt. Higher values increase prompt "
|
||||
f"adherence but may reduce diversity. 1.0 means no effect "
|
||||
f"(default: {DEFAULT_AUDIO_GUIDER_PARAMS.cfg_scale})."
|
||||
f"(default: {audio_guider.cfg_scale})."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-stg-guidance-scale",
|
||||
type=float,
|
||||
default=DEFAULT_AUDIO_GUIDER_PARAMS.stg_scale,
|
||||
default=audio_guider.stg_scale,
|
||||
help=(
|
||||
f"Audio STG (Spatio-Temporal Guidance) scale controlling how strongly "
|
||||
f"the model reacts to the perturbation of the audio modality. Higher values increase "
|
||||
f"the effect but may reduce quality. 0.0 means no effect "
|
||||
f"(default: {DEFAULT_AUDIO_GUIDER_PARAMS.stg_scale})."
|
||||
f"(default: {audio_guider.stg_scale})."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-rescale-scale",
|
||||
type=float,
|
||||
default=DEFAULT_AUDIO_GUIDER_PARAMS.rescale_scale,
|
||||
default=audio_guider.rescale_scale,
|
||||
help=(
|
||||
f"Audio rescale scale controlling how strongly "
|
||||
f"the model rescales the audio modality after applying other guidance. "
|
||||
f"Experimental. 0.0 means no effect (default: {DEFAULT_AUDIO_GUIDER_PARAMS.rescale_scale})."
|
||||
f"Experimental. 0.0 means no effect (default: {audio_guider.rescale_scale})."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-stg-blocks",
|
||||
type=int,
|
||||
nargs="*",
|
||||
default=DEFAULT_AUDIO_GUIDER_PARAMS.stg_blocks,
|
||||
help=(f"Which transformer blocks to perturb for Audio STG. Default: {DEFAULT_AUDIO_GUIDER_PARAMS.stg_blocks}."),
|
||||
default=audio_guider.stg_blocks,
|
||||
help=(f"Which transformer blocks to perturb for Audio STG. Default: {audio_guider.stg_blocks}."),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--v2a-guidance-scale",
|
||||
type=float,
|
||||
default=DEFAULT_AUDIO_GUIDER_PARAMS.modality_scale,
|
||||
default=audio_guider.modality_scale,
|
||||
help=(
|
||||
f"V2A (Video-to-Audio) guidance scale controlling how strongly "
|
||||
f"the model reacts to the perturbation of the video-to-audio cross-attention. Higher values may increase "
|
||||
f"lipsync quality. 1.0 means no effect (default: {DEFAULT_AUDIO_GUIDER_PARAMS.modality_scale})."
|
||||
f"lipsync quality. 1.0 means no effect (default: {audio_guider.modality_scale})."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-skip-step",
|
||||
type=int,
|
||||
default=DEFAULT_AUDIO_GUIDER_PARAMS.skip_step,
|
||||
default=audio_guider.skip_step,
|
||||
help=(
|
||||
"Audio skip step N controls periodic skipping during the audio diffusion process: "
|
||||
"only steps where step_index % (N + 1) == 0 are processed, all others are skipped "
|
||||
f"(e.g., 0 = no skipping; 1 = skip every other step; 2 = skip 2 of every 3 steps; "
|
||||
f"default: {DEFAULT_AUDIO_GUIDER_PARAMS.skip_step})."
|
||||
f"default: {audio_guider.skip_step})."
|
||||
),
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
def default_2_stage_arg_parser() -> argparse.ArgumentParser:
|
||||
parser = default_1_stage_arg_parser()
|
||||
parser.set_defaults(height=DEFAULT_2_STAGE_HEIGHT, width=DEFAULT_2_STAGE_WIDTH)
|
||||
def default_2_stage_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argparse.ArgumentParser:
|
||||
parser = default_1_stage_arg_parser(params=params)
|
||||
parser.set_defaults(height=params.stage_2_height, width=params.stage_2_width)
|
||||
# Update help text to reflect 2-stage defaults
|
||||
for action in parser._actions:
|
||||
if "--height" in action.option_strings:
|
||||
action.help = (
|
||||
f"Height of the generated video in pixels, should be divisible by 64 "
|
||||
f"(default: {DEFAULT_2_STAGE_HEIGHT})."
|
||||
f"(default: {params.stage_2_height})."
|
||||
)
|
||||
if "--width" in action.option_strings:
|
||||
action.help = (
|
||||
f"Width of the generated video in pixels, should be divisible by 64 (default: {DEFAULT_2_STAGE_WIDTH})."
|
||||
f"Width of the generated video in pixels, should be divisible by 64 (default: {params.stage_2_width})."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--distilled-lora",
|
||||
@@ -405,19 +457,19 @@ def default_2_stage_arg_parser() -> argparse.ArgumentParser:
|
||||
return parser
|
||||
|
||||
|
||||
def default_2_stage_distilled_arg_parser() -> argparse.ArgumentParser:
|
||||
parser = basic_arg_parser()
|
||||
parser.set_defaults(height=DEFAULT_2_STAGE_HEIGHT, width=DEFAULT_2_STAGE_WIDTH)
|
||||
def default_2_stage_distilled_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argparse.ArgumentParser:
|
||||
parser = basic_arg_parser(params=params, distilled=True)
|
||||
parser.set_defaults(height=params.stage_2_height, width=params.stage_2_width)
|
||||
# Update help text to reflect 2-stage defaults
|
||||
for action in parser._actions:
|
||||
if "--height" in action.option_strings:
|
||||
action.help = (
|
||||
f"Height of the generated video in pixels, should be divisible by 64 "
|
||||
f"(default: {DEFAULT_2_STAGE_HEIGHT})."
|
||||
f"(default: {params.stage_2_height})."
|
||||
)
|
||||
if "--width" in action.option_strings:
|
||||
action.help = (
|
||||
f"Width of the generated video in pixels, should be divisible by 64 (default: {DEFAULT_2_STAGE_WIDTH})."
|
||||
f"Width of the generated video in pixels, should be divisible by 64 (default: {params.stage_2_width})."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--spatial-upsampler-path",
|
||||
|
||||
@@ -1,12 +1,17 @@
|
||||
import logging
|
||||
from dataclasses import dataclass, field, replace
|
||||
|
||||
from safetensors import safe_open
|
||||
|
||||
from ltx_core.components.guiders import MultiModalGuiderParams
|
||||
from ltx_core.types import SpatioTemporalScaleFactors
|
||||
|
||||
# =============================================================================
|
||||
# Diffusion Schedule
|
||||
# =============================================================================
|
||||
|
||||
# Noise schedule for the distilled pipeline. These sigma values control noise
|
||||
# levels at each denoising step and were tuned to match the distillation process.
|
||||
from ltx_core.components.guiders import MultiModalGuiderParams
|
||||
from ltx_core.types import SpatioTemporalScaleFactors
|
||||
|
||||
DISTILLED_SIGMA_VALUES = [1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0]
|
||||
|
||||
# Reduced schedule for super-resolution stage 2 (subset of distilled values)
|
||||
@@ -14,64 +19,88 @@ STAGE_2_DISTILLED_SIGMA_VALUES = [0.909375, 0.725, 0.421875, 0.0]
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Video Generation Defaults
|
||||
# Pipeline Parameters
|
||||
# =============================================================================
|
||||
|
||||
DEFAULT_SEED = 10
|
||||
DEFAULT_1_STAGE_HEIGHT = 512
|
||||
DEFAULT_1_STAGE_WIDTH = 768
|
||||
DEFAULT_2_STAGE_HEIGHT = DEFAULT_1_STAGE_HEIGHT * 2
|
||||
DEFAULT_2_STAGE_WIDTH = DEFAULT_1_STAGE_WIDTH * 2
|
||||
DEFAULT_NUM_FRAMES = 121
|
||||
DEFAULT_FRAME_RATE = 24.0
|
||||
DEFAULT_NUM_INFERENCE_STEPS = 40
|
||||
DEFAULT_VIDEO_GUIDER_PARAMS = MultiModalGuiderParams(
|
||||
cfg_scale=3.0,
|
||||
stg_scale=1.0,
|
||||
rescale_scale=0.7,
|
||||
modality_scale=3.0,
|
||||
skip_step=0,
|
||||
stg_blocks=[29],
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PipelineParams:
|
||||
seed: int = 10
|
||||
stage_1_height: int = 512
|
||||
stage_1_width: int = 768
|
||||
num_frames: int = 121
|
||||
frame_rate: float = 24.0
|
||||
num_inference_steps: int = 40
|
||||
video_guider_params: MultiModalGuiderParams = field(
|
||||
default_factory=lambda: MultiModalGuiderParams(
|
||||
cfg_scale=3.0,
|
||||
stg_scale=1.0,
|
||||
rescale_scale=0.7,
|
||||
modality_scale=3.0,
|
||||
skip_step=0,
|
||||
stg_blocks=[29],
|
||||
)
|
||||
)
|
||||
audio_guider_params: MultiModalGuiderParams = field(
|
||||
default_factory=lambda: MultiModalGuiderParams(
|
||||
cfg_scale=7.0,
|
||||
stg_scale=1.0,
|
||||
rescale_scale=0.7,
|
||||
modality_scale=3.0,
|
||||
skip_step=0,
|
||||
stg_blocks=[29],
|
||||
)
|
||||
)
|
||||
|
||||
@property
|
||||
def stage_2_height(self) -> int:
|
||||
return int(self.stage_1_height * 2)
|
||||
|
||||
@property
|
||||
def stage_2_width(self) -> int:
|
||||
return int(self.stage_1_width * 2)
|
||||
|
||||
|
||||
# Default params for LTX-2.0 non-distilled models. These can be overridden by detecting from checkpoint metadata.
|
||||
LTX_2_PARAMS = PipelineParams()
|
||||
|
||||
# Default params for LTX-2.3 non-distilled models. These override some of the LTX-2.0 defaults.
|
||||
LTX_2_3_PARAMS = replace(
|
||||
LTX_2_PARAMS,
|
||||
num_inference_steps=30,
|
||||
video_guider_params=replace(LTX_2_PARAMS.video_guider_params, stg_blocks=[28]),
|
||||
audio_guider_params=replace(LTX_2_PARAMS.audio_guider_params, stg_blocks=[28]),
|
||||
)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Audio
|
||||
# =============================================================================
|
||||
|
||||
DEFAULT_AUDIO_GUIDER_PARAMS = MultiModalGuiderParams(
|
||||
cfg_scale=7.0,
|
||||
stg_scale=1.0,
|
||||
rescale_scale=0.7,
|
||||
modality_scale=3.0,
|
||||
skip_step=0,
|
||||
stg_blocks=[29],
|
||||
)
|
||||
AUDIO_SAMPLE_RATE = 24000
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# LoRA
|
||||
# =============================================================================
|
||||
|
||||
DEFAULT_LORA_STRENGTH = 1.0
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Video VAE Architecture
|
||||
# =============================================================================
|
||||
|
||||
DEFAULT_IMAGE_CRF = 33
|
||||
VIDEO_SCALE_FACTORS = SpatioTemporalScaleFactors.default()
|
||||
VIDEO_LATENT_CHANNELS = 128
|
||||
|
||||
_LTX_2_3_MODEL_VERSION_PREFIX = "2.3"
|
||||
|
||||
# =============================================================================
|
||||
# Image Preprocessing
|
||||
# =============================================================================
|
||||
|
||||
# CRF (Constant Rate Factor) for H.264 encoding used in image conditioning.
|
||||
# Lower = higher quality, 0 = lossless. This mimics compression artifacts.
|
||||
DEFAULT_IMAGE_CRF = 33
|
||||
def detect_params(checkpoint_path: str) -> PipelineParams:
|
||||
"""Detect pipeline params from checkpoint metadata.
|
||||
Reads the ``model_version`` field from the safetensors config metadata.
|
||||
Returns ``LTX_2_3_PARAMS`` when the version starts with "2.3",
|
||||
otherwise falls back to ``LTX_2_PARAMS``.
|
||||
"""
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
try:
|
||||
with safe_open(checkpoint_path, framework="pt") as f:
|
||||
metadata = f.metadata() or {}
|
||||
version = metadata.get("model_version", "")
|
||||
except Exception:
|
||||
logger.warning("Could not read checkpoint metadata from %s, using LTX-2 defaults", checkpoint_path)
|
||||
return LTX_2_PARAMS
|
||||
|
||||
if version.startswith(_LTX_2_3_MODEL_VERSION_PREFIX):
|
||||
return LTX_2_3_PARAMS
|
||||
|
||||
logger.info("Using LTX_2_PARAMS for checkpoint (version=%s)", version or "unknown")
|
||||
return LTX_2_PARAMS
|
||||
|
||||
|
||||
# =============================================================================
|
||||
|
||||
@@ -3,9 +3,8 @@ import logging
|
||||
from dataclasses import replace
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from ltx_core.components.guiders import MultiModalGuider
|
||||
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderFactory
|
||||
from ltx_core.components.noisers import Noiser
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol, GuiderProtocol
|
||||
from ltx_core.conditioning import (
|
||||
@@ -21,10 +20,10 @@ from ltx_core.guidance.perturbations import (
|
||||
)
|
||||
from ltx_core.model.transformer import Modality, X0Model
|
||||
from ltx_core.model.video_vae import VideoEncoder
|
||||
from ltx_core.text_encoders.gemma import GemmaTextEncoderModelBase
|
||||
from ltx_core.text_encoders.gemma import GemmaTextEncoder
|
||||
from ltx_core.tools import AudioLatentTools, LatentTools, VideoLatentTools
|
||||
from ltx_core.types import AudioLatentShape, LatentState, VideoLatentShape, VideoPixelShape
|
||||
from ltx_core.utils import to_denoised, to_velocity
|
||||
from ltx_pipelines.utils.args import ImageConditioningInput
|
||||
from ltx_pipelines.utils.media_io import decode_image, load_image_conditioning, resize_aspect_ratio_preserving
|
||||
from ltx_pipelines.utils.types import (
|
||||
DenoisingFunc,
|
||||
@@ -46,7 +45,7 @@ def cleanup_memory() -> None:
|
||||
|
||||
|
||||
def image_conditionings_by_replacing_latent(
|
||||
images: list[tuple[str, int, float]],
|
||||
images: list[ImageConditioningInput],
|
||||
height: int,
|
||||
width: int,
|
||||
video_encoder: VideoEncoder,
|
||||
@@ -54,20 +53,21 @@ def image_conditionings_by_replacing_latent(
|
||||
device: torch.device,
|
||||
) -> list[ConditioningItem]:
|
||||
conditionings = []
|
||||
for image_path, frame_idx, strength in images:
|
||||
for img in images:
|
||||
image = load_image_conditioning(
|
||||
image_path=image_path,
|
||||
image_path=img.path,
|
||||
height=height,
|
||||
width=width,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
crf=img.crf,
|
||||
)
|
||||
encoded_image = video_encoder(image)
|
||||
conditionings.append(
|
||||
VideoConditionByLatentIndex(
|
||||
latent=encoded_image,
|
||||
strength=strength,
|
||||
latent_idx=frame_idx,
|
||||
strength=img.strength,
|
||||
latent_idx=img.frame_idx,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -75,7 +75,7 @@ def image_conditionings_by_replacing_latent(
|
||||
|
||||
|
||||
def image_conditionings_by_adding_guiding_latent(
|
||||
images: list[tuple[str, int, float]],
|
||||
images: list[ImageConditioningInput],
|
||||
height: int,
|
||||
width: int,
|
||||
video_encoder: VideoEncoder,
|
||||
@@ -83,131 +83,22 @@ def image_conditionings_by_adding_guiding_latent(
|
||||
device: torch.device,
|
||||
) -> list[ConditioningItem]:
|
||||
conditionings = []
|
||||
for image_path, frame_idx, strength in images:
|
||||
for img in images:
|
||||
image = load_image_conditioning(
|
||||
image_path=image_path,
|
||||
image_path=img.path,
|
||||
height=height,
|
||||
width=width,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
crf=img.crf,
|
||||
)
|
||||
encoded_image = video_encoder(image)
|
||||
conditionings.append(
|
||||
VideoConditionByKeyframeIndex(keyframes=encoded_image, frame_idx=frame_idx, strength=strength)
|
||||
VideoConditionByKeyframeIndex(keyframes=encoded_image, frame_idx=img.frame_idx, strength=img.strength)
|
||||
)
|
||||
return conditionings
|
||||
|
||||
|
||||
def euler_denoising_loop(
|
||||
sigmas: torch.Tensor,
|
||||
video_state: LatentState,
|
||||
audio_state: LatentState,
|
||||
stepper: DiffusionStepProtocol,
|
||||
denoise_fn: DenoisingFunc,
|
||||
) -> tuple[LatentState, LatentState]:
|
||||
"""
|
||||
Perform the joint audio-video denoising loop over a diffusion schedule.
|
||||
This function iterates over all but the final value in ``sigmas`` and, at
|
||||
each diffusion step, calls ``denoise_fn`` to obtain denoised video and
|
||||
audio latents. The denoised latents are post-processed with their
|
||||
respective denoise masks and clean latents, then passed to ``stepper`` to
|
||||
advance the noisy latents one step along the diffusion schedule.
|
||||
### Parameters
|
||||
sigmas:
|
||||
A 1D tensor of noise levels (diffusion sigmas) defining the sampling
|
||||
schedule. All steps except the last element are iterated over.
|
||||
video_state:
|
||||
The current video :class:`LatentState`, containing the noisy latent,
|
||||
its clean reference latent, and the denoising mask.
|
||||
audio_state:
|
||||
The current audio :class:`LatentState`, analogous to ``video_state``
|
||||
but for the audio modality.
|
||||
stepper:
|
||||
An implementation of :class:`DiffusionStepProtocol` that updates a
|
||||
latent given the current latent, its denoised estimate, the full
|
||||
``sigmas`` schedule, and the current step index.
|
||||
denoise_fn:
|
||||
A callable implementing :class:`DenoisingFunc`. It is invoked as
|
||||
``denoise_fn(video_state, audio_state, sigmas, step_index)`` and must
|
||||
return a tuple ``(denoised_video, denoised_audio)``, where each element
|
||||
is a tensor with the same shape as the corresponding latent.
|
||||
### Returns
|
||||
tuple[LatentState, LatentState]
|
||||
A pair ``(video_state, audio_state)`` containing the final video and
|
||||
audio latent states after completing the denoising loop.
|
||||
"""
|
||||
for step_idx, _ in enumerate(tqdm(sigmas[:-1])):
|
||||
denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_idx)
|
||||
|
||||
denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent)
|
||||
denoised_audio = post_process_latent(denoised_audio, audio_state.denoise_mask, audio_state.clean_latent)
|
||||
|
||||
video_state = replace(video_state, latent=stepper.step(video_state.latent, denoised_video, sigmas, step_idx))
|
||||
audio_state = replace(audio_state, latent=stepper.step(audio_state.latent, denoised_audio, sigmas, step_idx))
|
||||
|
||||
return (video_state, audio_state)
|
||||
|
||||
|
||||
def gradient_estimating_euler_denoising_loop(
|
||||
sigmas: torch.Tensor,
|
||||
video_state: LatentState,
|
||||
audio_state: LatentState,
|
||||
stepper: DiffusionStepProtocol,
|
||||
denoise_fn: DenoisingFunc,
|
||||
ge_gamma: float = 2.0,
|
||||
) -> tuple[LatentState, LatentState]:
|
||||
"""
|
||||
Perform the joint audio-video denoising loop using gradient-estimation sampling.
|
||||
This function is similar to :func:`euler_denoising_loop`, but applies
|
||||
gradient estimation to improve the denoised estimates by tracking velocity
|
||||
changes across steps. See the referenced function for detailed parameter
|
||||
documentation.
|
||||
### Parameters
|
||||
ge_gamma:
|
||||
Gradient estimation coefficient controlling the velocity correction term.
|
||||
Default is 2.0. Paper: https://openreview.net/pdf?id=o2ND9v0CeK
|
||||
sigmas, video_state, audio_state, stepper, denoise_fn:
|
||||
See :func:`euler_denoising_loop` for parameter descriptions.
|
||||
### Returns
|
||||
tuple[LatentState, LatentState]
|
||||
See :func:`euler_denoising_loop` for return value description.
|
||||
"""
|
||||
|
||||
previous_audio_velocity = None
|
||||
previous_video_velocity = None
|
||||
|
||||
def update_velocity_and_sample(
|
||||
noisy_sample: torch.Tensor, denoised_sample: torch.Tensor, sigma: float, previous_velocity: torch.Tensor | None
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
current_velocity = to_velocity(noisy_sample, sigma, denoised_sample)
|
||||
if previous_velocity is not None:
|
||||
delta_v = current_velocity - previous_velocity
|
||||
total_velocity = ge_gamma * delta_v + previous_velocity
|
||||
denoised_sample = to_denoised(noisy_sample, total_velocity, sigma)
|
||||
return current_velocity, denoised_sample
|
||||
|
||||
for step_idx, _ in enumerate(tqdm(sigmas[:-1])):
|
||||
denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_idx)
|
||||
|
||||
denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent)
|
||||
denoised_audio = post_process_latent(denoised_audio, audio_state.denoise_mask, audio_state.clean_latent)
|
||||
|
||||
if sigmas[step_idx + 1] == 0:
|
||||
return replace(video_state, latent=denoised_video), replace(audio_state, latent=denoised_audio)
|
||||
|
||||
previous_video_velocity, denoised_video = update_velocity_and_sample(
|
||||
video_state.latent, denoised_video, sigmas[step_idx], previous_video_velocity
|
||||
)
|
||||
previous_audio_velocity, denoised_audio = update_velocity_and_sample(
|
||||
audio_state.latent, denoised_audio, sigmas[step_idx], previous_audio_velocity
|
||||
)
|
||||
|
||||
video_state = replace(video_state, latent=stepper.step(video_state.latent, denoised_video, sigmas, step_idx))
|
||||
audio_state = replace(audio_state, latent=stepper.step(audio_state.latent, denoised_audio, sigmas, step_idx))
|
||||
|
||||
return (video_state, audio_state)
|
||||
|
||||
|
||||
def noise_video_state(
|
||||
output_shape: VideoPixelShape,
|
||||
noiser: Noiser,
|
||||
@@ -313,7 +204,10 @@ def post_process_latent(denoised: torch.Tensor, denoise_mask: torch.Tensor, clea
|
||||
|
||||
|
||||
def modality_from_latent_state(
|
||||
state: LatentState, context: torch.Tensor, sigma: float | torch.Tensor, enabled: bool = True
|
||||
state: LatentState,
|
||||
context: torch.Tensor,
|
||||
sigma: torch.Tensor,
|
||||
enabled: bool = True,
|
||||
) -> Modality:
|
||||
"""Create a Modality from a latent state.
|
||||
Constructs a Modality object with the latent state's data, timesteps derived
|
||||
@@ -322,10 +216,12 @@ def modality_from_latent_state(
|
||||
return Modality(
|
||||
enabled=enabled,
|
||||
latent=state.latent,
|
||||
sigma=sigma,
|
||||
timesteps=timesteps_from_mask(state.denoise_mask, sigma),
|
||||
positions=state.positions,
|
||||
context=context,
|
||||
context_mask=None,
|
||||
attention_mask=state.attention_mask,
|
||||
)
|
||||
|
||||
|
||||
@@ -389,10 +285,10 @@ def multi_modal_guider_denoising_func(
|
||||
v_context: torch.Tensor,
|
||||
a_context: torch.Tensor,
|
||||
transformer: X0Model,
|
||||
*,
|
||||
last_denoised_video: torch.Tensor | None = None,
|
||||
last_denoised_audio: torch.Tensor | None = None,
|
||||
) -> DenoisingFunc:
|
||||
last_denoised_video = None
|
||||
last_denoised_audio = None
|
||||
|
||||
def guider_denoising_step(
|
||||
video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
@@ -490,6 +386,43 @@ def multi_modal_guider_denoising_func(
|
||||
return guider_denoising_step
|
||||
|
||||
|
||||
def multi_modal_guider_factory_denoising_func(
|
||||
video_guider_factory: MultiModalGuiderFactory,
|
||||
audio_guider_factory: MultiModalGuiderFactory | None,
|
||||
v_context: torch.Tensor,
|
||||
a_context: torch.Tensor,
|
||||
transformer: X0Model,
|
||||
) -> DenoisingFunc:
|
||||
"""Resolve guiders per step via factory.build_from_sigma, then multi_modal_guider_denoising_func."""
|
||||
last_denoised_video: torch.Tensor | None = None
|
||||
last_denoised_audio: torch.Tensor | None = None
|
||||
sigma_vals_cached: list[float] | None = None
|
||||
|
||||
def guider_denoising_step(
|
||||
video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
nonlocal last_denoised_video, last_denoised_audio, sigma_vals_cached
|
||||
if sigma_vals_cached is None:
|
||||
sigma_vals_cached = sigmas.detach().cpu().tolist()
|
||||
sigma_val = sigma_vals_cached[step_index]
|
||||
video_guider = video_guider_factory.build_from_sigma(sigma_val)
|
||||
audio_guider = (audio_guider_factory or video_guider_factory).build_from_sigma(sigma_val)
|
||||
denoise_fn = multi_modal_guider_denoising_func(
|
||||
video_guider,
|
||||
audio_guider,
|
||||
v_context,
|
||||
a_context,
|
||||
transformer,
|
||||
last_denoised_video=last_denoised_video,
|
||||
last_denoised_audio=last_denoised_audio,
|
||||
)
|
||||
denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_index)
|
||||
last_denoised_video, last_denoised_audio = denoised_video, denoised_audio
|
||||
return denoised_video, denoised_audio
|
||||
|
||||
return guider_denoising_step
|
||||
|
||||
|
||||
def denoise_audio_video( # noqa: PLR0913
|
||||
output_shape: VideoPixelShape,
|
||||
conditionings: list[ConditioningItem],
|
||||
@@ -540,6 +473,57 @@ def denoise_audio_video( # noqa: PLR0913
|
||||
return video_state, audio_state
|
||||
|
||||
|
||||
def denoise_video_only( # noqa: PLR0913
|
||||
output_shape: VideoPixelShape,
|
||||
conditionings: list[ConditioningItem],
|
||||
noiser: Noiser,
|
||||
sigmas: torch.Tensor,
|
||||
stepper: DiffusionStepProtocol,
|
||||
denoising_loop_fn: DenoisingLoopFunc,
|
||||
components: PipelineComponents,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
noise_scale: float = 1.0,
|
||||
initial_video_latent: torch.Tensor | None = None,
|
||||
initial_audio_latent: torch.Tensor | None = None,
|
||||
) -> LatentState:
|
||||
video_state, video_tools = noise_video_state(
|
||||
output_shape=output_shape,
|
||||
noiser=noiser,
|
||||
conditionings=conditionings,
|
||||
components=components,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
noise_scale=noise_scale,
|
||||
initial_latent=initial_video_latent,
|
||||
)
|
||||
|
||||
audio_state, _ = noise_audio_state(
|
||||
output_shape=output_shape,
|
||||
noiser=noiser,
|
||||
conditionings=[],
|
||||
components=components,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
noise_scale=0.0,
|
||||
initial_latent=initial_audio_latent,
|
||||
)
|
||||
|
||||
audio_state = replace(audio_state, denoise_mask=torch.zeros_like(audio_state.denoise_mask))
|
||||
|
||||
video_state, audio_state = denoising_loop_fn(
|
||||
sigmas,
|
||||
video_state,
|
||||
audio_state,
|
||||
stepper,
|
||||
)
|
||||
|
||||
video_state = video_tools.clear_conditioning(video_state)
|
||||
video_state = video_tools.unpatchify(video_state)
|
||||
|
||||
return video_state
|
||||
|
||||
|
||||
_UNICODE_REPLACEMENTS = str.maketrans("\u2018\u2019\u201c\u201d\u2014\u2013\u00a0\u2032\u2212", "''\"\"-- '-")
|
||||
|
||||
|
||||
@@ -555,7 +539,7 @@ def clean_response(text: str) -> str:
|
||||
|
||||
|
||||
def generate_enhanced_prompt(
|
||||
text_encoder: GemmaTextEncoderModelBase,
|
||||
text_encoder: GemmaTextEncoder,
|
||||
prompt: str,
|
||||
image_path: str | None = None,
|
||||
image_long_side: int = 896,
|
||||
|
||||
@@ -12,6 +12,7 @@ from PIL import Image
|
||||
from torch._prims_common import DeviceLikeType
|
||||
from tqdm import tqdm
|
||||
|
||||
from ltx_core.types import Audio
|
||||
from ltx_pipelines.utils.constants import DEFAULT_IMAGE_CRF
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -79,14 +80,19 @@ def normalize_latent(latent: torch.Tensor, device: torch.device, dtype: torch.dt
|
||||
|
||||
|
||||
def load_image_conditioning(
|
||||
image_path: str, height: int, width: int, dtype: torch.dtype, device: torch.device
|
||||
image_path: str,
|
||||
height: int,
|
||||
width: int,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
crf: int = DEFAULT_IMAGE_CRF,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Loads an image from a path and preprocesses it for conditioning.
|
||||
Note: The image is resized to the nearest multiple of 2 for compatibility with video codecs.
|
||||
"""
|
||||
image = decode_image(image_path=image_path)
|
||||
image = preprocess(image=image)
|
||||
image = preprocess(image=image, crf=crf)
|
||||
image = torch.tensor(image, dtype=torch.float32, device=device)
|
||||
image = resize_and_center_crop(image, height, width)
|
||||
image = normalize_latent(image, device, dtype)
|
||||
@@ -115,9 +121,8 @@ def decode_image(image_path: str) -> np.ndarray:
|
||||
return np_array
|
||||
|
||||
|
||||
def _write_audio(
|
||||
container: av.container.Container, audio_stream: av.audio.AudioStream, samples: torch.Tensor, audio_sample_rate: int
|
||||
) -> None:
|
||||
def _write_audio(container: av.container.Container, audio_stream: av.audio.AudioStream, audio: Audio) -> None:
|
||||
samples = audio.waveform
|
||||
if samples.ndim == 1:
|
||||
samples = samples[:, None]
|
||||
|
||||
@@ -137,7 +142,7 @@ def _write_audio(
|
||||
format="s16",
|
||||
layout="stereo",
|
||||
)
|
||||
frame_in.sample_rate = audio_sample_rate
|
||||
frame_in.sample_rate = audio.sampling_rate
|
||||
|
||||
_resample_audio(container, audio_stream, frame_in)
|
||||
|
||||
@@ -185,8 +190,7 @@ def _resample_audio(
|
||||
def encode_video(
|
||||
video: torch.Tensor | Iterator[torch.Tensor],
|
||||
fps: int,
|
||||
audio: torch.Tensor | None,
|
||||
audio_sample_rate: int | None,
|
||||
audio: Audio | None,
|
||||
output_path: str,
|
||||
video_chunks_number: int,
|
||||
) -> None:
|
||||
@@ -204,10 +208,7 @@ def encode_video(
|
||||
stream.pix_fmt = "yuv420p"
|
||||
|
||||
if audio is not None:
|
||||
if audio_sample_rate is None:
|
||||
raise ValueError("audio_sample_rate is required when audio is provided")
|
||||
|
||||
audio_stream = _prepare_audio_stream(container, audio_sample_rate)
|
||||
audio_stream = _prepare_audio_stream(container, audio.sampling_rate)
|
||||
|
||||
def all_tiles(
|
||||
first_chunk: torch.Tensor, tiles_generator: Generator[tuple[torch.Tensor, int], None, None]
|
||||
@@ -227,27 +228,114 @@ def encode_video(
|
||||
container.mux(packet)
|
||||
|
||||
if audio is not None:
|
||||
_write_audio(container, audio_stream, audio, audio_sample_rate)
|
||||
_write_audio(container, audio_stream, audio)
|
||||
|
||||
container.close()
|
||||
logger.info(f"Video saved to {output_path}")
|
||||
|
||||
|
||||
def decode_audio_from_file(path: str, device: torch.device) -> torch.Tensor | None:
|
||||
_INT_FORMAT_MAX: dict[str, float] = {
|
||||
"u8": 128.0,
|
||||
"u8p": 128.0,
|
||||
"s16": 32768.0,
|
||||
"s16p": 32768.0,
|
||||
"s32": 2147483648.0,
|
||||
"s32p": 2147483648.0,
|
||||
}
|
||||
|
||||
|
||||
def _audio_frame_to_float(frame: av.AudioFrame) -> np.ndarray:
|
||||
"""Convert an audio frame to a float32 ndarray with values in [-1, 1] and shape (channels, samples)."""
|
||||
fmt = frame.format.name
|
||||
arr = frame.to_ndarray().astype(np.float32)
|
||||
if fmt in _INT_FORMAT_MAX:
|
||||
arr = arr / _INT_FORMAT_MAX[fmt]
|
||||
if not frame.format.is_planar:
|
||||
# Interleaved formats have shape (1, samples * channels) — reshape to (channels, samples).
|
||||
channels = len(frame.layout.channels)
|
||||
arr = arr.reshape(-1, channels).T
|
||||
return arr
|
||||
|
||||
|
||||
def get_videostream_metadata(path: str) -> tuple[float, int, int, int]:
|
||||
"""Read video stream metadata: (fps, num_frames, width, height).
|
||||
If frame count is missing in the container, decodes the stream to count frames.
|
||||
"""
|
||||
container = av.open(path)
|
||||
try:
|
||||
audio = []
|
||||
audio_stream = next(s for s in container.streams if s.type == "audio")
|
||||
for frame in container.decode(audio_stream):
|
||||
audio.append(torch.tensor(frame.to_ndarray(), dtype=torch.float32, device=device).unsqueeze(0))
|
||||
container.close()
|
||||
audio = torch.cat(audio)
|
||||
except StopIteration:
|
||||
audio = None
|
||||
video_stream = next(s for s in container.streams if s.type == "video")
|
||||
fps = float(video_stream.average_rate)
|
||||
num_frames = video_stream.frames or 0
|
||||
if num_frames == 0:
|
||||
num_frames = sum(1 for _ in container.decode(video_stream))
|
||||
width = video_stream.codec_context.width
|
||||
height = video_stream.codec_context.height
|
||||
return fps, num_frames, width, height
|
||||
finally:
|
||||
container.close()
|
||||
|
||||
return audio
|
||||
|
||||
def decode_audio_from_file(
|
||||
path: str, device: torch.device, start_time: float = 0.0, max_duration: float | None = None
|
||||
) -> Audio | None:
|
||||
"""Decodes audio from a file, optionally seeking to a start time and limiting duration.
|
||||
Args:
|
||||
path: Path to the audio/video file containing an audio stream.
|
||||
device: Device to place the resulting tensor on.
|
||||
start_time: Start time in seconds to begin reading audio from.
|
||||
max_duration: Maximum audio duration in seconds. If None, reads to end of stream.
|
||||
Returns:
|
||||
An Audio object with waveform of shape (1, channels, samples), or None if no audio stream.
|
||||
"""
|
||||
container = av.open(path)
|
||||
try:
|
||||
audio_stream = next(s for s in container.streams if s.type == "audio")
|
||||
except StopIteration:
|
||||
container.close()
|
||||
return None
|
||||
|
||||
sample_rate = audio_stream.rate
|
||||
start_pts = int(start_time / audio_stream.time_base)
|
||||
end_time = start_time + max_duration if max_duration else audio_stream.duration * audio_stream.time_base
|
||||
container.seek(start_pts, stream=audio_stream)
|
||||
|
||||
samples = []
|
||||
first_frame_time = None
|
||||
for frame in container.decode(audio=0):
|
||||
if frame.pts is None:
|
||||
continue
|
||||
frame_time = float(frame.pts * audio_stream.time_base)
|
||||
frame_end = frame_time + frame.samples / frame.sample_rate
|
||||
if frame_end < start_time:
|
||||
continue
|
||||
if frame_time > end_time:
|
||||
break
|
||||
if first_frame_time is None:
|
||||
first_frame_time = frame_time
|
||||
samples.append(_audio_frame_to_float(frame))
|
||||
|
||||
container.close()
|
||||
|
||||
if not samples:
|
||||
return None
|
||||
|
||||
audio = np.concatenate(samples, axis=-1)
|
||||
|
||||
# Trim samples that fall outside the requested [start_time, start_time + max_duration] window.
|
||||
# Audio codecs decode in fixed-size frames whose boundaries may not align with the requested
|
||||
# time range, so the first frame can start before start_time and the last frame can end after
|
||||
# start_time + max_duration.
|
||||
skip_samples = round((start_time - first_frame_time) * sample_rate)
|
||||
if skip_samples > 0:
|
||||
audio = audio[..., skip_samples:]
|
||||
|
||||
if max_duration is not None:
|
||||
max_samples = round(max_duration * sample_rate)
|
||||
audio = audio[..., :max_samples]
|
||||
|
||||
waveform = torch.from_numpy(audio).to(device).unsqueeze(0)
|
||||
|
||||
return Audio(waveform=waveform, sampling_rate=sample_rate)
|
||||
|
||||
|
||||
def decode_video_from_file(path: str, frame_cap: int, device: DeviceLikeType) -> Generator[torch.Tensor]:
|
||||
|
||||
@@ -8,9 +8,12 @@ from ltx_core.loader.registry import DummyRegistry, Registry
|
||||
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder as Builder
|
||||
from ltx_core.model.audio_vae import (
|
||||
AUDIO_VAE_DECODER_COMFY_KEYS_FILTER,
|
||||
AUDIO_VAE_ENCODER_COMFY_KEYS_FILTER,
|
||||
VOCODER_COMFY_KEYS_FILTER,
|
||||
AudioDecoder,
|
||||
AudioDecoderConfigurator,
|
||||
AudioEncoder,
|
||||
AudioEncoderConfigurator,
|
||||
Vocoder,
|
||||
VocoderConfigurator,
|
||||
)
|
||||
@@ -31,11 +34,11 @@ from ltx_core.model.video_vae import (
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.text_encoders.gemma import (
|
||||
AV_GEMMA_TEXT_ENCODER_KEY_OPS,
|
||||
AVGemmaTextEncoderModel,
|
||||
AVGemmaTextEncoderModelConfigurator,
|
||||
GEMMA_MODEL_OPS,
|
||||
GemmaTextEncoder,
|
||||
GemmaTextEncoderConfigurator,
|
||||
module_ops_from_gemma_root,
|
||||
)
|
||||
from ltx_core.text_encoders.gemma.encoders.av_encoder import GEMMA_MODEL_OPS
|
||||
from ltx_core.utils import find_matching_file
|
||||
|
||||
|
||||
@@ -131,6 +134,13 @@ class ModelLedger:
|
||||
registry=self.registry,
|
||||
)
|
||||
|
||||
self.audio_encoder_builder = Builder[AudioEncoder](
|
||||
model_path=self.checkpoint_path,
|
||||
model_class_configurator=AudioEncoderConfigurator,
|
||||
model_sd_ops=AUDIO_VAE_ENCODER_COMFY_KEYS_FILTER,
|
||||
registry=self.registry,
|
||||
)
|
||||
|
||||
self.audio_decoder_builder = Builder(
|
||||
model_path=self.checkpoint_path,
|
||||
model_class_configurator=AudioDecoderConfigurator,
|
||||
@@ -152,7 +162,7 @@ class ModelLedger:
|
||||
|
||||
self.text_encoder_builder = Builder(
|
||||
model_path=(str(self.checkpoint_path), *weight_paths),
|
||||
model_class_configurator=AVGemmaTextEncoderModelConfigurator,
|
||||
model_class_configurator=GemmaTextEncoderConfigurator,
|
||||
model_sd_ops=AV_GEMMA_TEXT_ENCODER_KEY_OPS,
|
||||
registry=self.registry,
|
||||
module_ops=(GEMMA_MODEL_OPS, *module_ops),
|
||||
@@ -225,7 +235,7 @@ class ModelLedger:
|
||||
|
||||
return self.vae_encoder_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval()
|
||||
|
||||
def text_encoder(self) -> AVGemmaTextEncoderModel:
|
||||
def text_encoder(self) -> GemmaTextEncoder:
|
||||
if not hasattr(self, "text_encoder_builder"):
|
||||
raise ValueError(
|
||||
"Text encoder not initialized. Please provide a checkpoint path and gemma root path to the "
|
||||
@@ -234,6 +244,14 @@ class ModelLedger:
|
||||
|
||||
return self.text_encoder_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval()
|
||||
|
||||
def audio_encoder(self) -> AudioEncoder:
|
||||
if not hasattr(self, "audio_encoder_builder"):
|
||||
raise ValueError(
|
||||
"Audio encoder not initialized. Please provide a checkpoint path to the ModelLedger constructor."
|
||||
)
|
||||
|
||||
return self.audio_encoder_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval()
|
||||
|
||||
def audio_decoder(self) -> AudioDecoder:
|
||||
if not hasattr(self, "audio_decoder_builder"):
|
||||
raise ValueError(
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
import math
|
||||
|
||||
|
||||
def phi(j: int, neg_h: float) -> float:
|
||||
"""
|
||||
Compute φⱼ(z) where z = -h (negative step size in log-space)
|
||||
φ₁(z) = (e^z - 1) / z
|
||||
φ₂(z) = (e^z - 1 - z) / z²
|
||||
φⱼ(z) = (e^z - Σₖ₌₀^(j-1) zᵏ/k!) / zʲ
|
||||
These functions naturally appear when solving:
|
||||
dx/dt = A*x + g(x,t) (linear drift + nonlinear part)
|
||||
"""
|
||||
if abs(neg_h) < 1e-10:
|
||||
# Taylor series for small h to avoid division by zero
|
||||
# φⱼ(0) = 1/j!
|
||||
return 1.0 / math.factorial(j)
|
||||
|
||||
# Compute the "remainder" sum: Σₖ₌₀^(j-1) z^k/k!
|
||||
remainder = sum(neg_h**k / math.factorial(k) for k in range(j))
|
||||
|
||||
# φⱼ(z) = (e^z - remainder) / z^j
|
||||
return (math.exp(neg_h) - remainder) / (neg_h**j)
|
||||
|
||||
|
||||
def get_res2s_coefficients(h: float, phi_cache: dict, c2: float = 0.5) -> tuple[float, float, float]:
|
||||
"""
|
||||
Compute res_2s Runge-Kutta coefficients for a given step size.
|
||||
Args:
|
||||
h: Step size in log-space = log(sigma / sigma_next)
|
||||
phi_cache: Dictionary to cache phi function results. Cache key: (j, neg_h)
|
||||
c2: Substep position (default 0.5 = midpoint)
|
||||
Returns:
|
||||
a21: Coefficient for computing intermediate x
|
||||
b1, b2: Coefficients for final combination
|
||||
"""
|
||||
|
||||
def get_phi(j: int, neg_h: float) -> float:
|
||||
"""Get phi value with caching."""
|
||||
cache_key = (j, neg_h)
|
||||
if cache_key in phi_cache:
|
||||
return phi_cache[cache_key]
|
||||
result = phi(j, neg_h)
|
||||
phi_cache[cache_key] = result
|
||||
return result
|
||||
|
||||
# Substep coefficient: how much of ε₁ to use for intermediate point
|
||||
# a21 = c2 * φ₁(-h*c2)
|
||||
neg_h_c2 = -h * c2
|
||||
phi_1_c2 = get_phi(1, neg_h_c2)
|
||||
a21 = c2 * phi_1_c2
|
||||
|
||||
# Final combination weights
|
||||
# b2 = φ₂(-h) / c2
|
||||
neg_h_full = -h
|
||||
phi_2_full = get_phi(2, neg_h_full)
|
||||
b2 = phi_2_full / c2
|
||||
|
||||
# b1 = φ₁(-h) - b2
|
||||
phi_1_full = get_phi(1, neg_h_full)
|
||||
b1 = phi_1_full - b2
|
||||
|
||||
return a21, b1, b2
|
||||
@@ -0,0 +1,363 @@
|
||||
import logging
|
||||
from dataclasses import replace
|
||||
from functools import partial
|
||||
from typing import Callable
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from ltx_core.components.diffusion_steps import Res2sDiffusionStep
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.utils import to_denoised, to_velocity
|
||||
from ltx_pipelines.utils.helpers import post_process_latent, timesteps_from_mask
|
||||
from ltx_pipelines.utils.res2s import get_res2s_coefficients
|
||||
from ltx_pipelines.utils.types import DenoisingFunc, LatentState
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def euler_denoising_loop(
|
||||
sigmas: torch.Tensor,
|
||||
video_state: LatentState,
|
||||
audio_state: LatentState,
|
||||
stepper: DiffusionStepProtocol,
|
||||
denoise_fn: DenoisingFunc,
|
||||
) -> tuple[LatentState, LatentState]:
|
||||
"""
|
||||
Perform the joint audio-video denoising loop over a diffusion schedule.
|
||||
This function iterates over all but the final value in ``sigmas`` and, at
|
||||
each diffusion step, calls ``denoise_fn`` to obtain denoised video and
|
||||
audio latents. The denoised latents are post-processed with their
|
||||
respective denoise masks and clean latents, then passed to ``stepper`` to
|
||||
advance the noisy latents one step along the diffusion schedule.
|
||||
### Parameters
|
||||
sigmas:
|
||||
A 1D tensor of noise levels (diffusion sigmas) defining the sampling
|
||||
schedule. All steps except the last element are iterated over.
|
||||
video_state:
|
||||
The current video :class:`LatentState`, containing the noisy latent,
|
||||
its clean reference latent, and the denoising mask.
|
||||
audio_state:
|
||||
The current audio :class:`LatentState`, analogous to ``video_state``
|
||||
but for the audio modality.
|
||||
stepper:
|
||||
An implementation of :class:`DiffusionStepProtocol` that updates a
|
||||
latent given the current latent, its denoised estimate, the full
|
||||
``sigmas`` schedule, and the current step index.
|
||||
denoise_fn:
|
||||
A callable implementing :class:`DenoisingFunc`. It is invoked as
|
||||
``denoise_fn(video_state, audio_state, sigmas, step_index)`` and must
|
||||
return a tuple ``(denoised_video, denoised_audio)``, where each element
|
||||
is a tensor with the same shape as the corresponding latent.
|
||||
### Returns
|
||||
tuple[LatentState, LatentState]
|
||||
A pair ``(video_state, audio_state)`` containing the final video and
|
||||
audio latent states after completing the denoising loop.
|
||||
"""
|
||||
for step_idx, _ in enumerate(tqdm(sigmas[:-1])):
|
||||
denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_idx)
|
||||
|
||||
denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent)
|
||||
denoised_audio = post_process_latent(denoised_audio, audio_state.denoise_mask, audio_state.clean_latent)
|
||||
|
||||
video_state = replace(video_state, latent=stepper.step(video_state.latent, denoised_video, sigmas, step_idx))
|
||||
audio_state = replace(audio_state, latent=stepper.step(audio_state.latent, denoised_audio, sigmas, step_idx))
|
||||
|
||||
return (video_state, audio_state)
|
||||
|
||||
|
||||
def gradient_estimating_euler_denoising_loop(
|
||||
sigmas: torch.Tensor,
|
||||
video_state: LatentState,
|
||||
audio_state: LatentState,
|
||||
stepper: DiffusionStepProtocol,
|
||||
denoise_fn: DenoisingFunc,
|
||||
ge_gamma: float = 2.0,
|
||||
) -> tuple[LatentState, LatentState]:
|
||||
"""
|
||||
Perform the joint audio-video denoising loop using gradient-estimation sampling.
|
||||
This function is similar to :func:`euler_denoising_loop`, but applies
|
||||
gradient estimation to improve the denoised estimates by tracking velocity
|
||||
changes across steps. See the referenced function for detailed parameter
|
||||
documentation.
|
||||
### Parameters
|
||||
ge_gamma:
|
||||
Gradient estimation coefficient controlling the velocity correction term.
|
||||
Default is 2.0. Paper: https://openreview.net/pdf?id=o2ND9v0CeK
|
||||
sigmas, video_state, audio_state, stepper, denoise_fn:
|
||||
See :func:`euler_denoising_loop` for parameter descriptions.
|
||||
### Returns
|
||||
tuple[LatentState, LatentState]
|
||||
See :func:`euler_denoising_loop` for return value description.
|
||||
"""
|
||||
|
||||
previous_audio_velocity = None
|
||||
previous_video_velocity = None
|
||||
|
||||
def update_velocity_and_sample(
|
||||
noisy_sample: torch.Tensor, denoised_sample: torch.Tensor, sigma: float, previous_velocity: torch.Tensor | None
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
current_velocity = to_velocity(noisy_sample, sigma, denoised_sample)
|
||||
if previous_velocity is not None:
|
||||
delta_v = current_velocity - previous_velocity
|
||||
total_velocity = ge_gamma * delta_v + previous_velocity
|
||||
denoised_sample = to_denoised(noisy_sample, total_velocity, sigma)
|
||||
return current_velocity, denoised_sample
|
||||
|
||||
for step_idx, _ in enumerate(tqdm(sigmas[:-1])):
|
||||
denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_idx)
|
||||
|
||||
denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent)
|
||||
denoised_audio = post_process_latent(denoised_audio, audio_state.denoise_mask, audio_state.clean_latent)
|
||||
|
||||
if sigmas[step_idx + 1] == 0:
|
||||
return replace(video_state, latent=denoised_video), replace(audio_state, latent=denoised_audio)
|
||||
|
||||
previous_video_velocity, denoised_video = update_velocity_and_sample(
|
||||
video_state.latent, denoised_video, sigmas[step_idx], previous_video_velocity
|
||||
)
|
||||
previous_audio_velocity, denoised_audio = update_velocity_and_sample(
|
||||
audio_state.latent, denoised_audio, sigmas[step_idx], previous_audio_velocity
|
||||
)
|
||||
|
||||
video_state = replace(video_state, latent=stepper.step(video_state.latent, denoised_video, sigmas, step_idx))
|
||||
audio_state = replace(audio_state, latent=stepper.step(audio_state.latent, denoised_audio, sigmas, step_idx))
|
||||
|
||||
return (video_state, audio_state)
|
||||
|
||||
|
||||
def _channelwise_normalize(x: torch.Tensor) -> torch.Tensor:
|
||||
return x.sub_(x.mean(dim=(-2, -1), keepdim=True)).div_(x.std(dim=(-2, -1), keepdim=True))
|
||||
|
||||
|
||||
def _get_new_noise(x: torch.Tensor, generator: torch.Generator) -> torch.Tensor:
|
||||
noise = torch.randn(x.shape, generator=generator, dtype=torch.float64, device=generator.device)
|
||||
noise = (noise - noise.mean()) / noise.std()
|
||||
return _channelwise_normalize(noise)
|
||||
|
||||
|
||||
def _inject_sde_noise(
|
||||
state: LatentState,
|
||||
sample: torch.Tensor,
|
||||
denoised_sample: torch.Tensor,
|
||||
step_noise_generator: torch.Generator,
|
||||
new_noise_fn: Callable[[torch.Tensor, torch.Generator], torch.Tensor],
|
||||
stepper: DiffusionStepProtocol,
|
||||
sigmas: torch.Tensor,
|
||||
step_idx: int,
|
||||
legacy_mode: bool = False,
|
||||
) -> torch.Tensor:
|
||||
sigmas_copy = sigmas.clone()
|
||||
new_noise = new_noise_fn(state.latent, step_noise_generator)
|
||||
if not legacy_mode:
|
||||
timesteps = timesteps_from_mask(state.denoise_mask.double(), sigmas_copy[step_idx].double())
|
||||
next_timesteps = timesteps_from_mask(state.denoise_mask.double(), sigmas_copy[step_idx + 1].double())
|
||||
sigmas = torch.stack([timesteps, next_timesteps])
|
||||
step_idx = 0
|
||||
x_next = stepper.step(
|
||||
sample=sample,
|
||||
denoised_sample=denoised_sample,
|
||||
sigmas=sigmas,
|
||||
step_index=step_idx,
|
||||
noise=new_noise,
|
||||
)
|
||||
|
||||
if legacy_mode:
|
||||
x_next = post_process_latent(x_next, state.denoise_mask, state.clean_latent)
|
||||
|
||||
return x_next
|
||||
|
||||
|
||||
def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915
|
||||
sigmas: torch.Tensor,
|
||||
video_state: LatentState,
|
||||
audio_state: LatentState,
|
||||
stepper: DiffusionStepProtocol,
|
||||
denoise_fn: DenoisingFunc,
|
||||
noise_seed: int = -1,
|
||||
noise_seed_substep: int | None = None,
|
||||
bongmath: bool = True,
|
||||
bongmath_max_iter: int = 100,
|
||||
new_noise_fn: Callable[[torch.Tensor, torch.Generator], torch.Tensor] = _get_new_noise,
|
||||
model_dtype: torch.dtype = torch.bfloat16,
|
||||
legacy_mode: bool = True,
|
||||
) -> tuple[LatentState, LatentState]:
|
||||
"""
|
||||
Joint audio-video denoising loop using the res_2s second-order sampler.
|
||||
Iterates over the diffusion schedule with a two-stage Runge-Kutta step:
|
||||
evaluates the denoiser at the current point and at a midpoint (with SDE
|
||||
noise), then combines both with RK coefficients. Supports anchor-point
|
||||
refinement (bong iteration) and optional SDE noise injection. Requires
|
||||
:class:`Res2sDiffusionStep` as ``stepper``.
|
||||
### Parameters
|
||||
sigmas:
|
||||
A 1D tensor of noise levels defining the sampling schedule.
|
||||
video_state:
|
||||
Current video :class:`LatentState` (noisy latent, clean reference, mask).
|
||||
audio_state:
|
||||
Current audio :class:`LatentState`, same structure as ``video_state``.
|
||||
stepper:
|
||||
Must be an instance of :class:`Res2sDiffusionStep`; performs SDE step
|
||||
with noise injection.
|
||||
denoise_fn:
|
||||
Callable ``(video_state, audio_state, sigmas, step_index)`` returning
|
||||
``(denoised_video, denoised_audio)``.
|
||||
noise_seed:
|
||||
Seed for step-level SDE noise; substep seed defaults to ``noise_seed + 10000``.
|
||||
noise_seed_substep:
|
||||
Optional seed for substep SDE noise; if None, derived from ``noise_seed``.
|
||||
bongmath:
|
||||
Whether to run iterative anchor refinement (bong iteration) when step size is small.
|
||||
bongmath_max_iter:
|
||||
Max iterations for bong refinement when enabled.
|
||||
new_noise_fn:
|
||||
Callable ``(latent, generator) -> noise`` for SDE injection; default
|
||||
uses normalized channel-wise Gaussian noise.
|
||||
model_dtype:
|
||||
Dtype for latent state updates (e.g. bfloat16).
|
||||
### Returns
|
||||
tuple[LatentState, LatentState]
|
||||
Final ``(video_state, audio_state)`` after the denoising loop.
|
||||
"""
|
||||
# Initialize noise generators with different seeds
|
||||
if noise_seed_substep is None:
|
||||
noise_seed_substep = noise_seed + 10000 # Offset to ensure different seeds
|
||||
step_noise_generator = torch.Generator(device=video_state.latent.device).manual_seed(noise_seed)
|
||||
substep_noise_generator = torch.Generator(device=video_state.latent.device).manual_seed(noise_seed_substep)
|
||||
sde_noise_injecting_fn = partial(
|
||||
_inject_sde_noise, stepper=stepper, new_noise_fn=new_noise_fn, legacy_mode=legacy_mode
|
||||
)
|
||||
step_noise_injecting_fn = partial(sde_noise_injecting_fn, step_noise_generator=step_noise_generator)
|
||||
substep_noise_injecting_fn = partial(sde_noise_injecting_fn, step_noise_generator=substep_noise_generator)
|
||||
|
||||
if not isinstance(stepper, Res2sDiffusionStep):
|
||||
raise ValueError("stepper must be an instance of Res2sDiffusionStep")
|
||||
|
||||
n_full_steps = len(sigmas) - 1
|
||||
# inject minimal sigma value to avoid division by zero
|
||||
if sigmas[-1] == 0:
|
||||
sigmas = torch.cat([sigmas[:-1], torch.tensor([0.0011, 0.0], device=sigmas.device)], dim=0)
|
||||
# Compute step sizes in hyperbolic space
|
||||
hs = -torch.log(sigmas[1:].double().cpu() / (sigmas[:-1].double().cpu()))
|
||||
|
||||
# Initialize phi cache for reuse across loop iterations
|
||||
# Cache key: (j, neg_h) where j is phi order and neg_h is negative step value
|
||||
phi_cache = {}
|
||||
c2 = 0.5 # Midpoint for res_2s
|
||||
|
||||
# Progress bar shows only full two-stage steps; final (sigma_next==0) step is done silently
|
||||
|
||||
for step_idx in tqdm(range(n_full_steps)):
|
||||
sigma = sigmas[step_idx].double()
|
||||
sigma_next = sigmas[step_idx + 1].double()
|
||||
|
||||
# Initialize anchor point
|
||||
x_anchor_video = video_state.latent.clone().double()
|
||||
x_anchor_audio = audio_state.latent.clone().double()
|
||||
|
||||
# ====================================================================
|
||||
# STAGE 1: Evaluate at current point
|
||||
# ====================================================================
|
||||
denoised_video_1, denoised_audio_1 = denoise_fn(video_state, audio_state, sigmas, step_idx)
|
||||
denoised_video_1 = post_process_latent(denoised_video_1, video_state.denoise_mask, video_state.clean_latent)
|
||||
denoised_audio_1 = post_process_latent(denoised_audio_1, audio_state.denoise_mask, audio_state.clean_latent)
|
||||
|
||||
h = hs[step_idx].item()
|
||||
|
||||
# Compute RK coefficients (pass phi_cache for caching)
|
||||
a21, b1, b2 = get_res2s_coefficients(h, phi_cache, c2)
|
||||
|
||||
# Compute substep sigma, sqrt is a hardcode for c2 = 0.5
|
||||
sub_sigma = torch.sqrt(sigma * sigma_next)
|
||||
|
||||
# ====================================================================
|
||||
# Compute substep x using RK coefficient a21
|
||||
# ====================================================================
|
||||
eps_1_video = denoised_video_1.double() - x_anchor_video
|
||||
eps_1_audio = denoised_audio_1.double() - x_anchor_audio
|
||||
|
||||
x_mid_video = x_anchor_video.double() + h * a21 * eps_1_video
|
||||
x_mid_audio = x_anchor_audio.double() + h * a21 * eps_1_audio
|
||||
|
||||
# ====================================================================
|
||||
# SDE noise injection at substep
|
||||
# ====================================================================
|
||||
x_mid_video = substep_noise_injecting_fn(
|
||||
state=video_state,
|
||||
sample=x_anchor_video,
|
||||
denoised_sample=x_mid_video,
|
||||
sigmas=torch.stack([sigma, sub_sigma]),
|
||||
step_idx=0,
|
||||
)
|
||||
x_mid_audio = substep_noise_injecting_fn(
|
||||
state=audio_state,
|
||||
sample=x_anchor_audio,
|
||||
denoised_sample=x_mid_audio,
|
||||
sigmas=torch.stack([sigma, sub_sigma]),
|
||||
step_idx=0,
|
||||
)
|
||||
# ====================================================================
|
||||
# ITERATIVE REFINEMENT (Bong Iteration) - Stabilize anchor point
|
||||
# ====================================================================
|
||||
if bongmath and h < 0.5 and sigma > 0.03:
|
||||
for _ in range(bongmath_max_iter):
|
||||
x_anchor_video = x_mid_video - h * a21 * eps_1_video
|
||||
eps_1_video = denoised_video_1.double() - x_anchor_video
|
||||
x_anchor_audio = x_mid_audio - h * a21 * eps_1_audio
|
||||
eps_1_audio = denoised_audio_1.double() - x_anchor_audio
|
||||
|
||||
# ====================================================================
|
||||
# STAGE 2: Evaluate at substep point (WITH NOISE)
|
||||
# ====================================================================
|
||||
mid_video_state = replace(video_state, latent=x_mid_video.to(model_dtype))
|
||||
mid_audio_state = replace(audio_state, latent=x_mid_audio.to(model_dtype))
|
||||
|
||||
denoised_video_2, denoised_audio_2 = denoise_fn(
|
||||
video_state=mid_video_state,
|
||||
audio_state=mid_audio_state,
|
||||
sigmas=torch.stack([sub_sigma]).to(sigmas.device),
|
||||
step_index=0,
|
||||
)
|
||||
denoised_video_2 = post_process_latent(denoised_video_2, video_state.denoise_mask, video_state.clean_latent)
|
||||
denoised_audio_2 = post_process_latent(denoised_audio_2, audio_state.denoise_mask, audio_state.clean_latent)
|
||||
|
||||
# ====================================================================
|
||||
# FINAL COMBINATION: Compute x_next using RK coefficients
|
||||
# ====================================================================
|
||||
eps_2_video = denoised_video_2.double() - x_anchor_video
|
||||
eps_2_audio = denoised_audio_2.double() - x_anchor_audio
|
||||
|
||||
x_next_video = x_anchor_video + h * (b1 * eps_1_video + b2 * eps_2_video)
|
||||
x_next_audio = x_anchor_audio + h * (b1 * eps_1_audio + b2 * eps_2_audio)
|
||||
|
||||
# ====================================================================
|
||||
# SDE NOISE INJECTION AT STEP LEVEL
|
||||
# ====================================================================
|
||||
x_next_video = step_noise_injecting_fn(
|
||||
state=video_state,
|
||||
sample=x_anchor_video,
|
||||
denoised_sample=x_next_video,
|
||||
sigmas=sigmas,
|
||||
step_idx=step_idx,
|
||||
)
|
||||
x_next_audio = step_noise_injecting_fn(
|
||||
state=audio_state,
|
||||
sample=x_anchor_audio,
|
||||
denoised_sample=x_next_audio,
|
||||
sigmas=sigmas,
|
||||
step_idx=step_idx,
|
||||
)
|
||||
|
||||
# Update states
|
||||
video_state = replace(video_state, latent=x_next_video.to(model_dtype))
|
||||
audio_state = replace(audio_state, latent=x_next_audio.to(model_dtype))
|
||||
|
||||
# Final step if we need to fully remove the noise
|
||||
if sigmas[-1] == 0:
|
||||
denoised_video_1, denoised_audio_1 = denoise_fn(video_state, audio_state, sigmas, n_full_steps)
|
||||
denoised_video_1 = post_process_latent(denoised_video_1, video_state.denoise_mask, video_state.clean_latent)
|
||||
denoised_audio_1 = post_process_latent(denoised_audio_1, audio_state.denoise_mask, audio_state.clean_latent)
|
||||
video_state = replace(video_state, latent=denoised_video_1.to(model_dtype))
|
||||
audio_state = replace(audio_state, latent=denoised_audio_1.to(model_dtype))
|
||||
|
||||
return video_state, audio_state
|
||||
Reference in New Issue
Block a user