Automated PR - 2026-03-05
This commit is contained in:
@@ -1,7 +1,7 @@
|
|||||||
# LTX-2
|
# LTX-2
|
||||||
|
|
||||||
[](https://ltx.io)
|
[](https://ltx.io)
|
||||||
[](https://huggingface.co/Lightricks/LTX-2)
|
[](https://huggingface.co/Lightricks/LTX-2.3)
|
||||||
[](https://app.ltx.studio/ltx-2-playground/i2v)
|
[](https://app.ltx.studio/ltx-2-playground/i2v)
|
||||||
[](https://arxiv.org/abs/2601.03233)
|
[](https://arxiv.org/abs/2601.03233)
|
||||||
[](https://discord.gg/ltxplatform)
|
[](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
|
|
||||||
|
|||||||
+42
-35
@@ -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,
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
+61
-47
@@ -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(
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user