diff --git a/src/diffusers/pipelines/ltx2/export_utils.py b/src/diffusers/pipelines/ltx2/export_utils.py index b2206a79842c..be004023f3d2 100644 --- a/src/diffusers/pipelines/ltx2/export_utils.py +++ b/src/diffusers/pipelines/ltx2/export_utils.py @@ -13,14 +13,18 @@ # See the License for the specific language governing permissions and # limitations under the License. +import math +import os from fractions import Fraction from pathlib import Path from typing import Callable import numpy as np import torch +import torch.nn.functional as F -from ...utils import is_av_available +from ...utils import is_av_available, is_openexr_available +from .image_processor import ACESCG_TO_REC2020, REC709_TO_REC2020 _CAN_USE_AV = is_av_available() @@ -108,3 +112,335 @@ def simple_tone_map(x: np.ndarray) -> np.ndarray: container.mux(packet) finally: container.close() + + +# ARIB STD-B67 HLG OETF constants, as used by the LTX-2 reference (`colour` `CONSTANTS_ARIBSTDB67`). +_HLG_A = 0.17883277 +_HLG_B = 0.28466892 +_HLG_C = 0.55991073 + +# FFmpeg colour tags (`AVCOL_PRI_BT2020`, `AVCOL_TRC_ARIB_STD_B67`, `AVCOL_SPC_BT2020_NCL`, `AVCOL_RANGE_MPEG`). +_AV_COLOR_PRIMARIES_BT2020 = 9 +_AV_COLOR_TRC_ARIB_STD_B67 = 18 +_AV_COLORSPACE_BT2020_NCL = 9 +_AV_COLOR_RANGE_MPEG = 1 + +# Full-range RGB -> Y'CbCr matrix with the BT.2020 non-constant-luminance weights (Kr = 0.2627, Kb = 0.0593). The +# reference computes it as the float64 inverse of `colour.matrix_YCbCr(WEIGHTS_YCBCR["ITU-R BT.2020"])` stored as +# float32; these are those float32 values. +_RGB_TO_YCBCR_BT2020 = ( + (0.26269999146461487, 0.6779999732971191, 0.059300001710653305), + (-0.13963006436824799, -0.3603699505329132, 0.5), + (0.5, -0.45978569984436035, -0.04021429643034935), +) + +_HLG_PRIMARIES_TO_REC2020 = {"rec709": REC709_TO_REC2020, "acescg": ACESCG_TO_REC2020} + + +def _hlg_inverse_oetf(signal: float) -> float: + r"""Inverse HLG OETF (ITU-R BT.2100 reference constants): HLG signal `[0, 1]` -> scene-linear `[0, 1]`.""" + a = _HLG_A + b = 1.0 - 4.0 * a + c = 0.5 - a * math.log(4.0 * a) + linear = (signal / 0.5) ** 2 if signal <= 0.5 else math.exp((signal - c) / a) + b + return linear / 12.0 + + +def _hlg_oetf(x: torch.Tensor) -> torch.Tensor: + r"""HLG OETF (ARIB STD-B67): scene-linear `[0, 1]` -> HLG signal, clamped to `[0, 1]`.""" + return torch.where( + x <= 1.0 / 12.0, + torch.sqrt((3.0 * x).clamp(min=0.0)), + _HLG_A * torch.log((12.0 * x - _HLG_B).clamp(min=1e-12)) + _HLG_C, + ).clamp(0.0, 1.0) + + +def _linear_to_hlg_signal( + rgb_linear: torch.Tensor, primaries_matrix: torch.Tensor, white_x: float, roll_k: float +) -> torch.Tensor: + r"""Scene-linear `(..., 3, H, W)` RGB -> Rec.2020 HLG signal `[0, 1]`, with diffuse white mapped to `white_x`.""" + lin = torch.nan_to_num( + torch.einsum("...chw,dc->...dhw", rgb_linear, primaries_matrix).clamp(min=0.0), + nan=0.0, + neginf=0.0, + ) + # Diffuse white (linear 1.0) maps to `white_x`; highlights roll off exponentially toward 1.0. + x = torch.where( + lin <= 1.0, + lin * white_x, + 1.0 - (1.0 - white_x) * torch.exp(-roll_k * (lin - 1.0)), + ) + return _hlg_oetf(x) + + +def _rgb_to_yuv420p10_bt2020_limited(rgb: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + r""" + Float RGB `[0, 1]` `(F, 3, H, W)` -> planar 10-bit limited-range BT.2020 NCL Y, U, V code values (4:2:0). + + Returns `int32` tensors with values in `[0, 1023]`. + """ + _, _, height, width = rgb.shape + if height % 2 != 0 or width % 2 != 0: + raise ValueError(f"HLG export requires an even frame height and width, got {height}x{width}.") + matrix = torch.as_tensor(_RGB_TO_YCBCR_BT2020, dtype=torch.float32).to(device=rgb.device, dtype=rgb.dtype) + yuv = (rgb.movedim(-3, -1).flatten(-3, -2) @ matrix.T).unflatten(-2, (height, width)).movedim(-1, -3) + y = yuv[:, :1] + uv = F.avg_pool2d(yuv[:, 1:3].contiguous(), kernel_size=2, stride=2) + # Limited ("MPEG") range at 10 bits: Y = (219 * E' + 16) * 4, Cb/Cr = (224 * E' + 128) * 4. + y = y * (219 * 4) + 16 * 4 + uv = uv * (224 * 4) + 128 * 4 + y = y[:, 0].round().clamp(0, 1023).to(torch.int32) + u = uv[:, 0].round().clamp(0, 1023).to(torch.int32) + v = uv[:, 1].round().clamp(0, 1023).to(torch.int32) + return y, u, v + + +def _x265_params(threads: int, width: int, height: int) -> str: + r"""libx265 `x265-params` for a BT.2020 HLG `hvc1` MP4, as in the reference implementation.""" + params = ( + "colorprim=bt2020:transfer=arib-std-b67:colormatrix=bt2020nc:range=limited:" + f"repeat-headers=1:info=0:pools={threads}" + ) + if width <= 32 and height <= 32: + # With a single CTU per axis, frame-threading and B-frames can flush packets that the MP4 muxer rejects on + # very short clips. + return f"{params}:frame-threads=1:bframes=0:lookahead=0" + return f"{params}:frame-threads=4" + + +def encode_hdr_tensor_to_hlg_mp4( + frames: torch.Tensor | np.ndarray, + output_mp4: str | Path, + frame_rate: float, + primaries: str = "rec709", + white_signal: float = 0.75, + rolloff_k: float | None = None, + crf: int = 12, + preset: str = "ultrafast", + thread_count: int = 0, + device: str | torch.device | None = None, +) -> None: + r""" + Encodes a scene-linear HDR tensor to a BT.2020 / HLG / 10-bit HEVC `.mp4` file, following the LTX-2 reference HDR + export. + + Each frame is converted to Rec.2020 primaries; diffuse white (linear `1.0`) is mapped to the HLG signal + `white_signal` (`0.75` is the ITU-R BT.2408 HDR reference white) and brighter values roll off exponentially toward + the HLG peak, with a slope that is continuous at diffuse white. The ARIB STD-B67 OETF is then applied, followed by + a conversion to 10-bit limited-range BT.2020 non-constant-luminance Y'CbCr 4:2:0. The result is encoded with + `libx265` (`yuv420p10le`, `hvc1` tag) and tagged as BT.2020 primaries, ARIB STD-B67 transfer, BT.2020 NCL matrix + and limited range. No mastering-display or content-light-level metadata is written. + + Requires a PyAV build whose FFmpeg includes `libx265`. + + Args: + frames (`torch.Tensor` or `np.ndarray`): + Scene-linear HDR frames of shape `(F, H, W, 3)` with values in `[0, ∞)`, for example the output of + [`LTX2VideoHDRProcessor.postprocess_hdr_video`] for a single video. `H` and `W` must be even. + output_mp4 (`str` or `pathlib.Path`): + Output MP4 path. + frame_rate (`float`): + Frame rate for the output video. + primaries (`str`, *optional*, defaults to `"rec709"`): + Primaries of `frames`: `"rec709"` or `"acescg"`. + white_signal (`float`, *optional*, defaults to `0.75`): + HLG signal level that diffuse white (linear `1.0`) is mapped to. + rolloff_k (`float`, *optional*): + Exponential highlight roll-off rate. Defaults to `white_x / (1 - white_x)`, where `white_x` is the + scene-linear value of `white_signal`, which makes the mapping C1-continuous at linear `1.0`. + crf (`int`, *optional*, defaults to `12`): + libx265 CRF quality factor. Lower values produce higher quality. + preset (`str`, *optional*, defaults to `"ultrafast"`): + libx265 preset. + thread_count (`int`, *optional*, defaults to `0`): + libx265 thread pool size. `0` uses the number of CPUs, capped at 16. + device (`str` or `torch.device`, *optional*): + Device for the colour conversion. Defaults to the device of `frames` (CPU for NumPy input). + """ + if "libx265" not in av.codecs_available: + raise RuntimeError( + "HLG export requires the `libx265` encoder, but the FFmpeg build used by PyAV does not include it." + ) + if primaries not in _HLG_PRIMARIES_TO_REC2020: + raise ValueError(f"Unsupported primaries {primaries!r}. Expected 'rec709' or 'acescg'.") + if not 0.0 < white_signal < 1.0: + raise ValueError(f"`white_signal` must be in (0, 1), got {white_signal}.") + + frames = torch.as_tensor(frames) if isinstance(frames, np.ndarray) else frames.detach() + if frames.ndim != 4 or frames.shape[-1] != 3: + raise ValueError(f"Expected `frames` of shape (F, H, W, 3), got {tuple(frames.shape)}.") + num_frames, height, width, _ = frames.shape + if num_frames == 0: + raise ValueError("No HDR frames to encode.") + if height % 2 != 0 or width % 2 != 0: + raise ValueError(f"HLG export requires an even frame height and width, got {height}x{width}.") + + device = frames.device if device is None else torch.device(device) + primaries_matrix = torch.as_tensor(_HLG_PRIMARIES_TO_REC2020[primaries], dtype=torch.float32).to(device) + white_x = _hlg_inverse_oetf(white_signal) + roll_k = rolloff_k if rolloff_k is not None else white_x / (1.0 - white_x) + threads = thread_count if thread_count > 0 else max(1, min(os.cpu_count() or 8, 16)) + + output_mp4 = Path(output_mp4) + container = av.open(str(output_mp4), mode="w", options={"movflags": "+faststart"}) + try: + stream = container.add_stream("libx265", rate=Fraction(frame_rate).limit_denominator(1000)) + stream.width = width + stream.height = height + stream.pix_fmt = "yuv420p10le" + stream.codec_tag = "hvc1" + stream.options = {"crf": str(crf), "preset": preset, "x265-params": _x265_params(threads, width, height)} + codec_context = stream.codec_context + codec_context.thread_count = threads + codec_context.thread_type = "FRAME" + codec_context.color_primaries = _AV_COLOR_PRIMARIES_BT2020 + codec_context.color_trc = _AV_COLOR_TRC_ARIB_STD_B67 + codec_context.colorspace = _AV_COLORSPACE_BT2020_NCL + codec_context.color_range = _AV_COLOR_RANGE_MPEG + + for index in range(num_frames): + rgb = frames[index : index + 1].to(device=device, dtype=torch.float32).movedim(-1, -3) + hlg = _linear_to_hlg_signal(rgb, primaries_matrix, white_x, roll_k) + planes_yuv = [plane[0].cpu().numpy().astype(np.uint16) for plane in _rgb_to_yuv420p10_bt2020_limited(hlg)] + + frame = av.VideoFrame(width, height, "yuv420p10le") + for plane, src in zip(frame.planes, planes_yuv): + dest = np.frombuffer(plane, dtype=np.uint16).reshape(plane.height, plane.line_size // 2) + dest[:, : src.shape[1]] = src + frame.colorspace = _AV_COLORSPACE_BT2020_NCL + frame.color_range = _AV_COLOR_RANGE_MPEG + for packet in stream.encode(frame): + container.mux(packet) + + for packet in stream.encode(): + container.mux(packet) + except BaseException: + container.close() + output_mp4.unlink(missing_ok=True) + raise + container.close() + + +# OpenEXR `chromaticities` (R, G, B and white point xy) and `colorSpace` tags of the reference EXR writer, per EXR +# colour space (`ltx_pipelines` `EXRColorSpace`). ACEScct frames are tagged with AP1 chromaticities. +_EXR_CHROMATICITIES = { + "rec709": (0.64, 0.33, 0.30, 0.60, 0.15, 0.06, 0.3127, 0.3290), + "acescg": (0.713, 0.293, 0.165, 0.830, 0.128, 0.044, 0.32168, 0.33767), +} +_EXR_COLORSPACES = { + "srgb_linear": ("rec709", "sRGB"), + "acescg": ("acescg", "ACEScg"), + "acescct": ("acescg", "ACEScct"), +} + + +def _import_openexr(): + if not is_openexr_available(): + raise ImportError("OpenEXR is required to write EXR frames. You can install it with `pip install OpenEXR`.") + import OpenEXR + + # `OpenEXR.File` was added in OpenEXR 3.3; older bindings only expose the legacy `OutputFile` API. + if not hasattr(OpenEXR, "File"): + raise ImportError( + "OpenEXR>=3.3 is required to write EXR frames. You can upgrade it with `pip install -U OpenEXR`." + ) + return OpenEXR + + +def save_exr_frame( + frame: torch.Tensor | np.ndarray, + output_exr: str | Path, + primaries: str = "rec709", + color_space: str = "sRGB", + half: bool = True, +) -> None: + r""" + Saves a single RGB frame as an OpenEXR file tagged with its colour space, following the LTX-2 reference EXR writer: + a scanline image with `R`, `G` and `B` channels, ZIP compression, and `chromaticities` and `colorSpace` header + attributes. + + Requires the `OpenEXR` package (`pip install OpenEXR`, version 3.3 or later). + + Args: + frame (`torch.Tensor` or `np.ndarray`): + Float frame of shape `(H, W, 3)` or `(3, H, W)`. + output_exr (`str` or `pathlib.Path`): + Output EXR path. + primaries (`str`, *optional*, defaults to `"rec709"`): + Colour primaries written to the `chromaticities` attribute: `"rec709"` or `"acescg"` (AP1). + color_space (`str`, *optional*, defaults to `"sRGB"`): + Value of the `colorSpace` string attribute, which describes the encoding (e.g. `"sRGB"` or `"ACEScg"` for + scene-linear values, `"ACEScct"` for log codes). + half (`bool`, *optional*, defaults to `True`): + Write 16-bit half floats. When `False`, 32-bit floats are written, unless `frame` is already float16. + """ + OpenEXR = _import_openexr() + if primaries not in _EXR_CHROMATICITIES: + raise ValueError(f"Unsupported primaries {primaries!r}. Expected 'rec709' or 'acescg'.") + + if isinstance(frame, torch.Tensor): + use_half = half or frame.dtype == torch.float16 + frame = frame.detach().cpu().float().numpy() + else: + use_half = half or frame.dtype == np.float16 + frame = np.asarray(frame, dtype=np.float32) + if frame.ndim == 3 and frame.shape[0] == 3: + frame = frame.transpose(1, 2, 0) + if frame.ndim != 3 or frame.shape[-1] != 3: + raise ValueError(f"Expected `frame` of shape (H, W, 3) or (3, H, W), got {frame.shape}.") + frame = frame.astype(np.float16 if use_half else np.float32) + + header = { + "type": OpenEXR.scanlineimage, + "compression": OpenEXR.ZIP_COMPRESSION, + "chromaticities": _EXR_CHROMATICITIES[primaries], + "colorSpace": color_space, + } + channels = {name: np.ascontiguousarray(frame[..., index]) for index, name in enumerate("RGB")} + with OpenEXR.File(header, channels) as exr_file: + exr_file.write(str(output_exr)) + + +def export_to_exr_sequence( + frames: torch.Tensor | np.ndarray, + output_dir: str | Path, + exr_colorspace: str = "acescg", + half: bool = True, +) -> list[str]: + r""" + Saves HDR frames as a directory of OpenEXR files named `frame_00000.exr`, `frame_00001.exr`, ..., following the + LTX-2 reference HDR export. Each file is written with [`~pipelines.ltx2.export_utils.save_exr_frame`]. + + Requires the `OpenEXR` package (`pip install OpenEXR`, version 3.3 or later). + + Args: + frames (`torch.Tensor` or `np.ndarray`): + Frames of shape `(F, H, W, 3)` in the colour space given by `exr_colorspace`, for example the output of + [`LTX2VideoHDRProcessor.postprocess_hdr_video`] for a single video with the matching `output_colorspace`. + output_dir (`str` or `pathlib.Path`): + Output directory. It is created if it does not exist. + exr_colorspace (`str`, *optional*, defaults to `"acescg"`): + Colour space of `frames`, which sets the EXR tags: `"acescg"` (scene-linear ACEScg, AP1 chromaticities, + `colorSpace="ACEScg"`), `"srgb_linear"` (scene-linear Rec.709, Rec.709 chromaticities, `colorSpace="sRGB"`) + or `"acescct"` (ACEScct codes, AP1 chromaticities, `colorSpace="ACEScct"`). + half (`bool`, *optional*, defaults to `True`): + Write 16-bit half floats. + + Returns: + `list[str]`: Paths of the written EXR files. + """ + _import_openexr() + if exr_colorspace not in _EXR_COLORSPACES: + raise ValueError(f"Unsupported EXR colorspace {exr_colorspace!r}. Expected one of {tuple(_EXR_COLORSPACES)}.") + if frames.ndim != 4 or frames.shape[-1] != 3: + raise ValueError(f"Expected `frames` of shape (F, H, W, 3), got {tuple(frames.shape)}.") + primaries, color_space = _EXR_COLORSPACES[exr_colorspace] + + output_dir = Path(output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + paths = [] + for index, frame in enumerate(frames): + path = output_dir / f"frame_{index:05d}.exr" + save_exr_frame(frame, path, primaries=primaries, color_space=color_space, half=half) + paths.append(str(path)) + return paths diff --git a/src/diffusers/pipelines/ltx2/image_processor.py b/src/diffusers/pipelines/ltx2/image_processor.py index a25660073943..feef8cfb6d20 100644 --- a/src/diffusers/pipelines/ltx2/image_processor.py +++ b/src/diffusers/pipelines/ltx2/image_processor.py @@ -13,10 +13,12 @@ # limitations under the License. import numpy as np +import PIL.Image import torch import torch.nn.functional as F from ...configuration_utils import register_to_config +from ...image_processor import is_valid_image, is_valid_image_imagelist from ...utils import logging from ...video_processor import VideoProcessor @@ -24,6 +26,194 @@ logger = logging.get_logger(__name__) # pylint: disable=invalid-name +# ACEScct constants (Academy S-2016-001), ported from `ltx_core.hdr` in the Lightricks LTX-2 reference. +ACESCCT_A = 10.5402377416545 +ACESCCT_B = 0.0729055341958355 +ACESCCT_X_BRK = 0.0078125 +ACESCCT_Y_BRK = 0.155251141552511 +ACESCCT_LOG_M = 17.52 +ACESCCT_LOG_B = 9.72 + +# IEC 61966-2-1 sRGB EOTF constants. +_SRGB_A = 0.055 +_SRGB_LINEAR_THRESHOLD = 0.04045 +_SRGB_LINEAR_SLOPE = 12.92 +_SRGB_GAMMA = 2.4 + +# Linear RGB -> RGB primaries matrices (Bradford chromatic adaptation), applied as `out[d] = sum_c M[d][c] * in[c]`. +# The LTX-2 reference (`ltx_core.color.primaries`) computes them at import time with colour-science +# (`colour.matrix_RGB_to_RGB(..., chromatic_adaptation_transform="Bradford")`) and stores them as float32; +# `REC709_TO_ACESCG` is the float64 inverse of the float32 `ACESCG_TO_REC709`. The values below are those float32 +# matrices (colour-science 0.4.4), written out in full so that no colour-science dependency is needed. +ACESCG_TO_REC709 = ( + (1.7050509452819824, -0.6217921376228333, -0.08325887471437454), + (-0.13025641441345215, 1.1408047676086426, -0.010548318736255169), + (-0.02400335669517517, -0.1289689689874649, 1.1529723405838013), +) +REC709_TO_ACESCG = ( + (0.6130974292755127, 0.33952316641807556, 0.04737945273518562), + (0.07019372284412384, 0.9163538813591003, 0.013452397659420967), + (0.0206155925989151, 0.10956976562738419, 0.8698146343231201), +) +REC709_TO_REC2020 = ( + (0.6274039149284363, 0.3292830288410187, 0.04331306740641594), + (0.06909728795289993, 0.9195404052734375, 0.011362315155565739), + (0.016391439363360405, 0.08801330626010895, 0.8955952525138855), +) +ACESCG_TO_REC2020 = ( + (1.025824785232544, -0.020053191110491753, -0.005771556869149208), + (-0.0022343695163726807, 1.0045864582061768, -0.002352132461965084), + (-0.0050133513286709785, -0.025290071964263916, 1.0303034782409668), +) + +# Input colour spaces accepted by the ACEScct input transform (`ltx_pipelines` `HDRICLoraInputColorSpace`). +ACESCCT_INPUT_COLORSPACES = ("srgb_gamma", "srgb", "acescg", "acescct") +# Output colour spaces of the ACEScct output transform: scene-linear ACEScg (AP1), scene-linear Rec.709, or the raw +# ACEScct codes. +ACESCCT_OUTPUT_COLORSPACES = ("rec709", "acescg", "acescct") + + +def apply_primaries_matrix(video: torch.Tensor, matrix) -> torch.Tensor: + r""" + Apply a 3x3 linear primaries matrix to an RGB video or image tensor. + + The channel axis is `dim=1` for 5D `(B, C, F, H, W)` inputs and `dim=-3` otherwise (`(..., C, H, W)`), matching the + reference implementation. + + Args: + video (`torch.Tensor`): + Linear RGB tensor of shape `(B, 3, F, H, W)` or `(..., 3, H, W)`. + matrix (`torch.Tensor` or nested `tuple` of `float`): + The 3x3 matrix `M`, applied as `out[d] = sum_c M[d][c] * in[c]`. + + Returns: + `torch.Tensor`: The converted tensor, with the same shape, device and dtype as `video`. + """ + matrix = torch.as_tensor(matrix, dtype=torch.float32).to(device=video.device, dtype=video.dtype) + if video.ndim == 5: + return torch.einsum("bcfhw,dc->bdfhw", video, matrix) + return torch.einsum("...chw,dc->...dhw", video, matrix) + + +def srgb_eotf_to_linear(srgb: torch.Tensor) -> torch.Tensor: + r""" + sRGB-encoded `[0, 1]` code values to display-linear Rec.709 light (IEC 61966-2-1 EOTF). + + Inputs are cast to float32 and clamped to `[0, 1]` first. + + Args: + srgb (`torch.Tensor`): sRGB-encoded values. + + Returns: + `torch.Tensor`: Linear Rec.709 values in `[0, 1]`, as float32. + """ + x = torch.clamp(srgb.float(), 0.0, 1.0) + return torch.where( + x <= _SRGB_LINEAR_THRESHOLD, + x / _SRGB_LINEAR_SLOPE, + torch.pow((x + _SRGB_A) / (1.0 + _SRGB_A), _SRGB_GAMMA), + ) + + +def acescct_encode(linear_acescg: torch.Tensor) -> torch.Tensor: + r""" + Encode scene-linear ACEScg (AP1) values to ACEScct `[0, 1]`. + + Follows the reference implementation, which differs from the ACEScct specification in two places: negative inputs + are clamped to `0` before encoding, and the output is clamped to `[0, 1]` (so linear values above `2 ** (17.52 - + 9.72) ~= 222.86` are clipped). + + Args: + linear_acescg (`torch.Tensor`): Scene-linear ACEScg values. + + Returns: + `torch.Tensor`: ACEScct codes in `[0, 1]`. + """ + x = torch.clamp(linear_acescg, min=0.0) + log_part = (torch.log2(torch.clamp(x, min=1e-12)) + ACESCCT_LOG_B) / ACESCCT_LOG_M + lin_part = ACESCCT_A * x + ACESCCT_B + return torch.clamp(torch.where(x > ACESCCT_X_BRK, log_part, lin_part), 0.0, 1.0) + + +def acescct_decode(acescct: torch.Tensor) -> torch.Tensor: + r""" + Decode ACEScct codes to scene-linear ACEScg (AP1) values. + + The input is clamped to `[0, 1]` first. The linear toe is kept as is, so a code of `0` decodes to a small negative + value (`-0.0069169`), as in the reference implementation. + + Args: + acescct (`torch.Tensor`): ACEScct codes. + + Returns: + `torch.Tensor`: Scene-linear ACEScg values. + """ + ct = torch.clamp(acescct, 0.0, 1.0) + lin_from_log = torch.pow(2.0, ct * ACESCCT_LOG_M - ACESCCT_LOG_B) + lin_from_lin = (ct - ACESCCT_B) / ACESCCT_A + return torch.where(ct > ACESCCT_Y_BRK, lin_from_log, lin_from_lin) + + +def to_acescct(video: torch.Tensor, input_colorspace: str = "srgb_gamma") -> torch.Tensor: + r""" + Input transform of the LTX-2.5 SDR-To-HDR IC-LoRA: map RGB values to the ACEScct `[0, 1]` working space. + + The supported input colour spaces are those of the reference implementation: + + - `"srgb_gamma"`: sRGB-encoded Rec.709 in `[0, 1]` (e.g. an 8-bit video divided by 255). The sRGB EOTF is applied, + then the Rec.709 -> AP1 matrix, then the ACEScct encoding. + - `"srgb"`: scene-linear Rec.709 (e.g. a linear EXR plate). Same as `"srgb_gamma"` without the EOTF. Note that the + reference also applies this mode to 8-bit video divided by 255, i.e. without linearizing it. + - `"acescg"`: scene-linear ACEScg (AP1). Only the ACEScct encoding is applied. + - `"acescct"`: values that are already ACEScct codes. They are only clamped to `[0, 1]`. + + Args: + video (`torch.Tensor`): + RGB tensor of shape `(B, 3, F, H, W)` or `(..., 3, H, W)`. + input_colorspace (`str`, *optional*, defaults to `"srgb_gamma"`): + One of `"srgb_gamma"`, `"srgb"`, `"acescg"` or `"acescct"`. + + Returns: + `torch.Tensor`: ACEScct codes in `[0, 1]`, as float32, with the same shape as `video`. + """ + if input_colorspace not in ACESCCT_INPUT_COLORSPACES: + raise ValueError( + f"Unsupported input colorspace {input_colorspace!r}. Expected one of {ACESCCT_INPUT_COLORSPACES}." + ) + video = video.float() + if input_colorspace == "acescct": + return video.clamp(0.0, 1.0) + if input_colorspace == "srgb_gamma": + video = srgb_eotf_to_linear(video) + if input_colorspace in ("srgb_gamma", "srgb"): + video = apply_primaries_matrix(video, REC709_TO_ACESCG) + return acescct_encode(video.clamp(min=0.0)) + + +def acescct_to_linear(acescct: torch.Tensor, output_colorspace: str = "rec709") -> torch.Tensor: + r""" + Output transform of the LTX-2.5 SDR-To-HDR IC-LoRA: map ACEScct codes to scene-linear HDR. + + The codes are decoded to linear ACEScg, converted to the requested primaries, then clamped to `>= 0`. The clamp + happens after the primaries matrix, as in the reference implementation, so out-of-gamut Rec.709 values are clipped. + + Args: + acescct (`torch.Tensor`): + ACEScct codes of shape `(B, 3, F, H, W)` or `(..., 3, H, W)`. + output_colorspace (`str`, *optional*, defaults to `"rec709"`): + `"rec709"` for scene-linear Rec.709 primaries, or `"acescg"` for scene-linear ACEScg (AP1) primaries. + + Returns: + `torch.Tensor`: Scene-linear HDR values in `[0, inf)`, as float32. + """ + if output_colorspace not in ("rec709", "acescg"): + raise ValueError(f"Unsupported output colorspace {output_colorspace!r}. Expected 'rec709' or 'acescg'.") + linear_acescg = acescct_decode(acescct.float()) + if output_colorspace == "rec709": + linear_acescg = apply_primaries_matrix(linear_acescg, ACESCG_TO_REC709) + return linear_acescg.clamp(min=0.0) + + class LTX2VideoHDRProcessor(VideoProcessor): r""" Video processor for the LTX-2 HDR IC-LoRA pipeline. @@ -37,13 +227,19 @@ class LTX2VideoHDRProcessor(VideoProcessor): - `postprocess_hdr_video`: applies the LogC3 inverse transform to the VAE's decoded output, mapping `[0, 1]` → linear HDR `[0, ∞)`. + With `hdr_transform="acescct"` (LTX-2.5 SDR-To-HDR IC-LoRA), `preprocess_reference_video_hdr` additionally applies + an input transform to the ACEScct working space (see [`~pipelines.ltx2.image_processor.to_acescct`]) and + `postprocess_hdr_video` decodes ACEScct to scene-linear HDR (see + [`~pipelines.ltx2.image_processor.acescct_to_linear`]). + Args: vae_scale_factor (`int`, *optional*, defaults to `32`): VAE (spatial) scale factor for the LTX-2 video VAE. resample (`str`, *optional*, defaults to `"bilinear"`): Resampling filter used by the base [`VaeImageProcessor`] for PIL/tensor resizing. hdr_transform (`str`, *optional*, defaults to `"logc3"`): - HDR transform identifier. Only `"logc3"` (ARRI EI 800) is currently supported. + HDR transform identifier. `"logc3"` (ARRI LogC3 EI 800, LTX-2.3 HDR) or `"acescct"` (ACEScct, LTX-2.5 + SDR-To-HDR). """ # LogC3 (ARRI EI 800) coefficients, ported from `ltx_core.hdr.LogC3`. @@ -67,8 +263,8 @@ def __init__( vae_scale_factor=vae_scale_factor, resample=resample, ) - if hdr_transform != "logc3": - raise ValueError(f"Unsupported HDR transform {hdr_transform!r}. Only 'logc3' is supported.") + if hdr_transform not in ("logc3", "acescct"): + raise ValueError(f"Unsupported HDR transform {hdr_transform!r}. Expected 'logc3' or 'acescct'.") @classmethod def _logc3_decompress(cls, logc: torch.Tensor) -> torch.Tensor: @@ -115,32 +311,105 @@ def _resize_and_reflect_pad_video(video: torch.Tensor, height: int, width: int) return video + @staticmethod + def _video_to_float_tensor(video) -> tuple[torch.Tensor, bool]: + r""" + Convert a video input to a float32 `(B, C, F, H, W)` tensor without resizing or normalizing it. + + Accepts the layouts of [`VideoProcessor.preprocess_video`]. Integer inputs (PIL images, `uint8` arrays or + tensors) are divided by 255; floating point inputs are kept as is, so scene-linear values above 1 survive. + + Returns: + `tuple[torch.Tensor, bool]`: The video tensor, and whether the input was integer-valued. + """ + if isinstance(video, (np.ndarray, torch.Tensor)) and video.ndim == 5: + videos = list(video) + elif isinstance(video, list) and is_valid_image(video[0]) or is_valid_image_imagelist(video): + videos = [video] + elif isinstance(video, list) and is_valid_image_imagelist(video[0]): + videos = video + else: + raise ValueError( + "Input is in incorrect format. Currently, we only support numpy.ndarray, torch.Tensor, PIL.Image.Image" + ) + + tensors = [] + is_integer = True + for frames in videos: + if isinstance(frames, list) and isinstance(frames[0], PIL.Image.Image): + frames = np.stack([np.array(frame.convert("RGB")) for frame in frames], axis=0) + elif isinstance(frames, list) and isinstance(frames[0], np.ndarray): + frames = np.stack(frames, axis=0) + elif isinstance(frames, list) and isinstance(frames[0], torch.Tensor): + frames = torch.stack(frames, dim=0) + if isinstance(frames, np.ndarray): + # NumPy frames are channels-last: (F, H, W, C) -> (F, C, H, W). + frames = torch.from_numpy(np.ascontiguousarray(frames)).permute(0, 3, 1, 2) + if frames.is_floating_point(): + is_integer = False + frames = frames.float() + else: + frames = frames.float() / 255.0 + tensors.append(frames) + # (B, F, C, H, W) -> (B, C, F, H, W) + return torch.stack(tensors, dim=0).permute(0, 2, 1, 3, 4), is_integer + def preprocess_reference_video_hdr( self, video, height: int, width: int, + input_colorspace: str | None = None, ) -> torch.Tensor: r""" Preprocess a reference (SDR) video for HDR IC-LoRA conditioning. - Runs the input through the standard video preprocessing (normalization to `[-1, 1]`) without resizing, then - applies reflect-pad resize to the target dimensions. For LDR inputs this is numerically equivalent to - `load_video_conditioning_hdr` in the reference implementation (since `LogC3.compress_ldr` is an identity clamp - on `[0, 1]` inputs). + With `hdr_transform="logc3"`, runs the input through the standard video preprocessing (normalization to `[-1, + 1]`) without resizing, then applies reflect-pad resize to the target dimensions. For LDR inputs this is + numerically equivalent to `load_video_conditioning_hdr` in the reference implementation (since + `LogC3.compress_ldr` is an identity clamp on `[0, 1]` inputs). + + With `hdr_transform="acescct"`, the input is mapped to ACEScct `[0, 1]` with + [`~pipelines.ltx2.image_processor.to_acescct`], reflect-pad resized, then mapped to `[-1, 1]`. Integer inputs + (PIL images, `uint8` arrays or tensors) are divided by 255 and transformed before resizing, as the reference + does for MP4/MOV inputs; floating point inputs are used as is and resized before the transform, as the + reference does for EXR inputs. The two orders only differ when the video is downscaled. Args: video: Input accepted by `VideoProcessor.preprocess_video` (list of PIL images, 4D/5D tensor/array, etc.). height (`int`), width (`int`): Target spatial dimensions. + input_colorspace (`str`, *optional*): + Colour space of `video` for `hdr_transform="acescct"`: `"srgb_gamma"` (default), `"srgb"`, `"acescg"` + or `"acescct"`. See [`~pipelines.ltx2.image_processor.to_acescct`]. Must be `None` for + `hdr_transform="logc3"`. Returns: `torch.Tensor`: Preprocessed video of shape `(B, C, F, height, width)` with values in `[-1, 1]`. """ - video = self.preprocess_video(video, height=None, width=None) # (B, C, F, src_h, src_w) in [-1, 1] - video = self._resize_and_reflect_pad_video(video, height, width) - return video + if self.config.hdr_transform == "logc3": + if input_colorspace is not None: + raise ValueError("`input_colorspace` is only supported with `hdr_transform='acescct'`.") + video = self.preprocess_video(video, height=None, width=None) # (B, C, F, src_h, src_w) in [-1, 1] + video = self._resize_and_reflect_pad_video(video, height, width) + return video - def postprocess_hdr_video(self, video: torch.Tensor, output_type: str = "np") -> torch.Tensor | np.ndarray: + input_colorspace = "srgb_gamma" if input_colorspace is None else input_colorspace + if input_colorspace not in ACESCCT_INPUT_COLORSPACES: + raise ValueError( + f"Unsupported input colorspace {input_colorspace!r}. Expected one of {ACESCCT_INPUT_COLORSPACES}." + ) + video, is_integer = self._video_to_float_tensor(video) + if is_integer: + video = to_acescct(video, input_colorspace) + video = self._resize_and_reflect_pad_video(video, height, width) + else: + video = self._resize_and_reflect_pad_video(video, height, width) + video = to_acescct(video, input_colorspace) + return video * 2.0 - 1.0 + + def postprocess_hdr_video( + self, video: torch.Tensor, output_type: str = "np", output_colorspace: str | None = None + ) -> torch.Tensor | np.ndarray: r""" Postprocess the VAE's decoded output to linear HDR. @@ -149,9 +418,15 @@ def postprocess_hdr_video(self, video: torch.Tensor, output_type: str = "np") -> VAE decoded output in VAE range `[-1, 1]`, shape `(B, C, F, H, W)`. output_type (`str`, *optional*, defaults to `"np"`): Output type of post-processed video tensor; should be in `["np", "pt"]`. + output_colorspace (`str`, *optional*): + Output colour space for `hdr_transform="acescct"`: `"rec709"` (default, scene-linear Rec.709), + `"acescg"` (scene-linear ACEScg) or `"acescct"` (the decoded ACEScct codes in `[0, 1]`, without + decoding them to linear). See [`~pipelines.ltx2.image_processor.acescct_to_linear`]. Must be `None` for + `hdr_transform="logc3"`, whose output is linear in the primaries of the input. Returns: - Returns linear HDR video with values in `[0, ∞)`, depending on `output_type`: + Returns linear HDR video with values in `[0, ∞)` (or ACEScct codes in `[0, 1]` with + `output_colorspace="acescct"`), depending on `output_type`: - `output_type="pt"`: `torch.Tensor` with shape `(B, F, H, W, C)` and dtype `float32`. - `output_type="np"`: `np.ndarray` with shape `(B, F, H, W, C)` and dtype `float32`. """ @@ -163,8 +438,20 @@ def postprocess_hdr_video(self, video: torch.Tensor, output_type: str = "np") -> output_type = "np" video = self.denormalize(video.float()) - # Apply the inverse transform function to get linear HDR light - video = self._logc3_decompress(video) + if self.config.hdr_transform == "logc3": + if output_colorspace is not None: + raise ValueError("`output_colorspace` is only supported with `hdr_transform='acescct'`.") + # Apply the inverse transform function to get linear HDR light + video = self._logc3_decompress(video) + else: + output_colorspace = "rec709" if output_colorspace is None else output_colorspace + if output_colorspace not in ACESCCT_OUTPUT_COLORSPACES: + raise ValueError( + f"Unsupported output colorspace {output_colorspace!r}. Expected one of " + f"{ACESCCT_OUTPUT_COLORSPACES}." + ) + if output_colorspace != "acescct": + video = acescct_to_linear(video, output_colorspace) # Permute to channels-last: [B, C, F, H, W] --> [B, F, H, W, C] video = video = video.permute(0, 2, 3, 4, 1).contiguous() diff --git a/src/diffusers/utils/__init__.py b/src/diffusers/utils/__init__.py index b3051dfcb9d1..6efe38dea81e 100644 --- a/src/diffusers/utils/__init__.py +++ b/src/diffusers/utils/__init__.py @@ -97,6 +97,7 @@ is_nvidia_modelopt_version, is_onnx_available, is_opencv_available, + is_openexr_available, is_optimum_quanto_available, is_optimum_quanto_version, is_outlines_available, diff --git a/src/diffusers/utils/import_utils.py b/src/diffusers/utils/import_utils.py index d2cf394cd9a7..2bf0e06e3279 100644 --- a/src/diffusers/utils/import_utils.py +++ b/src/diffusers/utils/import_utils.py @@ -220,6 +220,7 @@ def _is_package_available(pkg_name: str, get_dist_name: bool = False) -> tuple[b _sdnq_available, _sdnq_version = _is_package_available("sdnq") _flashpack_available, _flashpack_version = _is_package_available("flashpack") _av_available, _av_version = _is_package_available("av") +_openexr_available, _openexr_version = _is_package_available("OpenEXR") def is_torch_available(): @@ -422,6 +423,10 @@ def is_av_available(): return _av_available +def is_openexr_available(): + return _openexr_available + + # docstyle-ignore INFLECT_IMPORT_ERROR = """ {0} requires the inflect library but it was not found in your environment. You can install it with pip: `pip install diff --git a/tests/pipelines/ltx2/test_ltx2_hdr_color.py b/tests/pipelines/ltx2/test_ltx2_hdr_color.py new file mode 100644 index 000000000000..225f1c0b7d87 --- /dev/null +++ b/tests/pipelines/ltx2/test_ltx2_hdr_color.py @@ -0,0 +1,470 @@ +# Copyright 2026 The HuggingFace Team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Colour transforms and HDR export utilities of the LTX-2.5 SDR-To-HDR IC-LoRA. + +Expected values were computed independently in float64 with colour-science 0.4.4 (`log_encoding_ACEScct`, +`log_decoding_ACEScct`, `cctf_decoding(..., function="sRGB")`, `matrix_RGB_to_RGB(..., "Bradford")`, +`oetf_BT2100_HLG` / `oetf_inverse_BT2100_HLG`), not with the code under test. +""" + +import numpy as np +import PIL.Image +import pytest +import torch + +from diffusers.pipelines.ltx2.image_processor import ( + ACESCCT_A, + ACESCCT_B, + ACESCG_TO_REC709, + ACESCG_TO_REC2020, + REC709_TO_ACESCG, + REC709_TO_REC2020, + LTX2VideoHDRProcessor, + acescct_decode, + acescct_encode, + acescct_to_linear, + apply_primaries_matrix, + srgb_eotf_to_linear, + to_acescct, +) +from diffusers.utils import is_av_available, is_openexr_available + + +# colour-science 0.4.4, float64: `colour.matrix_RGB_to_RGB(src, dst, chromatic_adaptation_transform="Bradford")`. +COLOUR_ACESCG_TO_REC709 = [ + [1.7050509926579835, -0.6217921206570048, -0.08325887200097871], + [-0.13025641750704345, 1.1408047365754013, -0.010548319068358038], + [-0.024003356804618025, -0.12896897606497054, 1.1529723328695887], +] +COLOUR_REC709_TO_ACESCG = [ + [0.6130974024, 0.3395231462, 0.0473794514], + [0.0701937225, 0.9163538791, 0.0134523985], + [0.0206155929, 0.1095697729, 0.8698146342], +] +COLOUR_REC709_TO_REC2020 = [ + [0.6274038959, 0.3292830384, 0.0433130657], + [0.0690972894, 0.9195403951, 0.0113623156], + [0.0163914389, 0.0880133079, 0.8955952532], +] +COLOUR_ACESCG_TO_REC2020 = [ + [1.0258247477, -0.0200531908, -0.0057715568], + [-0.0022343695, 1.0045865019, -0.0023521324], + [-0.0050133515, -0.0252900718, 1.0303034233], +] + + +def _logc3_decompress_reference(t: float) -> float: + # ARRI LogC3 EI 800 decoding, written out independently of the processor. + a, b, c, d, e, f, cut = 5.555556, 0.052272, 0.247190, 0.385537, 5.367655, 0.092809, 0.010591 + if t >= e * cut + f: + return (10.0 ** ((t - d) / c) - b) / a + return (t - f) / e + + +class TestACEScct: + def test_encode_known_values(self): + linear = torch.tensor([0.0, 0.0078125, 0.18, 1.0, 16.0, 222.0], dtype=torch.float64) + expected = torch.tensor( + [ + 0.0729055341958355, + 0.1552511415525113, + 0.4135884024924423, + 0.5547945205479452, + 0.7831050228310503, + 0.9996812709103942, + ], + dtype=torch.float64, + ) + torch.testing.assert_close(acescct_encode(linear), expected, rtol=0, atol=1e-12) + + def test_encode_clamps_like_the_reference(self): + # Negative inputs are clamped to 0 and the output is clamped to [0, 1] (the ACEScct spec does neither). + out = acescct_encode(torch.tensor([-1.0, 1000.0], dtype=torch.float64)) + torch.testing.assert_close(out, torch.tensor([ACESCCT_B, 1.0], dtype=torch.float64), rtol=0, atol=1e-12) + + def test_decode_known_values(self): + codes = torch.tensor( + [0.0729055341958355, 0.155251141552511, 0.4135884024924423, 1.0, 0.0, 0.5], dtype=torch.float64 + ) + expected = torch.tensor( + [0.0, 0.0078125, 0.18, 222.8609442038076, -0.006916877586898862, 0.5140569133280329], + dtype=torch.float64, + ) + torch.testing.assert_close(acescct_decode(codes), expected, rtol=1e-12, atol=1e-12) + + def test_decode_clamps_input(self): + out = acescct_decode(torch.tensor([-0.5, 1.5], dtype=torch.float64)) + expected = torch.tensor([-ACESCCT_B / ACESCCT_A, 222.8609442038076], dtype=torch.float64) + torch.testing.assert_close(out, expected, rtol=1e-12, atol=1e-12) + + def test_round_trip_linear(self): + linear = torch.cat([torch.linspace(0.0, 0.01, 101), torch.logspace(-2, np.log10(222.0), 200)]).double() + torch.testing.assert_close(acescct_decode(acescct_encode(linear)), linear, rtol=1e-12, atol=1e-15) + + def test_round_trip_codes(self): + # The valid code domain starts at the code of linear 0; codes below decode to negative linear values, + # which the encoder clamps. + codes = torch.linspace(ACESCCT_B, 1.0, 1001, dtype=torch.float64) + torch.testing.assert_close(acescct_encode(acescct_decode(codes)), codes, rtol=0, atol=1e-12) + + +class TestPrimaries: + @pytest.mark.parametrize( + "matrix, expected", + [ + (ACESCG_TO_REC709, COLOUR_ACESCG_TO_REC709), + (REC709_TO_ACESCG, COLOUR_REC709_TO_ACESCG), + (REC709_TO_REC2020, COLOUR_REC709_TO_REC2020), + (ACESCG_TO_REC2020, COLOUR_ACESCG_TO_REC2020), + ], + ) + def test_matrix_values(self, matrix, expected): + # The stored matrices are float32, as in the reference. + torch.testing.assert_close( + torch.tensor(matrix, dtype=torch.float64), torch.tensor(expected, dtype=torch.float64), rtol=0, atol=1e-7 + ) + + @pytest.mark.parametrize("matrix", [ACESCG_TO_REC709, REC709_TO_ACESCG, REC709_TO_REC2020, ACESCG_TO_REC2020]) + def test_white_is_preserved(self, matrix): + row_sums = torch.tensor(matrix, dtype=torch.float64).sum(dim=1) + torch.testing.assert_close(row_sums, torch.ones(3, dtype=torch.float64), rtol=0, atol=1e-6) + + def test_rec709_acescg_are_inverse(self): + product = torch.tensor(REC709_TO_ACESCG, dtype=torch.float64) @ torch.tensor( + ACESCG_TO_REC709, dtype=torch.float64 + ) + torch.testing.assert_close(product, torch.eye(3, dtype=torch.float64), rtol=0, atol=1e-6) + + def test_apply_primaries_matrix_layouts(self): + generator = torch.Generator().manual_seed(0) + video = torch.rand(2, 3, 4, 5, 6, generator=generator) + expected = torch.einsum("dc,bcfhw->bdfhw", torch.tensor(REC709_TO_ACESCG), video) + torch.testing.assert_close(apply_primaries_matrix(video, REC709_TO_ACESCG), expected) + # Non-5D inputs use `dim=-3` as the channel axis. + frames = video[0].permute(1, 0, 2, 3) # (F, C, H, W) + torch.testing.assert_close(apply_primaries_matrix(frames, REC709_TO_ACESCG), expected[0].permute(1, 0, 2, 3)) + + +class TestInputTransform: + # Neutral greys, a saturated red and an arbitrary colour, as (N, 3, 1, 1) images. + pixels = torch.tensor([[0, 0, 0], [0.04045] * 3, [0.5] * 3, [1, 1, 1], [1.0, 0.0, 0.0], [0.2, 0.6, 0.9]]) + + def _as_images(self, values): + return torch.as_tensor(values, dtype=torch.float32)[:, :, None, None] + + def test_srgb_eotf(self): + out = srgb_eotf_to_linear(torch.tensor([0.0, 0.04045, 0.5, 1.0, -0.2, 1.3])) + expected = torch.tensor([0.0, 0.0031308072830676845, 0.21404114048223255, 1.0, 0.0, 1.0]) + assert out.dtype == torch.float32 + torch.testing.assert_close(out, expected, rtol=0, atol=1e-7) + + def test_srgb_gamma(self): + expected = [ + [0.0729055341958355] * 3, + [0.10590498888079533, 0.10590498734413856, 0.1059049870368072], + [0.4278516036650092, 0.42785159983049315, 0.42785159906358994], + [0.5547945245358419, 0.5547945207013258, 0.5547945199344226], + [0.5145084623593209, 0.3360437118624257, 0.23515295334335573], + [0.4068006341802645, 0.45696458686011593, 0.5277994813544838], + ] + out = to_acescct(self._as_images(self.pixels), "srgb_gamma") + torch.testing.assert_close(out, self._as_images(expected), rtol=0, atol=1e-6) + + def test_srgb_linear(self): + expected = [ + [0.0729055341958355] * 3, + [0.2906554556363287, 0.2906554518018126, 0.2906554510349094], + [0.49771689896506566, 0.4977168951305496, 0.49771689436364636], + [0.5547945245358419, 0.5547945207013258, 0.5547945199344226], + [0.5145084623593209, 0.3360437118624257, 0.23515295334335573], + [0.47269375324230106, 0.509362790853311, 0.5416727756641857], + ] + out = to_acescct(self._as_images(self.pixels), "srgb") + torch.testing.assert_close(out, self._as_images(expected), rtol=0, atol=1e-6) + + def test_acescg_and_acescct(self): + linear = torch.tensor([0.0, 0.18, 1.0, 16.0, -1.0, 1000.0]).view(2, 3, 1, 1) + expected = torch.tensor( + [0.0729055341958355, 0.4135884024924423, 0.5547945205479452, 0.7831050228310503, ACESCCT_B, 1.0] + ).view(2, 3, 1, 1) + torch.testing.assert_close(to_acescct(linear, "acescg"), expected, rtol=0, atol=1e-6) + codes = torch.tensor([-0.5, 0.25, 0.5, 0.75, 1.0, 2.0]).view(2, 3, 1, 1) + torch.testing.assert_close(to_acescct(codes, "acescct"), codes.clamp(0.0, 1.0), rtol=0, atol=0) + + def test_invalid_colorspace(self): + with pytest.raises(ValueError, match="Unsupported input colorspace"): + to_acescct(torch.zeros(1, 3, 1, 1), "rec2020") + + +class TestOutputTransform: + codes = torch.tensor([[0.0, 0.0, 0.0], [0.5, 0.5, 0.5], [0.7, 0.4, 0.2], [1.0, 1.0, 1.0]])[:, :, None, None] + + def test_acescg(self): + expected = torch.tensor( + [ + [0.0, 0.0, 0.0], # the negative toe (-0.0069169) is clamped to 0 + [0.5140569133280329] * 3, + [5.83203751526631, 0.1526183140836417, 0.013452331067795364], + [222.8609442038076] * 3, + ] + )[:, :, None, None] + torch.testing.assert_close(acescct_to_linear(self.codes, "acescg"), expected, rtol=1e-6, atol=1e-7) + + def test_rec709(self): + expected = torch.tensor( + [ + [0.0, 0.0, 0.0], + [0.5140568788578308, 0.5140569310418868, 0.5140569209880779], + [9.847904184623355, 0.0, 0.0], # out-of-gamut values are clipped after the matrix + [222.8609292598168, 222.86095188335844, 222.86094752469444], + ] + )[:, :, None, None] + torch.testing.assert_close(acescct_to_linear(self.codes, "rec709"), expected, rtol=1e-6, atol=1e-7) + + def test_round_trip_through_working_space(self): + # sRGB code -> ACEScct -> Rec.709 linear equals the sRGB EOTF on in-gamut colours. + generator = torch.Generator().manual_seed(0) + srgb = torch.rand(1, 3, 2, 8, 8, generator=generator) + linear = acescct_to_linear(to_acescct(srgb, "srgb_gamma"), "rec709") + torch.testing.assert_close(linear, srgb_eotf_to_linear(srgb), rtol=1e-4, atol=1e-6) + + def test_invalid_colorspace(self): + with pytest.raises(ValueError, match="Unsupported output colorspace"): + acescct_to_linear(self.codes, "acescct") + + +class TestLTX2VideoHDRProcessor: + def test_logc3_is_the_default_and_unchanged(self): + processor = LTX2VideoHDRProcessor() + assert processor.config.hdr_transform == "logc3" + + # Preprocessing: plain [-1, 1] normalization of the sRGB codes, then reflect-padding. + video = (np.arange(2 * 32 * 64 * 3) % 251).astype(np.uint8).reshape(2, 32, 64, 3) + out = processor.preprocess_reference_video_hdr([PIL.Image.fromarray(frame) for frame in video], 32, 80) + expected = torch.from_numpy(video).permute(0, 3, 1, 2).float() / 255.0 * 2.0 - 1.0 # (F, C, H, W) + expected = torch.nn.functional.pad(expected, (0, 16, 0, 0), mode="reflect") + torch.testing.assert_close(out, expected.permute(1, 0, 2, 3)[None]) + + # Postprocessing: LogC3 decoding of the [0, 1] codes, no primaries conversion, channels-last. + codes = torch.tensor([0.0, 0.05, 0.391007, 0.6, 0.9, 1.0]) + decoded = processor.postprocess_hdr_video((codes * 2.0 - 1.0).view(1, 3, 2, 1, 1), output_type="pt") + expected = torch.tensor([_logc3_decompress_reference(float(t)) for t in codes]).view(1, 3, 2, 1, 1) + assert decoded.shape == (1, 2, 1, 1, 3) + torch.testing.assert_close(decoded, expected.permute(0, 2, 3, 4, 1), rtol=1e-5, atol=1e-6) + + def test_logc3_rejects_colorspace_options(self): + processor = LTX2VideoHDRProcessor() + with pytest.raises(ValueError, match="input_colorspace"): + processor.preprocess_reference_video_hdr(np.zeros((1, 4, 4, 3), dtype=np.uint8), 4, 4, "srgb") + with pytest.raises(ValueError, match="output_colorspace"): + processor.postprocess_hdr_video(torch.zeros(1, 3, 1, 4, 4), "pt", "rec709") + + def test_invalid_transform(self): + with pytest.raises(ValueError, match="Unsupported HDR transform"): + LTX2VideoHDRProcessor(hdr_transform="pq") + + def test_acescct_preprocess_srgb_gamma(self): + processor = LTX2VideoHDRProcessor(hdr_transform="acescct") + # Black and white 8-bit frames: sRGB 0 -> ACEScct 0.0729055 -> -0.8541889, sRGB 1 -> 0.5547945 -> 0.1095890. + video = np.zeros((2, 20, 24, 3), dtype=np.uint8) + video[1] = 255 + out = processor.preprocess_reference_video_hdr(video, 32, 32) + assert out.shape == (1, 3, 2, 32, 32) + torch.testing.assert_close(out[:, :, 0], torch.full((1, 3, 32, 32), -0.854188931608329), rtol=0, atol=1e-6) + torch.testing.assert_close(out[:, :, 1], torch.full((1, 3, 32, 32), 0.10958904109589), rtol=0, atol=1e-6) + + def test_acescct_preprocess_matches_functional_transform(self): + processor = LTX2VideoHDRProcessor(hdr_transform="acescct") + generator = torch.Generator().manual_seed(0) + # Integer input: transform, then reflect-pad. + frames = (torch.rand(3, 3, 20, 24, generator=generator) * 255).to(torch.uint8) # (F, C, H, W) + out = processor.preprocess_reference_video_hdr(frames, 32, 32, input_colorspace="srgb_gamma") + working = to_acescct(frames.permute(1, 0, 2, 3)[None].float() / 255.0, "srgb_gamma") + expected = processor._resize_and_reflect_pad_video(working, 32, 32) * 2.0 - 1.0 + torch.testing.assert_close(out, expected, rtol=0, atol=0) + # Float input (e.g. an EXR plate, values above 1): reflect-pad, then transform. + linear = torch.rand(1, 3, 20, 24, 3, generator=generator).numpy() * 50.0 # (B, F, H, W, C) + out = processor.preprocess_reference_video_hdr(linear, 32, 32, input_colorspace="acescg") + padded = processor._resize_and_reflect_pad_video(torch.from_numpy(linear).permute(0, 4, 1, 2, 3), 32, 32) + torch.testing.assert_close(out, to_acescct(padded, "acescg") * 2.0 - 1.0, rtol=0, atol=0) + assert out.min() >= -1.0 and out.max() <= 1.0 + + def test_acescct_postprocess(self): + processor = LTX2VideoHDRProcessor(hdr_transform="acescct") + codes = torch.tensor([[0.0, 0.0, 0.0], [0.5, 0.5, 0.5], [0.7, 0.4, 0.2], [1.0, 1.0, 1.0]]) + decoded = (codes * 2.0 - 1.0).T.reshape(1, 3, 4, 1, 1) # (B, C, F, H, W) in [-1, 1] + + out = processor.postprocess_hdr_video(decoded, output_type="pt") + assert out.shape == (1, 4, 1, 1, 3) + expected = acescct_to_linear(codes.T.reshape(1, 3, 4, 1, 1), "rec709").permute(0, 2, 3, 4, 1) + torch.testing.assert_close(out, expected) + + out = processor.postprocess_hdr_video(decoded, output_type="np", output_colorspace="acescg") + assert isinstance(out, np.ndarray) and out.dtype == np.float32 + np.testing.assert_allclose(out[0, 2, 0, 0], [5.83203751526631, 0.1526183140836417, 0.013452331067795364], 1e-6) + + out = processor.postprocess_hdr_video(decoded, output_type="pt", output_colorspace="acescct") + torch.testing.assert_close(out[0, :, 0, 0], codes, rtol=0, atol=1e-7) + + +class TestHLG: + @pytest.fixture(autouse=True) + def _export_utils(self): + if not is_av_available(): + pytest.skip("PyAV is not installed.") + from diffusers.pipelines.ltx2 import export_utils + + self.export_utils = export_utils + + def test_inverse_oetf_and_rolloff(self): + white_x = self.export_utils._hlg_inverse_oetf(0.75) + assert white_x == pytest.approx(0.26496256042100724, abs=1e-15) + assert white_x / (1.0 - white_x) == pytest.approx(0.36047491753994165, abs=1e-15) + assert self.export_utils._hlg_inverse_oetf(0.5) == pytest.approx(1.0 / 12.0, abs=1e-15) + assert self.export_utils._hlg_inverse_oetf(0.3) == pytest.approx(0.03, abs=1e-15) + + def test_signal_known_values(self): + white_x = self.export_utils._hlg_inverse_oetf(0.75) + identity = torch.eye(3) + # Neutral Rec.709 linear 0.18 / 1.0 / 4.0 / 100.0 (Rec.709 -> Rec.2020 keeps neutrals neutral). + linear = torch.tensor([0.18, 1.0, 4.0, 100.0]).view(4, 1, 1, 1).expand(4, 3, 1, 1) + signal = self.export_utils._linear_to_hlg_signal( + linear, torch.tensor(REC709_TO_REC2020), white_x, white_x / (1.0 - white_x) + ) + expected = torch.tensor([0.3782588830779046, 0.75, 0.947280748538359, 0.9999999950661305]) + torch.testing.assert_close(signal[:, 0, 0, 0], expected, rtol=0, atol=2e-6) + # A coloured pixel goes through the Rec.709 -> Rec.2020 matrix first. + signal = self.export_utils._linear_to_hlg_signal( + torch.tensor([2.0, 0.5, 0.05]).view(1, 3, 1, 1), + torch.tensor(REC709_TO_REC2020), + white_x, + white_x / (1.0 - white_x), + ) + expected = torch.tensor([0.8139141545892317, 0.6460072658555785, 0.3108599919641505]) + torch.testing.assert_close(signal[0, :, 0, 0], expected, rtol=0, atol=2e-6) + # Negative and NaN values are mapped to 0. + bad = torch.tensor([-1.0, float("nan"), float("-inf")]).view(1, 3, 1, 1) + signal = self.export_utils._linear_to_hlg_signal(bad, identity, white_x, 1.0) + assert torch.equal(signal, torch.zeros_like(signal)) + + def test_yuv_code_levels(self): + rgb = torch.zeros(2, 3, 2, 2) + rgb[1] = 1.0 + y, u, v = self.export_utils._rgb_to_yuv420p10_bt2020_limited(rgb) + assert y.shape == (2, 2, 2) and u.shape == (2, 1, 1) and v.shape == (2, 1, 1) + assert y[0].unique().tolist() == [64] and y[1].unique().tolist() == [940] + assert u.unique().tolist() == [512] and v.unique().tolist() == [512] + # Pure Rec.2020 red: Y = 0.2627, Cb = -0.5 * 0.2627 / (1 - 0.0593), Cr = 0.5. + red = torch.zeros(1, 3, 2, 2) + red[:, 0] = 1.0 + y, u, v = self.export_utils._rgb_to_yuv420p10_bt2020_limited(red) + assert y.unique().tolist() == [round((219 * 0.2627 + 16) * 4)] + assert u.unique().tolist() == [round((224 * -0.5 * 0.2627 / (1 - 0.0593) + 128) * 4)] + assert v.unique().tolist() == [round((224 * 0.5 + 128) * 4)] + + def test_encode_hlg_mp4(self, tmp_path): + import av + + if "libx265" not in av.codecs_available: + pytest.skip("The FFmpeg build used by PyAV does not include libx265.") + + # Diffuse white (linear 1.0) maps to the HLG signal 0.75 -> Y = round((219 * 0.75 + 16) * 4) = 721. + frames = torch.ones(9, 64, 96, 3) + output = tmp_path / "hlg.mp4" + self.export_utils.encode_hdr_tensor_to_hlg_mp4(frames, output, frame_rate=24.0, thread_count=1) + + with av.open(str(output)) as container: + stream = container.streams.video[0] + context = stream.codec_context + assert context.name == "hevc" + assert context.pix_fmt == "yuv420p10le" + assert stream.codec_tag == "hvc1" + assert (stream.width, stream.height) == (96, 64) + assert stream.average_rate == 24 + assert context.color_primaries == 9 # BT.2020 + assert context.color_trc == 18 # ARIB STD-B67 (HLG) + assert context.colorspace == 9 # BT.2020 NCL + assert context.color_range == 1 # limited + decoded = list(container.decode(video=0)) + assert len(decoded) == 9 + y_plane = decoded[0].planes[0] + y = np.frombuffer(y_plane, dtype=np.uint16).reshape(y_plane.height, y_plane.line_size // 2)[:, :96] + assert abs(int(np.median(y)) - 721) <= 2 + + def test_encode_hlg_mp4_rejects_odd_sizes(self, tmp_path): + import av + + if "libx265" not in av.codecs_available: + pytest.skip("The FFmpeg build used by PyAV does not include libx265.") + with pytest.raises(ValueError, match="even"): + self.export_utils.encode_hdr_tensor_to_hlg_mp4(torch.ones(1, 63, 64, 3), tmp_path / "x.mp4", 24.0) + assert not (tmp_path / "x.mp4").exists() + + +class TestEXR: + @pytest.fixture(autouse=True) + def _export_utils(self): + if not is_av_available(): + pytest.skip("PyAV is not installed.") + if not is_openexr_available(): + pytest.skip("OpenEXR is not installed.") + from diffusers.pipelines.ltx2 import export_utils + + self.export_utils = export_utils + + @staticmethod + def _read(path): + import OpenEXR + + with OpenEXR.File(str(path), separate_channels=True) as exr_file: + header = dict(exr_file.header()) + channels = {name: channel.pixels.copy() for name, channel in exr_file.channels().items()} + return header, channels + + @pytest.mark.parametrize( + "exr_colorspace, chromaticities, color_space", + [ + ("acescg", (0.713, 0.293, 0.165, 0.830, 0.128, 0.044, 0.32168, 0.33767), "ACEScg"), + ("acescct", (0.713, 0.293, 0.165, 0.830, 0.128, 0.044, 0.32168, 0.33767), "ACEScct"), + ("srgb_linear", (0.64, 0.33, 0.30, 0.60, 0.15, 0.06, 0.3127, 0.3290), "sRGB"), + ], + ) + def test_export_sequence_round_trip(self, tmp_path, exr_colorspace, chromaticities, color_space): + import OpenEXR + + generator = torch.Generator().manual_seed(0) + frames = torch.rand(2, 6, 10, 3, generator=generator) * 100.0 + paths = self.export_utils.export_to_exr_sequence(frames, tmp_path / "exr", exr_colorspace=exr_colorspace) + assert [p.split("/")[-1] for p in paths] == ["frame_00000.exr", "frame_00001.exr"] + + for path, frame in zip(paths, frames): + header, channels = self._read(path) + assert set(channels) == {"R", "G", "B"} + assert header["compression"] == OpenEXR.ZIP_COMPRESSION + assert header["colorSpace"] == color_space + np.testing.assert_allclose(header["chromaticities"], chromaticities, rtol=0, atol=1e-7) + for index, name in enumerate("RGB"): + assert channels[name].dtype == np.float16 + assert channels[name].shape == (6, 10) + np.testing.assert_array_equal(channels[name], frame[..., index].numpy().astype(np.float16)) + + def test_save_frame_channels_first_and_full_float(self, tmp_path): + frame = torch.arange(3 * 4 * 2, dtype=torch.float32).reshape(3, 4, 2) * 1e3 # (C, H, W) + self.export_utils.save_exr_frame(frame, tmp_path / "frame.exr", primaries="acescg", half=False) + header, channels = self._read(tmp_path / "frame.exr") + assert header["colorSpace"] == "sRGB" + for index, name in enumerate("RGB"): + assert channels[name].dtype == np.float32 + np.testing.assert_array_equal(channels[name], frame[index].numpy())