From 57f0bad4d3ca8e22945cd60932dab66b0fca409d Mon Sep 17 00:00:00 2001 From: christopher5106 Date: Tue, 6 Oct 2026 17:44:36 +0200 Subject: [PATCH 1/4] [LTX-2.5] Add ACEScct colour transforms and HLG/EXR export for SDR-To-HDR Port the colour pipeline of the LTX-2.5 SDR-To-HDR IC-LoRA from the Lightricks LTX-2 reference (ltx_core.hdr, ltx_core.color, ltx_pipelines media_io): - image_processor: sRGB EOTF, Bradford Rec.709/AP1/Rec.2020 matrices, ACEScct encode/decode, the srgb_gamma/srgb/acescg/acescct input transforms and the ACEScct -> ACEScg/Rec.709 linear output transform, as pure torch functions. LTX2VideoHDRProcessor gains hdr_transform="acescct" with input_colorspace / output_colorspace options; the LogC3 default is unchanged. - export_utils: encode_hdr_tensor_to_hlg_mp4 (BT.2020 HLG 10-bit HEVC via PyAV/libx265) and save_exr_frame / export_to_exr_sequence (half-float ZIP EXR with chromaticities and colorSpace), behind a new optional is_openexr_available() guard. - tests: exact values computed with colour-science, round trips, HLG stream properties and EXR read-back. Co-Authored-By: Claude Opus 5.5 --- src/diffusers/pipelines/ltx2/export_utils.py | 338 ++++++++++++- .../pipelines/ltx2/image_processor.py | 315 +++++++++++- src/diffusers/utils/__init__.py | 1 + src/diffusers/utils/import_utils.py | 5 + tests/pipelines/ltx2/test_ltx2_hdr_color.py | 470 ++++++++++++++++++ 5 files changed, 1114 insertions(+), 15 deletions(-) create mode 100644 tests/pipelines/ltx2/test_ltx2_hdr_color.py diff --git a/src/diffusers/pipelines/ltx2/export_utils.py b/src/diffusers/pipelines/ltx2/export_utils.py index b2206a79842c..be004023f3d2 100644 --- a/src/diffusers/pipelines/ltx2/export_utils.py +++ b/src/diffusers/pipelines/ltx2/export_utils.py @@ -13,14 +13,18 @@ # See the License for the specific language governing permissions and # limitations under the License. +import math +import os from fractions import Fraction from pathlib import Path from typing import Callable import numpy as np import torch +import torch.nn.functional as F -from ...utils import is_av_available +from ...utils import is_av_available, is_openexr_available +from .image_processor import ACESCG_TO_REC2020, REC709_TO_REC2020 _CAN_USE_AV = is_av_available() @@ -108,3 +112,335 @@ def simple_tone_map(x: np.ndarray) -> np.ndarray: container.mux(packet) finally: container.close() + + +# ARIB STD-B67 HLG OETF constants, as used by the LTX-2 reference (`colour` `CONSTANTS_ARIBSTDB67`). +_HLG_A = 0.17883277 +_HLG_B = 0.28466892 +_HLG_C = 0.55991073 + +# FFmpeg colour tags (`AVCOL_PRI_BT2020`, `AVCOL_TRC_ARIB_STD_B67`, `AVCOL_SPC_BT2020_NCL`, `AVCOL_RANGE_MPEG`). +_AV_COLOR_PRIMARIES_BT2020 = 9 +_AV_COLOR_TRC_ARIB_STD_B67 = 18 +_AV_COLORSPACE_BT2020_NCL = 9 +_AV_COLOR_RANGE_MPEG = 1 + +# Full-range RGB -> Y'CbCr matrix with the BT.2020 non-constant-luminance weights (Kr = 0.2627, Kb = 0.0593). The +# reference computes it as the float64 inverse of `colour.matrix_YCbCr(WEIGHTS_YCBCR["ITU-R BT.2020"])` stored as +# float32; these are those float32 values. +_RGB_TO_YCBCR_BT2020 = ( + (0.26269999146461487, 0.6779999732971191, 0.059300001710653305), + (-0.13963006436824799, -0.3603699505329132, 0.5), + (0.5, -0.45978569984436035, -0.04021429643034935), +) + +_HLG_PRIMARIES_TO_REC2020 = {"rec709": REC709_TO_REC2020, "acescg": ACESCG_TO_REC2020} + + +def _hlg_inverse_oetf(signal: float) -> float: + r"""Inverse HLG OETF (ITU-R BT.2100 reference constants): HLG signal `[0, 1]` -> scene-linear `[0, 1]`.""" + a = _HLG_A + b = 1.0 - 4.0 * a + c = 0.5 - a * math.log(4.0 * a) + linear = (signal / 0.5) ** 2 if signal <= 0.5 else math.exp((signal - c) / a) + b + return linear / 12.0 + + +def _hlg_oetf(x: torch.Tensor) -> torch.Tensor: + r"""HLG OETF (ARIB STD-B67): scene-linear `[0, 1]` -> HLG signal, clamped to `[0, 1]`.""" + return torch.where( + x <= 1.0 / 12.0, + torch.sqrt((3.0 * x).clamp(min=0.0)), + _HLG_A * torch.log((12.0 * x - _HLG_B).clamp(min=1e-12)) + _HLG_C, + ).clamp(0.0, 1.0) + + +def _linear_to_hlg_signal( + rgb_linear: torch.Tensor, primaries_matrix: torch.Tensor, white_x: float, roll_k: float +) -> torch.Tensor: + r"""Scene-linear `(..., 3, H, W)` RGB -> Rec.2020 HLG signal `[0, 1]`, with diffuse white mapped to `white_x`.""" + lin = torch.nan_to_num( + torch.einsum("...chw,dc->...dhw", rgb_linear, primaries_matrix).clamp(min=0.0), + nan=0.0, + neginf=0.0, + ) + # Diffuse white (linear 1.0) maps to `white_x`; highlights roll off exponentially toward 1.0. + x = torch.where( + lin <= 1.0, + lin * white_x, + 1.0 - (1.0 - white_x) * torch.exp(-roll_k * (lin - 1.0)), + ) + return _hlg_oetf(x) + + +def _rgb_to_yuv420p10_bt2020_limited(rgb: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + r""" + Float RGB `[0, 1]` `(F, 3, H, W)` -> planar 10-bit limited-range BT.2020 NCL Y, U, V code values (4:2:0). + + Returns `int32` tensors with values in `[0, 1023]`. + """ + _, _, height, width = rgb.shape + if height % 2 != 0 or width % 2 != 0: + raise ValueError(f"HLG export requires an even frame height and width, got {height}x{width}.") + matrix = torch.as_tensor(_RGB_TO_YCBCR_BT2020, dtype=torch.float32).to(device=rgb.device, dtype=rgb.dtype) + yuv = (rgb.movedim(-3, -1).flatten(-3, -2) @ matrix.T).unflatten(-2, (height, width)).movedim(-1, -3) + y = yuv[:, :1] + uv = F.avg_pool2d(yuv[:, 1:3].contiguous(), kernel_size=2, stride=2) + # Limited ("MPEG") range at 10 bits: Y = (219 * E' + 16) * 4, Cb/Cr = (224 * E' + 128) * 4. + y = y * (219 * 4) + 16 * 4 + uv = uv * (224 * 4) + 128 * 4 + y = y[:, 0].round().clamp(0, 1023).to(torch.int32) + u = uv[:, 0].round().clamp(0, 1023).to(torch.int32) + v = uv[:, 1].round().clamp(0, 1023).to(torch.int32) + return y, u, v + + +def _x265_params(threads: int, width: int, height: int) -> str: + r"""libx265 `x265-params` for a BT.2020 HLG `hvc1` MP4, as in the reference implementation.""" + params = ( + "colorprim=bt2020:transfer=arib-std-b67:colormatrix=bt2020nc:range=limited:" + f"repeat-headers=1:info=0:pools={threads}" + ) + if width <= 32 and height <= 32: + # With a single CTU per axis, frame-threading and B-frames can flush packets that the MP4 muxer rejects on + # very short clips. + return f"{params}:frame-threads=1:bframes=0:lookahead=0" + return f"{params}:frame-threads=4" + + +def encode_hdr_tensor_to_hlg_mp4( + frames: torch.Tensor | np.ndarray, + output_mp4: str | Path, + frame_rate: float, + primaries: str = "rec709", + white_signal: float = 0.75, + rolloff_k: float | None = None, + crf: int = 12, + preset: str = "ultrafast", + thread_count: int = 0, + device: str | torch.device | None = None, +) -> None: + r""" + Encodes a scene-linear HDR tensor to a BT.2020 / HLG / 10-bit HEVC `.mp4` file, following the LTX-2 reference HDR + export. + + Each frame is converted to Rec.2020 primaries; diffuse white (linear `1.0`) is mapped to the HLG signal + `white_signal` (`0.75` is the ITU-R BT.2408 HDR reference white) and brighter values roll off exponentially toward + the HLG peak, with a slope that is continuous at diffuse white. The ARIB STD-B67 OETF is then applied, followed by + a conversion to 10-bit limited-range BT.2020 non-constant-luminance Y'CbCr 4:2:0. The result is encoded with + `libx265` (`yuv420p10le`, `hvc1` tag) and tagged as BT.2020 primaries, ARIB STD-B67 transfer, BT.2020 NCL matrix + and limited range. No mastering-display or content-light-level metadata is written. + + Requires a PyAV build whose FFmpeg includes `libx265`. + + Args: + frames (`torch.Tensor` or `np.ndarray`): + Scene-linear HDR frames of shape `(F, H, W, 3)` with values in `[0, ∞)`, for example the output of + [`LTX2VideoHDRProcessor.postprocess_hdr_video`] for a single video. `H` and `W` must be even. + output_mp4 (`str` or `pathlib.Path`): + Output MP4 path. + frame_rate (`float`): + Frame rate for the output video. + primaries (`str`, *optional*, defaults to `"rec709"`): + Primaries of `frames`: `"rec709"` or `"acescg"`. + white_signal (`float`, *optional*, defaults to `0.75`): + HLG signal level that diffuse white (linear `1.0`) is mapped to. + rolloff_k (`float`, *optional*): + Exponential highlight roll-off rate. Defaults to `white_x / (1 - white_x)`, where `white_x` is the + scene-linear value of `white_signal`, which makes the mapping C1-continuous at linear `1.0`. + crf (`int`, *optional*, defaults to `12`): + libx265 CRF quality factor. Lower values produce higher quality. + preset (`str`, *optional*, defaults to `"ultrafast"`): + libx265 preset. + thread_count (`int`, *optional*, defaults to `0`): + libx265 thread pool size. `0` uses the number of CPUs, capped at 16. + device (`str` or `torch.device`, *optional*): + Device for the colour conversion. Defaults to the device of `frames` (CPU for NumPy input). + """ + if "libx265" not in av.codecs_available: + raise RuntimeError( + "HLG export requires the `libx265` encoder, but the FFmpeg build used by PyAV does not include it." + ) + if primaries not in _HLG_PRIMARIES_TO_REC2020: + raise ValueError(f"Unsupported primaries {primaries!r}. Expected 'rec709' or 'acescg'.") + if not 0.0 < white_signal < 1.0: + raise ValueError(f"`white_signal` must be in (0, 1), got {white_signal}.") + + frames = torch.as_tensor(frames) if isinstance(frames, np.ndarray) else frames.detach() + if frames.ndim != 4 or frames.shape[-1] != 3: + raise ValueError(f"Expected `frames` of shape (F, H, W, 3), got {tuple(frames.shape)}.") + num_frames, height, width, _ = frames.shape + if num_frames == 0: + raise ValueError("No HDR frames to encode.") + if height % 2 != 0 or width % 2 != 0: + raise ValueError(f"HLG export requires an even frame height and width, got {height}x{width}.") + + device = frames.device if device is None else torch.device(device) + primaries_matrix = torch.as_tensor(_HLG_PRIMARIES_TO_REC2020[primaries], dtype=torch.float32).to(device) + white_x = _hlg_inverse_oetf(white_signal) + roll_k = rolloff_k if rolloff_k is not None else white_x / (1.0 - white_x) + threads = thread_count if thread_count > 0 else max(1, min(os.cpu_count() or 8, 16)) + + output_mp4 = Path(output_mp4) + container = av.open(str(output_mp4), mode="w", options={"movflags": "+faststart"}) + try: + stream = container.add_stream("libx265", rate=Fraction(frame_rate).limit_denominator(1000)) + stream.width = width + stream.height = height + stream.pix_fmt = "yuv420p10le" + stream.codec_tag = "hvc1" + stream.options = {"crf": str(crf), "preset": preset, "x265-params": _x265_params(threads, width, height)} + codec_context = stream.codec_context + codec_context.thread_count = threads + codec_context.thread_type = "FRAME" + codec_context.color_primaries = _AV_COLOR_PRIMARIES_BT2020 + codec_context.color_trc = _AV_COLOR_TRC_ARIB_STD_B67 + codec_context.colorspace = _AV_COLORSPACE_BT2020_NCL + codec_context.color_range = _AV_COLOR_RANGE_MPEG + + for index in range(num_frames): + rgb = frames[index : index + 1].to(device=device, dtype=torch.float32).movedim(-1, -3) + hlg = _linear_to_hlg_signal(rgb, primaries_matrix, white_x, roll_k) + planes_yuv = [plane[0].cpu().numpy().astype(np.uint16) for plane in _rgb_to_yuv420p10_bt2020_limited(hlg)] + + frame = av.VideoFrame(width, height, "yuv420p10le") + for plane, src in zip(frame.planes, planes_yuv): + dest = np.frombuffer(plane, dtype=np.uint16).reshape(plane.height, plane.line_size // 2) + dest[:, : src.shape[1]] = src + frame.colorspace = _AV_COLORSPACE_BT2020_NCL + frame.color_range = _AV_COLOR_RANGE_MPEG + for packet in stream.encode(frame): + container.mux(packet) + + for packet in stream.encode(): + container.mux(packet) + except BaseException: + container.close() + output_mp4.unlink(missing_ok=True) + raise + container.close() + + +# OpenEXR `chromaticities` (R, G, B and white point xy) and `colorSpace` tags of the reference EXR writer, per EXR +# colour space (`ltx_pipelines` `EXRColorSpace`). ACEScct frames are tagged with AP1 chromaticities. +_EXR_CHROMATICITIES = { + "rec709": (0.64, 0.33, 0.30, 0.60, 0.15, 0.06, 0.3127, 0.3290), + "acescg": (0.713, 0.293, 0.165, 0.830, 0.128, 0.044, 0.32168, 0.33767), +} +_EXR_COLORSPACES = { + "srgb_linear": ("rec709", "sRGB"), + "acescg": ("acescg", "ACEScg"), + "acescct": ("acescg", "ACEScct"), +} + + +def _import_openexr(): + if not is_openexr_available(): + raise ImportError("OpenEXR is required to write EXR frames. You can install it with `pip install OpenEXR`.") + import OpenEXR + + # `OpenEXR.File` was added in OpenEXR 3.3; older bindings only expose the legacy `OutputFile` API. + if not hasattr(OpenEXR, "File"): + raise ImportError( + "OpenEXR>=3.3 is required to write EXR frames. You can upgrade it with `pip install -U OpenEXR`." + ) + return OpenEXR + + +def save_exr_frame( + frame: torch.Tensor | np.ndarray, + output_exr: str | Path, + primaries: str = "rec709", + color_space: str = "sRGB", + half: bool = True, +) -> None: + r""" + Saves a single RGB frame as an OpenEXR file tagged with its colour space, following the LTX-2 reference EXR writer: + a scanline image with `R`, `G` and `B` channels, ZIP compression, and `chromaticities` and `colorSpace` header + attributes. + + Requires the `OpenEXR` package (`pip install OpenEXR`, version 3.3 or later). + + Args: + frame (`torch.Tensor` or `np.ndarray`): + Float frame of shape `(H, W, 3)` or `(3, H, W)`. + output_exr (`str` or `pathlib.Path`): + Output EXR path. + primaries (`str`, *optional*, defaults to `"rec709"`): + Colour primaries written to the `chromaticities` attribute: `"rec709"` or `"acescg"` (AP1). + color_space (`str`, *optional*, defaults to `"sRGB"`): + Value of the `colorSpace` string attribute, which describes the encoding (e.g. `"sRGB"` or `"ACEScg"` for + scene-linear values, `"ACEScct"` for log codes). + half (`bool`, *optional*, defaults to `True`): + Write 16-bit half floats. When `False`, 32-bit floats are written, unless `frame` is already float16. + """ + OpenEXR = _import_openexr() + if primaries not in _EXR_CHROMATICITIES: + raise ValueError(f"Unsupported primaries {primaries!r}. Expected 'rec709' or 'acescg'.") + + if isinstance(frame, torch.Tensor): + use_half = half or frame.dtype == torch.float16 + frame = frame.detach().cpu().float().numpy() + else: + use_half = half or frame.dtype == np.float16 + frame = np.asarray(frame, dtype=np.float32) + if frame.ndim == 3 and frame.shape[0] == 3: + frame = frame.transpose(1, 2, 0) + if frame.ndim != 3 or frame.shape[-1] != 3: + raise ValueError(f"Expected `frame` of shape (H, W, 3) or (3, H, W), got {frame.shape}.") + frame = frame.astype(np.float16 if use_half else np.float32) + + header = { + "type": OpenEXR.scanlineimage, + "compression": OpenEXR.ZIP_COMPRESSION, + "chromaticities": _EXR_CHROMATICITIES[primaries], + "colorSpace": color_space, + } + channels = {name: np.ascontiguousarray(frame[..., index]) for index, name in enumerate("RGB")} + with OpenEXR.File(header, channels) as exr_file: + exr_file.write(str(output_exr)) + + +def export_to_exr_sequence( + frames: torch.Tensor | np.ndarray, + output_dir: str | Path, + exr_colorspace: str = "acescg", + half: bool = True, +) -> list[str]: + r""" + Saves HDR frames as a directory of OpenEXR files named `frame_00000.exr`, `frame_00001.exr`, ..., following the + LTX-2 reference HDR export. Each file is written with [`~pipelines.ltx2.export_utils.save_exr_frame`]. + + Requires the `OpenEXR` package (`pip install OpenEXR`, version 3.3 or later). + + Args: + frames (`torch.Tensor` or `np.ndarray`): + Frames of shape `(F, H, W, 3)` in the colour space given by `exr_colorspace`, for example the output of + [`LTX2VideoHDRProcessor.postprocess_hdr_video`] for a single video with the matching `output_colorspace`. + output_dir (`str` or `pathlib.Path`): + Output directory. It is created if it does not exist. + exr_colorspace (`str`, *optional*, defaults to `"acescg"`): + Colour space of `frames`, which sets the EXR tags: `"acescg"` (scene-linear ACEScg, AP1 chromaticities, + `colorSpace="ACEScg"`), `"srgb_linear"` (scene-linear Rec.709, Rec.709 chromaticities, `colorSpace="sRGB"`) + or `"acescct"` (ACEScct codes, AP1 chromaticities, `colorSpace="ACEScct"`). + half (`bool`, *optional*, defaults to `True`): + Write 16-bit half floats. + + Returns: + `list[str]`: Paths of the written EXR files. + """ + _import_openexr() + if exr_colorspace not in _EXR_COLORSPACES: + raise ValueError(f"Unsupported EXR colorspace {exr_colorspace!r}. Expected one of {tuple(_EXR_COLORSPACES)}.") + if frames.ndim != 4 or frames.shape[-1] != 3: + raise ValueError(f"Expected `frames` of shape (F, H, W, 3), got {tuple(frames.shape)}.") + primaries, color_space = _EXR_COLORSPACES[exr_colorspace] + + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + paths = [] + for index, frame in enumerate(frames): + path = output_dir / f"frame_{index:05d}.exr" + save_exr_frame(frame, path, primaries=primaries, color_space=color_space, half=half) + paths.append(str(path)) + return paths diff --git a/src/diffusers/pipelines/ltx2/image_processor.py b/src/diffusers/pipelines/ltx2/image_processor.py index a25660073943..feef8cfb6d20 100644 --- a/src/diffusers/pipelines/ltx2/image_processor.py +++ b/src/diffusers/pipelines/ltx2/image_processor.py @@ -13,10 +13,12 @@ # limitations under the License. import numpy as np +import PIL.Image import torch import torch.nn.functional as F from ...configuration_utils import register_to_config +from ...image_processor import is_valid_image, is_valid_image_imagelist from ...utils import logging from ...video_processor import VideoProcessor @@ -24,6 +26,194 @@ logger = logging.get_logger(__name__) # pylint: disable=invalid-name +# ACEScct constants (Academy S-2016-001), ported from `ltx_core.hdr` in the Lightricks LTX-2 reference. +ACESCCT_A = 10.5402377416545 +ACESCCT_B = 0.0729055341958355 +ACESCCT_X_BRK = 0.0078125 +ACESCCT_Y_BRK = 0.155251141552511 +ACESCCT_LOG_M = 17.52 +ACESCCT_LOG_B = 9.72 + +# IEC 61966-2-1 sRGB EOTF constants. +_SRGB_A = 0.055 +_SRGB_LINEAR_THRESHOLD = 0.04045 +_SRGB_LINEAR_SLOPE = 12.92 +_SRGB_GAMMA = 2.4 + +# Linear RGB -> RGB primaries matrices (Bradford chromatic adaptation), applied as `out[d] = sum_c M[d][c] * in[c]`. +# The LTX-2 reference (`ltx_core.color.primaries`) computes them at import time with colour-science +# (`colour.matrix_RGB_to_RGB(..., chromatic_adaptation_transform="Bradford")`) and stores them as float32; +# `REC709_TO_ACESCG` is the float64 inverse of the float32 `ACESCG_TO_REC709`. The values below are those float32 +# matrices (colour-science 0.4.4), written out in full so that no colour-science dependency is needed. +ACESCG_TO_REC709 = ( + (1.7050509452819824, -0.6217921376228333, -0.08325887471437454), + (-0.13025641441345215, 1.1408047676086426, -0.010548318736255169), + (-0.02400335669517517, -0.1289689689874649, 1.1529723405838013), +) +REC709_TO_ACESCG = ( + (0.6130974292755127, 0.33952316641807556, 0.04737945273518562), + (0.07019372284412384, 0.9163538813591003, 0.013452397659420967), + (0.0206155925989151, 0.10956976562738419, 0.8698146343231201), +) +REC709_TO_REC2020 = ( + (0.6274039149284363, 0.3292830288410187, 0.04331306740641594), + (0.06909728795289993, 0.9195404052734375, 0.011362315155565739), + (0.016391439363360405, 0.08801330626010895, 0.8955952525138855), +) +ACESCG_TO_REC2020 = ( + (1.025824785232544, -0.020053191110491753, -0.005771556869149208), + (-0.0022343695163726807, 1.0045864582061768, -0.002352132461965084), + (-0.0050133513286709785, -0.025290071964263916, 1.0303034782409668), +) + +# Input colour spaces accepted by the ACEScct input transform (`ltx_pipelines` `HDRICLoraInputColorSpace`). +ACESCCT_INPUT_COLORSPACES = ("srgb_gamma", "srgb", "acescg", "acescct") +# Output colour spaces of the ACEScct output transform: scene-linear ACEScg (AP1), scene-linear Rec.709, or the raw +# ACEScct codes. +ACESCCT_OUTPUT_COLORSPACES = ("rec709", "acescg", "acescct") + + +def apply_primaries_matrix(video: torch.Tensor, matrix) -> torch.Tensor: + r""" + Apply a 3x3 linear primaries matrix to an RGB video or image tensor. + + The channel axis is `dim=1` for 5D `(B, C, F, H, W)` inputs and `dim=-3` otherwise (`(..., C, H, W)`), matching the + reference implementation. + + Args: + video (`torch.Tensor`): + Linear RGB tensor of shape `(B, 3, F, H, W)` or `(..., 3, H, W)`. + matrix (`torch.Tensor` or nested `tuple` of `float`): + The 3x3 matrix `M`, applied as `out[d] = sum_c M[d][c] * in[c]`. + + Returns: + `torch.Tensor`: The converted tensor, with the same shape, device and dtype as `video`. + """ + matrix = torch.as_tensor(matrix, dtype=torch.float32).to(device=video.device, dtype=video.dtype) + if video.ndim == 5: + return torch.einsum("bcfhw,dc->bdfhw", video, matrix) + return torch.einsum("...chw,dc->...dhw", video, matrix) + + +def srgb_eotf_to_linear(srgb: torch.Tensor) -> torch.Tensor: + r""" + sRGB-encoded `[0, 1]` code values to display-linear Rec.709 light (IEC 61966-2-1 EOTF). + + Inputs are cast to float32 and clamped to `[0, 1]` first. + + Args: + srgb (`torch.Tensor`): sRGB-encoded values. + + Returns: + `torch.Tensor`: Linear Rec.709 values in `[0, 1]`, as float32. + """ + x = torch.clamp(srgb.float(), 0.0, 1.0) + return torch.where( + x <= _SRGB_LINEAR_THRESHOLD, + x / _SRGB_LINEAR_SLOPE, + torch.pow((x + _SRGB_A) / (1.0 + _SRGB_A), _SRGB_GAMMA), + ) + + +def acescct_encode(linear_acescg: torch.Tensor) -> torch.Tensor: + r""" + Encode scene-linear ACEScg (AP1) values to ACEScct `[0, 1]`. + + Follows the reference implementation, which differs from the ACEScct specification in two places: negative inputs + are clamped to `0` before encoding, and the output is clamped to `[0, 1]` (so linear values above `2 ** (17.52 - + 9.72) ~= 222.86` are clipped). + + Args: + linear_acescg (`torch.Tensor`): Scene-linear ACEScg values. + + Returns: + `torch.Tensor`: ACEScct codes in `[0, 1]`. + """ + x = torch.clamp(linear_acescg, min=0.0) + log_part = (torch.log2(torch.clamp(x, min=1e-12)) + ACESCCT_LOG_B) / ACESCCT_LOG_M + lin_part = ACESCCT_A * x + ACESCCT_B + return torch.clamp(torch.where(x > ACESCCT_X_BRK, log_part, lin_part), 0.0, 1.0) + + +def acescct_decode(acescct: torch.Tensor) -> torch.Tensor: + r""" + Decode ACEScct codes to scene-linear ACEScg (AP1) values. + + The input is clamped to `[0, 1]` first. The linear toe is kept as is, so a code of `0` decodes to a small negative + value (`-0.0069169`), as in the reference implementation. + + Args: + acescct (`torch.Tensor`): ACEScct codes. + + Returns: + `torch.Tensor`: Scene-linear ACEScg values. + """ + ct = torch.clamp(acescct, 0.0, 1.0) + lin_from_log = torch.pow(2.0, ct * ACESCCT_LOG_M - ACESCCT_LOG_B) + lin_from_lin = (ct - ACESCCT_B) / ACESCCT_A + return torch.where(ct > ACESCCT_Y_BRK, lin_from_log, lin_from_lin) + + +def to_acescct(video: torch.Tensor, input_colorspace: str = "srgb_gamma") -> torch.Tensor: + r""" + Input transform of the LTX-2.5 SDR-To-HDR IC-LoRA: map RGB values to the ACEScct `[0, 1]` working space. + + The supported input colour spaces are those of the reference implementation: + + - `"srgb_gamma"`: sRGB-encoded Rec.709 in `[0, 1]` (e.g. an 8-bit video divided by 255). The sRGB EOTF is applied, + then the Rec.709 -> AP1 matrix, then the ACEScct encoding. + - `"srgb"`: scene-linear Rec.709 (e.g. a linear EXR plate). Same as `"srgb_gamma"` without the EOTF. Note that the + reference also applies this mode to 8-bit video divided by 255, i.e. without linearizing it. + - `"acescg"`: scene-linear ACEScg (AP1). Only the ACEScct encoding is applied. + - `"acescct"`: values that are already ACEScct codes. They are only clamped to `[0, 1]`. + + Args: + video (`torch.Tensor`): + RGB tensor of shape `(B, 3, F, H, W)` or `(..., 3, H, W)`. + input_colorspace (`str`, *optional*, defaults to `"srgb_gamma"`): + One of `"srgb_gamma"`, `"srgb"`, `"acescg"` or `"acescct"`. + + Returns: + `torch.Tensor`: ACEScct codes in `[0, 1]`, as float32, with the same shape as `video`. + """ + if input_colorspace not in ACESCCT_INPUT_COLORSPACES: + raise ValueError( + f"Unsupported input colorspace {input_colorspace!r}. Expected one of {ACESCCT_INPUT_COLORSPACES}." + ) + video = video.float() + if input_colorspace == "acescct": + return video.clamp(0.0, 1.0) + if input_colorspace == "srgb_gamma": + video = srgb_eotf_to_linear(video) + if input_colorspace in ("srgb_gamma", "srgb"): + video = apply_primaries_matrix(video, REC709_TO_ACESCG) + return acescct_encode(video.clamp(min=0.0)) + + +def acescct_to_linear(acescct: torch.Tensor, output_colorspace: str = "rec709") -> torch.Tensor: + r""" + Output transform of the LTX-2.5 SDR-To-HDR IC-LoRA: map ACEScct codes to scene-linear HDR. + + The codes are decoded to linear ACEScg, converted to the requested primaries, then clamped to `>= 0`. The clamp + happens after the primaries matrix, as in the reference implementation, so out-of-gamut Rec.709 values are clipped. + + Args: + acescct (`torch.Tensor`): + ACEScct codes of shape `(B, 3, F, H, W)` or `(..., 3, H, W)`. + output_colorspace (`str`, *optional*, defaults to `"rec709"`): + `"rec709"` for scene-linear Rec.709 primaries, or `"acescg"` for scene-linear ACEScg (AP1) primaries. + + Returns: + `torch.Tensor`: Scene-linear HDR values in `[0, inf)`, as float32. + """ + if output_colorspace not in ("rec709", "acescg"): + raise ValueError(f"Unsupported output colorspace {output_colorspace!r}. Expected 'rec709' or 'acescg'.") + linear_acescg = acescct_decode(acescct.float()) + if output_colorspace == "rec709": + linear_acescg = apply_primaries_matrix(linear_acescg, ACESCG_TO_REC709) + return linear_acescg.clamp(min=0.0) + + class LTX2VideoHDRProcessor(VideoProcessor): r""" Video processor for the LTX-2 HDR IC-LoRA pipeline. @@ -37,13 +227,19 @@ class LTX2VideoHDRProcessor(VideoProcessor): - `postprocess_hdr_video`: applies the LogC3 inverse transform to the VAE's decoded output, mapping `[0, 1]` → linear HDR `[0, ∞)`. + With `hdr_transform="acescct"` (LTX-2.5 SDR-To-HDR IC-LoRA), `preprocess_reference_video_hdr` additionally applies + an input transform to the ACEScct working space (see [`~pipelines.ltx2.image_processor.to_acescct`]) and + `postprocess_hdr_video` decodes ACEScct to scene-linear HDR (see + [`~pipelines.ltx2.image_processor.acescct_to_linear`]). + Args: vae_scale_factor (`int`, *optional*, defaults to `32`): VAE (spatial) scale factor for the LTX-2 video VAE. resample (`str`, *optional*, defaults to `"bilinear"`): Resampling filter used by the base [`VaeImageProcessor`] for PIL/tensor resizing. hdr_transform (`str`, *optional*, defaults to `"logc3"`): - HDR transform identifier. Only `"logc3"` (ARRI EI 800) is currently supported. + HDR transform identifier. `"logc3"` (ARRI LogC3 EI 800, LTX-2.3 HDR) or `"acescct"` (ACEScct, LTX-2.5 + SDR-To-HDR). """ # LogC3 (ARRI EI 800) coefficients, ported from `ltx_core.hdr.LogC3`. @@ -67,8 +263,8 @@ def __init__( vae_scale_factor=vae_scale_factor, resample=resample, ) - if hdr_transform != "logc3": - raise ValueError(f"Unsupported HDR transform {hdr_transform!r}. Only 'logc3' is supported.") + if hdr_transform not in ("logc3", "acescct"): + raise ValueError(f"Unsupported HDR transform {hdr_transform!r}. Expected 'logc3' or 'acescct'.") @classmethod def _logc3_decompress(cls, logc: torch.Tensor) -> torch.Tensor: @@ -115,32 +311,105 @@ def _resize_and_reflect_pad_video(video: torch.Tensor, height: int, width: int) return video + @staticmethod + def _video_to_float_tensor(video) -> tuple[torch.Tensor, bool]: + r""" + Convert a video input to a float32 `(B, C, F, H, W)` tensor without resizing or normalizing it. + + Accepts the layouts of [`VideoProcessor.preprocess_video`]. Integer inputs (PIL images, `uint8` arrays or + tensors) are divided by 255; floating point inputs are kept as is, so scene-linear values above 1 survive. + + Returns: + `tuple[torch.Tensor, bool]`: The video tensor, and whether the input was integer-valued. + """ + if isinstance(video, (np.ndarray, torch.Tensor)) and video.ndim == 5: + videos = list(video) + elif isinstance(video, list) and is_valid_image(video[0]) or is_valid_image_imagelist(video): + videos = [video] + elif isinstance(video, list) and is_valid_image_imagelist(video[0]): + videos = video + else: + raise ValueError( + "Input is in incorrect format. Currently, we only support numpy.ndarray, torch.Tensor, PIL.Image.Image" + ) + + tensors = [] + is_integer = True + for frames in videos: + if isinstance(frames, list) and isinstance(frames[0], PIL.Image.Image): + frames = np.stack([np.array(frame.convert("RGB")) for frame in frames], axis=0) + elif isinstance(frames, list) and isinstance(frames[0], np.ndarray): + frames = np.stack(frames, axis=0) + elif isinstance(frames, list) and isinstance(frames[0], torch.Tensor): + frames = torch.stack(frames, dim=0) + if isinstance(frames, np.ndarray): + # NumPy frames are channels-last: (F, H, W, C) -> (F, C, H, W). + frames = torch.from_numpy(np.ascontiguousarray(frames)).permute(0, 3, 1, 2) + if frames.is_floating_point(): + is_integer = False + frames = frames.float() + else: + frames = frames.float() / 255.0 + tensors.append(frames) + # (B, F, C, H, W) -> (B, C, F, H, W) + return torch.stack(tensors, dim=0).permute(0, 2, 1, 3, 4), is_integer + def preprocess_reference_video_hdr( self, video, height: int, width: int, + input_colorspace: str | None = None, ) -> torch.Tensor: r""" Preprocess a reference (SDR) video for HDR IC-LoRA conditioning. - Runs the input through the standard video preprocessing (normalization to `[-1, 1]`) without resizing, then - applies reflect-pad resize to the target dimensions. For LDR inputs this is numerically equivalent to - `load_video_conditioning_hdr` in the reference implementation (since `LogC3.compress_ldr` is an identity clamp - on `[0, 1]` inputs). + With `hdr_transform="logc3"`, runs the input through the standard video preprocessing (normalization to `[-1, + 1]`) without resizing, then applies reflect-pad resize to the target dimensions. For LDR inputs this is + numerically equivalent to `load_video_conditioning_hdr` in the reference implementation (since + `LogC3.compress_ldr` is an identity clamp on `[0, 1]` inputs). + + With `hdr_transform="acescct"`, the input is mapped to ACEScct `[0, 1]` with + [`~pipelines.ltx2.image_processor.to_acescct`], reflect-pad resized, then mapped to `[-1, 1]`. Integer inputs + (PIL images, `uint8` arrays or tensors) are divided by 255 and transformed before resizing, as the reference + does for MP4/MOV inputs; floating point inputs are used as is and resized before the transform, as the + reference does for EXR inputs. The two orders only differ when the video is downscaled. Args: video: Input accepted by `VideoProcessor.preprocess_video` (list of PIL images, 4D/5D tensor/array, etc.). height (`int`), width (`int`): Target spatial dimensions. + input_colorspace (`str`, *optional*): + Colour space of `video` for `hdr_transform="acescct"`: `"srgb_gamma"` (default), `"srgb"`, `"acescg"` + or `"acescct"`. See [`~pipelines.ltx2.image_processor.to_acescct`]. Must be `None` for + `hdr_transform="logc3"`. Returns: `torch.Tensor`: Preprocessed video of shape `(B, C, F, height, width)` with values in `[-1, 1]`. """ - video = self.preprocess_video(video, height=None, width=None) # (B, C, F, src_h, src_w) in [-1, 1] - video = self._resize_and_reflect_pad_video(video, height, width) - return video + if self.config.hdr_transform == "logc3": + if input_colorspace is not None: + raise ValueError("`input_colorspace` is only supported with `hdr_transform='acescct'`.") + video = self.preprocess_video(video, height=None, width=None) # (B, C, F, src_h, src_w) in [-1, 1] + video = self._resize_and_reflect_pad_video(video, height, width) + return video - def postprocess_hdr_video(self, video: torch.Tensor, output_type: str = "np") -> torch.Tensor | np.ndarray: + input_colorspace = "srgb_gamma" if input_colorspace is None else input_colorspace + if input_colorspace not in ACESCCT_INPUT_COLORSPACES: + raise ValueError( + f"Unsupported input colorspace {input_colorspace!r}. Expected one of {ACESCCT_INPUT_COLORSPACES}." + ) + video, is_integer = self._video_to_float_tensor(video) + if is_integer: + video = to_acescct(video, input_colorspace) + video = self._resize_and_reflect_pad_video(video, height, width) + else: + video = self._resize_and_reflect_pad_video(video, height, width) + video = to_acescct(video, input_colorspace) + return video * 2.0 - 1.0 + + def postprocess_hdr_video( + self, video: torch.Tensor, output_type: str = "np", output_colorspace: str | None = None + ) -> torch.Tensor | np.ndarray: r""" Postprocess the VAE's decoded output to linear HDR. @@ -149,9 +418,15 @@ def postprocess_hdr_video(self, video: torch.Tensor, output_type: str = "np") -> VAE decoded output in VAE range `[-1, 1]`, shape `(B, C, F, H, W)`. output_type (`str`, *optional*, defaults to `"np"`): Output type of post-processed video tensor; should be in `["np", "pt"]`. + output_colorspace (`str`, *optional*): + Output colour space for `hdr_transform="acescct"`: `"rec709"` (default, scene-linear Rec.709), + `"acescg"` (scene-linear ACEScg) or `"acescct"` (the decoded ACEScct codes in `[0, 1]`, without + decoding them to linear). See [`~pipelines.ltx2.image_processor.acescct_to_linear`]. Must be `None` for + `hdr_transform="logc3"`, whose output is linear in the primaries of the input. Returns: - Returns linear HDR video with values in `[0, ∞)`, depending on `output_type`: + Returns linear HDR video with values in `[0, ∞)` (or ACEScct codes in `[0, 1]` with + `output_colorspace="acescct"`), depending on `output_type`: - `output_type="pt"`: `torch.Tensor` with shape `(B, F, H, W, C)` and dtype `float32`. - `output_type="np"`: `np.ndarray` with shape `(B, F, H, W, C)` and dtype `float32`. """ @@ -163,8 +438,20 @@ def postprocess_hdr_video(self, video: torch.Tensor, output_type: str = "np") -> output_type = "np" video = self.denormalize(video.float()) - # Apply the inverse transform function to get linear HDR light - video = self._logc3_decompress(video) + if self.config.hdr_transform == "logc3": + if output_colorspace is not None: + raise ValueError("`output_colorspace` is only supported with `hdr_transform='acescct'`.") + # Apply the inverse transform function to get linear HDR light + video = self._logc3_decompress(video) + else: + output_colorspace = "rec709" if output_colorspace is None else output_colorspace + if output_colorspace not in ACESCCT_OUTPUT_COLORSPACES: + raise ValueError( + f"Unsupported output colorspace {output_colorspace!r}. Expected one of " + f"{ACESCCT_OUTPUT_COLORSPACES}." + ) + if output_colorspace != "acescct": + video = acescct_to_linear(video, output_colorspace) # Permute to channels-last: [B, C, F, H, W] --> [B, F, H, W, C] video = video = video.permute(0, 2, 3, 4, 1).contiguous() diff --git a/src/diffusers/utils/__init__.py b/src/diffusers/utils/__init__.py index b3051dfcb9d1..6efe38dea81e 100644 --- a/src/diffusers/utils/__init__.py +++ b/src/diffusers/utils/__init__.py @@ -97,6 +97,7 @@ is_nvidia_modelopt_version, is_onnx_available, is_opencv_available, + is_openexr_available, is_optimum_quanto_available, is_optimum_quanto_version, is_outlines_available, diff --git a/src/diffusers/utils/import_utils.py b/src/diffusers/utils/import_utils.py index d2cf394cd9a7..2bf0e06e3279 100644 --- a/src/diffusers/utils/import_utils.py +++ b/src/diffusers/utils/import_utils.py @@ -220,6 +220,7 @@ def _is_package_available(pkg_name: str, get_dist_name: bool = False) -> tuple[b _sdnq_available, _sdnq_version = _is_package_available("sdnq") _flashpack_available, _flashpack_version = _is_package_available("flashpack") _av_available, _av_version = _is_package_available("av") +_openexr_available, _openexr_version = _is_package_available("OpenEXR") def is_torch_available(): @@ -422,6 +423,10 @@ def is_av_available(): return _av_available +def is_openexr_available(): + return _openexr_available + + # docstyle-ignore INFLECT_IMPORT_ERROR = """ {0} requires the inflect library but it was not found in your environment. You can install it with pip: `pip install diff --git a/tests/pipelines/ltx2/test_ltx2_hdr_color.py b/tests/pipelines/ltx2/test_ltx2_hdr_color.py new file mode 100644 index 000000000000..225f1c0b7d87 --- /dev/null +++ b/tests/pipelines/ltx2/test_ltx2_hdr_color.py @@ -0,0 +1,470 @@ +# Copyright 2026 The HuggingFace Team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Colour transforms and HDR export utilities of the LTX-2.5 SDR-To-HDR IC-LoRA. + +Expected values were computed independently in float64 with colour-science 0.4.4 (`log_encoding_ACEScct`, +`log_decoding_ACEScct`, `cctf_decoding(..., function="sRGB")`, `matrix_RGB_to_RGB(..., "Bradford")`, +`oetf_BT2100_HLG` / `oetf_inverse_BT2100_HLG`), not with the code under test. +""" + +import numpy as np +import PIL.Image +import pytest +import torch + +from diffusers.pipelines.ltx2.image_processor import ( + ACESCCT_A, + ACESCCT_B, + ACESCG_TO_REC709, + ACESCG_TO_REC2020, + REC709_TO_ACESCG, + REC709_TO_REC2020, + LTX2VideoHDRProcessor, + acescct_decode, + acescct_encode, + acescct_to_linear, + apply_primaries_matrix, + srgb_eotf_to_linear, + to_acescct, +) +from diffusers.utils import is_av_available, is_openexr_available + + +# colour-science 0.4.4, float64: `colour.matrix_RGB_to_RGB(src, dst, chromatic_adaptation_transform="Bradford")`. +COLOUR_ACESCG_TO_REC709 = [ + [1.7050509926579835, -0.6217921206570048, -0.08325887200097871], + [-0.13025641750704345, 1.1408047365754013, -0.010548319068358038], + [-0.024003356804618025, -0.12896897606497054, 1.1529723328695887], +] +COLOUR_REC709_TO_ACESCG = [ + [0.6130974024, 0.3395231462, 0.0473794514], + [0.0701937225, 0.9163538791, 0.0134523985], + [0.0206155929, 0.1095697729, 0.8698146342], +] +COLOUR_REC709_TO_REC2020 = [ + [0.6274038959, 0.3292830384, 0.0433130657], + [0.0690972894, 0.9195403951, 0.0113623156], + [0.0163914389, 0.0880133079, 0.8955952532], +] +COLOUR_ACESCG_TO_REC2020 = [ + [1.0258247477, -0.0200531908, -0.0057715568], + [-0.0022343695, 1.0045865019, -0.0023521324], + [-0.0050133515, -0.0252900718, 1.0303034233], +] + + +def _logc3_decompress_reference(t: float) -> float: + # ARRI LogC3 EI 800 decoding, written out independently of the processor. + a, b, c, d, e, f, cut = 5.555556, 0.052272, 0.247190, 0.385537, 5.367655, 0.092809, 0.010591 + if t >= e * cut + f: + return (10.0 ** ((t - d) / c) - b) / a + return (t - f) / e + + +class TestACEScct: + def test_encode_known_values(self): + linear = torch.tensor([0.0, 0.0078125, 0.18, 1.0, 16.0, 222.0], dtype=torch.float64) + expected = torch.tensor( + [ + 0.0729055341958355, + 0.1552511415525113, + 0.4135884024924423, + 0.5547945205479452, + 0.7831050228310503, + 0.9996812709103942, + ], + dtype=torch.float64, + ) + torch.testing.assert_close(acescct_encode(linear), expected, rtol=0, atol=1e-12) + + def test_encode_clamps_like_the_reference(self): + # Negative inputs are clamped to 0 and the output is clamped to [0, 1] (the ACEScct spec does neither). + out = acescct_encode(torch.tensor([-1.0, 1000.0], dtype=torch.float64)) + torch.testing.assert_close(out, torch.tensor([ACESCCT_B, 1.0], dtype=torch.float64), rtol=0, atol=1e-12) + + def test_decode_known_values(self): + codes = torch.tensor( + [0.0729055341958355, 0.155251141552511, 0.4135884024924423, 1.0, 0.0, 0.5], dtype=torch.float64 + ) + expected = torch.tensor( + [0.0, 0.0078125, 0.18, 222.8609442038076, -0.006916877586898862, 0.5140569133280329], + dtype=torch.float64, + ) + torch.testing.assert_close(acescct_decode(codes), expected, rtol=1e-12, atol=1e-12) + + def test_decode_clamps_input(self): + out = acescct_decode(torch.tensor([-0.5, 1.5], dtype=torch.float64)) + expected = torch.tensor([-ACESCCT_B / ACESCCT_A, 222.8609442038076], dtype=torch.float64) + torch.testing.assert_close(out, expected, rtol=1e-12, atol=1e-12) + + def test_round_trip_linear(self): + linear = torch.cat([torch.linspace(0.0, 0.01, 101), torch.logspace(-2, np.log10(222.0), 200)]).double() + torch.testing.assert_close(acescct_decode(acescct_encode(linear)), linear, rtol=1e-12, atol=1e-15) + + def test_round_trip_codes(self): + # The valid code domain starts at the code of linear 0; codes below decode to negative linear values, + # which the encoder clamps. + codes = torch.linspace(ACESCCT_B, 1.0, 1001, dtype=torch.float64) + torch.testing.assert_close(acescct_encode(acescct_decode(codes)), codes, rtol=0, atol=1e-12) + + +class TestPrimaries: + @pytest.mark.parametrize( + "matrix, expected", + [ + (ACESCG_TO_REC709, COLOUR_ACESCG_TO_REC709), + (REC709_TO_ACESCG, COLOUR_REC709_TO_ACESCG), + (REC709_TO_REC2020, COLOUR_REC709_TO_REC2020), + (ACESCG_TO_REC2020, COLOUR_ACESCG_TO_REC2020), + ], + ) + def test_matrix_values(self, matrix, expected): + # The stored matrices are float32, as in the reference. + torch.testing.assert_close( + torch.tensor(matrix, dtype=torch.float64), torch.tensor(expected, dtype=torch.float64), rtol=0, atol=1e-7 + ) + + @pytest.mark.parametrize("matrix", [ACESCG_TO_REC709, REC709_TO_ACESCG, REC709_TO_REC2020, ACESCG_TO_REC2020]) + def test_white_is_preserved(self, matrix): + row_sums = torch.tensor(matrix, dtype=torch.float64).sum(dim=1) + torch.testing.assert_close(row_sums, torch.ones(3, dtype=torch.float64), rtol=0, atol=1e-6) + + def test_rec709_acescg_are_inverse(self): + product = torch.tensor(REC709_TO_ACESCG, dtype=torch.float64) @ torch.tensor( + ACESCG_TO_REC709, dtype=torch.float64 + ) + torch.testing.assert_close(product, torch.eye(3, dtype=torch.float64), rtol=0, atol=1e-6) + + def test_apply_primaries_matrix_layouts(self): + generator = torch.Generator().manual_seed(0) + video = torch.rand(2, 3, 4, 5, 6, generator=generator) + expected = torch.einsum("dc,bcfhw->bdfhw", torch.tensor(REC709_TO_ACESCG), video) + torch.testing.assert_close(apply_primaries_matrix(video, REC709_TO_ACESCG), expected) + # Non-5D inputs use `dim=-3` as the channel axis. + frames = video[0].permute(1, 0, 2, 3) # (F, C, H, W) + torch.testing.assert_close(apply_primaries_matrix(frames, REC709_TO_ACESCG), expected[0].permute(1, 0, 2, 3)) + + +class TestInputTransform: + # Neutral greys, a saturated red and an arbitrary colour, as (N, 3, 1, 1) images. + pixels = torch.tensor([[0, 0, 0], [0.04045] * 3, [0.5] * 3, [1, 1, 1], [1.0, 0.0, 0.0], [0.2, 0.6, 0.9]]) + + def _as_images(self, values): + return torch.as_tensor(values, dtype=torch.float32)[:, :, None, None] + + def test_srgb_eotf(self): + out = srgb_eotf_to_linear(torch.tensor([0.0, 0.04045, 0.5, 1.0, -0.2, 1.3])) + expected = torch.tensor([0.0, 0.0031308072830676845, 0.21404114048223255, 1.0, 0.0, 1.0]) + assert out.dtype == torch.float32 + torch.testing.assert_close(out, expected, rtol=0, atol=1e-7) + + def test_srgb_gamma(self): + expected = [ + [0.0729055341958355] * 3, + [0.10590498888079533, 0.10590498734413856, 0.1059049870368072], + [0.4278516036650092, 0.42785159983049315, 0.42785159906358994], + [0.5547945245358419, 0.5547945207013258, 0.5547945199344226], + [0.5145084623593209, 0.3360437118624257, 0.23515295334335573], + [0.4068006341802645, 0.45696458686011593, 0.5277994813544838], + ] + out = to_acescct(self._as_images(self.pixels), "srgb_gamma") + torch.testing.assert_close(out, self._as_images(expected), rtol=0, atol=1e-6) + + def test_srgb_linear(self): + expected = [ + [0.0729055341958355] * 3, + [0.2906554556363287, 0.2906554518018126, 0.2906554510349094], + [0.49771689896506566, 0.4977168951305496, 0.49771689436364636], + [0.5547945245358419, 0.5547945207013258, 0.5547945199344226], + [0.5145084623593209, 0.3360437118624257, 0.23515295334335573], + [0.47269375324230106, 0.509362790853311, 0.5416727756641857], + ] + out = to_acescct(self._as_images(self.pixels), "srgb") + torch.testing.assert_close(out, self._as_images(expected), rtol=0, atol=1e-6) + + def test_acescg_and_acescct(self): + linear = torch.tensor([0.0, 0.18, 1.0, 16.0, -1.0, 1000.0]).view(2, 3, 1, 1) + expected = torch.tensor( + [0.0729055341958355, 0.4135884024924423, 0.5547945205479452, 0.7831050228310503, ACESCCT_B, 1.0] + ).view(2, 3, 1, 1) + torch.testing.assert_close(to_acescct(linear, "acescg"), expected, rtol=0, atol=1e-6) + codes = torch.tensor([-0.5, 0.25, 0.5, 0.75, 1.0, 2.0]).view(2, 3, 1, 1) + torch.testing.assert_close(to_acescct(codes, "acescct"), codes.clamp(0.0, 1.0), rtol=0, atol=0) + + def test_invalid_colorspace(self): + with pytest.raises(ValueError, match="Unsupported input colorspace"): + to_acescct(torch.zeros(1, 3, 1, 1), "rec2020") + + +class TestOutputTransform: + codes = torch.tensor([[0.0, 0.0, 0.0], [0.5, 0.5, 0.5], [0.7, 0.4, 0.2], [1.0, 1.0, 1.0]])[:, :, None, None] + + def test_acescg(self): + expected = torch.tensor( + [ + [0.0, 0.0, 0.0], # the negative toe (-0.0069169) is clamped to 0 + [0.5140569133280329] * 3, + [5.83203751526631, 0.1526183140836417, 0.013452331067795364], + [222.8609442038076] * 3, + ] + )[:, :, None, None] + torch.testing.assert_close(acescct_to_linear(self.codes, "acescg"), expected, rtol=1e-6, atol=1e-7) + + def test_rec709(self): + expected = torch.tensor( + [ + [0.0, 0.0, 0.0], + [0.5140568788578308, 0.5140569310418868, 0.5140569209880779], + [9.847904184623355, 0.0, 0.0], # out-of-gamut values are clipped after the matrix + [222.8609292598168, 222.86095188335844, 222.86094752469444], + ] + )[:, :, None, None] + torch.testing.assert_close(acescct_to_linear(self.codes, "rec709"), expected, rtol=1e-6, atol=1e-7) + + def test_round_trip_through_working_space(self): + # sRGB code -> ACEScct -> Rec.709 linear equals the sRGB EOTF on in-gamut colours. + generator = torch.Generator().manual_seed(0) + srgb = torch.rand(1, 3, 2, 8, 8, generator=generator) + linear = acescct_to_linear(to_acescct(srgb, "srgb_gamma"), "rec709") + torch.testing.assert_close(linear, srgb_eotf_to_linear(srgb), rtol=1e-4, atol=1e-6) + + def test_invalid_colorspace(self): + with pytest.raises(ValueError, match="Unsupported output colorspace"): + acescct_to_linear(self.codes, "acescct") + + +class TestLTX2VideoHDRProcessor: + def test_logc3_is_the_default_and_unchanged(self): + processor = LTX2VideoHDRProcessor() + assert processor.config.hdr_transform == "logc3" + + # Preprocessing: plain [-1, 1] normalization of the sRGB codes, then reflect-padding. + video = (np.arange(2 * 32 * 64 * 3) % 251).astype(np.uint8).reshape(2, 32, 64, 3) + out = processor.preprocess_reference_video_hdr([PIL.Image.fromarray(frame) for frame in video], 32, 80) + expected = torch.from_numpy(video).permute(0, 3, 1, 2).float() / 255.0 * 2.0 - 1.0 # (F, C, H, W) + expected = torch.nn.functional.pad(expected, (0, 16, 0, 0), mode="reflect") + torch.testing.assert_close(out, expected.permute(1, 0, 2, 3)[None]) + + # Postprocessing: LogC3 decoding of the [0, 1] codes, no primaries conversion, channels-last. + codes = torch.tensor([0.0, 0.05, 0.391007, 0.6, 0.9, 1.0]) + decoded = processor.postprocess_hdr_video((codes * 2.0 - 1.0).view(1, 3, 2, 1, 1), output_type="pt") + expected = torch.tensor([_logc3_decompress_reference(float(t)) for t in codes]).view(1, 3, 2, 1, 1) + assert decoded.shape == (1, 2, 1, 1, 3) + torch.testing.assert_close(decoded, expected.permute(0, 2, 3, 4, 1), rtol=1e-5, atol=1e-6) + + def test_logc3_rejects_colorspace_options(self): + processor = LTX2VideoHDRProcessor() + with pytest.raises(ValueError, match="input_colorspace"): + processor.preprocess_reference_video_hdr(np.zeros((1, 4, 4, 3), dtype=np.uint8), 4, 4, "srgb") + with pytest.raises(ValueError, match="output_colorspace"): + processor.postprocess_hdr_video(torch.zeros(1, 3, 1, 4, 4), "pt", "rec709") + + def test_invalid_transform(self): + with pytest.raises(ValueError, match="Unsupported HDR transform"): + LTX2VideoHDRProcessor(hdr_transform="pq") + + def test_acescct_preprocess_srgb_gamma(self): + processor = LTX2VideoHDRProcessor(hdr_transform="acescct") + # Black and white 8-bit frames: sRGB 0 -> ACEScct 0.0729055 -> -0.8541889, sRGB 1 -> 0.5547945 -> 0.1095890. + video = np.zeros((2, 20, 24, 3), dtype=np.uint8) + video[1] = 255 + out = processor.preprocess_reference_video_hdr(video, 32, 32) + assert out.shape == (1, 3, 2, 32, 32) + torch.testing.assert_close(out[:, :, 0], torch.full((1, 3, 32, 32), -0.854188931608329), rtol=0, atol=1e-6) + torch.testing.assert_close(out[:, :, 1], torch.full((1, 3, 32, 32), 0.10958904109589), rtol=0, atol=1e-6) + + def test_acescct_preprocess_matches_functional_transform(self): + processor = LTX2VideoHDRProcessor(hdr_transform="acescct") + generator = torch.Generator().manual_seed(0) + # Integer input: transform, then reflect-pad. + frames = (torch.rand(3, 3, 20, 24, generator=generator) * 255).to(torch.uint8) # (F, C, H, W) + out = processor.preprocess_reference_video_hdr(frames, 32, 32, input_colorspace="srgb_gamma") + working = to_acescct(frames.permute(1, 0, 2, 3)[None].float() / 255.0, "srgb_gamma") + expected = processor._resize_and_reflect_pad_video(working, 32, 32) * 2.0 - 1.0 + torch.testing.assert_close(out, expected, rtol=0, atol=0) + # Float input (e.g. an EXR plate, values above 1): reflect-pad, then transform. + linear = torch.rand(1, 3, 20, 24, 3, generator=generator).numpy() * 50.0 # (B, F, H, W, C) + out = processor.preprocess_reference_video_hdr(linear, 32, 32, input_colorspace="acescg") + padded = processor._resize_and_reflect_pad_video(torch.from_numpy(linear).permute(0, 4, 1, 2, 3), 32, 32) + torch.testing.assert_close(out, to_acescct(padded, "acescg") * 2.0 - 1.0, rtol=0, atol=0) + assert out.min() >= -1.0 and out.max() <= 1.0 + + def test_acescct_postprocess(self): + processor = LTX2VideoHDRProcessor(hdr_transform="acescct") + codes = torch.tensor([[0.0, 0.0, 0.0], [0.5, 0.5, 0.5], [0.7, 0.4, 0.2], [1.0, 1.0, 1.0]]) + decoded = (codes * 2.0 - 1.0).T.reshape(1, 3, 4, 1, 1) # (B, C, F, H, W) in [-1, 1] + + out = processor.postprocess_hdr_video(decoded, output_type="pt") + assert out.shape == (1, 4, 1, 1, 3) + expected = acescct_to_linear(codes.T.reshape(1, 3, 4, 1, 1), "rec709").permute(0, 2, 3, 4, 1) + torch.testing.assert_close(out, expected) + + out = processor.postprocess_hdr_video(decoded, output_type="np", output_colorspace="acescg") + assert isinstance(out, np.ndarray) and out.dtype == np.float32 + np.testing.assert_allclose(out[0, 2, 0, 0], [5.83203751526631, 0.1526183140836417, 0.013452331067795364], 1e-6) + + out = processor.postprocess_hdr_video(decoded, output_type="pt", output_colorspace="acescct") + torch.testing.assert_close(out[0, :, 0, 0], codes, rtol=0, atol=1e-7) + + +class TestHLG: + @pytest.fixture(autouse=True) + def _export_utils(self): + if not is_av_available(): + pytest.skip("PyAV is not installed.") + from diffusers.pipelines.ltx2 import export_utils + + self.export_utils = export_utils + + def test_inverse_oetf_and_rolloff(self): + white_x = self.export_utils._hlg_inverse_oetf(0.75) + assert white_x == pytest.approx(0.26496256042100724, abs=1e-15) + assert white_x / (1.0 - white_x) == pytest.approx(0.36047491753994165, abs=1e-15) + assert self.export_utils._hlg_inverse_oetf(0.5) == pytest.approx(1.0 / 12.0, abs=1e-15) + assert self.export_utils._hlg_inverse_oetf(0.3) == pytest.approx(0.03, abs=1e-15) + + def test_signal_known_values(self): + white_x = self.export_utils._hlg_inverse_oetf(0.75) + identity = torch.eye(3) + # Neutral Rec.709 linear 0.18 / 1.0 / 4.0 / 100.0 (Rec.709 -> Rec.2020 keeps neutrals neutral). + linear = torch.tensor([0.18, 1.0, 4.0, 100.0]).view(4, 1, 1, 1).expand(4, 3, 1, 1) + signal = self.export_utils._linear_to_hlg_signal( + linear, torch.tensor(REC709_TO_REC2020), white_x, white_x / (1.0 - white_x) + ) + expected = torch.tensor([0.3782588830779046, 0.75, 0.947280748538359, 0.9999999950661305]) + torch.testing.assert_close(signal[:, 0, 0, 0], expected, rtol=0, atol=2e-6) + # A coloured pixel goes through the Rec.709 -> Rec.2020 matrix first. + signal = self.export_utils._linear_to_hlg_signal( + torch.tensor([2.0, 0.5, 0.05]).view(1, 3, 1, 1), + torch.tensor(REC709_TO_REC2020), + white_x, + white_x / (1.0 - white_x), + ) + expected = torch.tensor([0.8139141545892317, 0.6460072658555785, 0.3108599919641505]) + torch.testing.assert_close(signal[0, :, 0, 0], expected, rtol=0, atol=2e-6) + # Negative and NaN values are mapped to 0. + bad = torch.tensor([-1.0, float("nan"), float("-inf")]).view(1, 3, 1, 1) + signal = self.export_utils._linear_to_hlg_signal(bad, identity, white_x, 1.0) + assert torch.equal(signal, torch.zeros_like(signal)) + + def test_yuv_code_levels(self): + rgb = torch.zeros(2, 3, 2, 2) + rgb[1] = 1.0 + y, u, v = self.export_utils._rgb_to_yuv420p10_bt2020_limited(rgb) + assert y.shape == (2, 2, 2) and u.shape == (2, 1, 1) and v.shape == (2, 1, 1) + assert y[0].unique().tolist() == [64] and y[1].unique().tolist() == [940] + assert u.unique().tolist() == [512] and v.unique().tolist() == [512] + # Pure Rec.2020 red: Y = 0.2627, Cb = -0.5 * 0.2627 / (1 - 0.0593), Cr = 0.5. + red = torch.zeros(1, 3, 2, 2) + red[:, 0] = 1.0 + y, u, v = self.export_utils._rgb_to_yuv420p10_bt2020_limited(red) + assert y.unique().tolist() == [round((219 * 0.2627 + 16) * 4)] + assert u.unique().tolist() == [round((224 * -0.5 * 0.2627 / (1 - 0.0593) + 128) * 4)] + assert v.unique().tolist() == [round((224 * 0.5 + 128) * 4)] + + def test_encode_hlg_mp4(self, tmp_path): + import av + + if "libx265" not in av.codecs_available: + pytest.skip("The FFmpeg build used by PyAV does not include libx265.") + + # Diffuse white (linear 1.0) maps to the HLG signal 0.75 -> Y = round((219 * 0.75 + 16) * 4) = 721. + frames = torch.ones(9, 64, 96, 3) + output = tmp_path / "hlg.mp4" + self.export_utils.encode_hdr_tensor_to_hlg_mp4(frames, output, frame_rate=24.0, thread_count=1) + + with av.open(str(output)) as container: + stream = container.streams.video[0] + context = stream.codec_context + assert context.name == "hevc" + assert context.pix_fmt == "yuv420p10le" + assert stream.codec_tag == "hvc1" + assert (stream.width, stream.height) == (96, 64) + assert stream.average_rate == 24 + assert context.color_primaries == 9 # BT.2020 + assert context.color_trc == 18 # ARIB STD-B67 (HLG) + assert context.colorspace == 9 # BT.2020 NCL + assert context.color_range == 1 # limited + decoded = list(container.decode(video=0)) + assert len(decoded) == 9 + y_plane = decoded[0].planes[0] + y = np.frombuffer(y_plane, dtype=np.uint16).reshape(y_plane.height, y_plane.line_size // 2)[:, :96] + assert abs(int(np.median(y)) - 721) <= 2 + + def test_encode_hlg_mp4_rejects_odd_sizes(self, tmp_path): + import av + + if "libx265" not in av.codecs_available: + pytest.skip("The FFmpeg build used by PyAV does not include libx265.") + with pytest.raises(ValueError, match="even"): + self.export_utils.encode_hdr_tensor_to_hlg_mp4(torch.ones(1, 63, 64, 3), tmp_path / "x.mp4", 24.0) + assert not (tmp_path / "x.mp4").exists() + + +class TestEXR: + @pytest.fixture(autouse=True) + def _export_utils(self): + if not is_av_available(): + pytest.skip("PyAV is not installed.") + if not is_openexr_available(): + pytest.skip("OpenEXR is not installed.") + from diffusers.pipelines.ltx2 import export_utils + + self.export_utils = export_utils + + @staticmethod + def _read(path): + import OpenEXR + + with OpenEXR.File(str(path), separate_channels=True) as exr_file: + header = dict(exr_file.header()) + channels = {name: channel.pixels.copy() for name, channel in exr_file.channels().items()} + return header, channels + + @pytest.mark.parametrize( + "exr_colorspace, chromaticities, color_space", + [ + ("acescg", (0.713, 0.293, 0.165, 0.830, 0.128, 0.044, 0.32168, 0.33767), "ACEScg"), + ("acescct", (0.713, 0.293, 0.165, 0.830, 0.128, 0.044, 0.32168, 0.33767), "ACEScct"), + ("srgb_linear", (0.64, 0.33, 0.30, 0.60, 0.15, 0.06, 0.3127, 0.3290), "sRGB"), + ], + ) + def test_export_sequence_round_trip(self, tmp_path, exr_colorspace, chromaticities, color_space): + import OpenEXR + + generator = torch.Generator().manual_seed(0) + frames = torch.rand(2, 6, 10, 3, generator=generator) * 100.0 + paths = self.export_utils.export_to_exr_sequence(frames, tmp_path / "exr", exr_colorspace=exr_colorspace) + assert [p.split("/")[-1] for p in paths] == ["frame_00000.exr", "frame_00001.exr"] + + for path, frame in zip(paths, frames): + header, channels = self._read(path) + assert set(channels) == {"R", "G", "B"} + assert header["compression"] == OpenEXR.ZIP_COMPRESSION + assert header["colorSpace"] == color_space + np.testing.assert_allclose(header["chromaticities"], chromaticities, rtol=0, atol=1e-7) + for index, name in enumerate("RGB"): + assert channels[name].dtype == np.float16 + assert channels[name].shape == (6, 10) + np.testing.assert_array_equal(channels[name], frame[..., index].numpy().astype(np.float16)) + + def test_save_frame_channels_first_and_full_float(self, tmp_path): + frame = torch.arange(3 * 4 * 2, dtype=torch.float32).reshape(3, 4, 2) * 1e3 # (C, H, W) + self.export_utils.save_exr_frame(frame, tmp_path / "frame.exr", primaries="acescg", half=False) + header, channels = self._read(tmp_path / "frame.exr") + assert header["colorSpace"] == "sRGB" + for index, name in enumerate("RGB"): + assert channels[name].dtype == np.float32 + np.testing.assert_array_equal(channels[name], frame[index].numpy()) From 271ff026ef5af071cd53c0fe2248f3400020f59a Mon Sep 17 00:00:00 2001 From: christopher5106 Date: Tue, 6 Oct 2026 20:09:56 +0200 Subject: [PATCH 2/4] [LTX-2.5] Support the SDR-To-HDR IC-LoRA in LTX2HDRPipeline Run the Lightricks LTX-2.5 SDR-To-HDR IC-LoRA (ltx_pipelines.hdr_ic_lora.HDRICLoraPipeline, without its optional seam keyframes) when the pipeline is built with hdr_transform="acescct": - the reference video goes through the ACEScct input transform (new `input_colorspace`, default "srgb_gamma"), is reflect-padded up to a multiple of the VAE spatial compression ratio and VAE-encoded in float32; the decoded video is cropped back to height x width; - conditioning comes from precomputed `connector_video_embeds` only (a 2D `video_context` is accepted): no prompt, no text encoder call, `connector_audio_embeds` optional; - video-only denoising: `isolate_modalities=True` with a single placeholder audio token, so the audio stream cannot reach the video; - the distilled schedule used verbatim (DISTILLED_SIGMA_VALUES, no mu shift), no CFG/STG/modality guidance, RoPE frame rate capped at 30 fps, first target latent frame marked for the keyframe position embedding; - decoding in float32 with the optional `diffusion_decoder` component (LTX-2.5) or the VAE, then `postprocess_hdr_video(output_colorspace=...)` (new argument, default "rec709"). `hdr_transform` is now registered in the pipeline config so it survives save/load. The LogC3 (LTX-2.3) path is unchanged: its outputs are bit-identical to the parent commit. Co-Authored-By: Claude Opus 5.5 --- .../pipelines/ltx2/pipeline_ltx2_hdr_lora.py | 445 +++++++++++++++--- tests/pipelines/ltx2/test_ltx2_hdr.py | 2 +- .../ltx2/test_ltx2_hdr_sdr_to_hdr.py | 387 +++++++++++++++ 3 files changed, 771 insertions(+), 63 deletions(-) create mode 100644 tests/pipelines/ltx2/test_ltx2_hdr_sdr_to_hdr.py diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py index d9d95cdc0a0b..dd5c18f0dd6f 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py @@ -14,6 +14,8 @@ import copy import inspect +import math +from contextlib import contextmanager from dataclasses import dataclass from typing import Any, Callable @@ -29,7 +31,7 @@ from ...callbacks import MultiPipelineCallbacks, PipelineCallback from ...loaders import FromSingleFileMixin, LTX2LoraLoaderMixin -from ...models.autoencoders import AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video +from ...models.autoencoders import AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video, LTX2VideoDiffusionDecoderModel from ...models.transformers import LTX2VideoTransformer3DModel from ...schedulers import FlowMatchEulerDiscreteScheduler from ...utils import is_torch_xla_available, logging, replace_example_docstring @@ -38,6 +40,7 @@ from .connectors import LTX2TextConnectors from .image_processor import LTX2VideoHDRProcessor from .pipeline_output import LTX2PipelineOutput +from .utils import DISTILLED_SIGMA_VALUES, SNAP_CONDITIONING_FPS_ABOVE from .vocoder import LTX2Vocoder, LTX2VocoderWithBWE @@ -50,6 +53,10 @@ logger = logging.get_logger(__name__) # pylint: disable=invalid-name +# The LTX-2.5 SDR-To-HDR IC-LoRA conditions RoPE on 30 fps for sources above 30 fps; the playback frame rate is +# unchanged (`_conditioning_fps` in `ltx_pipelines/hdr_ic_lora.py`). +SDR_TO_HDR_MAX_CONDITIONING_FPS = 30.0 + @dataclass class LTX2HDRReferenceCondition: @@ -122,6 +129,68 @@ class LTX2HDRReferenceCondition: >>> # A custom tone-mapper can be specified via the `tone_mapping_fn` argument. >>> encode_hdr_tensor_to_mp4(hdr_video[0], "ltx2_hdr_lora_output.mp4", frame_rate=24.0) ``` + + LTX-2.5 SDR-To-HDR IC-LoRA (ACEScct). The LoRA ships a precomputed scene embedding, so the text encoder is not + needed, and the clip keeps its own resolution and frame count: + + ```py + >>> import torch + >>> from huggingface_hub import hf_hub_download + >>> from safetensors.torch import load_file + >>> from diffusers import LTX2HDRPipeline + >>> from diffusers.models.autoencoders.ltx2_diffusion_decoder import LTX2VideoVaeNeighborhoodNattenProcessor + >>> from diffusers.pipelines.ltx2.export_utils import encode_hdr_tensor_to_hlg_mp4, export_to_exr_sequence + >>> from diffusers.pipelines.ltx2.image_processor import acescct_to_linear + >>> from diffusers.pipelines.ltx2.pipeline_ltx2_hdr_lora import LTX2HDRReferenceCondition + >>> from diffusers.utils import load_video + + >>> # The VAE and the diffusion decoder run in float32, as in the reference (the pipeline would otherwise + >>> # upcast them for each call). + >>> pipe = LTX2HDRPipeline.from_pretrained( + ... "Lightricks/LTX-2.5-Diffusers", + ... hdr_transform="acescct", + ... text_encoder=None, + ... tokenizer=None, + ... connectors=None, + ... dtype={"default": torch.bfloat16, "vae": torch.float32, "diffusion_decoder": torch.float32}, + ... ) + >>> pipe.enable_model_cpu_offload() + >>> # The diffusion decoder needs NATTEN (`pip install kernels`) and tiling at video resolutions. + >>> pipe.diffusion_decoder.set_attn_processor(LTX2VideoVaeNeighborhoodNattenProcessor()) + >>> pipe.diffusion_decoder.enable_tiling() + + >>> repo_id = "Lightricks/LTX-2.5-22b-IC-LoRA-SDR-To-HDR" + >>> pipe.load_lora_weights( + ... repo_id, adapter_name="sdr_to_hdr", weight_name="ltx-2.5-22b-ic-lora-sdr-to-hdr-1.0.safetensors" + ... ) + >>> pipe.set_adapters("sdr_to_hdr", 1.0) + >>> scene_embeds = load_file(hf_hub_download(repo_id, "ltx-2.5-22b-ic-lora-sdr-to-hdr-scene-emb.safetensors")) + + >>> # An 8-bit sRGB clip with 8k + 1 frames. Any size works: it is reflect-padded to a multiple of 32 and the + >>> # output is cropped back. + >>> sdr_video = load_video("/path/to/sdr.mp4") + >>> frame_rate = 24.0 + >>> acescct = pipe( + ... reference_conditions=LTX2HDRReferenceCondition(frames=sdr_video), + ... connector_video_embeds=scene_embeds["video_context"], + ... input_colorspace="srgb_gamma", + ... output_colorspace="acescct", + ... height=sdr_video[0].height, + ... width=sdr_video[0].width, + ... num_frames=len(sdr_video), + ... frame_rate=frame_rate, + ... generator=torch.Generator("cuda").manual_seed(42), + ... ).frames[0] + + + >>> def to_linear(frames, colorspace): # (F, H, W, 3) ACEScct codes -> scene-linear HDR + ... return acescct_to_linear(frames.permute(0, 3, 1, 2), colorspace).permute(0, 2, 3, 1) + + + >>> # A BT.2020 HLG master from scene-linear Rec.709, and an ACEScg EXR sequence. + >>> encode_hdr_tensor_to_hlg_mp4(to_linear(acescct, "rec709"), "ltx2_sdr_to_hdr.mp4", frame_rate=frame_rate) + >>> export_to_exr_sequence(to_linear(acescct, "acescg"), "ltx2_sdr_to_hdr_exr", exr_colorspace="acescg") + ``` """ @@ -256,6 +325,20 @@ class LTX2HDRPipeline(DiffusionPipeline, FromSingleFileMixin, LTX2LoraLoaderMixi values to avoid wasted compute. - No frame-level keyframe conditioning (the reference HDR pipeline does not support this). + With `hdr_transform="acescct"`, the pipeline runs the LTX-2.5 SDR-To-HDR IC-LoRA the way the reference + `HDRICLoraPipeline` does (without its optional seam keyframes): + + - the reference video is mapped to ACEScct (see `input_colorspace`), reflect-padded to a multiple of the VAE's + spatial compression ratio and VAE-encoded in float32; the output is cropped back to `height` x `width`; + - the transformer is conditioned on precomputed `connector_video_embeds` only. No prompt is encoded, so + `text_encoder`, `tokenizer` and `connectors` can be loaded as `None`; + - the denoising is video-only: audio-to-video cross-attention is disabled, so the placeholder audio stream has no + influence on the video; + - a single distilled stage (`DISTILLED_SIGMA_VALUES` by default, used as given), without CFG, STG or modality + guidance, and with RoPE conditioned on at most 30 fps; + - the latents are decoded in float32, by `diffusion_decoder` when it is loaded and by `vae` otherwise, then + converted from ACEScct to scene-linear HDR (see `output_colorspace`). + Two-stage inference is supported through separate calls to `__call__`: - **Stage 1**: generate video latents at target resolution with HDR IC-LoRA conditioning (`output_type="latent"`). @@ -282,12 +365,18 @@ class LTX2HDRPipeline(DiffusionPipeline, FromSingleFileMixin, LTX2LoraLoaderMixi Transformer backbone. vocoder ([`LTX2Vocoder`] or [`LTX2VocoderWithBWE`]): Vocoder. Required for transformer compatibility; its outputs are discarded. + audio_scheduler ([`FlowMatchEulerDiscreteScheduler`], *optional*): + Scheduler for the (discarded) audio stream. Defaults to a copy of `scheduler`. + diffusion_decoder ([`LTX2VideoDiffusionDecoderModel`], *optional*): + The LTX-2.5 diffusion video decoder. Only used with `hdr_transform="acescct"`, where it replaces `vae` for + decoding, as in the reference implementation. hdr_transform (`str`, *optional*, defaults to `"logc3"`): - HDR transform identifier applied during postprocessing. Currently only `"logc3"` is supported. + HDR transform of the IC-LoRA: `"logc3"` (ARRI LogC3, LTX-2.3 HDR IC-LoRA) or `"acescct"` (ACEScct, LTX-2.5 + SDR-To-HDR IC-LoRA). See [`LTX2VideoHDRProcessor`]. """ - model_cpu_offload_seq = "text_encoder->connectors->transformer->vae->audio_vae->vocoder" - _optional_components = ["audio_scheduler"] + model_cpu_offload_seq = "text_encoder->connectors->transformer->vae->diffusion_decoder->audio_vae->vocoder" + _optional_components = ["audio_scheduler", "diffusion_decoder"] _callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"] def __init__( @@ -301,11 +390,13 @@ def __init__( transformer: LTX2VideoTransformer3DModel, vocoder: LTX2Vocoder | LTX2VocoderWithBWE, audio_scheduler: FlowMatchEulerDiscreteScheduler | None = None, + diffusion_decoder: LTX2VideoDiffusionDecoderModel | None = None, hdr_transform: str = "logc3", ): super().__init__() self.register_modules( + diffusion_decoder=diffusion_decoder, vae=vae, audio_vae=audio_vae, text_encoder=text_encoder, @@ -351,6 +442,7 @@ def __init__( vae_scale_factor=self.vae_spatial_compression_ratio, hdr_transform=hdr_transform, ) + self.register_to_config(hdr_transform=hdr_transform) self.tokenizer_max_length = ( self.tokenizer.model_max_length if getattr(self, "tokenizer", None) is not None else 1024 @@ -569,6 +661,78 @@ def check_inputs( " block indices at which to apply STG in `spatio_temporal_guidance_blocks`" ) + def check_sdr_to_hdr_inputs( + self, + prompt, + num_frames, + reference_conditions, + callback_on_step_end_tensor_inputs=None, + prompt_embeds=None, + negative_prompt_embeds=None, + connector_video_embeds=None, + latents=None, + guidance_scale=1.0, + stg_scale=0.0, + modality_scale=1.0, + ): + r"""Input checks for `hdr_transform="acescct"` (LTX-2.5 SDR-To-HDR IC-LoRA).""" + if callback_on_step_end_tensor_inputs is not None and not all( + k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs + ): + raise ValueError( + f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}" + ) + + # The reference runs no text encoder: the transformer only sees the IC-LoRA's precomputed video context. + if prompt is not None or prompt_embeds is not None or negative_prompt_embeds is not None: + raise ValueError( + "`hdr_transform='acescct'` conditions on `connector_video_embeds` only; `prompt`, `prompt_embeds` and" + " `negative_prompt_embeds` are not supported." + ) + if connector_video_embeds is None: + raise ValueError("`hdr_transform='acescct'` requires `connector_video_embeds`.") + if connector_video_embeds.ndim not in (2, 3): + raise ValueError( + "`connector_video_embeds` must have shape `(sequence_length, dim)` or `(batch_size, sequence_length," + f" dim)`, but got {tuple(connector_video_embeds.shape)}." + ) + + # The reference denoiser makes a single unguided forward per step. + if guidance_scale > 1.0 or stg_scale > 0.0 or modality_scale > 1.0: + raise ValueError( + "`hdr_transform='acescct'` runs the distilled model without guidance; `guidance_scale` and" + f" `modality_scale` must be <= 1 and `stg_scale` 0, but got {guidance_scale}, {modality_scale} and" + f" {stg_scale}." + ) + + if not reference_conditions: + raise ValueError("`hdr_transform='acescct'` requires a reference video in `reference_conditions`.") + + if (num_frames - 1) % self.vae_temporal_compression_ratio != 0: + raise ValueError( + f"`num_frames` must be of the form {self.vae_temporal_compression_ratio}k + 1 with" + f" `hdr_transform='acescct'`, but got {num_frames}." + ) + + if latents is not None and latents.ndim != 5: + raise ValueError( + f"Only unpacked (5D) video latents of shape `[batch_size, latent_channels, latent_frames," + f" latent_height, latent_width] are supported, but got {latents.ndim} dims." + ) + + @contextmanager + def _float32(self, module: torch.nn.Module): + r"""Run `module` in float32 for the duration of the context, then restore its dtype.""" + dtype = module.dtype + if dtype == torch.float32: + yield module + return + module.to(dtype=torch.float32) + try: + yield module + finally: + module.to(dtype=dtype) + @staticmethod # Copied from diffusers.pipelines.ltx2.pipeline_ltx2.LTX2Pipeline._pack_latents def _pack_latents(latents: torch.Tensor, patch_size: int = 1, patch_size_t: int = 1) -> torch.Tensor: @@ -708,6 +872,7 @@ def prepare_latents( device: torch.device | None = None, generator: torch.Generator | None = None, latents: torch.Tensor | None = None, + input_colorspace: str | None = None, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None, int, torch.Tensor | None]: r""" Prepare noisy video latents, applying HDR IC-LoRA reference-video conditioning. @@ -727,6 +892,9 @@ def prepare_latents( no reference conditions are provided. - `num_ref_tokens`: count of reference tokens at the END of `latents`. - `ref_cross_mask`: always `None` for HDR LoRA (no cross-attention masking support). + + `input_colorspace` is forwarded to [`LTX2VideoHDRProcessor.preprocess_reference_video_hdr`] and is only + supported with `hdr_transform="acescct"`. """ latent_height = height // self.vae_spatial_compression_ratio latent_width = width // self.vae_spatial_compression_ratio @@ -782,6 +950,7 @@ def prepare_latents( dtype=dtype, device=device, generator=generator[0] if isinstance(generator, list) else generator, + input_colorspace=input_colorspace, ) num_ref_tokens = ref_latents_packed.shape[1] @@ -860,13 +1029,18 @@ def _encode_reference_conditions( dtype: torch.dtype | None = None, device: torch.device | None = None, generator: torch.Generator | None = None, + input_colorspace: str | None = None, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]: """Encode HDR IC-LoRA reference videos into `(reference_latents, reference_coords, reference_cross_mask)`. Shared encoding core used by both `prepare_latents` (which folds reference tokens into the main noisy sequence) and the back-compat shim `prepare_reference_latents`. HDR LoRA does not currently support cross-attention masking for reference tokens, so the third return is always `None`. + + With `hdr_transform="acescct"` the reference is mapped to ACEScct (`input_colorspace`), must cover `num_frames` + frames, and is encoded with the VAE in float32 (`hdr_ic_lora.py:189, 427-431` in the reference). """ + is_sdr_to_hdr = self.hdr_video_processor.config.hdr_transform == "acescct" ref_height = height // reference_downscale_factor ref_width = width // reference_downscale_factor @@ -894,11 +1068,23 @@ def _encode_reference_conditions( # HDR-specific preprocessing: reflect-pad resize (vs center-crop in the standard IC-LoRA pipeline). # For LDR reference videos the numerical output of `preprocess_reference_video_hdr` is identical to the # standard [-1, 1] normalization since LogC3's `compress_ldr` is an identity clamp. - ref_pixels = self.hdr_video_processor.preprocess_reference_video_hdr(video_like, ref_height, ref_width) + ref_pixels = self.hdr_video_processor.preprocess_reference_video_hdr( + video_like, ref_height, ref_width, input_colorspace=input_colorspace + ) ref_pixels = ref_pixels[:, :, :num_frames] - ref_pixels = ref_pixels.to(dtype=self.vae.dtype, device=device) - - ref_latent = retrieve_latents(self.vae.encode(ref_pixels), generator=generator, sample_mode="argmax") + if is_sdr_to_hdr: + if ref_pixels.shape[2] < num_frames: + raise ValueError( + f"The reference video has {ref_pixels.shape[2]} frames, fewer than `num_frames={num_frames}`." + ) + with self._float32(self.vae): + ref_pixels = ref_pixels.to(dtype=torch.float32, device=device) + ref_latent = retrieve_latents( + self.vae.encode(ref_pixels), generator=generator, sample_mode="argmax" + ) + else: + ref_pixels = ref_pixels.to(dtype=self.vae.dtype, device=device) + ref_latent = retrieve_latents(self.vae.encode(ref_pixels), generator=generator, sample_mode="argmax") ref_latent = self._normalize_latents(ref_latent, self.vae.latents_mean, self.vae.latents_std).to( device=device, dtype=dtype ) @@ -1049,6 +1235,8 @@ def __call__( negative_prompt: str | list[str] | None = None, reference_conditions: LTX2HDRReferenceCondition | list[LTX2HDRReferenceCondition] | None = None, reference_downscale_factor: int = 1, + input_colorspace: str | None = None, + output_colorspace: str | None = None, height: int = 512, width: int = 768, num_frames: int = 121, @@ -1094,18 +1282,31 @@ def __call__( reference_downscale_factor (`int`, *optional*, defaults to `1`): Ratio between target and reference video resolutions. IC-LoRA models trained with downscaled reference videos store this factor in their safetensors metadata. + input_colorspace (`str`, *optional*): + Colour space of the reference video with `hdr_transform="acescct"`: `"srgb_gamma"` (default, + sRGB-encoded video such as an 8-bit MP4), `"srgb"` (scene-linear Rec.709), `"acescg"` (scene-linear + ACEScg) or `"acescct"`. See [`~pipelines.ltx2.image_processor.to_acescct`]. Not supported with + `hdr_transform="logc3"`. + output_colorspace (`str`, *optional*): + Colour space of the output with `hdr_transform="acescct"`: `"rec709"` (default, scene-linear Rec.709), + `"acescg"` (scene-linear ACEScg) or `"acescct"` (the decoded ACEScct codes in `[0, 1]`). See + [`LTX2VideoHDRProcessor.postprocess_hdr_video`]. Not supported with `hdr_transform="logc3"`. height (`int`, *optional*, defaults to `512`): - Output video height in pixels. Must be divisible by 32. + Output video height in pixels. Must be divisible by 32 with `hdr_transform="logc3"`. With + `hdr_transform="acescct"` any height works: the reference is reflect-padded up to a multiple of the VAE + spatial compression ratio and the decoded video is cropped back to `height`. width (`int`, *optional*, defaults to `768`): - Output video width in pixels. Must be divisible by 32. + Output video width in pixels, with the same constraints as `height`. num_frames (`int`, *optional*, defaults to `121`): Number of frames to generate. Must satisfy `(n - 1) % 8 == 0`. frame_rate (`float`, *optional*, defaults to `24.0`): - Output frame rate (used for temporal positional encoding). + Output frame rate (used for temporal positional encoding). With `hdr_transform="acescct"`, frame rates + above 30 are conditioned as 30 fps, as in the reference implementation. num_inference_steps (`int`, *optional*, defaults to `8`): Number of denoising steps. Default matches the distilled model schedule. sigmas (`List[float]`, *optional*): - Custom sigma schedule. Overrides `num_inference_steps` when set. + Custom sigma schedule. Overrides `num_inference_steps` when set. With `hdr_transform="acescct"` it + defaults to `DISTILLED_SIGMA_VALUES` and is used as given, without a resolution-dependent shift. timesteps (`List[float]`, *optional*): Custom timesteps schedule. Overrides `num_inference_steps` when set. guidance_scale (`float`, *optional*, defaults to `1.0`): @@ -1136,10 +1337,12 @@ def __call__( Attention mask for `negative_prompt_embeds`. connector_video_embeds (`torch.Tensor`, *optional*): Optional pre-computed connector outputs for the video modality. Used by the HDR LoRA pipeline; if - supplied, will override any `prompt`/`prompt_embeds`. + supplied, will override any `prompt`/`prompt_embeds`. Required with `hdr_transform="acescct"`, where a + 2D `(sequence_length, dim)` tensor (the layout of the IC-LoRA's `video_context`) is also accepted. connector_audio_embeds (`torch.Tensor`, *optional*): Optional pre-computed connector outputs for the audio modality. Used by the HDR LoRA pipeline; if - supplied, will override any `prompt`/`prompt_embeds`. + supplied, will override any `prompt`/`prompt_embeds`. Optional with `hdr_transform="acescct"`: the + audio stream is isolated from the video there, so it only feeds the discarded audio output. decode_timestep (`float` or `list[float]`, defaults to `0.0`): VAE-decode timestep conditioning (only used by VAE configs with `timestep_conditioning=True`). decode_noise_scale (`float` or `list[float]`, *optional*): @@ -1149,7 +1352,8 @@ def __call__( output_type (`str`, *optional*, defaults to `"pt"`): One of `"pt"`, `"np"`, or `"latent"`. `"pt"` returns a linear HDR torch tensor in `[0, ∞)` of shape `(batch_size, num_frames, height, width, channels)`; `"np"` returns the equivalent `float32` NumPy - array; `"latent"` returns the raw denoised latents (skip the HDR decode). + array; `"latent"` returns the raw denoised latents (skip the HDR decode). With + `hdr_transform="acescct"` the latents cover the padded size and are not cropped. return_dict (`bool`, *optional*, defaults to `True`): Whether to return an [`LTX2PipelineOutput`] instead of a plain tuple. attention_kwargs (`dict`, *optional*): @@ -1171,22 +1375,47 @@ def __call__( if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)): callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs + if reference_conditions is not None and not isinstance(reference_conditions, list): + reference_conditions = [reference_conditions] + + # `hdr_transform="acescct"` is the LTX-2.5 SDR-To-HDR IC-LoRA (`ltx_pipelines.hdr_ic_lora.HDRICLoraPipeline`). + is_sdr_to_hdr = self.hdr_video_processor.config.hdr_transform == "acescct" + # 1. Check inputs - self.check_inputs( - prompt=prompt, - height=height, - width=width, - callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs, - prompt_embeds=prompt_embeds, - negative_prompt_embeds=negative_prompt_embeds, - prompt_attention_mask=prompt_attention_mask, - negative_prompt_attention_mask=negative_prompt_attention_mask, - connector_video_embeds=connector_video_embeds, - connector_audio_embeds=connector_audio_embeds, - latents=latents, - spatio_temporal_guidance_blocks=spatio_temporal_guidance_blocks, - stg_scale=stg_scale, - ) + if is_sdr_to_hdr: + self.check_sdr_to_hdr_inputs( + prompt=prompt, + num_frames=num_frames, + reference_conditions=reference_conditions, + callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + connector_video_embeds=connector_video_embeds, + latents=latents, + guidance_scale=guidance_scale, + stg_scale=stg_scale, + modality_scale=modality_scale, + ) + else: + if input_colorspace is not None or output_colorspace is not None: + raise ValueError( + "`input_colorspace` and `output_colorspace` are only supported with `hdr_transform='acescct'`." + ) + self.check_inputs( + prompt=prompt, + height=height, + width=width, + callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + prompt_attention_mask=prompt_attention_mask, + negative_prompt_attention_mask=negative_prompt_attention_mask, + connector_video_embeds=connector_video_embeds, + connector_audio_embeds=connector_audio_embeds, + latents=latents, + spatio_temporal_guidance_blocks=spatio_temporal_guidance_blocks, + stg_scale=stg_scale, + ) # Video-only guidance state. self._guidance_scale = guidance_scale @@ -1199,6 +1428,12 @@ def __call__( self._current_timestep = None # 2. Define call parameters + if is_sdr_to_hdr and connector_video_embeds.ndim == 2: + # The IC-LoRA's `video_context` is stored without a batch dimension. + connector_video_embeds = connector_video_embeds.unsqueeze(0) + if is_sdr_to_hdr and connector_audio_embeds is not None and connector_audio_embeds.ndim == 2: + connector_audio_embeds = connector_audio_embeds.unsqueeze(0) + if prompt is not None and isinstance(prompt, str): batch_size = 1 elif prompt is not None and isinstance(prompt, list): @@ -1208,8 +1443,23 @@ def __call__( else: batch_size = connector_video_embeds.shape[0] - if reference_conditions is not None and not isinstance(reference_conditions, list): - reference_conditions = [reference_conditions] + if is_sdr_to_hdr: + # Generate at the size rounded up to the VAE's spatial compression ratio and crop back after decoding + # (`align_resolution(..., ResizeMode.REFLECT_PAD, divisor=32)` and `decoded[:, :crop_h, :crop_w]` in + # `hdr_ic_lora.py:271-273, 520`). + output_height, output_width = height, width + height = math.ceil(height / self.vae_spatial_compression_ratio) * self.vae_spatial_compression_ratio + width = math.ceil(width / self.vae_spatial_compression_ratio) * self.vae_spatial_compression_ratio + # RoPE is conditioned on at most 30 fps (`_conditioning_fps`, `hdr_ic_lora.py:74-78`). The playback rate + # `frame_rate` still sizes the (discarded) audio stream. + conditioning_frame_rate = ( + SDR_TO_HDR_MAX_CONDITIONING_FPS if frame_rate > SNAP_CONDITIONING_FPS_ABOVE else frame_rate + ) + # The full distilled schedule, used verbatim (`DEFAULT_DENOISE_SIGMAS`, `hdr_ic_lora.py:66`). + if sigmas is None and timesteps is None: + sigmas = DISTILLED_SIGMA_VALUES + else: + conditioning_frame_rate = frame_rate if noise_scale is None: noise_scale = sigmas[0] if sigmas is not None else 1.0 @@ -1217,7 +1467,30 @@ def __call__( device = self._execution_device # 3. Prepare text embeddings - if connector_video_embeds is None or connector_audio_embeds is None: + if is_sdr_to_hdr: + # No text encoder: the precomputed video context is the only conditioning, attended without a mask + # (`SimpleDenoiser(video_context, None)`, `hdr_ic_lora.py:467-485`). + effective_batch_size = batch_size * num_videos_per_prompt + connector_prompt_embeds = connector_video_embeds.to(device=device, dtype=self.transformer.dtype) + connector_prompt_embeds = connector_prompt_embeds.repeat_interleave(num_videos_per_prompt, dim=0) + if connector_audio_embeds is not None: + connector_audio_prompt_embeds = connector_audio_embeds.to(device=device, dtype=self.transformer.dtype) + connector_audio_prompt_embeds = connector_audio_prompt_embeds.repeat_interleave( + num_videos_per_prompt, dim=0 + ) + else: + # The reference passes no audio context at all; with the modalities isolated this placeholder only + # reaches the audio stream. + audio_context_dim = ( + self.transformer.config.caption_channels + if self.transformer.config.use_prompt_embeddings + else self.transformer.config.audio_cross_attention_dim + ) + connector_audio_prompt_embeds = torch.zeros( + (effective_batch_size, 1, audio_context_dim), device=device, dtype=self.transformer.dtype + ) + connector_attention_mask = None + elif connector_video_embeds is None or connector_audio_embeds is None: ( prompt_embeds, prompt_attention_mask, @@ -1268,12 +1541,13 @@ def __call__( height=height, width=width, num_frames=num_frames, - frame_rate=frame_rate, + frame_rate=conditioning_frame_rate, noise_scale=noise_scale, dtype=torch.float32, device=device, generator=generator, latents=latents, + input_colorspace=input_colorspace, ) # Track the base (non-reference) token count so we can trim the appended reference tokens off # `latents` before unpack/decode at the end. @@ -1283,23 +1557,34 @@ def __call__( # 5. Prepare audio latents. Audio is discarded at the end, but the transformer's audio branch still runs so # we need well-formed audio inputs. Audio guidance is fixed so no extra audio-only forward passes fire. - duration_s = num_frames / frame_rate - audio_latents_per_second = ( - self.audio_sampling_rate / self.audio_hop_length / float(self.audio_vae_temporal_compression_ratio) - ) - audio_num_frames = round(duration_s * audio_latents_per_second) - - audio_latents = self.prepare_audio_latents( - batch_size * num_videos_per_prompt, - num_channels_latents=self.audio_latent_channels, - audio_latent_length=audio_num_frames, - num_mel_bins=self.audio_mel_bins, - noise_scale=noise_scale, - dtype=torch.float32, - device=device, - generator=generator, - latents=None, - ) + if is_sdr_to_hdr: + # The reference is video-only. The audio stream is isolated from the video below, so a single + # placeholder audio token is enough and does not draw from `generator`. + audio_num_frames = 1 + latent_mel_bins = self.audio_mel_bins // self.audio_vae_mel_compression_ratio + audio_latents = torch.zeros( + (batch_size * num_videos_per_prompt, audio_num_frames, self.audio_latent_channels * latent_mel_bins), + device=device, + dtype=torch.float32, + ) + else: + duration_s = num_frames / frame_rate + audio_latents_per_second = ( + self.audio_sampling_rate / self.audio_hop_length / float(self.audio_vae_temporal_compression_ratio) + ) + audio_num_frames = round(duration_s * audio_latents_per_second) + + audio_latents = self.prepare_audio_latents( + batch_size * num_videos_per_prompt, + num_channels_latents=self.audio_latent_channels, + audio_latent_length=audio_num_frames, + num_mel_bins=self.audio_mel_bins, + noise_scale=noise_scale, + dtype=torch.float32, + device=device, + generator=generator, + latents=None, + ) # 6. Prepare timesteps sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) if sigmas is None else sigmas @@ -1310,6 +1595,8 @@ def __call__( self.scheduler.config.get("base_shift", 0.95), self.scheduler.config.get("max_shift", 2.05), ) + # The SDR-To-HDR schedule is used verbatim, so no resolution-dependent shift is passed to the scheduler. + scheduler_kwargs = {} if is_sdr_to_hdr else {"mu": mu} if self.audio_scheduler is not None: audio_scheduler = self.audio_scheduler else: @@ -1320,7 +1607,7 @@ def __call__( device, timesteps, sigmas=sigmas, - mu=mu, + **scheduler_kwargs, ) timesteps, num_inference_steps = retrieve_timesteps( self.scheduler, @@ -1328,14 +1615,19 @@ def __call__( device, timesteps, sigmas=sigmas, - mu=mu, + **scheduler_kwargs, ) num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) self._num_timesteps = len(timesteps) # 7. Prepare positional coordinates video_coords = self.transformer.rope.prepare_video_coords( - latents.shape[0], latent_num_frames, latent_height, latent_width, latents.device, fps=frame_rate + latents.shape[0], + latent_num_frames, + latent_height, + latent_width, + latents.device, + fps=conditioning_frame_rate, ) if appended_coords is not None: # Expand appended_coords to effective batch size (to [B, 3, num_extra_tokens, 2]) @@ -1348,6 +1640,17 @@ def __call__( video_coords = video_coords.repeat((2,) + (1,) * (video_coords.ndim - 1)) audio_coords = audio_coords.repeat((2,) + (1,) * (audio_coords.ndim - 1)) + # The causal VAE encodes the first latent frame from a single pixel frame. The reference marks those target + # tokens for the keyframe position embedding, and never the reference tokens (`_first_frame_keyframes_mask` + # in `ltx_core/tools.py`, `reference_video_cond.py:110-112`). Models without that embedding ignore it. + video_keyframes_mask = None + if is_sdr_to_hdr: + tokens_per_latent_frame = (latent_height // self.transformer_spatial_patch_size) * ( + latent_width // self.transformer_spatial_patch_size + ) + video_keyframes_mask = torch.zeros((latents.shape[0], latents.shape[1], 1), device=device) + video_keyframes_mask[:, :tokens_per_latent_frame] = 1.0 + # 8. Denoising loop video_seq_len = latents.shape[1] @@ -1391,15 +1694,19 @@ def __call__( num_frames=latent_num_frames, height=latent_height, width=latent_width, - fps=frame_rate, + fps=conditioning_frame_rate, audio_num_frames=audio_num_frames, video_coords=video_coords, audio_coords=audio_coords, - isolate_modalities=False, + # The reference runs video only (`VideoAudio(video=...)`, `hdr_ic_lora.py:477-485`): turning + # off the audio-to-video and video-to-audio cross-attention keeps the placeholder audio + # stream out of the video. + isolate_modalities=is_sdr_to_hdr, spatio_temporal_guidance_blocks=None, perturbation_mask=None, use_cross_timestep=use_cross_timestep, attention_kwargs=attention_kwargs, + video_keyframes_mask=video_keyframes_mask, return_dict=False, ) noise_pred_video = noise_pred_video.float() @@ -1578,7 +1885,8 @@ def __call__( ) video = latents else: - latents = latents.to(connector_prompt_embeds.dtype) + # The SDR-To-HDR reference decodes in float32 (`vae_dtype`, `hdr_ic_lora.py:189, 512-518`). + latents = latents.to(torch.float32 if is_sdr_to_hdr else connector_prompt_embeds.dtype) if not self.vae.config.timestep_conditioning: timestep = None @@ -1600,12 +1908,25 @@ def __call__( latents = self._denormalize_latents( latents, self.vae.latents_mean, self.vae.latents_std, self.vae.config.scaling_factor ) - latents = latents.to(self.vae.dtype) + if is_sdr_to_hdr: + # Decoded ACEScct codes in [-1, 1], cropped back to the requested size. + diffusion_decoder = getattr(self, "diffusion_decoder", None) + if diffusion_decoder is not None: + with self._float32(diffusion_decoder): + decoded = diffusion_decoder.decode(latents, generator=generator, return_dict=False)[0] + else: + with self._float32(self.vae): + decoded = self.vae.decode(latents, timestep, return_dict=False)[0] + decoded = decoded[:, :, :, :output_height, :output_width] + else: + latents = latents.to(self.vae.dtype) - # VAE decode returns a video tensor in the VAE's native range ([-1, 1]). - decoded = self.vae.decode(latents, timestep, return_dict=False)[0] - # HDR postprocess: LogC3 decompress → linear HDR [0, ∞). Always float32 for HDR fidelity. - video = self.hdr_video_processor.postprocess_hdr_video(decoded, output_type=output_type) + # VAE decode returns a video tensor in the VAE's native range ([-1, 1]). + decoded = self.vae.decode(latents, timestep, return_dict=False)[0] + # HDR postprocess: LogC3 or ACEScct decompress → linear HDR [0, ∞). Always float32 for HDR fidelity. + video = self.hdr_video_processor.postprocess_hdr_video( + decoded, output_type=output_type, output_colorspace=output_colorspace + ) # Audio is always None for this video-only pipeline. self.maybe_free_model_hooks() diff --git a/tests/pipelines/ltx2/test_ltx2_hdr.py b/tests/pipelines/ltx2/test_ltx2_hdr.py index 6e9d95591763..8bd745343c2c 100644 --- a/tests/pipelines/ltx2/test_ltx2_hdr.py +++ b/tests/pipelines/ltx2/test_ltx2_hdr.py @@ -55,7 +55,7 @@ class LTX2HDRPipelineTesterConfig(LTX2BaseTesterConfig): optional_input_params = frozenset( ["num_inference_steps", "num_videos_per_prompt", "generator", "latents", "output_type", "return_dict"] ) - unset_components = ("audio_scheduler",) + unset_components = ("audio_scheduler", "diffusion_decoder") def get_dummy_inputs(self): generator = self.get_generator(0) diff --git a/tests/pipelines/ltx2/test_ltx2_hdr_sdr_to_hdr.py b/tests/pipelines/ltx2/test_ltx2_hdr_sdr_to_hdr.py new file mode 100644 index 000000000000..916b3bd883e0 --- /dev/null +++ b/tests/pipelines/ltx2/test_ltx2_hdr_sdr_to_hdr.py @@ -0,0 +1,387 @@ +# Copyright 2026 The HuggingFace Team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""`LTX2HDRPipeline` with `hdr_transform="acescct"`: the LTX-2.5 SDR-To-HDR IC-LoRA path.""" + +from unittest import mock + +import pytest +import torch + +from diffusers import LTX2HDRPipeline, LTX2VideoDiffusionDecoderModel +from diffusers.pipelines.ltx2 import LTX2HDRReferenceCondition +from diffusers.pipelines.ltx2.utils import DISTILLED_SIGMA_VALUES +from diffusers.utils.import_utils import is_peft_available + +from ...testing_utils import enable_full_determinism, torch_device +from ..testing_utils.common import BasePipelineOutputMixin +from .test_ltx2_hdr import LTX2HDRPipelineTesterConfig + + +enable_full_determinism() + + +# The dummy VAE compresses by 2 in space and time, so 31 x 29 is padded to 32 x 30 and 5 frames is "2k + 1". +HEIGHT, WIDTH, NUM_FRAMES = 31, 29, 5 +PADDED_HEIGHT, PADDED_WIDTH = 32, 30 +# The dummy transformer uses LTX-2.0-style caption projections, whose input width is the tiny Gemma's hidden size. +CONTEXT_DIM = 32 + + +def get_dummy_diffusion_decoder(): + """A tiny `LTX2VideoDiffusionDecoderModel` matching the dummy VAE: 4 latent channels, x2 in space and time.""" + torch.manual_seed(0) + return LTX2VideoDiffusionDecoderModel( + out_channels=3, + latent_channels=4, + patch_size=2, + decoder_head_dim=16, + decoder_stage_channels=(32, 16, 16, 16, 16), + decoder_stage_depths=(1, 1, 1, 1, 1), + decoder_stage_kernels=((3, 3, 3),) * 4, + decoder_upsample_strides=((1, 1, 1), (2, 1, 1), (1, 1, 1), (1, 1, 1)), + decoder_upsample_channel_reductions=(2, 1, 1, 1), + decoder_stage5_kernel=(3, 3, 3), + decoder_t_emb_dim=32, + spatial_compression_ratio=2, + temporal_compression_ratio=2, + ) + + +class TestLTX2HDRPipelineSDRToHDR(LTX2HDRPipelineTesterConfig, BasePipelineOutputMixin): + def get_sdr_to_hdr_pipeline(self, text_components: bool = False, diffusion_decoder: bool = False): + components = self.get_dummy_components() + if not text_components: + # The reference runs no text encoder, so the pipeline must work without one. + components.update(text_encoder=None, tokenizer=None, connectors=None) + if diffusion_decoder: + components["diffusion_decoder"] = get_dummy_diffusion_decoder() + pipe = LTX2HDRPipeline(**components, hdr_transform="acescct") + pipe.set_progress_bar_config(disable=True) + return pipe.to(torch_device) + + def get_sdr_to_hdr_inputs(self, **overrides): + generator = self.get_generator(0) + # An 8-bit sRGB clip whose size is not a multiple of the VAE's spatial compression ratio. + frames = torch.randint(0, 256, (NUM_FRAMES, HEIGHT, WIDTH, 3), generator=generator, dtype=torch.uint8) + inputs = { + "reference_conditions": LTX2HDRReferenceCondition(frames=frames.numpy()), + # Stored without a batch dimension, like the IC-LoRA's `video_context`. + "connector_video_embeds": torch.randn(7, CONTEXT_DIM, generator=generator), + "height": HEIGHT, + "width": WIDTH, + "num_frames": NUM_FRAMES, + "frame_rate": 24.0, + "sigmas": [1.0, 0.5], + "generator": generator, + "output_type": "pt", + } + inputs.update(overrides) + return inputs + + def test_end_to_end(self): + pipe = self.get_sdr_to_hdr_pipeline() + assert pipe.text_encoder is None and pipe.tokenizer is None and pipe.connectors is None + + video = pipe(**self.get_sdr_to_hdr_inputs()).frames + assert video.shape == (1, NUM_FRAMES, HEIGHT, WIDTH, 3) + assert video.dtype == torch.float32 + assert torch.isfinite(video).all() + assert video.min() >= 0.0 + + codes = pipe(**self.get_sdr_to_hdr_inputs(output_colorspace="acescct")).frames + assert codes.min() >= 0.0 and codes.max() <= 1.0 + + def test_end_to_end_diffusion_decoder(self): + pipe = self.get_sdr_to_hdr_pipeline(diffusion_decoder=True) + with mock.patch.object(pipe.vae, "decode", side_effect=AssertionError("the VAE must not decode")): + video = pipe(**self.get_sdr_to_hdr_inputs()).frames + assert video.shape == (1, NUM_FRAMES, HEIGHT, WIDTH, 3) + assert torch.isfinite(video).all() + + def test_no_text_encoder_call(self): + pipe = self.get_sdr_to_hdr_pipeline(text_components=True) + forbidden = AssertionError("the SDR-To-HDR path must not encode a prompt") + with ( + mock.patch.object(pipe, "encode_prompt", side_effect=forbidden), + mock.patch.object(pipe.text_encoder, "forward", side_effect=forbidden), + mock.patch.object(pipe.connectors, "forward", side_effect=forbidden), + ): + pipe(**self.get_sdr_to_hdr_inputs()) + + @pytest.mark.parametrize( + "overrides, match", + [ + ({"prompt": "a robot dancing"}, "prompt"), + ({"connector_video_embeds": None}, "connector_video_embeds"), + ({"guidance_scale": 3.0}, "guidance"), + ({"stg_scale": 1.0}, "guidance"), + ({"modality_scale": 2.0}, "guidance"), + ({"reference_conditions": None}, "reference"), + ({"num_frames": 4}, "num_frames"), + ({"num_frames": 7}, "fewer than"), + ], + ) + def test_invalid_inputs(self, overrides, match): + pipe = self.get_sdr_to_hdr_pipeline() + with pytest.raises(ValueError, match=match): + pipe(**self.get_sdr_to_hdr_inputs(**overrides)) + + def test_batched_scene_embedding_matches_unbatched(self): + pipe = self.get_sdr_to_hdr_pipeline() + inputs = self.get_sdr_to_hdr_inputs(output_type="latent") + unbatched = pipe(**inputs).frames + inputs = self.get_sdr_to_hdr_inputs(output_type="latent") + inputs["connector_video_embeds"] = inputs["connector_video_embeds"][None] + batched = pipe(**inputs).frames + assert torch.equal(unbatched, batched) + + @staticmethod + def _perturb_audio(pipe, seed): + """Feed the transformer random audio latents and audio context instead of what the pipeline built.""" + forward = pipe.transformer.forward + + def perturbed_forward(*args, **kwargs): + generator = torch.Generator().manual_seed(seed) + for name in ("audio_hidden_states", "audio_encoder_hidden_states"): + value = kwargs[name] + kwargs[name] = torch.randn(value.shape, generator=generator).to(value.device, value.dtype) * 3.0 + return forward(*args, **kwargs) + + return mock.patch.object(pipe.transformer, "forward", side_effect=perturbed_forward) + + def test_audio_does_not_influence_video(self): + pipe = self.get_sdr_to_hdr_pipeline() + outputs = [] + for seed in (0, 1): + with self._perturb_audio(pipe, seed): + outputs.append(pipe(**self.get_sdr_to_hdr_inputs(output_type="latent")).frames) + assert torch.equal(outputs[0], outputs[1]) + + # Passing the IC-LoRA's `audio_context` (the reference ignores it) does not change the video either. + with_audio_context = pipe( + **self.get_sdr_to_hdr_inputs(output_type="latent", connector_audio_embeds=torch.randn(7, CONTEXT_DIM)) + ).frames + without_audio_context = pipe(**self.get_sdr_to_hdr_inputs(output_type="latent")).frames + assert torch.equal(with_audio_context, without_audio_context) + + def test_audio_influences_video_with_logc3(self): + # Sanity check of `_perturb_audio`: the LogC3 path keeps audio-to-video cross-attention. + pipe = self.get_pipeline().to(torch_device) + outputs = [] + for seed in (0, 1): + with self._perturb_audio(pipe, seed): + inputs = self.get_dummy_inputs() + inputs["output_type"] = "latent" + outputs.append(pipe(**inputs).frames) + assert not torch.allclose(outputs[0], outputs[1]) + + def test_vae_runs_in_float32(self): + pipe = self.get_sdr_to_hdr_pipeline().to(dtype=torch.bfloat16) + seen = {} + + def spy(name, fn): + def wrapped(x, *args, **kwargs): + seen[name] = (x.dtype, pipe.vae.dtype) + return fn(x, *args, **kwargs) + + return wrapped + + with ( + mock.patch.object(pipe.vae, "encode", side_effect=spy("encode", pipe.vae.encode)), + mock.patch.object(pipe.vae, "decode", side_effect=spy("decode", pipe.vae.decode)), + ): + video = pipe(**self.get_sdr_to_hdr_inputs()).frames + + assert seen == {"encode": (torch.float32, torch.float32), "decode": (torch.float32, torch.float32)} + # The VAE is handed back in the dtype it was loaded in. + assert pipe.vae.dtype == torch.bfloat16 + assert pipe.transformer.dtype == torch.bfloat16 + assert video.dtype == torch.float32 + + def test_diffusion_decoder_runs_in_float32(self): + pipe = self.get_sdr_to_hdr_pipeline(diffusion_decoder=True).to(dtype=torch.bfloat16) + decoder = pipe.diffusion_decoder + seen = {} + decode = decoder.decode + + def spy(z, *args, **kwargs): + seen["decode"] = (z.dtype, decoder.dtype) + return decode(z, *args, **kwargs) + + with mock.patch.object(decoder, "decode", side_effect=spy): + pipe(**self.get_sdr_to_hdr_inputs()) + assert seen["decode"] == (torch.float32, torch.float32) + assert decoder.dtype == torch.bfloat16 + + def test_reflect_pad_and_crop_back(self): + pipe = self.get_sdr_to_hdr_pipeline() + seen = {} + encode, decode = pipe.vae.encode, pipe.vae.decode + + def encode_spy(x, *args, **kwargs): + seen["encoded_pixels"] = x.detach().clone() + return encode(x, *args, **kwargs) + + def decode_spy(*args, **kwargs): + output = decode(*args, **kwargs) + seen["decoded_shape"] = tuple(output[0].shape) + return output + + with ( + mock.patch.object(pipe.vae, "encode", side_effect=encode_spy), + mock.patch.object(pipe.vae, "decode", side_effect=decode_spy), + ): + video = pipe(**self.get_sdr_to_hdr_inputs()).frames + + pixels = seen["encoded_pixels"] + assert pixels.shape == (1, 3, NUM_FRAMES, PADDED_HEIGHT, PADDED_WIDTH) + # Bottom/right reflect padding: the padded row and column mirror the ones before the last source row/column. + assert torch.equal(pixels[..., HEIGHT, :], pixels[..., HEIGHT - 2, :]) + assert torch.equal(pixels[..., :, WIDTH], pixels[..., :, WIDTH - 2]) + # Input transform: 8-bit sRGB maps into ACEScct, so black is not -1 and white is not +1 in VAE range. + assert pixels.min() >= 2 * 0.0729055341958355 - 1 - 1e-6 + assert pixels.max() <= 2 * 0.5547945205479452 - 1 + 1e-6 + + assert seen["decoded_shape"] == (1, 3, NUM_FRAMES, PADDED_HEIGHT, PADDED_WIDTH) + assert video.shape == (1, NUM_FRAMES, HEIGHT, WIDTH, 3) + + latents = pipe(**self.get_sdr_to_hdr_inputs(output_type="latent")).frames + assert latents.shape[-2:] == (PADDED_HEIGHT // 2, PADDED_WIDTH // 2) + + @pytest.mark.parametrize("frame_rate, expected_rope_fps", [(24.0, 24.0), (30.0, 30.0), (50.0, 30.0), (60.0, 30.0)]) + def test_rope_frame_rate_is_capped_at_30(self, frame_rate, expected_rope_fps): + pipe = self.get_sdr_to_hdr_pipeline() + prepare_video_coords = pipe.transformer.rope.prepare_video_coords + with mock.patch.object( + pipe.transformer.rope, "prepare_video_coords", side_effect=prepare_video_coords + ) as coords_spy: + pipe(**self.get_sdr_to_hdr_inputs(frame_rate=frame_rate, output_type="latent")) + # The reference latents and the generated latents share the same time base. + assert [call.kwargs["fps"] for call in coords_spy.call_args_list] == [expected_rope_fps] * 2 + + def test_frame_rates_above_30_condition_like_30(self): + pipe = self.get_sdr_to_hdr_pipeline() + outputs = { + frame_rate: pipe(**self.get_sdr_to_hdr_inputs(frame_rate=frame_rate, output_type="latent")).frames + for frame_rate in (24.0, 30.0, 60.0) + } + assert torch.equal(outputs[30.0], outputs[60.0]) + assert not torch.equal(outputs[24.0], outputs[30.0]) + + def test_distilled_sigmas_by_default(self): + pipe = self.get_sdr_to_hdr_pipeline() + inputs = self.get_sdr_to_hdr_inputs(output_type="latent") + inputs.pop("sigmas") + forward = pipe.transformer.forward + with mock.patch.object(pipe.transformer, "forward", side_effect=forward) as spy: + pipe(**inputs) + # One unguided forward per distilled sigma, each at that exact sigma. + sigmas = [call.kwargs["sigma"].item() / 1000.0 for call in spy.call_args_list] + assert sigmas == pytest.approx(DISTILLED_SIGMA_VALUES, abs=1e-6) + + def test_first_latent_frame_is_marked_as_keyframe(self): + pipe = self.get_sdr_to_hdr_pipeline() + forward = pipe.transformer.forward + with mock.patch.object(pipe.transformer, "forward", side_effect=forward) as spy: + pipe(**self.get_sdr_to_hdr_inputs(output_type="latent")) + kwargs = spy.call_args.kwargs + assert kwargs["isolate_modalities"] is True + mask = kwargs["video_keyframes_mask"] + latent_frames = (NUM_FRAMES - 1) // 2 + 1 + tokens_per_frame = (PADDED_HEIGHT // 2) * (PADDED_WIDTH // 2) + # [generated | reference] tokens: only the first latent frame of the generated video is marked. (The dummy + # VAE only downsamples in space, so its reference latents keep all 5 frames.) + assert mask.shape == (1, kwargs["hidden_states"].shape[1], 1) + assert mask.shape[1] > latent_frames * tokens_per_frame + assert mask[:, :tokens_per_frame].eq(1).all() + assert mask[:, tokens_per_frame:].eq(0).all() + + def test_hdr_transform_survives_save_load(self, tmp_path): + pipe = self.get_sdr_to_hdr_pipeline(text_components=True) + pipe.save_pretrained(str(tmp_path), safe_serialization=False) + loaded = LTX2HDRPipeline.from_pretrained(str(tmp_path)) + assert loaded.hdr_video_processor.config.hdr_transform == "acescct" + + @pytest.mark.skipif(not is_peft_available(), reason="PEFT is required for LoRA loading.") + def test_sdr_to_hdr_lora_key_format(self): + # The LTX-2.5 SDR-To-HDR IC-LoRA stores ComfyUI-style keys: `diffusion_model.transformer_blocks.{i}.` followed + # by these 20 suffixes (rank 128, no alphas), for the video self-attention, text cross-attention and FFN. + suffixes = [ + f"{module}.lora_{side}.weight" + for module in ( + "attn1.to_q", + "attn1.to_k", + "attn1.to_v", + "attn1.to_out.0", + "attn2.to_q", + "attn2.to_k", + "attn2.to_v", + "attn2.to_out.0", + "ff.net.0.proj", + "ff.net.2", + ) + for side in ("A", "B") + ] + pipe = self.get_sdr_to_hdr_pipeline() + modules = dict(pipe.transformer.named_modules()) + generator = torch.Generator().manual_seed(0) + state_dict, expected = {}, set() + for block in range(len(pipe.transformer.transformer_blocks)): + for suffix in suffixes: + module_name = f"transformer_blocks.{block}.{suffix.rsplit('.lora_', 1)[0]}" + linear = modules[module_name] + shape = (4, linear.in_features) if ".lora_A." in suffix else (linear.out_features, 4) + state_dict[f"diffusion_model.transformer_blocks.{block}.{suffix}"] = torch.randn( + shape, generator=generator + ) + expected.add(module_name) + + base = pipe(**self.get_sdr_to_hdr_inputs(output_type="latent")).frames + pipe.load_lora_weights(state_dict, adapter_name="sdr_to_hdr") + pipe.set_adapters("sdr_to_hdr", 1.0) + + from peft.tuners.tuners_utils import BaseTunerLayer + + adapted = {name for name, module in pipe.transformer.named_modules() if isinstance(module, BaseTunerLayer)} + assert adapted == expected + with_lora = pipe(**self.get_sdr_to_hdr_inputs(output_type="latent")).frames + assert not torch.allclose(base, with_lora) + + +class TestLTX2HDRPipelineLogC3Unchanged(LTX2HDRPipelineTesterConfig, BasePipelineOutputMixin): + def test_default_hdr_transform_is_logc3(self): + pipe = self.get_pipeline() + assert pipe.hdr_video_processor.config.hdr_transform == "logc3" + assert pipe.config.hdr_transform == "logc3" + + @pytest.mark.parametrize("argument", ["input_colorspace", "output_colorspace"]) + def test_colorspaces_rejected_with_logc3(self, argument): + pipe = self.get_pipeline().to(torch_device) + inputs = self.get_dummy_inputs() + inputs[argument] = "srgb_gamma" if argument == "input_colorspace" else "acescg" + with pytest.raises(ValueError, match="only supported with `hdr_transform='acescct'`"): + pipe(**inputs) + + def test_logc3_output_slice(self): + # Recorded on the parent commit, before the SDR-To-HDR path was added. + pipe = self.get_pipeline().to(torch_device) + inputs = self.get_dummy_inputs() + inputs["output_type"] = "latent" + latents = pipe(**inputs).frames.flatten().cpu() + expected = torch.tensor( + [-1.4695, -0.8892, -1.27, 1.4034, 0.5713, 1.7472, 0.2216, 0.9364] + + [-0.6652, -0.8569, 0.755, 0.1796, 0.1709, -0.1783, -0.1346, -0.4674] + ) + assert torch.allclose(torch.cat([latents[:8], latents[-8:]]), expected, atol=1e-3) From da6aa638e9b629c39a7e69f3d2aef2d169402300 Mon Sep 17 00:00:00 2001 From: christopher5106 Date: Tue, 6 Oct 2026 19:51:40 +0200 Subject: [PATCH 3/4] [LTX-2.5] Keyframe-aware diffusion decoding Port Lightricks' keyframe-aware DiffVAE decode (LTX-2 @ 9ec55f9f, `ltx_core/model/video_vae/{keyframes,diffusion_video_decoder}.py`, `transformer/fallback_na/joint_eager.py`) to `LTX2VideoDiffusionDecoderModel`. - `decoder_keyframe_type_embedding` config flag (default off) creates the learned keyframe tag `decoder.type_emb` `(latent_channels,)`; existing checkpoints load unchanged. - `decode(..., keyframe_latents=, keyframe_frame_indices=)`: planes are tagged, share `conv_in` and every stage with the video, and attend jointly with it (each video position sees its 2 nearest planes, each plane its 2 nearest frames, one softmax), with per-stage plane times. Tiled decode keeps the planes inside each temporal tile plus the nearest on each side, with times rebased on the tile. Without planes the decode is bit-identical to before. - Joint attention is a port of the reference's pure-torch backend, bitwise equal to it; full-decoder parity on random tiny weights is within 2.4e-6. - Converter carries `decoder.type_emb` and sets the flag when present (the current original checkpoint has it, so the strict load used to fail). - `dfr_layout.resolve_seam_positions`: seam keyframes of a single-window clip (24/32 segments clipped to the clip, high-quality doubling), equal to the reference for every 8k+1 frame count up to 2001. Co-Authored-By: Claude Opus 5.5 --- .../en/api/models/ltx2_diffusion_decoder.md | 22 + scripts/convert_ltx2_to_diffusers.py | 4 + .../autoencoders/ltx2_diffusion_decoder.py | 930 +++++++++++++++++- src/diffusers/pipelines/ltx2/dfr_layout.py | 32 + .../test_models_ltx2_diffusion_decoder.py | 268 +++++ tests/pipelines/ltx2/test_ltx2_hdr_seams.py | 74 ++ 6 files changed, 1311 insertions(+), 19 deletions(-) create mode 100644 tests/pipelines/ltx2/test_ltx2_hdr_seams.py diff --git a/docs/source/en/api/models/ltx2_diffusion_decoder.md b/docs/source/en/api/models/ltx2_diffusion_decoder.md index 9a4059166267..9f33ae5607c7 100644 --- a/docs/source/en/api/models/ltx2_diffusion_decoder.md +++ b/docs/source/en/api/models/ltx2_diffusion_decoder.md @@ -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 diff --git a/scripts/convert_ltx2_to_diffusers.py b/scripts/convert_ltx2_to_diffusers.py index 33b91790ef1c..69d66198cbe3 100644 --- a/scripts/convert_ltx2_to_diffusers.py +++ b/scripts/convert_ltx2_to_diffusers.py @@ -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) diff --git a/src/diffusers/models/autoencoders/ltx2_diffusion_decoder.py b/src/diffusers/models/autoencoders/ltx2_diffusion_decoder.py index 41388991e0b4..5a238a5c2b5d 100644 --- a/src/diffusers/models/autoencoders/ltx2_diffusion_decoder.py +++ b/src/diffusers/models/autoencoders/ltx2_diffusion_decoder.py @@ -90,6 +90,505 @@ def mask_mod(batch_idx, head_idx, q_idx, kv_idx): return create_block_mask(mask_mod, B=None, H=None, Q_LEN=seq_len, KV_LEN=seq_len, device=device) +# -------------------------------------------------------------------------------------------------------------------- +# Keyframe-aware decoding +# +# A keyframe decode carries a second stream through the decoder: a stack of single-frame latent *planes* `(B, P, H, +# W, C)` whose plane axis sits in the video's temporal slot. Every weight is shared between the two streams, and they +# only mix inside one joint neighborhood-attention softmax, where each video query also sees the same spatial window +# on its two nearest planes and each plane query also sees its two nearest video frames. +# -------------------------------------------------------------------------------------------------------------------- + +# Keyframe planes visible to one video query, and video frames visible to one plane query. +_KEYFRAME_CONTEXT_SLOTS = 2 + + +def _keyframe_stage_times(pixel_frame_indices: torch.Tensor, remaining_time_stride: int) -> torch.Tensor: + """Chunk-center position of each keyframe plane in the temporal units of one decoder stage. + + A stage whose remaining temporal upsampling is `r` has cells covering `r` pixel frames each, except cell 0 which + covers only pixel frame 0 (the causal first frame). So `t(0) = 0` and `t(f) = (f + (r - 1) / 2) / r`, the center of + the cell holding `f`. In the diffusion stage `r == 1`, making the times the raw pixel indices. + """ + frames = pixel_frame_indices.to(torch.float32) + times = (frames + (remaining_time_stride - 1) / 2) / remaining_time_stride + return torch.where(frames == 0, torch.zeros_like(times), times) + + +def _keyframe_planes_for_tile(pixel_frame_indices: torch.Tensor, frame_lo: int, frame_hi: int) -> torch.Tensor: + """`(P,)` bool: the planes a tile spanning pixel frames `[frame_lo, frame_hi]` (inclusive) has to carry. + + Every plane inside the span plus the nearest plane on each side outside it. The two outside planes are what keep a + tiled decode consistent with a whole one: a frame near a tile edge ranks its planes by temporal distance, so + dropping the closest plane beyond the edge would make it attend to a farther one instead. + """ + indices = pixel_frame_indices.to(torch.int64) + keep = (indices >= frame_lo) & (indices <= frame_hi) + before = indices < frame_lo + if bool(before.any()): + keep[int(torch.where(before, indices, torch.full_like(indices, -1)).argmax())] = True + after = indices > frame_hi + if bool(after.any()): + sentinel = int(indices.max()) + 1 + keep[int(torch.where(after, indices, torch.full_like(indices, sentinel)).argmin())] = True + return keep + + +def _nearest_slots(query_times: torch.Tensor, candidate_times: torch.Tensor, num_slots: int) -> torch.Tensor: + """`(Q, num_slots)` candidate indices ranked by `(|dt|, index)`, `-1` where there are fewer candidates.""" + distances = (query_times[:, None] - candidate_times[None, :]).abs().to(torch.float32) + # A stable sort breaks distance ties by ascending candidate index. + order = torch.argsort(distances, dim=-1, stable=True) + take = min(num_slots, candidate_times.shape[0]) + chosen = order[:, :take] + if take < num_slots: + pad = torch.full((chosen.shape[0], num_slots - take), -1, dtype=chosen.dtype, device=chosen.device) + chosen = torch.cat([chosen, pad], dim=1) + return chosen + + +def _upsample_keyframe_planes(upsample: nn.Module, hidden_states: torch.Tensor) -> torch.Tensor: + """Spatially upsample keyframe planes `(B, P, H, W, C)` with the video stream's upsampler. + + Each plane goes through as its own one-frame clip with the leading frame always dropped, so a temporal stride of 2 + expands it to two frames and takes it back to one: the plane count never changes, only `H` and `W` grow. + """ + batch_size, num_planes = hidden_states.shape[:2] + flat = hidden_states.reshape(batch_size * num_planes, 1, *hidden_states.shape[2:]) + upsampled = upsample(flat, drop_leading_frame=True) + return upsampled.reshape(batch_size, num_planes, *upsampled.shape[2:]) + + +# Joint neighborhood attention. Queries are grouped into `(bt, bh, bw)` bricks and many bricks share one +# `scaled_dot_product_attention` call as its batch dimension. All queries in a brick share one gathered key slab, the +# visible-key pattern is one mask shared by every brick, and per-key validity (outside the volume, an empty slot) +# rides in an extra key channel that adds `_JOINT_DEAD_KEY` to the score of a dead key. +_JOINT_DEAD_KEY = -1.0e4 +# Keep the query/key head dim a multiple of this, so SDPA can keep a fused kernel once the bias channel is added. +_JOINT_HEAD_DIM_ALIGN = 8 +_JOINT_BRICK_QUERIES = 64 +_JOINT_BRICK_DEPTH = 4 +# Transient memory budget of one staging pass and one key/value block. +_JOINT_WORKSPACE_BYTES = 256 * 1024**2 +# Peak-to-staging multipliers for SDPA kernels that keep the scores on chip, and for those that materialize them. +_JOINT_STAGING_FACTOR_FUSED = 4.75 +_JOINT_STAGING_FACTOR_MATERIALIZED = 22.1 + + +def _joint_key_channels(head_dim: int) -> int: + return -(-(head_dim + 1) // _JOINT_HEAD_DIM_ALIGN) * _JOINT_HEAD_DIM_ALIGN + + +def _joint_window(kernel: int) -> tuple[int, int]: + """`(lo, hi)` halo of one axis: a centered window of `kernel` positions.""" + lo = kernel // 2 + return lo, kernel - lo - 1 + + +class _JointGeometry: + """Brick decomposition of one `(H, W)` grid, plus the padding and slab extents it implies.""" + + def __init__(self, height: int, width: int, kernel: tuple[int, int, int], brick: tuple[int, int, int]): + kernel_t, kernel_h, kernel_w = kernel + lo_h, hi_h = _joint_window(kernel_h) + lo_w, hi_w = _joint_window(kernel_w) + self.height, self.width = height, width + self.brick = brick + self.kernel = kernel + self.grid = (-(-height // brick[1]), -(-width // brick[2])) + self.span_t = brick[0] + kernel_t - 1 + self.span = (brick[1] + kernel_h - 1, brick[2] + kernel_w - 1) + self.pad_h = (lo_h, hi_h + self.grid[0] * brick[1] - height) + self.pad_w = (lo_w, hi_w + self.grid[1] * brick[2] - width) + self.pad_t = _joint_window(kernel_t) + self.queries = brick[0] * brick[1] * brick[2] + self.footprint = self.span[0] * self.span[1] + self.padded_height = height + sum(self.pad_h) + self.padded_width = width + sum(self.pad_w) + + def row_extent(self, rows: int) -> int: + return (rows - 1) * self.brick[1] + self.span[0] + + +class _JointSchedule: + """How the loops are cut so transient memory stays within `_JOINT_WORKSPACE_BYTES`.""" + + def __init__( + self, + geometry: _JointGeometry, + blocks: int, + heads: int, + head_dim: int, + axis_bricks: int, + element_size: int, + factor: float, + ): + channels = _joint_key_channels(head_dim) + keys = blocks * geometry.footprint + pair_bytes = geometry.grid[1] * heads * keys * (channels + head_dim) * element_size + pairs = max(1, int(_JOINT_WORKSPACE_BYTES / max(pair_bytes * factor, 1.0))) + if pairs >= geometry.grid[0]: + self.group_axis = min(axis_bricks, max(1, pairs // geometry.grid[0])) + self.group_rows = geometry.grid[0] + else: + self.group_axis = 1 + self.group_rows = pairs + staged = geometry.padded_height * geometry.padded_width * heads * (channels + head_dim) * element_size + per_axis_brick = staged * geometry.brick[0] + self.stage_axis = min(axis_bricks, max(self.group_axis, _JOINT_WORKSPACE_BYTES // max(per_axis_brick, 1))) + + +def _joint_banded(queries: int, span: int, kernel: int, device: torch.device) -> torch.Tensor: + key = torch.arange(span, device=device)[None, :] + query = torch.arange(queries, device=device)[:, None] + return (key >= query) & (key < query + kernel) + + +def _joint_mask(geometry: _JointGeometry, num_slots: int, device: torch.device) -> torch.Tensor: + """`(1, 1, Nq, Nk)` visibility shared by every brick: the video slab, then `num_slots` plane slabs. + + Plane keys carry no temporal condition, which is what lets a whole brick share one set of planes. + """ + brick_t, brick_h, brick_w = geometry.brick + kernel_t, kernel_h, kernel_w = geometry.kernel + spatial = ( + _joint_banded(brick_h, geometry.span[0], kernel_h, device)[:, None, :, None] + & _joint_banded(brick_w, geometry.span[1], kernel_w, device)[None, :, None, :] + ).reshape(brick_h * brick_w, geometry.footprint) + temporal = _joint_banded(brick_t, geometry.span_t, kernel_t, device) + video = (temporal[:, None, :, None] & spatial[None, :, None, :]).reshape( + geometry.queries, geometry.span_t * geometry.footprint + ) + planes = ( + spatial[None, :, None, :] + .expand(brick_t, brick_h * brick_w, num_slots, geometry.footprint) + .reshape(geometry.queries, num_slots * geometry.footprint) + ) + return torch.cat([video, planes], dim=1)[None, None].contiguous() + + +def _joint_stage( + x: torch.Tensor, geometry: _JointGeometry, pad_t: tuple[int, int], with_bias_channel: bool +) -> torch.Tensor: + """`(B, A, H, W, heads, head_dim)` to a padded head-major `(B, heads, A + pad, Hp, Wp, C)`.""" + batch, axis, height, width, heads, head_dim = x.shape + channels = _joint_key_channels(head_dim) if with_bias_channel else head_dim + out = x.new_zeros((batch, heads, axis + sum(pad_t), geometry.padded_height, geometry.padded_width, channels)) + if with_bias_channel: + out[..., head_dim] = _JOINT_DEAD_KEY + live = out[ + :, + :, + pad_t[0] : pad_t[0] + axis, + geometry.pad_h[0] : geometry.pad_h[0] + height, + geometry.pad_w[0] : geometry.pad_w[0] + width, + ] + live[..., :head_dim] = x.permute(0, 4, 1, 2, 3, 5) + if with_bias_channel: + live[..., head_dim] = 0.0 + return out + + +def _joint_slabs( + staged: torch.Tensor, geometry: _JointGeometry, bricks: int, rows: int, blocks: int, group_stride: int +) -> torch.Tensor: + """Overlapping brick slabs as a view, `(B, bricks, rows, Gw, heads, blocks, eh, ew, C)`.""" + batch, heads = staged.shape[0], staged.shape[1] + stride_b, stride_nh, stride_a, stride_h, stride_w, _ = staged.stride() + return staged.as_strided( + (batch, bricks, rows, geometry.grid[1], heads, blocks, *geometry.span, staged.shape[-1]), + ( + stride_b, + group_stride * stride_a, + geometry.brick[1] * stride_h, + geometry.brick[2] * stride_w, + stride_nh, + stride_a, + stride_h, + stride_w, + 1, + ), + ) + + +def _joint_query_bricks(x: torch.Tensor, geometry: _JointGeometry, bricks: int, rows: int) -> torch.Tensor: + """`(B, A, h, W, heads, head_dim)` to `(B * bricks * rows * Gw, heads, Nq, C)`, with a constant bias channel.""" + batch, axis, height, width, heads, head_dim = x.shape + brick_t, brick_h, brick_w = geometry.brick + pad_t, pad_h, pad_w = bricks * brick_t - axis, rows * brick_h - height, geometry.grid[1] * brick_w - width + if pad_t or pad_h or pad_w: + x = F.pad(x, (0, 0, 0, 0, 0, pad_w, 0, pad_h, 0, pad_t)) + bricked = ( + x.reshape(batch, bricks, brick_t, rows, brick_h, geometry.grid[1], brick_w, heads, head_dim) + .permute(0, 1, 3, 5, 7, 2, 4, 6, 8) + .reshape(batch * bricks * rows * geometry.grid[1], heads, geometry.queries, head_dim) + ) + out = bricked.new_zeros((*bricked.shape[:-1], _joint_key_channels(head_dim))) + out[..., :head_dim] = bricked + out[..., head_dim] = 1.0 + return out + + +def _joint_unbrick( + attended: torch.Tensor, geometry: _JointGeometry, batch: int, bricks: int, rows: int, extent: tuple[int, int] +) -> torch.Tensor: + brick_t, brick_h, brick_w = geometry.brick + heads, head_dim = attended.shape[1], attended.shape[3] + plane = ( + attended.reshape(batch, bricks, rows, geometry.grid[1], heads, brick_t, brick_h, brick_w, head_dim) + .permute(0, 1, 5, 2, 6, 3, 7, 4, 8) + .reshape(batch, bricks * brick_t, rows * brick_h, geometry.grid[1] * brick_w, heads, head_dim) + ) + return plane[:, : extent[0], : extent[1], : geometry.width] + + +def _joint_with_null(slots: torch.Tensor, null_index: int) -> torch.Tensor: + return torch.where(slots < 0, torch.full_like(slots, null_index), slots) + + +def _joint_slot_runs(slots: torch.Tensor) -> list[tuple[int, int]]: + """Maximal `[start, stop)` runs of leading-axis positions whose slot row is identical.""" + rows = slots.tolist() + runs = [] + start = 0 + for index in range(1, len(rows)): + if rows[index] != rows[start]: + runs.append((start, index)) + start = index + runs.append((start, len(rows))) + return runs + + +def _joint_attend_group( + query_slice: torch.Tensor, + key_views: tuple[torch.Tensor, ...], + value_views: tuple[torch.Tensor, ...], + geometry: _JointGeometry, + shape: tuple[int, int], + mask: torch.Tensor, +) -> torch.Tensor: + """Gather one `(bricks, brick rows)` block's keys, attend, and un-brick the result.""" + bricks, rows = shape + batch = query_slice.shape[0] + heads, head_dim = query_slice.shape[4], query_slice.shape[5] + blocks = sum(view.shape[5] for view in key_views) + channels = _joint_key_channels(head_dim) + keys = query_slice.new_empty((batch, bricks, rows, geometry.grid[1], heads, blocks, *geometry.span, channels)) + values = query_slice.new_empty((batch, bricks, rows, geometry.grid[1], heads, blocks, *geometry.span, head_dim)) + start = 0 + for key_view, value_view in zip(key_views, value_views): + stop = start + key_view.shape[5] + keys[:, :, :, :, :, start:stop].copy_(key_view) + values[:, :, :, :, :, start:stop].copy_(value_view) + start = stop + count = batch * bricks * rows * geometry.grid[1] + # `scale=1.0`: the query is already scaled in `project_qkv`. + attended = F.scaled_dot_product_attention( + _joint_query_bricks(query_slice, geometry, bricks, rows), + keys.view(count, heads, blocks * geometry.footprint, channels), + values.view(count, heads, blocks * geometry.footprint, head_dim), + attn_mask=mask, + scale=1.0, + ) + return _joint_unbrick(attended, geometry, batch, bricks, rows, (query_slice.shape[1], query_slice.shape[2])) + + +def _joint_row_groups(geometry: _JointGeometry, schedule: _JointSchedule) -> list[tuple[int, slice, slice]]: + """`(rows, staged H slice, output H slice)` per group of brick rows.""" + brick_h = geometry.brick[1] + groups = [] + for row in range(0, geometry.grid[0], schedule.group_rows): + rows = min(schedule.group_rows, geometry.grid[0] - row) + groups.append( + ( + rows, + slice(row * brick_h, row * brick_h + geometry.row_extent(rows)), + slice(row * brick_h, min((row + rows) * brick_h, geometry.height)), + ) + ) + return groups + + +def _joint_video_query_pass( + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + keyframe_key: torch.Tensor, + keyframe_value: torch.Tensor, + slots: torch.Tensor, + geometry: _JointGeometry, + factor: float, +) -> torch.Tensor: + """Video queries: their local `Kt x Kh x Kw` window plus the same `Kh x Kw` window on their nearest planes.""" + num_frames, heads, head_dim = query.shape[1], query.shape[4], query.shape[5] + brick_t = geometry.brick[0] + lo_t, hi_t = geometry.pad_t + num_slots = slots.shape[1] + + plane_keys = _joint_stage(keyframe_key, geometry, (0, 0), with_bias_channel=True) + plane_values = _joint_stage(keyframe_value, geometry, (0, 0), with_bias_channel=False) + # One all-dead plane appended for empty (`-1`) slots to point at. + null_shape = (*plane_keys.shape[:2], 1, *plane_keys.shape[3:]) + null_key = plane_keys.new_zeros(null_shape) + null_key[..., head_dim] = _JOINT_DEAD_KEY + plane_keys = torch.cat([plane_keys, null_key], dim=2) + plane_values = torch.cat([plane_values, plane_values.new_zeros((*null_shape[:-1], head_dim))], dim=2) + slot_table = _joint_with_null(slots, keyframe_key.shape[1]) + + mask = _joint_mask(geometry, num_slots, query.device) + schedule = _JointSchedule( + geometry, geometry.span_t + num_slots, heads, head_dim, -(-num_frames // brick_t), query.element_size(), factor + ) + row_groups = _joint_row_groups(geometry, schedule) + + out = torch.empty_like(query) + for run_start, run_stop in _joint_slot_runs(slot_table): + # Every brick inside a run sees the same planes, so they are gathered once per run. + planes = plane_keys.index_select(2, slot_table[run_start]) + plane_vals = plane_values.index_select(2, slot_table[run_start]) + run_bricks = -(-(run_stop - run_start) // brick_t) + for staged_brick in range(0, run_bricks, schedule.stage_axis): + staged_bricks = min(schedule.stage_axis, run_bricks - staged_brick) + first = run_start + staged_brick * brick_t + last = first + staged_bricks * brick_t + source = slice(max(0, first - lo_t), min(num_frames, last + hi_t)) + pad_t = (max(0, lo_t - first), max(0, last + hi_t - num_frames)) + window_keys = _joint_stage(key[:, source], geometry, pad_t, with_bias_channel=True) + window_values = _joint_stage(value[:, source], geometry, pad_t, with_bias_channel=False) + + for brick in range(staged_brick, staged_brick + staged_bricks, schedule.group_axis): + count = min(schedule.group_axis, staged_brick + staged_bricks - brick) + start = run_start + brick * brick_t + stop = min(start + count * brick_t, run_stop) + offset = (brick - staged_brick) * brick_t + for rows, key_rows, out_rows in row_groups: + out[:, start:stop, out_rows] = _joint_attend_group( + query[:, start:stop, out_rows], + ( + _joint_slabs( + window_keys[:, :, offset:, key_rows], geometry, count, rows, geometry.span_t, brick_t + ), + _joint_slabs(planes[:, :, :, key_rows], geometry, count, rows, num_slots, 0), + ), + ( + _joint_slabs( + window_values[:, :, offset:, key_rows], geometry, count, rows, geometry.span_t, brick_t + ), + _joint_slabs(plane_vals[:, :, :, key_rows], geometry, count, rows, num_slots, 0), + ), + geometry, + (count, rows), + mask, + ) + return out + + +def _joint_keyframe_query_pass( + keyframe_query: torch.Tensor, + keyframe_key: torch.Tensor, + keyframe_value: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + slots: torch.Tensor, + geometry: _JointGeometry, + factor: float, +) -> torch.Tensor: + """Plane queries: the `Kh x Kw` window on their own plane plus the same window on their nearest video frames.""" + num_planes, heads, head_dim = keyframe_query.shape[1], keyframe_query.shape[4], keyframe_query.shape[5] + num_slots = slots.shape[1] + num_frames = key.shape[1] + flat = _JointGeometry(geometry.height, geometry.width, (1, *geometry.kernel[1:]), (1, *geometry.brick[1:])) + + # Stage only the frames some plane points at; `unique` doubles as the remap of the slot table. + wanted, inverse = torch.unique(_joint_with_null(slots, num_frames).reshape(-1), return_inverse=True) + frame_keys = _joint_stage( + key.index_select(1, wanted.clamp(max=num_frames - 1)), flat, (0, 0), with_bias_channel=True + ) + frame_values = _joint_stage( + value.index_select(1, wanted.clamp(max=num_frames - 1)), flat, (0, 0), with_bias_channel=False + ) + # An empty slot was clamped onto a real frame above; kill it here. + frame_keys[:, :, wanted == num_frames, ..., head_dim] = _JOINT_DEAD_KEY + own_keys = _joint_stage(keyframe_key, flat, (0, 0), with_bias_channel=True) + own_values = _joint_stage(keyframe_value, flat, (0, 0), with_bias_channel=False) + slot_table = inverse.reshape(num_planes, num_slots) + mask = _joint_mask(flat, num_slots, keyframe_query.device) + schedule = _JointSchedule(flat, 1 + num_slots, heads, head_dim, num_planes, keyframe_query.element_size(), factor) + row_groups = _joint_row_groups(flat, schedule) + + out = torch.empty_like(keyframe_query) + for start in range(0, num_planes, schedule.group_axis): + stop = min(start + schedule.group_axis, num_planes) + count = stop - start + picked = slot_table[start:stop].reshape(-1) + frames = frame_keys.index_select(2, picked) + frame_vals = frame_values.index_select(2, picked) + for rows, key_rows, out_rows in row_groups: + out[:, start:stop, out_rows] = _joint_attend_group( + keyframe_query[:, start:stop, out_rows], + ( + _joint_slabs(own_keys[:, :, start:, key_rows], flat, count, rows, 1, 1), + _joint_slabs(frames[:, :, :, key_rows], flat, count, rows, num_slots, num_slots), + ), + ( + _joint_slabs(own_values[:, :, start:, key_rows], flat, count, rows, 1, 1), + _joint_slabs(frame_vals[:, :, :, key_rows], flat, count, rows, num_slots, num_slots), + ), + flat, + (count, rows), + mask, + ) + return out + + +def _joint_neighborhood_attention( + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + keyframe_query: torch.Tensor, + keyframe_key: torch.Tensor, + keyframe_value: torch.Tensor, + keyframe_times: torch.Tensor, + kernel_size: tuple[int, int, int], +) -> tuple[torch.Tensor, torch.Tensor]: + """Joint neighborhood attention over a video volume and a stack of keyframe planes, one softmax per query. + + A video query at `(t, h, w)` attends to its `Kt x Kh x Kw` video window plus the `Kh x Kw` window at the same `(h, + w)` on each of its `_KEYFRAME_CONTEXT_SLOTS` nearest planes, ranked by `|t - keyframe_times|` whatever `Kt` is. A + plane query attends to the `Kh x Kw` window on its own plane (there is no plane-to-plane attention) plus the same + window on each of its nearest video frames. + + Unlike [`LTX2VideoVaeNeighborhoodAttnProcessor`] and NATTEN, whose windows shift inward at the volume border, every + window here is centered and *clamped*: offsets that fall outside the volume are masked out. That is also why no + axis needs to be at least its kernel size. + + Args: + query, key, value: `(B, T, H, W, heads, head_dim)` video stream, query already scaled. + keyframe_query, keyframe_key, keyframe_value: `(B, P, H, W, heads, head_dim)` keyframe stream. + keyframe_times: `(P,)` plane positions, in the same temporal units and origin as the video stream's RoPE. + kernel_size: `(Kt, Kh, Kw)`. + + Returns: + `(video_out, keyframe_out)`, each shaped like its stream's query. + """ + num_frames, height, width = query.shape[1], query.shape[2], query.shape[3] + keyframe_times = keyframe_times.to(device=query.device, dtype=torch.float32) + frame_times = torch.arange(num_frames, dtype=torch.float32, device=query.device) + video_slots = _nearest_slots(frame_times, keyframe_times, _KEYFRAME_CONTEXT_SLOTS) + keyframe_slots = _nearest_slots(keyframe_times, frame_times, _KEYFRAME_CONTEXT_SLOTS) + + side = max(1, round(math.sqrt(_JOINT_BRICK_QUERIES))) + brick = (min(_JOINT_BRICK_DEPTH, num_frames), min(side, height), min(side, width)) + geometry = _JointGeometry(height, width, kernel_size, brick) + # Only CUDA keeps the score block on chip for this broadcast mask; elsewhere SDPA materializes it. + factor = _JOINT_STAGING_FACTOR_FUSED if query.device.type == "cuda" else _JOINT_STAGING_FACTOR_MATERIALIZED + video_out = _joint_video_query_pass(query, key, value, keyframe_key, keyframe_value, video_slots, geometry, factor) + keyframe_out = _joint_keyframe_query_pass( + keyframe_query, keyframe_key, keyframe_value, key, value, keyframe_slots, geometry, factor + ) + return video_out, keyframe_out + + class LTX2VideoVaeRotaryPosEmbed3D(nn.Module): """Absolute 3D rotary embedding for the diffusion decoder's neighborhood attention. @@ -134,14 +633,21 @@ def _rotate_axis(self, x: torch.Tensor, positions: torch.Tensor, inv_freqs: torc rotated = torch.stack([even * cos - odd * sin, even * sin + odd * cos], dim=-1) return rotated.reshape(x.shape).to(out_dtype) - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - """`hidden_states`: `(B, T, H, W, heads, head_dim)`.""" + def forward(self, hidden_states: torch.Tensor, positions_t: torch.Tensor | None = None) -> torch.Tensor: + """`hidden_states`: `(B, T, H, W, heads, head_dim)`. + + `positions_t` overrides the integer positions on the first axis. Keyframe planes pass their (possibly + fractional) times there, so that both streams of a keyframe decode share one temporal origin. + """ dim_t, dim_h, _ = self.rope_dim_split num_frames, height, width = hidden_states.shape[1:4] device = hidden_states.device inv_t, inv_h, inv_w = (self._inv_freqs(dim, device) for dim in self.rope_dim_split) - positions_t = torch.arange(num_frames, dtype=torch.float32, device=device) + if positions_t is None: + positions_t = torch.arange(num_frames, dtype=torch.float32, device=device) + else: + positions_t = positions_t.to(device=device, dtype=torch.float32) positions_h = torch.arange(height, dtype=torch.float32, device=device) positions_w = torch.arange(width, dtype=torch.float32, device=device) rotated_t = self._rotate_axis(hidden_states[..., :dim_t], positions_t, inv_t, axis=1) @@ -274,11 +780,14 @@ def __init__( self.rope = LTX2VideoVaeRotaryPosEmbed3D(head_dim, base=rope_base) self.set_processor(self._default_processor_cls()) - def project_qkv(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + def project_qkv( + self, hidden_states: torch.Tensor, positions_t: torch.Tensor | None = None + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Q/K/V as `(B, T, H, W, heads, head_dim)`, RMS-normed, query pre-scaled, then rotated. The query carries the `1 / sqrt(head_dim)` factor here so both processors can ask their attention backend for - `scale=1.0` — this is the order the reference uses (norm, scale, then rotate). + `scale=1.0` — this is the order the reference uses (norm, scale, then rotate). `positions_t` overrides the + temporal RoPE positions, see [`LTX2VideoVaeRotaryPosEmbed3D`]. """ batch_size, num_frames, height, width, _ = hidden_states.shape shape = (batch_size, num_frames, height, width, self.heads, self.head_dim) @@ -289,7 +798,7 @@ def project_qkv(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch. query = self.norm_q(query) key = self.norm_k(key) query = query * self.scale - return self.rope(query), self.rope(key), value + return self.rope(query, positions_t), self.rope(key, positions_t), value def build_block_mask(self, hidden_states: torch.Tensor): """The flex `BlockMask` for this grid, or `None` if the processor does not read one. @@ -315,6 +824,35 @@ def forward(self, hidden_states: torch.Tensor, block_mask=None) -> torch.Tensor: ) return self.processor(self, hidden_states, block_mask) + def forward_with_keyframes( + self, hidden_states: torch.Tensor, keyframe_hidden_states: torch.Tensor, keyframe_times: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + """Joint attention over the video `(B, T, H, W, C)` and the keyframe planes `(B, P, H, W, C)`. + + This bypasses the attention processor: neither FlexAttention's neighborhood mask nor NATTEN expresses the joint + window, so a keyframe decode always runs [`_joint_neighborhood_attention`], whatever processor is set. + `keyframe_times` are the planes' `(P,)` temporal RoPE positions, in the video stream's tile-local origin. + """ + batch_size, num_frames, height, width, channels = hidden_states.shape + num_planes = keyframe_hidden_states.shape[1] + query, key, value = self.project_qkv(hidden_states) + keyframe_query, keyframe_key, keyframe_value = self.project_qkv(keyframe_hidden_states, keyframe_times) + hidden_states, keyframe_hidden_states = _joint_neighborhood_attention( + query.contiguous(), + key.contiguous(), + value.contiguous(), + keyframe_query.contiguous(), + keyframe_key.contiguous(), + keyframe_value.contiguous(), + keyframe_times, + self.kernel_size, + ) + hidden_states = self.to_out[0](hidden_states.reshape(batch_size, num_frames, height, width, channels)) + keyframe_hidden_states = self.to_out[0]( + keyframe_hidden_states.reshape(batch_size, num_planes, height, width, channels) + ) + return hidden_states, keyframe_hidden_states + # Tokens per tile in `LTX2VideoVaeSwiGLU`, matching the reference decoder's own default. `w_gate(x)` and # `w_up(x)` are both hidden-width and their product makes a third, so evaluating a whole video at once @@ -377,6 +915,19 @@ def forward(self, hidden_states: torch.Tensor, block_mask=None) -> torch.Tensor: hidden_states = hidden_states + self.mlp(self.norm2(hidden_states)) return hidden_states + def forward_with_keyframes( + self, hidden_states: torch.Tensor, keyframe_hidden_states: torch.Tensor, keyframe_times: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + """Both streams through the same weights; they only meet inside the joint attention softmax.""" + attn_output, keyframe_attn_output = self.attn.forward_with_keyframes( + self.norm1(hidden_states), self.norm1(keyframe_hidden_states), keyframe_times + ) + hidden_states = hidden_states + attn_output + keyframe_hidden_states = keyframe_hidden_states + keyframe_attn_output + hidden_states = hidden_states + self.mlp(self.norm2(hidden_states)) + keyframe_hidden_states = keyframe_hidden_states + self.mlp(self.norm2(keyframe_hidden_states)) + return hidden_states, keyframe_hidden_states + class LTX2VideoVaeAdaLNZero(nn.Module): """Shared AdaLN-Zero modulation: a timestep embedding to seven `(B, 1, 1, 1, C)` chunks. @@ -439,6 +990,35 @@ def forward( hidden_states = hidden_states + self.mlp(self.norm2(hidden_states) * (1 + scale_mlp) + shift_mlp) return hidden_states + def forward_with_keyframes( + self, + hidden_states: torch.Tensor, + latent_context: torch.Tensor, + keyframe_hidden_states: torch.Tensor, + keyframe_latent_context: torch.Tensor, + modulation: tuple[torch.Tensor, ...], + keyframe_times: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Both streams through the same weights and the same modulation, each with its own context.""" + scale_msa, shift_msa, _, scale_mlp, shift_mlp, _, _ = [ + modulation[i] + self.scale_shift_table[i].view(1, 1, 1, 1, -1) for i in range(self.num_mod_params) + ] + + hidden_states = hidden_states + self.context_proj(latent_context) + keyframe_hidden_states = keyframe_hidden_states + self.context_proj(keyframe_latent_context) + attn_output, keyframe_attn_output = self.attn.forward_with_keyframes( + self.norm1(hidden_states) * (1 + scale_msa) + shift_msa, + self.norm1(keyframe_hidden_states) * (1 + scale_msa) + shift_msa, + keyframe_times, + ) + hidden_states = hidden_states + attn_output + keyframe_hidden_states = keyframe_hidden_states + keyframe_attn_output + hidden_states = hidden_states + self.mlp(self.norm2(hidden_states) * (1 + scale_mlp) + shift_mlp) + keyframe_hidden_states = keyframe_hidden_states + self.mlp( + self.norm2(keyframe_hidden_states) * (1 + scale_mlp) + shift_mlp + ) + return hidden_states, keyframe_hidden_states + class LTX2VideoVaePixelShuffleUpsampler(nn.Module): """Linear channel expansion followed by a channels-last pixel shuffle. @@ -499,6 +1079,7 @@ def __init__( timestep_scale_multiplier: float = 1000.0, model_output_type: str = "x0", default_num_inference_steps: int = 1, + keyframe_type_embedding: bool = False, ): super().__init__() if model_output_type not in ("x0", "v"): @@ -527,6 +1108,15 @@ def __init__( self.trailing_pad_latent_frames = (stage_kernels[0][0] // 2) * 2 self.conv_in = nn.Linear(in_channels, stage_channels[0], bias=True) + # The learned tag of the keyframe stream, added to the keyframe latents right before the shared `conv_in`. It + # is the only keyframe-specific weight, so a checkpoint trained without keyframes simply does not have one. + self.type_emb = nn.Parameter(torch.zeros(in_channels)) if keyframe_type_embedding else None + # Temporal upsampling still to come at each stage input, plus 1 for the diffusion stage: the divisor of + # `_keyframe_stage_times`. (8, 8, 4, 2, 1) for the production strides. + self.keyframe_time_strides = tuple( + math.prod(stride[0] for stride in upsample_strides[stage_idx:]) + for stage_idx in range(len(upsample_strides) + 1) + ) self.det_stages = nn.ModuleList() self.upsamples = nn.ModuleList() @@ -662,13 +1252,197 @@ def denoise(self, latent_context: torch.Tensor, x_t: torch.Tensor, num_inference x_t = (x_t_fp32 - dt * model_out).to(x_t.dtype) return x_t + def forward_stages_1_to_3_with_keyframes( + self, hidden_states: torch.Tensor, keyframe_hidden_states: torch.Tensor, keyframe_frame_indices: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + """Keyframe counterpart of [`forward_stages_1_to_3`], carrying both streams. + + `keyframe_hidden_states` are the denormalized keyframe latents `(B, C, P, H, W)`, one latent frame per plane, + on the same spatial grid as `hidden_states`. They are tagged with `type_emb`, then share `conv_in` and every + stage with the video. The trailing ghost frames are a temporal border workaround of the video stream only. The + keyframe stream comes back channels-last, `(B, P, H_4, W_4, C_4)`. + """ + num_pad = self.trailing_pad_latent_frames + if num_pad > 0: + trailing = hidden_states[:, :, -1:].expand(-1, -1, num_pad, -1, -1) + hidden_states = torch.cat([hidden_states, trailing], dim=2) + + hidden_states = self.conv_in(hidden_states.permute(0, 2, 3, 4, 1)) + keyframe_hidden_states = keyframe_hidden_states.permute(0, 2, 3, 4, 1) + if self.type_emb is not None: + keyframe_hidden_states = keyframe_hidden_states + self.type_emb.view(1, 1, 1, 1, -1) + keyframe_hidden_states = self.conv_in(keyframe_hidden_states) + for stage_idx, (blocks, upsample) in enumerate(zip(self.det_stages[:-1], self.upsamples[:-1])): + keyframe_times = _keyframe_stage_times(keyframe_frame_indices, self.keyframe_time_strides[stage_idx]) + for block in blocks: + hidden_states, keyframe_hidden_states = block.forward_with_keyframes( + hidden_states, keyframe_hidden_states, keyframe_times + ) + hidden_states = upsample(hidden_states) + keyframe_hidden_states = _upsample_keyframe_planes(upsample, keyframe_hidden_states) + return hidden_states, keyframe_hidden_states + + def forward_stage_4_with_keyframes( + self, + hidden_states: torch.Tensor, + keyframe_hidden_states: torch.Tensor, + keyframe_frame_indices: torch.Tensor, + drop_leading_frame: bool = True, + crop_trailing_ghost: bool = True, + stage_4_time_origin: float = 0.0, + pixel_time_origin: float = 0.0, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Keyframe counterpart of [`forward_stage_4`]. Returns the video context, the keyframe context, and the + planes' times in the diffusion stage. + + The video stream's RoPE positions are 0-based within a tile at every stage, so the plane times are rebased on + the tile's first frame. That takes two origins at two scales, both taken from the tile and not derived from one + another (the causal first frame makes `pixel_time_origin == stride * stage_4_time_origin` an off-by-one trap): + `stage_4_time_origin` is the tile's first cell entering this stage, `pixel_time_origin` its first pixel frame. + Both are 0 for an untiled decode. `keyframe_frame_indices` stay global pixel frames. + """ + keyframe_times = ( + _keyframe_stage_times(keyframe_frame_indices, self.keyframe_time_strides[-2]) - stage_4_time_origin + ) + for block in self.det_stages[-1]: + hidden_states, keyframe_hidden_states = block.forward_with_keyframes( + hidden_states, keyframe_hidden_states, keyframe_times + ) + hidden_states = self.upsamples[-1](hidden_states, drop_leading_frame=drop_leading_frame) + keyframe_hidden_states = _upsample_keyframe_planes(self.upsamples[-1], keyframe_hidden_states) + + num_pad = self.trailing_pad_latent_frames + if crop_trailing_ghost and num_pad > 0: + hidden_states = hidden_states[:, : -num_pad * self.temporal_compression_ratio] + keyframe_times = ( + _keyframe_stage_times(keyframe_frame_indices, self.keyframe_time_strides[-1]) - pixel_time_origin + ) + return hidden_states, keyframe_hidden_states, keyframe_times + + def forward_diffusion_step_with_keyframes( + self, + latent_context: torch.Tensor, + x_t: torch.Tensor, + keyframe_latent_context: torch.Tensor, + keyframe_x_t: torch.Tensor, + timestep: torch.Tensor, + keyframe_times: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + """One stage-5 step over both streams. Returns `(video_prediction, keyframe_prediction)` in pixel space. + + The keyframe stream is a second pixel diffusion stream, one pixel frame per plane: its own noised pixels + through the shared `conv_in_x_t`, its own context, the same modulation. It is stepped along with the video so + the hidden states the joint attention reads stay at the noise level the decoder was trained on, and is then + discarded. + """ + t_emb = self.t_embedder( + self.timestep_scale_multiplier * timestep, + resolution=None, + aspect_ratio=None, + batch_size=timestep.shape[0], + hidden_dtype=latent_context.dtype, + ) + modulation = self.shared_adaln(t_emb) + + hidden_states = self.conv_in_x_t(_patchify(x_t, self.patch_size).permute(0, 2, 3, 4, 1)) + keyframe_hidden_states = self.conv_in_x_t(_patchify(keyframe_x_t, self.patch_size).permute(0, 2, 3, 4, 1)) + for block in self.diff_blocks: + hidden_states, keyframe_hidden_states = block.forward_with_keyframes( + hidden_states, + latent_context, + keyframe_hidden_states, + keyframe_latent_context, + modulation, + keyframe_times, + ) + + outputs = [] + for states in (hidden_states, keyframe_hidden_states): + states = self.conv_out(self.norm_out(states)) + outputs.append(_unpatchify(states.permute(0, 4, 1, 2, 3).contiguous(), self.patch_size)) + return outputs[0], outputs[1] + + def denoise_with_keyframes( + self, + latent_context: torch.Tensor, + x_t: torch.Tensor, + keyframe_latent_context: torch.Tensor, + keyframe_x_t: torch.Tensor, + keyframe_times: torch.Tensor, + num_inference_steps: int, + ) -> torch.Tensor: + """Keyframe counterpart of [`denoise`]: both streams through the same Euler loop, only the video returned.""" + batch_size = latent_context.shape[0] + timesteps = torch.linspace( + 1.0, 1.0 / num_inference_steps, num_inference_steps, device=latent_context.device, dtype=torch.float32 + ) + + if num_inference_steps == 1 and self.model_output_type == "x0": + return self.forward_diffusion_step_with_keyframes( + latent_context, + x_t, + keyframe_latent_context, + keyframe_x_t, + timesteps[:1].expand(batch_size), + keyframe_times, + )[0] + + for step_idx in range(num_inference_steps): + t_now = timesteps[step_idx].expand(batch_size) + t_next = timesteps[step_idx + 1] if step_idx + 1 < num_inference_steps else torch.zeros_like(t_now) + model_outputs = self.forward_diffusion_step_with_keyframes( + latent_context, x_t, keyframe_latent_context, keyframe_x_t, t_now, keyframe_times + ) + sigma = t_now.view(-1, *([1] * (x_t.ndim - 1))) + dt = (t_now - t_next).view(-1, *([1] * (x_t.ndim - 1))) + updated = [] + for sample, model_out in zip((x_t, keyframe_x_t), model_outputs): + sample_fp32, model_out = sample.float(), model_out.float() + if self.model_output_type == "x0": + model_out = (sample_fp32 - model_out) / sigma + updated.append((sample_fp32 - dt * model_out).to(sample.dtype)) + x_t, keyframe_x_t = updated + return x_t + + def _decode_with_keyframes( + self, + hidden_states: torch.Tensor, + keyframe_hidden_states: torch.Tensor, + keyframe_frame_indices: torch.Tensor, + generator: torch.Generator | None, + num_inference_steps: int, + ) -> torch.Tensor: + features, keyframe_features = self.forward_stages_1_to_3_with_keyframes( + hidden_states, keyframe_hidden_states, keyframe_frame_indices + ) + latent_context, keyframe_latent_context, keyframe_times = self.forward_stage_4_with_keyframes( + features, keyframe_features, keyframe_frame_indices + ) + batch_size, num_frames, height, width = latent_context.shape[:4] + pixel_shape = (batch_size, self.out_channels, num_frames, height * self.patch_size, width * self.patch_size) + x_t = randn_tensor(pixel_shape, generator=generator, device=hidden_states.device, dtype=hidden_states.dtype) + # The keyframe stream draws its own noise, after the video's: its planes are not part of the video canvas. + keyframe_shape = (*pixel_shape[:2], keyframe_latent_context.shape[1], *pixel_shape[3:]) + keyframe_x_t = randn_tensor( + keyframe_shape, generator=generator, device=hidden_states.device, dtype=hidden_states.dtype + ) + return self.denoise_with_keyframes( + latent_context, x_t, keyframe_latent_context, keyframe_x_t, keyframe_times, num_inference_steps + ) + def forward( self, hidden_states: torch.Tensor, generator: torch.Generator | None = None, num_inference_steps: int | None = None, + keyframe_hidden_states: torch.Tensor | None = None, + keyframe_frame_indices: torch.Tensor | None = None, ) -> torch.Tensor: num_inference_steps = num_inference_steps or self.default_num_inference_steps + if keyframe_hidden_states is not None: + return self._decode_with_keyframes( + hidden_states, keyframe_hidden_states, keyframe_frame_indices, generator, num_inference_steps + ) latent_context = self.forward_stage_4(self.forward_stages_1_to_3(hidden_states)) # The context grid is the stage-5 token grid, so the pixel canvas is its shape times the patch size — # temporally that is the causal (T - 1) * ratio + 1 mapping of the LTX-2 latent space. @@ -712,6 +1486,12 @@ class LTX2VideoDiffusionDecoderModel(ModelMixin, AttentionMixin, ConfigMixin): The latent statistics are carried here as buffers so the decode pipeline can denormalize without loading a second autoencoder just for two vectors. + Decoding can be anchored on *keyframe planes*: single-frame latents at known pixel frames, each encoded as a + standalone one-frame clip, passed as `keyframe_latents` / `keyframe_frame_indices` to [`decode`]. Every video + position then also attends to the same spatial window on its two nearest planes. Checkpoints trained for it carry a + learned tag added to the plane latents, `decoder.type_emb`, which `decoder_keyframe_type_embedding=True` creates; + without it the planes are decoded untagged. + This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented for all models (such as downloading or saving). """ @@ -741,6 +1521,7 @@ def __init__( decoder_num_inference_steps: int = 1, spatial_compression_ratio: int = 32, temporal_compression_ratio: int = 8, + decoder_keyframe_type_embedding: bool = False, ) -> None: super().__init__() @@ -760,6 +1541,7 @@ def __init__( timestep_scale_multiplier=decoder_timestep_scale_multiplier, model_output_type=decoder_model_output_type, default_num_inference_steps=decoder_num_inference_steps, + keyframe_type_embedding=decoder_keyframe_type_embedding, ) self.spatial_compression_ratio = spatial_compression_ratio @@ -857,6 +1639,8 @@ def tiled_decode( z: torch.Tensor, generator: torch.Generator | None = None, num_inference_steps: int | None = None, + keyframe_latents: torch.Tensor | None = None, + keyframe_frame_indices: torch.Tensor | None = None, ) -> torch.Tensor: r"""Decode a batch of latents with the last deterministic stage and the diffusion stage running per tile. @@ -865,7 +1649,13 @@ def tiled_decode( px spatially and 2 frames temporally for the production config). Temporal tiles follow the causal frame mapping: the tile containing t=0 drops the temporal upsample's duplicate leading frame and only the tile containing the video end carries the NATTEN border padding. + + With keyframe planes, each temporal tile carries the planes inside its pixel-frame span plus the nearest one on + each side of it, with their times rebased on the tile's first cell and first pixel frame, and each spatial tile + crops the planes with the same slices as the video. """ + if keyframe_latents is not None: + keyframe_frame_indices = self._check_keyframes(z, keyframe_latents, keyframe_frame_indices) decoder = self.decoder num_inference_steps = num_inference_steps or decoder.default_num_inference_steps batch_size = z.shape[0] @@ -890,7 +1680,12 @@ def tiled_decode( ) ] - features = decoder.forward_stages_1_to_3(z) + if keyframe_latents is None: + features = decoder.forward_stages_1_to_3(z) + else: + features, keyframe_features = decoder.forward_stages_1_to_3_with_keyframes( + z, keyframe_latents, keyframe_frame_indices + ) # The trailing ghost frames replicate through the earlier stages' temporal upsamples, whose composed # mapping is affine with slope equal to the product of their strides. ghost_frames = decoder.trailing_pad_latent_frames * math.prod(up.stride[0] for up in decoder.upsamples[:-1]) @@ -923,15 +1718,34 @@ def tiled_decode( is_trailing = t1 == num_frames # The tile containing the video end takes the ghost frames with it into stage 4. feature_t1 = features.shape[1] if is_trailing else t1 + # A non-origin tile keeps the duplicate leading frame, placing its first cell one pixel frame earlier + # than `t0 * scale_t` — the causal 1-then-`scale_t` frame mapping. + pixel_t0 = t0 * scale_t - (1 if not is_origin and scale_t == 2 else 0) + if keyframe_latents is not None: + tile_pixel_frames = (t1 - t0) * scale_t - (1 if is_origin and scale_t == 2 else 0) + keep = _keyframe_planes_for_tile(keyframe_frame_indices, pixel_t0, pixel_t0 + tile_pixel_frames - 1) + tile_keyframe_features = keyframe_features[:, keep.to(keyframe_features.device)] + tile_keyframe_indices = keyframe_frame_indices[keep] rows = [] for h0, h1 in height_tiles: row = [] for w0, w1 in width_tiles: - context = decoder.forward_stage_4( - features[:, t0:feature_t1, h0:h1, w0:w1], - drop_leading_frame=is_origin, - crop_trailing_ghost=is_trailing, - ) + if keyframe_latents is None: + context = decoder.forward_stage_4( + features[:, t0:feature_t1, h0:h1, w0:w1], + drop_leading_frame=is_origin, + crop_trailing_ghost=is_trailing, + ) + else: + context, keyframe_context, keyframe_times = decoder.forward_stage_4_with_keyframes( + features[:, t0:feature_t1, h0:h1, w0:w1], + tile_keyframe_features[:, :, h0:h1, w0:w1], + tile_keyframe_indices, + drop_leading_frame=is_origin, + crop_trailing_ghost=is_trailing, + stage_4_time_origin=float(t0), + pixel_time_origin=float(pixel_t0), + ) tile_pixel_shape = ( batch_size, decoder.out_channels, @@ -942,9 +1756,6 @@ def tiled_decode( if single_step_x0: x_t = randn_tensor(tile_pixel_shape, generator=generator, device=z.device, dtype=z.dtype) else: - # A non-origin tile keeps the duplicate leading frame, placing its first cell one pixel - # frame earlier than `t0 * scale_t` — the causal 1-then-`scale_t` frame mapping. - pixel_t0 = t0 * scale_t - (1 if not is_origin and scale_t == 2 else 0) x_t = x_t_full[ :, :, @@ -952,7 +1763,21 @@ def tiled_decode( h0 * scale_h : h0 * scale_h + tile_pixel_shape[3], w0 * scale_w : w0 * scale_w + tile_pixel_shape[4], ] - row.append(decoder.denoise(context, x_t, num_inference_steps)) + if keyframe_latents is None: + row.append(decoder.denoise(context, x_t, num_inference_steps)) + continue + # The keyframe stream always draws its own noise, after the video's. + keyframe_x_t = randn_tensor( + (*tile_pixel_shape[:2], keyframe_context.shape[1], *tile_pixel_shape[3:]), + generator=generator, + device=z.device, + dtype=z.dtype, + ) + row.append( + decoder.denoise_with_keyframes( + context, x_t, keyframe_context, keyframe_x_t, keyframe_times, num_inference_steps + ) + ) rows.append(row) result_rows = [] @@ -985,6 +1810,38 @@ def tiled_decode( result.append(group) return torch.cat(result, dim=2) + def _check_keyframes( + self, z: torch.Tensor, keyframe_latents: torch.Tensor, keyframe_frame_indices: torch.Tensor | None + ) -> torch.Tensor: + """Validate a keyframe input and return its frame indices as a 1-D int64 tensor.""" + if keyframe_frame_indices is None: + raise ValueError("`keyframe_frame_indices` is required when `keyframe_latents` is passed.") + keyframe_frame_indices = torch.as_tensor(keyframe_frame_indices, dtype=torch.int64).cpu() + if keyframe_latents.ndim != 5: + raise ValueError(f"`keyframe_latents` must be (B, C, P, H, W), got {tuple(keyframe_latents.shape)}.") + num_planes = keyframe_latents.shape[2] + if num_planes == 0: + raise ValueError("`keyframe_latents` needs at least one plane; omit it for a plain decode.") + if keyframe_latents.shape[:2] != z.shape[:2] or keyframe_latents.shape[3:] != z.shape[3:]: + raise ValueError( + f"`keyframe_latents` {tuple(keyframe_latents.shape)} must match the batch size, channels, height and " + f"width of `z` {tuple(z.shape)}." + ) + if keyframe_frame_indices.shape != (num_planes,): + raise ValueError( + f"`keyframe_frame_indices` must be 1-D with one pixel frame per plane ({num_planes}), got shape " + f"{tuple(keyframe_frame_indices.shape)}." + ) + if int(keyframe_frame_indices.min()) < 0: + raise ValueError("`keyframe_frame_indices` must be non-negative pixel frame indices.") + if self.decoder.type_emb is None: + logger.warning( + "Decoding with keyframe planes on a checkpoint without a keyframe tag (`decoder.type_emb`): the planes " + "enter the decoder untagged. Checkpoints trained for keyframe decoding are converted with " + "`decoder_keyframe_type_embedding=True`." + ) + return keyframe_frame_indices + @apply_forward_hook def decode( self, @@ -992,12 +1849,20 @@ def decode( generator: torch.Generator | None = None, num_inference_steps: int | None = None, return_dict: bool = True, + keyframe_latents: torch.Tensor | None = None, + keyframe_frame_indices: torch.Tensor | None = None, ) -> DecoderOutput | torch.Tensor: """Decode a batch of latents. `z` is expected to be denormalized already (the pipeline applies `latents_mean` / `latents_std`), matching [`AutoencoderKLLTX2Video`]. This decoder denoises, so pass `generator` for reproducibility. + + `keyframe_latents` `(B, C, P, H, W)`, denormalized like `z`, anchors the decode on `P` keyframe planes at the + pixel frames `keyframe_frame_indices` `(P,)`. Each plane must be the latent of a standalone one-frame clip. + Planes may lie outside the decoded clip: they are ranked by temporal distance. """ + if keyframe_latents is not None: + keyframe_frame_indices = self._check_keyframes(z, keyframe_latents, keyframe_frame_indices) tile_latent_min_height = self.tile_sample_min_height // self.spatial_compression_ratio tile_latent_min_width = self.tile_sample_min_width // self.spatial_compression_ratio tile_latent_min_num_frames = self.tile_sample_min_num_frames // self.temporal_compression_ratio @@ -1006,9 +1871,21 @@ def decode( or z.shape[3] > tile_latent_min_height or z.shape[4] > tile_latent_min_width ): - decoded = self.tiled_decode(z, generator=generator, num_inference_steps=num_inference_steps) + decoded = self.tiled_decode( + z, + generator=generator, + num_inference_steps=num_inference_steps, + keyframe_latents=keyframe_latents, + keyframe_frame_indices=keyframe_frame_indices, + ) else: - decoded = self.decoder(z, generator=generator, num_inference_steps=num_inference_steps) + decoded = self.decoder( + z, + generator=generator, + num_inference_steps=num_inference_steps, + keyframe_hidden_states=keyframe_latents, + keyframe_frame_indices=keyframe_frame_indices, + ) if not return_dict: return (decoded,) @@ -1020,6 +1897,8 @@ def forward( generator: torch.Generator | None = None, num_inference_steps: int | None = None, return_dict: bool = True, + keyframe_latents: torch.Tensor | None = None, + keyframe_frame_indices: torch.Tensor | None = None, ) -> DecoderOutput | tuple[torch.Tensor]: r""" Args: @@ -1032,8 +1911,21 @@ def forward( Number of denoising steps. Defaults to the decoder's `decoder_num_inference_steps` config value. return_dict (`bool`, *optional*, defaults to `True`): Whether to return a [`~models.autoencoders.vae.DecoderOutput`] instead of a plain tuple. + keyframe_latents (`torch.Tensor`, *optional*): + Keyframe planes of shape `(B, C, P, H, W)`, denormalized like `z` and on its latent grid, one latent + frame per plane, each encoded as a standalone one-frame clip. Every video position also attends to the + same spatial window on its two nearest planes. + keyframe_frame_indices (`torch.Tensor`, *optional*): + The `(P,)` pixel frame of each plane in the decoded video. Required with `keyframe_latents`. Returns: [`~models.autoencoders.vae.DecoderOutput`] or `tuple` """ - return self.decode(z, generator=generator, num_inference_steps=num_inference_steps, return_dict=return_dict) + return self.decode( + z, + generator=generator, + num_inference_steps=num_inference_steps, + return_dict=return_dict, + keyframe_latents=keyframe_latents, + keyframe_frame_indices=keyframe_frame_indices, + ) diff --git a/src/diffusers/pipelines/ltx2/dfr_layout.py b/src/diffusers/pipelines/ltx2/dfr_layout.py index d1b0384c7bd3..9fd321bda52c 100644 --- a/src/diffusers/pipelines/ltx2/dfr_layout.py +++ b/src/diffusers/pipelines/ltx2/dfr_layout.py @@ -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: diff --git a/tests/models/autoencoders/test_models_ltx2_diffusion_decoder.py b/tests/models/autoencoders/test_models_ltx2_diffusion_decoder.py index df4b9536af4a..7d56ed84a1da 100644 --- a/tests/models/autoencoders/test_models_ltx2_diffusion_decoder.py +++ b/tests/models/autoencoders/test_models_ltx2_diffusion_decoder.py @@ -240,3 +240,271 @@ def test_natten_processor_decodes(self): assert output.shape == (inputs["z"].shape[0], *self.output_shape) assert torch.isfinite(output).all(), "NATTEN decode produced NaN/inf values" + + +def _dense_joint_attention(query, key, value, keyframe_query, keyframe_key, keyframe_value, keyframe_times, kernel): + """Joint attention written out as one dense masked softmax over every video and plane token. + + An independent statement of the visibility rule `_joint_neighborhood_attention` implements with query bricks: a + centered window clamped (not shifted) at the volume border, the same spatial window on the two planes nearest a + frame, and the same spatial window on the two frames nearest a plane. + """ + batch_size, num_frames, height, width, heads, head_dim = query.shape + num_planes = keyframe_query.shape[1] + lo = [k // 2 for k in kernel] + hi = [k - l - 1 for k, l in zip(kernel, lo)] + frame_times = torch.arange(num_frames, dtype=torch.float32) + video_slots = torch.argsort((frame_times[:, None] - keyframe_times[None]).abs(), dim=-1, stable=True)[:, :2] + plane_slots = torch.argsort((keyframe_times[:, None] - frame_times[None]).abs(), dim=-1, stable=True)[:, :2] + + def grid(length): + t, h, w = torch.meshgrid(torch.arange(length), torch.arange(height), torch.arange(width), indexing="ij") + return t.reshape(-1), h.reshape(-1), w.reshape(-1) + + vt, vh, vw = grid(num_frames) + pp, ph, pw = grid(num_planes) + is_plane = torch.cat([torch.zeros_like(vt, dtype=torch.bool), torch.ones_like(pp, dtype=torch.bool)]) + axis0, hh, ww = torch.cat([vt, pp]), torch.cat([vh, ph]), torch.cat([vw, pw]) + dh, dw = hh[None] - hh[:, None], ww[None] - ww[:, None] + spatial = (dh >= -lo[1]) & (dh <= hi[1]) & (dw >= -lo[2]) & (dw <= hi[2]) + dt = axis0[None] - axis0[:, None] + video_video = ~is_plane[:, None] & ~is_plane[None] & (dt >= -lo[0]) & (dt <= hi[0]) + query_slots = torch.where(is_plane, 0, axis0.clamp(max=num_frames - 1)) + video_plane = ( + ~is_plane[:, None] & is_plane[None] & (video_slots[query_slots][:, None, :] == axis0[None, :, None]).any(-1) + ) + plane_rows = torch.where(is_plane, axis0.clamp(max=num_planes - 1), 0) + plane_video = ( + is_plane[:, None] & ~is_plane[None] & (plane_slots[plane_rows][:, None, :] == axis0[None, :, None]).any(-1) + ) + plane_plane = is_plane[:, None] & is_plane[None] & (dt == 0) + visible = spatial & (video_video | video_plane | plane_video | plane_plane) + + def tokens(video, planes): + flat = torch.cat( + [video.reshape(batch_size, -1, heads, head_dim), planes.reshape(batch_size, -1, heads, head_dim)], 1 + ) + return flat.transpose(1, 2) + + scores = tokens(query, keyframe_query) @ tokens(key, keyframe_key).transpose(-1, -2) + attended = scores.masked_fill(~visible, float("-inf")).softmax(-1) @ tokens(value, keyframe_value) + attended = attended.transpose(1, 2) + num_video = num_frames * height * width + return ( + attended[:, :num_video].reshape(query.shape), + attended[:, num_video:].reshape(keyframe_query.shape), + ) + + +class TestLTX2VideoDiffusionDecoderModelKeyframes(LTX2VideoDiffusionDecoderModelTesterConfig): + """Keyframe-aware decoding: a second stream of single-frame planes, attended jointly with the video.""" + + def get_model(self, keyframe_type_embedding=True): + torch.manual_seed(0) + model = self.model_class(**self.get_init_dict(), decoder_keyframe_type_embedding=keyframe_type_embedding) + with torch.no_grad(): + # Random, not default-initialized: the zero `type_emb` and `scale_shift_table` would hide wiring errors. + for parameter in model.parameters(): + parameter.copy_(torch.randn(parameter.shape, generator=self.generator) * 0.3) + return model.to(torch_device).eval() + + def get_latents(self, num_planes=2, latent_frames=3): + generator = self.generator + latents = randn_tensor((1, 8, latent_frames, 3, 4), generator=generator, device=torch_device) + keyframes = randn_tensor((1, 8, num_planes, 3, 4), generator=generator, device=torch_device) + return latents, keyframes + + def decode(self, model, latents, **kwargs): + generator = torch.Generator("cpu").manual_seed(0) + with torch.no_grad(): + return model.decode(latents, generator=generator, return_dict=False, **kwargs)[0] + + def test_type_emb_is_opt_in(self): + default = self.model_class(**self.get_init_dict()) + assert default.decoder.type_emb is None + assert "decoder.type_emb" not in default.state_dict() + + model = self.model_class(**self.get_init_dict(), decoder_keyframe_type_embedding=True) + assert model.state_dict()["decoder.type_emb"].shape == (self.get_init_dict()["latent_channels"],) + assert torch.equal(model.decoder.type_emb, torch.zeros_like(model.decoder.type_emb)) + + def test_type_emb_save_load_round_trip(self, tmp_path): + model = self.get_model() + model.save_pretrained(tmp_path / "keyframes") + loaded, info = self.model_class.from_pretrained(tmp_path / "keyframes", output_loading_info=True) + assert loaded.config.decoder_keyframe_type_embedding + assert not info["missing_keys"] and not info["unexpected_keys"] + assert torch.equal(loaded.decoder.type_emb.to(torch_device), model.decoder.type_emb) + + # A checkpoint without the tag keeps loading cleanly with the default config. + self.get_model(keyframe_type_embedding=False).save_pretrained(tmp_path / "plain") + loaded, info = self.model_class.from_pretrained(tmp_path / "plain", output_loading_info=True) + assert loaded.decoder.type_emb is None + assert not info["missing_keys"] and not info["unexpected_keys"] + + def test_plain_decode_is_unchanged(self): + """Without planes the decode must not depend on whether the model carries a tag, nor on the new arguments.""" + latents, _ = self.get_latents() + model = self.get_model() + plain = self.decode(model, latents) + assert torch.equal(self.decode(model, latents, keyframe_latents=None, keyframe_frame_indices=None), plain) + + untagged = self.model_class(**self.get_init_dict()).to(torch_device).eval() + state_dict = {k: v for k, v in model.state_dict().items() if k != "decoder.type_emb"} + untagged.load_state_dict(state_dict, strict=True) + assert torch.equal(self.decode(untagged, latents), plain) + + def test_keyframe_decode_depends_on_planes_and_tag(self): + latents, keyframes = self.get_latents() + model = self.get_model() + indices = torch.tensor([8, 16]) + plain = self.decode(model, latents) + with_planes = self.decode(model, latents, keyframe_latents=keyframes, keyframe_frame_indices=indices) + assert with_planes.shape == plain.shape + assert torch.isfinite(with_planes).all() + assert not torch.allclose(with_planes, plain) + assert torch.equal( + self.decode(model, latents, keyframe_latents=keyframes, keyframe_frame_indices=indices), with_planes + ) + + other = self.decode(model, latents, keyframe_latents=keyframes.flip(2), keyframe_frame_indices=indices) + assert not torch.allclose(other, with_planes) + + with torch.no_grad(): + model.decoder.type_emb.add_(1.0) + retagged = self.decode(model, latents, keyframe_latents=keyframes, keyframe_frame_indices=indices) + assert not torch.allclose(retagged, with_planes) + + def test_planes_outside_the_two_nearest_are_invisible(self): + """Every video position attends to its two nearest planes only, so a third, farther plane cannot matter.""" + latents, keyframes = self.get_latents(num_planes=3) + model = self.get_model() + indices = torch.tensor([8, 16, 400]) + reference = self.decode(model, latents, keyframe_latents=keyframes, keyframe_frame_indices=indices) + + far = keyframes.clone() + far[:, :, 2] = torch.randn_like(far[:, :, 2]) + assert torch.equal( + self.decode(model, latents, keyframe_latents=far, keyframe_frame_indices=indices), reference + ) + + near = keyframes.clone() + near[:, :, 0] = torch.randn_like(near[:, :, 0]) + assert not torch.allclose( + self.decode(model, latents, keyframe_latents=near, keyframe_frame_indices=indices), reference + ) + + @pytest.mark.parametrize( + "num_frames, num_planes, kernel", + [(5, 2, (3, 3, 3)), (3, 1, (3, 5, 5)), (2, 3, (5, 3, 7))], + ) + def test_joint_attention_matches_dense_softmax(self, num_frames, num_planes, kernel): + generator = torch.Generator().manual_seed(0) + shape = (2, num_frames, 6, 7, 2, 16) + query, key, value = (torch.randn(shape, generator=generator) for _ in range(3)) + plane_shape = (2, num_planes, 6, 7, 2, 16) + keyframe_query, keyframe_key, keyframe_value = ( + torch.randn(plane_shape, generator=generator) for _ in range(3) + ) + keyframe_times = torch.rand(num_planes, generator=generator) * (num_frames + 2) - 1 + + ours = ltx2_diffusion_decoder._joint_neighborhood_attention( + query, key, value, keyframe_query, keyframe_key, keyframe_value, keyframe_times, kernel + ) + dense = _dense_joint_attention( + query, key, value, keyframe_query, keyframe_key, keyframe_value, keyframe_times, kernel + ) + for got, expected in zip(ours, dense): + assert torch.allclose(got, expected, atol=1e-5, rtol=1e-5), (got - expected).abs().max() + + def test_keyframe_geometry(self): + indices = torch.tensor([0, 8, 16]) + # t(0) = 0 and t(f) = (f + (r - 1) / 2) / r: the center of the cell holding frame f. + assert ltx2_diffusion_decoder._keyframe_stage_times(indices, 8).tolist() == [0.0, 1.4375, 2.4375] + assert ltx2_diffusion_decoder._keyframe_stage_times(indices, 2).tolist() == [0.0, 4.25, 8.25] + assert ltx2_diffusion_decoder._keyframe_stage_times(indices, 1).tolist() == [0.0, 8.0, 16.0] + + # Nearest two by |dt|, ties to the lower index, `-1` when there are fewer candidates. + slots = ltx2_diffusion_decoder._nearest_slots(torch.arange(4.0), torch.tensor([1.0, 3.0]), 2) + assert slots.tolist() == [[0, 1], [0, 1], [0, 1], [1, 0]] + slots = ltx2_diffusion_decoder._nearest_slots(torch.arange(2.0), torch.tensor([5.0]), 2) + assert slots.tolist() == [[0, -1], [0, -1]] + + # A tile keeps its planes plus the nearest one on each side: [56, 64] has none inside, keeps 48 and 96. + planes = torch.tensor([0, 48, 96, 144]) + assert ltx2_diffusion_decoder._keyframe_planes_for_tile(planes, 56, 64).tolist() == [False, True, True, False] + assert ltx2_diffusion_decoder._keyframe_planes_for_tile(planes, 0, 100).tolist() == [True, True, True, True] + assert ltx2_diffusion_decoder._keyframe_planes_for_tile(planes, 150, 200).tolist() == [ + False, + False, + False, + True, + ] + + def test_single_covering_tile_matches_untiled_with_keyframes(self): + """One tile spanning the video must reproduce the untiled keyframe decode bit for bit.""" + latents, keyframes = self.get_latents(num_planes=3) + model = self.get_model() + indices = torch.tensor([8, 16, 40]) # 40 lies past the 17-frame clip: kept as the nearest plane after it + for num_inference_steps in (None, 3): + untiled = self.decode( + model, + latents, + num_inference_steps=num_inference_steps, + keyframe_latents=keyframes, + keyframe_frame_indices=indices, + ) + generator = torch.Generator("cpu").manual_seed(0) + with torch.no_grad(): + tiled = model.tiled_decode( + latents, + generator=generator, + num_inference_steps=num_inference_steps, + keyframe_latents=keyframes, + keyframe_frame_indices=indices, + ) + assert torch.equal(tiled, untiled), (tiled - untiled).abs().max() + + def test_tiled_decode_with_splits_and_keyframes(self): + latents, keyframes = self.get_latents(num_planes=2, latent_frames=4) + model = self.get_model() + indices = torch.tensor([8, 24]) + untiled = self.decode(model, latents, keyframe_latents=keyframes, keyframe_frame_indices=indices) + # Cells are 2 frames x 4 px here: temporal tiles (0, 4), (3, 7), (6, 13) over the 13-cell grid, and two tiles + # over each spatial axis. + model.enable_tiling( + tile_sample_min_num_frames=8, + tile_sample_stride_num_frames=6, + tile_sample_min_height=32, + tile_sample_stride_height=24, + tile_sample_min_width=40, + tile_sample_stride_width=32, + ) + for num_inference_steps in (None, 3): + tiled = self.decode( + model, + latents, + num_inference_steps=num_inference_steps, + keyframe_latents=keyframes, + keyframe_frame_indices=indices, + ) + assert tiled.shape == untiled.shape + assert torch.isfinite(tiled).all() + model.disable_tiling() + assert torch.equal( + self.decode(model, latents, keyframe_latents=keyframes, keyframe_frame_indices=indices), untiled + ) + + def test_keyframe_inputs_are_validated(self): + latents, keyframes = self.get_latents() + model = self.get_model() + with pytest.raises(ValueError, match="keyframe_frame_indices"): + model.decode(latents, keyframe_latents=keyframes) + with pytest.raises(ValueError, match="one pixel frame per plane"): + model.decode(latents, keyframe_latents=keyframes, keyframe_frame_indices=[8]) + with pytest.raises(ValueError, match="non-negative"): + model.decode(latents, keyframe_latents=keyframes, keyframe_frame_indices=[-8, 8]) + with pytest.raises(ValueError, match="must match"): + model.decode(latents, keyframe_latents=keyframes[..., :2], keyframe_frame_indices=[8, 16]) + with pytest.raises(ValueError, match="at least one plane"): + model.decode(latents, keyframe_latents=keyframes[:, :, :0], keyframe_frame_indices=[]) diff --git a/tests/pipelines/ltx2/test_ltx2_hdr_seams.py b/tests/pipelines/ltx2/test_ltx2_hdr_seams.py new file mode 100644 index 000000000000..635439b1dc43 --- /dev/null +++ b/tests/pipelines/ltx2/test_ltx2_hdr_seams.py @@ -0,0 +1,74 @@ +# Copyright 2026 The HuggingFace Team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Seam keyframe positions of the LTX-2.5 SDR-To-HDR IC-LoRA. + +Expected values were produced by the reference implementation's own seam function (Lightricks/LTX-2, `hdr_ic_lora.py` +`_dfr_seam_pixel_positions` on top of `dfr_helpers/layout.py` `resolve_canvas`), not by the code under test. +""" + +import pytest + +from diffusers.pipelines.ltx2.dfr_layout import resolve_canvas, resolve_seam_positions + + +@pytest.mark.parametrize( + "num_frames, expected", + [ + (1, []), + (9, []), + (17, []), + (25, [24]), + (33, [32]), + (49, [24, 48]), + (57, [32]), + (97, [32, 64, 96]), + (105, [24, 48, 72, 96]), + (121, [24, 48, 72, 96, 120]), + (129, [32, 64, 96, 128]), + (193, [32, 64, 96, 128, 160, 192]), + (201, [24, 48, 72, 96, 120, 144, 168, 192]), + (257, [32, 64, 96, 128, 160, 192, 224, 256]), + ], +) +def test_seam_positions_match_reference(num_frames, expected): + assert resolve_seam_positions(num_frames) == expected + # High quality runs on the frame-doubled `2N - 1` grid: positions are computed on `N`, then doubled. + assert resolve_seam_positions(num_frames, high_quality=True) == [2 * position for position in expected] + + +def test_seam_positions_are_the_canvas_keyframes_inside_the_clip(): + # 105 frames pick 24-frame segments and pad the canvas to 121; the seams stop short of the padding, leaving an + # 8-frame tail segment. + canvas, segment, positions = resolve_canvas(105) + assert (canvas, segment, positions) == (121, 24, [24, 48, 72, 96, 120]) + assert resolve_seam_positions(105) == [24, 48, 72, 96] + + for num_frames in range(9, 2002, 8): + _, _, positions = resolve_canvas(num_frames) + assert resolve_seam_positions(num_frames) == [p for p in positions if p < num_frames] + + +def test_single_frame_has_no_seams_where_the_canvas_rejects_it(): + # The one 8k+1 frame count below 9 is a single frame: the reference returns no seams for it, while the DFR canvas + # (which needs at least one latent step) rejects it. + assert resolve_seam_positions(1) == [] + with pytest.raises(ValueError): + resolve_canvas(1) + + +@pytest.mark.parametrize("num_frames", [10, 96, 98]) +def test_seam_positions_reject_off_grid_frame_counts(num_frames): + with pytest.raises(ValueError): + resolve_seam_positions(num_frames) From cb28112557a399ee2f63a4d6c68ade29fe411ac5 Mon Sep 17 00:00:00 2001 From: christopher5106 Date: Tue, 6 Oct 2026 22:08:45 +0200 Subject: [PATCH 4/4] [LTX-2.5] Seam keyframes for the SDR-To-HDR IC-LoRA in LTX2HDRPipeline Port the seam keyframes of Lightricks' `HDRICLoraPipeline` (LTX-2 @ 9ec55f9f, `ltx_pipelines/hdr_ic_lora.py`) to the `hdr_transform="acescct"` path: - `keyframe_strength` (default 0.95, `None` = plain IC-LoRA): seams from `resolve_seam_positions`; a clip without seams warns and runs plain. Every seam gets a guide (the ACEScct source frame VAE-encoded alone in float32, tiled above 512x768 when VAE tiling is on, held at `keyframe_strength`) and a generated slot (mask 0, same one-pixel-frame RoPE span), appended as [video | reference | guides | slots] like the reference. The slots and the first latent frame carry the keyframe position embedding; seams need a transformer with `use_keyframes_abs_pos_embedding` and a single reference video. - Guide velocities are converted to x0 with each token's own timestep, as the reference X0 model does. - After denoising, the slots are cut out, denormalized like the video and passed to `diffusion_decoder.decode` as `keyframe_latents` / `keyframe_frame_indices`. Without a diffusion decoder the pipeline warns and the VAE decodes the video without them. - `high_quality_hdr`: frame-doubled source, `2N - 1` generated frames, doubled seams, every second frame kept. - Reuses the DFR helpers (`_prepare_keyframe_coords`, `_unpack_video_and_slots`) through "Copied from". `keyframe_strength=None` reproduces the parent commit bit for bit; the LogC3 path ignores `keyframe_strength` and rejects `high_quality_hdr`. On tiny shapes with a stand-in encoder and velocity model, the token sequence, RoPE positions, keyframe marker, per-token timesteps, the 8-step Euler trajectory and the extracted slots match the reference builders to within 6e-8 (video tokens and slots bitwise). Co-Authored-By: Claude Opus 5.5 --- .../pipelines/ltx2/pipeline_ltx2_hdr_lora.py | 408 ++++++++++++++++-- .../ltx2/test_ltx2_hdr_sdr_to_hdr.py | 251 ++++++++++- 2 files changed, 620 insertions(+), 39 deletions(-) diff --git a/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py b/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py index dd5c18f0dd6f..c62fcaad08c8 100644 --- a/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py +++ b/src/diffusers/pipelines/ltx2/pipeline_ltx2_hdr_lora.py @@ -32,12 +32,14 @@ from ...callbacks import MultiPipelineCallbacks, PipelineCallback from ...loaders import FromSingleFileMixin, LTX2LoraLoaderMixin from ...models.autoencoders import AutoencoderKLLTX2Audio, AutoencoderKLLTX2Video, LTX2VideoDiffusionDecoderModel +from ...models.autoencoders.vae import DiagonalGaussianDistribution from ...models.transformers import LTX2VideoTransformer3DModel from ...schedulers import FlowMatchEulerDiscreteScheduler from ...utils import is_torch_xla_available, logging, replace_example_docstring from ...utils.torch_utils import randn_tensor from ..pipeline_utils import DiffusionPipeline from .connectors import LTX2TextConnectors +from .dfr_layout import resolve_seam_positions from .image_processor import LTX2VideoHDRProcessor from .pipeline_output import LTX2PipelineOutput from .utils import DISTILLED_SIGMA_VALUES, SNAP_CONDITIONING_FPS_ABOVE @@ -56,6 +58,9 @@ # The LTX-2.5 SDR-To-HDR IC-LoRA conditions RoPE on 30 fps for sources above 30 fps; the playback frame rate is # unchanged (`_conditioning_fps` in `ltx_pipelines/hdr_ic_lora.py`). SDR_TO_HDR_MAX_CONDITIONING_FPS = 30.0 +# Frame area above which the SDR-To-HDR seam guides are VAE-encoded tile by tile, when VAE tiling is enabled +# (`TILED_VAE_ENCODE_PIXEL_THRESHOLD` in `ltx_pipelines/hdr_ic_lora.py`). +SDR_TO_HDR_TILED_ENCODE_PIXEL_THRESHOLD = 512 * 768 @dataclass @@ -170,6 +175,8 @@ class LTX2HDRReferenceCondition: >>> # output is cropped back. >>> sdr_video = load_video("/path/to/sdr.mp4") >>> frame_rate = 24.0 + >>> # Seam keyframes are on by default (`keyframe_strength=0.95`) and anchor the diffusion decoder; pass + >>> # `keyframe_strength=None` for the plain IC-LoRA, or `high_quality_hdr=True` to generate at twice the frames. >>> acescct = pipe( ... reference_conditions=LTX2HDRReferenceCondition(frames=sdr_video), ... connector_video_embeds=scene_embeds["video_context"], @@ -323,10 +330,10 @@ class LTX2HDRPipeline(DiffusionPipeline, FromSingleFileMixin, LTX2LoraLoaderMixi - Video-only (no audio output). The transformer's audio branch is still run since the diffusers transformer API requires audio inputs, but the decoded audio is discarded and audio-specific guidance scales are fixed to no-op values to avoid wasted compute. - - No frame-level keyframe conditioning (the reference HDR pipeline does not support this). + - No user-supplied frame-level keyframe conditioning (the reference HDR pipelines do not support this). With `hdr_transform="acescct"`, the pipeline runs the LTX-2.5 SDR-To-HDR IC-LoRA the way the reference - `HDRICLoraPipeline` does (without its optional seam keyframes): + `HDRICLoraPipeline` does: - the reference video is mapped to ACEScct (see `input_colorspace`), reflect-padded to a multiple of the VAE's spatial compression ratio and VAE-encoded in float32; the output is cropped back to `height` x `width`; @@ -336,6 +343,12 @@ class LTX2HDRPipeline(DiffusionPipeline, FromSingleFileMixin, LTX2LoraLoaderMixi influence on the video; - a single distilled stage (`DISTILLED_SIGMA_VALUES` by default, used as given), without CFG, STG or modality guidance, and with RoPE conditioned on at most 30 fps; + - seam keyframes (`keyframe_strength`, on by default): at every 24- or 32-frame segment seam of the clip (see + [`~pipelines.ltx2.dfr_layout.resolve_seam_positions`]), the source frame is VAE-encoded on its own and appended + as a one-frame guide at `keyframe_strength`, together with an empty generated keyframe slot at the same position. + The slots are denoised with the video, then anchor its decode by `diffusion_decoder`; + - optionally (`high_quality_hdr`), every source frame is doubled, `2 * num_frames - 1` frames are generated and + every second one is kept; - the latents are decoded in float32, by `diffusion_decoder` when it is loaded and by `vae` otherwise, then converted from ACEScct to scene-linear HDR (see `output_colorspace`). @@ -674,6 +687,7 @@ def check_sdr_to_hdr_inputs( guidance_scale=1.0, stg_scale=0.0, modality_scale=1.0, + keyframe_strength=None, ): r"""Input checks for `hdr_transform="acescct"` (LTX-2.5 SDR-To-HDR IC-LoRA).""" if callback_on_step_end_tensor_inputs is not None and not all( @@ -720,6 +734,25 @@ def check_sdr_to_hdr_inputs( f" latent_height, latent_width] are supported, but got {latents.ndim} dims." ) + if keyframe_strength is not None: + if not 0.0 <= keyframe_strength <= 1.0: + raise ValueError(f"`keyframe_strength` must be in [0, 1] or `None`, but got {keyframe_strength}.") + seams = resolve_seam_positions(num_frames, temporal_compression_ratio=self.vae_temporal_compression_ratio) + if seams and len(reference_conditions) != 1: + raise ValueError( + "Seam keyframes are encoded from the source video, so `keyframe_strength` needs exactly one" + f" reference video, but got {len(reference_conditions)}. Pass `keyframe_strength=None` to run" + " without seam keyframes." + ) + # `assert_stage_supports_generated_keyframes` in the reference: a generated slot is only told apart from a + # video frame by the learned keyframe embedding. + if seams and not self.transformer.config.use_keyframes_abs_pos_embedding: + raise ValueError( + "Seam keyframes generate keyframe slots, which needs a transformer whose config sets" + " `use_keyframes_abs_pos_embedding` (LTX-2.5). Pass `keyframe_strength=None` to run without seam" + " keyframes." + ) + @contextmanager def _float32(self, module: torch.nn.Module): r"""Run `module` in float32 for the duration of the context, then restore its dtype.""" @@ -857,6 +890,103 @@ def _unpack_audio_latents( latents = latents.unflatten(2, (-1, num_mel_bins)).transpose(1, 2) return latents + # Copied from diffusers.pipelines.ltx2.pipeline_ltx2_condition.LTX2ConditionPipeline._prepare_keyframe_coords + def _prepare_keyframe_coords( + self, + keyframe_latent_num_frames: int, + keyframe_latent_height: int, + keyframe_latent_width: int, + pixel_frame_idx: int, + num_pixel_frames: int, + fps: float, + device: torch.device, + ) -> torch.Tensor: + """ + Compute positional coordinates for a keyframe condition being appended as extra tokens. + + Mirrors `VideoConditionByKeyframeIndex.apply_to` in the reference implementation: + - Latent coords scaled to pixel space *without* the causal fix (since non-zero-index keyframes don't need the + first-frame causal adjustment). + - Temporal axis offset by `pixel_frame_idx` (the pixel-space index at which the keyframe appears). + - For single-pixel-frame keyframes, the per-patch temporal extent is clamped to `[idx, idx + 1)` so the + keyframe occupies a single pixel timestep rather than the VAE-scaled range. + - Temporal coords divided by `fps` to produce seconds. + """ + patch_size = self.transformer_spatial_patch_size + patch_size_t = self.transformer_temporal_patch_size + scale_factors = ( + self.vae_temporal_compression_ratio, + self.vae_spatial_compression_ratio, + self.vae_spatial_compression_ratio, + ) + + grid_f = torch.arange( + start=0, end=keyframe_latent_num_frames, step=patch_size_t, dtype=torch.float32, device=device + ) + grid_h = torch.arange(start=0, end=keyframe_latent_height, step=patch_size, dtype=torch.float32, device=device) + grid_w = torch.arange(start=0, end=keyframe_latent_width, step=patch_size, dtype=torch.float32, device=device) + grid = torch.meshgrid(grid_f, grid_h, grid_w, indexing="ij") + grid = torch.stack(grid, dim=0) + + patch_size_delta = torch.tensor((patch_size_t, patch_size, patch_size), dtype=grid.dtype, device=device) + patch_ends = grid + patch_size_delta.view(3, 1, 1, 1) + + latent_coords = torch.stack([grid, patch_ends], dim=-1) # [3, N_F, N_H, N_W, 2] + latent_coords = latent_coords.flatten(1, 3) # [3, num_patches, 2] + latent_coords = latent_coords.unsqueeze(0) # [1, 3, num_patches, 2] + + scale_tensor = torch.tensor(scale_factors, device=device, dtype=latent_coords.dtype) + broadcast_shape = [1] * latent_coords.ndim + broadcast_shape[1] = -1 + pixel_coords = latent_coords * scale_tensor.view(*broadcast_shape) + + # No causal fix: keyframe coords place the keyframe at `pixel_frame_idx` without the first-frame adjustment. + pixel_coords[:, 0, :, :] = pixel_coords[:, 0, :, :] + pixel_frame_idx + + if num_pixel_frames == 1: + # Single-pixel-frame keyframe: clamp temporal extent to [idx, idx + 1). + pixel_coords[:, 0, :, 1:] = pixel_coords[:, 0, :, :1] + 1 + + pixel_coords[:, 0, :, :] = pixel_coords[:, 0, :, :] / fps + + return pixel_coords + + # Copied from diffusers.pipelines.ltx2.pipeline_ltx2_dfr.LTX2DFRPipeline._unpack_video_latents + def _unpack_video_latents(self, tokens: torch.Tensor, num_frames: int, height: int, width: int) -> torch.Tensor: + """Unpack a `(batch_size, tokens, channels)` block onto this transformer's patch grid.""" + return self._unpack_latents( + tokens, + num_frames, + height, + width, + self.transformer_spatial_patch_size, + self.transformer_temporal_patch_size, + ) + + # Copied from diffusers.pipelines.ltx2.pipeline_ltx2_dfr.LTX2DFRPipeline._unpack_video_and_slots + def _unpack_video_and_slots( + self, + packed: torch.Tensor, + num_frames: int, + height: int, + width: int, + slot_token_slice: slice | None, + num_slots: int, + ) -> tuple[torch.Tensor, torch.Tensor | None]: + latent_num_frames = (num_frames - 1) // self.vae_temporal_compression_ratio + 1 + latent_height = height // self.vae_spatial_compression_ratio + latent_width = width // self.vae_spatial_compression_ratio + video = self._unpack_video_latents( + packed[:, : latent_num_frames * latent_height * latent_width], + latent_num_frames, + latent_height, + latent_width, + ) + keyframes = None + if slot_token_slice is not None and num_slots: + keyframes = self._unpack_video_latents(packed[:, slot_token_slice], num_slots, latent_height, latent_width) + return video, keyframes + def prepare_latents( self, reference_conditions: list[LTX2HDRReferenceCondition] | None = None, @@ -873,28 +1003,39 @@ def prepare_latents( generator: torch.Generator | None = None, latents: torch.Tensor | None = None, input_colorspace: str | None = None, + seam_frame_indices: list[int] | None = None, + keyframe_strength: float | None = None, + high_quality_hdr: bool = False, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None, int, torch.Tensor | None]: r""" Prepare noisy video latents, applying HDR IC-LoRA reference-video conditioning. - Builds a packed latent sequence in the order `[base | reference]`: + Builds a packed latent sequence in the order `[base | reference | guides | slots]`, the order in which the + reference `HDRICLoraPipeline` applies its conditionings: - Base: either fresh noise (Stage 1, `latents=None`) or pre-existing upsampled latents (Stage 2). - Reference: HDR-encoded reference-video tokens appended with per-token `conditioning_mask = strength`, - following the same pattern as [`LTX2InContextPipeline.prepare_latents`]. (HDR LoRA does not currently take - per-frame `conditions`, so there is no first-frame / keyframe block in between.) + following the same pattern as [`LTX2InContextPipeline.prepare_latents`]. + - Guides (`hdr_transform="acescct"` with `seam_frame_indices`): per seam, one latent frame encoded from the + source frame alone, with `conditioning_mask = keyframe_strength` and a RoPE time span of exactly one pixel + frame at the seam (`VideoConditionByKeyframeIndex`). + - Slots (same condition): per seam, one latent frame of empty tokens with `conditioning_mask = 0` and the + guide's coordinates, contiguous (`VideoGeneratedKeyframeSlots`). The model generates their content. Returns a 6-tuple matching [`LTX2InContextPipeline.prepare_latents`]: - - `latents`: packed noisy latents `(B, base + n_ref, C)`. - - `conditioning_mask`: `(B, seq_len, 1)` with `strength` at reference positions, `0` elsewhere. - - `clean_latents`: clean reference values at reference positions (zeros elsewhere); same shape as - `latents`. - - `appended_coords`: `[1, 3, n_ref, 2]` reference coordinates to concat onto `video_coords`, or `None` when - no reference conditions are provided. - - `num_ref_tokens`: count of reference tokens at the END of `latents`. + - `latents`: packed noisy latents `(B, base + n_appended, C)`. + - `conditioning_mask`: `(B, seq_len, 1)` with each appended token's strength, `0` elsewhere. + - `clean_latents`: clean conditioning values at reference and guide positions (zeros elsewhere); same shape + as `latents`. + - `appended_coords`: `[1, 3, n_appended, 2]` coordinates to concat onto `video_coords`, or `None` when no + reference conditions are provided. + - `num_appended_tokens`: count of the reference, guide and slot tokens at the END of `latents`. The slots + are the last `len(seam_frame_indices)` latent frames' worth of them. - `ref_cross_mask`: always `None` for HDR LoRA (no cross-attention masking support). `input_colorspace` is forwarded to [`LTX2VideoHDRProcessor.preprocess_reference_video_hdr`] and is only - supported with `hdr_transform="acescct"`. + supported with `hdr_transform="acescct"`, as are `seam_frame_indices`, `keyframe_strength` and + `high_quality_hdr` (`num_frames` is then the doubled `2 * N - 1` frame count; see + [`~LTX2HDRPipeline._encode_reference_conditions`]). """ latent_height = height // self.vae_spatial_compression_ratio latent_width = width // self.vae_spatial_compression_ratio @@ -940,7 +1081,7 @@ def prepare_latents( ref_coords: torch.Tensor | None = None num_ref_tokens = 0 if reference_conditions is not None and len(reference_conditions) > 0: - ref_latents_packed, ref_coords, _ = self._encode_reference_conditions( + ref_latents_packed, ref_coords, _, guide_latents = self._encode_reference_conditions( reference_conditions=reference_conditions, num_frames=num_frames, height=height, @@ -951,6 +1092,8 @@ def prepare_latents( device=device, generator=generator[0] if isinstance(generator, list) else generator, input_colorspace=input_colorspace, + high_quality_hdr=high_quality_hdr, + guide_frame_indices=seam_frame_indices, ) num_ref_tokens = ref_latents_packed.shape[1] @@ -972,8 +1115,64 @@ def prepare_latents( conditioning_mask = torch.cat([conditioning_mask, ref_mask_full], dim=1) clean_latents = torch.cat([clean_latents, ref_latents_packed_b], dim=1) - # HDR LoRA has no keyframe conditions, so the only appended tokens are reference tokens. appended_coords = ref_coords + num_appended_tokens = num_ref_tokens + if seam_frame_indices and ref_coords is not None: + latent_height_t = latent_height // self.transformer_spatial_patch_size + latent_width_t = latent_width // self.transformer_spatial_patch_size + appended_coords = [ref_coords] + + # Seam guides (`_keyframe_conditionings_from_pixel_frames`, `hdr_ic_lora.py:109-143`). Like the reference + # tokens above, their clean content also goes into `latents`, so the noising below leaves them at + # `keyframe_strength * clean + (1 - keyframe_strength) * noise`. + for position, guide_latent in zip(seam_frame_indices, guide_latents): + guide_tokens = self._pack_latents( + guide_latent, self.transformer_spatial_patch_size, self.transformer_temporal_patch_size + ).expand(batch_size, -1, -1) + latents = torch.cat([latents, guide_tokens], dim=1) + clean_latents = torch.cat([clean_latents, guide_tokens], dim=1) + conditioning_mask = torch.cat( + [ + conditioning_mask, + conditioning_mask.new_full((batch_size, guide_tokens.shape[1], 1), keyframe_strength), + ], + dim=1, + ) + appended_coords.append( + self._prepare_keyframe_coords( + keyframe_latent_num_frames=1, + keyframe_latent_height=guide_latent.shape[3], + keyframe_latent_width=guide_latent.shape[4], + pixel_frame_idx=position, + num_pixel_frames=1, + fps=frame_rate, + device=device, + ) + ) + + # Generated keyframe slots (`VideoGeneratedKeyframeSlots`, `keyframe_slots.py:71-174`): one empty latent + # frame per seam, fully denoised, at the guide's coordinates. + num_slot_tokens = len(seam_frame_indices) * latent_height_t * latent_width_t + slot_tokens = latents.new_zeros((batch_size, num_slot_tokens, latents.shape[2])) + latents = torch.cat([latents, slot_tokens], dim=1) + clean_latents = torch.cat([clean_latents, slot_tokens], dim=1) + conditioning_mask = torch.cat( + [conditioning_mask, conditioning_mask.new_zeros((batch_size, num_slot_tokens, 1))], dim=1 + ) + appended_coords.extend( + self._prepare_keyframe_coords( + keyframe_latent_num_frames=1, + keyframe_latent_height=latent_height, + keyframe_latent_width=latent_width, + pixel_frame_idx=position, + num_pixel_frames=1, + fps=frame_rate, + device=device, + ) + for position in seam_frame_indices + ) + appended_coords = torch.cat(appended_coords, dim=2) + num_appended_tokens = latents.shape[1] - base_seq_len # The conditioning_mask values have the following semantics: # - mask=0: fully noise tokens (e.g. noisy latents) @@ -983,7 +1182,7 @@ def prepare_latents( scaled_mask = (1.0 - conditioning_mask) * noise_scale # noise to initial noise level `noise_scale` latents = noise * scaled_mask + latents * (1 - scaled_mask) - return latents, conditioning_mask, clean_latents, appended_coords, num_ref_tokens, None + return latents, conditioning_mask, clean_latents, appended_coords, num_appended_tokens, None # Copied from diffusers.pipelines.ltx2.pipeline_ltx2_condition.LTX2ConditionPipeline.prepare_audio_latents def prepare_audio_latents( @@ -1030,15 +1229,22 @@ def _encode_reference_conditions( device: torch.device | None = None, generator: torch.Generator | None = None, input_colorspace: str | None = None, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]: - """Encode HDR IC-LoRA reference videos into `(reference_latents, reference_coords, reference_cross_mask)`. + high_quality_hdr: bool = False, + guide_frame_indices: list[int] | None = None, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None, list[torch.Tensor] | None]: + """Encode HDR IC-LoRA reference videos into `(reference_latents, reference_coords, reference_cross_mask, + guide_latents)`. Shared encoding core used by both `prepare_latents` (which folds reference tokens into the main noisy sequence) and the back-compat shim `prepare_reference_latents`. HDR LoRA does not currently support cross-attention masking for reference tokens, so the third return is always `None`. With `hdr_transform="acescct"` the reference is mapped to ACEScct (`input_colorspace`), must cover `num_frames` - frames, and is encoded with the VAE in float32 (`hdr_ic_lora.py:189, 427-431` in the reference). + frames, and is encoded with the VAE in float32 (`hdr_ic_lora.py:189, 427-431` in the reference). With + `high_quality_hdr`, `num_frames` is the doubled `2 * N - 1` count: the first `N` source frames are each + repeated twice and trimmed to `num_frames` before encoding (`hdr_ic_lora.py:322-325`). `guide_frame_indices` + (pixel frames of that same ACEScct clip) are encoded one frame at a time by + [`~LTX2HDRPipeline._encode_seam_guides`] and returned, normalized, as `guide_latents`; it is `None` otherwise. """ is_sdr_to_hdr = self.hdr_video_processor.config.hdr_transform == "acescct" ref_height = height // reference_downscale_factor @@ -1054,6 +1260,9 @@ def _encode_reference_conditions( all_ref_latents = [] all_ref_coords = [] + guide_latents = None + # Source frames read from the reference video: `N` of the `2 * N - 1` generated in high-quality mode. + num_source_frames = (num_frames + 1) // 2 if high_quality_hdr else num_frames for ref_cond in reference_conditions: if isinstance(ref_cond.frames, PIL.Image.Image): @@ -1071,17 +1280,22 @@ def _encode_reference_conditions( ref_pixels = self.hdr_video_processor.preprocess_reference_video_hdr( video_like, ref_height, ref_width, input_colorspace=input_colorspace ) - ref_pixels = ref_pixels[:, :, :num_frames] + ref_pixels = ref_pixels[:, :, :num_source_frames] if is_sdr_to_hdr: - if ref_pixels.shape[2] < num_frames: + if ref_pixels.shape[2] < num_source_frames: raise ValueError( - f"The reference video has {ref_pixels.shape[2]} frames, fewer than `num_frames={num_frames}`." + f"The reference video has {ref_pixels.shape[2]} frames, fewer than" + f" `num_frames={num_source_frames}`." ) + if high_quality_hdr: + ref_pixels = ref_pixels.repeat_interleave(2, dim=2)[:, :, :num_frames] with self._float32(self.vae): ref_pixels = ref_pixels.to(dtype=torch.float32, device=device) ref_latent = retrieve_latents( self.vae.encode(ref_pixels), generator=generator, sample_mode="argmax" ) + if guide_frame_indices: + guide_latents = self._encode_seam_guides(ref_pixels, guide_frame_indices, dtype, device) else: ref_pixels = ref_pixels.to(dtype=self.vae.dtype, device=device) ref_latent = retrieve_latents(self.vae.encode(ref_pixels), generator=generator, sample_mode="argmax") @@ -1113,7 +1327,36 @@ def _encode_reference_conditions( reference_latents = torch.cat(all_ref_latents, dim=1) reference_coords = torch.cat(all_ref_coords, dim=2) - return reference_latents, reference_coords, None + return reference_latents, reference_coords, None, guide_latents + + def _encode_seam_guides( + self, + pixels: torch.Tensor, + frame_indices: list[int], + dtype: torch.dtype | None = None, + device: torch.device | None = None, + ) -> list[torch.Tensor]: + r""" + VAE-encode each seam frame of an ACEScct `(1, C, F, H, W)` clip in `[-1, 1]` as a standalone one-frame clip. + + Mirrors `_keyframe_conditionings_from_pixel_frames` (`hdr_ic_lora.py:109-143`): the VAE is causal, so a frame + encoded alone is a single "first frame" latent, unlike the same frame inside the clip. As in the reference, + frames larger than `SDR_TO_HDR_TILED_ENCODE_PIXEL_THRESHOLD` pixels are encoded with `tiled_encode` when VAE + tiling is enabled ([`AutoencoderKLLTX2Video.enable_tiling`]), and without tiling otherwise. Call it with the + VAE in float32. Returns one normalized `(1, C, 1, H', W')` latent per index. + """ + num_frames, height, width = pixels.shape[2:] + use_tiled = self.vae.use_tiling and height * width > SDR_TO_HDR_TILED_ENCODE_PIXEL_THRESHOLD + guide_latents = [] + for frame_index in frame_indices: + if not 0 <= frame_index < num_frames: + raise ValueError(f"Seam frame {frame_index} is outside the {num_frames}-frame reference clip.") + frame = pixels[:, :, frame_index : frame_index + 1] + moments = self.vae.tiled_encode(frame) if use_tiled else self.vae.encoder(frame) + latent = DiagonalGaussianDistribution(moments).mode() + latent = self._normalize_latents(latent, self.vae.latents_mean, self.vae.latents_std) + guide_latents.append(latent.to(device=device, dtype=dtype)) + return guide_latents def prepare_reference_latents( self, @@ -1137,7 +1380,7 @@ def prepare_reference_latents( Returns a 3-tuple `(reference_latents, reference_coords, reference_denoise_factors)` with the same shapes as [`LTX2InContextPipeline.prepare_reference_latents`]. """ - reference_latents, reference_coords, _ = self._encode_reference_conditions( + reference_latents, reference_coords, _, _ = self._encode_reference_conditions( reference_conditions=reference_conditions, height=height, width=width, @@ -1237,6 +1480,8 @@ def __call__( reference_downscale_factor: int = 1, input_colorspace: str | None = None, output_colorspace: str | None = None, + keyframe_strength: float | None = 0.95, + high_quality_hdr: bool = False, height: int = 512, width: int = 768, num_frames: int = 121, @@ -1291,6 +1536,20 @@ def __call__( Colour space of the output with `hdr_transform="acescct"`: `"rec709"` (default, scene-linear Rec.709), `"acescg"` (scene-linear ACEScg) or `"acescct"` (the decoded ACEScct codes in `[0, 1]`). See [`LTX2VideoHDRProcessor.postprocess_hdr_video`]. Not supported with `hdr_transform="logc3"`. + keyframe_strength (`float`, *optional*, defaults to `0.95`): + Seam keyframes of `hdr_transform="acescct"`, as the reference CLI runs them by default (its + `--keyframe-strength`; `None` matches `--no-keyframes`). At every seam of + [`~pipelines.ltx2.dfr_layout.resolve_seam_positions`], the source frame is encoded on its own and + appended as a guide held at this strength, next to a generated keyframe slot at the same position; the + slots then anchor the decode by `diffusion_decoder`. Needs a transformer with + `use_keyframes_abs_pos_embedding` (LTX-2.5) and a single reference video. `None` runs the plain + IC-LoRA. Clips without a seam (fewer than 25 frames) run the plain IC-LoRA with a warning. Ignored with + `hdr_transform="logc3"`. + high_quality_hdr (`bool`, *optional*, defaults to `False`): + With `hdr_transform="acescct"`: repeat every source frame, generate `2 * num_frames - 1` frames (seams + doubled onto that grid) and keep every second decoded frame, as the reference `--high-quality` does. + Reduces temporal artifacts for about twice the cost. `output_type="latent"` returns the latents of the + doubled clip. Not supported with `hdr_transform="logc3"`. height (`int`, *optional*, defaults to `512`): Output video height in pixels. Must be divisible by 32 with `hdr_transform="logc3"`. With `hdr_transform="acescct"` any height works: the reference is reflect-padded up to a multiple of the VAE @@ -1353,7 +1612,8 @@ def __call__( One of `"pt"`, `"np"`, or `"latent"`. `"pt"` returns a linear HDR torch tensor in `[0, ∞)` of shape `(batch_size, num_frames, height, width, channels)`; `"np"` returns the equivalent `float32` NumPy array; `"latent"` returns the raw denoised latents (skip the HDR decode). With - `hdr_transform="acescct"` the latents cover the padded size and are not cropped. + `hdr_transform="acescct"` the latents cover the padded size and are not cropped, and the generated seam + keyframe slots are not returned. return_dict (`bool`, *optional*, defaults to `True`): Whether to return an [`LTX2PipelineOutput`] instead of a plain tuple. attention_kwargs (`dict`, *optional*): @@ -1395,12 +1655,15 @@ def __call__( guidance_scale=guidance_scale, stg_scale=stg_scale, modality_scale=modality_scale, + keyframe_strength=keyframe_strength, ) else: if input_colorspace is not None or output_colorspace is not None: raise ValueError( "`input_colorspace` and `output_colorspace` are only supported with `hdr_transform='acescct'`." ) + if high_quality_hdr: + raise ValueError("`high_quality_hdr` is only supported with `hdr_transform='acescct'`.") self.check_inputs( prompt=prompt, height=height, @@ -1458,8 +1721,34 @@ def __call__( # The full distilled schedule, used verbatim (`DEFAULT_DENOISE_SIGMAS`, `hdr_ic_lora.py:66`). if sigmas is None and timesteps is None: sigmas = DISTILLED_SIGMA_VALUES + + # Seam keyframes (`dfr_seam_roles`, `hdr_ic_lora.py:81-106, 284-302`): every seam gets both a guide and a + # generated slot. The positions are found on the source clip and doubled in high-quality mode, which + # generates `2 * num_frames - 1` frames from the frame-doubled source (`hdr_ic_lora.py:280-282`). + seam_frame_indices = [] + if keyframe_strength is not None: + seam_frame_indices = resolve_seam_positions( + num_frames, + high_quality=high_quality_hdr, + temporal_compression_ratio=self.vae_temporal_compression_ratio, + ) + if not seam_frame_indices: + logger.warning( + f"`keyframe_strength={keyframe_strength}` but this {num_frames}-frame clip has no seam keyframes;" + " running the plain IC-LoRA." + ) + elif output_type != "latent" and getattr(self, "diffusion_decoder", None) is None: + logger.warning( + "Seam keyframes are decoded by the keyframe-aware `diffusion_decoder`, which is not loaded:" + " the generated keyframe slots still take part in denoising, but `vae` decodes the video without" + " them. Load the LTX-2.5 `diffusion_decoder` to decode as the reference does, or pass" + " `keyframe_strength=None`." + ) + if high_quality_hdr: + num_frames = 2 * num_frames - 1 else: conditioning_frame_rate = frame_rate + seam_frame_indices = [] if noise_scale is None: noise_scale = sigmas[0] if sigmas is not None else 1.0 @@ -1533,7 +1822,7 @@ def __call__( _, _, latent_num_frames, latent_height, latent_width = latents.shape num_channels_latents = self.transformer.config.in_channels - latents, conditioning_mask, clean_latents, appended_coords, num_ref_tokens, _ = self.prepare_latents( + latents, conditioning_mask, clean_latents, appended_coords, num_condition_tokens, _ = self.prepare_latents( reference_conditions=reference_conditions, reference_downscale_factor=reference_downscale_factor, batch_size=batch_size * num_videos_per_prompt, @@ -1548,11 +1837,14 @@ def __call__( generator=generator, latents=latents, input_colorspace=input_colorspace, + seam_frame_indices=seam_frame_indices, + keyframe_strength=keyframe_strength, + high_quality_hdr=high_quality_hdr, ) - # Track the base (non-reference) token count so we can trim the appended reference tokens off - # `latents` before unpack/decode at the end. - base_token_count = latents.shape[1] - num_ref_tokens - if self.do_classifier_free_guidance and num_ref_tokens > 0: + # Track the base (non-conditioning) token count so we can trim the appended reference, guide and slot tokens + # off `latents` before unpack/decode at the end. + base_token_count = latents.shape[1] - num_condition_tokens + if self.do_classifier_free_guidance and num_condition_tokens > 0: conditioning_mask = torch.cat([conditioning_mask, conditioning_mask]) # 5. Prepare audio latents. Audio is discarded at the end, but the transformer's audio branch still runs so @@ -1650,6 +1942,12 @@ def __call__( ) video_keyframes_mask = torch.zeros((latents.shape[0], latents.shape[1], 1), device=device) video_keyframes_mask[:, :tokens_per_latent_frame] = 1.0 + # The generated slots carry the same marker; the guides do not (`keyframe_slots.py:121`, + # `keyframe_cond.py:84-86`). The slots are the last tokens of the sequence. + num_slot_tokens = len(seam_frame_indices) * tokens_per_latent_frame + slot_token_slice = slice(latents.shape[1] - num_slot_tokens, latents.shape[1]) if num_slot_tokens else None + if slot_token_slice is not None: + video_keyframes_mask[:, slot_token_slice] = 1.0 # 8. Denoising loop video_seq_len = latents.shape[1] @@ -1669,7 +1967,7 @@ def __call__( audio_latent_model_input = audio_latent_model_input.to(connector_prompt_embeds.dtype) timestep_scalar = t.expand(latent_model_input.shape[0]) - if num_ref_tokens > 0: + if num_condition_tokens > 0: video_timestep = timestep_scalar.unsqueeze(-1) * (1 - conditioning_mask.squeeze(-1)) else: video_timestep = timestep_scalar.unsqueeze(-1).expand(-1, video_seq_len) @@ -1727,7 +2025,7 @@ def __call__( video_pos_ids = video_coords.chunk(2, dim=0)[0] audio_pos_ids = audio_coords.chunk(2, dim=0)[0] timestep_scalar_single = timestep_scalar.chunk(2, dim=0)[0] - if num_ref_tokens > 0: + if num_condition_tokens > 0: video_timestep_single = video_timestep.chunk(2, dim=0)[0] else: video_timestep_single = timestep_scalar_single.unsqueeze(-1).expand(-1, video_seq_len) @@ -1742,13 +2040,21 @@ def __call__( audio_pos_ids = audio_coords timestep_scalar_single = timestep_scalar - if num_ref_tokens > 0: + if num_condition_tokens > 0: video_timestep_single = video_timestep else: video_timestep_single = timestep_scalar.unsqueeze(-1).expand(-1, video_seq_len) audio_timestep_single = audio_timestep - noise_pred_video = self.convert_velocity_to_x0(latents, noise_pred_video, i, self.scheduler) + if is_sdr_to_hdr and num_condition_tokens > 0: + # The reference converts the velocity with each token's own timestep, `sigma * denoise_mask` + # (`to_denoised(latent, v, timesteps)`, `ltx_core/model/transformer/model.py:598`). This only + # matters for the partially-clean seam guides: fully denoised tokens use `sigma`, and clean + # tokens are replaced by their clean value below. + token_sigmas = self.scheduler.sigmas[i] * (1 - conditioning_mask) + noise_pred_video = latents - noise_pred_video * token_sigmas + else: + noise_pred_video = self.convert_velocity_to_x0(latents, noise_pred_video, i, self.scheduler) # --- STG forward pass (video only — audio output discarded) --- if self.do_spatio_temporal_guidance: @@ -1836,7 +2142,7 @@ def __call__( noise_pred_video = noise_pred_video_g # Apply the conditioning mask to apply the reference conditions at the specified strength. - if num_ref_tokens > 0: + if num_condition_tokens > 0: bsz = noise_pred_video.size(0) denoised_sample_cond = ( noise_pred_video * (1 - conditioning_mask[:bsz]) @@ -1868,7 +2174,13 @@ def __call__( xm.mark_step() # 9. Decode - # Trim any appended reference tokens from the latents to recover the generated video only. + # The generated seam keyframe slots, `(B, C, num_seams, H, W)` (`clear_conditioning`, `ltx_core/tools.py:88-117`). + slot_latents = None + if seam_frame_indices: + _, slot_latents = self._unpack_video_and_slots( + latents, num_frames, height, width, slot_token_slice, len(seam_frame_indices) + ) + # Trim the appended reference, guide and slot tokens from the latents to recover the generated video only. latents = latents[:, :base_token_count] latents = self._unpack_latents( latents, @@ -1911,13 +2223,35 @@ def __call__( if is_sdr_to_hdr: # Decoded ACEScct codes in [-1, 1], cropped back to the requested size. diffusion_decoder = getattr(self, "diffusion_decoder", None) + keyframe_kwargs = {} + if slot_latents is not None: + # `decode_keyframes_from_slots` (`ltx_pipelines/utils/helpers.py:539-559`): keep the slots that land + # inside the generated clip (all of them here, the seams are clipped to it), denormalized like the + # video, at their pixel frames. + kept = [(index, p) for index, p in enumerate(seam_frame_indices) if 0 <= p < num_frames] + slot_latents = self._denormalize_latents( + slot_latents[:, :, [index for index, _ in kept]].to(torch.float32), + self.vae.latents_mean, + self.vae.latents_std, + self.vae.config.scaling_factor, + ) + keyframe_kwargs = { + "keyframe_latents": slot_latents, + "keyframe_frame_indices": torch.tensor([p for _, p in kept], dtype=torch.long), + } if diffusion_decoder is not None: with self._float32(diffusion_decoder): - decoded = diffusion_decoder.decode(latents, generator=generator, return_dict=False)[0] + decoded = diffusion_decoder.decode( + latents, generator=generator, return_dict=False, **keyframe_kwargs + )[0] else: + # The VAE cannot take keyframe planes; the slots only shaped the denoising (warned above). with self._float32(self.vae): decoded = self.vae.decode(latents, timestep, return_dict=False)[0] decoded = decoded[:, :, :, :output_height, :output_width] + if high_quality_hdr: + # Undo the frame doubling (`decoded[::2]`, `hdr_ic_lora.py:521-522`). + decoded = decoded[:, :, ::2] else: latents = latents.to(self.vae.dtype) diff --git a/tests/pipelines/ltx2/test_ltx2_hdr_sdr_to_hdr.py b/tests/pipelines/ltx2/test_ltx2_hdr_sdr_to_hdr.py index 916b3bd883e0..b65c7ecf42f4 100644 --- a/tests/pipelines/ltx2/test_ltx2_hdr_sdr_to_hdr.py +++ b/tests/pipelines/ltx2/test_ltx2_hdr_sdr_to_hdr.py @@ -21,12 +21,15 @@ from diffusers import LTX2HDRPipeline, LTX2VideoDiffusionDecoderModel from diffusers.pipelines.ltx2 import LTX2HDRReferenceCondition +from diffusers.pipelines.ltx2 import pipeline_ltx2_hdr_lora as hdr_module from diffusers.pipelines.ltx2.utils import DISTILLED_SIGMA_VALUES +from diffusers.utils import logging from diffusers.utils.import_utils import is_peft_available -from ...testing_utils import enable_full_determinism, torch_device +from ...testing_utils import CaptureLogger, enable_full_determinism, torch_device from ..testing_utils.common import BasePipelineOutputMixin from .test_ltx2_hdr import LTX2HDRPipelineTesterConfig +from .testing_utils import DFR_TRANSFORMER_KWARGS, get_ltx2_dummy_components enable_full_determinism() @@ -130,7 +133,8 @@ def test_no_text_encoder_call(self): ({"modality_scale": 2.0}, "guidance"), ({"reference_conditions": None}, "reference"), ({"num_frames": 4}, "num_frames"), - ({"num_frames": 7}, "fewer than"), + # 7 frames has a seam with the dummy VAE, which this transformer could not generate. + ({"num_frames": 7, "keyframe_strength": None}, "fewer than"), ], ) def test_invalid_inputs(self, overrides, match): @@ -308,6 +312,20 @@ def test_first_latent_frame_is_marked_as_keyframe(self): assert mask[:, :tokens_per_frame].eq(1).all() assert mask[:, tokens_per_frame:].eq(0).all() + def test_keyframe_strength_none_matches_parent_commit(self): + # 25 frames would have seam keyframes; `keyframe_strength=None` must give back the plain IC-LoRA run exactly. + # Recorded on the parent commit, before seam keyframes were added. + pipe = self.get_sdr_to_hdr_pipeline() + inputs = self.get_sdr_to_hdr_inputs(output_type="latent", keyframe_strength=None) + frames = torch.randint(0, 256, (25, HEIGHT, WIDTH, 3), generator=inputs["generator"], dtype=torch.uint8) + inputs.update(reference_conditions=LTX2HDRReferenceCondition(frames=frames.numpy()), num_frames=25) + latents = pipe(**inputs).frames.flatten().cpu() + expected = torch.tensor( + [-0.7471, 1.8191, -1.6666, 0.8206, 1.2447, -0.8798, -1.1223, 0.6885] + + [0.6787, -0.5253, 0.2984, -0.374, -0.6168, -0.3119, 0.5331, -0.7598] + ) + assert torch.allclose(torch.cat([latents[:8], latents[-8:]]), expected, atol=1e-3) + def test_hdr_transform_survives_save_load(self, tmp_path): pipe = self.get_sdr_to_hdr_pipeline(text_components=True) pipe.save_pretrained(str(tmp_path), safe_serialization=False) @@ -360,6 +378,219 @@ def test_sdr_to_hdr_lora_key_format(self): assert not torch.allclose(base, with_lora) +# Seam positions scale with the VAE's temporal compression (segments of 3 or 4 latent frames): with the dummy x2 VAE, +# 7 frames has one seam at 6 (the x8 VAE's 25 frames -> [24]), 25 frames has [8, 16, 24], 5 frames has none. +ONE_SEAM_FRAMES, ONE_SEAM = 7, [6] +THREE_SEAMS_FRAMES, THREE_SEAMS = 25, [8, 16, 24] +TOKENS_PER_FRAME = (PADDED_HEIGHT // 2) * (PADDED_WIDTH // 2) + + +class TestLTX2HDRPipelineSeamKeyframes(LTX2HDRPipelineTesterConfig, BasePipelineOutputMixin): + def get_seam_pipeline(self, diffusion_decoder: bool = True, keyframe_embedding: bool = True): + # The DFR transformer config: the keyframe position embedding, and RoPE on the dummy VAE's x2 grid. + transformer_kwargs = DFR_TRANSFORMER_KWARGS if keyframe_embedding else {"vae_scale_factors": (2, 2, 2)} + components = get_ltx2_dummy_components( + unset_components=self.unset_components, transformer_kwargs=transformer_kwargs + ) + components.update(text_encoder=None, tokenizer=None, connectors=None) + if diffusion_decoder: + components["diffusion_decoder"] = get_dummy_diffusion_decoder() + pipe = LTX2HDRPipeline(**components, hdr_transform="acescct") + pipe.set_progress_bar_config(disable=True) + return pipe.to(torch_device) + + def get_seam_inputs(self, num_frames: int = ONE_SEAM_FRAMES, **overrides): + generator = self.get_generator(0) + frames = torch.randint(0, 256, (num_frames, HEIGHT, WIDTH, 3), generator=generator, dtype=torch.uint8) + inputs = { + "reference_conditions": LTX2HDRReferenceCondition(frames=frames.numpy()), + "connector_video_embeds": torch.randn(7, CONTEXT_DIM, generator=generator), + "height": HEIGHT, + "width": WIDTH, + "num_frames": num_frames, + "frame_rate": 24.0, + "sigmas": [1.0, 0.5], + "generator": generator, + "output_type": "pt", + } + inputs.update(overrides) + return inputs + + @staticmethod + def spy_transformer(pipe): + return mock.patch.object(pipe.transformer, "forward", side_effect=pipe.transformer.forward) + + def test_one_seam_token_layout(self): + pipe = self.get_seam_pipeline() + with self.spy_transformer(pipe) as spy: + pipe(**self.get_seam_inputs(output_type="latent")) + kwargs = spy.call_args_list[0].kwargs + t = TOKENS_PER_FRAME + num_base = ((ONE_SEAM_FRAMES - 1) // 2 + 1) * t + # The dummy VAE only downsamples in space, so its reference latents keep every frame. + num_reference = ONE_SEAM_FRAMES * t + # [base | reference | guide | slot], the order the reference applies its conditionings in. + assert kwargs["hidden_states"].shape[1] == num_base + num_reference + t + t + guide = slice(num_base + num_reference, num_base + num_reference + t) + slot = slice(guide.stop, guide.stop + t) + + mask = kwargs["video_keyframes_mask"][0, :, 0] + assert mask[:t].eq(1).all() and mask[t : guide.start].eq(0).all() + assert mask[guide].eq(0).all() and mask[slot].eq(1).all() + + # Per-token timesteps: the guide is held at strength 0.95, the slot is fully denoised. + sigma = kwargs["sigma"][0] + timestep = kwargs["timestep"][0] + assert torch.allclose(timestep[guide], sigma * (1 - 0.95), atol=1e-4) + assert torch.allclose(timestep[slot], sigma.expand(t)) + assert timestep[num_base : guide.start].eq(0).all() + + # Guide and slot share their RoPE coordinates: one pixel frame, [6, 7) / fps in time. + coords = kwargs["video_coords"][0] + assert torch.equal(coords[:, guide], coords[:, slot]) + assert torch.allclose(coords[0, slot, 0], torch.tensor(6 / 24.0, device=coords.device)) + assert torch.allclose(coords[0, slot, 1], torch.tensor(7 / 24.0, device=coords.device)) + + def test_guide_is_a_standalone_one_frame_encode(self): + pipe = self.get_seam_pipeline() + calls = [] + encoder_forward = pipe.vae.encoder.forward + + def encoder_spy(x, *args, **kwargs): + calls.append(x.detach().clone()) + return encoder_forward(x, *args, **kwargs) + + with mock.patch.object(pipe.vae.encoder, "forward", side_effect=encoder_spy): + pipe(**self.get_seam_inputs(output_type="latent")) + clip, guide = calls + assert clip.shape[2] == ONE_SEAM_FRAMES and clip.dtype == torch.float32 + assert torch.equal(guide, clip[:, :, ONE_SEAM[0] : ONE_SEAM[0] + 1]) + + def test_slots_are_passed_to_the_decoder(self): + pipe = self.get_seam_pipeline() + final = {} + + def callback(pipe, i, t, callback_kwargs): + final["latents"] = callback_kwargs["latents"] + return {} + + decode = pipe.diffusion_decoder.decode + with ( + mock.patch.object(pipe.diffusion_decoder, "decode", side_effect=decode) as spy, + mock.patch.object(pipe.vae, "decode", side_effect=AssertionError("the VAE must not decode")), + ): + video = pipe(**self.get_seam_inputs(THREE_SEAMS_FRAMES, callback_on_step_end=callback)).frames + assert video.shape == (1, THREE_SEAMS_FRAMES, HEIGHT, WIDTH, 3) + + kwargs = spy.call_args.kwargs + assert kwargs["keyframe_frame_indices"].tolist() == THREE_SEAMS + keyframes = kwargs["keyframe_latents"] + assert keyframes.shape == (1, 4, len(THREE_SEAMS), PADDED_HEIGHT // 2, PADDED_WIDTH // 2) + assert keyframes.dtype == torch.float32 + # The slots are the last tokens of the denoised sequence; the dummy VAE's statistics make denormalization + # the identity. + slot_tokens = final["latents"][:, -len(THREE_SEAMS) * TOKENS_PER_FRAME :] + expected = pipe._unpack_video_latents(slot_tokens, len(THREE_SEAMS), PADDED_HEIGHT // 2, PADDED_WIDTH // 2) + assert torch.allclose(keyframes, expected.float()) + # The reference, guide and slot tokens are dropped from the decoded video. + assert spy.call_args.args[0].shape == (1, 4, (THREE_SEAMS_FRAMES - 1) // 2 + 1, 16, 15) + + def test_seams_change_the_output(self): + pipe = self.get_seam_pipeline() + with_seams = pipe(**self.get_seam_inputs(output_type="latent")).frames + without = pipe(**self.get_seam_inputs(output_type="latent", keyframe_strength=None)).frames + assert with_seams.shape == without.shape + assert not torch.allclose(with_seams, without) + with self.spy_transformer(pipe) as spy: + pipe(**self.get_seam_inputs(output_type="latent", keyframe_strength=None)) + num_base = ((ONE_SEAM_FRAMES - 1) // 2 + 1) * TOKENS_PER_FRAME + assert spy.call_args.kwargs["hidden_states"].shape[1] == num_base + ONE_SEAM_FRAMES * TOKENS_PER_FRAME + + def test_no_seams_runs_plain_ic_lora(self): + # 5 frames has no seam: the default `keyframe_strength` warns and runs the plain IC-LoRA, even on a transformer + # without the keyframe embedding. + pipe = self.get_seam_pipeline(keyframe_embedding=False) + logger = logging.get_logger(hdr_module.__name__) + with CaptureLogger(logger) as captured: + default = pipe(**self.get_seam_inputs(5, output_type="latent")).frames + assert "no seam keyframes" in captured.out + plain = pipe(**self.get_seam_inputs(5, output_type="latent", keyframe_strength=None)).frames + assert torch.equal(default, plain) + + def test_seams_without_diffusion_decoder_warn_and_decode_with_the_vae(self): + pipe = self.get_seam_pipeline(diffusion_decoder=False) + logger = logging.get_logger(hdr_module.__name__) + with CaptureLogger(logger) as captured: + video = pipe(**self.get_seam_inputs()).frames + assert "`diffusion_decoder`, which is not loaded" in captured.out + assert video.shape == (1, ONE_SEAM_FRAMES, HEIGHT, WIDTH, 3) + assert torch.isfinite(video).all() + + def test_high_quality_hdr(self): + pipe = self.get_seam_pipeline() + encoded = [] + encoder_forward = pipe.vae.encoder.forward + + def encoder_spy(x, *args, **kwargs): + encoded.append(x.detach().clone()) + return encoder_forward(x, *args, **kwargs) + + decode = pipe.diffusion_decoder.decode + with ( + mock.patch.object(pipe.vae.encoder, "forward", side_effect=encoder_spy), + mock.patch.object(pipe.diffusion_decoder, "decode", side_effect=decode) as spy, + ): + video = pipe(**self.get_seam_inputs(high_quality_hdr=True)).frames + + generated = 2 * ONE_SEAM_FRAMES - 1 + clip, guide = encoded + # Every source frame is doubled, then trimmed to 2N - 1 frames. + assert clip.shape[2] == generated + assert torch.equal(clip[:, :, 0:-1:2], clip[:, :, 1::2]) + # The seam is doubled onto that grid, so its guide is still source frame 6. + assert torch.equal(guide, clip[:, :, 12:13]) + assert spy.call_args.kwargs["keyframe_frame_indices"].tolist() == [2 * p for p in ONE_SEAM] + assert spy.call_args.args[0].shape[2] == (generated - 1) // 2 + 1 + # Every second decoded frame is kept. + assert video.shape == (1, ONE_SEAM_FRAMES, HEIGHT, WIDTH, 3) + + def test_tiled_guide_encode_above_the_threshold(self, monkeypatch): + pipe = self.get_seam_pipeline() + pipe.vae.enable_tiling() + pixels = torch.zeros((1, 3, ONE_SEAM_FRAMES, PADDED_HEIGHT, PADDED_WIDTH), device=torch_device) + with mock.patch.object(pipe.vae, "tiled_encode", side_effect=pipe.vae.tiled_encode) as tiled: + pipe._encode_seam_guides(pixels, ONE_SEAM) + assert tiled.call_count == 0 # 32 x 30 is below 512 x 768 + monkeypatch.setattr(hdr_module, "SDR_TO_HDR_TILED_ENCODE_PIXEL_THRESHOLD", 16 * 16) + pipe._encode_seam_guides(pixels, ONE_SEAM) + assert tiled.call_count == 1 + + @pytest.mark.parametrize( + "overrides, match", + [ + ({"keyframe_strength": 1.5}, "keyframe_strength"), + ({"keyframe_strength": -0.1}, "keyframe_strength"), + ], + ) + def test_invalid_keyframe_strength(self, overrides, match): + pipe = self.get_seam_pipeline() + with pytest.raises(ValueError, match=match): + pipe(**self.get_seam_inputs(**overrides)) + + def test_seams_need_the_keyframe_embedding(self): + pipe = self.get_seam_pipeline(keyframe_embedding=False) + with pytest.raises(ValueError, match="use_keyframes_abs_pos_embedding"): + pipe(**self.get_seam_inputs()) + pipe(**self.get_seam_inputs(keyframe_strength=None, output_type="latent")) + + def test_seams_need_a_single_reference_video(self): + pipe = self.get_seam_pipeline() + inputs = self.get_seam_inputs() + inputs["reference_conditions"] = [inputs["reference_conditions"]] * 2 + with pytest.raises(ValueError, match="exactly one reference video"): + pipe(**inputs) + + class TestLTX2HDRPipelineLogC3Unchanged(LTX2HDRPipelineTesterConfig, BasePipelineOutputMixin): def test_default_hdr_transform_is_logc3(self): pipe = self.get_pipeline() @@ -374,6 +605,22 @@ def test_colorspaces_rejected_with_logc3(self, argument): with pytest.raises(ValueError, match="only supported with `hdr_transform='acescct'`"): pipe(**inputs) + def test_high_quality_hdr_rejected_with_logc3(self): + pipe = self.get_pipeline().to(torch_device) + inputs = self.get_dummy_inputs() + inputs["high_quality_hdr"] = True + with pytest.raises(ValueError, match="only supported with `hdr_transform='acescct'`"): + pipe(**inputs) + + def test_keyframe_strength_ignored_with_logc3(self): + pipe = self.get_pipeline().to(torch_device) + outputs = [] + for keyframe_strength in (0.95, None): + inputs = self.get_dummy_inputs() + inputs.update(output_type="latent", keyframe_strength=keyframe_strength) + outputs.append(pipe(**inputs).frames) + assert torch.equal(outputs[0], outputs[1]) + def test_logc3_output_slice(self): # Recorded on the parent commit, before the SDR-To-HDR path was added. pipe = self.get_pipeline().to(torch_device)