Automated PR - 2026-03-05

This commit is contained in:
sync-bot
2026-03-05 15:47:20 +00:00
parent 3b6d09d7b6
commit d230aec5cd
29 changed files with 739 additions and 540 deletions
+13 -14
View File
@@ -1,7 +1,7 @@
# LTX-2 # LTX-2
[![Website](https://img.shields.io/badge/Website-LTX-181717?logo=google-chrome)](https://ltx.io) [![Website](https://img.shields.io/badge/Website-LTX-181717?logo=google-chrome)](https://ltx.io)
[![Model](https://img.shields.io/badge/HuggingFace-Model-orange?logo=huggingface)](https://huggingface.co/Lightricks/LTX-2) [![Model](https://img.shields.io/badge/HuggingFace-Model-orange?logo=huggingface)](https://huggingface.co/Lightricks/LTX-2.3)
[![Demo](https://img.shields.io/badge/Demo-Try%20Now-brightgreen?logo=vercel)](https://app.ltx.studio/ltx-2-playground/i2v) [![Demo](https://img.shields.io/badge/Demo-Try%20Now-brightgreen?logo=vercel)](https://app.ltx.studio/ltx-2-playground/i2v)
[![Paper](https://img.shields.io/badge/Paper-PDF-EC1C24?logo=adobeacrobatreader&logoColor=white)](https://arxiv.org/abs/2601.03233) [![Paper](https://img.shields.io/badge/Paper-PDF-EC1C24?logo=adobeacrobatreader&logoColor=white)](https://arxiv.org/abs/2601.03233)
[![Discord](https://img.shields.io/badge/Join-Discord-5865F2?logo=discord)](https://discord.gg/ltxplatform) [![Discord](https://img.shields.io/badge/Join-Discord-5865F2?logo=discord)](https://discord.gg/ltxplatform)
@@ -26,30 +26,29 @@ source .venv/bin/activate
### Required Models ### Required Models
Download the following models from the [LTX-2 HuggingFace repository](https://huggingface.co/Lightricks/LTX-2): Download the following models from the [LTX-2.3 HuggingFace repository](https://huggingface.co/Lightricks/LTX-2.3):
**LTX-2 Model Checkpoint** (choose and download one of the following) **LTX-2.3 Model Checkpoint** (choose and download one of the following)
* [`ltx-2-19b-dev-fp8.safetensors`](https://huggingface.co/Lightricks/LTX-2/blob/main/ltx-2-19b-dev-fp8.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2/resolve/main/ltx-2-19b-dev-fp8.safetensors) * [`ltx-2.3-22b-dev.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-22b-dev.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-22b-dev.safetensors)
* [`ltx-2.3-22b-distilled.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-22b-distilled.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-22b-distilled.safetensors)
* [`ltx-2-19b-dev.safetensors`](https://huggingface.co/Lightricks/LTX-2/blob/main/ltx-2-19b-dev.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2/resolve/main/ltx-2-19b-dev.safetensors)
* [`ltx-2-19b-distilled.safetensors`](https://huggingface.co/Lightricks/LTX-2/blob/main/ltx-2-19b-distilled.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2/resolve/main/ltx-2-19b-distilled.safetensors)
* [`ltx-2-19b-distilled-fp8.safetensors`](https://huggingface.co/Lightricks/LTX-2/blob/main/ltx-2-19b-distilled-fp8.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2/resolve/main/ltx-2-19b-distilled-fp8.safetensors)
**Spatial Upscaler** - Required for current two-stage pipeline implementations in this repository **Spatial Upscaler** - Required for current two-stage pipeline implementations in this repository
* [`ltx-2-spatial-upscaler-x2-1.0.safetensors`](https://huggingface.co/Lightricks/LTX-2/blob/main/ltx-2-spatial-upscaler-x2-1.0.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2/resolve/main/ltx-2-spatial-upscaler-x2-1.0.safetensors) * [`ltx-2.3-spatial-upscaler-x2-1.0.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-spatial-upscaler-x2-1.0.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-spatial-upscaler-x2-1.0.safetensors)
* [`ltx-2.3-spatial-upscaler-x1.5-1.0.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-spatial-upscaler-x1.5-1.0.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-spatial-upscaler-x1.5-1.0.safetensors)
**Temporal Upscaler** - Supported by the model and will be required for future pipeline implementations **Temporal Upscaler** - Supported by the model and will be required for future pipeline implementations
* [`ltx-2-temporal-upscaler-x2-1.0.safetensors`](https://huggingface.co/Lightricks/LTX-2/blob/main/ltx-2-temporal-upscaler-x2-1.0.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2/resolve/main/ltx-2-temporal-upscaler-x2-1.0.safetensors) * [`ltx-2.3-temporal-upscaler-x2-1.0.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-temporal-upscaler-x2-1.0.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-temporal-upscaler-x2-1.0.safetensors)
**Distilled LoRA** - Required for current two-stage pipeline implementations in this repository (except DistilledPipeline and ICLoraPipeline) **Distilled LoRA** - Required for current two-stage pipeline implementations in this repository (except DistilledPipeline and ICLoraPipeline)
* [`ltx-2-19b-distilled-lora-384.safetensors`](https://huggingface.co/Lightricks/LTX-2/blob/main/ltx-2-19b-distilled-lora-384.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2/resolve/main/ltx-2-19b-distilled-lora-384.safetensors) * [`ltx-2.3-22b-distilled-lora-384.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-22b-distilled-lora-384.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-22b-distilled-lora-384.safetensors)
**Gemma Text Encoder** (download all assets from the repository) **Gemma Text Encoder** (download all assets from the repository)
* [`Gemma 3`](https://huggingface.co/google/gemma-3-12b-it-qat-q4_0-unquantized/tree/main) * [`Gemma 3`](https://huggingface.co/google/gemma-3-12b-it-qat-q4_0-unquantized/tree/main)
**LoRAs** **LoRAs**
* [`LTX-2-19b-IC-LoRA-Canny-Control`](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Canny-Control) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Canny-Control/resolve/main/ltx-2-19b-ic-lora-canny-control.safetensors) * [`LTX-2.3-22b-IC-LoRA-Union-Control`](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-Union-Control) - [Download](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-Union-Control/resolve/main/ltx-2.3-22b-ic-lora-union-control-ref0.5.safetensors)
* [`LTX-2-19b-IC-LoRA-Depth-Control`](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Depth-Control) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Depth-Control/resolve/main/ltx-2-19b-ic-lora-depth-control.safetensors) * [`LTX-2.3-22b-IC-LoRA-Inpainting`](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-Inpainting) - [Download](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-Inpainting/resolve/main/ltx-2.3-22b-ic-lora-inpainting.safetensors)
* [`LTX-2.3-22b-IC-LoRA-Motion-Track-Control`](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-Motion-Track-Control) - [Download](https://huggingface.co/Lightricks/LTX-2.3-22b-IC-LoRA-Motion-Track-Control/resolve/main/ltx-2.3-22b-ic-lora-motion-track-control-ref0.5.safetensors)
* [`LTX-2-19b-IC-LoRA-Detailer`](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Detailer) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Detailer/resolve/main/ltx-2-19b-ic-lora-detailer.safetensors) * [`LTX-2-19b-IC-LoRA-Detailer`](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Detailer) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Detailer/resolve/main/ltx-2-19b-ic-lora-detailer.safetensors)
* [`LTX-2-19b-IC-LoRA-Pose-Control`](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Pose-Control) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Pose-Control/resolve/main/ltx-2-19b-ic-lora-pose-control.safetensors) * [`LTX-2-19b-IC-LoRA-Pose-Control`](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Pose-Control) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-IC-LoRA-Pose-Control/resolve/main/ltx-2-19b-ic-lora-pose-control.safetensors)
* [`LTX-2-19b-LoRA-Camera-Control-Dolly-In`](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Dolly-In) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Dolly-In/resolve/main/ltx-2-19b-lora-camera-control-dolly-in.safetensors) * [`LTX-2-19b-LoRA-Camera-Control-Dolly-In`](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Dolly-In) - [Download](https://huggingface.co/Lightricks/LTX-2-19b-LoRA-Camera-Control-Dolly-In/resolve/main/ltx-2-19b-lora-camera-control-dolly-in.safetensors)
@@ -63,7 +62,7 @@ Download the following models from the [LTX-2 HuggingFace repository](https://hu
### Available Pipelines ### Available Pipelines
* **[TI2VidTwoStagesPipeline](packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages.py)** - Production-quality text/image-to-video with 2x upsampling (recommended) * **[TI2VidTwoStagesPipeline](packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages.py)** - Production-quality text/image-to-video with 2x upsampling (recommended)
* **[TI2VidTwoStagesRes2sPipeline](packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages_res2s.py)** - Same two-stage flow as above but uses the res_2s second-order sampler (fewer steps, different quality/speed trade-off) * **[TI2VidTwoStagesHQPipeline](packages/ltx-pipelines/src/ltx_pipelines/ti2vid_two_stages_hq.py)** - Same two-stage flow as above but uses the res_2s second-order sampler (fewer steps, better quality)
* **[TI2VidOneStagePipeline](packages/ltx-pipelines/src/ltx_pipelines/ti2vid_one_stage.py)** - Single-stage generation for quick prototyping * **[TI2VidOneStagePipeline](packages/ltx-pipelines/src/ltx_pipelines/ti2vid_one_stage.py)** - Single-stage generation for quick prototyping
* **[DistilledPipeline](packages/ltx-pipelines/src/ltx_pipelines/distilled.py)** - Fastest inference with 8 predefined sigmas * **[DistilledPipeline](packages/ltx-pipelines/src/ltx_pipelines/distilled.py)** - Fastest inference with 8 predefined sigmas
* **[ICLoraPipeline](packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py)** - Video-to-video and image-to-video transformations (uses distilled model.) * **[ICLoraPipeline](packages/ltx-pipelines/src/ltx_pipelines/ic_lora.py)** - Video-to-video and image-to-video transformations (uses distilled model.)
@@ -57,7 +57,10 @@ class SafetensorsModelStateDictLoader(StateDictLoader):
def metadata(self, path: str) -> dict: def metadata(self, path: str) -> dict:
with safetensors.safe_open(path, framework="pt") as f: with safetensors.safe_open(path, framework="pt") as f:
return json.loads(f.metadata()["config"]) meta = f.metadata()
if meta is None or "config" not in meta:
return {}
return json.loads(meta["config"])
def load(self, path: str | list[str], sd_ops: SDOps | None = None, device: torch.device | None = None) -> StateDict: def load(self, path: str | list[str], sd_ops: SDOps | None = None, device: torch.device | None = None) -> StateDict:
return self.weight_loader.load(path, sd_ops, device) return self.weight_loader.load(path, sd_ops, device)
@@ -128,12 +128,12 @@ def _build_caption_projections(
) -> tuple[torch.nn.Module | None, torch.nn.Module | None]: ) -> tuple[torch.nn.Module | None, torch.nn.Module | None]:
"""Build caption projections for the transformer when projection is NOT in the text encoder. """Build caption projections for the transformer when projection is NOT in the text encoder.
19B models: projection is in the transformer (caption_proj_before_connector=False). 19B models: projection is in the transformer (caption_proj_before_connector=False).
20B models: projection is in the text encoder, so no projections are created here. 22B models: projection is in the text encoder, so no projections are created here.
Args: Args:
config: Full model config dict (must contain "transformer" key). config: Full model config dict (must contain "transformer" key).
is_av: Whether this is an audio-video model. When False, audio projection is skipped. is_av: Whether this is an audio-video model. When False, audio projection is skipped.
Returns: Returns:
Tuple of (video_caption_projection, audio_caption_projection), both None for 20B models. Tuple of (video_caption_projection, audio_caption_projection), both None for 22B models.
""" """
transformer_config = config.get("transformer", {}) transformer_config = config.get("transformer", {})
if transformer_config.get("caption_proj_before_connector", False): if transformer_config.get("caption_proj_before_connector", False):
@@ -140,19 +140,20 @@ class BasicAVTransformerBlock(torch.nn.Module):
batch_size: int, batch_size: int,
scale_shift_timestep: torch.Tensor, scale_shift_timestep: torch.Tensor,
gate_timestep: torch.Tensor, gate_timestep: torch.Tensor,
scale_shift_indices: slice,
num_scale_shift_values: int = 4, num_scale_shift_values: int = 4,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
scale_shift_ada_values = self.get_ada_values( scale_shift_ada_values = self.get_ada_values(
scale_shift_table[:num_scale_shift_values, :], batch_size, scale_shift_timestep, slice(None, None) scale_shift_table[:num_scale_shift_values, :], batch_size, scale_shift_timestep, scale_shift_indices
) )
gate_ada_values = self.get_ada_values( gate_ada_values = self.get_ada_values(
scale_shift_table[num_scale_shift_values:, :], batch_size, gate_timestep, slice(None, None) scale_shift_table[num_scale_shift_values:, :], batch_size, gate_timestep, slice(None, None)
) )
scale_shift_chunks = [t.squeeze(2) for t in scale_shift_ada_values] scale, shift = (t.squeeze(2) for t in scale_shift_ada_values)
gate_ada_values = [t.squeeze(2) for t in gate_ada_values] (gate,) = (t.squeeze(2) for t in gate_ada_values)
return (*scale_shift_chunks, *gate_ada_values) return scale, shift, gate
def _apply_text_cross_attention( def _apply_text_cross_attention(
self, self,
@@ -287,37 +288,26 @@ class BasicAVTransformerBlock(torch.nn.Module):
vx_norm3 = rms_norm(vx, eps=self.norm_eps) vx_norm3 = rms_norm(vx, eps=self.norm_eps)
ax_norm3 = rms_norm(ax, eps=self.norm_eps) ax_norm3 = rms_norm(ax, eps=self.norm_eps)
( if run_a2v and not perturbations.all_in_batch(PerturbationType.SKIP_A2V_CROSS_ATTN, self.idx):
scale_ca_audio_hidden_states_a2v, scale_ca_video_a2v, shift_ca_video_a2v, gate_out_a2v = self.get_av_ca_ada_values(
shift_ca_audio_hidden_states_a2v,
scale_ca_audio_hidden_states_v2a,
shift_ca_audio_hidden_states_v2a,
gate_out_v2a,
) = self.get_av_ca_ada_values(
self.scale_shift_table_a2v_ca_audio,
ax.shape[0],
audio.cross_scale_shift_timestep,
audio.cross_gate_timestep,
)
(
scale_ca_video_hidden_states_a2v,
shift_ca_video_hidden_states_a2v,
scale_ca_video_hidden_states_v2a,
shift_ca_video_hidden_states_v2a,
gate_out_a2v,
) = self.get_av_ca_ada_values(
self.scale_shift_table_a2v_ca_video, self.scale_shift_table_a2v_ca_video,
vx.shape[0], vx.shape[0],
video.cross_scale_shift_timestep, video.cross_scale_shift_timestep,
video.cross_gate_timestep, video.cross_gate_timestep,
slice(0, 2),
) )
vx_scaled = vx_norm3 * (1 + scale_ca_video_a2v) + shift_ca_video_a2v
del scale_ca_video_a2v, shift_ca_video_a2v
if run_a2v and not perturbations.all_in_batch(PerturbationType.SKIP_A2V_CROSS_ATTN, self.idx): scale_ca_audio_a2v, shift_ca_audio_a2v, _ = self.get_av_ca_ada_values(
vx_scaled = vx_norm3 * (1 + scale_ca_video_hidden_states_a2v) + shift_ca_video_hidden_states_a2v self.scale_shift_table_a2v_ca_audio,
del scale_ca_video_hidden_states_a2v, shift_ca_video_hidden_states_a2v ax.shape[0],
ax_scaled = ax_norm3 * (1 + scale_ca_audio_hidden_states_a2v) + shift_ca_audio_hidden_states_a2v audio.cross_scale_shift_timestep,
del scale_ca_audio_hidden_states_a2v, shift_ca_audio_hidden_states_a2v audio.cross_gate_timestep,
slice(0, 2),
)
ax_scaled = ax_norm3 * (1 + scale_ca_audio_a2v) + shift_ca_audio_a2v
del scale_ca_audio_a2v, shift_ca_audio_a2v
a2v_mask = perturbations.mask_like(PerturbationType.SKIP_A2V_CROSS_ATTN, self.idx, vx) a2v_mask = perturbations.mask_like(PerturbationType.SKIP_A2V_CROSS_ATTN, self.idx, vx)
vx = vx + ( vx = vx + (
self.audio_to_video_attn( self.audio_to_video_attn(
@@ -330,11 +320,26 @@ class BasicAVTransformerBlock(torch.nn.Module):
* a2v_mask * a2v_mask
) )
del gate_out_a2v, a2v_mask, vx_scaled, ax_scaled del gate_out_a2v, a2v_mask, vx_scaled, ax_scaled
if run_v2a and not perturbations.all_in_batch(PerturbationType.SKIP_V2A_CROSS_ATTN, self.idx): if run_v2a and not perturbations.all_in_batch(PerturbationType.SKIP_V2A_CROSS_ATTN, self.idx):
ax_scaled = ax_norm3 * (1 + scale_ca_audio_hidden_states_v2a) + shift_ca_audio_hidden_states_v2a scale_ca_audio_v2a, shift_ca_audio_v2a, gate_out_v2a = self.get_av_ca_ada_values(
del scale_ca_audio_hidden_states_v2a, shift_ca_audio_hidden_states_v2a self.scale_shift_table_a2v_ca_audio,
vx_scaled = vx_norm3 * (1 + scale_ca_video_hidden_states_v2a) + shift_ca_video_hidden_states_v2a ax.shape[0],
del scale_ca_video_hidden_states_v2a, shift_ca_video_hidden_states_v2a audio.cross_scale_shift_timestep,
audio.cross_gate_timestep,
slice(2, 4),
)
ax_scaled = ax_norm3 * (1 + scale_ca_audio_v2a) + shift_ca_audio_v2a
del scale_ca_audio_v2a, shift_ca_audio_v2a
scale_ca_video_v2a, shift_ca_video_v2a, _ = self.get_av_ca_ada_values(
self.scale_shift_table_a2v_ca_video,
vx.shape[0],
video.cross_scale_shift_timestep,
video.cross_gate_timestep,
slice(2, 4),
)
vx_scaled = vx_norm3 * (1 + scale_ca_video_v2a) + shift_ca_video_v2a
del scale_ca_video_v2a, shift_ca_video_v2a
v2a_mask = perturbations.mask_like(PerturbationType.SKIP_V2A_CROSS_ATTN, self.idx, ax) v2a_mask = perturbations.mask_like(PerturbationType.SKIP_V2A_CROSS_ATTN, self.idx, ax)
ax = ax + ( ax = ax + (
self.video_to_audio_attn( self.video_to_audio_attn(
@@ -347,6 +352,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
* v2a_mask * v2a_mask
) )
del gate_out_v2a, v2a_mask, ax_scaled, vx_scaled del gate_out_v2a, v2a_mask, ax_scaled, vx_scaled
del vx_norm3, ax_norm3 del vx_norm3, ax_norm3
if run_vx: if run_vx:
@@ -1,25 +1,33 @@
"""Gemma text encoder components.""" """Gemma text encoder components."""
from ltx_core.text_encoders.gemma.embeddings_processor import (
EmbeddingsProcessor,
EmbeddingsProcessorOutput,
convert_to_additive_mask,
)
from ltx_core.text_encoders.gemma.encoders.base_encoder import ( from ltx_core.text_encoders.gemma.encoders.base_encoder import (
GemmaEncoderOutput,
GemmaTextEncoder, GemmaTextEncoder,
encode_text,
module_ops_from_gemma_root, module_ops_from_gemma_root,
) )
from ltx_core.text_encoders.gemma.encoders.encoder_configurator import ( from ltx_core.text_encoders.gemma.encoders.encoder_configurator import (
AV_GEMMA_TEXT_ENCODER_KEY_OPS, EMBEDDINGS_PROCESSOR_KEY_OPS,
GEMMA_LLM_KEY_OPS,
GEMMA_MODEL_OPS, GEMMA_MODEL_OPS,
VIDEO_ONLY_GEMMA_TEXT_ENCODER_KEY_OPS, VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS,
EmbeddingsProcessorConfigurator,
GemmaTextEncoderConfigurator, GemmaTextEncoderConfigurator,
) )
__all__ = [ __all__ = [
"AV_GEMMA_TEXT_ENCODER_KEY_OPS", "EMBEDDINGS_PROCESSOR_KEY_OPS",
"GEMMA_LLM_KEY_OPS",
"GEMMA_MODEL_OPS", "GEMMA_MODEL_OPS",
"VIDEO_ONLY_GEMMA_TEXT_ENCODER_KEY_OPS", "VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS",
"GemmaEncoderOutput", "EmbeddingsProcessor",
"EmbeddingsProcessorConfigurator",
"EmbeddingsProcessorOutput",
"GemmaTextEncoder", "GemmaTextEncoder",
"GemmaTextEncoderConfigurator", "GemmaTextEncoderConfigurator",
"encode_text", "convert_to_additive_mask",
"module_ops_from_gemma_root", "module_ops_from_gemma_root",
] ]
@@ -1,9 +1,24 @@
from typing import NamedTuple
import torch import torch
from torch import nn from torch import nn
from ltx_core.text_encoders.gemma.embeddings_connector import Embeddings1DConnector from ltx_core.text_encoders.gemma.embeddings_connector import Embeddings1DConnector
class EmbeddingsProcessorOutput(NamedTuple):
video_encoding: torch.Tensor
audio_encoding: torch.Tensor | None
attention_mask: torch.Tensor
def convert_to_additive_mask(attention_mask: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
"""Convert binary attention mask to additive form for transformer masking."""
return (attention_mask.to(torch.int64) - 1).to(dtype).reshape(
(attention_mask.shape[0], 1, -1, attention_mask.shape[-1])
) * torch.finfo(dtype).max
def _to_binary_mask(encoded: torch.Tensor, encoded_mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: def _to_binary_mask(encoded: torch.Tensor, encoded_mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Convert connector output mask to binary mask and apply to encoded tensor.""" """Convert connector output mask to binary mask and apply to encoded tensor."""
binary_mask = (encoded_mask < 0.000001).to(torch.int64) binary_mask = (encoded_mask < 0.000001).to(torch.int64)
@@ -13,12 +28,21 @@ def _to_binary_mask(encoded: torch.Tensor, encoded_mask: torch.Tensor) -> tuple[
class EmbeddingsProcessor(nn.Module): class EmbeddingsProcessor(nn.Module):
"""Wraps video connector + optional audio connector. """Wraps feature extractor + video connector + optional audio connector.
Returns (video_encoded, audio_encoded | None, binary_mask). Can operate in two modes:
1. create_embeddings(): Takes pre-computed features + additive mask (backward compat, used by trainer)
2. process_hidden_states(): Takes raw Gemma hidden states, runs feature extraction + connectors
""" """
def __init__(self, video_connector: Embeddings1DConnector, audio_connector: Embeddings1DConnector | None = None): def __init__(
self,
*,
feature_extractor: nn.Module | None = None,
video_connector: Embeddings1DConnector,
audio_connector: Embeddings1DConnector | None = None,
):
super().__init__() super().__init__()
self.feature_extractor = feature_extractor
self.video_connector = video_connector self.video_connector = video_connector
self.audio_connector = audio_connector self.audio_connector = audio_connector
@@ -41,3 +65,25 @@ class EmbeddingsProcessor(nn.Module):
audio_encoded, _ = self.audio_connector(audio_features, additive_attention_mask) audio_encoded, _ = self.audio_connector(audio_features, additive_attention_mask)
return video_encoded, audio_encoded, binary_mask.squeeze(-1) return video_encoded, audio_encoded, binary_mask.squeeze(-1)
def process_hidden_states(
self,
hidden_states: tuple[torch.Tensor, ...],
attention_mask: torch.Tensor,
padding_side: str = "left",
) -> EmbeddingsProcessorOutput:
"""Full pipeline: feature extraction -> connectors -> final embeddings.
Args:
hidden_states: Raw Gemma hidden states (tuple of tensors per layer).
attention_mask: Binary attention mask [B, seq_len].
padding_side: Padding side used during tokenization.
Returns:
EmbeddingsProcessorOutput with video_encoding, audio_encoding, and attention_mask.
"""
if self.feature_extractor is None:
raise ValueError("feature_extractor is required for process_hidden_states()")
video_feats, audio_feats = self.feature_extractor(hidden_states, attention_mask, padding_side)
additive_mask = convert_to_additive_mask(attention_mask, video_feats.dtype)
video_enc, audio_enc, binary_mask = self.create_embeddings(video_feats, audio_feats, additive_mask)
return EmbeddingsProcessorOutput(video_enc, audio_enc, binary_mask)
@@ -1,33 +1,22 @@
import functools import functools
from pathlib import Path from pathlib import Path
from typing import NamedTuple
import torch import torch
from transformers import AutoImageProcessor, Gemma3ForConditionalGeneration, Gemma3Processor from transformers import AutoImageProcessor, Gemma3ForConditionalGeneration, Gemma3Processor
from ltx_core.loader.module_ops import ModuleOps from ltx_core.loader.module_ops import ModuleOps
from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessor
from ltx_core.text_encoders.gemma.tokenizer import LTXVGemmaTokenizer from ltx_core.text_encoders.gemma.tokenizer import LTXVGemmaTokenizer
from ltx_core.utils import find_matching_file from ltx_core.utils import find_matching_file
class GemmaEncoderOutput(NamedTuple):
video_encoding: torch.Tensor
audio_encoding: torch.Tensor | None
attention_mask: torch.Tensor
class GemmaTextEncoder(torch.nn.Module): class GemmaTextEncoder(torch.nn.Module):
"""Unified Gemma text encoder with 3-block pipeline. """Pure Gemma text encoder — runs the LLM and returns raw hidden states.
Block 1: Gemma model (runs LLM, gets hidden states) Prompt enhancement (generate) is also supported since the full
Block 2: Feature extractor Gemma3ForConditionalGeneration model (including lm_head) is loaded.
Block 3: Embeddings processor (connector with optional audio)
""" """
def __init__( def __init__(
self, self,
feature_extractor: torch.nn.Module,
embeddings_processor: EmbeddingsProcessor,
model: Gemma3ForConditionalGeneration | None = None, model: Gemma3ForConditionalGeneration | None = None,
tokenizer: LTXVGemmaTokenizer | None = None, tokenizer: LTXVGemmaTokenizer | None = None,
processor: Gemma3Processor | None = None, processor: Gemma3Processor | None = None,
@@ -37,39 +26,25 @@ class GemmaTextEncoder(torch.nn.Module):
self.model = model self.model = model
self.tokenizer = tokenizer self.tokenizer = tokenizer
self.processor = processor self.processor = processor
self.feature_extractor = feature_extractor.to(dtype=dtype) self._dtype = dtype
self.embeddings_processor = embeddings_processor.to(dtype=dtype)
def _convert_to_additive_mask(self, attention_mask: torch.Tensor, dtype: torch.dtype) -> torch.Tensor: def encode(
return (attention_mask.to(torch.int64) - 1).to(dtype).reshape( self,
(attention_mask.shape[0], 1, -1, attention_mask.shape[-1]) text: str,
) * torch.finfo(dtype).max padding_side: str = "left", # noqa: ARG002
) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]:
def precompute( """Run Gemma LLM and return raw hidden states + attention mask.
self, text: str, padding_side: str = "left" Calls the inner model (self.model.model) to skip lm_head logits computation (~500 MiB saving).
) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor]: Returns:
"""Blocks 1+2: Gemma model -> feature extraction. (hidden_states, attention_mask) where hidden_states is a tuple of per-layer tensors.
Used by process_captions.py for offline precomputation.
Returns (video_features, audio_features | None, attention_mask).
""" """
# Block 1: Run Gemma
token_pairs = self.tokenizer.tokenize_with_weights(text)["gemma"] token_pairs = self.tokenizer.tokenize_with_weights(text)["gemma"]
input_ids = torch.tensor([[t[0] for t in token_pairs]], device=self.model.device) input_ids = torch.tensor([[t[0] for t in token_pairs]], device=self.model.device)
attention_mask = torch.tensor([[w[1] for w in token_pairs]], device=self.model.device) attention_mask = torch.tensor([[w[1] for w in token_pairs]], device=self.model.device)
outputs = self.model(input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=True) outputs = self.model.model(input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=True)
hidden_states = outputs.hidden_states
# Block 2: Feature extraction del outputs
video_feats, audio_feats = self.feature_extractor(outputs.hidden_states, attention_mask, padding_side) return hidden_states, attention_mask
return video_feats, audio_feats, attention_mask
def forward(self, text: str, padding_side: str = "left") -> GemmaEncoderOutput:
"""Full pipeline: precompute -> embeddings processor."""
video_feats, audio_feats, attention_mask = self.precompute(text, padding_side)
additive_mask = self._convert_to_additive_mask(attention_mask, video_feats.dtype)
video_enc, audio_enc, binary_mask = self.embeddings_processor.create_embeddings(
video_feats, audio_feats, additive_mask
)
return GemmaEncoderOutput(video_enc, audio_enc, binary_mask)
# --- Prompt enhancement methods --- # --- Prompt enhancement methods ---
@@ -225,15 +200,3 @@ def module_ops_from_gemma_root(gemma_root: str) -> tuple[ModuleOps, ...]:
mutator=load_processor, mutator=load_processor,
) )
return (tokenizer_load_ops, processor_load_ops) return (tokenizer_load_ops, processor_load_ops)
def encode_text(text_encoder: GemmaTextEncoder, prompts: list[str]) -> list[tuple[torch.Tensor, torch.Tensor]]:
"""Encode a list of prompts using the provided Gemma text encoder.
Returns:
List of tuples, each containing (v_context, a_context) tensors for each prompt.
"""
result = []
for prompt in prompts:
v_context, a_context, _ = text_encoder(prompt)
result.append((v_context, a_context))
return result
@@ -22,31 +22,32 @@ from ltx_core.text_encoders.gemma.feature_extractor import (
class GemmaTextEncoderConfigurator(ModelConfigurator[GemmaTextEncoder]): class GemmaTextEncoderConfigurator(ModelConfigurator[GemmaTextEncoder]):
@classmethod @classmethod
def from_config(cls, config: dict) -> GemmaTextEncoder: def from_config(cls, config: dict) -> GemmaTextEncoder: # noqa: ARG003
transformer_config = config.get("transformer", {})
gemma_config = Gemma3Config.from_dict(GEMMA3_CONFIG_FOR_LTX.to_dict()) gemma_config = Gemma3Config.from_dict(GEMMA3_CONFIG_FOR_LTX.to_dict())
with torch.device("meta"): with torch.device("meta"):
model = Gemma3ForConditionalGeneration(gemma_config) model = Gemma3ForConditionalGeneration(gemma_config)
return GemmaTextEncoder(model=model)
class EmbeddingsProcessorConfigurator(ModelConfigurator[EmbeddingsProcessor]):
@classmethod
def from_config(cls, config: dict) -> EmbeddingsProcessor:
transformer_config = config.get("transformer", {})
# Create video embeddings connector (always needed) # Create video embeddings connector (always needed)
video_connector = Embeddings1DConnectorConfigurator.from_config(config) video_connector = Embeddings1DConnectorConfigurator.from_config(config)
# Create audio embeddings connector # Create audio embeddings connector
audio_connector = AudioEmbeddings1DConnectorConfigurator.from_config(config) audio_connector = AudioEmbeddings1DConnectorConfigurator.from_config(config)
# Create embeddings processor with both connectors # Create feature extractor
embeddings_processor = EmbeddingsProcessor(
video_connector=video_connector,
audio_connector=audio_connector,
)
feature_extractor = _create_feature_extractor(transformer_config) feature_extractor = _create_feature_extractor(transformer_config)
return GemmaTextEncoder( return EmbeddingsProcessor(
video_connector=video_connector,
audio_connector=audio_connector,
feature_extractor=feature_extractor, feature_extractor=feature_extractor,
embeddings_processor=embeddings_processor,
model=model,
) )
@@ -97,8 +98,31 @@ def _create_feature_extractor(transformer_config: dict) -> torch.nn.Module:
) )
AV_GEMMA_TEXT_ENCODER_KEY_OPS = ( # --- Split SDOps: Gemma LLM keys vs Embeddings Processor keys ---
SDOps("AV_GEMMA_TEXT_ENCODER_KEY_OPS")
GEMMA_LLM_KEY_OPS = (
SDOps("GEMMA_LLM_KEY_OPS")
# 1. Map language model layers (note the double .model prefix)
.with_matching(prefix="language_model.model.")
.with_replacement("language_model.model.", "model.model.language_model.")
# 2. Map the Vision Tower
.with_matching(prefix="vision_tower.")
.with_replacement("vision_tower.", "model.model.vision_tower.")
# 3. Map the Multi-Modal Projector
.with_matching(prefix="multi_modal_projector.")
.with_replacement("multi_modal_projector.", "model.model.multi_modal_projector.")
# 4. Duplicate embed_tokens to lm_head (needed for prompt enhancement via generate())
.with_kv_operation(
operation=lambda key, value: [
KeyValueOperationResult(key, value),
KeyValueOperationResult("model.lm_head.weight", value),
],
key_prefix="model.model.language_model.embed_tokens.weight",
)
)
EMBEDDINGS_PROCESSOR_KEY_OPS = (
SDOps("EMBEDDINGS_PROCESSOR_KEY_OPS")
# 1. Map the feature extractor (V1: aggregate_embed inside feature_extractor) # 1. Map the feature extractor (V1: aggregate_embed inside feature_extractor)
.with_matching(prefix="text_embedding_projection.aggregate_embed.") .with_matching(prefix="text_embedding_projection.aggregate_embed.")
.with_replacement("text_embedding_projection.aggregate_embed.", "feature_extractor.aggregate_embed.") .with_replacement("text_embedding_projection.aggregate_embed.", "feature_extractor.aggregate_embed.")
@@ -109,30 +133,13 @@ AV_GEMMA_TEXT_ENCODER_KEY_OPS = (
.with_replacement("text_embedding_projection.audio_aggregate_embed.", "feature_extractor.audio_aggregate_embed.") .with_replacement("text_embedding_projection.audio_aggregate_embed.", "feature_extractor.audio_aggregate_embed.")
# 2. Map the connectors # 2. Map the connectors
.with_matching(prefix="model.diffusion_model.video_embeddings_connector.") .with_matching(prefix="model.diffusion_model.video_embeddings_connector.")
.with_replacement("model.diffusion_model.video_embeddings_connector.", "embeddings_processor.video_connector.") .with_replacement("model.diffusion_model.video_embeddings_connector.", "video_connector.")
.with_matching(prefix="model.diffusion_model.audio_embeddings_connector.") .with_matching(prefix="model.diffusion_model.audio_embeddings_connector.")
.with_replacement("model.diffusion_model.audio_embeddings_connector.", "embeddings_processor.audio_connector.") .with_replacement("model.diffusion_model.audio_embeddings_connector.", "audio_connector.")
# 3. Map language model layers (note the double .model prefix)
.with_matching(prefix="language_model.model.")
.with_replacement("language_model.model.", "model.model.language_model.")
# 4. Map the Vision Tower
.with_matching(prefix="vision_tower.")
.with_replacement("vision_tower.", "model.model.vision_tower.")
# 5. Map the Multi-Modal Projector
.with_matching(prefix="multi_modal_projector.")
.with_replacement("multi_modal_projector.", "model.model.multi_modal_projector.")
.with_kv_operation(
operation=lambda key, value: [
KeyValueOperationResult(key, value),
KeyValueOperationResult("model.lm_head.weight", value),
],
key_prefix="model.model.language_model.embed_tokens.weight",
)
) )
VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS = (
VIDEO_ONLY_GEMMA_TEXT_ENCODER_KEY_OPS = ( SDOps("VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS")
SDOps("VIDEO_ONLY_GEMMA_TEXT_ENCODER_KEY_OPS")
# 1. Map the feature extractor (V1: aggregate_embed inside feature_extractor) # 1. Map the feature extractor (V1: aggregate_embed inside feature_extractor)
.with_matching(prefix="text_embedding_projection.aggregate_embed.") .with_matching(prefix="text_embedding_projection.aggregate_embed.")
.with_replacement("text_embedding_projection.aggregate_embed.", "feature_extractor.aggregate_embed.") .with_replacement("text_embedding_projection.aggregate_embed.", "feature_extractor.aggregate_embed.")
@@ -110,7 +110,7 @@ class FeatureExtractorV1(nn.Module):
class FeatureExtractorV2(nn.Module): class FeatureExtractorV2(nn.Module):
"""20B: per-token RMS norm → rescale → dual aggregate embeds""" """22B: per-token RMS norm → rescale → dual aggregate embeds"""
def __init__( def __init__(
self, self,
+5 -5
View File
@@ -56,7 +56,7 @@ python -m ltx_pipelines.ti2vid_two_stages --help
Available pipeline modules: Available pipeline modules:
- `ltx_pipelines.ti2vid_two_stages` - Two-stage text/image-to-video (recommended). - `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_two_stages_hq` - Two-stage text/image-to-video (different sampler, better quality).
- `ltx_pipelines.ti2vid_one_stage` - Single-stage text/image-to-video. - `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.distilled` - Fast text/image-to-video pipeline using only the distilled model.
- `ltx_pipelines.ic_lora` - Video-to-video with IC-LoRA. - `ltx_pipelines.ic_lora` - Video-to-video with IC-LoRA.
@@ -94,14 +94,14 @@ Do you need to condition on existing images/videos?
└─ YES → Use DistilledPipeline (with 8 predefined sigmas) └─ 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), [`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). > **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), [`TI2VidTwoStagesHQPipeline`](src/ltx_pipelines/ti2vid_two_stages_hq.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 ### Features Comparison
| Pipeline | Stages | [Multimodal Guidance](#%EF%B8%8F-multimodal-guidance) | Upsampling | Conditioning | Best For | | Pipeline | Stages | [Multimodal Guidance](#%EF%B8%8F-multimodal-guidance) | Upsampling | Conditioning | Best For |
| -------- | ------ | --- | ---------- | ------------- | -------- | | -------- | ------ | --- | ---------- | ------------- | -------- |
| **TI2VidTwoStagesPipeline** | 2 | ✅ | ✅ | Image | **Production quality** (recommended) | | **TI2VidTwoStagesPipeline** | 2 | ✅ | ✅ | Image | **Production quality** (recommended) |
| **TI2VidTwoStagesRes2sPipeline** | 2 | ✅ | ✅ | Image | Same as above, res_2s sampler (fewer steps) | | **TI2VidTwoStagesHQPipeline** | 2 | ✅ | ✅ | Image | Same as above, res_2s sampler (higher quality) |
| **TI2VidOneStagePipeline** | 1 | ✅ | ❌ | Image | Educational, prototyping | | **TI2VidOneStagePipeline** | 1 | ✅ | ❌ | Image | Educational, prototyping |
| **DistilledPipeline** | 2 | ❌ | ✅ | Image | Fastest inference (8 sigmas) | | **DistilledPipeline** | 2 | ❌ | ✅ | Image | Fastest inference (8 sigmas) |
| **ICLoraPipeline** | 2 | ✅ | ✅ | Image + Video | Video-to-video transformations | | **ICLoraPipeline** | 2 | ✅ | ✅ | Image + Video | Video-to-video transformations |
@@ -125,11 +125,11 @@ Two-stage generation: Stage 1 generates low-resolution video with [multimodal gu
--- ---
### 2. TI2VidTwoStagesRes2sPipeline ### 2. TI2VidTwoStagesHQPipeline
**Best for:** Same two-stage text/image-to-video as TI2VidTwoStagesPipeline but with a different sampler and step count. **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) **Source**: [`src/ltx_pipelines/ti2vid_two_stages_hq.py`](src/ltx_pipelines/ti2vid_two_stages_hq.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. 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.
@@ -14,7 +14,6 @@ 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 TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy 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_core.types import Audio, AudioLatentShape, LatentState, VideoPixelShape
from ltx_pipelines.utils import ModelLedger from ltx_pipelines.utils import ModelLedger
from ltx_pipelines.utils.args import default_2_stage_arg_parser from ltx_pipelines.utils.args import default_2_stage_arg_parser
@@ -24,10 +23,10 @@ from ltx_pipelines.utils.constants import (
from ltx_pipelines.utils.helpers import ( from ltx_pipelines.utils.helpers import (
assert_resolution, assert_resolution,
cleanup_memory, cleanup_memory,
combined_image_conditionings,
denoise_video_only, denoise_video_only,
generate_enhanced_prompt, encode_prompts,
get_device, get_device,
image_conditionings_by_replacing_latent,
multi_modal_guider_denoising_func, multi_modal_guider_denoising_func,
simple_denoising_func, simple_denoising_func,
) )
@@ -69,7 +68,7 @@ class A2VidPipelineTwoStage:
quantization=quantization, quantization=quantization,
) )
self.stage_2_model_ledger = self.stage_1_model_ledger.with_loras( self.stage_2_model_ledger = self.stage_1_model_ledger.with_additional_loras(
loras=distilled_lora, loras=distilled_lora,
) )
@@ -103,16 +102,14 @@ class A2VidPipelineTwoStage:
stepper = EulerDiffusionStep() stepper = EulerDiffusionStep()
dtype = torch.bfloat16 dtype = torch.bfloat16
text_encoder = self.stage_1_model_ledger.text_encoder() ctx_p, ctx_n = encode_prompts(
if enhance_prompt: [prompt, negative_prompt],
prompt = generate_enhanced_prompt(text_encoder, prompt, images[0][0] if len(images) > 0 else None) self.stage_1_model_ledger,
context_p, context_n = encode_text(text_encoder, prompts=[prompt, negative_prompt]) enhance_first_prompt=enhance_prompt,
v_context_p, a_context_p = context_p enhance_prompt_image=images[0][0] if len(images) > 0 else None,
v_context_n, _ = context_n )
v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
torch.cuda.synchronize() v_context_n, _ = ctx_n.video_encoding, ctx_n.audio_encoding
del text_encoder
cleanup_memory()
# Encode audio. # Encode audio.
decoded_audio = decode_audio_from_file(audio_path, self.device, audio_start_time, audio_max_duration) decoded_audio = decode_audio_from_file(audio_path, self.device, audio_start_time, audio_max_duration)
@@ -120,8 +117,28 @@ class A2VidPipelineTwoStage:
audio_shape = AudioLatentShape.from_duration(batch=1, duration=num_frames / frame_rate, channels=8, mel_bins=16) 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] encoded_audio_latent = encoded_audio_latent[:, :, : audio_shape.frames]
# Stage 1: encode image conditionings with the VAE encoder, then free it
# before loading the transformer to reduce peak VRAM.
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 = combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
torch.cuda.synchronize()
del video_encoder
cleanup_memory() cleanup_memory()
# Stage 1: Initial low resolution video generation with audio conditioning.
transformer = self.stage_1_model_ledger.transformer() transformer = self.stage_1_model_ledger.transformer()
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device) sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
@@ -147,24 +164,6 @@ class A2VidPipelineTwoStage:
), ),
) )
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( video_state = denoise_video_only(
output_shape=stage_1_output_shape, output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings, conditionings=stage_1_conditionings,
@@ -183,12 +182,23 @@ class A2VidPipelineTwoStage:
cleanup_memory() cleanup_memory()
# Stage 2: Upsample and refine the video at higher resolution with distilled LoRA. # Stage 2: Upsample and refine the video at higher resolution with distilled LoRA.
video_encoder = self.stage_1_model_ledger.video_encoder()
upscaled_video_latent = upsample_video( upscaled_video_latent = upsample_video(
latent=video_state.latent[:1], latent=video_state.latent[:1],
video_encoder=video_encoder, video_encoder=video_encoder,
upsampler=self.stage_2_model_ledger.spatial_upsampler(), upsampler=self.stage_2_model_ledger.spatial_upsampler(),
) )
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = combined_image_conditionings(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
del video_encoder
torch.cuda.synchronize() torch.cuda.synchronize()
cleanup_memory() cleanup_memory()
@@ -210,15 +220,6 @@ class A2VidPipelineTwoStage:
), ),
) )
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( video_state = denoise_video_only(
output_shape=stage_2_output_shape, output_shape=stage_2_output_shape,
conditionings=stage_2_conditionings, conditionings=stage_2_conditionings,
@@ -236,7 +237,6 @@ class A2VidPipelineTwoStage:
torch.cuda.synchronize() torch.cuda.synchronize()
del transformer del transformer
del video_encoder
cleanup_memory() cleanup_memory()
decoded_video = vae_decode_video( decoded_video = vae_decode_video(
@@ -278,7 +278,7 @@ def main() -> None:
distilled_lora=args.distilled_lora, distilled_lora=args.distilled_lora,
spatial_upsampler_path=args.spatial_upsampler_path, spatial_upsampler_path=args.spatial_upsampler_path,
gemma_root=args.gemma_root, gemma_root=args.gemma_root,
loras=args.lora, loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization, quantization=args.quantization,
) )
tiling_config = TilingConfig.default() tiling_config = TilingConfig.default()
@@ -12,7 +12,6 @@ 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 TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy from ltx_core.quantization import QuantizationPolicy
from ltx_core.text_encoders.gemma import encode_text
from ltx_core.types import Audio, LatentState, VideoPixelShape from ltx_core.types import Audio, LatentState, VideoPixelShape
from ltx_pipelines.utils import ModelLedger, euler_denoising_loop from ltx_pipelines.utils import ModelLedger, euler_denoising_loop
from ltx_pipelines.utils.args import ( from ltx_pipelines.utils.args import (
@@ -28,10 +27,10 @@ from ltx_pipelines.utils.constants import (
from ltx_pipelines.utils.helpers import ( from ltx_pipelines.utils.helpers import (
assert_resolution, assert_resolution,
cleanup_memory, cleanup_memory,
combined_image_conditionings,
denoise_audio_video, denoise_audio_video,
generate_enhanced_prompt, encode_prompts,
get_device, get_device,
image_conditionings_by_replacing_latent,
simple_denoising_func, simple_denoising_func,
) )
from ltx_pipelines.utils.media_io import encode_video from ltx_pipelines.utils.media_io import encode_video
@@ -93,15 +92,13 @@ class DistilledPipeline:
stepper = EulerDiffusionStep() stepper = EulerDiffusionStep()
dtype = torch.bfloat16 dtype = torch.bfloat16
text_encoder = self.model_ledger.text_encoder() (ctx_p,) = encode_prompts(
if enhance_prompt: [prompt],
prompt = generate_enhanced_prompt(text_encoder, prompt, images[0][0] if len(images) > 0 else None) self.model_ledger,
context_p = encode_text(text_encoder, prompts=[prompt])[0] enhance_first_prompt=enhance_prompt,
video_context, audio_context = context_p enhance_prompt_image=images[0][0] if len(images) > 0 else None,
)
torch.cuda.synchronize() video_context, audio_context = ctx_p.video_encoding, ctx_p.audio_encoding
del text_encoder
cleanup_memory()
# Stage 1: Initial low resolution video generation. # Stage 1: Initial low resolution video generation.
video_encoder = self.model_ledger.video_encoder() video_encoder = self.model_ledger.video_encoder()
@@ -130,7 +127,7 @@ class DistilledPipeline:
height=height // 2, height=height // 2,
fps=frame_rate, fps=frame_rate,
) )
stage_1_conditionings = image_conditionings_by_replacing_latent( stage_1_conditionings = combined_image_conditionings(
images=images, images=images,
height=stage_1_output_shape.height, height=stage_1_output_shape.height,
width=stage_1_output_shape.width, width=stage_1_output_shape.width,
@@ -161,7 +158,7 @@ class DistilledPipeline:
stage_2_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device) stage_2_sigmas = torch.Tensor(STAGE_2_DISTILLED_SIGMA_VALUES).to(self.device)
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate) stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = image_conditionings_by_replacing_latent( stage_2_conditionings = combined_image_conditionings(
images=images, images=images,
height=stage_2_output_shape.height, height=stage_2_output_shape.height,
width=stage_2_output_shape.width, width=stage_2_output_shape.width,
@@ -209,7 +206,7 @@ def main() -> None:
distilled_checkpoint_path=args.distilled_checkpoint_path, distilled_checkpoint_path=args.distilled_checkpoint_path,
spatial_upsampler_path=args.spatial_upsampler_path, spatial_upsampler_path=args.spatial_upsampler_path,
gemma_root=args.gemma_root, gemma_root=args.gemma_root,
loras=args.lora, loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization, quantization=args.quantization,
) )
tiling_config = TilingConfig.default() tiling_config = TilingConfig.default()
@@ -19,17 +19,16 @@ from ltx_core.model.upsampler import upsample_video
from ltx_core.model.video_vae import TilingConfig, VideoEncoder, get_video_chunks_number from ltx_core.model.video_vae import TilingConfig, VideoEncoder, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy from ltx_core.quantization import QuantizationPolicy
from ltx_core.text_encoders.gemma import encode_text
from ltx_core.types import Audio, LatentState, VideoLatentShape, VideoPixelShape from ltx_core.types import Audio, LatentState, VideoLatentShape, VideoPixelShape
from ltx_pipelines.utils import ( from ltx_pipelines.utils import (
ModelLedger, ModelLedger,
assert_resolution, assert_resolution,
cleanup_memory, cleanup_memory,
combined_image_conditionings,
denoise_audio_video, denoise_audio_video,
encode_prompts,
euler_denoising_loop, euler_denoising_loop,
generate_enhanced_prompt,
get_device, get_device,
image_conditionings_by_replacing_latent,
simple_denoising_func, simple_denoising_func,
) )
from ltx_pipelines.utils.args import ( from ltx_pipelines.utils.args import (
@@ -169,20 +168,37 @@ class ICLoraPipeline:
stepper = EulerDiffusionStep() stepper = EulerDiffusionStep()
dtype = torch.bfloat16 dtype = torch.bfloat16
text_encoder = self.stage_1_model_ledger.text_encoder() (ctx_p,) = encode_prompts(
[prompt],
if enhance_prompt: self.stage_1_model_ledger,
prompt = generate_enhanced_prompt( enhance_first_prompt=enhance_prompt,
text_encoder, prompt, images[0][0] if len(images) > 0 else None, seed=seed enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
) )
video_context, audio_context = encode_text(text_encoder, prompts=[prompt])[0] video_context, audio_context = ctx_p.video_encoding, ctx_p.audio_encoding
torch.cuda.synchronize()
del text_encoder
cleanup_memory()
# Stage 1: Initial low resolution video generation. # Stage 1: Initial low resolution video generation.
stage_1_output_shape = VideoPixelShape(
batch=1,
frames=num_frames,
width=width // 2,
height=height // 2,
fps=frame_rate,
)
# Encode conditionings before loading transformer to reduce peak VRAM
video_encoder = self.stage_1_model_ledger.video_encoder() video_encoder = self.stage_1_model_ledger.video_encoder()
stage_1_conditionings = self._create_conditionings(
images=images,
video_conditioning=video_conditioning,
height=stage_1_output_shape.height,
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,
)
transformer = self.stage_1_model_ledger.transformer() transformer = self.stage_1_model_ledger.transformer()
stage_1_sigmas = torch.Tensor(DISTILLED_SIGMA_VALUES).to(self.device) stage_1_sigmas = torch.Tensor(DISTILLED_SIGMA_VALUES).to(self.device)
@@ -201,25 +217,6 @@ class ICLoraPipeline:
), ),
) )
stage_1_output_shape = VideoPixelShape(
batch=1,
frames=num_frames,
width=width // 2,
height=height // 2,
fps=frame_rate,
)
stage_1_conditionings = self._create_conditionings(
images=images,
video_conditioning=video_conditioning,
height=stage_1_output_shape.height,
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( video_state, audio_state = denoise_audio_video(
output_shape=stage_1_output_shape, output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings, conditionings=stage_1_conditionings,
@@ -278,7 +275,7 @@ class ICLoraPipeline:
) )
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate) stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = image_conditionings_by_replacing_latent( stage_2_conditionings = combined_image_conditionings(
images=images, images=images,
height=stage_2_output_shape.height, height=stage_2_output_shape.height,
width=stage_2_output_shape.width, width=stage_2_output_shape.width,
@@ -340,7 +337,7 @@ class ICLoraPipeline:
Returns: Returns:
List of conditioning items. IC-LoRA conditionings are appended last. List of conditioning items. IC-LoRA conditionings are appended last.
""" """
conditionings = image_conditionings_by_replacing_latent( conditionings = combined_image_conditionings(
images=images, images=images,
height=height, height=height,
width=width, width=width,
@@ -510,7 +507,7 @@ def main() -> None:
distilled_checkpoint_path=args.distilled_checkpoint_path, distilled_checkpoint_path=args.distilled_checkpoint_path,
spatial_upsampler_path=args.spatial_upsampler_path, spatial_upsampler_path=args.spatial_upsampler_path,
gemma_root=args.gemma_root, gemma_root=args.gemma_root,
loras=args.lora, loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization, quantization=args.quantization,
) )
tiling_config = TilingConfig.default() tiling_config = TilingConfig.default()
@@ -18,7 +18,6 @@ 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 TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy from ltx_core.quantization import QuantizationPolicy
from ltx_core.text_encoders.gemma import encode_text
from ltx_core.types import Audio, LatentState, VideoPixelShape from ltx_core.types import Audio, LatentState, VideoPixelShape
from ltx_pipelines.utils import ModelLedger from ltx_pipelines.utils import ModelLedger
from ltx_pipelines.utils.args import ImageConditioningInput, default_2_stage_arg_parser, detect_checkpoint_path from ltx_pipelines.utils.args import ImageConditioningInput, default_2_stage_arg_parser, detect_checkpoint_path
@@ -27,7 +26,7 @@ from ltx_pipelines.utils.helpers import (
assert_resolution, assert_resolution,
cleanup_memory, cleanup_memory,
denoise_audio_video, denoise_audio_video,
generate_enhanced_prompt, encode_prompts,
get_device, get_device,
image_conditionings_by_adding_guiding_latent, image_conditionings_by_adding_guiding_latent,
multi_modal_guider_factory_denoising_func, multi_modal_guider_factory_denoising_func,
@@ -71,7 +70,7 @@ class KeyframeInterpolationPipeline:
loras=loras, loras=loras,
quantization=quantization, quantization=quantization,
) )
self.stage_2_model_ledger = self.stage_1_model_ledger.with_loras( self.stage_2_model_ledger = self.stage_1_model_ledger.with_additional_loras(
loras=distilled_lora, loras=distilled_lora,
) )
self.pipeline_components = PipelineComponents( self.pipeline_components = PipelineComponents(
@@ -102,18 +101,15 @@ class KeyframeInterpolationPipeline:
stepper = EulerDiffusionStep() stepper = EulerDiffusionStep()
dtype = torch.bfloat16 dtype = torch.bfloat16
text_encoder = self.stage_1_model_ledger.text_encoder() ctx_p, ctx_n = encode_prompts(
if enhance_prompt: [prompt, negative_prompt],
prompt = generate_enhanced_prompt( self.stage_1_model_ledger,
text_encoder, prompt, images[0][0] if len(images) > 0 else None, seed=seed enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
) )
context_p, context_n = encode_text(text_encoder, prompts=[prompt, negative_prompt]) v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_p, a_context_p = context_p v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
v_context_n, a_context_n = context_n
torch.cuda.synchronize()
del text_encoder
cleanup_memory()
# Stage 1: Initial low resolution video generation. # Stage 1: Initial low resolution video generation.
video_encoder = self.stage_1_model_ledger.video_encoder() video_encoder = self.stage_1_model_ledger.video_encoder()
@@ -252,7 +248,7 @@ def main() -> None:
distilled_lora=args.distilled_lora, distilled_lora=args.distilled_lora,
spatial_upsampler_path=args.spatial_upsampler_path, spatial_upsampler_path=args.spatial_upsampler_path,
gemma_root=args.gemma_root, gemma_root=args.gemma_root,
loras=args.lora, loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization, quantization=args.quantization,
) )
tiling_config = TilingConfig.default() tiling_config = TilingConfig.default()
@@ -20,7 +20,6 @@ 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 TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy from ltx_core.quantization import QuantizationPolicy
from ltx_core.text_encoders.gemma import encode_text
from ltx_core.tools import LatentTools from ltx_core.tools import LatentTools
from ltx_core.types import ( from ltx_core.types import (
Audio, Audio,
@@ -30,10 +29,11 @@ from ltx_core.types import (
VideoPixelShape, VideoPixelShape,
) )
from ltx_pipelines.utils import ModelLedger from ltx_pipelines.utils import ModelLedger
from ltx_pipelines.utils.constants import DISTILLED_SIGMA_VALUES from ltx_pipelines.utils.args import QuantizationAction
from ltx_pipelines.utils.constants import DISTILLED_SIGMA_VALUES, detect_params
from ltx_pipelines.utils.helpers import ( from ltx_pipelines.utils.helpers import (
cleanup_memory, cleanup_memory,
generate_enhanced_prompt, encode_prompts,
get_device, get_device,
multi_modal_guider_denoising_func, multi_modal_guider_denoising_func,
noise_audio_state, noise_audio_state,
@@ -197,7 +197,6 @@ class RetakePipeline:
# Public entry point # # Public entry point #
# --------------------------------------------------------------------- # # --------------------------------------------------------------------- #
@torch.inference_mode()
def __call__( # noqa: PLR0913, PLR0915 def __call__( # noqa: PLR0913, PLR0915
self, self,
video_path: str, video_path: str,
@@ -214,6 +213,7 @@ class RetakePipeline:
regenerate_audio: bool = True, regenerate_audio: bool = True,
enhance_prompt: bool = False, enhance_prompt: bool = False,
distilled: bool = False, distilled: bool = False,
tiling_config: TilingConfig | None = None,
) -> tuple[Iterator[torch.Tensor], torch.Tensor]: ) -> tuple[Iterator[torch.Tensor], torch.Tensor]:
"""Regenerate ``[start_time, end_time]`` of the source video (retake). """Regenerate ``[start_time, end_time]`` of the source video (retake).
Parameters Parameters
@@ -321,22 +321,17 @@ class RetakePipeline:
del audio_encoder del audio_encoder
cleanup_memory() cleanup_memory()
text_encoder = self.model_ledger.text_encoder() prompts_to_encode = [prompt] if distilled else [prompt, negative_prompt]
if enhance_prompt: contexts = encode_prompts(
prompt = generate_enhanced_prompt(text_encoder, prompt, None, seed=effective_seed) prompts_to_encode,
self.model_ledger,
enhance_first_prompt=enhance_prompt,
enhance_prompt_seed=effective_seed,
)
if distilled: v_context_p, a_context_p = contexts[0].video_encoding, contexts[0].audio_encoding
# Distilled mode: single prompt, no negative if not distilled:
context_p = encode_text(text_encoder, prompts=[prompt])[0] v_context_n, a_context_n = contexts[1].video_encoding, contexts[1].audio_encoding
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() transformer = self.model_ledger.transformer()
@@ -412,7 +407,9 @@ class RetakePipeline:
del transformer del transformer
cleanup_memory() cleanup_memory()
decoded_video = vae_decode_video(video_state.latent, self.model_ledger.video_decoder(), generator=generator) decoded_video = vae_decode_video(
video_state.latent, self.model_ledger.video_decoder(), tiling_config, generator
)
decoded_audio = vae_decode_audio( decoded_audio = vae_decode_audio(
audio_state.latent, self.model_ledger.audio_decoder(), self.model_ledger.vocoder() audio_state.latent, self.model_ledger.audio_decoder(), self.model_ledger.vocoder()
) )
@@ -420,6 +417,7 @@ class RetakePipeline:
return decoded_video, decoded_audio return decoded_video, decoded_audio
@torch.inference_mode()
def main() -> None: def main() -> None:
"""CLI entry point for retake (regenerate a time region).""" """CLI entry point for retake (regenerate a time region)."""
logging.getLogger().setLevel(logging.INFO) logging.getLogger().setLevel(logging.INFO)
@@ -433,6 +431,15 @@ def main() -> None:
parser.add_argument("--gemma-root", type=str, required=True, help="Path to Gemma text encoder weights.") 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("--seed", type=int, default=42, help="Random seed. Use -1 for a random seed.")
parser.add_argument("--loras", nargs="*", default=[], help="LoRA paths (optional).") parser.add_argument("--loras", nargs="*", default=[], help="LoRA paths (optional).")
parser.add_argument(
"--quantization",
dest="quantization",
action=QuantizationAction,
nargs="+",
metavar=("POLICY", "AMAX_PATH"),
default=None,
help="Quantization policy: fp8-cast or fp8-scaled-mm [AMAX_PATH].",
)
args = parser.parse_args() args = parser.parse_args()
if args.start_time >= args.end_time: if args.start_time >= args.end_time:
@@ -452,16 +459,21 @@ def main() -> None:
pipeline = RetakePipeline( pipeline = RetakePipeline(
checkpoint_path=args.checkpoint_path, checkpoint_path=args.checkpoint_path,
gemma_root=args.gemma_root, gemma_root=args.gemma_root,
loras=args.loras or [], loras=tuple(args.loras) if args.loras else (),
quantization=args.quantization,
) )
params = detect_params(args.checkpoint_path)
tiling_config = TilingConfig.default()
video_iter, audio = pipeline( video_iter, audio = pipeline(
video_path=args.video_path, video_path=args.video_path,
prompt=args.prompt, prompt=args.prompt,
start_time=args.start_time, start_time=args.start_time,
end_time=args.end_time, end_time=args.end_time,
seed=args.seed, seed=args.seed,
video_guider_params=params.video_guider_params,
audio_guider_params=params.audio_guider_params,
tiling_config=tiling_config,
) )
tiling_config = TilingConfig.default()
video_chunks_number = get_video_chunks_number(num_frames, tiling_config) video_chunks_number = get_video_chunks_number(num_frames, tiling_config)
encode_video( encode_video(
video=video_iter, video=video_iter,
@@ -16,17 +16,16 @@ 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 decode_audio as vae_decode_audio
from ltx_core.model.video_vae import decode_video as vae_decode_video from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy from ltx_core.quantization import QuantizationPolicy
from ltx_core.text_encoders.gemma import encode_text
from ltx_core.types import Audio, LatentState, VideoPixelShape from ltx_core.types import Audio, LatentState, VideoPixelShape
from ltx_pipelines.utils import ( from ltx_pipelines.utils import (
ModelLedger, ModelLedger,
assert_resolution, assert_resolution,
cleanup_memory, cleanup_memory,
combined_image_conditionings,
denoise_audio_video, denoise_audio_video,
encode_prompts,
euler_denoising_loop, euler_denoising_loop,
generate_enhanced_prompt,
get_device, get_device,
image_conditionings_by_replacing_latent,
multi_modal_guider_factory_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.args import ImageConditioningInput, default_1_stage_arg_parser, detect_checkpoint_path
@@ -91,21 +90,32 @@ class TI2VidOneStagePipeline:
stepper = EulerDiffusionStep() stepper = EulerDiffusionStep()
dtype = torch.bfloat16 dtype = torch.bfloat16
text_encoder = self.model_ledger.text_encoder() ctx_p, ctx_n = encode_prompts(
if enhance_prompt: [prompt, negative_prompt],
prompt = generate_enhanced_prompt( self.model_ledger,
text_encoder, prompt, images[0][0] if len(images) > 0 else None, seed=seed enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
) )
context_p, context_n = encode_text(text_encoder, prompts=[prompt, negative_prompt]) v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_p, a_context_p = context_p v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
v_context_n, a_context_n = context_n
# Encode image conditionings with the VAE encoder, then free it
# before loading the transformer to reduce peak VRAM.
stage_1_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
video_encoder = self.model_ledger.video_encoder()
stage_1_conditionings = combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
torch.cuda.synchronize() torch.cuda.synchronize()
del text_encoder del video_encoder
cleanup_memory() cleanup_memory()
# Stage 1: Initial low resolution video generation.
video_encoder = self.model_ledger.video_encoder()
transformer = self.model_ledger.transformer() transformer = self.model_ledger.transformer()
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device) sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
@@ -135,16 +145,6 @@ class TI2VidOneStagePipeline:
), ),
) )
stage_1_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
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( video_state, audio_state = denoise_audio_video(
output_shape=stage_1_output_shape, output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings, conditionings=stage_1_conditionings,
@@ -178,7 +178,7 @@ def main() -> None:
pipeline = TI2VidOneStagePipeline( pipeline = TI2VidOneStagePipeline(
checkpoint_path=args.checkpoint_path, checkpoint_path=args.checkpoint_path,
gemma_root=args.gemma_root, gemma_root=args.gemma_root,
loras=args.lora, loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization, quantization=args.quantization,
) )
video, audio = pipeline( video, audio = pipeline(
@@ -18,17 +18,16 @@ 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 TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy from ltx_core.quantization import QuantizationPolicy
from ltx_core.text_encoders.gemma import encode_text
from ltx_core.types import Audio, LatentState, VideoPixelShape from ltx_core.types import Audio, LatentState, VideoPixelShape
from ltx_pipelines.utils import ( from ltx_pipelines.utils import (
ModelLedger, ModelLedger,
assert_resolution, assert_resolution,
cleanup_memory, cleanup_memory,
combined_image_conditionings,
denoise_audio_video, denoise_audio_video,
encode_prompts,
euler_denoising_loop, euler_denoising_loop,
generate_enhanced_prompt,
get_device, get_device,
image_conditionings_by_replacing_latent,
multi_modal_guider_factory_denoising_func, multi_modal_guider_factory_denoising_func,
simple_denoising_func, simple_denoising_func,
) )
@@ -71,7 +70,7 @@ class TI2VidTwoStagesPipeline:
quantization=quantization, quantization=quantization,
) )
self.stage_2_model_ledger = self.stage_1_model_ledger.with_loras( self.stage_2_model_ledger = self.stage_1_model_ledger.with_additional_loras(
loras=distilled_lora, loras=distilled_lora,
) )
@@ -103,21 +102,38 @@ class TI2VidTwoStagesPipeline:
stepper = EulerDiffusionStep() stepper = EulerDiffusionStep()
dtype = torch.bfloat16 dtype = torch.bfloat16
text_encoder = self.stage_1_model_ledger.text_encoder() ctx_p, ctx_n = encode_prompts(
if enhance_prompt: [prompt, negative_prompt],
prompt = generate_enhanced_prompt( self.stage_1_model_ledger,
text_encoder, prompt, images[0][0] if len(images) > 0 else None, seed=seed enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
) )
context_p, context_n = encode_text(text_encoder, prompts=[prompt, negative_prompt]) v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_p, a_context_p = context_p v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
v_context_n, a_context_n = context_n
# Stage 1: encode image conditionings with the VAE encoder, then free it
# before loading the transformer to reduce peak VRAM.
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 = combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
torch.cuda.synchronize() torch.cuda.synchronize()
del text_encoder del video_encoder
cleanup_memory() 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() transformer = self.stage_1_model_ledger.transformer()
sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device) sigmas = LTX2Scheduler().execute(steps=num_inference_steps).to(dtype=torch.float32, device=self.device)
@@ -144,21 +160,6 @@ class TI2VidTwoStagesPipeline:
), ),
) )
stage_1_output_shape = VideoPixelShape(
batch=1,
frames=num_frames,
width=width // 2,
height=height // 2,
fps=frame_rate,
)
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( video_state, audio_state = denoise_audio_video(
output_shape=stage_1_output_shape, output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings, conditionings=stage_1_conditionings,
@@ -176,12 +177,23 @@ class TI2VidTwoStagesPipeline:
cleanup_memory() cleanup_memory()
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA. # Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
video_encoder = self.stage_1_model_ledger.video_encoder()
upscaled_video_latent = upsample_video( upscaled_video_latent = upsample_video(
latent=video_state.latent[:1], latent=video_state.latent[:1],
video_encoder=video_encoder, video_encoder=video_encoder,
upsampler=self.stage_2_model_ledger.spatial_upsampler(), upsampler=self.stage_2_model_ledger.spatial_upsampler(),
) )
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = combined_image_conditionings(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
del video_encoder
torch.cuda.synchronize() torch.cuda.synchronize()
cleanup_memory() cleanup_memory()
@@ -203,15 +215,6 @@ class TI2VidTwoStagesPipeline:
), ),
) )
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( video_state, audio_state = denoise_audio_video(
output_shape=stage_2_output_shape, output_shape=stage_2_output_shape,
conditionings=stage_2_conditionings, conditionings=stage_2_conditionings,
@@ -229,7 +232,6 @@ class TI2VidTwoStagesPipeline:
torch.cuda.synchronize() torch.cuda.synchronize()
del transformer del transformer
del video_encoder
cleanup_memory() cleanup_memory()
decoded_video = vae_decode_video( decoded_video = vae_decode_video(
@@ -253,7 +255,7 @@ def main() -> None:
distilled_lora=args.distilled_lora, distilled_lora=args.distilled_lora,
spatial_upsampler_path=args.spatial_upsampler_path, spatial_upsampler_path=args.spatial_upsampler_path,
gemma_root=args.gemma_root, gemma_root=args.gemma_root,
loras=args.lora, loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization, quantization=args.quantization,
) )
tiling_config = TilingConfig.default() tiling_config = TilingConfig.default()
@@ -14,30 +14,29 @@ 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 TilingConfig, get_video_chunks_number
from ltx_core.model.video_vae import decode_video as vae_decode_video from ltx_core.model.video_vae import decode_video as vae_decode_video
from ltx_core.quantization import QuantizationPolicy from ltx_core.quantization import QuantizationPolicy
from ltx_core.text_encoders.gemma import encode_text
from ltx_core.tools import VideoLatentShape from ltx_core.tools import VideoLatentShape
from ltx_core.types import Audio, LatentState, VideoPixelShape from ltx_core.types import Audio, LatentState, VideoPixelShape
from ltx_pipelines.utils import ( from ltx_pipelines.utils import (
ModelLedger, ModelLedger,
assert_resolution, assert_resolution,
cleanup_memory, cleanup_memory,
combined_image_conditionings,
denoise_audio_video, denoise_audio_video,
generate_enhanced_prompt, encode_prompts,
get_device, get_device,
image_conditionings_by_replacing_latent,
multi_modal_guider_denoising_func, multi_modal_guider_denoising_func,
res2s_audio_video_denoising_loop, res2s_audio_video_denoising_loop,
simple_denoising_func, simple_denoising_func,
) )
from ltx_pipelines.utils.args import ImageConditioningInput, default_2_stage_arg_parser, detect_checkpoint_path from ltx_pipelines.utils.args import ImageConditioningInput, hq_2_stage_arg_parser
from ltx_pipelines.utils.constants import STAGE_2_DISTILLED_SIGMA_VALUES, detect_params from ltx_pipelines.utils.constants import LTX_2_3_HQ_PARAMS, STAGE_2_DISTILLED_SIGMA_VALUES
from ltx_pipelines.utils.media_io import encode_video from ltx_pipelines.utils.media_io import encode_video
from ltx_pipelines.utils.types import PipelineComponents from ltx_pipelines.utils.types import PipelineComponents
device = get_device() device = get_device()
class TI2VidTwoStagesRes2sPipeline: class TI2VidTwoStagesHQPipeline:
""" """
Two-stage text/image-to-video generation pipeline using the res_2s sampler. Two-stage text/image-to-video generation pipeline using the res_2s sampler.
Same structure as :class:`TI2VidTwoStagesPipeline`: stage 1 generates video at Same structure as :class:`TI2VidTwoStagesPipeline`: stage 1 generates video at
@@ -53,26 +52,38 @@ class TI2VidTwoStagesRes2sPipeline:
self, self,
checkpoint_path: str, checkpoint_path: str,
distilled_lora: list[LoraPathStrengthAndSDOps], distilled_lora: list[LoraPathStrengthAndSDOps],
distilled_lora_strength_stage_1: float,
distilled_lora_strength_stage_2: float,
spatial_upsampler_path: str, spatial_upsampler_path: str,
gemma_root: str, gemma_root: str,
loras: list[LoraPathStrengthAndSDOps], loras: tuple[LoraPathStrengthAndSDOps, ...],
device: str = device, device: str = device,
quantization: QuantizationPolicy | None = None, quantization: QuantizationPolicy | None = None,
): ):
self.device = device self.device = device
self.dtype = torch.bfloat16 self.dtype = torch.bfloat16
distilled_lora_stage_1 = LoraPathStrengthAndSDOps(
path=distilled_lora[0].path,
strength=distilled_lora_strength_stage_1,
sd_ops=distilled_lora[0].sd_ops,
)
distilled_lora_stage_2 = LoraPathStrengthAndSDOps(
path=distilled_lora[0].path,
strength=distilled_lora_strength_stage_2,
sd_ops=distilled_lora[0].sd_ops,
)
self.stage_1_model_ledger = ModelLedger( self.stage_1_model_ledger = ModelLedger(
dtype=self.dtype, dtype=self.dtype,
device=device, device=device,
checkpoint_path=checkpoint_path, checkpoint_path=checkpoint_path,
gemma_root_path=gemma_root, gemma_root_path=gemma_root,
spatial_upsampler_path=spatial_upsampler_path, spatial_upsampler_path=spatial_upsampler_path,
loras=loras, loras=(*loras, distilled_lora_stage_1),
quantization=quantization, quantization=quantization,
) )
self.stage_2_model_ledger = self.stage_1_model_ledger.with_loras( self.stage_2_model_ledger = self.stage_1_model_ledger.with_loras(
loras=distilled_lora, loras=(*loras, distilled_lora_stage_2),
) )
self.pipeline_components = PipelineComponents( self.pipeline_components = PipelineComponents(
@@ -103,23 +114,18 @@ class TI2VidTwoStagesRes2sPipeline:
noiser = GaussianNoiser(generator=generator) noiser = GaussianNoiser(generator=generator)
dtype = torch.bfloat16 dtype = torch.bfloat16
text_encoder = self.stage_1_model_ledger.text_encoder() ctx_p, ctx_n = encode_prompts(
if enhance_prompt: [prompt, negative_prompt],
prompt = generate_enhanced_prompt( self.stage_1_model_ledger,
text_encoder, prompt, images[0][0] if len(images) > 0 else None, seed=seed enhance_first_prompt=enhance_prompt,
enhance_prompt_image=images[0][0] if len(images) > 0 else None,
enhance_prompt_seed=seed,
) )
context_p, context_n = encode_text(text_encoder, prompts=[prompt, negative_prompt]) v_context_p, a_context_p = ctx_p.video_encoding, ctx_p.audio_encoding
v_context_p, a_context_p = context_p v_context_n, a_context_n = ctx_n.video_encoding, ctx_n.audio_encoding
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: encode image conditionings with the VAE encoder, then free it
# before loading the transformer to reduce peak VRAM.
stage_1_output_shape = VideoPixelShape( stage_1_output_shape = VideoPixelShape(
batch=1, batch=1,
frames=num_frames, frames=num_frames,
@@ -127,6 +133,21 @@ class TI2VidTwoStagesRes2sPipeline:
height=height // 2, height=height // 2,
fps=frame_rate, fps=frame_rate,
) )
video_encoder = self.stage_1_model_ledger.video_encoder()
stage_1_conditionings = combined_image_conditionings(
images=images,
height=stage_1_output_shape.height,
width=stage_1_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
torch.cuda.synchronize()
del video_encoder
cleanup_memory()
transformer = self.stage_1_model_ledger.transformer()
empty_latent = torch.empty(VideoLatentShape.from_pixel_shape(stage_1_output_shape).to_torch_shape()) empty_latent = torch.empty(VideoLatentShape.from_pixel_shape(stage_1_output_shape).to_torch_shape())
stepper = Res2sDiffusionStep() stepper = Res2sDiffusionStep()
sigmas = ( sigmas = (
@@ -158,14 +179,6 @@ class TI2VidTwoStagesRes2sPipeline:
), ),
) )
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( video_state, audio_state = denoise_audio_video(
output_shape=stage_1_output_shape, output_shape=stage_1_output_shape,
conditionings=stage_1_conditionings, conditionings=stage_1_conditionings,
@@ -183,13 +196,24 @@ class TI2VidTwoStagesRes2sPipeline:
cleanup_memory() cleanup_memory()
# Stage 2: Upsample and refine the video at higher resolution with distilled LORA. # Stage 2: Upsample and refine the video at higher resolution with distilled LORA.
video_encoder = self.stage_1_model_ledger.video_encoder()
upscaled_video_latent = upsample_video( upscaled_video_latent = upsample_video(
latent=video_state.latent[:1], latent=video_state.latent[:1],
video_encoder=video_encoder, video_encoder=video_encoder,
upsampler=self.stage_2_model_ledger.spatial_upsampler(), upsampler=self.stage_2_model_ledger.spatial_upsampler(),
) )
stage_2_output_shape = VideoPixelShape(batch=1, frames=num_frames, width=width, height=height, fps=frame_rate)
stage_2_conditionings = combined_image_conditionings(
images=images,
height=stage_2_output_shape.height,
width=stage_2_output_shape.width,
video_encoder=video_encoder,
dtype=dtype,
device=self.device,
)
torch.cuda.synchronize() torch.cuda.synchronize()
del video_encoder
cleanup_memory() cleanup_memory()
transformer = self.stage_2_model_ledger.transformer() transformer = self.stage_2_model_ledger.transformer()
@@ -210,15 +234,6 @@ class TI2VidTwoStagesRes2sPipeline:
), ),
) )
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( video_state, audio_state = denoise_audio_video(
output_shape=stage_2_output_shape, output_shape=stage_2_output_shape,
conditionings=stage_2_conditionings, conditionings=stage_2_conditionings,
@@ -236,7 +251,6 @@ class TI2VidTwoStagesRes2sPipeline:
torch.cuda.synchronize() torch.cuda.synchronize()
del transformer del transformer
del video_encoder
cleanup_memory() cleanup_memory()
decoded_video = vae_decode_video( decoded_video = vae_decode_video(
@@ -251,16 +265,16 @@ class TI2VidTwoStagesRes2sPipeline:
@torch.inference_mode() @torch.inference_mode()
def main() -> None: def main() -> None:
logging.getLogger().setLevel(logging.INFO) logging.getLogger().setLevel(logging.INFO)
checkpoint_path = detect_checkpoint_path() parser = hq_2_stage_arg_parser(params=LTX_2_3_HQ_PARAMS)
params = detect_params(checkpoint_path)
parser = default_2_stage_arg_parser(params=params)
args = parser.parse_args() args = parser.parse_args()
pipeline = TI2VidTwoStagesRes2sPipeline( pipeline = TI2VidTwoStagesHQPipeline(
checkpoint_path=args.checkpoint_path, checkpoint_path=args.checkpoint_path,
distilled_lora=args.distilled_lora, distilled_lora=args.distilled_lora,
distilled_lora_strength_stage_1=args.distilled_lora_strength_stage_1,
distilled_lora_strength_stage_2=args.distilled_lora_strength_stage_2,
spatial_upsampler_path=args.spatial_upsampler_path, spatial_upsampler_path=args.spatial_upsampler_path,
gemma_root=args.gemma_root, gemma_root=args.gemma_root,
loras=args.lora, loras=tuple(args.lora) if args.lora else (),
quantization=args.quantization, quantization=args.quantization,
) )
tiling_config = TilingConfig.default() tiling_config = TilingConfig.default()
@@ -1,10 +1,11 @@
from ltx_pipelines.utils.helpers import ( from ltx_pipelines.utils.helpers import (
assert_resolution, assert_resolution,
cleanup_memory, cleanup_memory,
combined_image_conditionings,
denoise_audio_video, denoise_audio_video,
encode_prompts,
generate_enhanced_prompt, generate_enhanced_prompt,
get_device, get_device,
image_conditionings_by_replacing_latent,
multi_modal_guider_denoising_func, multi_modal_guider_denoising_func,
multi_modal_guider_factory_denoising_func, multi_modal_guider_factory_denoising_func,
simple_denoising_func, simple_denoising_func,
@@ -20,12 +21,13 @@ __all__ = [
"ModelLedger", "ModelLedger",
"assert_resolution", "assert_resolution",
"cleanup_memory", "cleanup_memory",
"combined_image_conditionings",
"denoise_audio_video", "denoise_audio_video",
"encode_prompts",
"euler_denoising_loop", "euler_denoising_loop",
"generate_enhanced_prompt", "generate_enhanced_prompt",
"get_device", "get_device",
"gradient_estimating_euler_denoising_loop", "gradient_estimating_euler_denoising_loop",
"image_conditionings_by_replacing_latent",
"multi_modal_guider_denoising_func", "multi_modal_guider_denoising_func",
"multi_modal_guider_factory_denoising_func", "multi_modal_guider_factory_denoising_func",
"res2s_audio_video_denoising_loop", "res2s_audio_video_denoising_loop",
@@ -8,6 +8,7 @@ from ltx_pipelines.utils.constants import (
DEFAULT_IMAGE_CRF, DEFAULT_IMAGE_CRF,
DEFAULT_LORA_STRENGTH, DEFAULT_LORA_STRENGTH,
DEFAULT_NEGATIVE_PROMPT, DEFAULT_NEGATIVE_PROMPT,
LTX_2_3_HQ_PARAMS,
LTX_2_3_PARAMS, LTX_2_3_PARAMS,
PipelineParams, PipelineParams,
) )
@@ -457,6 +458,23 @@ def default_2_stage_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argpa
return parser return parser
def hq_2_stage_arg_parser(params: PipelineParams = LTX_2_3_HQ_PARAMS) -> argparse.ArgumentParser:
parser = default_2_stage_arg_parser(params=params)
parser.add_argument(
"--distilled-lora-strength-stage-1",
type=float,
default=0.25,
help=(f"Strength of the distilled LoRA used in the first stage (default: {0.25})."),
)
parser.add_argument(
"--distilled-lora-strength-stage-2",
type=float,
default=0.5,
help=(f"Strength of the distilled LoRA used in the second stage (default: {0.5})."),
)
return parser
def default_2_stage_distilled_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argparse.ArgumentParser: def default_2_stage_distilled_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argparse.ArgumentParser:
parser = basic_arg_parser(params=params, distilled=True) parser = basic_arg_parser(params=params, distilled=True)
parser.set_defaults(height=params.stage_2_height, width=params.stage_2_width) parser.set_defaults(height=params.stage_2_height, width=params.stage_2_width)
@@ -71,6 +71,27 @@ LTX_2_3_PARAMS = replace(
video_guider_params=replace(LTX_2_PARAMS.video_guider_params, stg_blocks=[28]), 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_guider_params=replace(LTX_2_PARAMS.audio_guider_params, stg_blocks=[28]),
) )
LTX_2_3_HQ_PARAMS = PipelineParams(
num_inference_steps=15,
stage_1_height=1088 // 2,
stage_1_width=1920 // 2,
video_guider_params=MultiModalGuiderParams(
cfg_scale=3.0,
stg_scale=0.0,
rescale_scale=0.45,
modality_scale=3.0,
skip_step=0,
stg_blocks=[],
),
audio_guider_params=MultiModalGuiderParams(
cfg_scale=7.0,
stg_scale=0.0,
rescale_scale=1.0,
modality_scale=3.0,
skip_step=0,
stg_blocks=[],
),
)
DEFAULT_LORA_STRENGTH = 1.0 DEFAULT_LORA_STRENGTH = 1.0
DEFAULT_IMAGE_CRF = 33 DEFAULT_IMAGE_CRF = 33
@@ -21,6 +21,7 @@ from ltx_core.guidance.perturbations import (
from ltx_core.model.transformer import Modality, X0Model from ltx_core.model.transformer import Modality, X0Model
from ltx_core.model.video_vae import VideoEncoder from ltx_core.model.video_vae import VideoEncoder
from ltx_core.text_encoders.gemma import GemmaTextEncoder from ltx_core.text_encoders.gemma import GemmaTextEncoder
from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessorOutput
from ltx_core.tools import AudioLatentTools, LatentTools, VideoLatentTools from ltx_core.tools import AudioLatentTools, LatentTools, VideoLatentTools
from ltx_core.types import AudioLatentShape, LatentState, VideoLatentShape, VideoPixelShape from ltx_core.types import AudioLatentShape, LatentState, VideoLatentShape, VideoPixelShape
from ltx_pipelines.utils.args import ImageConditioningInput from ltx_pipelines.utils.args import ImageConditioningInput
@@ -44,6 +45,84 @@ def cleanup_memory() -> None:
torch.cuda.synchronize() torch.cuda.synchronize()
def encode_prompts(
prompts: list[str],
model_ledger: object,
*,
enhance_prompt_image: str | None = None,
enhance_prompt_seed: int = 42,
enhance_first_prompt: bool = False,
) -> list[EmbeddingsProcessorOutput]:
"""Encode prompts through Gemma → embeddings processor, freeing each after use.
Loads the text encoder from *model_ledger*, optionally enhances the first
prompt, encodes all *prompts*, frees the text encoder, then loads the
embeddings processor to produce the final outputs. Because the text encoder
is loaded and freed entirely within this function, there are no lingering
references that could prevent GPU memory reclamation.
Args:
prompts: Text prompts to encode.
model_ledger: ModelLedger instance (used to load text encoder and embeddings processor).
enhance_prompt_image: Optional image path for prompt enhancement.
enhance_prompt_seed: Seed for prompt enhancement (default 42).
enhance_first_prompt: If True, enhance ``prompts[0]`` before encoding.
Returns:
List of EmbeddingsProcessorOutput, one per prompt.
"""
text_encoder = model_ledger.text_encoder()
if enhance_first_prompt:
prompts = list(prompts)
prompts[0] = generate_enhanced_prompt(text_encoder, prompts[0], enhance_prompt_image, seed=enhance_prompt_seed)
raw_outputs = [text_encoder.encode(p) for p in prompts]
torch.cuda.synchronize()
del text_encoder
cleanup_memory()
embeddings_processor = model_ledger.gemma_embeddings_processor()
results: list[EmbeddingsProcessorOutput] = [
embeddings_processor.process_hidden_states(hs, mask) for hs, mask in raw_outputs
]
del embeddings_processor
cleanup_memory()
return results
def combined_image_conditionings(
images: list[ImageConditioningInput],
height: int,
width: int,
video_encoder: VideoEncoder,
dtype: torch.dtype,
device: torch.device,
) -> list[ConditioningItem]:
"""Create a list of conditionings by replacing the latent at the first frame with the encoded image if present
and using other encoded images as the keyframe conditionings."""
conditionings = []
for img in images:
image = load_image_conditioning(
image_path=img.path,
height=height,
width=width,
dtype=dtype,
device=device,
crf=img.crf,
)
encoded_image = video_encoder(image)
if img.frame_idx == 0:
conditioning = VideoConditionByLatentIndex(
latent=encoded_image,
strength=img.strength,
latent_idx=0,
)
else:
conditioning = VideoConditionByKeyframeIndex(
keyframes=encoded_image,
strength=img.strength,
frame_idx=img.frame_idx,
)
conditionings.append(conditioning)
return conditionings
def image_conditionings_by_replacing_latent( def image_conditionings_by_replacing_latent(
images: list[ImageConditioningInput], images: list[ImageConditioningInput],
height: int, height: int,
@@ -33,8 +33,11 @@ from ltx_core.model.video_vae import (
) )
from ltx_core.quantization import QuantizationPolicy from ltx_core.quantization import QuantizationPolicy
from ltx_core.text_encoders.gemma import ( from ltx_core.text_encoders.gemma import (
AV_GEMMA_TEXT_ENCODER_KEY_OPS, EMBEDDINGS_PROCESSOR_KEY_OPS,
GEMMA_LLM_KEY_OPS,
GEMMA_MODEL_OPS, GEMMA_MODEL_OPS,
EmbeddingsProcessor,
EmbeddingsProcessorConfigurator,
GemmaTextEncoder, GemmaTextEncoder,
GemmaTextEncoderConfigurator, GemmaTextEncoderConfigurator,
module_ops_from_gemma_root, module_ops_from_gemma_root,
@@ -76,8 +79,8 @@ class ModelLedger:
:meth:`spatial_upsampler` method becomes available; otherwise calling it raises :meth:`spatial_upsampler` method becomes available; otherwise calling it raises
a :class:`ValueError`. a :class:`ValueError`.
loras: loras:
Optional collection of LoRA configurations (paths, strengths, and key operations) Tuple of LoRA configurations (path, strength, sd_ops) applied on top of the base
that are applied on top of the base transformer weights when building the model. transformer weights. Use ``()`` for none.
registry: registry:
Optional :class:`Registry` instance for weight caching across builders. Optional :class:`Registry` instance for weight caching across builders.
Defaults to :class:`DummyRegistry` which performs no cross-builder caching. Defaults to :class:`DummyRegistry` which performs no cross-builder caching.
@@ -85,8 +88,9 @@ class ModelLedger:
Optional :class:`QuantizationPolicy` controlling how transformer weights Optional :class:`QuantizationPolicy` controlling how transformer weights
are stored and how matmul is executed. Defaults to None, which means no quantization. are stored and how matmul is executed. Defaults to None, which means no quantization.
### Creating Variants ### Creating Variants
Use :meth:`with_loras` to create a new ``ModelLedger`` instance that includes Use :meth:`with_additional_loras` to create a new ``ModelLedger`` instance that
additional LoRA configurations while sharing the same registry for weight caching. includes additional LoRA configurations or :meth:`with_loras` to replace existing
lora configurations while sharing the same registry for weight caching.
""" """
def __init__( def __init__(
@@ -96,7 +100,7 @@ class ModelLedger:
checkpoint_path: str | None = None, checkpoint_path: str | None = None,
gemma_root_path: str | None = None, gemma_root_path: str | None = None,
spatial_upsampler_path: str | None = None, spatial_upsampler_path: str | None = None,
loras: LoraPathStrengthAndSDOps | None = None, loras: tuple[LoraPathStrengthAndSDOps, ...] = (),
registry: Registry | None = None, registry: Registry | None = None,
quantization: QuantizationPolicy | None = None, quantization: QuantizationPolicy | None = None,
): ):
@@ -105,7 +109,7 @@ class ModelLedger:
self.checkpoint_path = checkpoint_path self.checkpoint_path = checkpoint_path
self.gemma_root_path = gemma_root_path self.gemma_root_path = gemma_root_path
self.spatial_upsampler_path = spatial_upsampler_path self.spatial_upsampler_path = spatial_upsampler_path
self.loras = loras or () self.loras = loras
self.registry = registry or DummyRegistry() self.registry = registry or DummyRegistry()
self.quantization = quantization self.quantization = quantization
self.build_model_builders() self.build_model_builders()
@@ -155,15 +159,23 @@ class ModelLedger:
registry=self.registry, registry=self.registry,
) )
# Embeddings processor only needs the LTX checkpoint (no Gemma weights)
self.embeddings_processor_builder = Builder(
model_path=self.checkpoint_path,
model_class_configurator=EmbeddingsProcessorConfigurator,
model_sd_ops=EMBEDDINGS_PROCESSOR_KEY_OPS,
registry=self.registry,
)
if self.gemma_root_path is not None: if self.gemma_root_path is not None:
module_ops = module_ops_from_gemma_root(self.gemma_root_path) module_ops = module_ops_from_gemma_root(self.gemma_root_path)
model_folder = find_matching_file(self.gemma_root_path, "model*.safetensors").parent model_folder = find_matching_file(self.gemma_root_path, "model*.safetensors").parent
weight_paths = [str(p) for p in model_folder.rglob("*.safetensors")] weight_paths = [str(p) for p in model_folder.rglob("*.safetensors")]
self.text_encoder_builder = Builder( self.text_encoder_builder = Builder(
model_path=(str(self.checkpoint_path), *weight_paths), model_path=tuple(weight_paths),
model_class_configurator=GemmaTextEncoderConfigurator, model_class_configurator=GemmaTextEncoderConfigurator,
model_sd_ops=AV_GEMMA_TEXT_ENCODER_KEY_OPS, model_sd_ops=GEMMA_LLM_KEY_OPS,
registry=self.registry, registry=self.registry,
module_ops=(GEMMA_MODEL_OPS, *module_ops), module_ops=(GEMMA_MODEL_OPS, *module_ops),
) )
@@ -181,14 +193,19 @@ class ModelLedger:
else: else:
return torch.device("cpu") return torch.device("cpu")
def with_loras(self, loras: LoraPathStrengthAndSDOps) -> "ModelLedger": def with_additional_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> "ModelLedger":
"""Add new lora configurations to the existing ones."""
return self.with_loras((*self.loras, *loras))
def with_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> "ModelLedger":
"""Replace existing lora configurations with new ones."""
return ModelLedger( return ModelLedger(
dtype=self.dtype, dtype=self.dtype,
device=self.device, device=self.device,
checkpoint_path=self.checkpoint_path, checkpoint_path=self.checkpoint_path,
gemma_root_path=self.gemma_root_path, gemma_root_path=self.gemma_root_path,
spatial_upsampler_path=self.spatial_upsampler_path, spatial_upsampler_path=self.spatial_upsampler_path,
loras=(*self.loras, *loras), loras=loras,
registry=self.registry, registry=self.registry,
quantization=self.quantization, quantization=self.quantization,
) )
@@ -244,6 +261,18 @@ class ModelLedger:
return self.text_encoder_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval() return self.text_encoder_builder.build(device=self._target_device(), dtype=self.dtype).to(self.device).eval()
def gemma_embeddings_processor(self) -> EmbeddingsProcessor:
if not hasattr(self, "embeddings_processor_builder"):
raise ValueError(
"Embeddings processor not initialized. Please provide a checkpoint path to the ModelLedger constructor."
)
return (
self.embeddings_processor_builder.build(device=self._target_device(), dtype=self.dtype)
.to(self.device)
.eval()
)
def audio_encoder(self) -> AudioEncoder: def audio_encoder(self) -> AudioEncoder:
if not hasattr(self, "audio_encoder_builder"): if not hasattr(self, "audio_encoder_builder"):
raise ValueError( raise ValueError(
+8 -6
View File
@@ -14,7 +14,7 @@ This file provides guidance to AI coding assistants (Claude, Cursor, etc.) when
**Supported model versions:** **Supported model versions:**
- **LTX-2** (19B, initial audio-video model) - **LTX-2** (19B, initial audio-video model)
- **LTX-2.3** (20B, improved text conditioning and audio quality) - **LTX-2.3** (22B, improved text conditioning and audio quality)
Version detection is fully automatic — ltx-core reads the checkpoint config and selects the correct architecture Version detection is fully automatic — ltx-core reads the checkpoint config and selects the correct architecture
components. The trainer does not need version-specific code paths. components. The trainer does not need version-specific code paths.
@@ -80,9 +80,10 @@ packages/ltx-trainer/
- `ltx_trainer.model_loader` provides component loaders using `ltx-core` - `ltx_trainer.model_loader` provides component loaders using `ltx-core`
- Individual loaders: `load_transformer()`, `load_video_vae_encoder()`, `load_video_vae_decoder()`, - Individual loaders: `load_transformer()`, `load_video_vae_encoder()`, `load_video_vae_decoder()`,
`load_text_encoder()`, etc. `load_text_encoder()`, `load_embeddings_processor()`, etc.
- Combined loader: `load_model()` returns `LtxModelComponents` dataclass - Combined loader: `load_model()` returns `LtxModelComponents` dataclass
- Uses `SingleGPUModelBuilder` from ltx-core internally - Uses `SingleGPUModelBuilder` from ltx-core internally
- Text encoder and embeddings processor are loaded separately (the text encoder only needs Gemma weights; the embeddings processor only needs the LTX checkpoint)
- 8-bit text encoder loading via `gemma_8bit.py` (bitsandbytes) - 8-bit text encoder loading via `gemma_8bit.py` (bitsandbytes)
**Training Flow:** **Training Flow:**
@@ -146,7 +147,7 @@ Both model versions share the same latent space interface (see [Latent Space Con
The differences lie in how text conditioning and audio generation work. Version detection is automatic via checkpoint The differences lie in how text conditioning and audio generation work. Version detection is automatic via checkpoint
config — the trainer uses a unified API. config — the trainer uses a unified API.
| Component | LTX-2 (19B) | LTX-2.3 (20B) | | Component | LTX-2 (19B) | LTX-2.3 (22B) |
|-----------------------|---------------------------------------------------------------------------------|-----------------------------------------------------------------------------------------------------| |-----------------------|---------------------------------------------------------------------------------|-----------------------------------------------------------------------------------------------------|
| Feature extractor | `FeatureExtractorV1`: single `aggregate_embed`, same output for video and audio | `FeatureExtractorV2`: separate `video_aggregate_embed` + `audio_aggregate_embed`, per-token RMSNorm | | Feature extractor | `FeatureExtractorV1`: single `aggregate_embed`, same output for video and audio | `FeatureExtractorV2`: separate `video_aggregate_embed` + `audio_aggregate_embed`, per-token RMSNorm |
| Caption projection | Inside the transformer (`caption_projection`) | Inside the feature extractor (before connector) | | Caption projection | Inside the transformer (`caption_projection`) | Inside the feature extractor (before connector) |
@@ -309,7 +310,7 @@ Key classes:
- Implements distributed training with Accelerate - Implements distributed training with Accelerate
- Handles mixed precision, gradient accumulation, checkpointing - Handles mixed precision, gradient accumulation, checkpointing
- `_training_step()` applies embedding connectors then delegates to strategy - `_training_step()` applies embedding connectors then delegates to strategy
- `_load_text_encoder_and_cache_embeddings()` caches validation embeddings and unloads heavy components - `_load_text_encoder_and_cache_embeddings()` loads the text encoder + embeddings processor, caches validation embeddings, then unloads the Gemma LLM (keeps only the embeddings processor connectors for training)
- Uses training strategies for mode-specific logic - Uses training strategies for mode-specific logic
**`src/ltx_trainer/training_strategies/`** - Strategy pattern **`src/ltx_trainer/training_strategies/`** - Strategy pattern
@@ -334,7 +335,8 @@ Component loaders:
- `load_video_vae_decoder()``VideoDecoder` - `load_video_vae_decoder()``VideoDecoder`
- `load_audio_vae_decoder()``AudioDecoder` - `load_audio_vae_decoder()``AudioDecoder`
- `load_vocoder()``Vocoder` or `VocoderWithBWE` (auto-detected) - `load_vocoder()``Vocoder` or `VocoderWithBWE` (auto-detected)
- `load_text_encoder()``GemmaTextEncoder` (unified, handles V1/V2 automatically) - `load_text_encoder(gemma_model_path)``GemmaTextEncoder` (pure Gemma LLM, no checkpoint needed)
- `load_embeddings_processor(checkpoint_path)``EmbeddingsProcessor` (feature extractor + connectors)
- `load_model()``LtxModelComponents` (convenience wrapper) - `load_model()``LtxModelComponents` (convenience wrapper)
**`src/ltx_trainer/validation_sampler.py`** - Inference for validation **`src/ltx_trainer/validation_sampler.py`** - Inference for validation
@@ -506,7 +508,7 @@ packages/ltx-core/src/ltx_core/
│ ├── encoders/ │ ├── encoders/
│ │ ├── base_encoder.py # GemmaTextEncoder (unified 3-block pipeline) │ │ ├── base_encoder.py # GemmaTextEncoder (unified 3-block pipeline)
│ │ └── encoder_configurator.py # GemmaTextEncoderConfigurator, _create_feature_extractor │ │ └── encoder_configurator.py # GemmaTextEncoderConfigurator, _create_feature_extractor
│ ├── feature_extractor.py # FeatureExtractorV1 (19B), FeatureExtractorV2 (20B) │ ├── feature_extractor.py # FeatureExtractorV1 (19B), FeatureExtractorV2 (22B)
│ ├── embeddings_connector.py # Embeddings1DConnector, Embeddings1DConnectorConfigurator, │ ├── embeddings_connector.py # Embeddings1DConnector, Embeddings1DConnectorConfigurator,
│ │ # AudioEmbeddings1DConnectorConfigurator │ │ # AudioEmbeddings1DConnectorConfigurator
│ ├── embeddings_processor.py # EmbeddingsProcessor (wraps video + audio connectors) │ ├── embeddings_processor.py # EmbeddingsProcessor (wraps video + audio connectors)
@@ -34,7 +34,7 @@ from torch.utils.data import DataLoader, Dataset
from transformers.utils.logging import disable_progress_bar from transformers.utils.logging import disable_progress_bar
from ltx_trainer import logger from ltx_trainer import logger
from ltx_trainer.model_loader import load_text_encoder from ltx_trainer.model_loader import load_embeddings_processor, load_text_encoder
# Disable tokenizers parallelism to avoid warnings # Disable tokenizers parallelism to avoid warnings
os.environ["TOKENIZERS_PARALLELISM"] = "false" os.environ["TOKENIZERS_PARALLELISM"] = "false"
@@ -264,17 +264,21 @@ def compute_captions_embeddings( # noqa: PLR0913
output_path = Path(output_dir) output_path = Path(output_dir)
output_path.mkdir(parents=True, exist_ok=True) output_path.mkdir(parents=True, exist_ok=True)
# Load text encoder # Load text encoder and embeddings processor
with console.status("[bold]Loading Gemma text encoder...", spinner="dots"): with console.status("[bold]Loading Gemma text encoder...", spinner="dots"):
text_encoder = load_text_encoder( text_encoder = load_text_encoder(
model_path,
text_encoder_path, text_encoder_path,
device=device, device=device,
dtype=torch.bfloat16, dtype=torch.bfloat16,
load_in_8bit=load_in_8bit, load_in_8bit=load_in_8bit,
) )
embeddings_processor = load_embeddings_processor(
model_path,
device=device,
dtype=torch.bfloat16,
)
logger.info("Text encoder loaded successfully") logger.info("Text encoder and embeddings processor loaded successfully")
# TODO(batch-tokenization): The current Gemma tokenizer doesn't support batched tokenization. # TODO(batch-tokenization): The current Gemma tokenizer doesn't support batched tokenization.
if batch_size > 1: if batch_size > 1:
@@ -303,14 +307,16 @@ def compute_captions_embeddings( # noqa: PLR0913
) as progress: ) as progress:
task = progress.add_task("Processing captions", total=len(dataloader)) task = progress.add_task("Processing captions", total=len(dataloader))
for batch in dataloader: for batch in dataloader:
# Encode prompts using precompute() (returns video/audio features before connector) # Encode prompts using text_encoder.encode() + feature_extractor
# (returns video/audio features before connector).
# The connector is applied during training via embeddings_processor # The connector is applied during training via embeddings_processor
with torch.inference_mode(): with torch.inference_mode():
# TODO(batch-tokenization): When tokenizer supports batching, encode all prompts at once. # TODO(batch-tokenization): When tokenizer supports batching, encode all prompts at once.
# For now, process one at a time: # For now, process one at a time:
for i in range(len(batch["prompt"])): for i in range(len(batch["prompt"])):
video_prompt_embeds, audio_prompt_embeds, prompt_attention_mask = text_encoder.precompute( hidden_states, prompt_attention_mask = text_encoder.encode(batch["prompt"][i], padding_side="left")
batch["prompt"][i], padding_side="left" video_prompt_embeds, audio_prompt_embeds = embeddings_processor.feature_extractor(
hidden_states, prompt_attention_mask, "left"
) )
output_rel_path = Path(batch["output_path"][i]) output_rel_path = Path(batch["output_path"][i])
@@ -6,10 +6,7 @@ This module provides functionality for loading the Gemma text encoder in 8-bit p
using bitsandbytes, which significantly reduces GPU memory usage. using bitsandbytes, which significantly reduces GPU memory usage.
Example usage: Example usage:
from ltx_trainer.gemma_8bit import load_8bit_gemma from ltx_trainer.gemma_8bit import load_8bit_gemma
text_encoder = load_8bit_gemma( text_encoder = load_8bit_gemma(gemma_model_path="/path/to/gemma")
checkpoint_path="/path/to/ltx2.safetensors",
gemma_model_path="/path/to/gemma",
)
""" """
from __future__ import annotations from __future__ import annotations
@@ -21,34 +18,20 @@ from pathlib import Path
import torch import torch
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader
from ltx_core.text_encoders.gemma import AV_GEMMA_TEXT_ENCODER_KEY_OPS
from ltx_core.text_encoders.gemma.embeddings_connector import (
AudioEmbeddings1DConnectorConfigurator,
Embeddings1DConnectorConfigurator,
)
from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessor
from ltx_core.text_encoders.gemma.encoders.base_encoder import GemmaTextEncoder from ltx_core.text_encoders.gemma.encoders.base_encoder import GemmaTextEncoder
from ltx_core.text_encoders.gemma.encoders.encoder_configurator import _create_feature_extractor
from ltx_core.text_encoders.gemma.tokenizer import LTXVGemmaTokenizer from ltx_core.text_encoders.gemma.tokenizer import LTXVGemmaTokenizer
def load_8bit_gemma( def load_8bit_gemma(gemma_model_path: str | Path, dtype: torch.dtype = torch.bfloat16) -> GemmaTextEncoder:
checkpoint_path: str | Path,
gemma_model_path: str | Path,
dtype: torch.dtype = torch.bfloat16,
) -> GemmaTextEncoder:
"""Load the Gemma text encoder in 8-bit precision using bitsandbytes. """Load the Gemma text encoder in 8-bit precision using bitsandbytes.
This function bypasses ltx-core's standard loading path to enable 8-bit quantization Only the Gemma LLM backbone is loaded here. The embeddings processor
via the bitsandbytes library. The Gemma model is loaded with load_in_8bit=True and (feature extractor + connectors) should be loaded separately via
torch_dtype=bfloat16, while the feature extractor and connector weights are loaded :func:`ltx_trainer.model_loader.load_embeddings_processor`.
from the LTX-2 checkpoint.
Args: Args:
checkpoint_path: Path to the LTX-2 safetensors checkpoint file
gemma_model_path: Path to Gemma model directory gemma_model_path: Path to Gemma model directory
dtype: Data type for non-quantized model weights (feature extractor, connectors) dtype: Data type for non-quantized model weights
Returns: Returns:
Loaded GemmaTextEncoder with 8-bit quantized Gemma backbone GemmaTextEncoder with 8-bit quantized Gemma backbone
Raises: Raises:
ImportError: If bitsandbytes is not installed ImportError: If bitsandbytes is not installed
FileNotFoundError: If required model files are not found FileNotFoundError: If required model files are not found
@@ -60,7 +43,6 @@ def load_8bit_gemma(
"8-bit text encoder loading requires bitsandbytes. Install it with: uv pip install bitsandbytes" "8-bit text encoder loading requires bitsandbytes. Install it with: uv pip install bitsandbytes"
) from e ) from e
# Find paths within gemma_model_path
gemma_path = _find_gemma_subpath(gemma_model_path, "model*.safetensors") gemma_path = _find_gemma_subpath(gemma_model_path, "model*.safetensors")
tokenizer_path = _find_gemma_subpath(gemma_model_path, "tokenizer.model") tokenizer_path = _find_gemma_subpath(gemma_model_path, "tokenizer.model")
@@ -74,51 +56,14 @@ def load_8bit_gemma(
local_files_only=True, local_files_only=True,
) )
# Load tokenizer
tokenizer = LTXVGemmaTokenizer(tokenizer_path, 1024) tokenizer = LTXVGemmaTokenizer(tokenizer_path, 1024)
# Load config and weights from the LTX-2 checkpoint return GemmaTextEncoder(
loader = SafetensorsModelStateDictLoader()
config = loader.metadata(str(checkpoint_path))
sd = loader.load(str(checkpoint_path), sd_ops=AV_GEMMA_TEXT_ENCODER_KEY_OPS)
# Helper to extract state dict for a given prefix
def extract_state_dict(prefix: str) -> dict[str, torch.Tensor]:
return {k.replace(prefix, ""): v for k, v in sd.sd.items() if k.startswith(prefix)}
# Create and load video embeddings connector
embeddings_connector = Embeddings1DConnectorConfigurator.from_config(config)
embeddings_connector.load_state_dict(extract_state_dict("embeddings_processor.video_connector."))
embeddings_connector = embeddings_connector.to(device=gemma_model.device, dtype=dtype)
# Create and load audio embeddings connector
audio_embeddings_connector = AudioEmbeddings1DConnectorConfigurator.from_config(config)
audio_embeddings_connector.load_state_dict(extract_state_dict("embeddings_processor.audio_connector."))
audio_embeddings_connector = audio_embeddings_connector.to(device=gemma_model.device, dtype=dtype)
# Create embeddings processor
embeddings_processor = EmbeddingsProcessor(
video_connector=embeddings_connector,
audio_connector=audio_embeddings_connector,
)
transformer_config = config.get("transformer", {})
feature_extractor = _create_feature_extractor(transformer_config)
feature_extractor.load_state_dict(
{k.removeprefix("feature_extractor."): v for k, v in sd.sd.items() if k.startswith("feature_extractor.")},
)
feature_extractor = feature_extractor.to(device=gemma_model.device, dtype=dtype)
text_encoder = GemmaTextEncoder(
feature_extractor=feature_extractor,
embeddings_processor=embeddings_processor,
tokenizer=tokenizer, tokenizer=tokenizer,
model=gemma_model, model=gemma_model,
dtype=dtype, dtype=dtype,
) )
return text_encoder
def _find_gemma_subpath(root_path: str | Path, pattern: str) -> str: def _find_gemma_subpath(root_path: str | Path, pattern: str) -> str:
"""Find a file matching a glob pattern and return its parent directory.""" """Find a file matching a glob pattern and return its parent directory."""
@@ -8,7 +8,7 @@ Example usage:
# Load individual components # Load individual components
vae_encoder = load_video_vae_encoder("/path/to/checkpoint.safetensors", device="cuda") vae_encoder = load_video_vae_encoder("/path/to/checkpoint.safetensors", device="cuda")
vae_decoder = load_video_vae_decoder("/path/to/checkpoint.safetensors", device="cuda") vae_decoder = load_video_vae_decoder("/path/to/checkpoint.safetensors", device="cuda")
text_encoder = load_text_encoder("/path/to/checkpoint.safetensors", "/path/to/gemma", device="cuda") text_encoder = load_text_encoder("/path/to/gemma", device="cuda")
# Load all components at once # Load all components at once
components = load_model("/path/to/checkpoint.safetensors", text_encoder_path="/path/to/gemma") components = load_model("/path/to/checkpoint.safetensors", text_encoder_path="/path/to/gemma")
""" """
@@ -33,6 +33,7 @@ if TYPE_CHECKING:
from ltx_core.model.transformer import LTXModel from ltx_core.model.transformer import LTXModel
from ltx_core.model.video_vae import VideoDecoder, VideoEncoder from ltx_core.model.video_vae import VideoDecoder, VideoEncoder
from ltx_core.text_encoders.gemma import GemmaTextEncoder from ltx_core.text_encoders.gemma import GemmaTextEncoder
from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessor
def _to_torch_device(device: Device) -> torch.device: def _to_torch_device(device: Device) -> torch.device:
@@ -187,7 +188,6 @@ def load_vocoder(
def load_text_encoder( def load_text_encoder(
checkpoint_path: str | Path,
gemma_model_path: str | Path, gemma_model_path: str | Path,
device: Device = "cpu", device: Device = "cpu",
dtype: torch.dtype = torch.bfloat16, dtype: torch.dtype = torch.bfloat16,
@@ -195,15 +195,14 @@ def load_text_encoder(
) -> "GemmaTextEncoder": ) -> "GemmaTextEncoder":
"""Load the Gemma text encoder. """Load the Gemma text encoder.
Args: Args:
checkpoint_path: Path to the LTX-2 safetensors checkpoint file
gemma_model_path: Path to Gemma model directory gemma_model_path: Path to Gemma model directory
device: Device to load model on device: Device to load model on
dtype: Data type for model weights dtype: Data type for model weights
load_in_8bit: Whether to load the Gemma model in 8-bit precision using bitsandbytes. load_in_8bit: Whether to load the Gemma model in 8-bit precision using bitsandbytes.
When True, the model is loaded with device_map="auto" and the device argument When True, the model is loaded with device_map="auto" and the device argument
is ignored for the Gemma backbone (feature extractor still uses dtype). is ignored for the Gemma backbone.
Returns: Returns:
Loaded GemmaTextEncoder (unified encoder handling V1/V2/V3) Loaded GemmaTextEncoder
""" """
if not Path(gemma_model_path).is_dir(): if not Path(gemma_model_path).is_dir():
raise ValueError(f"Gemma model path is not a directory: {gemma_model_path}") raise ValueError(f"Gemma model path is not a directory: {gemma_model_path}")
@@ -212,12 +211,12 @@ def load_text_encoder(
if load_in_8bit: if load_in_8bit:
from ltx_trainer.gemma_8bit import load_8bit_gemma from ltx_trainer.gemma_8bit import load_8bit_gemma
return load_8bit_gemma(checkpoint_path, gemma_model_path, dtype) return load_8bit_gemma(gemma_model_path, dtype)
# Standard loading path # Standard loading path
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
from ltx_core.text_encoders.gemma import ( from ltx_core.text_encoders.gemma import (
AV_GEMMA_TEXT_ENCODER_KEY_OPS, GEMMA_LLM_KEY_OPS,
GEMMA_MODEL_OPS, GEMMA_MODEL_OPS,
GemmaTextEncoderConfigurator, GemmaTextEncoderConfigurator,
module_ops_from_gemma_root, module_ops_from_gemma_root,
@@ -230,15 +229,43 @@ def load_text_encoder(
gemma_weight_paths = [str(p) for p in gemma_model_folder.rglob("*.safetensors")] gemma_weight_paths = [str(p) for p in gemma_model_folder.rglob("*.safetensors")]
text_encoder = SingleGPUModelBuilder( text_encoder = SingleGPUModelBuilder(
model_path=(str(checkpoint_path), *gemma_weight_paths), model_path=tuple(gemma_weight_paths),
model_class_configurator=GemmaTextEncoderConfigurator, model_class_configurator=GemmaTextEncoderConfigurator,
model_sd_ops=AV_GEMMA_TEXT_ENCODER_KEY_OPS, model_sd_ops=GEMMA_LLM_KEY_OPS,
module_ops=(GEMMA_MODEL_OPS, *module_ops_from_gemma_root(str(gemma_model_path))), module_ops=(GEMMA_MODEL_OPS, *module_ops_from_gemma_root(str(gemma_model_path))),
).build(device=torch_device, dtype=dtype) ).build(device=torch_device, dtype=dtype)
return text_encoder return text_encoder
def load_embeddings_processor(
checkpoint_path: str | Path,
device: Device = "cpu",
dtype: torch.dtype = torch.bfloat16,
) -> "EmbeddingsProcessor":
"""Load the embeddings processor (feature extractor + video/audio connectors).
Args:
checkpoint_path: Path to the LTX-2 safetensors checkpoint file
device: Device to load model on
dtype: Data type for model weights
Returns:
Loaded EmbeddingsProcessor with feature extractor and connectors
"""
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
from ltx_core.text_encoders.gemma import (
EMBEDDINGS_PROCESSOR_KEY_OPS,
EmbeddingsProcessorConfigurator,
)
torch_device = _to_torch_device(device)
return SingleGPUModelBuilder(
model_path=str(checkpoint_path),
model_class_configurator=EmbeddingsProcessorConfigurator,
model_sd_ops=EMBEDDINGS_PROCESSOR_KEY_OPS,
).build(device=torch_device, dtype=dtype)
# ============================================================================= # =============================================================================
# Combined Component Loader # Combined Component Loader
# ============================================================================= # =============================================================================
@@ -337,7 +364,7 @@ def load_model(
if text_encoder_path is None: if text_encoder_path is None:
raise ValueError("text_encoder_path must be provided when with_text_encoder=True") raise ValueError("text_encoder_path must be provided when with_text_encoder=True")
logger.debug("Loading Gemma text encoder...") logger.debug("Loading Gemma text encoder...")
text_encoder = load_text_encoder(checkpoint_path, text_encoder_path, torch_device, dtype) text_encoder = load_text_encoder(text_encoder_path, torch_device, dtype)
# Create scheduler (stateless, no loading needed) # Create scheduler (stateless, no loading needed)
scheduler = LTX2Scheduler() scheduler = LTX2Scheduler()
+34 -24
View File
@@ -27,14 +27,15 @@ from torch.optim.lr_scheduler import (
from torch.utils.data import DataLoader from torch.utils.data import DataLoader
from torchvision.transforms import functional as F # noqa: N812 from torchvision.transforms import functional as F # noqa: N812
from ltx_core.text_encoders.gemma import convert_to_additive_mask
from ltx_trainer import logger from ltx_trainer import logger
from ltx_trainer.config import LtxTrainerConfig from ltx_trainer.config import LtxTrainerConfig
from ltx_trainer.config_display import print_config from ltx_trainer.config_display import print_config
from ltx_trainer.datasets import PrecomputedDataset from ltx_trainer.datasets import PrecomputedDataset
from ltx_trainer.gpu_utils import free_gpu_memory, free_gpu_memory_context, get_gpu_memory_gb from ltx_trainer.gpu_utils import free_gpu_memory, free_gpu_memory_context, get_gpu_memory_gb
from ltx_trainer.hf_hub_utils import push_to_hub from ltx_trainer.hf_hub_utils import push_to_hub
from ltx_trainer.model_loader import load_embeddings_processor, load_text_encoder
from ltx_trainer.model_loader import load_model as load_ltx_model from ltx_trainer.model_loader import load_model as load_ltx_model
from ltx_trainer.model_loader import load_text_encoder
from ltx_trainer.progress import TrainingProgress from ltx_trainer.progress import TrainingProgress
from ltx_trainer.quantization import quantize_model from ltx_trainer.quantization import quantize_model
from ltx_trainer.timestep_samplers import SAMPLERS from ltx_trainer.timestep_samplers import SAMPLERS
@@ -320,8 +321,8 @@ class LtxvTrainer:
audio_features = conditions["prompt_embeds"] audio_features = conditions["prompt_embeds"]
mask = conditions["prompt_attention_mask"] mask = conditions["prompt_attention_mask"]
additive_mask = self._text_encoder._convert_to_additive_mask(mask, video_features.dtype) additive_mask = convert_to_additive_mask(mask, video_features.dtype)
video_embeds, audio_embeds, attention_mask = self._text_encoder.embeddings_processor.create_embeddings( video_embeds, audio_embeds, attention_mask = self._embeddings_processor.create_embeddings(
video_features, audio_features, additive_mask video_features, audio_features, additive_mask
) )
@@ -346,26 +347,31 @@ class LtxvTrainer:
@free_gpu_memory_context(after=True) @free_gpu_memory_context(after=True)
def _load_text_encoder_and_cache_embeddings(self) -> list[CachedPromptEmbeddings] | None: def _load_text_encoder_and_cache_embeddings(self) -> list[CachedPromptEmbeddings] | None:
"""Load text encoder, computes and returns validation embeddings.""" """Load text encoder + embeddings processor, compute and cache validation embeddings."""
# This method: # This method:
# 1. Loads the text encoder on GPU # 1. Loads the pure Gemma text encoder on GPU
# 2. If validation prompts are configured, computes and caches their embeddings # 2. Loads the embeddings processor (feature extractor + connectors)
# 3. Unloads the heavy Gemma model while keeping the lightweight embedding connectors # 3. If validation prompts are configured, computes and caches their embeddings
# The text encoder is kept (as self._text_encoder) but with model/tokenizer/feature_extractor # 4. Unloads the Gemma model entirely, keeps the embeddings processor for training
# set to None. Only the embedding connectors remain for use during training.
# Load text encoder on GPU # Load text encoder (pure Gemma LLM) on GPU
logger.debug("Loading text encoder...") logger.debug("Loading text encoder...")
text_encoder = load_text_encoder(
self._text_encoder = load_text_encoder(
checkpoint_path=self._config.model.model_path,
gemma_model_path=self._config.model.text_encoder_path, gemma_model_path=self._config.model.text_encoder_path,
device="cuda", device="cuda",
dtype=torch.bfloat16, dtype=torch.bfloat16,
load_in_8bit=self._config.acceleration.load_text_encoder_in_8bit, load_in_8bit=self._config.acceleration.load_text_encoder_in_8bit,
) )
# Load embeddings processor (feature extractor + connectors)
logger.debug("Loading embeddings processor...")
self._embeddings_processor = load_embeddings_processor(
checkpoint_path=self._config.model.model_path,
device="cuda",
dtype=torch.bfloat16,
)
# Cache validation embeddings if prompts are configured # Cache validation embeddings if prompts are configured
cached_embeddings = None cached_embeddings = None
if self._config.validation.prompts: if self._config.validation.prompts:
@@ -373,22 +379,26 @@ class LtxvTrainer:
cached_embeddings = [] cached_embeddings = []
with torch.inference_mode(): with torch.inference_mode():
for prompt in self._config.validation.prompts: for prompt in self._config.validation.prompts:
v_ctx_pos, a_ctx_pos, _ = self._text_encoder(prompt) pos_hs, pos_mask = text_encoder.encode(prompt)
v_ctx_neg, a_ctx_neg, _ = self._text_encoder(self._config.validation.negative_prompt) pos_out = self._embeddings_processor.process_hidden_states(pos_hs, pos_mask)
neg_hs, neg_mask = text_encoder.encode(self._config.validation.negative_prompt)
neg_out = self._embeddings_processor.process_hidden_states(neg_hs, neg_mask)
cached_embeddings.append( cached_embeddings.append(
CachedPromptEmbeddings( CachedPromptEmbeddings(
video_context_positive=v_ctx_pos.cpu(), video_context_positive=pos_out.video_encoding.cpu(),
audio_context_positive=a_ctx_pos.cpu(), audio_context_positive=pos_out.audio_encoding.cpu(),
video_context_negative=v_ctx_neg.cpu() if v_ctx_neg is not None else None, video_context_negative=neg_out.video_encoding.cpu(),
audio_context_negative=a_ctx_neg.cpu() if a_ctx_neg is not None else None, audio_context_negative=(
neg_out.audio_encoding.cpu() if neg_out.audio_encoding is not None else None
),
) )
) )
# Unload heavy components to free VRAM, keeping only the embedding connectors # Unload Gemma model and feature extractor, keep only connectors for training
self._text_encoder.model = None del text_encoder
self._text_encoder.tokenizer = None self._embeddings_processor.feature_extractor = None
self._text_encoder.feature_extractor = None
logger.debug("Validation prompt embeddings cached. Gemma model unloaded") logger.debug("Validation prompt embeddings cached. Gemma model unloaded")
return cached_embeddings return cached_embeddings
@@ -426,7 +436,7 @@ class LtxvTrainer:
self._scheduler = components.scheduler self._scheduler = components.scheduler
self._audio_vae = components.audio_vae_decoder self._audio_vae = components.audio_vae_decoder
self._vocoder = components.vocoder self._vocoder = components.vocoder
# Note: self._text_encoder was set in _load_text_encoder_and_cache_embeddings # Note: self._embeddings_processor was set in _load_text_encoder_and_cache_embeddings
# Determine initial dtype based on training mode. # Determine initial dtype based on training mode.
# Note: For FSDP + LoRA, we'll cast to FP32 later in _prepare_models_for_training() # Note: For FSDP + LoRA, we'll cast to FP32 later in _prepare_models_for_training()
@@ -37,6 +37,7 @@ if TYPE_CHECKING:
from ltx_core.model.transformer import LTXModel from ltx_core.model.transformer import LTXModel
from ltx_core.model.video_vae import VideoDecoder, VideoEncoder from ltx_core.model.video_vae import VideoDecoder, VideoEncoder
from ltx_core.text_encoders.gemma import GemmaTextEncoder from ltx_core.text_encoders.gemma import GemmaTextEncoder
from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcessor
VIDEO_SCALE_FACTORS = SpatioTemporalScaleFactors.default() VIDEO_SCALE_FACTORS = SpatioTemporalScaleFactors.default()
@@ -128,21 +129,24 @@ class ValidationSampler:
audio_decoder: "AudioDecoder | None" = None, audio_decoder: "AudioDecoder | None" = None,
vocoder: "Vocoder | None" = None, vocoder: "Vocoder | None" = None,
sampling_context: SamplingContext | None = None, sampling_context: SamplingContext | None = None,
embeddings_processor: "EmbeddingsProcessor | None" = None,
): ):
"""Initialize the validation sampler. """Initialize the validation sampler.
Args: Args:
transformer: LTX-2 transformer model transformer: LTX-2 transformer model
vae_decoder: Video VAE decoder vae_decoder: Video VAE decoder
vae_encoder: Video VAE encoder (for image/video conditioning), can be None if not needed vae_encoder: Video VAE encoder (for image/video conditioning), can be None if not needed
text_encoder: Gemma text encoder with embeddings connector (optional if cached_embeddings in config) text_encoder: Gemma text encoder (optional if cached_embeddings in config)
audio_decoder: Optional audio VAE decoder (for audio generation) audio_decoder: Optional audio VAE decoder (for audio generation)
vocoder: Optional vocoder (for audio generation) vocoder: Optional vocoder (for audio generation)
sampling_context: Optional SamplingContext for progress display during denoising sampling_context: Optional SamplingContext for progress display during denoising
embeddings_processor: Optional embeddings processor (required if text_encoder provided)
""" """
self._transformer = transformer self._transformer = transformer
self._vae_decoder = vae_decoder self._vae_decoder = vae_decoder
self._vae_encoder = vae_encoder self._vae_encoder = vae_encoder
self._text_encoder = text_encoder self._text_encoder = text_encoder
self._embeddings_processor = embeddings_processor
self._audio_decoder = audio_decoder self._audio_decoder = audio_decoder
self._vocoder = vocoder self._vocoder = vocoder
self._sampling_context = sampling_context self._sampling_context = sampling_context
@@ -677,6 +681,8 @@ class ValidationSampler:
# Validate prompt embedding source # Validate prompt embedding source
if config.cached_embeddings is None and self._text_encoder is None: if config.cached_embeddings is None and self._text_encoder is None:
raise ValueError("Either text_encoder or config.cached_embeddings must be provided") raise ValueError("Either text_encoder or config.cached_embeddings must be provided")
if config.cached_embeddings is None and self._embeddings_processor is None:
raise ValueError("embeddings_processor is required when encoding prompts on-the-fly")
def _get_prompt_embeddings( def _get_prompt_embeddings(
self, config: GenerationConfig, device: torch.device self, config: GenerationConfig, device: torch.device
@@ -697,18 +703,22 @@ class ValidationSampler:
def _encode_prompts( def _encode_prompts(
self, config: GenerationConfig, device: torch.device self, config: GenerationConfig, device: torch.device
) -> tuple[Tensor, Tensor, Tensor | None, Tensor | None]: ) -> tuple[Tensor, Tensor, Tensor | None, Tensor | None]:
"""Encode positive and negative prompts using the text encoder.""" """Encode positive and negative prompts using the text encoder + embeddings processor."""
self._text_encoder.to(device) self._text_encoder.to(device)
v_ctx_pos, a_ctx_pos, _ = self._text_encoder(config.prompt) self._embeddings_processor.to(device)
pos_hs, pos_mask = self._text_encoder.encode(config.prompt)
pos_out = self._embeddings_processor.process_hidden_states(pos_hs, pos_mask)
v_ctx_pos, a_ctx_pos = pos_out.video_encoding, pos_out.audio_encoding
v_ctx_neg, a_ctx_neg = None, None v_ctx_neg, a_ctx_neg = None, None
if config.guidance_scale != 1.0: if config.guidance_scale != 1.0:
v_ctx_neg, a_ctx_neg, _ = self._text_encoder(config.negative_prompt) neg_hs, neg_mask = self._text_encoder.encode(config.negative_prompt)
neg_out = self._embeddings_processor.process_hidden_states(neg_hs, neg_mask)
v_ctx_neg, a_ctx_neg = neg_out.video_encoding, neg_out.audio_encoding
# Move the base Gemma model to CPU but keep embeddings connectors on GPU # Move the base Gemma model to CPU
# as this module is also used during training
self._text_encoder.model.to("cpu") self._text_encoder.model.to("cpu")
if self._text_encoder.feature_extractor is not None:
self._text_encoder.feature_extractor.to("cpu")
return v_ctx_pos, a_ctx_pos, v_ctx_neg, a_ctx_neg return v_ctx_pos, a_ctx_pos, v_ctx_neg, a_ctx_neg