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/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/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.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_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()) 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)