Automated PR - 2026-03-04

This commit is contained in:
sync-bot
2026-03-04 19:34:46 +00:00
parent 28c3c73fe5
commit 822ce3c4b1
73 changed files with 4984 additions and 1220 deletions
@@ -3,9 +3,8 @@ import logging
from dataclasses import replace
import torch
from tqdm import tqdm
from ltx_core.components.guiders import MultiModalGuider
from ltx_core.components.guiders import MultiModalGuider, MultiModalGuiderFactory
from ltx_core.components.noisers import Noiser
from ltx_core.components.protocols import DiffusionStepProtocol, GuiderProtocol
from ltx_core.conditioning import (
@@ -21,10 +20,10 @@ from ltx_core.guidance.perturbations import (
)
from ltx_core.model.transformer import Modality, X0Model
from ltx_core.model.video_vae import VideoEncoder
from ltx_core.text_encoders.gemma import GemmaTextEncoderModelBase
from ltx_core.text_encoders.gemma import GemmaTextEncoder
from ltx_core.tools import AudioLatentTools, LatentTools, VideoLatentTools
from ltx_core.types import AudioLatentShape, LatentState, VideoLatentShape, VideoPixelShape
from ltx_core.utils import to_denoised, to_velocity
from ltx_pipelines.utils.args import ImageConditioningInput
from ltx_pipelines.utils.media_io import decode_image, load_image_conditioning, resize_aspect_ratio_preserving
from ltx_pipelines.utils.types import (
DenoisingFunc,
@@ -46,7 +45,7 @@ def cleanup_memory() -> None:
def image_conditionings_by_replacing_latent(
images: list[tuple[str, int, float]],
images: list[ImageConditioningInput],
height: int,
width: int,
video_encoder: VideoEncoder,
@@ -54,20 +53,21 @@ def image_conditionings_by_replacing_latent(
device: torch.device,
) -> list[ConditioningItem]:
conditionings = []
for image_path, frame_idx, strength in images:
for img in images:
image = load_image_conditioning(
image_path=image_path,
image_path=img.path,
height=height,
width=width,
dtype=dtype,
device=device,
crf=img.crf,
)
encoded_image = video_encoder(image)
conditionings.append(
VideoConditionByLatentIndex(
latent=encoded_image,
strength=strength,
latent_idx=frame_idx,
strength=img.strength,
latent_idx=img.frame_idx,
)
)
@@ -75,7 +75,7 @@ def image_conditionings_by_replacing_latent(
def image_conditionings_by_adding_guiding_latent(
images: list[tuple[str, int, float]],
images: list[ImageConditioningInput],
height: int,
width: int,
video_encoder: VideoEncoder,
@@ -83,131 +83,22 @@ def image_conditionings_by_adding_guiding_latent(
device: torch.device,
) -> list[ConditioningItem]:
conditionings = []
for image_path, frame_idx, strength in images:
for img in images:
image = load_image_conditioning(
image_path=image_path,
image_path=img.path,
height=height,
width=width,
dtype=dtype,
device=device,
crf=img.crf,
)
encoded_image = video_encoder(image)
conditionings.append(
VideoConditionByKeyframeIndex(keyframes=encoded_image, frame_idx=frame_idx, strength=strength)
VideoConditionByKeyframeIndex(keyframes=encoded_image, frame_idx=img.frame_idx, strength=img.strength)
)
return conditionings
def euler_denoising_loop(
sigmas: torch.Tensor,
video_state: LatentState,
audio_state: LatentState,
stepper: DiffusionStepProtocol,
denoise_fn: DenoisingFunc,
) -> tuple[LatentState, LatentState]:
"""
Perform the joint audio-video denoising loop over a diffusion schedule.
This function iterates over all but the final value in ``sigmas`` and, at
each diffusion step, calls ``denoise_fn`` to obtain denoised video and
audio latents. The denoised latents are post-processed with their
respective denoise masks and clean latents, then passed to ``stepper`` to
advance the noisy latents one step along the diffusion schedule.
### Parameters
sigmas:
A 1D tensor of noise levels (diffusion sigmas) defining the sampling
schedule. All steps except the last element are iterated over.
video_state:
The current video :class:`LatentState`, containing the noisy latent,
its clean reference latent, and the denoising mask.
audio_state:
The current audio :class:`LatentState`, analogous to ``video_state``
but for the audio modality.
stepper:
An implementation of :class:`DiffusionStepProtocol` that updates a
latent given the current latent, its denoised estimate, the full
``sigmas`` schedule, and the current step index.
denoise_fn:
A callable implementing :class:`DenoisingFunc`. It is invoked as
``denoise_fn(video_state, audio_state, sigmas, step_index)`` and must
return a tuple ``(denoised_video, denoised_audio)``, where each element
is a tensor with the same shape as the corresponding latent.
### Returns
tuple[LatentState, LatentState]
A pair ``(video_state, audio_state)`` containing the final video and
audio latent states after completing the denoising loop.
"""
for step_idx, _ in enumerate(tqdm(sigmas[:-1])):
denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_idx)
denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent)
denoised_audio = post_process_latent(denoised_audio, audio_state.denoise_mask, audio_state.clean_latent)
video_state = replace(video_state, latent=stepper.step(video_state.latent, denoised_video, sigmas, step_idx))
audio_state = replace(audio_state, latent=stepper.step(audio_state.latent, denoised_audio, sigmas, step_idx))
return (video_state, audio_state)
def gradient_estimating_euler_denoising_loop(
sigmas: torch.Tensor,
video_state: LatentState,
audio_state: LatentState,
stepper: DiffusionStepProtocol,
denoise_fn: DenoisingFunc,
ge_gamma: float = 2.0,
) -> tuple[LatentState, LatentState]:
"""
Perform the joint audio-video denoising loop using gradient-estimation sampling.
This function is similar to :func:`euler_denoising_loop`, but applies
gradient estimation to improve the denoised estimates by tracking velocity
changes across steps. See the referenced function for detailed parameter
documentation.
### Parameters
ge_gamma:
Gradient estimation coefficient controlling the velocity correction term.
Default is 2.0. Paper: https://openreview.net/pdf?id=o2ND9v0CeK
sigmas, video_state, audio_state, stepper, denoise_fn:
See :func:`euler_denoising_loop` for parameter descriptions.
### Returns
tuple[LatentState, LatentState]
See :func:`euler_denoising_loop` for return value description.
"""
previous_audio_velocity = None
previous_video_velocity = None
def update_velocity_and_sample(
noisy_sample: torch.Tensor, denoised_sample: torch.Tensor, sigma: float, previous_velocity: torch.Tensor | None
) -> tuple[torch.Tensor, torch.Tensor]:
current_velocity = to_velocity(noisy_sample, sigma, denoised_sample)
if previous_velocity is not None:
delta_v = current_velocity - previous_velocity
total_velocity = ge_gamma * delta_v + previous_velocity
denoised_sample = to_denoised(noisy_sample, total_velocity, sigma)
return current_velocity, denoised_sample
for step_idx, _ in enumerate(tqdm(sigmas[:-1])):
denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_idx)
denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent)
denoised_audio = post_process_latent(denoised_audio, audio_state.denoise_mask, audio_state.clean_latent)
if sigmas[step_idx + 1] == 0:
return replace(video_state, latent=denoised_video), replace(audio_state, latent=denoised_audio)
previous_video_velocity, denoised_video = update_velocity_and_sample(
video_state.latent, denoised_video, sigmas[step_idx], previous_video_velocity
)
previous_audio_velocity, denoised_audio = update_velocity_and_sample(
audio_state.latent, denoised_audio, sigmas[step_idx], previous_audio_velocity
)
video_state = replace(video_state, latent=stepper.step(video_state.latent, denoised_video, sigmas, step_idx))
audio_state = replace(audio_state, latent=stepper.step(audio_state.latent, denoised_audio, sigmas, step_idx))
return (video_state, audio_state)
def noise_video_state(
output_shape: VideoPixelShape,
noiser: Noiser,
@@ -313,7 +204,10 @@ def post_process_latent(denoised: torch.Tensor, denoise_mask: torch.Tensor, clea
def modality_from_latent_state(
state: LatentState, context: torch.Tensor, sigma: float | torch.Tensor, enabled: bool = True
state: LatentState,
context: torch.Tensor,
sigma: torch.Tensor,
enabled: bool = True,
) -> Modality:
"""Create a Modality from a latent state.
Constructs a Modality object with the latent state's data, timesteps derived
@@ -322,10 +216,12 @@ def modality_from_latent_state(
return Modality(
enabled=enabled,
latent=state.latent,
sigma=sigma,
timesteps=timesteps_from_mask(state.denoise_mask, sigma),
positions=state.positions,
context=context,
context_mask=None,
attention_mask=state.attention_mask,
)
@@ -389,10 +285,10 @@ def multi_modal_guider_denoising_func(
v_context: torch.Tensor,
a_context: torch.Tensor,
transformer: X0Model,
*,
last_denoised_video: torch.Tensor | None = None,
last_denoised_audio: torch.Tensor | None = None,
) -> DenoisingFunc:
last_denoised_video = None
last_denoised_audio = None
def guider_denoising_step(
video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
) -> tuple[torch.Tensor, torch.Tensor]:
@@ -490,6 +386,43 @@ def multi_modal_guider_denoising_func(
return guider_denoising_step
def multi_modal_guider_factory_denoising_func(
video_guider_factory: MultiModalGuiderFactory,
audio_guider_factory: MultiModalGuiderFactory | None,
v_context: torch.Tensor,
a_context: torch.Tensor,
transformer: X0Model,
) -> DenoisingFunc:
"""Resolve guiders per step via factory.build_from_sigma, then multi_modal_guider_denoising_func."""
last_denoised_video: torch.Tensor | None = None
last_denoised_audio: torch.Tensor | None = None
sigma_vals_cached: list[float] | None = None
def guider_denoising_step(
video_state: LatentState, audio_state: LatentState, sigmas: torch.Tensor, step_index: int
) -> tuple[torch.Tensor, torch.Tensor]:
nonlocal last_denoised_video, last_denoised_audio, sigma_vals_cached
if sigma_vals_cached is None:
sigma_vals_cached = sigmas.detach().cpu().tolist()
sigma_val = sigma_vals_cached[step_index]
video_guider = video_guider_factory.build_from_sigma(sigma_val)
audio_guider = (audio_guider_factory or video_guider_factory).build_from_sigma(sigma_val)
denoise_fn = multi_modal_guider_denoising_func(
video_guider,
audio_guider,
v_context,
a_context,
transformer,
last_denoised_video=last_denoised_video,
last_denoised_audio=last_denoised_audio,
)
denoised_video, denoised_audio = denoise_fn(video_state, audio_state, sigmas, step_index)
last_denoised_video, last_denoised_audio = denoised_video, denoised_audio
return denoised_video, denoised_audio
return guider_denoising_step
def denoise_audio_video( # noqa: PLR0913
output_shape: VideoPixelShape,
conditionings: list[ConditioningItem],
@@ -540,6 +473,57 @@ def denoise_audio_video( # noqa: PLR0913
return video_state, audio_state
def denoise_video_only( # noqa: PLR0913
output_shape: VideoPixelShape,
conditionings: list[ConditioningItem],
noiser: Noiser,
sigmas: torch.Tensor,
stepper: DiffusionStepProtocol,
denoising_loop_fn: DenoisingLoopFunc,
components: PipelineComponents,
dtype: torch.dtype,
device: torch.device,
noise_scale: float = 1.0,
initial_video_latent: torch.Tensor | None = None,
initial_audio_latent: torch.Tensor | None = None,
) -> LatentState:
video_state, video_tools = noise_video_state(
output_shape=output_shape,
noiser=noiser,
conditionings=conditionings,
components=components,
dtype=dtype,
device=device,
noise_scale=noise_scale,
initial_latent=initial_video_latent,
)
audio_state, _ = noise_audio_state(
output_shape=output_shape,
noiser=noiser,
conditionings=[],
components=components,
dtype=dtype,
device=device,
noise_scale=0.0,
initial_latent=initial_audio_latent,
)
audio_state = replace(audio_state, denoise_mask=torch.zeros_like(audio_state.denoise_mask))
video_state, audio_state = denoising_loop_fn(
sigmas,
video_state,
audio_state,
stepper,
)
video_state = video_tools.clear_conditioning(video_state)
video_state = video_tools.unpatchify(video_state)
return video_state
_UNICODE_REPLACEMENTS = str.maketrans("\u2018\u2019\u201c\u201d\u2014\u2013\u00a0\u2032\u2212", "''\"\"-- '-")
@@ -555,7 +539,7 @@ def clean_response(text: str) -> str:
def generate_enhanced_prompt(
text_encoder: GemmaTextEncoderModelBase,
text_encoder: GemmaTextEncoder,
prompt: str,
image_path: str | None = None,
image_long_side: int = 896,