Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 22 additions & 0 deletions docs/source/en/api/models/ltx2_diffusion_decoder.md
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,28 @@ reproduce the untiled result exactly; the default tile and overlap sizes match t
Neighborhood attention rejects any grid smaller than its kernel, so a trailing remnant tile is merged into its
neighbor rather than decoded on its own.

## Keyframe-aware decoding

A decode can be anchored on *keyframe planes*: single-frame latents at known pixel frames, for example the generated
keyframe slots of the LTX-2.5 SDR-To-HDR IC-LoRA. Each plane must be the latent of a standalone one-frame clip
(the VAE is causal, so a frame encoded inside a clip is a different latent), denormalized like the video latents.

```python
video = decoder.decode(
latents, # (B, C, F, H, W)
generator=torch.Generator("cuda").manual_seed(0),
keyframe_latents=keyframe_latents, # (B, C, P, H, W)
keyframe_frame_indices=torch.tensor([24, 48, 72, 96]), # pixel frame of each plane
).sample
```

The planes go through the same weights as the video. In every attention layer each video position also attends to
the same spatial window on its two nearest planes, and each plane to the same window on its two nearest frames. This
joint attention runs on a built-in PyTorch implementation whatever attention processor is set, since neither
FlexAttention's block mask nor NATTEN expresses it. Checkpoints trained for keyframe decoding also carry a learned
tag added to the plane latents, which `decoder_keyframe_type_embedding=True` creates as `decoder.type_emb`. With
tiling enabled each temporal tile keeps the planes inside it plus the nearest plane on each side.

## LTX2VideoDiffusionDecoderModel

[[autodoc]] LTX2VideoDiffusionDecoderModel
Expand Down
4 changes: 4 additions & 0 deletions scripts/convert_ltx2_to_diffusers.py
Original file line number Diff line number Diff line change
Expand Up @@ -895,6 +895,10 @@ def get_ltx2_diffusion_video_vae_config(version: str) -> tuple[dict[str, Any], d
def convert_ltx2_diffusion_video_vae(original_state_dict: dict[str, Any], version: str) -> dict[str, Any]:
config, rename_dict, special_keys_remap = get_ltx2_diffusion_video_vae_config(version)
diffusers_config = config["diffusers_config"]
# Checkpoints trained for keyframe-aware decoding carry the keyframe stream's learned tag, `decoder.type_emb`
# (`(latent_channels,)`); older ones do not. It keeps its name, so only the config has to know it is there.
if "decoder.type_emb" in original_state_dict:
diffusers_config = {**diffusers_config, "decoder_keyframe_type_embedding": True}

with init_empty_weights():
vae = LTX2VideoDiffusionDecoderModel.from_config(diffusers_config)
Expand Down
930 changes: 911 additions & 19 deletions src/diffusers/models/autoencoders/ltx2_diffusion_decoder.py

Large diffs are not rendered by default.

32 changes: 32 additions & 0 deletions src/diffusers/pipelines/ltx2/dfr_layout.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,6 +144,38 @@ def resolve_canvas(num_frames: int, temporal_compression_ratio: int = 8) -> tupl
return content_padded + 1, segment, positions


def resolve_seam_positions(
num_frames: int, high_quality: bool = False, temporal_compression_ratio: int = 8
) -> list[int]:
"""
Keyframe seam positions of a clip decoded in one window: the [`resolve_canvas`] keyframes that fall inside it.

Unlike [`resolve_canvas`] the canvas is not padded: positions at or past `num_frames` are dropped, so the last
segment may be shorter than the others. This is where the LTX-2.5 SDR-to-HDR IC-LoRA places its seam keyframes.

Args:
num_frames (`int`):
Pixel frame count of the clip. Must satisfy `(num_frames - 1) % temporal_compression_ratio == 0`. A single
frame has no seam and returns `[]`, where [`resolve_canvas`] would reject it.
high_quality (`bool`, defaults to `False`):
Whether the clip runs on a frame-doubled `2 * num_frames - 1` grid. The positions are then computed on
`num_frames` and doubled onto that grid.
temporal_compression_ratio (`int`, defaults to `8`):
The VAE's temporal compression ratio.

Returns:
`list[int]`: the seam pixel frames, ascending; frame 0 is never one. 9 and 17 frames have none (the first seam
is at 24), 97 frames give `[32, 64, 96]` and 121 frames `[24, 48, 72, 96, 120]`.
"""
if num_frames < 2:
return []
_, _, positions = resolve_canvas(num_frames, temporal_compression_ratio)
positions = [position for position in positions if position < num_frames]
if high_quality:
positions = [2 * position for position in positions]
return positions


def pixel_to_latent_index(pixel_frame: int, temporal_compression_ratio: int = 8) -> int:
"""Map a pixel frame sitting on a latent border to its latent index."""
if pixel_frame < 0:
Expand Down
Loading
Loading