Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3bd2a885e4 | ||
|
|
6566a05dac | ||
|
|
169d4849b3 | ||
|
|
abeea25f77 | ||
|
|
db9ba98fbf | ||
|
|
8bb9a31292 | ||
|
|
0c33204bbd | ||
|
|
b076cd934e | ||
|
|
81872ee886 |
@@ -0,0 +1,90 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.entrypoints.upsample import (_prepare_video, _read_video,
|
||||
_write_video)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import UpsamplerLoader, VAELoader
|
||||
from fastvideo.models.upsamplers import upsample_video
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Input/output
|
||||
INPUT_VIDEO = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
|
||||
OUTPUT_VIDEO = "outputs_video/ltx2_upscale/ltx2_spatiotemporal_upscale_x2.mp4"
|
||||
|
||||
# Diffusers-style LTX-2 repo with upsamplers included
|
||||
MODEL_ID = "FastVideo/LTX2-Diffusers"
|
||||
|
||||
# Controls
|
||||
DOUBLE_FPS = True
|
||||
|
||||
|
||||
def main() -> None:
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
precision_str = "bf16" if torch.cuda.is_available() else "fp32"
|
||||
dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32
|
||||
|
||||
video, fps = _read_video(INPUT_VIDEO)
|
||||
video = _prepare_video(
|
||||
video,
|
||||
trim_frames=False,
|
||||
pad_frames=True,
|
||||
crop_multiple=32,
|
||||
)
|
||||
|
||||
model_root = maybe_download_model(MODEL_ID)
|
||||
vae_path = str(Path(model_root) / "vae")
|
||||
spatial_upsampler_path = str(Path(model_root) / "spatial_upsampler")
|
||||
temporal_upsampler_path = str(Path(model_root) / "temporal_upsampler")
|
||||
|
||||
args = FastVideoArgs(
|
||||
model_path=vae_path,
|
||||
pipeline_config=PipelineConfig(vae_precision=precision_str),
|
||||
vae_cpu_offload=False,
|
||||
)
|
||||
vae = VAELoader().load(vae_path, args).to(device=device, dtype=dtype)
|
||||
spatial_upsampler = UpsamplerLoader().load(
|
||||
spatial_upsampler_path, args).to(device=device, dtype=dtype)
|
||||
temporal_upsampler = UpsamplerLoader().load(
|
||||
temporal_upsampler_path, args).to(device=device, dtype=dtype)
|
||||
|
||||
if hasattr(vae.decoder, "decode_noise_scale"):
|
||||
vae.decoder.decode_noise_scale = 0.0
|
||||
|
||||
# [F, C, H, W] -> [B, C, F, H, W]
|
||||
video = video.unsqueeze(0).permute(0, 2, 1, 3, 4).to(
|
||||
device=device, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
latents = vae.encoder(video)
|
||||
up_latents = upsample_video(latents, vae.encoder,
|
||||
getattr(spatial_upsampler, "model",
|
||||
spatial_upsampler))
|
||||
up_latents = upsample_video(up_latents, vae.encoder,
|
||||
getattr(temporal_upsampler, "model",
|
||||
temporal_upsampler))
|
||||
|
||||
timestep_value = getattr(vae.decoder, "decode_timestep", 0.05)
|
||||
timestep = torch.full((video.shape[0], ),
|
||||
float(timestep_value),
|
||||
device=device,
|
||||
dtype=dtype)
|
||||
decoded = vae.decoder(up_latents, timestep=timestep)
|
||||
|
||||
# [B, C, F, H, W] -> [F, C, H, W]
|
||||
decoded = decoded[0].permute(1, 0, 2, 3).detach().cpu()
|
||||
|
||||
output_fps = fps * 2 if DOUBLE_FPS else fps
|
||||
Path(OUTPUT_VIDEO).parent.mkdir(parents=True, exist_ok=True)
|
||||
_write_video(decoded, OUTPUT_VIDEO, output_fps)
|
||||
logger.info("Spatiotemporal upsampled video saved to %s", OUTPUT_VIDEO)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,84 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.entrypoints.upsample import (_prepare_video, _read_video,
|
||||
_write_video)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import UpsamplerLoader, VAELoader
|
||||
from fastvideo.models.upsamplers import upsample_video
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Input/output
|
||||
INPUT_VIDEO = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
|
||||
OUTPUT_VIDEO = "outputs_video/ltx2_upscale/ltx2_temporal_upscale_x2.mp4"
|
||||
|
||||
# Diffusers-style LTX-2 repo with upsamplers included
|
||||
MODEL_ID = "FastVideo/LTX2-Diffusers"
|
||||
|
||||
# Controls
|
||||
DOUBLE_FPS = True
|
||||
|
||||
|
||||
def main() -> None:
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
precision_str = "bf16" if torch.cuda.is_available() else "fp32"
|
||||
dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32
|
||||
|
||||
video, fps = _read_video(INPUT_VIDEO)
|
||||
video = _prepare_video(
|
||||
video,
|
||||
trim_frames=False,
|
||||
pad_frames=True,
|
||||
crop_multiple=32,
|
||||
)
|
||||
|
||||
model_root = maybe_download_model(MODEL_ID)
|
||||
vae_path = str(Path(model_root) / "vae")
|
||||
temporal_upsampler_path = str(Path(model_root) / "temporal_upsampler")
|
||||
|
||||
args = FastVideoArgs(
|
||||
model_path=vae_path,
|
||||
pipeline_config=PipelineConfig(vae_precision=precision_str),
|
||||
vae_cpu_offload=False,
|
||||
)
|
||||
vae = VAELoader().load(vae_path, args).to(device=device, dtype=dtype)
|
||||
temporal_upsampler = UpsamplerLoader().load(
|
||||
temporal_upsampler_path, args).to(device=device, dtype=dtype)
|
||||
|
||||
if hasattr(vae.decoder, "decode_noise_scale"):
|
||||
vae.decoder.decode_noise_scale = 0.0
|
||||
|
||||
# [F, C, H, W] -> [B, C, F, H, W]
|
||||
video = video.unsqueeze(0).permute(0, 2, 1, 3, 4).to(
|
||||
device=device, dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
latents = vae.encoder(video)
|
||||
up_latents = upsample_video(latents, vae.encoder,
|
||||
getattr(temporal_upsampler, "model",
|
||||
temporal_upsampler))
|
||||
|
||||
timestep_value = getattr(vae.decoder, "decode_timestep", 0.05)
|
||||
timestep = torch.full((video.shape[0], ),
|
||||
float(timestep_value),
|
||||
device=device,
|
||||
dtype=dtype)
|
||||
decoded = vae.decoder(up_latents, timestep=timestep)
|
||||
|
||||
# [B, C, F, H, W] -> [F, C, H, W]
|
||||
decoded = decoded[0].permute(1, 0, 2, 3).detach().cpu()
|
||||
|
||||
output_fps = fps * 2 if DOUBLE_FPS else fps
|
||||
Path(OUTPUT_VIDEO).parent.mkdir(parents=True, exist_ok=True)
|
||||
_write_video(decoded, OUTPUT_VIDEO, output_fps)
|
||||
logger.info("Temporal upsampled video saved to %s", OUTPUT_VIDEO)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,51 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
|
||||
# HF model ID (downloaded automatically). Distilled repos include upsamplers + refine LoRA reference.
|
||||
MODEL_ID = "FastVideo/LTX2-Distilled-Diffusers"
|
||||
|
||||
OUTPUT_PATH = "outputs_video/ltx2_upscale/ltx2_two_stage.mp4"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
MODEL_ID,
|
||||
num_gpus=8,
|
||||
ltx2_refine_enabled=True,
|
||||
ltx2_refine_num_inference_steps=3,
|
||||
ltx2_refine_guidance_scale=1.0,
|
||||
ltx2_refine_add_noise=True,
|
||||
)
|
||||
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=OUTPUT_PATH,
|
||||
height=960,
|
||||
width=1664,
|
||||
num_frames=81,
|
||||
fps=24,
|
||||
seed=10,
|
||||
# DistilledPipeline uses the 8-step distilled schedule without CFG.
|
||||
num_inference_steps=8,
|
||||
guidance_scale=1.0,
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -16,5 +16,18 @@ class LTX2SamplingParam(SamplingParam):
|
||||
fps: int = 24
|
||||
num_inference_steps: int = 8
|
||||
guidance_scale: float = 1.0
|
||||
# No default negative_prompt for distilled models
|
||||
negative_prompt: str = ""
|
||||
|
||||
# Official LTX-2 negative prompt (used only when guidance_scale > 1)
|
||||
negative_prompt: str = (
|
||||
"blurry, out of focus, overexposed, underexposed, low contrast, washed out colors, excessive noise, "
|
||||
"grainy texture, poor lighting, flickering, motion blur, distorted proportions, unnatural skin tones, "
|
||||
"deformed facial features, asymmetrical face, missing facial features, extra limbs, disfigured hands, "
|
||||
"wrong hand count, artifacts around text, inconsistent perspective, camera shake, incorrect depth of "
|
||||
"field, background too sharp, background clutter, distracting reflections, harsh shadows, inconsistent "
|
||||
"lighting direction, color banding, cartoonish rendering, 3D CGI look, unrealistic materials, uncanny "
|
||||
"valley effect, incorrect ethnicity, wrong gender, exaggerated expressions, wrong gaze direction, "
|
||||
"mismatched lip sync, silent or muted audio, distorted voice, robotic voice, echo, background noise, "
|
||||
"off-sync audio, incorrect dialogue, added dialogue, repetitive speech, jittery movement, awkward "
|
||||
"pauses, incorrect timing, unnatural transitions, inconsistent framing, tilted camera, flat lighting, "
|
||||
"inconsistent tone, cinematic oversaturation, stylized filters, or AI artifacts."
|
||||
)
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.cli.generate import cmd_init as generate_cmd_init
|
||||
from fastvideo.entrypoints.cli.upsample import cmd_init as upsample_cmd_init
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
@@ -10,6 +11,7 @@ def cmd_init() -> list[CLISubcommand]:
|
||||
"""Initialize all commands from separate modules"""
|
||||
commands = []
|
||||
commands.extend(generate_cmd_init())
|
||||
commands.extend(upsample_cmd_init())
|
||||
return commands
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.upsample import upscale_video_file
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class UpsampleSubcommand(CLISubcommand):
|
||||
"""The `upsample` subcommand for the FastVideo CLI."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.name = "upsample"
|
||||
super().__init__()
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
upscale_video_file(
|
||||
input_video=args.input_video,
|
||||
output_video=args.output_video,
|
||||
vae_path=args.vae_path,
|
||||
upsampler_path=args.upsampler_path,
|
||||
precision=args.precision,
|
||||
device=args.device,
|
||||
max_frames=args.max_frames,
|
||||
trim_frames=args.trim_frames,
|
||||
pad_frames=args.pad_frames,
|
||||
crop_multiple=args.crop_multiple,
|
||||
output_fps=args.output_fps,
|
||||
)
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
if not os.path.exists(args.input_video):
|
||||
raise ValueError(f"Input video not found: {args.input_video}")
|
||||
if args.crop_multiple is not None and args.crop_multiple < 0:
|
||||
raise ValueError("crop_multiple must be >= 0")
|
||||
if args.max_frames is not None and args.max_frames <= 0:
|
||||
raise ValueError("max_frames must be positive")
|
||||
if args.trim_frames and args.pad_frames:
|
||||
raise ValueError(
|
||||
"Only one of --trim-frames or --pad-frames can be enabled")
|
||||
|
||||
def subparser_init(
|
||||
self,
|
||||
subparsers: argparse._SubParsersAction,
|
||||
) -> FlexibleArgumentParser:
|
||||
parser = subparsers.add_parser(
|
||||
"upsample",
|
||||
help="Upscale an existing video using the LTX-2 spatial upsampler",
|
||||
usage=
|
||||
("fastvideo upsample --input-video INPUT.mp4 --output-video OUTPUT.mp4 "
|
||||
"[--vae-path PATH] [--upsampler-path PATH]"),
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--input-video",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to the input video file",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-video",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to save the upscaled video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-path",
|
||||
type=str,
|
||||
default="converted/ltx2_diffusers/vae",
|
||||
help="Path to LTX-2 VAE weights (diffusers-style)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--upsampler-path",
|
||||
type=str,
|
||||
default="converted/ltx2_spatial_upscaler",
|
||||
help="Path to LTX-2 spatial upsampler weights",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--precision",
|
||||
type=str,
|
||||
default="bf16",
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision to use for VAE + upsampler",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Torch device string (e.g. cuda, cuda:0, cpu)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max-frames",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Maximum number of frames to read from the input video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--trim-frames",
|
||||
action=StoreBoolean,
|
||||
default=True,
|
||||
help="Trim frames to satisfy the 1+8k requirement",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pad-frames",
|
||||
action=StoreBoolean,
|
||||
default=False,
|
||||
help=
|
||||
"Pad frames to satisfy the 1+8k requirement (repeats last frame)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--crop-multiple",
|
||||
type=int,
|
||||
default=32,
|
||||
help=
|
||||
"Center-crop to make H/W divisible by this value (0 to disable)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-fps",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Override output video FPS (defaults to input FPS)",
|
||||
)
|
||||
|
||||
return cast(FlexibleArgumentParser, parser)
|
||||
|
||||
|
||||
def cmd_init() -> list[CLISubcommand]:
|
||||
return [UpsampleSubcommand()]
|
||||
@@ -0,0 +1,202 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Utilities for upscaling videos with LTX-2 spatial upsampler."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import av
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import UpsamplerLoader, VAELoader
|
||||
from fastvideo.models.upsamplers import upsample_video
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _read_video(path: str | Path,
|
||||
max_frames: int | None = None) -> tuple[torch.Tensor, float]:
|
||||
"""Read video frames via PyAV.
|
||||
|
||||
Returns a tensor of shape [F, C, H, W] in [0, 1] and the fps.
|
||||
"""
|
||||
path = Path(path)
|
||||
if not path.exists():
|
||||
raise FileNotFoundError(f"Input video not found: {path}")
|
||||
|
||||
frames: list[np.ndarray] = []
|
||||
with av.open(str(path)) as container:
|
||||
video_stream = container.streams.video[0]
|
||||
fps = float(video_stream.average_rate or video_stream.base_rate or 24)
|
||||
for frame in container.decode(video=0):
|
||||
if max_frames is not None and len(frames) >= max_frames:
|
||||
break
|
||||
frames.append(frame.to_ndarray(format="rgb24"))
|
||||
|
||||
if not frames:
|
||||
raise ValueError(f"No frames decoded from {path}")
|
||||
|
||||
frames_np = np.stack(frames, axis=0)
|
||||
video = torch.from_numpy(frames_np).float().div(255.0)
|
||||
return video.permute(0, 3, 1, 2), fps
|
||||
|
||||
|
||||
def _write_video(frames: torch.Tensor, output_path: str | Path,
|
||||
fps: float) -> None:
|
||||
"""Write frames [F, C, H, W] in [0, 1] to a video file."""
|
||||
output_path = Path(output_path)
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
frames = frames.clamp(0, 1)
|
||||
frames = (frames * 255.0).to(torch.uint8)
|
||||
frames_np = frames.permute(0, 2, 3, 1).cpu().numpy()
|
||||
imageio.mimsave(str(output_path), list(frames_np), fps=fps, format="mp4")
|
||||
|
||||
|
||||
def _prepare_video(
|
||||
video: torch.Tensor,
|
||||
*,
|
||||
trim_frames: bool,
|
||||
pad_frames: bool,
|
||||
crop_multiple: int,
|
||||
) -> torch.Tensor:
|
||||
"""Ensure frames count and resolution satisfy LTX-2 VAE constraints."""
|
||||
frames, _, height, width = video.shape
|
||||
|
||||
if trim_frames and pad_frames:
|
||||
raise ValueError(
|
||||
"Only one of trim_frames or pad_frames can be enabled.")
|
||||
|
||||
if trim_frames and ((frames - 1) % 8) != 0:
|
||||
valid_frames = 1 + 8 * ((frames - 1) // 8)
|
||||
if valid_frames < 1:
|
||||
raise ValueError("Video must have at least 1 frame.")
|
||||
if valid_frames != frames:
|
||||
logger.warning(
|
||||
"Trimming frames from %d to %d to satisfy 1+8k requirement.",
|
||||
frames,
|
||||
valid_frames,
|
||||
)
|
||||
video = video[:valid_frames]
|
||||
frames = valid_frames
|
||||
elif pad_frames and ((frames - 1) % 8) != 0:
|
||||
valid_frames = 1 + 8 * (((frames - 1) + 7) // 8)
|
||||
pad_count = valid_frames - frames
|
||||
if pad_count > 0:
|
||||
logger.warning(
|
||||
"Padding frames from %d to %d to satisfy 1+8k requirement.",
|
||||
frames,
|
||||
valid_frames,
|
||||
)
|
||||
pad = video[-1:].repeat(pad_count, 1, 1, 1)
|
||||
video = torch.cat([video, pad], dim=0)
|
||||
frames = valid_frames
|
||||
|
||||
if crop_multiple > 0:
|
||||
new_height = height - (height % crop_multiple)
|
||||
new_width = width - (width % crop_multiple)
|
||||
if new_height != height or new_width != width:
|
||||
top = max((height - new_height) // 2, 0)
|
||||
left = max((width - new_width) // 2, 0)
|
||||
logger.warning(
|
||||
"Center-cropping from %dx%d to %dx%d to be divisible by %d.",
|
||||
height,
|
||||
width,
|
||||
new_height,
|
||||
new_width,
|
||||
crop_multiple,
|
||||
)
|
||||
video = video[:, :, top:top + new_height, left:left + new_width]
|
||||
|
||||
return video
|
||||
|
||||
|
||||
def upscale_video_file(
|
||||
*,
|
||||
input_video: str | Path,
|
||||
output_video: str | Path,
|
||||
vae_path: str | Path,
|
||||
upsampler_path: str | Path,
|
||||
precision: str = "bf16",
|
||||
device: str | None = None,
|
||||
max_frames: int | None = None,
|
||||
trim_frames: bool = True,
|
||||
pad_frames: bool = False,
|
||||
crop_multiple: int = 32,
|
||||
output_fps: float | None = None,
|
||||
) -> None:
|
||||
"""Upscale an existing video using the LTX-2 spatial upsampler."""
|
||||
input_video = str(input_video)
|
||||
output_video = str(output_video)
|
||||
vae_path = str(vae_path)
|
||||
upsampler_path = str(upsampler_path)
|
||||
|
||||
video, fps = _read_video(input_video, max_frames=max_frames)
|
||||
original_frames = video.shape[0]
|
||||
video = _prepare_video(
|
||||
video,
|
||||
trim_frames=trim_frames,
|
||||
pad_frames=pad_frames,
|
||||
crop_multiple=crop_multiple,
|
||||
)
|
||||
final_frames = video.shape[0]
|
||||
|
||||
target_device = torch.device(device) if device else (torch.device(
|
||||
"cuda") if torch.cuda.is_available() else torch.device("cpu"))
|
||||
|
||||
precision = precision.lower()
|
||||
dtype = PRECISION_TO_TYPE.get(precision, torch.bfloat16)
|
||||
if target_device.type == "cpu" and dtype != torch.float32:
|
||||
logger.warning("CPU device selected; overriding precision to fp32.")
|
||||
dtype = torch.float32
|
||||
precision = "fp32"
|
||||
|
||||
args = FastVideoArgs(
|
||||
model_path=vae_path,
|
||||
pipeline_config=PipelineConfig(vae_precision=precision),
|
||||
vae_cpu_offload=False,
|
||||
)
|
||||
|
||||
vae_loader = VAELoader()
|
||||
upsampler_loader = UpsamplerLoader()
|
||||
|
||||
vae = vae_loader.load(vae_path, args).to(device=target_device, dtype=dtype)
|
||||
upsampler = upsampler_loader.load(upsampler_path,
|
||||
args).to(device=target_device,
|
||||
dtype=dtype)
|
||||
|
||||
if hasattr(vae.decoder, "decode_noise_scale"):
|
||||
vae.decoder.decode_noise_scale = 0.0
|
||||
|
||||
# [F, C, H, W] -> [B, C, F, H, W]
|
||||
video = video.unsqueeze(0).permute(0, 2, 1, 3, 4).to(device=target_device,
|
||||
dtype=dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
latents = vae.encoder(video)
|
||||
up_latents = upsample_video(latents, vae.encoder,
|
||||
getattr(upsampler, "model", upsampler))
|
||||
|
||||
timestep_value = getattr(vae.decoder, "decode_timestep", 0.05)
|
||||
timestep = torch.full((video.shape[0], ),
|
||||
float(timestep_value),
|
||||
device=target_device,
|
||||
dtype=dtype)
|
||||
decoded = vae.decoder(up_latents, timestep=timestep)
|
||||
|
||||
# [B, C, F, H, W] -> [F, C, H, W]
|
||||
decoded = decoded[0].permute(1, 0, 2, 3).detach().cpu()
|
||||
if pad_frames and final_frames != original_frames:
|
||||
decoded = decoded[:original_frames]
|
||||
final_fps = output_fps or fps
|
||||
_write_video(decoded, output_video, final_fps)
|
||||
logger.info("Upscaled video saved to %s", output_video)
|
||||
|
||||
|
||||
__all__ = ["upscale_video_file"]
|
||||
@@ -171,6 +171,39 @@ class FastVideoArgs:
|
||||
ltx2_vae_temporal_tile_size_in_frames: int | None = None
|
||||
ltx2_vae_temporal_tile_overlap_in_frames: int | None = None
|
||||
ltx2_initial_latent_path: str | None = None
|
||||
ltx2_audio_latent_path: str | None = None
|
||||
# Generic stage-2 refine args (preferred API). These map to LTX-2 refine args
|
||||
# for now, but keep the user-facing API model-agnostic.
|
||||
refine_enabled: bool | None = None
|
||||
refine_upsampler_path: str | None = None
|
||||
refine_transformer_path: str | None = None
|
||||
refine_lora_path: str | None = None
|
||||
refine_num_inference_steps: int | None = None
|
||||
refine_guidance_scale: float | None = None
|
||||
refine_add_noise: bool | None = None
|
||||
refine_noise_path: str | None = None
|
||||
refine_audio_noise_path: str | None = None
|
||||
ltx2_refine_enabled: bool = False
|
||||
ltx2_refine_upsampler_path: str | None = None
|
||||
ltx2_refine_transformer_path: str | None = None
|
||||
ltx2_refine_lora_path: str | None = None
|
||||
ltx2_refine_num_inference_steps: int = 3
|
||||
ltx2_refine_guidance_scale: float = 1.0
|
||||
ltx2_refine_add_noise: bool = True
|
||||
ltx2_refine_noise_path: str | None = None
|
||||
ltx2_refine_audio_noise_path: str | None = None
|
||||
|
||||
# Debugging (opt-in, minimal overhead when disabled)
|
||||
debug_stage_sums: bool = False
|
||||
debug_stage_sums_path: str | None = None
|
||||
debug_model_sums: bool = False
|
||||
debug_model_sums_path: str | None = None
|
||||
debug_model_detail: bool = False
|
||||
debug_model_detail_path: str | None = None
|
||||
debug_module_sums: bool = False
|
||||
debug_module_sums_path: str | None = None
|
||||
debug_module_sums_include: list[str] | None = None
|
||||
debug_module_sums_exclude: list[str] | None = None
|
||||
|
||||
# model paths for correct deallocation
|
||||
model_paths: dict[str, str] = field(default_factory=dict)
|
||||
@@ -211,6 +244,7 @@ class FastVideoArgs:
|
||||
self.moba_config_path, e)
|
||||
raise
|
||||
self._apply_ltx2_vae_overrides()
|
||||
self._resolve_refine_args()
|
||||
self.check_fastvideo_args()
|
||||
|
||||
def _apply_ltx2_vae_overrides(self) -> None:
|
||||
@@ -248,6 +282,27 @@ class FastVideoArgs:
|
||||
vae_config.ltx2_temporal_tile_overlap_in_frames = (
|
||||
self.ltx2_vae_temporal_tile_overlap_in_frames)
|
||||
|
||||
def _resolve_refine_args(self) -> None:
|
||||
"""Map generic refine_* args to LTX-2-specific refine fields."""
|
||||
if self.refine_enabled is not None:
|
||||
self.ltx2_refine_enabled = self.refine_enabled
|
||||
if self.refine_upsampler_path is not None:
|
||||
self.ltx2_refine_upsampler_path = self.refine_upsampler_path
|
||||
if self.refine_transformer_path is not None:
|
||||
self.ltx2_refine_transformer_path = self.refine_transformer_path
|
||||
if self.refine_lora_path is not None:
|
||||
self.ltx2_refine_lora_path = self.refine_lora_path
|
||||
if self.refine_num_inference_steps is not None:
|
||||
self.ltx2_refine_num_inference_steps = self.refine_num_inference_steps
|
||||
if self.refine_guidance_scale is not None:
|
||||
self.ltx2_refine_guidance_scale = self.refine_guidance_scale
|
||||
if self.refine_add_noise is not None:
|
||||
self.ltx2_refine_add_noise = self.refine_add_noise
|
||||
if self.refine_noise_path is not None:
|
||||
self.ltx2_refine_noise_path = self.refine_noise_path
|
||||
if self.refine_audio_noise_path is not None:
|
||||
self.ltx2_refine_audio_noise_path = self.refine_audio_noise_path
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
# Model and path configuration
|
||||
@@ -405,6 +460,136 @@ class FastVideoArgs:
|
||||
default=FastVideoArgs.ltx2_initial_latent_path,
|
||||
help="Path to load/save a precomputed LTX-2 initial latent.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-audio-latent-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.ltx2_audio_latent_path,
|
||||
help="Path to load/save a precomputed LTX-2 initial audio latent.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-refine-enabled",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.ltx2_refine_enabled,
|
||||
help=
|
||||
"Enable LTX-2 stage2 refinement (2x spatial upsample + distilled denoising).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-refine-upsampler-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.ltx2_refine_upsampler_path,
|
||||
help=
|
||||
"Path to the LTX-2 spatial upsampler weights (diffusers format).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-refine-transformer-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.ltx2_refine_transformer_path,
|
||||
help=
|
||||
"Optional path to a dedicated stage2 transformer (e.g., distilled LoRA weights).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-refine-lora-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.ltx2_refine_lora_path,
|
||||
help=
|
||||
"Optional LoRA path to apply only during LTX-2 refinement stage2.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-refine-num-inference-steps",
|
||||
type=int,
|
||||
default=FastVideoArgs.ltx2_refine_num_inference_steps,
|
||||
help="Number of refinement steps for stage2 denoising (default: 3).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-refine-guidance-scale",
|
||||
type=float,
|
||||
default=FastVideoArgs.ltx2_refine_guidance_scale,
|
||||
help="CFG guidance scale for refinement (1.0 disables CFG).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-refine-add-noise",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.ltx2_refine_add_noise,
|
||||
help="Add noise at sigma0 before stage2 denoising.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-refine-noise-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.ltx2_refine_noise_path,
|
||||
help="Path to load/save stage2 video noise before refinement.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ltx2-refine-audio-noise-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.ltx2_refine_audio_noise_path,
|
||||
help="Path to load/save stage2 audio noise before refinement.",
|
||||
)
|
||||
|
||||
# Debugging (opt-in)
|
||||
parser.add_argument(
|
||||
"--debug-stage-sums",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.debug_stage_sums,
|
||||
help="Log tensor sums after each pipeline stage (debug only).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--debug-stage-sums-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.debug_stage_sums_path,
|
||||
help="Path to write stage-level sum logs (appended).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--debug-model-sums",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.debug_model_sums,
|
||||
help="Enable model-level sum logging (e.g., LTX-2 transformer).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--debug-model-sums-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.debug_model_sums_path,
|
||||
help="Path to write model-level sum logs.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--debug-model-detail",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.debug_model_detail,
|
||||
help="Enable detailed model hooks for activation sums (debug only).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--debug-model-detail-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.debug_model_detail_path,
|
||||
help="Path to write detailed model hook logs.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--debug-module-sums",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.debug_module_sums,
|
||||
help="Enable recursive module-level output sum logging.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--debug-module-sums-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.debug_module_sums_path,
|
||||
help="Path to write recursive module-level sum logs.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--debug-module-sums-include",
|
||||
nargs="+",
|
||||
type=str,
|
||||
default=FastVideoArgs.debug_module_sums_include,
|
||||
help=
|
||||
"Optional list of substrings; only module names containing these will be logged.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--debug-module-sums-exclude",
|
||||
nargs="+",
|
||||
type=str,
|
||||
default=FastVideoArgs.debug_module_sums_exclude,
|
||||
help=
|
||||
"Optional list of substrings; module names containing these will be skipped.",
|
||||
)
|
||||
|
||||
# LoRA parameters (inference-time adapter loading)
|
||||
parser.add_argument(
|
||||
|
||||
@@ -13,7 +13,7 @@ from fastvideo.distributed import (get_local_torch_device, get_tp_rank,
|
||||
split_tensor_along_last_dim,
|
||||
tensor_model_parallel_all_gather,
|
||||
tensor_model_parallel_all_reduce)
|
||||
from fastvideo.layers.linear import (ColumnParallelLinear, LinearBase,
|
||||
from fastvideo.layers.linear import (ColumnParallelLinear,
|
||||
MergedColumnParallelLinear,
|
||||
QKVParallelLinear, ReplicatedLinear,
|
||||
RowParallelLinear)
|
||||
@@ -82,10 +82,9 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
lora_B_sliced = self.slice_lora_b_weights(
|
||||
lora_B.to(x, non_blocking=True))
|
||||
delta = x @ lora_A_sliced.T @ lora_B_sliced.T
|
||||
if self.lora_alpha != self.lora_rank:
|
||||
delta = delta * (
|
||||
self.lora_alpha / self.lora_rank # type: ignore
|
||||
) # type: ignore
|
||||
if (self.lora_alpha is not None and self.lora_rank is not None
|
||||
and self.lora_alpha != self.lora_rank):
|
||||
delta = delta * (self.lora_alpha / self.lora_rank)
|
||||
out, output_bias = self.base_layer(x)
|
||||
return out + delta, output_bias
|
||||
else:
|
||||
@@ -98,6 +97,31 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
|
||||
return B
|
||||
|
||||
|
||||
class TorchLinearWithLoRA(BaseLayerWithLoRA):
|
||||
"""LoRA wrapper for torch.nn.Linear modules."""
|
||||
|
||||
@torch.compile()
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
lora_A = self.lora_A
|
||||
lora_B = self.lora_B
|
||||
if isinstance(self.lora_B, DTensor):
|
||||
lora_B = self.lora_B.to_local()
|
||||
lora_A = self.lora_A.to_local()
|
||||
|
||||
if not self.merged and not self.disable_lora:
|
||||
lora_A_sliced = self.slice_lora_a_weights(
|
||||
lora_A.to(x, non_blocking=True))
|
||||
lora_B_sliced = self.slice_lora_b_weights(
|
||||
lora_B.to(x, non_blocking=True))
|
||||
delta = x @ lora_A_sliced.T @ lora_B_sliced.T
|
||||
if (self.lora_alpha is not None and self.lora_rank is not None
|
||||
and self.lora_alpha != self.lora_rank):
|
||||
delta = delta * (self.lora_alpha / self.lora_rank)
|
||||
out = self.base_layer(x)
|
||||
return out + delta
|
||||
return self.base_layer(x)
|
||||
|
||||
def set_lora_weights(self,
|
||||
A: torch.Tensor,
|
||||
B: torch.Tensor,
|
||||
@@ -369,7 +393,7 @@ def get_lora_layer(layer: nn.Module,
|
||||
lora_rank: int | None = None,
|
||||
lora_alpha: int | None = None,
|
||||
training_mode: bool = False) -> BaseLayerWithLoRA | None:
|
||||
supported_layer_types: dict[type[LinearBase], type[BaseLayerWithLoRA]] = {
|
||||
supported_layer_types: dict[type[nn.Module], type[BaseLayerWithLoRA]] = {
|
||||
# the order matters
|
||||
# VocabParallelEmbedding: VocabParallelEmbeddingWithLoRA,
|
||||
QKVParallelLinear: QKVParallelLinearWithLoRA,
|
||||
@@ -377,6 +401,7 @@ def get_lora_layer(layer: nn.Module,
|
||||
ColumnParallelLinear: ColumnParallelLinearWithLoRA,
|
||||
RowParallelLinear: RowParallelLinearWithLoRA,
|
||||
ReplicatedLinear: BaseLayerWithLoRA,
|
||||
nn.Linear: TorchLinearWithLoRA,
|
||||
}
|
||||
for src_layer_type, lora_layer_type in supported_layer_types.items():
|
||||
if isinstance(layer, src_layer_type): # pylint: disable=unidiomatic-typecheck
|
||||
|
||||
@@ -34,6 +34,7 @@ from fastvideo.models.loader.weight_utils import (
|
||||
safetensors_weights_iterator,
|
||||
)
|
||||
from fastvideo.models.registry import ModelRegistry
|
||||
from fastvideo.models.upsamplers.config_adapters import get_upsampler_config
|
||||
from fastvideo.utils import PRECISION_TO_TYPE, is_pin_memory_available
|
||||
from fastvideo.hooks.layerwise_offload import enable_layerwise_offload
|
||||
|
||||
@@ -84,6 +85,8 @@ class ComponentLoader(ABC):
|
||||
"audio_vae": (AudioDecoderLoader, "diffusers"),
|
||||
"audio_decoder": (AudioDecoderLoader, "diffusers"),
|
||||
"vocoder": (VocoderLoader, "diffusers"),
|
||||
"upsampler": (UpsamplerLoader, "diffusers"),
|
||||
"spatial_upsampler": (UpsamplerLoader, "diffusers"),
|
||||
"text_encoder": (TextEncoderLoader, "transformers"),
|
||||
"text_encoder_2": (TextEncoderLoader, "transformers"),
|
||||
"tokenizer": (TokenizerLoader, "transformers"),
|
||||
@@ -727,6 +730,34 @@ class VocoderLoader(ComponentLoader):
|
||||
return vocoder.eval()
|
||||
|
||||
|
||||
class UpsamplerLoader(ComponentLoader):
|
||||
"""Loader for LTX-2 spatial/temporal upsampler."""
|
||||
|
||||
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
|
||||
config = get_diffusers_config(model=model_path)
|
||||
class_name = config.pop("_class_name", None) or "LTX2LatentUpsampler"
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
target_device = get_local_torch_device()
|
||||
|
||||
precision = getattr(
|
||||
fastvideo_args.pipeline_config, "vae_precision", "bf16"
|
||||
)
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[precision]):
|
||||
upsampler = model_cls(config).to(target_device)
|
||||
|
||||
safetensors_list = glob.glob(
|
||||
os.path.join(str(model_path), "*.safetensors")
|
||||
)
|
||||
loaded: dict[str, torch.Tensor] = {}
|
||||
for sf_file in safetensors_list:
|
||||
loaded.update(safetensors_load_file(sf_file))
|
||||
|
||||
target_module = getattr(upsampler, "model", upsampler)
|
||||
target_module.load_state_dict(loaded, strict=False)
|
||||
return upsampler.eval()
|
||||
|
||||
|
||||
class TransformerLoader(ComponentLoader):
|
||||
"""Loader for transformer."""
|
||||
|
||||
@@ -896,14 +927,10 @@ class UpsamplerLoader(ComponentLoader):
|
||||
"Only diffusers format is supported."
|
||||
)
|
||||
|
||||
try:
|
||||
upsampler_cfg = deepcopy(fastvideo_args.pipeline_config.upsampler_config[0])
|
||||
upsampler_cfg.update_model_config(config_dict)
|
||||
except Exception as e:
|
||||
upsampler_cfg = deepcopy(fastvideo_args.pipeline_config.upsampler_config[1])
|
||||
upsampler_cfg.update_model_config(config_dict)
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
upsampler_cfg = get_upsampler_config(
|
||||
class_name, config_dict, fastvideo_args.pipeline_config
|
||||
)
|
||||
model = model_cls(upsampler_cfg)
|
||||
|
||||
target_device = get_local_torch_device()
|
||||
@@ -914,15 +941,21 @@ class UpsamplerLoader(ComponentLoader):
|
||||
os.path.join(str(model_path), "*.safetensors"))
|
||||
if not safetensors_list:
|
||||
raise ValueError(f"No safetensors files found in {model_path}")
|
||||
|
||||
|
||||
if len(safetensors_list) == 1:
|
||||
loaded = safetensors_load_file(safetensors_list[0])
|
||||
else:
|
||||
loaded = {}
|
||||
for sf_file in safetensors_list:
|
||||
loaded.update(safetensors_load_file(sf_file))
|
||||
|
||||
model.load_state_dict(loaded, strict=True)
|
||||
|
||||
# LTX2 latent upsamplers typically store weights without "model." prefix.
|
||||
target_module = getattr(model, "model", model)
|
||||
if loaded and all(k.startswith("model.") for k in loaded.keys()):
|
||||
stripped = {k[len("model.") :]: v for k, v in loaded.items()}
|
||||
target_module.load_state_dict(stripped, strict=True)
|
||||
else:
|
||||
target_module.load_state_dict(loaded, strict=True)
|
||||
|
||||
return model.eval()
|
||||
|
||||
@@ -1005,4 +1038,4 @@ class PipelineComponentLoader:
|
||||
)
|
||||
|
||||
# Load the module
|
||||
return loader.load(component_model_path, fastvideo_args)
|
||||
return loader.load(component_model_path, fastvideo_args)
|
||||
|
||||
@@ -84,6 +84,10 @@ _AUDIO_MODELS = {
|
||||
"LTX2Vocoder": ("audio", "ltx2_audio_vae", "LTX2Vocoder"),
|
||||
}
|
||||
|
||||
_UPSAMPLER_MODELS = {
|
||||
"LTX2LatentUpsampler": ("upsamplers", "ltx2_upsampler", "LTX2LatentUpsampler"),
|
||||
}
|
||||
|
||||
_SCHEDULERS = {
|
||||
"FlowMatchEulerDiscreteScheduler":
|
||||
("schedulers", "scheduling_flow_match_euler_discrete",
|
||||
@@ -111,6 +115,7 @@ _LEGACY_FAST_VIDEO_MODELS = {
|
||||
**_IMAGE_ENCODER_MODELS,
|
||||
**_VAE_MODELS,
|
||||
**_AUDIO_MODELS,
|
||||
**_UPSAMPLER_MODELS,
|
||||
**_SCHEDULERS,
|
||||
**_UPSAMPLERS,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.models.upsamplers.ltx2_upsampler import (
|
||||
BlurDownsample,
|
||||
LatentUpsampler,
|
||||
LatentUpsamplerConfigurator,
|
||||
LTX2LatentUpsampler,
|
||||
PixelShuffleND,
|
||||
ResBlock,
|
||||
SpatialRationalResampler,
|
||||
upsample_video,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"BlurDownsample",
|
||||
"LatentUpsampler",
|
||||
"LatentUpsamplerConfigurator",
|
||||
"LTX2LatentUpsampler",
|
||||
"PixelShuffleND",
|
||||
"ResBlock",
|
||||
"SpatialRationalResampler",
|
||||
"upsample_video",
|
||||
]
|
||||
@@ -0,0 +1,319 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 latent upsampler (spatial/temporal) implementation.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Any, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
class PixelShuffleND(nn.Module):
|
||||
"""N-dimensional pixel shuffle for upsampling."""
|
||||
|
||||
def __init__(self, dims: int, upscale_factors: Tuple[int, int, int] = (2, 2, 2)) -> None:
|
||||
super().__init__()
|
||||
if dims not in (1, 2, 3):
|
||||
raise ValueError("dims must be 1, 2, or 3")
|
||||
self.dims = dims
|
||||
self.upscale_factors = upscale_factors
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if self.dims == 3:
|
||||
return rearrange(
|
||||
x,
|
||||
"b (c p1 p2 p3) d h w -> b c (d p1) (h p2) (w p3)",
|
||||
p1=self.upscale_factors[0],
|
||||
p2=self.upscale_factors[1],
|
||||
p3=self.upscale_factors[2],
|
||||
)
|
||||
if self.dims == 2:
|
||||
return rearrange(
|
||||
x,
|
||||
"b (c p1 p2) h w -> b c (h p1) (w p2)",
|
||||
p1=self.upscale_factors[0],
|
||||
p2=self.upscale_factors[1],
|
||||
)
|
||||
if self.dims == 1:
|
||||
return rearrange(
|
||||
x,
|
||||
"b (c p1) f h w -> b c (f p1) h w",
|
||||
p1=self.upscale_factors[0],
|
||||
)
|
||||
raise ValueError(f"Unsupported dims: {self.dims}")
|
||||
|
||||
|
||||
class BlurDownsample(nn.Module):
|
||||
"""
|
||||
Anti-aliased spatial downsampling by integer stride using a fixed separable binomial kernel.
|
||||
Applies only on H,W. Works for dims=2 or dims=3 (per-frame).
|
||||
"""
|
||||
|
||||
def __init__(self, dims: int, stride: int, kernel_size: int = 5) -> None:
|
||||
super().__init__()
|
||||
if dims not in (2, 3):
|
||||
raise ValueError("dims must be 2 or 3")
|
||||
if stride < 1:
|
||||
raise ValueError("stride must be >= 1")
|
||||
if kernel_size < 3 or kernel_size % 2 != 1:
|
||||
raise ValueError("kernel_size must be an odd integer >= 3")
|
||||
|
||||
self.dims = dims
|
||||
self.stride = stride
|
||||
self.kernel_size = kernel_size
|
||||
|
||||
k = torch.tensor([math.comb(kernel_size - 1, idx) for idx in range(kernel_size)])
|
||||
k2d = k[:, None] @ k[None, :]
|
||||
k2d = (k2d / k2d.sum()).float()
|
||||
self.register_buffer("kernel", k2d[None, None, :, :])
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if self.stride == 1:
|
||||
return x
|
||||
if self.dims == 2:
|
||||
return self._apply_2d(x)
|
||||
b, _, f, _, _ = x.shape
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = self._apply_2d(x)
|
||||
h2, w2 = x.shape[-2:]
|
||||
return rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f, h=h2, w=w2)
|
||||
|
||||
def _apply_2d(self, x2d: torch.Tensor) -> torch.Tensor:
|
||||
c = x2d.shape[1]
|
||||
weight = self.kernel.expand(c, 1, self.kernel_size, self.kernel_size)
|
||||
return nn.functional.conv2d(
|
||||
x2d,
|
||||
weight=weight,
|
||||
bias=None,
|
||||
stride=self.stride,
|
||||
padding=self.kernel_size // 2,
|
||||
groups=c,
|
||||
)
|
||||
|
||||
|
||||
def _rational_for_scale(scale: float) -> Tuple[int, int]:
|
||||
mapping = {0.75: (3, 4), 1.5: (3, 2), 2.0: (2, 1), 4.0: (4, 1)}
|
||||
if float(scale) not in mapping:
|
||||
raise ValueError(f"Unsupported scale {scale}. Choose from {list(mapping.keys())}")
|
||||
return mapping[float(scale)]
|
||||
|
||||
|
||||
class SpatialRationalResampler(nn.Module):
|
||||
"""
|
||||
Fully-learned rational spatial scaling: up by 'num' via PixelShuffle, then
|
||||
anti-aliased downsample by 'den' using fixed blur + stride. Operates on H,W only.
|
||||
For dims==3, work per-frame for spatial scaling (temporal axis untouched).
|
||||
"""
|
||||
|
||||
def __init__(self, mid_channels: int, scale: float) -> None:
|
||||
super().__init__()
|
||||
self.scale = float(scale)
|
||||
self.num, self.den = _rational_for_scale(self.scale)
|
||||
self.conv = nn.Conv2d(mid_channels, (self.num**2) * mid_channels, kernel_size=3, padding=1)
|
||||
self.pixel_shuffle = PixelShuffleND(2, upscale_factors=(self.num, self.num))
|
||||
self.blur_down = BlurDownsample(dims=2, stride=self.den)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
b, _, f, _, _ = x.shape
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = self.conv(x)
|
||||
x = self.pixel_shuffle(x)
|
||||
x = self.blur_down(x)
|
||||
return rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f)
|
||||
|
||||
|
||||
class ResBlock(nn.Module):
|
||||
"""Residual block with two convolutional layers, group norm, and SiLU."""
|
||||
|
||||
def __init__(self, channels: int, mid_channels: Optional[int] = None, dims: int = 3) -> None:
|
||||
super().__init__()
|
||||
if mid_channels is None:
|
||||
mid_channels = channels
|
||||
|
||||
conv = nn.Conv2d if dims == 2 else nn.Conv3d
|
||||
|
||||
self.conv1 = conv(channels, mid_channels, kernel_size=3, padding=1)
|
||||
self.norm1 = nn.GroupNorm(32, mid_channels)
|
||||
self.conv2 = conv(mid_channels, channels, kernel_size=3, padding=1)
|
||||
self.norm2 = nn.GroupNorm(32, channels)
|
||||
self.activation = nn.SiLU()
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
residual = x
|
||||
x = self.conv1(x)
|
||||
x = self.norm1(x)
|
||||
x = self.activation(x)
|
||||
x = self.conv2(x)
|
||||
x = self.norm2(x)
|
||||
x = self.activation(x + residual)
|
||||
return x
|
||||
|
||||
|
||||
class LatentUpsampler(nn.Module):
|
||||
"""
|
||||
Model to upsample VAE latents spatially and/or temporally.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 128,
|
||||
mid_channels: int = 512,
|
||||
num_blocks_per_stage: int = 4,
|
||||
dims: int = 3,
|
||||
spatial_upsample: bool = True,
|
||||
temporal_upsample: bool = False,
|
||||
spatial_scale: float = 2.0,
|
||||
rational_resampler: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.in_channels = in_channels
|
||||
self.mid_channels = mid_channels
|
||||
self.num_blocks_per_stage = num_blocks_per_stage
|
||||
self.dims = dims
|
||||
self.spatial_upsample = spatial_upsample
|
||||
self.temporal_upsample = temporal_upsample
|
||||
self.spatial_scale = float(spatial_scale)
|
||||
self.rational_resampler = rational_resampler
|
||||
|
||||
conv = nn.Conv2d if dims == 2 else nn.Conv3d
|
||||
|
||||
self.initial_conv = conv(in_channels, mid_channels, kernel_size=3, padding=1)
|
||||
self.initial_norm = nn.GroupNorm(32, mid_channels)
|
||||
self.initial_activation = nn.SiLU()
|
||||
|
||||
self.res_blocks = nn.ModuleList([ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)])
|
||||
|
||||
if spatial_upsample and temporal_upsample:
|
||||
self.upsampler = nn.Sequential(
|
||||
nn.Conv3d(mid_channels, 8 * mid_channels, kernel_size=3, padding=1),
|
||||
PixelShuffleND(3),
|
||||
)
|
||||
elif spatial_upsample:
|
||||
if rational_resampler:
|
||||
self.upsampler = SpatialRationalResampler(mid_channels=mid_channels, scale=self.spatial_scale)
|
||||
else:
|
||||
self.upsampler = nn.Sequential(
|
||||
nn.Conv2d(mid_channels, 4 * mid_channels, kernel_size=3, padding=1),
|
||||
PixelShuffleND(2),
|
||||
)
|
||||
elif temporal_upsample:
|
||||
self.upsampler = nn.Sequential(
|
||||
nn.Conv3d(mid_channels, 2 * mid_channels, kernel_size=3, padding=1),
|
||||
PixelShuffleND(1),
|
||||
)
|
||||
else:
|
||||
raise ValueError("Either spatial_upsample or temporal_upsample must be True")
|
||||
|
||||
self.post_upsample_res_blocks = nn.ModuleList(
|
||||
[ResBlock(mid_channels, dims=dims) for _ in range(num_blocks_per_stage)]
|
||||
)
|
||||
|
||||
self.final_conv = conv(mid_channels, in_channels, kernel_size=3, padding=1)
|
||||
|
||||
def forward(self, latent: torch.Tensor) -> torch.Tensor:
|
||||
b, _, f, _, _ = latent.shape
|
||||
|
||||
if self.dims == 2:
|
||||
x = rearrange(latent, "b c f h w -> (b f) c h w")
|
||||
x = self.initial_conv(x)
|
||||
x = self.initial_norm(x)
|
||||
x = self.initial_activation(x)
|
||||
|
||||
for block in self.res_blocks:
|
||||
x = block(x)
|
||||
|
||||
x = self.upsampler(x)
|
||||
|
||||
for block in self.post_upsample_res_blocks:
|
||||
x = block(x)
|
||||
|
||||
x = self.final_conv(x)
|
||||
x = rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f)
|
||||
else:
|
||||
x = self.initial_conv(latent)
|
||||
x = self.initial_norm(x)
|
||||
x = self.initial_activation(x)
|
||||
|
||||
for block in self.res_blocks:
|
||||
x = block(x)
|
||||
|
||||
if self.temporal_upsample:
|
||||
x = self.upsampler(x)
|
||||
x = x[:, :, 1:, :, :]
|
||||
elif isinstance(self.upsampler, SpatialRationalResampler):
|
||||
x = self.upsampler(x)
|
||||
else:
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = self.upsampler(x)
|
||||
x = rearrange(x, "(b f) c h w -> b c f h w", b=b, f=f)
|
||||
|
||||
for block in self.post_upsample_res_blocks:
|
||||
x = block(x)
|
||||
|
||||
x = self.final_conv(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class LatentUpsamplerConfigurator:
|
||||
"""Configurator for LatentUpsampler from a config dict."""
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> LatentUpsampler:
|
||||
cfg = dict(config)
|
||||
cfg.pop("_class_name", None)
|
||||
if "upsampler" in cfg and isinstance(cfg["upsampler"], dict):
|
||||
cfg = cfg["upsampler"]
|
||||
|
||||
return LatentUpsampler(
|
||||
in_channels=cfg.get("in_channels", 128),
|
||||
mid_channels=cfg.get("mid_channels", 512),
|
||||
num_blocks_per_stage=cfg.get("num_blocks_per_stage", 4),
|
||||
dims=cfg.get("dims", 3),
|
||||
spatial_upsample=cfg.get("spatial_upsample", True),
|
||||
temporal_upsample=cfg.get("temporal_upsample", False),
|
||||
spatial_scale=cfg.get("spatial_scale", 2.0),
|
||||
rational_resampler=cfg.get("rational_resampler", False),
|
||||
)
|
||||
|
||||
|
||||
class LTX2LatentUpsampler(nn.Module):
|
||||
"""Public wrapper for the LTX-2 latent upsampler."""
|
||||
|
||||
def __init__(self, config: dict[str, Any]):
|
||||
super().__init__()
|
||||
self.model: LatentUpsampler = LatentUpsamplerConfigurator.from_config(config)
|
||||
|
||||
def forward(self, latent: torch.Tensor) -> torch.Tensor:
|
||||
return self.model(latent)
|
||||
|
||||
|
||||
def upsample_video(latent: torch.Tensor, video_encoder: Any, upsampler: LatentUpsampler) -> torch.Tensor:
|
||||
"""
|
||||
Upsample a latent tensor with normalization based on the video encoder's per-channel statistics.
|
||||
"""
|
||||
if not hasattr(video_encoder, "per_channel_statistics"):
|
||||
raise ValueError("video_encoder must expose per_channel_statistics for normalization")
|
||||
stats = video_encoder.per_channel_statistics
|
||||
latent = stats.un_normalize(latent)
|
||||
latent = upsampler(latent)
|
||||
latent = stats.normalize(latent)
|
||||
return latent
|
||||
|
||||
|
||||
__all__ = [
|
||||
"PixelShuffleND",
|
||||
"BlurDownsample",
|
||||
"SpatialRationalResampler",
|
||||
"ResBlock",
|
||||
"LatentUpsampler",
|
||||
"LatentUpsamplerConfigurator",
|
||||
"LTX2LatentUpsampler",
|
||||
"upsample_video",
|
||||
]
|
||||
@@ -17,7 +17,9 @@ import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.models.vaes.common import DiagonalGaussianDistribution
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# =============================================================================
|
||||
# Enums
|
||||
@@ -1285,6 +1287,7 @@ class VideoEncoder(nn.Module):
|
||||
|
||||
def forward(self, sample: torch.Tensor) -> torch.Tensor:
|
||||
frames_count = sample.shape[2]
|
||||
logger.info(f"Frames count: {frames_count}")
|
||||
if ((frames_count - 1) % 8) != 0:
|
||||
raise ValueError(
|
||||
"Invalid number of frames: Encode input must have 1 + 8 * x frames "
|
||||
|
||||
@@ -11,17 +11,17 @@ from transformers import AutoTokenizer
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import PipelineComponentLoader
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages import (DecodingStage, InputValidationStage,
|
||||
LTX2AudioDecodingStage,
|
||||
LTX2DenoisingStage,
|
||||
LTX2LatentPreparationStage,
|
||||
LTX2TextEncodingStage)
|
||||
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
|
||||
from fastvideo.pipelines.stages import (
|
||||
DecodingStage, InputValidationStage, LTX2AudioDecodingStage,
|
||||
LTX2DenoisingStage, LTX2LatentPreparationStage, LTX2RefineInitStage,
|
||||
LTX2RefineLoRAStage, LTX2UpsampleStage, STAGE_2_DISTILLED_SIGMA_VALUES,
|
||||
LTX2TextEncodingStage)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LTX2Pipeline(ComposedPipelineBase):
|
||||
class LTX2Pipeline(LoRAPipeline):
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder",
|
||||
@@ -33,6 +33,8 @@ class LTX2Pipeline(ComposedPipelineBase):
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
refine_enabled = fastvideo_args.ltx2_refine_enabled
|
||||
|
||||
self.add_stage(
|
||||
stage_name="input_validation_stage",
|
||||
stage=InputValidationStage(),
|
||||
@@ -46,6 +48,12 @@ class LTX2Pipeline(ComposedPipelineBase):
|
||||
),
|
||||
)
|
||||
|
||||
if refine_enabled:
|
||||
self.add_stage(
|
||||
stage_name="ltx2_refine_init_stage",
|
||||
stage=LTX2RefineInitStage(),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="latent_preparation_stage",
|
||||
stage=LTX2LatentPreparationStage(
|
||||
@@ -58,6 +66,54 @@ class LTX2Pipeline(ComposedPipelineBase):
|
||||
transformer=self.get_module("transformer"), ),
|
||||
)
|
||||
|
||||
if refine_enabled:
|
||||
stage2_sigmas = STAGE_2_DISTILLED_SIGMA_VALUES
|
||||
stage2_steps = fastvideo_args.ltx2_refine_num_inference_steps
|
||||
expected_steps = len(stage2_sigmas) - 1
|
||||
if stage2_steps != expected_steps:
|
||||
logger.warning(
|
||||
"ltx2_refine_num_inference_steps=%s does not match distilled schedule; "
|
||||
"using %s steps to align with stage2 sigmas.",
|
||||
stage2_steps,
|
||||
expected_steps,
|
||||
)
|
||||
stage2_steps = expected_steps
|
||||
|
||||
transformer_refine = self.get_module("transformer_refine",
|
||||
self.get_module("transformer"))
|
||||
|
||||
self.add_stage(
|
||||
stage_name="ltx2_upsample_stage",
|
||||
stage=LTX2UpsampleStage(
|
||||
upsampler=self.get_module("spatial_upsampler"),
|
||||
vae=self.get_module("vae"),
|
||||
transformer=transformer_refine,
|
||||
sigmas=stage2_sigmas,
|
||||
add_noise=fastvideo_args.ltx2_refine_add_noise,
|
||||
),
|
||||
)
|
||||
|
||||
if fastvideo_args.ltx2_refine_lora_path:
|
||||
self.add_stage(
|
||||
stage_name="ltx2_refine_lora_stage",
|
||||
stage=LTX2RefineLoRAStage(
|
||||
pipeline=self,
|
||||
lora_path=fastvideo_args.ltx2_refine_lora_path,
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="ltx2_refine_denoising_stage",
|
||||
stage=LTX2DenoisingStage(
|
||||
transformer=transformer_refine,
|
||||
sigmas_override=stage2_sigmas,
|
||||
num_inference_steps_override=stage2_steps,
|
||||
force_guidance_scale=fastvideo_args.
|
||||
ltx2_refine_guidance_scale,
|
||||
initial_audio_latents_key="ltx2_audio_latents",
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="audio_decoding_stage",
|
||||
stage=LTX2AudioDecodingStage(
|
||||
@@ -72,6 +128,31 @@ class LTX2Pipeline(ComposedPipelineBase):
|
||||
)
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
if fastvideo_args.debug_model_sums:
|
||||
os.environ["LTX2_PIPELINE_DEBUG_LOG"] = "1"
|
||||
if fastvideo_args.debug_model_sums_path:
|
||||
os.environ[
|
||||
"LTX2_PIPELINE_DEBUG_PATH"] = fastvideo_args.debug_model_sums_path
|
||||
else:
|
||||
logger.warning(
|
||||
"debug_model_sums is enabled but debug_model_sums_path is not set; no model sums will be logged."
|
||||
)
|
||||
else:
|
||||
os.environ.pop("LTX2_PIPELINE_DEBUG_LOG", None)
|
||||
os.environ.pop("LTX2_PIPELINE_DEBUG_PATH", None)
|
||||
if fastvideo_args.debug_model_detail:
|
||||
os.environ["LTX2_DEBUG_DETAIL"] = "1"
|
||||
if fastvideo_args.debug_model_detail_path:
|
||||
os.environ[
|
||||
"LTX2_PIPELINE_DEBUG_DETAIL_PATH"] = fastvideo_args.debug_model_detail_path
|
||||
else:
|
||||
logger.warning(
|
||||
"debug_model_detail is enabled but debug_model_detail_path is not set; no detailed hooks will be logged."
|
||||
)
|
||||
else:
|
||||
os.environ.pop("LTX2_DEBUG_DETAIL", None)
|
||||
os.environ.pop("LTX2_PIPELINE_DEBUG_DETAIL_PATH", None)
|
||||
|
||||
tokenizer = self.get_module("tokenizer")
|
||||
if tokenizer is not None:
|
||||
tokenizer.padding_side = "left"
|
||||
@@ -86,6 +167,51 @@ class LTX2Pipeline(ComposedPipelineBase):
|
||||
model_index = self._load_config(self.model_path)
|
||||
logger.info("Loading pipeline modules from config: %s", model_index)
|
||||
|
||||
# Apply optional FastVideo-specific refine defaults embedded in model_index.json.
|
||||
def _resolve_refine_path(value: str | None) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
if os.path.isabs(value):
|
||||
return value
|
||||
candidate = os.path.join(self.model_path, value)
|
||||
if os.path.exists(candidate):
|
||||
return candidate
|
||||
return value
|
||||
|
||||
if model_index.get("fastvideo_refine_enabled") is True:
|
||||
if fastvideo_args.refine_enabled is None:
|
||||
fastvideo_args.ltx2_refine_enabled = True
|
||||
if fastvideo_args.refine_upsampler_path is None and fastvideo_args.ltx2_refine_upsampler_path is None:
|
||||
fastvideo_args.ltx2_refine_upsampler_path = _resolve_refine_path(
|
||||
model_index.get("fastvideo_refine_upsampler_path"))
|
||||
if fastvideo_args.ltx2_refine_upsampler_path is None and "spatial_upsampler" in model_index:
|
||||
fastvideo_args.ltx2_refine_upsampler_path = _resolve_refine_path(
|
||||
"spatial_upsampler")
|
||||
if fastvideo_args.refine_transformer_path is None and fastvideo_args.ltx2_refine_transformer_path is None:
|
||||
fastvideo_args.ltx2_refine_transformer_path = _resolve_refine_path(
|
||||
model_index.get("fastvideo_refine_transformer_path"))
|
||||
if fastvideo_args.refine_lora_path is None and fastvideo_args.ltx2_refine_lora_path is None:
|
||||
fastvideo_args.ltx2_refine_lora_path = _resolve_refine_path(
|
||||
model_index.get("fastvideo_refine_lora_path"))
|
||||
if fastvideo_args.refine_num_inference_steps is None and model_index.get(
|
||||
"fastvideo_refine_num_inference_steps") is not None:
|
||||
fastvideo_args.ltx2_refine_num_inference_steps = int(
|
||||
model_index["fastvideo_refine_num_inference_steps"])
|
||||
if fastvideo_args.refine_guidance_scale is None and model_index.get(
|
||||
"fastvideo_refine_guidance_scale") is not None:
|
||||
fastvideo_args.ltx2_refine_guidance_scale = float(
|
||||
model_index["fastvideo_refine_guidance_scale"])
|
||||
if fastvideo_args.refine_add_noise is None and model_index.get(
|
||||
"fastvideo_refine_add_noise") is not None:
|
||||
fastvideo_args.ltx2_refine_add_noise = bool(
|
||||
model_index["fastvideo_refine_add_noise"])
|
||||
if fastvideo_args.refine_noise_path is None and fastvideo_args.ltx2_refine_noise_path is None:
|
||||
fastvideo_args.ltx2_refine_noise_path = _resolve_refine_path(
|
||||
model_index.get("fastvideo_refine_noise_path"))
|
||||
if fastvideo_args.refine_audio_noise_path is None and fastvideo_args.ltx2_refine_audio_noise_path is None:
|
||||
fastvideo_args.ltx2_refine_audio_noise_path = _resolve_refine_path(
|
||||
model_index.get("fastvideo_refine_audio_noise_path"))
|
||||
|
||||
model_index.pop("_class_name")
|
||||
model_index.pop("_diffusers_version")
|
||||
model_index.pop("workload_type", None)
|
||||
@@ -144,6 +270,50 @@ class LTX2Pipeline(ComposedPipelineBase):
|
||||
raise ValueError(
|
||||
f"Required module {module_name} was not loaded properly")
|
||||
|
||||
if fastvideo_args.ltx2_refine_enabled:
|
||||
upsampler_path = fastvideo_args.ltx2_refine_upsampler_path
|
||||
if upsampler_path is None:
|
||||
raise ValueError(
|
||||
"ltx2_refine_enabled is True but ltx2_refine_upsampler_path was not provided."
|
||||
)
|
||||
if not os.path.isdir(upsampler_path):
|
||||
raise ValueError(
|
||||
"ltx2_refine_upsampler_path must be a directory containing Diffusers-style "
|
||||
f"upsampler weights; got {upsampler_path}")
|
||||
config_path = os.path.join(upsampler_path, "config.json")
|
||||
if not os.path.exists(config_path):
|
||||
raise ValueError(
|
||||
"ltx2_refine_upsampler_path must contain a Diffusers config.json; "
|
||||
f"missing {config_path}")
|
||||
if loaded_modules is not None and "spatial_upsampler" in loaded_modules:
|
||||
modules["spatial_upsampler"] = loaded_modules[
|
||||
"spatial_upsampler"]
|
||||
else:
|
||||
modules[
|
||||
"spatial_upsampler"] = PipelineComponentLoader.load_module(
|
||||
module_name="spatial_upsampler",
|
||||
component_model_path=upsampler_path,
|
||||
transformers_or_diffusers="diffusers",
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
logger.info("Loaded module spatial_upsampler from %s",
|
||||
upsampler_path)
|
||||
|
||||
if loaded_modules is not None and "transformer_refine" in loaded_modules:
|
||||
modules["transformer_refine"] = loaded_modules[
|
||||
"transformer_refine"]
|
||||
elif fastvideo_args.ltx2_refine_transformer_path:
|
||||
modules[
|
||||
"transformer_refine"] = PipelineComponentLoader.load_module(
|
||||
module_name="transformer",
|
||||
component_model_path=fastvideo_args.
|
||||
ltx2_refine_transformer_path,
|
||||
transformers_or_diffusers="diffusers",
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
logger.info("Loaded module transformer_refine from %s",
|
||||
fastvideo_args.ltx2_refine_transformer_path)
|
||||
|
||||
return modules
|
||||
|
||||
|
||||
|
||||
@@ -155,6 +155,8 @@ class ComposedPipelineBase(ABC):
|
||||
transformer_module_2, **compile_kwargs)
|
||||
logger.info("Torch Compile enabled for DiT")
|
||||
|
||||
self._maybe_attach_module_sum_hooks()
|
||||
|
||||
if not self.fastvideo_args.training_mode:
|
||||
logger.info("Creating pipeline stages...")
|
||||
self.create_pipeline_stages(self.fastvideo_args)
|
||||
@@ -224,6 +226,100 @@ class ComposedPipelineBase(ABC):
|
||||
def add_module(self, module_name: str, module: Any):
|
||||
self.modules[module_name] = module
|
||||
|
||||
def _maybe_attach_module_sum_hooks(self) -> None:
|
||||
args = self.fastvideo_args
|
||||
if args is None or not getattr(args, "debug_module_sums", False):
|
||||
return
|
||||
if getattr(self, "_debug_module_sums_attached", False):
|
||||
return
|
||||
|
||||
log_path = getattr(args, "debug_module_sums_path", None)
|
||||
if not log_path:
|
||||
log_path = os.path.join("outputs", "debug", "module_sums.log")
|
||||
logger.warning(
|
||||
"debug_module_sums is enabled but debug_module_sums_path is not set; "
|
||||
"defaulting to %s",
|
||||
log_path,
|
||||
)
|
||||
|
||||
include = getattr(args, "debug_module_sums_include", None) or []
|
||||
exclude = getattr(args, "debug_module_sums_exclude", None) or []
|
||||
|
||||
def _matches(name: str) -> bool:
|
||||
include_ok = (not include) or any(key in name for key in include)
|
||||
exclude_ok = (not exclude) or not any(key in name
|
||||
for key in exclude)
|
||||
return include_ok and exclude_ok
|
||||
|
||||
def _sum_output(value: object) -> float | None:
|
||||
if isinstance(value, torch.Tensor):
|
||||
return float(value.detach().sum(dtype=torch.float32).item())
|
||||
if isinstance(value, dict):
|
||||
total = 0.0
|
||||
found = False
|
||||
for item in value.values():
|
||||
summed = _sum_output(item)
|
||||
if summed is not None:
|
||||
total += summed
|
||||
found = True
|
||||
return total if found else None
|
||||
if isinstance(value, list | tuple):
|
||||
total = 0.0
|
||||
found = False
|
||||
for item in value:
|
||||
summed = _sum_output(item)
|
||||
if summed is not None:
|
||||
total += summed
|
||||
found = True
|
||||
return total if found else None
|
||||
return None
|
||||
|
||||
def _write_line(line: str) -> None:
|
||||
log_dir = os.path.dirname(log_path)
|
||||
if log_dir:
|
||||
os.makedirs(log_dir, exist_ok=True)
|
||||
with open(log_path, "a", encoding="utf-8") as f:
|
||||
f.write(line + "\n")
|
||||
|
||||
handles: list[torch.utils.hooks.RemovableHandle] = []
|
||||
|
||||
for root_name, module in self.modules.items():
|
||||
if module is None or not isinstance(module, torch.nn.Module):
|
||||
continue
|
||||
for name, submodule in module.named_modules():
|
||||
full_name = root_name if not name else f"{root_name}.{name}"
|
||||
if not _matches(full_name):
|
||||
continue
|
||||
if getattr(submodule, "_fastvideo_module_sum_hooked", False):
|
||||
continue
|
||||
# Only hook modules with direct parameters to avoid excessive noise.
|
||||
is_root = name == ""
|
||||
if not is_root and not any(
|
||||
True for _ in submodule.parameters(recurse=False)):
|
||||
continue
|
||||
|
||||
def _hook_factory(module_name: str, module_type: str):
|
||||
|
||||
def _hook(_module, _inputs, outputs): # noqa: ANN001
|
||||
summed = _sum_output(outputs)
|
||||
if summed is None:
|
||||
return
|
||||
line = (f"fastvideo:module={module_name} "
|
||||
f"class={module_type} out_sum={summed:.6f}")
|
||||
_write_line(line)
|
||||
|
||||
return _hook
|
||||
|
||||
handle = submodule.register_forward_hook(
|
||||
_hook_factory(full_name, submodule.__class__.__name__))
|
||||
handles.append(handle)
|
||||
submodule._fastvideo_module_sum_hooked = True
|
||||
|
||||
self._debug_module_sums_attached = True
|
||||
self._debug_module_sum_handles = handles
|
||||
logger.info("Attached %s module sum hooks for recursive logging",
|
||||
len(handles))
|
||||
|
||||
def _load_config(self, model_path: str) -> dict[str, Any]:
|
||||
model_path = maybe_download_model(self.model_path)
|
||||
self.model_path = model_path
|
||||
|
||||
@@ -30,6 +30,9 @@ from fastvideo.pipelines.stages.ltx2_denoising import LTX2DenoisingStage
|
||||
from fastvideo.pipelines.stages.ltx2_latent_preparation import (
|
||||
LTX2LatentPreparationStage)
|
||||
from fastvideo.pipelines.stages.ltx2_text_encoding import LTX2TextEncodingStage
|
||||
from fastvideo.pipelines.stages.ltx2_refine import (
|
||||
LTX2RefineInitStage, LTX2RefineLoRAStage, LTX2UpsampleStage,
|
||||
STAGE_2_DISTILLED_SIGMA_VALUES)
|
||||
from fastvideo.pipelines.stages.matrixgame_denoising import (
|
||||
MatrixGameCausalDenoisingStage)
|
||||
from fastvideo.pipelines.stages.hyworld_denoising import HYWorldDenoisingStage
|
||||
@@ -58,6 +61,10 @@ __all__ = [
|
||||
"Cosmos25AutoLatentPreparationStage",
|
||||
"LTX2LatentPreparationStage",
|
||||
"LTX2AudioDecodingStage",
|
||||
"LTX2RefineInitStage",
|
||||
"LTX2RefineLoRAStage",
|
||||
"LTX2UpsampleStage",
|
||||
"STAGE_2_DISTILLED_SIGMA_VALUES",
|
||||
"ConditioningStage",
|
||||
"DenoisingStage",
|
||||
"DmdDenoisingStage",
|
||||
|
||||
@@ -6,6 +6,7 @@ This module defines the abstract base classes for pipeline stages that can be
|
||||
composed to create complete diffusion pipelines.
|
||||
"""
|
||||
|
||||
import os
|
||||
import time
|
||||
import traceback
|
||||
from abc import ABC, abstractmethod
|
||||
@@ -21,6 +22,28 @@ from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _sum_tensor(value: object) -> float | None:
|
||||
if isinstance(value, torch.Tensor):
|
||||
return float(value.detach().sum(dtype=torch.float32).item())
|
||||
if isinstance(value, list | tuple):
|
||||
total = 0.0
|
||||
found = False
|
||||
for item in value:
|
||||
if isinstance(item, torch.Tensor):
|
||||
total += item.detach().sum(dtype=torch.float32).item()
|
||||
found = True
|
||||
return total if found else None
|
||||
return None
|
||||
|
||||
|
||||
def _write_debug_line(path: str, line: str) -> None:
|
||||
log_dir = os.path.dirname(path)
|
||||
if log_dir:
|
||||
os.makedirs(log_dir, exist_ok=True)
|
||||
with open(path, "a", encoding="utf-8") as f:
|
||||
f.write(line + "\n")
|
||||
|
||||
|
||||
class StageVerificationError(Exception):
|
||||
"""Exception raised when stage verification fails."""
|
||||
pass
|
||||
@@ -152,6 +175,7 @@ class PipelineStage(ABC):
|
||||
|
||||
try:
|
||||
result = self.forward(batch, fastvideo_args)
|
||||
self._maybe_log_stage_sums(stage_name, result, fastvideo_args)
|
||||
execution_time = time.perf_counter() - start_time
|
||||
logger.info("[%s] Execution completed in %s ms", stage_name,
|
||||
execution_time * 1000)
|
||||
@@ -167,6 +191,7 @@ class PipelineStage(ABC):
|
||||
else:
|
||||
# Direct execution (current behavior)
|
||||
result = self.forward(batch, fastvideo_args)
|
||||
self._maybe_log_stage_sums(stage_name, result, fastvideo_args)
|
||||
|
||||
if enable_verification:
|
||||
# Post-execution output verification
|
||||
@@ -180,6 +205,49 @@ class PipelineStage(ABC):
|
||||
|
||||
return result
|
||||
|
||||
def _maybe_log_stage_sums(
|
||||
self,
|
||||
stage_name: str,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> None:
|
||||
if not getattr(fastvideo_args, "debug_stage_sums", False):
|
||||
return
|
||||
|
||||
entries: dict[str, float] = {}
|
||||
|
||||
def _add(name: str, value: object) -> None:
|
||||
summed = _sum_tensor(value)
|
||||
if summed is not None:
|
||||
entries[name] = summed
|
||||
|
||||
_add("latents", batch.latents)
|
||||
_add("image_latent", batch.image_latent)
|
||||
_add("video_latent", batch.video_latent)
|
||||
_add("noise_pred", batch.noise_pred)
|
||||
_add("output", batch.output)
|
||||
_add("prompt_embeds", batch.prompt_embeds)
|
||||
_add("negative_prompt_embeds", batch.negative_prompt_embeds)
|
||||
|
||||
if isinstance(batch.extra, dict):
|
||||
for key, value in batch.extra.items():
|
||||
summed = _sum_tensor(value)
|
||||
if summed is not None:
|
||||
entries[f"extra.{key}"] = summed
|
||||
|
||||
if not entries:
|
||||
return
|
||||
|
||||
batch.logging_info.add_stage_metric(stage_name, "sums", entries)
|
||||
line = "fastvideo:stage={}".format(stage_name) + " " + " ".join(
|
||||
f"{key}={value:.6f}" for key, value in entries.items())
|
||||
|
||||
path = getattr(fastvideo_args, "debug_stage_sums_path", None)
|
||||
if path:
|
||||
_write_debug_line(path, line)
|
||||
else:
|
||||
logger.info("%s", line)
|
||||
|
||||
@abstractmethod
|
||||
def forward(
|
||||
self,
|
||||
|
||||
@@ -7,6 +7,7 @@ from __future__ import annotations
|
||||
|
||||
import math
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from tqdm.auto import tqdm
|
||||
@@ -80,9 +81,21 @@ def _ltx2_sigmas(
|
||||
class LTX2DenoisingStage(PipelineStage):
|
||||
"""Run the LTX-2 denoising loop over the sigma schedule."""
|
||||
|
||||
def __init__(self, transformer) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
transformer,
|
||||
*,
|
||||
sigmas_override: list[float] | None = None,
|
||||
num_inference_steps_override: int | None = None,
|
||||
force_guidance_scale: float | None = None,
|
||||
initial_audio_latents_key: str | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.sigmas_override = sigmas_override
|
||||
self.num_inference_steps_override = num_inference_steps_override
|
||||
self.force_guidance_scale = force_guidance_scale
|
||||
self.initial_audio_latents_key = initial_audio_latents_key
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -98,8 +111,14 @@ class LTX2DenoisingStage(PipelineStage):
|
||||
|
||||
neg_prompt_embeds = None
|
||||
neg_prompt_mask = None
|
||||
guidance_scale = batch.guidance_scale
|
||||
use_cfg = batch.do_classifier_free_guidance
|
||||
if self.force_guidance_scale is not None:
|
||||
guidance_scale = self.force_guidance_scale
|
||||
use_cfg = guidance_scale > 1.0
|
||||
|
||||
# Only load negative prompts if CFG is actually enabled
|
||||
if batch.do_classifier_free_guidance:
|
||||
if use_cfg:
|
||||
assert batch.negative_prompt_embeds is not None, (
|
||||
"CFG is enabled but negative_prompt_embeds is None")
|
||||
neg_prompt_embeds = batch.negative_prompt_embeds[0]
|
||||
@@ -121,22 +140,31 @@ class LTX2DenoisingStage(PipelineStage):
|
||||
) and not fastvideo_args.disable_autocast and (
|
||||
not disable_autocast)
|
||||
|
||||
# Use official distilled sigma schedule for 8 steps (distilled models)
|
||||
use_distilled_sigmas = os.getenv("LTX2_USE_DISTILLED_SIGMAS",
|
||||
"1") == "1"
|
||||
if use_distilled_sigmas and batch.num_inference_steps == 8:
|
||||
num_steps = self.num_inference_steps_override or batch.num_inference_steps
|
||||
if self.sigmas_override is not None:
|
||||
sigmas = torch.tensor(
|
||||
DISTILLED_SIGMA_VALUES,
|
||||
self.sigmas_override,
|
||||
device=latents.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
logger.info("[LTX2] Using official distilled sigma schedule")
|
||||
logger.info("[LTX2] Using overridden sigma schedule")
|
||||
else:
|
||||
sigmas = _ltx2_sigmas(
|
||||
steps=batch.num_inference_steps,
|
||||
latent=None,
|
||||
device=latents.device,
|
||||
)
|
||||
# Use official distilled sigma schedule for 8 steps (distilled models)
|
||||
use_distilled_sigmas = os.getenv("LTX2_USE_DISTILLED_SIGMAS",
|
||||
"1") == "1"
|
||||
if use_distilled_sigmas and num_steps == 8:
|
||||
sigmas = torch.tensor(
|
||||
DISTILLED_SIGMA_VALUES,
|
||||
device=latents.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
logger.info("[LTX2] Using official distilled sigma schedule")
|
||||
else:
|
||||
sigmas = _ltx2_sigmas(
|
||||
steps=num_steps,
|
||||
latent=None,
|
||||
device=latents.device,
|
||||
)
|
||||
if hasattr(self.transformer, "patchifier"):
|
||||
video_shape = VideoLatentShape.from_torch_shape(latents.shape)
|
||||
token_count = self.transformer.patchifier.get_token_count(
|
||||
@@ -144,7 +172,7 @@ class LTX2DenoisingStage(PipelineStage):
|
||||
else:
|
||||
token_count = 1
|
||||
timestep_template = torch.ones(
|
||||
(latents.shape[0], token_count),
|
||||
(latents.shape[0], token_count, 1),
|
||||
device=latents.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
@@ -154,7 +182,15 @@ class LTX2DenoisingStage(PipelineStage):
|
||||
audio_context_n = audio_neg_embeds[0] if audio_neg_embeds else None
|
||||
audio_latents = None
|
||||
audio_timestep_template = None
|
||||
if audio_context_p is not None:
|
||||
if self.initial_audio_latents_key is not None:
|
||||
audio_latents = batch.extra.get(self.initial_audio_latents_key)
|
||||
if audio_latents is not None:
|
||||
audio_timestep_template = torch.ones(
|
||||
(latents.shape[0], audio_latents.shape[2], 1),
|
||||
device=latents.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
if audio_context_p is not None and audio_latents is None:
|
||||
fps_value = batch.fps
|
||||
if isinstance(fps_value, list):
|
||||
fps_value = fps_value[0] if fps_value else None
|
||||
@@ -170,53 +206,69 @@ class LTX2DenoisingStage(PipelineStage):
|
||||
hop_length=DEFAULT_LTX2_AUDIO_HOP_LENGTH,
|
||||
audio_latent_downsample_factor=DEFAULT_LTX2_AUDIO_DOWNSAMPLE,
|
||||
)
|
||||
audio_generator = None
|
||||
if fastvideo_args.ltx2_initial_latent_path and batch.seed is not None:
|
||||
audio_generator = torch.Generator(
|
||||
device=latents.device).manual_seed(batch.seed)
|
||||
elif batch.generator is not None:
|
||||
if isinstance(batch.generator, list):
|
||||
audio_generator = batch.generator[0]
|
||||
else:
|
||||
audio_generator = batch.generator
|
||||
if audio_generator is not None and audio_generator.device.type != latents.device.type:
|
||||
if batch.seed is None:
|
||||
audio_generator = torch.Generator(device=latents.device)
|
||||
else:
|
||||
audio_generator = torch.Generator(
|
||||
device=latents.device).manual_seed(batch.seed)
|
||||
audio_patch_shape = (
|
||||
expected_shape = (
|
||||
audio_shape.batch,
|
||||
audio_shape.channels,
|
||||
audio_shape.frames,
|
||||
audio_shape.channels * audio_shape.mel_bins,
|
||||
audio_shape.mel_bins,
|
||||
)
|
||||
audio_latents_patch = torch.randn(
|
||||
audio_patch_shape,
|
||||
generator=audio_generator,
|
||||
audio_latent_path = fastvideo_args.ltx2_audio_latent_path
|
||||
audio_latents = self._load_audio_latents(
|
||||
audio_latent_path,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
if hasattr(self.transformer, "audio_patchifier"):
|
||||
audio_latents = self.transformer.audio_patchifier.unpatchify(
|
||||
audio_latents_patch, audio_shape)
|
||||
else:
|
||||
audio_latents = audio_latents_patch.view(
|
||||
expected_shape=expected_shape,
|
||||
) if audio_latent_path else None
|
||||
if audio_latents is None:
|
||||
audio_generator = None
|
||||
if fastvideo_args.ltx2_initial_latent_path and batch.seed is not None:
|
||||
audio_generator = torch.Generator(
|
||||
device=latents.device).manual_seed(batch.seed)
|
||||
elif batch.generator is not None:
|
||||
if isinstance(batch.generator, list):
|
||||
audio_generator = batch.generator[0]
|
||||
else:
|
||||
audio_generator = batch.generator
|
||||
if audio_generator is not None and audio_generator.device.type != latents.device.type:
|
||||
if batch.seed is None:
|
||||
audio_generator = torch.Generator(device=latents.device)
|
||||
else:
|
||||
audio_generator = torch.Generator(
|
||||
device=latents.device).manual_seed(batch.seed)
|
||||
audio_patch_shape = (
|
||||
audio_shape.batch,
|
||||
audio_shape.frames,
|
||||
audio_shape.channels,
|
||||
audio_shape.mel_bins,
|
||||
).permute(0, 2, 1, 3).contiguous()
|
||||
audio_shape.channels * audio_shape.mel_bins,
|
||||
)
|
||||
audio_latents_patch = torch.randn(
|
||||
audio_patch_shape,
|
||||
generator=audio_generator,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
if hasattr(self.transformer, "audio_patchifier"):
|
||||
audio_latents = self.transformer.audio_patchifier.unpatchify(
|
||||
audio_latents_patch, audio_shape)
|
||||
else:
|
||||
audio_latents = audio_latents_patch.view(
|
||||
audio_shape.batch,
|
||||
audio_shape.frames,
|
||||
audio_shape.channels,
|
||||
audio_shape.mel_bins,
|
||||
).permute(0, 2, 1, 3).contiguous()
|
||||
if audio_latent_path:
|
||||
self._save_audio_latents(audio_latent_path, audio_latents)
|
||||
audio_timestep_template = torch.ones(
|
||||
(latents.shape[0], audio_shape.frames),
|
||||
(latents.shape[0], audio_shape.frames, 1),
|
||||
device=latents.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
logger.info(
|
||||
"[LTX2] Denoising start: steps=%d dtype=%s guidance=%s "
|
||||
"sigmas_shape=%s latents_shape=%s",
|
||||
batch.num_inference_steps,
|
||||
num_steps,
|
||||
target_dtype,
|
||||
batch.guidance_scale,
|
||||
guidance_scale,
|
||||
tuple(sigmas.shape),
|
||||
tuple(latents.shape),
|
||||
)
|
||||
@@ -253,7 +305,7 @@ class LTX2DenoisingStage(PipelineStage):
|
||||
pos_audio = None
|
||||
|
||||
# Only run negative pass if CFG is enabled
|
||||
if batch.do_classifier_free_guidance:
|
||||
if use_cfg:
|
||||
neg_outputs = self.transformer(
|
||||
hidden_states=latents.to(target_dtype),
|
||||
encoder_hidden_states=neg_prompt_embeds,
|
||||
@@ -268,10 +320,10 @@ class LTX2DenoisingStage(PipelineStage):
|
||||
else:
|
||||
neg_denoised = neg_outputs
|
||||
neg_audio = None
|
||||
pos_denoised = pos_denoised + (batch.guidance_scale - 1) * (
|
||||
pos_denoised = pos_denoised + (guidance_scale - 1) * (
|
||||
pos_denoised - neg_denoised)
|
||||
if pos_audio is not None and neg_audio is not None:
|
||||
pos_audio = pos_audio + (batch.guidance_scale -
|
||||
pos_audio = pos_audio + (guidance_scale -
|
||||
1) * (pos_audio - neg_audio)
|
||||
|
||||
sigma_value = sigma.to(torch.float32) if isinstance(
|
||||
@@ -297,6 +349,43 @@ class LTX2DenoisingStage(PipelineStage):
|
||||
logger.info("[LTX2] Denoising done.")
|
||||
return batch
|
||||
|
||||
def _load_audio_latents(
|
||||
self,
|
||||
latent_path: str | None,
|
||||
*,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
expected_shape: tuple[int, ...],
|
||||
) -> torch.Tensor | None:
|
||||
if not latent_path:
|
||||
return None
|
||||
path = Path(latent_path)
|
||||
if not path.exists():
|
||||
return None
|
||||
payload = torch.load(path, map_location=device)
|
||||
if isinstance(payload, dict):
|
||||
latent = (payload.get("audio_latent") or payload.get("latent")
|
||||
or payload.get("audio"))
|
||||
else:
|
||||
latent = payload
|
||||
if not torch.is_tensor(latent):
|
||||
raise TypeError(f"Expected tensor audio latent in {path}")
|
||||
if tuple(latent.shape) != tuple(expected_shape):
|
||||
raise ValueError(
|
||||
f"Audio latent shape mismatch for {path}: expected {expected_shape}, got {tuple(latent.shape)}"
|
||||
)
|
||||
logger.info("[LTX2] Loaded audio latent from %s", path)
|
||||
return latent.to(device=device, dtype=dtype)
|
||||
|
||||
def _save_audio_latents(self, latent_path: str,
|
||||
latents: torch.Tensor) -> None:
|
||||
path = Path(latent_path)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if path.exists():
|
||||
return
|
||||
torch.save({"audio_latent": latents.detach().cpu()}, path)
|
||||
logger.info("[LTX2] Saved audio latent to %s", path)
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
Latent preparation stage for LTX-2 pipelines.
|
||||
"""
|
||||
|
||||
import math
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
@@ -11,6 +12,7 @@ from diffusers.utils.torch_utils import randn_tensor
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.ltx2 import VideoLatentShape
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
@@ -61,6 +63,26 @@ class LTX2LatentPreparationStage(PipelineStage):
|
||||
dtype = batch.prompt_embeds[0].dtype
|
||||
device = get_local_torch_device()
|
||||
generator = batch.generator
|
||||
if generator is not None:
|
||||
if isinstance(generator, list):
|
||||
if generator and generator[0].device.type != device.type:
|
||||
seeds = batch.seeds
|
||||
if seeds is None and batch.seed is not None:
|
||||
seeds = [batch.seed + i for i in range(len(generator))]
|
||||
if seeds is not None:
|
||||
generator = [
|
||||
torch.Generator(device=device).manual_seed(seed)
|
||||
for seed in seeds
|
||||
]
|
||||
batch.generator = generator
|
||||
else:
|
||||
if generator.device.type != device.type:
|
||||
if batch.seed is not None:
|
||||
generator = torch.Generator(device=device).manual_seed(
|
||||
batch.seed)
|
||||
else:
|
||||
generator = torch.Generator(device=device)
|
||||
batch.generator = generator
|
||||
latents = batch.latents
|
||||
num_frames = latent_num_frames if latent_num_frames is not None else batch.num_frames
|
||||
height = batch.height
|
||||
@@ -95,16 +117,16 @@ class LTX2LatentPreparationStage(PipelineStage):
|
||||
if loaded_latents is not None:
|
||||
latents = loaded_latents
|
||||
else:
|
||||
latents = randn_tensor(
|
||||
shape,
|
||||
latents = self._generate_initial_latents(
|
||||
shape=shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
self._save_initial_latent(latent_path, latents)
|
||||
else:
|
||||
latents = randn_tensor(
|
||||
shape,
|
||||
latents = self._generate_initial_latents(
|
||||
shape=shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
@@ -116,6 +138,39 @@ class LTX2LatentPreparationStage(PipelineStage):
|
||||
batch.raw_latent_shape = shape
|
||||
return batch
|
||||
|
||||
def _generate_initial_latents(
|
||||
self,
|
||||
shape: tuple[int, int, int, int, int],
|
||||
*,
|
||||
generator: torch.Generator | list[torch.Generator] | None,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
patchifier = getattr(self.transformer, "patchifier", None)
|
||||
if patchifier is None:
|
||||
return randn_tensor(
|
||||
shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
video_shape = VideoLatentShape.from_torch_shape(shape)
|
||||
patch_volume = math.prod(patchifier.patch_size)
|
||||
token_count = patchifier.get_token_count(video_shape)
|
||||
patch_shape = (
|
||||
shape[0],
|
||||
token_count,
|
||||
shape[1] * patch_volume,
|
||||
)
|
||||
patch_noise = randn_tensor(
|
||||
patch_shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
return patchifier.unpatchify(patch_noise, video_shape)
|
||||
|
||||
def _adjust_video_length(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> int | None:
|
||||
if not fastvideo_args.pipeline_config.vae_config.use_temporal_scaling_frames:
|
||||
|
||||
@@ -0,0 +1,355 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
LTX-2 refinement stages for 2x spatial upscaling + distilled denoising.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
import weakref
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.ltx2 import AudioLatentShape, VideoLatentShape
|
||||
from fastvideo.models.upsamplers import upsample_video
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Reduced schedule for super-resolution stage 2 (subset of distilled values)
|
||||
# From LTX-2/packages/ltx-pipelines/src/ltx_pipelines/utils/constants.py
|
||||
STAGE_2_DISTILLED_SIGMA_VALUES = [0.909375, 0.725, 0.421875, 0.0]
|
||||
|
||||
|
||||
class LTX2RefineInitStage(PipelineStage):
|
||||
"""Prepare stage-1 resolution for LTX-2 2x spatial refinement."""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
if not fastvideo_args.ltx2_refine_enabled:
|
||||
return batch
|
||||
|
||||
height = batch.height
|
||||
width = batch.width
|
||||
if height is None or width is None:
|
||||
raise ValueError(
|
||||
"Height and width must be provided for LTX-2 refinement.")
|
||||
if isinstance(height, list) or isinstance(width, list):
|
||||
raise ValueError("LTX-2 refinement expects scalar height/width.")
|
||||
|
||||
if height % 2 != 0 or width % 2 != 0:
|
||||
raise ValueError(
|
||||
"LTX-2 refinement requires even height/width so stage1 can be half resolution."
|
||||
)
|
||||
|
||||
spatial_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
|
||||
stage1_height = height // 2
|
||||
stage1_width = width // 2
|
||||
if stage1_height % spatial_ratio != 0 or stage1_width % spatial_ratio != 0:
|
||||
raise ValueError(
|
||||
f"LTX-2 refinement requires height/width divisible by {2 * spatial_ratio} "
|
||||
f"(got {height}x{width}).")
|
||||
|
||||
batch.extra["ltx2_refine_target_height"] = height
|
||||
batch.extra["ltx2_refine_target_width"] = width
|
||||
batch.height = stage1_height
|
||||
batch.width = stage1_width
|
||||
|
||||
logger.info(
|
||||
"[LTX2] Refinement enabled: stage1=%dx%d stage2=%dx%d",
|
||||
stage1_width,
|
||||
stage1_height,
|
||||
width,
|
||||
height,
|
||||
)
|
||||
return batch
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
# if fastvideo_args.ltx2_refine_enabled:
|
||||
# result.add_check(
|
||||
# "height", batch.height, [V.is_int, V.positive_int])
|
||||
# result.add_check(
|
||||
# "width", batch.width, [V.is_int, V.positive_int])
|
||||
return result
|
||||
|
||||
|
||||
class LTX2UpsampleStage(PipelineStage):
|
||||
"""Upsample stage-1 latents to stage-2 resolution and add refinement noise."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
upsampler: Any,
|
||||
vae: Any,
|
||||
transformer: Any | None = None,
|
||||
sigmas: list[float] | None = None,
|
||||
add_noise: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.upsampler = upsampler
|
||||
self.vae = vae
|
||||
self.transformer = transformer
|
||||
self.sigmas = sigmas or STAGE_2_DISTILLED_SIGMA_VALUES
|
||||
self.add_noise = add_noise
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
if not fastvideo_args.ltx2_refine_enabled:
|
||||
return batch
|
||||
|
||||
if batch.latents is None:
|
||||
raise ValueError(
|
||||
"Latents must be available before LTX-2 upsample stage.")
|
||||
|
||||
latents = batch.latents
|
||||
orig_dtype = latents.dtype
|
||||
orig_device = latents.device
|
||||
if isinstance(self.upsampler, torch.nn.Module):
|
||||
first_param = next(self.upsampler.parameters(), None)
|
||||
if first_param is not None:
|
||||
if first_param.device != orig_device:
|
||||
latents = latents.to(device=first_param.device)
|
||||
if first_param.dtype != latents.dtype:
|
||||
latents = latents.to(dtype=first_param.dtype)
|
||||
if latents.dtype != orig_dtype or latents.device != orig_device:
|
||||
logger.info(
|
||||
"[LTX2] Cast latents to %s on %s for upsampler.",
|
||||
latents.dtype,
|
||||
latents.device,
|
||||
)
|
||||
|
||||
target_height = batch.extra.get("ltx2_refine_target_height")
|
||||
target_width = batch.extra.get("ltx2_refine_target_width")
|
||||
if target_height is None or target_width is None:
|
||||
raise ValueError("Missing target resolution for LTX-2 refinement.")
|
||||
|
||||
video_encoder = getattr(self.vae, "encoder", None)
|
||||
if video_encoder is None:
|
||||
raise ValueError(
|
||||
"LTX-2 VAE encoder is required for latent upsampling.")
|
||||
|
||||
upsampler_module = getattr(self.upsampler, "model", self.upsampler)
|
||||
latents = upsample_video(latents, video_encoder, upsampler_module)
|
||||
if latents.dtype != orig_dtype or latents.device != orig_device:
|
||||
latents = latents.to(device=orig_device, dtype=orig_dtype)
|
||||
|
||||
sigma0 = float(self.sigmas[0]) if self.sigmas else 1.0
|
||||
if self.add_noise:
|
||||
patchifier = getattr(self.transformer, "patchifier", None)
|
||||
if patchifier is not None:
|
||||
video_shape = VideoLatentShape.from_torch_shape(latents.shape)
|
||||
latents_patch = patchifier.patchify(latents)
|
||||
noise_shape = latents_patch.shape
|
||||
noise_path = fastvideo_args.ltx2_refine_noise_path
|
||||
noise = self._load_noise(
|
||||
noise_path,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype,
|
||||
expected_shape=noise_shape,
|
||||
alternate_shape=latents.shape,
|
||||
) if noise_path else None
|
||||
if noise is None:
|
||||
noise = randn_tensor(
|
||||
noise_shape,
|
||||
generator=batch.generator,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
if noise_path:
|
||||
self._save_noise(noise_path, noise)
|
||||
elif noise.shape == latents.shape:
|
||||
noise = patchifier.patchify(noise)
|
||||
noised_patch = noise * sigma0 + latents_patch * (1.0 - sigma0)
|
||||
latents = patchifier.unpatchify(noised_patch, video_shape)
|
||||
else:
|
||||
noise_path = fastvideo_args.ltx2_refine_noise_path
|
||||
noise = self._load_noise(
|
||||
noise_path,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype,
|
||||
expected_shape=latents.shape,
|
||||
) if noise_path else None
|
||||
if noise is None:
|
||||
noise = randn_tensor(
|
||||
latents.shape,
|
||||
generator=batch.generator,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
if noise_path:
|
||||
self._save_noise(noise_path, noise)
|
||||
# Match LTX-2 GaussianNoiser: noise * sigma + latent * (1 - sigma).
|
||||
latents = noise * sigma0 + latents * (1.0 - sigma0)
|
||||
|
||||
audio_latents = batch.extra.get("ltx2_audio_latents")
|
||||
if audio_latents is not None:
|
||||
audio_latents = audio_latents.to(device=latents.device)
|
||||
if self.add_noise:
|
||||
audio_patchifier = getattr(self.transformer, "audio_patchifier",
|
||||
None)
|
||||
if audio_patchifier is not None:
|
||||
audio_shape = AudioLatentShape.from_torch_shape(
|
||||
audio_latents.shape)
|
||||
audio_patch = audio_patchifier.patchify(audio_latents)
|
||||
audio_noise_shape = audio_patch.shape
|
||||
audio_noise_path = fastvideo_args.ltx2_refine_audio_noise_path
|
||||
audio_noise = self._load_noise(
|
||||
audio_noise_path,
|
||||
device=audio_latents.device,
|
||||
dtype=audio_latents.dtype,
|
||||
expected_shape=audio_noise_shape,
|
||||
alternate_shape=audio_latents.shape,
|
||||
) if audio_noise_path else None
|
||||
if audio_noise is None:
|
||||
audio_noise = randn_tensor(
|
||||
audio_noise_shape,
|
||||
generator=batch.generator,
|
||||
device=audio_latents.device,
|
||||
dtype=audio_latents.dtype,
|
||||
)
|
||||
if audio_noise_path:
|
||||
self._save_noise(audio_noise_path, audio_noise)
|
||||
elif audio_noise.shape == audio_latents.shape:
|
||||
audio_noise = audio_patchifier.patchify(audio_noise)
|
||||
audio_noised_patch = audio_noise * sigma0 + audio_patch * (
|
||||
1.0 - sigma0)
|
||||
audio_latents = audio_patchifier.unpatchify(
|
||||
audio_noised_patch, audio_shape)
|
||||
else:
|
||||
audio_noise_path = fastvideo_args.ltx2_refine_audio_noise_path
|
||||
audio_noise = self._load_noise(
|
||||
audio_noise_path,
|
||||
device=audio_latents.device,
|
||||
dtype=audio_latents.dtype,
|
||||
expected_shape=audio_latents.shape,
|
||||
) if audio_noise_path else None
|
||||
if audio_noise is None:
|
||||
audio_noise = randn_tensor(
|
||||
audio_latents.shape,
|
||||
generator=batch.generator,
|
||||
device=audio_latents.device,
|
||||
dtype=audio_latents.dtype,
|
||||
)
|
||||
if audio_noise_path:
|
||||
self._save_noise(audio_noise_path, audio_noise)
|
||||
# Same noise mixing as video latents for distilled refinement.
|
||||
audio_latents = audio_noise * sigma0 + audio_latents * (
|
||||
1.0 - sigma0)
|
||||
batch.extra["ltx2_audio_latents"] = audio_latents
|
||||
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = latents.shape
|
||||
batch.height = target_height
|
||||
batch.width = target_width
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
if fastvideo_args.ltx2_refine_enabled:
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
def _load_noise(
|
||||
self,
|
||||
noise_path: str | None,
|
||||
*,
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
expected_shape: torch.Size | tuple[int, ...],
|
||||
alternate_shape: torch.Size | tuple[int, ...] | None = None,
|
||||
) -> torch.Tensor | None:
|
||||
if not noise_path:
|
||||
return None
|
||||
path = Path(noise_path)
|
||||
if not path.exists():
|
||||
return None
|
||||
payload = torch.load(path, map_location=device)
|
||||
if isinstance(payload, dict):
|
||||
noise = (payload.get("noise") or payload.get("latent_noise")
|
||||
or payload.get("latent") or payload.get("video_noise"))
|
||||
else:
|
||||
noise = payload
|
||||
if not torch.is_tensor(noise):
|
||||
raise TypeError(f"Expected tensor noise in {path}")
|
||||
noise_shape = tuple(noise.shape)
|
||||
if noise_shape != tuple(expected_shape) and (alternate_shape is None
|
||||
or noise_shape
|
||||
!= tuple(alternate_shape)):
|
||||
raise ValueError(
|
||||
f"Noise shape mismatch for {path}: expected {tuple(expected_shape)}, got {noise_shape}"
|
||||
)
|
||||
logger.info("[LTX2] Loaded refine noise from %s", path)
|
||||
return noise.to(device=device, dtype=dtype)
|
||||
|
||||
def _save_noise(self, noise_path: str, noise: torch.Tensor) -> None:
|
||||
path = Path(noise_path)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if path.exists():
|
||||
return
|
||||
torch.save({"noise": noise.detach().cpu()}, path)
|
||||
logger.info("[LTX2] Saved refine noise to %s", path)
|
||||
|
||||
|
||||
class LTX2RefineLoRAStage(PipelineStage):
|
||||
"""Apply a refinement-specific LoRA before stage-2 denoising."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
pipeline: Any,
|
||||
lora_path: str | None,
|
||||
lora_nickname: str = "ltx2_refine",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self._pipeline_ref = weakref.ref(
|
||||
pipeline) if pipeline is not None else None
|
||||
self._lora_path = lora_path
|
||||
self._lora_nickname = lora_nickname
|
||||
self._applied = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
if not fastvideo_args.ltx2_refine_enabled:
|
||||
return batch
|
||||
lora_path = fastvideo_args.ltx2_refine_lora_path or self._lora_path
|
||||
if not lora_path or self._applied:
|
||||
return batch
|
||||
|
||||
pipeline = self._pipeline_ref(
|
||||
) if self._pipeline_ref is not None else None
|
||||
if pipeline is None or not hasattr(pipeline, "set_lora_adapter"):
|
||||
raise ValueError(
|
||||
"LTX2 refinement LoRA requested but pipeline does not support LoRA adapters."
|
||||
)
|
||||
|
||||
pipeline.set_lora_adapter(self._lora_nickname, lora_path)
|
||||
self._applied = True
|
||||
logger.info("[LTX2] Applied refinement LoRA from %s", lora_path)
|
||||
return batch
|
||||
|
||||
|
||||
__all__ = [
|
||||
"STAGE_2_DISTILLED_SIGMA_VALUES",
|
||||
"LTX2RefineInitStage",
|
||||
"LTX2UpsampleStage",
|
||||
"LTX2RefineLoRAStage",
|
||||
]
|
||||
@@ -0,0 +1,112 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Convert LTX-2 latent upsampler weights to FastVideo format.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import load_file, save_file
|
||||
|
||||
|
||||
def _find_shards(model_path: Path) -> list[Path]:
|
||||
if model_path.is_file():
|
||||
return [model_path]
|
||||
|
||||
index_files = list(model_path.glob("*.safetensors.index.json"))
|
||||
if index_files:
|
||||
with index_files[0].open("r", encoding="utf-8") as f:
|
||||
index = json.load(f)
|
||||
return sorted({model_path / shard for shard in index["weight_map"].values()})
|
||||
return sorted(Path(p) for p in glob.glob(str(model_path / "*.safetensors")))
|
||||
|
||||
|
||||
def _read_metadata_config(path: Path) -> dict:
|
||||
with safe_open(str(path), framework="pt") as f:
|
||||
metadata = f.metadata()
|
||||
if not metadata or "config" not in metadata:
|
||||
raise ValueError(f"Missing config metadata in {path}")
|
||||
return json.loads(metadata["config"])
|
||||
|
||||
|
||||
def _load_weights(shards: list[Path]) -> dict:
|
||||
weights: dict = {}
|
||||
for shard in shards:
|
||||
weights.update(load_file(str(shard)))
|
||||
return weights
|
||||
|
||||
|
||||
def _map_keys(weights: dict, add_model_prefix: bool) -> dict:
|
||||
remapped = {}
|
||||
for key, value in weights.items():
|
||||
new_key = key
|
||||
if not add_model_prefix and new_key.startswith("model."):
|
||||
new_key = new_key[len("model.") :]
|
||||
if add_model_prefix and not new_key.startswith("model."):
|
||||
new_key = f"model.{new_key}"
|
||||
remapped[new_key] = value
|
||||
return remapped
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Convert LTX-2 spatial upsampler weights to FastVideo format"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--source",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to the official LTX-2 upsampler safetensors file or directory",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Output directory for converted weights",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--class-name",
|
||||
type=str,
|
||||
default="LTX2LatentUpsampler",
|
||||
help="_class_name to write into config.json",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--add-model-prefix",
|
||||
action="store_true",
|
||||
help="Prefix all weights with 'model.' for loading into wrapper modules",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
source_path = Path(args.source)
|
||||
shards = _find_shards(source_path)
|
||||
if not shards:
|
||||
raise FileNotFoundError(f"No safetensors found in {source_path}")
|
||||
|
||||
config = _read_metadata_config(shards[0])
|
||||
config["_class_name"] = args.class_name
|
||||
|
||||
weights = _load_weights(shards)
|
||||
weights = _map_keys(weights, add_model_prefix=args.add_model_prefix)
|
||||
|
||||
output_dir = Path(args.output)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
output_weights = output_dir / "model.safetensors"
|
||||
save_file(weights, str(output_weights))
|
||||
print(f"Saved weights to {output_weights}")
|
||||
|
||||
config_path = output_dir / "config.json"
|
||||
with config_path.open("w", encoding="utf-8") as f:
|
||||
json.dump(config, f, indent=2)
|
||||
f.write("\n")
|
||||
print(f"Saved config to {config_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -2,8 +2,13 @@ from huggingface_hub import HfApi
|
||||
|
||||
api = HfApi()
|
||||
|
||||
repo_id = "FastVideo/LTX2-Distilled-LoRA"
|
||||
|
||||
# Create the repo if it doesn't exist
|
||||
api.create_repo(repo_id=repo_id, repo_type="model", exist_ok=True)
|
||||
|
||||
api.upload_folder(
|
||||
folder_path="Wan2.2-TI2V-5B-Diffusers",
|
||||
repo_id="FastVideo/FastWan2.2-TI2V-5B-Diffusers",
|
||||
folder_path="/mnt/user_storage/dev/FastVideo/ltx_distilled_lora",
|
||||
repo_id=repo_id,
|
||||
repo_type="model",
|
||||
)
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
if str(repo_root) not in sys.path:
|
||||
sys.path.insert(0, str(repo_root))
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
|
||||
from fastvideo.pipelines.stages.ltx2_denoising import _ltx2_sigmas
|
||||
|
||||
|
||||
def test_ltx2_sigma_schedule_parity():
|
||||
try:
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 import failed: {exc}")
|
||||
|
||||
scheduler = LTX2Scheduler()
|
||||
device = torch.device("cpu")
|
||||
|
||||
for steps in (4, 8, 40):
|
||||
ref = scheduler.execute(steps=steps, latent=None).to(torch.float32)
|
||||
ours = _ltx2_sigmas(steps=steps, latent=None, device=device)
|
||||
assert_close(ref, ours, atol=1e-6, rtol=1e-6)
|
||||
|
||||
latent = torch.randn(1, 128, 4, 8, 8)
|
||||
ref = scheduler.execute(steps=10, latent=latent).to(torch.float32)
|
||||
ours = _ltx2_sigmas(steps=10, latent=latent, device=device)
|
||||
assert_close(ref, ours, atol=1e-6, rtol=1e-6)
|
||||
@@ -0,0 +1,379 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import tempfile
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
if str(repo_root) not in sys.path:
|
||||
sys.path.insert(0, str(repo_root))
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
def _log_tensor_stats(label: str, tensor: torch.Tensor) -> None:
|
||||
tensor_f32 = tensor.float()
|
||||
print(
|
||||
f"[LTX2 TWO STAGE] {label}: shape={tuple(tensor.shape)} "
|
||||
f"dtype={tensor.dtype} device={tensor.device} "
|
||||
f"min={tensor_f32.min().item():.6f} max={tensor_f32.max().item():.6f} "
|
||||
f"mean={tensor_f32.mean().item():.6f} sum={tensor_f32.sum().item():.6f}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="LTX-2 two-stage parity test requires CUDA.",
|
||||
)
|
||||
def test_ltx2_two_stage_parity():
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
if str(repo_root) not in sys.path:
|
||||
sys.path.insert(0, str(repo_root))
|
||||
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
ltx_pipelines_path = repo_root / "LTX-2" / "packages" / "ltx-pipelines" / "src"
|
||||
if ltx_pipelines_path.exists() and str(ltx_pipelines_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_pipelines_path))
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
os.environ.setdefault("LTX2_REFERENCE_ATTN", "pytorch")
|
||||
torch.backends.cuda.enable_flash_sdp(False)
|
||||
torch.backends.cuda.enable_mem_efficient_sdp(False)
|
||||
torch.backends.cuda.enable_math_sdp(True)
|
||||
|
||||
diffusers_path = os.getenv("LTX2_DIFFUSERS_PATH", "converted/ltx2_diffusers")
|
||||
gemma_model_path = os.path.join(diffusers_path, "text_encoder", "gemma")
|
||||
official_path = os.getenv(
|
||||
"LTX2_OFFICIAL_PATH",
|
||||
"official_ltx_weights/ltx-2-19b-distilled.safetensors",
|
||||
)
|
||||
upsampler_official_path = os.getenv(
|
||||
"LTX2_UPSAMPLER_OFFICIAL_PATH",
|
||||
"official_ltx_weights/ltx-2-spatial-upscaler-x2-1.0.safetensors",
|
||||
)
|
||||
upsampler_path = os.getenv(
|
||||
"LTX2_UPSAMPLER_PATH",
|
||||
"converted/ltx2_spatial_upscaler",
|
||||
)
|
||||
lora_path = os.getenv(
|
||||
"LTX2_REFINER_LORA_PATH",
|
||||
"official_ltx_weights/ltx-2-19b-distilled-lora-384.safetensors",
|
||||
)
|
||||
|
||||
if not os.path.isdir(diffusers_path):
|
||||
pytest.skip(f"Missing LTX-2 diffusers repo at {diffusers_path}")
|
||||
if not os.path.isfile(os.path.join(diffusers_path, "model_index.json")):
|
||||
pytest.skip("Missing model_index.json in diffusers path")
|
||||
if not gemma_model_path or not os.path.isdir(gemma_model_path):
|
||||
pytest.skip("Gemma weights not found in text_encoder/gemma.")
|
||||
if not os.path.isfile(official_path):
|
||||
pytest.skip(f"Missing LTX-2 official weights at {official_path}")
|
||||
if not os.path.isfile(upsampler_official_path):
|
||||
pytest.skip(f"Missing official upsampler at {upsampler_official_path}")
|
||||
if not os.path.isdir(upsampler_path):
|
||||
pytest.skip(f"Missing FastVideo upsampler at {upsampler_path}")
|
||||
if not os.path.isfile(lora_path):
|
||||
pytest.skip(f"Missing distilled LoRA at {lora_path}")
|
||||
|
||||
try:
|
||||
from ltx_pipelines.ti2vid_two_stages import TI2VidTwoStagesPipeline
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.sd_ops import LTXV_LORA_COMFY_RENAMING_MAP
|
||||
from ltx_core.model.transformer import attention as ltx_attention
|
||||
from ltx_core.model.transformer.attention import Attention, AttentionFunction
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 pipeline import failed: {exc}")
|
||||
|
||||
# Force reference attention to use PyTorch SDPA
|
||||
ltx_attention.memory_efficient_attention = None
|
||||
ltx_attention.flash_attn_interface = None
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers."
|
||||
negative_prompt = "low quality, blurry, distorted, artifacts, jpeg compression"
|
||||
seed = 42
|
||||
height = 64
|
||||
width = 64
|
||||
num_frames = 9
|
||||
fps = 12.0
|
||||
steps = 4
|
||||
guidance_scale = 4.0
|
||||
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
def _env_flag(name: str) -> bool:
|
||||
return os.getenv(name, "0").lower() in ("1", "true", "yes")
|
||||
|
||||
debug_kwargs = {}
|
||||
if _env_flag("FASTVIDEO_DEBUG_STAGE_SUMS"):
|
||||
debug_kwargs["debug_stage_sums"] = True
|
||||
debug_kwargs["debug_stage_sums_path"] = os.getenv(
|
||||
"FASTVIDEO_DEBUG_STAGE_SUMS_PATH")
|
||||
if _env_flag("FASTVIDEO_DEBUG_MODULE_SUMS"):
|
||||
debug_kwargs["debug_module_sums"] = True
|
||||
debug_kwargs["debug_module_sums_path"] = os.getenv(
|
||||
"FASTVIDEO_DEBUG_MODULE_SUMS_PATH")
|
||||
include = os.getenv("FASTVIDEO_DEBUG_MODULE_SUMS_INCLUDE")
|
||||
exclude = os.getenv("FASTVIDEO_DEBUG_MODULE_SUMS_EXCLUDE")
|
||||
if include:
|
||||
debug_kwargs["debug_module_sums_include"] = include.split(",")
|
||||
if exclude:
|
||||
debug_kwargs["debug_module_sums_exclude"] = exclude.split(",")
|
||||
if _env_flag("FASTVIDEO_DEBUG_MODEL_SUMS"):
|
||||
debug_kwargs["debug_model_sums"] = True
|
||||
debug_kwargs["debug_model_sums_path"] = os.getenv(
|
||||
"FASTVIDEO_DEBUG_MODEL_SUMS_PATH")
|
||||
if _env_flag("FASTVIDEO_DEBUG_MODEL_DETAIL"):
|
||||
debug_kwargs["debug_model_detail"] = True
|
||||
debug_kwargs["debug_model_detail_path"] = os.getenv(
|
||||
"FASTVIDEO_DEBUG_MODEL_DETAIL_PATH")
|
||||
|
||||
ref_debug = _env_flag("FASTVIDEO_DEBUG_REF_SUMS")
|
||||
ref_debug_path = os.getenv("FASTVIDEO_DEBUG_REF_SUMS_PATH")
|
||||
|
||||
def _ref_log(line: str) -> None:
|
||||
if ref_debug_path:
|
||||
os.makedirs(os.path.dirname(ref_debug_path), exist_ok=True)
|
||||
with open(ref_debug_path, "a", encoding="utf-8") as f:
|
||||
f.write(line + "\n")
|
||||
else:
|
||||
print(line)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
# FastVideo two-stage
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
diffusers_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
pin_cpu_memory=False,
|
||||
ltx2_vae_tiling=False,
|
||||
ltx2_refine_enabled=True,
|
||||
ltx2_refine_upsampler_path=upsampler_path,
|
||||
ltx2_refine_lora_path=lora_path,
|
||||
ltx2_refine_num_inference_steps=3,
|
||||
ltx2_refine_guidance_scale=1.0,
|
||||
ltx2_refine_add_noise=True,
|
||||
**debug_kwargs,
|
||||
)
|
||||
# Disable decoder noise for parity if possible.
|
||||
if hasattr(generator, "executor") and hasattr(generator.executor, "pipeline"):
|
||||
pipeline = generator.executor.pipeline
|
||||
if hasattr(pipeline, "modules") and "vae" in pipeline.modules:
|
||||
decoder = getattr(pipeline.modules["vae"], "decoder", None)
|
||||
if decoder is not None and hasattr(decoder, "decode_noise_scale"):
|
||||
decoder.decode_noise_scale = 0.0
|
||||
|
||||
result = generator.generate_video(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
output_path=os.path.join(tmpdir, "fastvideo_two_stage"),
|
||||
save_video=False,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
fps=fps,
|
||||
num_inference_steps=steps,
|
||||
guidance_scale=guidance_scale,
|
||||
seed=seed,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
fastvideo_out = result["samples"].to(dtype=torch.float32).cpu()
|
||||
_log_tensor_stats("fastvideo_video", fastvideo_out)
|
||||
del result
|
||||
torch.cuda.synchronize()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Reference two-stage
|
||||
if ref_debug:
|
||||
from ltx_pipelines.utils import helpers as ltx_helpers
|
||||
import ltx_pipelines.ti2vid_two_stages as ltx_two_stage
|
||||
from ltx_core.model import upsampler as ltx_upsampler
|
||||
from ltx_core.tools import AudioLatentTools, VideoLatentTools
|
||||
from ltx_core.text_encoders import gemma as ltx_gemma
|
||||
|
||||
original_denoise = ltx_helpers.denoise_audio_video
|
||||
original_upsample = ltx_upsampler.upsample_video
|
||||
original_create_noised_state = ltx_helpers.create_noised_state
|
||||
original_encode_text = ltx_gemma.encode_text
|
||||
|
||||
def _ref_sum(tensor: torch.Tensor | None) -> float:
|
||||
if tensor is None:
|
||||
return 0.0
|
||||
return float(tensor.detach().sum(dtype=torch.float32).item())
|
||||
|
||||
def _wrapped_upsample(*args, **kwargs):
|
||||
out = original_upsample(*args, **kwargs)
|
||||
_ref_log(
|
||||
f"reference:upsample_latent_sum={_ref_sum(out):.6f} "
|
||||
f"shape={tuple(out.shape)}")
|
||||
return out
|
||||
|
||||
def _wrapped_denoise(
|
||||
*,
|
||||
output_shape,
|
||||
conditionings,
|
||||
noiser,
|
||||
sigmas,
|
||||
stepper,
|
||||
denoising_loop_fn,
|
||||
components,
|
||||
dtype,
|
||||
device,
|
||||
noise_scale=1.0,
|
||||
initial_video_latent=None,
|
||||
initial_audio_latent=None,
|
||||
skip_video_noiser=False,
|
||||
skip_audio_noiser=False,
|
||||
):
|
||||
stage = "stage2" if output_shape.width == width else "stage1"
|
||||
_ref_log(
|
||||
f"reference:{stage}:noise_scale={float(noise_scale):.6f} "
|
||||
f"init_video_sum={_ref_sum(initial_video_latent):.6f} "
|
||||
f"init_audio_sum={_ref_sum(initial_audio_latent):.6f}")
|
||||
video_state, audio_state = original_denoise(
|
||||
output_shape=output_shape,
|
||||
conditionings=conditionings,
|
||||
noiser=noiser,
|
||||
sigmas=sigmas,
|
||||
stepper=stepper,
|
||||
denoising_loop_fn=denoising_loop_fn,
|
||||
components=components,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
noise_scale=noise_scale,
|
||||
initial_video_latent=initial_video_latent,
|
||||
initial_audio_latent=initial_audio_latent,
|
||||
skip_video_noiser=skip_video_noiser,
|
||||
skip_audio_noiser=skip_audio_noiser,
|
||||
)
|
||||
_ref_log(
|
||||
f"reference:{stage}:video_latent_sum={_ref_sum(video_state.latent):.6f} "
|
||||
f"audio_latent_sum={_ref_sum(audio_state.latent):.6f}")
|
||||
return video_state, audio_state
|
||||
|
||||
def _wrapped_create_noised_state(
|
||||
tools,
|
||||
conditionings,
|
||||
noiser,
|
||||
dtype,
|
||||
device,
|
||||
noise_scale=1.0,
|
||||
initial_latent=None,
|
||||
skip_noiser=False,
|
||||
):
|
||||
state = original_create_noised_state(
|
||||
tools=tools,
|
||||
conditionings=conditionings,
|
||||
noiser=noiser,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
noise_scale=noise_scale,
|
||||
initial_latent=initial_latent,
|
||||
skip_noiser=skip_noiser,
|
||||
)
|
||||
stage = "stage2" if initial_latent is not None else "stage1"
|
||||
kind = "video" if isinstance(tools, VideoLatentTools) else "audio"
|
||||
_ref_log(
|
||||
f"reference:{stage}:{kind}:noised_latent_sum={_ref_sum(state.latent):.6f} "
|
||||
f"noise_scale={float(noise_scale):.6f}")
|
||||
return state
|
||||
|
||||
def _wrapped_encode_text(text_encoder, prompts):
|
||||
context_p, context_n = original_encode_text(text_encoder, prompts=prompts)
|
||||
v_context_p, a_context_p = context_p
|
||||
v_context_n, a_context_n = context_n
|
||||
_ref_log(
|
||||
"reference:text:"
|
||||
f"v_pos_sum={_ref_sum(v_context_p):.6f} "
|
||||
f"a_pos_sum={_ref_sum(a_context_p):.6f} "
|
||||
f"v_neg_sum={_ref_sum(v_context_n):.6f} "
|
||||
f"a_neg_sum={_ref_sum(a_context_n):.6f}")
|
||||
return context_p, context_n
|
||||
|
||||
ltx_helpers.denoise_audio_video = _wrapped_denoise
|
||||
ltx_upsampler.upsample_video = _wrapped_upsample
|
||||
ltx_helpers.create_noised_state = _wrapped_create_noised_state
|
||||
ltx_two_stage.denoise_audio_video = _wrapped_denoise
|
||||
ltx_two_stage.upsample_video = _wrapped_upsample
|
||||
ltx_gemma.encode_text = _wrapped_encode_text
|
||||
ltx_two_stage.encode_text = _wrapped_encode_text
|
||||
|
||||
ref_pipeline = TI2VidTwoStagesPipeline(
|
||||
checkpoint_path=official_path,
|
||||
distilled_lora=[
|
||||
LoraPathStrengthAndSDOps(
|
||||
lora_path,
|
||||
1.0,
|
||||
LTXV_LORA_COMFY_RENAMING_MAP,
|
||||
)
|
||||
],
|
||||
spatial_upsampler_path=upsampler_official_path,
|
||||
gemma_root=gemma_model_path,
|
||||
loras=[],
|
||||
device=device,
|
||||
fp8transformer=False,
|
||||
)
|
||||
|
||||
original_text_encoder = ref_pipeline.stage_1_model_ledger.text_encoder
|
||||
original_transformer = ref_pipeline.stage_1_model_ledger.transformer
|
||||
original_video_decoder = ref_pipeline.stage_2_model_ledger.video_decoder
|
||||
|
||||
def _patched_text_encoder():
|
||||
encoder = original_text_encoder()
|
||||
if hasattr(encoder, "model") and hasattr(encoder.model, "config"):
|
||||
if hasattr(encoder.model.config, "attn_implementation"):
|
||||
encoder.model.config.attn_implementation = "sdpa"
|
||||
if hasattr(encoder.model.config, "_attn_implementation"):
|
||||
encoder.model.config._attn_implementation = "sdpa"
|
||||
for module in encoder.modules():
|
||||
if isinstance(module, Attention):
|
||||
module.attention_function = AttentionFunction.PYTORCH
|
||||
return encoder
|
||||
|
||||
def _patched_video_decoder():
|
||||
decoder = original_video_decoder()
|
||||
if hasattr(decoder, "decode_noise_scale"):
|
||||
decoder.decode_noise_scale = 0.0
|
||||
return decoder
|
||||
|
||||
ref_pipeline.stage_1_model_ledger.text_encoder = _patched_text_encoder
|
||||
ref_pipeline.stage_2_model_ledger.video_decoder = _patched_video_decoder
|
||||
ref_pipeline.stage_1_model_ledger.transformer = original_transformer
|
||||
|
||||
with torch.no_grad():
|
||||
ref_video_iter, _ = ref_pipeline(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
seed=seed,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
frame_rate=fps,
|
||||
num_inference_steps=steps,
|
||||
cfg_guidance_scale=guidance_scale,
|
||||
images=[],
|
||||
enhance_prompt=False,
|
||||
)
|
||||
ref_chunks = list(ref_video_iter)
|
||||
ref_video = torch.cat(
|
||||
[chunk if torch.is_tensor(chunk) else torch.from_numpy(chunk) for chunk in ref_chunks],
|
||||
dim=0,
|
||||
)
|
||||
ref_video = ref_video.to(torch.float32) / 255.0
|
||||
ref_video = ref_video.permute(3, 0, 1, 2).unsqueeze(0)
|
||||
ref_video = ref_video.cpu()
|
||||
_log_tensor_stats("reference_video", ref_video)
|
||||
|
||||
assert ref_video.shape == fastvideo_out.shape
|
||||
assert_close(ref_video, fastvideo_out, atol=2 / 255, rtol=1e-3)
|
||||
@@ -9,8 +9,11 @@ from torch.testing import assert_close
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29513")
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
if str(repo_root) not in sys.path:
|
||||
sys.path.insert(0, str(repo_root))
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
@@ -32,7 +35,7 @@ def _read_transformer_config(config: dict) -> dict:
|
||||
def _infer_patch_params(in_channels: int) -> tuple[int, int]:
|
||||
patch_size = 1
|
||||
num_channels_latents = 128
|
||||
for candidate in (8, 16, 32, 64, 128):
|
||||
for candidate in (128, 64, 32, 16, 8):
|
||||
if in_channels % candidate != 0:
|
||||
continue
|
||||
patch_volume = in_channels // candidate
|
||||
@@ -161,17 +164,27 @@ def test_ltx2_transformer_parity():
|
||||
pytest.skip(f"FastVideo converted weights not found at {fastvideo_path}")
|
||||
|
||||
try:
|
||||
from ltx_core.components.patchifiers import VideoLatentPatchifier
|
||||
from ltx_core.components.patchifiers import (VideoLatentPatchifier,
|
||||
get_pixel_coords)
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
|
||||
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader
|
||||
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
|
||||
from ltx_core.model.transformer import (LTXModelConfigurator,
|
||||
LTXV_MODEL_COMFY_RENAMING_MAP)
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
from ltx_core.types import VideoLatentShape
|
||||
from ltx_core.types import VideoLatentShape, VIDEO_SCALE_FACTORS
|
||||
from ltx_core.utils import to_denoised
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 import failed: {exc}")
|
||||
|
||||
# Force reference attention to use PyTorch SDPA for parity with FastVideo.
|
||||
try:
|
||||
from ltx_core.model.transformer import attention as ltx_attention
|
||||
ltx_attention.memory_efficient_attention = None
|
||||
ltx_attention.flash_attn_interface = None
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
config_loader = SafetensorsModelStateDictLoader()
|
||||
metadata = config_loader.metadata(str(official_path))
|
||||
transformer_config = _read_transformer_config(metadata)
|
||||
@@ -304,11 +317,21 @@ def test_ltx2_transformer_parity():
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
timestep = torch.tensor([500], device=device, dtype=precision)
|
||||
|
||||
video_shape = VideoLatentShape.from_torch_shape(hidden_states.shape)
|
||||
positions = patchifier.get_patch_grid_bounds(video_shape, device=hidden_states.device)
|
||||
latent_coords = patchifier.get_patch_grid_bounds(video_shape, device=hidden_states.device)
|
||||
positions = get_pixel_coords(
|
||||
latent_coords=latent_coords,
|
||||
scale_factors=VIDEO_SCALE_FACTORS,
|
||||
causal_fix=True,
|
||||
)
|
||||
latents = patchifier.patchify(hidden_states)
|
||||
timestep = torch.full(
|
||||
(batch_size, latents.shape[1], 1),
|
||||
500,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
|
||||
video = Modality(
|
||||
enabled=True,
|
||||
@@ -325,6 +348,7 @@ def test_ltx2_transformer_parity():
|
||||
audio=None,
|
||||
perturbations=BatchedPerturbationConfig.empty(batch_size),
|
||||
)
|
||||
ref_out = to_denoised(latents, ref_out, timestep)
|
||||
ref_out = patchifier.unpatchify(ref_out, output_shape=video_shape)
|
||||
print(f"[LTX2 TEST] Reference model output shape: {ref_out.shape}")
|
||||
with set_forward_context(
|
||||
|
||||
@@ -14,6 +14,8 @@ os.environ.setdefault("MASTER_PORT", "29513")
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
if str(repo_root) not in sys.path:
|
||||
sys.path.insert(0, str(repo_root))
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
@@ -55,17 +57,33 @@ def test_ltx2_transformer_audio_parity():
|
||||
pytest.skip(f"FastVideo converted weights not found at {fastvideo_path}")
|
||||
|
||||
try:
|
||||
from ltx_core.components.patchifiers import AudioPatchifier, VideoLatentPatchifier
|
||||
from ltx_core.components.patchifiers import (
|
||||
AudioPatchifier,
|
||||
VideoLatentPatchifier,
|
||||
get_pixel_coords,
|
||||
)
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
|
||||
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader
|
||||
from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder
|
||||
from ltx_core.model.transformer import (LTXModelConfigurator,
|
||||
LTXV_MODEL_COMFY_RENAMING_MAP)
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
from ltx_core.types import AudioLatentShape, VideoLatentShape
|
||||
from ltx_core.types import (
|
||||
AudioLatentShape,
|
||||
VideoLatentShape,
|
||||
VIDEO_SCALE_FACTORS,
|
||||
)
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 import failed: {exc}")
|
||||
|
||||
# Force reference attention to use PyTorch SDPA for parity with FastVideo.
|
||||
try:
|
||||
from ltx_core.model.transformer import attention as ltx_attention
|
||||
ltx_attention.memory_efficient_attention = None
|
||||
ltx_attention.flash_attn_interface = None
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Load config from metadata using same approach as test_ltx2.py
|
||||
config_loader = SafetensorsModelStateDictLoader()
|
||||
metadata = config_loader.metadata(str(official_path))
|
||||
@@ -201,11 +219,22 @@ def test_ltx2_transformer_audio_parity():
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
timestep = torch.tensor([500], device=device, dtype=precision)
|
||||
timestep = None
|
||||
|
||||
video_shape = VideoLatentShape.from_torch_shape(hidden_states.shape)
|
||||
positions = patchifier.get_patch_grid_bounds(video_shape, device=hidden_states.device)
|
||||
latent_coords = patchifier.get_patch_grid_bounds(video_shape, device=hidden_states.device)
|
||||
positions = get_pixel_coords(
|
||||
latent_coords=latent_coords,
|
||||
scale_factors=VIDEO_SCALE_FACTORS,
|
||||
causal_fix=True,
|
||||
)
|
||||
latents = patchifier.patchify(hidden_states)
|
||||
timestep = torch.full(
|
||||
(batch_size, latents.shape[1], 1),
|
||||
500,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
|
||||
audio_frames = 16
|
||||
audio_channels = 8
|
||||
@@ -221,6 +250,12 @@ def test_ltx2_transformer_audio_parity():
|
||||
audio_shape = AudioLatentShape.from_torch_shape(audio_latents.shape)
|
||||
audio_positions = audio_patchifier.get_patch_grid_bounds(audio_shape, device=audio_latents.device)
|
||||
audio_tokens = audio_patchifier.patchify(audio_latents)
|
||||
audio_timestep = torch.full(
|
||||
(batch_size, audio_tokens.shape[1], 1),
|
||||
500,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
|
||||
video = Modality(
|
||||
enabled=True,
|
||||
@@ -233,7 +268,7 @@ def test_ltx2_transformer_audio_parity():
|
||||
audio = Modality(
|
||||
enabled=True,
|
||||
latent=audio_tokens,
|
||||
timesteps=timestep,
|
||||
timesteps=audio_timestep,
|
||||
positions=audio_positions,
|
||||
context=encoder_hidden_states,
|
||||
context_mask=None,
|
||||
@@ -250,7 +285,7 @@ def test_ltx2_transformer_audio_parity():
|
||||
fastvideo_audio = FastVideoModality(
|
||||
enabled=True,
|
||||
latent=audio_tokens,
|
||||
timesteps=timestep,
|
||||
timesteps=audio_timestep,
|
||||
positions=audio_positions,
|
||||
context=encoder_hidden_states,
|
||||
context_mask=None,
|
||||
|
||||
@@ -0,0 +1,228 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29513")
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
if str(repo_root) not in sys.path:
|
||||
sys.path.insert(0, str(repo_root))
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
|
||||
from fastvideo.configs.models.dits import LTX2VideoConfig
|
||||
from fastvideo.configs.pipelines import LTX2T2VConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.basic.ltx2.ltx2_pipeline import LTX2Pipeline
|
||||
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
|
||||
from .test_ltx2 import _infer_patch_params, _read_transformer_config
|
||||
|
||||
|
||||
def test_ltx2_transformer_lora_parity():
|
||||
torch.manual_seed(42)
|
||||
diffusers_root = Path(
|
||||
os.getenv("LTX2_DIFFUSERS_PATH", "converted/ltx2_diffusers")
|
||||
)
|
||||
official_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_OFFICIAL_PATH",
|
||||
"official_ltx_weights/ltx-2-19b-distilled.safetensors",
|
||||
)
|
||||
)
|
||||
lora_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_LORA_PATH",
|
||||
"official_ltx_weights/ltx-2-19b-distilled-lora-384.safetensors",
|
||||
)
|
||||
)
|
||||
if not official_path.exists():
|
||||
pytest.skip(f"LTX-2 official weights not found at {official_path}")
|
||||
if not lora_path.exists():
|
||||
pytest.skip(f"LTX-2 distilled LoRA not found at {lora_path}")
|
||||
if not diffusers_root.exists():
|
||||
pytest.skip(f"LTX-2 diffusers weights not found at {diffusers_root}")
|
||||
|
||||
try:
|
||||
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 import failed: {exc}")
|
||||
|
||||
config_loader = SafetensorsModelStateDictLoader()
|
||||
metadata = config_loader.metadata(str(official_path))
|
||||
transformer_config = _read_transformer_config(metadata)
|
||||
|
||||
config = LTX2VideoConfig()
|
||||
cfg = config.arch_config
|
||||
cfg.num_attention_heads = transformer_config.get("num_attention_heads",
|
||||
cfg.num_attention_heads)
|
||||
cfg.attention_head_dim = transformer_config.get("attention_head_dim",
|
||||
cfg.attention_head_dim)
|
||||
cfg.num_layers = transformer_config.get("num_layers", cfg.num_layers)
|
||||
cfg.cross_attention_dim = transformer_config.get(
|
||||
"cross_attention_dim", cfg.cross_attention_dim)
|
||||
cfg.caption_channels = transformer_config.get("caption_channels",
|
||||
cfg.caption_channels)
|
||||
cfg.norm_eps = transformer_config.get("norm_eps", cfg.norm_eps)
|
||||
cfg.attention_type = transformer_config.get("attention_type",
|
||||
cfg.attention_type)
|
||||
cfg.positional_embedding_theta = transformer_config.get(
|
||||
"positional_embedding_theta", cfg.positional_embedding_theta)
|
||||
cfg.positional_embedding_max_pos = transformer_config.get(
|
||||
"positional_embedding_max_pos", cfg.positional_embedding_max_pos)
|
||||
cfg.timestep_scale_multiplier = transformer_config.get(
|
||||
"timestep_scale_multiplier", cfg.timestep_scale_multiplier)
|
||||
cfg.use_middle_indices_grid = transformer_config.get(
|
||||
"use_middle_indices_grid", cfg.use_middle_indices_grid)
|
||||
cfg.rope_type = transformer_config.get("rope_type", cfg.rope_type)
|
||||
cfg.double_precision_rope = transformer_config.get(
|
||||
"double_precision_rope",
|
||||
transformer_config.get("frequencies_precision", "")
|
||||
== "float64",
|
||||
)
|
||||
cfg.audio_num_attention_heads = transformer_config.get(
|
||||
"audio_num_attention_heads", cfg.audio_num_attention_heads)
|
||||
cfg.audio_attention_head_dim = transformer_config.get(
|
||||
"audio_attention_head_dim", cfg.audio_attention_head_dim)
|
||||
cfg.audio_in_channels = transformer_config.get("audio_in_channels",
|
||||
cfg.audio_in_channels)
|
||||
cfg.audio_out_channels = transformer_config.get("audio_out_channels",
|
||||
cfg.audio_out_channels)
|
||||
cfg.audio_cross_attention_dim = transformer_config.get(
|
||||
"audio_cross_attention_dim", cfg.audio_cross_attention_dim)
|
||||
cfg.audio_positional_embedding_max_pos = transformer_config.get(
|
||||
"audio_positional_embedding_max_pos",
|
||||
cfg.audio_positional_embedding_max_pos,
|
||||
)
|
||||
cfg.av_ca_timestep_scale_multiplier = transformer_config.get(
|
||||
"av_ca_timestep_scale_multiplier", cfg.av_ca_timestep_scale_multiplier)
|
||||
cfg.in_channels = transformer_config.get("in_channels", cfg.in_channels)
|
||||
cfg.out_channels = transformer_config.get("out_channels", cfg.out_channels)
|
||||
|
||||
patch_size, num_channels_latents = _infer_patch_params(cfg.in_channels)
|
||||
cfg.patch_size = (1, patch_size, patch_size)
|
||||
cfg.num_channels_latents = num_channels_latents
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("LTX-2 LoRA parity test requires CUDA for attention backends.")
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
|
||||
pipeline_config = LTX2T2VConfig()
|
||||
pipeline_config.dit_config = config
|
||||
pipeline_config.dit_precision = precision_str
|
||||
|
||||
args = FastVideoArgs(
|
||||
model_path=str(diffusers_root),
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
dit_layerwise_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
lora_path=str(lora_path),
|
||||
lora_nickname="refine",
|
||||
pipeline_config=pipeline_config,
|
||||
)
|
||||
args.device = device
|
||||
|
||||
fastvideo_pipeline = LTX2Pipeline(
|
||||
model_path=str(diffusers_root),
|
||||
fastvideo_args=args,
|
||||
required_config_modules=["transformer"],
|
||||
)
|
||||
fastvideo_model = fastvideo_pipeline.modules["transformer"].to(
|
||||
device=device, dtype=precision)
|
||||
fastvideo_model.eval()
|
||||
|
||||
lora_layers = fastvideo_pipeline.lora_layers.get("transformer")
|
||||
assert lora_layers is not None, "LoRA layers were not created."
|
||||
applied_layers = [
|
||||
layer for _, layer in lora_layers.all_lora_layers()
|
||||
if not layer.disable_lora
|
||||
]
|
||||
assert applied_layers, "LoRA adapter did not match any layers."
|
||||
|
||||
fastvideo_state = fastvideo_model.state_dict()
|
||||
|
||||
def _canonical(name: str) -> str:
|
||||
while name.startswith("model."):
|
||||
name = name[len("model."):]
|
||||
return name
|
||||
|
||||
fastvideo_weights: dict[str, torch.Tensor] = {}
|
||||
for key, value in fastvideo_state.items():
|
||||
if key.endswith(".base_layer.weight"):
|
||||
trimmed = key[:-len(".base_layer.weight")]
|
||||
fastvideo_weights[_canonical(trimmed)] = value
|
||||
continue
|
||||
if key.endswith(".weight") and ".base_layer." not in key:
|
||||
fastvideo_weights[_canonical(key[:-len(".weight")])] = value
|
||||
|
||||
def _iter_lora_layers(path: Path) -> list[str]:
|
||||
from safetensors import safe_open
|
||||
|
||||
with safe_open(str(path), framework="pt") as f:
|
||||
layers = set()
|
||||
for key in f.keys():
|
||||
if key.endswith(".lora_A.weight"):
|
||||
layer = key.replace("diffusion_model.", "").replace(
|
||||
".lora_A.weight", "")
|
||||
layers.add(layer)
|
||||
return sorted(layers)
|
||||
|
||||
def _get_tensor(path: Path, key: str) -> torch.Tensor | None:
|
||||
from safetensors import safe_open
|
||||
|
||||
with safe_open(str(path), framework="pt") as f:
|
||||
if key not in f.keys():
|
||||
return None
|
||||
return f.get_tensor(key)
|
||||
|
||||
lora_layers = _iter_lora_layers(lora_path)
|
||||
assert lora_layers, "No LoRA layers found in the adapter."
|
||||
|
||||
checked = 0
|
||||
for layer in lora_layers:
|
||||
base_key = f"model.diffusion_model.{layer}.weight"
|
||||
lora_a_key = f"diffusion_model.{layer}.lora_A.weight"
|
||||
lora_b_key = f"diffusion_model.{layer}.lora_B.weight"
|
||||
fast_weight = fastvideo_weights.get(layer)
|
||||
if fast_weight is None:
|
||||
continue
|
||||
base_weight = _get_tensor(official_path, base_key)
|
||||
lora_a = _get_tensor(lora_path, lora_a_key)
|
||||
lora_b = _get_tensor(lora_path, lora_b_key)
|
||||
if base_weight is None or lora_a is None or lora_b is None:
|
||||
continue
|
||||
expected = (base_weight.to(torch.bfloat16) +
|
||||
torch.matmul(lora_b.to(torch.bfloat16),
|
||||
lora_a.to(torch.bfloat16)))
|
||||
actual = fast_weight.detach().cpu().to(torch.bfloat16)
|
||||
assert_close(actual.float(), expected.float(), atol=5e-4, rtol=5e-4)
|
||||
checked += 1
|
||||
if checked >= 10:
|
||||
break
|
||||
|
||||
assert checked > 0, "No matching LoRA layers found for parity checks."
|
||||
|
||||
del fastvideo_model
|
||||
del fastvideo_pipeline
|
||||
LoRAPipeline.lora_layers.clear()
|
||||
LoRAPipeline.lora_adapters.clear()
|
||||
LoRAPipeline.cur_adapter_name = ""
|
||||
LoRAPipeline.cur_adapter_path = ""
|
||||
LoRAPipeline.lora_initialized = False
|
||||
import gc
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
@@ -0,0 +1 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -0,0 +1,104 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import load_file
|
||||
from torch.testing import assert_close
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
if str(repo_root) not in sys.path:
|
||||
sys.path.insert(0, str(repo_root))
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
|
||||
from fastvideo.configs.pipelines import PipelineConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.models.loader.component_loader import UpsamplerLoader
|
||||
|
||||
|
||||
def _load_metadata(path: Path) -> dict:
|
||||
with safe_open(str(path), framework="pt") as f:
|
||||
meta = f.metadata()
|
||||
if not meta or "config" not in meta:
|
||||
raise KeyError("Missing config metadata in safetensors file.")
|
||||
return json.loads(meta["config"])
|
||||
|
||||
|
||||
def test_ltx2_upsampler_parity():
|
||||
official_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_UPSAMPLER_OFFICIAL_PATH",
|
||||
"official_ltx_weights/ltx-2-spatial-upscaler-x2-1.0.safetensors",
|
||||
)
|
||||
)
|
||||
fastvideo_path = Path(
|
||||
os.getenv(
|
||||
"LTX2_UPSAMPLER_PATH",
|
||||
"converted/ltx2_spatial_upscaler",
|
||||
)
|
||||
)
|
||||
|
||||
if not official_path.exists():
|
||||
pytest.skip(f"LTX-2 upsampler weights not found at {official_path}")
|
||||
if not fastvideo_path.exists():
|
||||
pytest.skip(f"FastVideo upsampler weights not found at {fastvideo_path}")
|
||||
|
||||
try:
|
||||
from ltx_core.model.upsampler import LatentUpsamplerConfigurator
|
||||
except ImportError as exc:
|
||||
pytest.skip(f"LTX-2 import failed: {exc}")
|
||||
|
||||
config = _load_metadata(official_path)
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16 if torch.cuda.is_available() else torch.float32
|
||||
precision_str = "bf16" if torch.cuda.is_available() else "fp32"
|
||||
|
||||
ref_model = LatentUpsamplerConfigurator.from_config(config).to(
|
||||
device=device, dtype=precision
|
||||
)
|
||||
ref_weights = load_file(str(official_path))
|
||||
ref_model.load_state_dict(ref_weights, strict=False)
|
||||
|
||||
args = FastVideoArgs(
|
||||
model_path=str(fastvideo_path),
|
||||
pipeline_config=PipelineConfig(vae_precision=precision_str),
|
||||
)
|
||||
loader = UpsamplerLoader()
|
||||
fastvideo_model = loader.load(str(fastvideo_path), args).to(
|
||||
device=device, dtype=precision
|
||||
)
|
||||
|
||||
ref_model.eval()
|
||||
fastvideo_model.eval()
|
||||
|
||||
batch = 1
|
||||
channels = config.get("in_channels", 128)
|
||||
frames = 3
|
||||
height = 8
|
||||
width = 8
|
||||
latent = torch.randn(
|
||||
batch,
|
||||
channels,
|
||||
frames,
|
||||
height,
|
||||
width,
|
||||
device=device,
|
||||
dtype=precision,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
ref_out = ref_model(latent)
|
||||
fast_out = fastvideo_model(latent)
|
||||
|
||||
assert ref_out.shape == fast_out.shape
|
||||
assert ref_out.dtype == fast_out.dtype
|
||||
assert torch.isfinite(ref_out).all(), "Reference upsampler produced non-finite output."
|
||||
assert torch.isfinite(fast_out).all(), "FastVideo upsampler produced non-finite output."
|
||||
assert_close(ref_out, fast_out, atol=1e-4, rtol=1e-4)
|
||||
@@ -8,6 +8,8 @@ import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
if str(repo_root) not in sys.path:
|
||||
sys.path.insert(0, str(repo_root))
|
||||
ltx_core_path = repo_root / "LTX-2" / "packages" / "ltx-core" / "src"
|
||||
if ltx_core_path.exists() and str(ltx_core_path) not in sys.path:
|
||||
sys.path.insert(0, str(ltx_core_path))
|
||||
|
||||
Reference in New Issue
Block a user