LTX-2-3-hdr / ltx-2-internal /conftest.py
linoyts's picture
linoyts HF Staff
Switch to HDR IC-LoRA app with on-the-fly PromptEncoder + bundled ltx-2-internal
8077a04 verified
Raw
History Blame Contribute Delete
8.12 kB
from __future__ import annotations
import gc
import os
from pathlib import Path
from typing import TYPE_CHECKING, Callable
import av
import pytest
import torch
import torch.nn.functional as F
from torch._prims_common import DeviceLikeType
if TYPE_CHECKING:
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
from ltx_core.model.transformer import Modality
torch.use_deterministic_algorithms(True)
os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"
MODELS_PATH = Path(os.getenv("MODELS_PATH", "/models"))
CHECKPOINTS_DIR = MODELS_PATH / "comfyui_models" / "checkpoints"
LORAS_DIR = MODELS_PATH / "comfyui_models" / "loras"
GEMMA_ROOT = MODELS_PATH / "comfyui_models" / "text_encoders" / "gemma-3-12b-it-qat-q4_0-unquantized_readout_proj"
GEMMA_FLATTENED_ROOT = MODELS_PATH / "comfyui_models" / "text_encoders" / "gemma-3-12b-it-qat-q4_0-unquantized"
LTX_2_0_CHECKPOINT_PATH = CHECKPOINTS_DIR / "ltx-2-19b-dev.safetensors"
LTX_2_0_CHECKPOINT_FP8_PATH = CHECKPOINTS_DIR / "ltx-2-19b-dev-fp8.safetensors"
LTX_2_3_CHECKPOINT_PATH = CHECKPOINTS_DIR / "ltx-2.3-22b-dev.safetensors"
LTX_2_0_DISTILLED_CHECKPOINT_PATH = CHECKPOINTS_DIR / "ltx-2-19b-distilled.safetensors"
LTX_2_3_DISTILLED_CHECKPOINT_PATH = CHECKPOINTS_DIR / "ltx-2.3-22b-distilled.safetensors"
LTX_2_0_SPATIAL_UPSAMPLER_PATH = CHECKPOINTS_DIR / "ltx2-spatial-upscaler-x2-1.0.bf16.safetensors"
LTX_2_3_SPATIAL_UPSAMPLER_PATH = (
MODELS_PATH / "comfyui_models" / "latent_upscale_models" / "ltx-2.3-spatial-upscaler-x2-1.0.safetensors"
)
LTX_2_0_DISTILLED_LORA_PATH = LORAS_DIR / "ltxv" / "ltx2" / "ltx-av-distilled-from-42500-lora-384_comfy.safetensors"
LTX_2_3_DISTILLED_LORA_PATH = LORAS_DIR / "ltxv" / "ltx2" / "ltx-2.3-22b-distilled-lora-384.safetensors"
LTX_2_3_IC_LORA_PATH = LORAS_DIR / "ltxv" / "ltx2" / "ltx-2.3-22b-ic-lora-union-control-ref0.5.safetensors"
LTX_2_3_HDR_IC_LORA_PATH = LORAS_DIR / "ltxv" / "ltx2" / "ltx-2.3-22b-ic-lora-hdr.safetensors"
LTX_2_3_HDR_TEXT_EMBEDDINGS_PATH = LORAS_DIR / "ltxv" / "ltx2" / "ltx-2.3-22b-ic-lora-hdr-scene-emb.pt"
def _psnr(pred: torch.Tensor, target: torch.Tensor, max_val: float = 1.0, eps: float = 1e-8) -> torch.Tensor:
"""
Compute Peak Signal-to-Noise Ratio (PSNR) between two images (or batches of images).
Args:
pred: Predicted image tensor, shape (..., H, W) or (..., C, H, W)
target: Ground truth image tensor, same shape as `pred`
max_val: Maximum possible pixel value of the images.
For images in [0, 1] use 1.0, for [0, 255] use 255.0, etc.
eps: Small value to avoid log of zero.
Returns:
psnr: PSNR value (in dB).
"""
if pred.shape != target.shape:
raise ValueError(f"Shape mismatch: pred {pred.shape}, target {target.shape}")
pred = pred.float()
target = target.float()
if pred.dim() > 1:
mse = F.mse_loss(pred, target, reduction="none")
dims = list(range(mse.dim()))
mse = mse.mean(dim=dims)
else:
mse = F.mse_loss(pred, target, reduction="mean")
psnr_val = 10.0 * torch.log10((max_val**2) / (mse + eps))
return psnr_val
def _psnr_per_frame(pred: torch.Tensor, target: torch.Tensor, max_val: float = 1.0, eps: float = 1e-8) -> torch.Tensor:
"""
Compute per-frame PSNR between two video tensors.
Args:
pred: Predicted video tensor, shape (T, C, H, W) or higher
target: Ground truth video tensor, same shape as `pred`
max_val: Maximum possible pixel value.
eps: Small value to avoid log of zero.
Returns:
psnr: Per-frame PSNR values (in dB).
"""
if pred.shape != target.shape:
raise ValueError(f"Shape mismatch: pred {pred.shape}, target {target.shape}")
pred = pred.float()
target = target.float()
if pred.dim() < 4:
raise ValueError("Expected at least 4D tensor (T, C, H, W) for per-frame PSNR.")
mse = F.mse_loss(pred, target, reduction="none")
mse = mse.mean(dim=(-3, -2, -1))
psnr_val = 10.0 * torch.log10((max_val**2) / (mse + eps))
return psnr_val
def _decode_video_from_file(path: str, device: DeviceLikeType) -> tuple[torch.Tensor, torch.Tensor | None]:
container = av.open(path)
try:
video_stream = next(s for s in container.streams if s.type == "video")
audio_stream = next((s for s in container.streams if s.type == "audio"), None)
frames = []
audio = [] if audio_stream else None
streams_to_decode = [video_stream]
if audio_stream:
streams_to_decode.append(audio_stream)
for frame in container.decode(*streams_to_decode):
if isinstance(frame, av.VideoFrame):
tensor = torch.tensor(frame.to_rgb().to_ndarray(), dtype=torch.uint8, device=device).unsqueeze(0)
frames.append(tensor)
elif isinstance(frame, av.AudioFrame):
audio.append(torch.tensor(frame.to_ndarray(), dtype=torch.float32, device=device).unsqueeze(0))
if audio:
audio = torch.cat(audio)
finally:
container.close()
return torch.cat(frames), audio
class MockTransformer:
"""Mock transformer that tracks calls and returns the input latents.
Records ``call_count`` and ``batch_sizes`` for assertions. Output is
simply ``video.latent`` / ``audio.latent`` (identity), which is
deterministic and split-safe for adapter equivalence tests.
"""
def __init__(self) -> None:
self.call_count = 0
self.batch_sizes: list[int] = []
@staticmethod
def _validate_modality(modality: Modality) -> None:
"""Check that all tensor fields share the same batch dimension.
The real transformer's ``_prepare_timestep`` does
``view(batch_size, ...)`` on sigma embeddings, which crashes if
sigma's batch dim doesn't match latent's. This check catches
such mismatches without needing real model weights.
"""
batch_size = modality.latent.shape[0]
for name, tensor in [
("sigma", modality.sigma),
("timesteps", modality.timesteps),
("positions", modality.positions),
("context", modality.context),
]:
if tensor.shape[0] != batch_size:
msg = f"Modality.{name} batch dim {tensor.shape[0]} != latent batch dim {batch_size}"
raise ValueError(msg)
def __call__(
self,
video: Modality | None = None,
audio: Modality | None = None,
perturbations: BatchedPerturbationConfig | None = None, # noqa: ARG002
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
self.call_count += 1
if video is not None:
self._validate_modality(video)
if audio is not None:
self._validate_modality(audio)
b = video.latent.shape[0] if video is not None else audio.latent.shape[0]
self.batch_sizes.append(b)
return (
video.latent if video is not None else None,
audio.latent if audio is not None else None,
)
@pytest.fixture
def mock_transformer() -> MockTransformer:
"""Return a fresh MockTransformer instance."""
return MockTransformer()
@pytest.fixture(autouse=True)
def _cleanup_cuda_memory() -> None:
"""Free GPU memory before and after each test to prevent OOM across test modules."""
gc.collect()
torch.cuda.empty_cache()
yield
gc.collect()
torch.cuda.empty_cache()
@pytest.fixture
def psnr() -> Callable[[torch.Tensor, torch.Tensor, float, float], torch.Tensor]:
"""Fixture that returns the PSNR function."""
return _psnr
@pytest.fixture
def psnr_per_frame() -> Callable[[torch.Tensor, torch.Tensor, float, float], torch.Tensor]:
"""Fixture that returns the per-frame PSNR function."""
return _psnr_per_frame
@pytest.fixture
def decode_video_from_file() -> Callable[[str, DeviceLikeType], tuple[torch.Tensor, torch.Tensor | None]]:
"""Fixture that returns the function to decode a video from a file."""
return _decode_video_from_file