Compare commits

..
Author SHA1 Message Date
Will Lin 3bd2a885e4 update 2026-02-06 15:41:09 -08:00
Will Lin 6566a05dac update 2026-02-06 15:07:40 -08:00
Will Lin 169d4849b3 update 2026-02-06 15:07:07 -08:00
Will Lin abeea25f77 update 2026-02-06 14:40:45 -08:00
Will Lin db9ba98fbf fix 2026-02-06 14:40:45 -08:00
Will Lin 8bb9a31292 revert 2026-02-06 14:40:45 -08:00
Will Lin 0c33204bbd update 2026-02-06 14:40:44 -08:00
Will Lin b076cd934e test 2026-02-06 14:40:44 -08:00
Will Lin 81872ee886 update 2026-02-06 14:40:44 -08:00
51 changed files with 3191 additions and 509 deletions
+1 -1
View File
@@ -1,5 +1,5 @@
| **[Documentation](https://hao-ai-lab.github.io/FastVideo)** | **[Quick Start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/)** | **[Weekly Dev Meeting](https://github.com/hao-ai-lab/FastVideo/discussions/982)** | 🟣💬 **[Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ)** |
| **[Documentation](https://hao-ai-lab.github.io/FastVideo)** | **[Quick Start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/)** | **[Weekly Dev Meeting](https://github.com/hao-ai-lab/FastVideo/discussions/982)** | 🟣💬 **[Slack**](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ) |
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
@@ -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()
@@ -4,7 +4,7 @@ export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
@@ -14,6 +14,7 @@ NUM_GPUS=1
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_crush_smol"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "wan_ode_init_crush_smol"
--max_train_steps 6000
--train_batch_size 1
@@ -33,7 +34,7 @@ parallel_args=(
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
--hsdp_shard_dim 1
)
# Model arguments
@@ -50,17 +51,20 @@ dataset_args=(
# Validation arguments
validation_args=(
--log-visualization
--visualization-steps 100
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--learning_rate 6e-6
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 0.01
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -1,6 +1,4 @@
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig, WanVideoConfig
@@ -78,14 +76,8 @@ class MatrixGameWanVideoArchConfig(WanVideoArchConfig):
image_dim: int = 1280
def _is_transformer_block(param_name: str, module: torch.nn.Module) -> bool:
return bool("blocks" in param_name and param_name.split(".")[-1].isdigit())
@dataclass
class MatrixGameWanVideoConfig(WanVideoConfig):
arch_config: MatrixGameWanVideoArchConfig = field(
default_factory=MatrixGameWanVideoArchConfig)
prefix: str = "Wan"
_compile_conditions: list = field(
default_factory=lambda: [_is_transformer_block])
+15 -2
View File
@@ -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."
)
+2
View File
@@ -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
+134
View File
@@ -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()]
+202
View File
@@ -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"]
+2 -41
View File
@@ -9,7 +9,6 @@ diffusion models.
import math
import os
import re
import threading
import time
from copy import deepcopy
from typing import Any
@@ -32,21 +31,6 @@ from fastvideo.worker.executor import Executor
logger = init_logger(__name__)
def _infer_latent_batch_size(batch: ForwardBatch) -> int:
if isinstance(batch.prompt, list):
latent_batch_size = len(batch.prompt)
elif batch.prompt is not None:
latent_batch_size = 1
elif batch.prompt_embeds is not None and len(batch.prompt_embeds) > 0:
latent_batch_size = batch.prompt_embeds[0].shape[0]
else:
raise ValueError(
"Cannot infer batch size from batch; no prompt or prompt_embeds found"
)
latent_batch_size *= batch.num_videos_per_prompt
return latent_batch_size
class VideoGenerator:
"""
A unified class for generating videos using diffusion models.
@@ -388,31 +372,8 @@ class VideoGenerator:
# Run inference
start_time = time.perf_counter()
# Execute forward pass in a new thread for non-blocking tensor allocation
result_container = {}
def execute_forward_thread():
result_container['output_batch'] = self.executor.execute_forward(
batch, fastvideo_args)
thread = threading.Thread(target=execute_forward_thread)
thread.start()
latent_batch_size = _infer_latent_batch_size(batch)
samples = torch.empty((latent_batch_size, 3, sampling_param.num_frames,
sampling_param.height, sampling_param.width),
device='cpu',
pin_memory=fastvideo_args.pin_cpu_memory)
thread.join()
output_batch = result_container['output_batch']
if output_batch.output.shape == samples.shape:
samples.copy_(output_batch.output)
else:
logger.warning(
"Output shape %s does not match expected shape %s; use slow path",
output_batch.output.shape, samples.shape)
samples = output_batch.output.cpu()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch.output
logging_info = output_batch.logging_info
gen_time = time.perf_counter() - start_time
+185 -4
View 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(
@@ -916,7 +1101,6 @@ class TrainingArgs(FastVideoArgs):
training_state_checkpointing_steps: int = 0 # for resuming training
weight_only_checkpointing_steps: int = 0 # for inference
log_visualization: bool = False
visualization_steps: int = 0
# simulate generator forward to match inference
simulate_generator_forward: bool = False
warp_denoising_step: bool = False
@@ -1080,9 +1264,6 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--log-validation",
action=StoreBoolean,
help="Whether to log validation results")
parser.add_argument("--visualization-steps",
type=int,
help="Number of visualization steps")
parser.add_argument("--tracker-project-name",
type=str,
help="Project name for tracking")
+31 -6
View File
@@ -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
+2 -3
View File
@@ -33,7 +33,7 @@ from fastvideo.layers.visual_embedding import (PatchEmbed)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.platforms import AttentionBackendEnum
logger = init_logger(__name__)
class CausalWanSelfAttention(nn.Module):
@@ -286,8 +286,6 @@ class CausalWanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -454,6 +452,7 @@ class CausalWanTransformer3DModel(BaseDiT):
This function will be run for num_frame times.
Process the latent frames one by one (1560 tokens each)
"""
from fastvideo.platforms import current_platform
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
+44 -11
View File
@@ -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)
+5
View File
@@ -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,
}
+23
View File
@@ -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",
]
+3
View File
@@ -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 "
+1 -1
View File
@@ -277,7 +277,7 @@ def load_video(
if convert_method is not None:
pil_images = convert_method(pil_images)
return (pil_images, original_fps) if return_fps else pil_images
return pil_images, original_fps if return_fps else pil_images
def get_default_height_width(
+177 -7
View File
@@ -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
+120 -62
View File
@@ -101,58 +101,6 @@ class ComposedPipelineBase(ABC):
module.requires_grad_(True)
module.train()
@staticmethod
def _compile_with_conditions(
module: torch.nn.Module,
compile_kwargs: dict[str, Any],
) -> int:
"""Compile submodules that match module._compile_conditions."""
compile_conditions = getattr(module, "_compile_conditions", None)
if not compile_conditions:
return 0
compiled_count = 0
for name, submodule in module.named_modules():
if not name:
continue
if any(cond(name, submodule) for cond in compile_conditions):
submodule.forward = torch.compile(submodule.forward,
**compile_kwargs)
compiled_count += 1
return compiled_count
def _maybe_compile_pipeline_module(
self,
module_name: str,
fsdp_module_cls: type | None,
compile_kwargs: dict[str, Any],
) -> None:
if module_name not in self.modules:
return
module = self.modules[module_name]
if fsdp_module_cls is not None and isinstance(module, fsdp_module_cls):
logger.info(
"%s is already FSDP-wrapped; skipping torch.compile in pipeline",
module_name.capitalize(),
)
return
compiled_count = self._compile_with_conditions(module, compile_kwargs)
if compiled_count > 0:
logger.info(
"Enabled torch.compile for %d submodules in %s via _compile_conditions with kwargs=%s",
compiled_count,
module_name,
compile_kwargs,
)
return
# Backward-compatible fallback: compile full module if no condition matched.
logger.info("Enabling torch.compile for %s with kwargs=%s", module_name,
compile_kwargs)
self.modules[module_name] = torch.compile(module, **compile_kwargs)
def post_init(self) -> None:
assert self.fastvideo_args is not None, "fastvideo_args must be set"
if self.post_init_called:
@@ -168,6 +116,7 @@ class ComposedPipelineBase(ABC):
self.initialize_pipeline(self.fastvideo_args)
if self.fastvideo_args.enable_torch_compile:
transformer_module = self.modules["transformer"]
if self.fastvideo_args.training_mode:
logger.info(
"Torch Compile enabled via FSDP loader for training; skipping additional pipeline compile"
@@ -181,18 +130,33 @@ class ComposedPipelineBase(ABC):
fsdp_module_cls = None
compile_kwargs = self.fastvideo_args.torch_compile_kwargs or {}
self._maybe_compile_pipeline_module(
module_name="transformer",
fsdp_module_cls=fsdp_module_cls,
compile_kwargs=compile_kwargs,
)
self._maybe_compile_pipeline_module(
module_name="transformer_2",
fsdp_module_cls=fsdp_module_cls,
compile_kwargs=compile_kwargs,
)
if fsdp_module_cls is not None and isinstance(
transformer_module, fsdp_module_cls):
logger.info(
"Transformer is already FSDP-wrapped; skipping torch.compile in pipeline"
)
else:
logger.info("Enabling torch.compile for DiT with kwargs=%s",
compile_kwargs)
self.modules["transformer"] = torch.compile(
transformer_module, **compile_kwargs)
if "transformer_2" in self.modules:
transformer_module_2 = self.modules["transformer_2"]
if fsdp_module_cls is not None and isinstance(
transformer_module_2, fsdp_module_cls):
logger.info(
"Transformer_2 is already FSDP-wrapped; skipping torch.compile in pipeline"
)
else:
logger.info(
"Enabling torch.compile for Transformer_2 with kwargs=%s",
compile_kwargs)
self.modules["transformer_2"] = torch.compile(
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)
@@ -262,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
+7
View File
@@ -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",
+68
View File
@@ -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,
+2 -2
View File
@@ -242,8 +242,8 @@ class DecodingStage(PipelineStage):
decoded_frames = self.decode(cur_latent, fastvideo_args)
batch.trajectory_decoded.append(decoded_frames.cpu().float())
# Convert to float32 for compatibility
frames = frames.to(torch.float32)
# Convert to CPU float32 for compatibility
frames = frames.cpu().float()
# Crop padding if this is a LongCat refinement
if hasattr(batch, 'num_cond_frames_added') and hasattr(
-1
View File
@@ -314,7 +314,6 @@ class DenoisingStage(PipelineStage):
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
else:
t_expand = t.repeat(latent_model_input.shape[0])
t_expand = t_expand.to(get_local_torch_device())
use_meanflow = getattr(self.transformer.config, "use_meanflow",
False)
+139 -50
View File
@@ -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:
+355
View File
@@ -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",
]
+2 -5
View File
@@ -43,10 +43,7 @@ from fastvideo.training.training_utils import (
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
@@ -1270,7 +1267,7 @@ class DistillationPipeline(TrainingPipeline):
with torch.no_grad():
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output.cpu()
samples = output_batch.output
if self.rank_in_sp_group != 0:
continue
@@ -18,10 +18,7 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
+26 -13
View File
@@ -17,7 +17,6 @@ from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
WanCausalDMDPipeline)
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
@@ -58,17 +57,15 @@ class ODEInitTrainingPipeline(TrainingPipeline):
self.vae.requires_grad_(False)
self.timestep_shift = self.training_args.pipeline_config.flow_shift
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
self.noise_scheduler = SelfForcingFlowMatchScheduler(
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
training=True)
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
logger.info("dmd_denoising_steps: %s",
self.training_args.pipeline_config.dmd_denoising_steps)
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250, 0],
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
dtype=torch.long,
device=get_local_torch_device())
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
@@ -164,12 +161,27 @@ class ODEInitTrainingPipeline(TrainingPipeline):
device, dtype=torch.bfloat16)
training_batch.infos = infos
return training_batch, trajectory_latents[:, :, :self.training_args.
num_latent_t].to(
device,
dtype=torch.bfloat16
), trajectory_timesteps.to(
device)
return training_batch, trajectory_latents.to(
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
## TEMP Used for loading the sf .pt files directly
"""
self.manual_idx = self.manual_idx % 155
path = f"/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_pt_vidprom_1000/{self.manual_idx:05d}.pt"
logger.info("path: %s", path)
self.manual_idx += 1
# path = "/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_single_full/00000.pt"
b = torch.load(path)
training_batch.encoder_hidden_states = b["text_embedding"][0].unsqueeze(
0).to(device, dtype=torch.bfloat16)
trajectory_latents = b["ode_latent"].to(device, dtype=torch.bfloat16)
logger.info("trajectory_latents: %s", trajectory_latents.shape)
logger.info("encoder_hidden_states: %s",
training_batch.encoder_hidden_states.shape)
assert trajectory_latents.shape[1] <= 10, "trajectory_latents.shape[1] must be <= 10"
return training_batch, trajectory_latents.to(
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
"""
def _get_timestep(self,
min_timestep: int,
@@ -213,7 +225,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
if self._cached_closest_idx_per_dmd is None:
self._cached_closest_idx_per_dmd = torch.tensor(
[0, 12, 24, 36, S - 1], dtype=torch.long).cpu()
[0, 12, 24, 36], dtype=torch.long).cpu()
# [0, 1, 2, 3], dtype=torch.long).cpu()
logger.info("self._cached_closest_idx_per_dmd: %s",
self._cached_closest_idx_per_dmd)
@@ -355,7 +367,8 @@ class ODEInitTrainingPipeline(TrainingPipeline):
assert latent_key in latents_vis_dict and latents_vis_dict[
latent_key] is not None
latent = latents_vis_dict[latent_key]
pixel_latent = self.decoding_stage.decode(latent, training_args)
pixel_latent = self.validation_pipeline.decoding_stage.decode(
latent, training_args)
video = pixel_latent.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
@@ -30,10 +30,7 @@ from fastvideo.profiler import profile_region
logger = init_logger(__name__)
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
class SelfForcingDistillationPipeline(DistillationPipeline):
+34 -47
View File
@@ -13,18 +13,16 @@ import numpy as np
import torch
import torch.distributed as dist
import torchvision
from diffusers import FlowMatchEulerDiscreteScheduler
from einops import rearrange
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
import fastvideo.envs as envs
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder)
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
except Exception:
pass
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder)
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import build_parquet_map_style_dataloader
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
@@ -50,12 +48,8 @@ from fastvideo.training.training_utils import (
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
set_random_seed, shallow_asdict)
try:
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
except Exception:
vsa_available = False
vmoba_available = False
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
logger = init_logger(__name__)
@@ -114,7 +108,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Set random seeds for deterministic training
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed + self.global_rank)
set_random_seed(self.seed)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
self.transformer = apply_activation_checkpointing(
@@ -594,15 +588,15 @@ class TrainingPipeline(LoRAPipeline, ABC):
round(num_trainable_params / 1e9, 3))
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed + self.global_rank)
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
self.noise_gen_cuda = torch.Generator(
device=current_platform.device_name).manual_seed(self.seed +
self.global_rank)
device=current_platform.device_name).manual_seed(self.seed)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed + self.global_rank)
logger.info("Initialized random seeds with seed: %s",
self.seed + self.global_rank)
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
@@ -667,31 +661,26 @@ class TrainingPipeline(LoRAPipeline, ABC):
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
}
try:
metrics["batch_size"] = int(
training_batch.raw_latent_shape[0])
metrics["batch_size"] = int(training_batch.raw_latent_shape[0])
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
seq_len = (
training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] // patch_h) * (
training_batch.raw_latent_shape[4] // patch_w)
if training_batch.encoder_hidden_states is not None:
context_len = int(
training_batch.encoder_hidden_states.shape[1])
else:
context_len = 0
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] //
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
if training_batch.encoder_hidden_states is not None:
context_len = int(
training_batch.encoder_hidden_states.shape[1])
else:
context_len = 0
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
arch_config = self.training_args.pipeline_config.dit_config.arch_config
arch_config = self.training_args.pipeline_config.dit_config.arch_config
metrics["hidden_dim"] = arch_config.hidden_size
metrics["num_layers"] = arch_config.num_layers
metrics["ffn_dim"] = arch_config.ffn_dim
except Exception:
pass
metrics["hidden_dim"] = arch_config.hidden_size
metrics["num_layers"] = arch_config.num_layers
metrics["ffn_dim"] = arch_config.ffn_dim
self.tracker.log(metrics, step)
if step % self.training_args.training_state_checkpointing_steps == 0:
@@ -704,14 +693,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.noise_random_generator)
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_visualization and step % self.training_args.visualization_steps == 0:
self.visualize_intermediate_latents(training_batch,
self.training_args, step)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
with self.profiler_controller.region(
"profiler_region_training_validation"):
if self.training_args.log_visualization:
self.visualize_intermediate_latents(
training_batch, self.training_args, step)
self._log_validation(self.transformer, self.training_args,
step)
gpu_memory_usage = current_platform.get_torch_device(
@@ -870,7 +857,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Run validation inference
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output.cpu()
samples = output_batch.output
if self.rank_in_sp_group != 0:
continue
+51 -191
View File
@@ -257,6 +257,8 @@ def save_distillation_checkpoint(
if generator_scheduler is not None:
generator_states["scheduler"] = SchedulerWrapper(
generator_scheduler)
if generator_ema is not None:
generator_states["ema"] = generator_ema.state_dict()
generator_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
"generator")
@@ -288,6 +290,8 @@ def save_distillation_checkpoint(
if generator_scheduler_2 is not None:
generator_2_states["scheduler"] = SchedulerWrapper(
generator_scheduler_2)
if generator_ema_2 is not None:
generator_2_states["ema"] = generator_ema_2.state_dict()
generator_2_dcp_dir = os.path.join(save_dir,
"distributed_checkpoint",
@@ -413,67 +417,6 @@ def save_distillation_checkpoint(
rank,
local_main_process_only=False)
# Persist EMA separately to avoid shape mismatches across ranks.
# Supports:
# - mode="rank0_full": save consolidated EMA only on rank 0
# - mode="local_shard": save per-rank EMA shard for each rank
try:
if generator_ema is not None and getattr(generator_ema, "mode",
None) == "rank0_full":
_save_rank0_full_ema_safetensors(generator_ema,
generator_transformer, rank,
save_dir, "generator_ema")
elif generator_ema is not None and getattr(generator_ema, "mode",
None) == "local_shard":
# Save per-rank shard
ema_dir_shard = os.path.join(save_dir, "ema_local_shard")
os.makedirs(ema_dir_shard, exist_ok=True)
ema_shard_path = os.path.join(ema_dir_shard,
f"generator_ema_rank{rank}.pt")
torch.save(generator_ema.state_dict(), ema_shard_path)
logger.info(
"rank: %s, saved generator EMA shard (local_shard) to %s",
rank,
ema_shard_path,
local_main_process_only=False)
# Also consolidate EMA to a single full-state file on rank 0 by applying EMA to the model and gathering
_consolidate_local_shard_ema_and_save_safetensors(
generator_ema, generator_transformer, rank, save_dir,
"generator_ema")
except Exception as e:
logger.warning("rank: %s, failed saving EMA separately: %s", rank,
str(e))
try:
if generator_ema_2 is not None and getattr(generator_ema_2, "mode",
None) == "rank0_full":
_save_rank0_full_ema_safetensors(generator_ema_2,
generator_transformer_2, rank,
save_dir, "generator_ema_2")
elif generator_ema_2 is not None and getattr(generator_ema_2, "mode",
None) == "local_shard":
# Save per-rank shard for EMA_2
ema_dir_shard_2 = os.path.join(save_dir, "ema_local_shard")
os.makedirs(ema_dir_shard_2, exist_ok=True)
ema2_shard_path = os.path.join(ema_dir_shard_2,
f"generator_ema_2_rank{rank}.pt")
torch.save(generator_ema_2.state_dict(), ema2_shard_path)
logger.info(
"rank: %s, saved generator_2 EMA shard (local_shard) to %s",
rank,
ema2_shard_path,
local_main_process_only=False)
# Also consolidate EMA_2 to a single full-state file on rank 0
_consolidate_local_shard_ema_and_save_safetensors(
generator_ema_2, generator_transformer_2, rank, save_dir,
"generator_ema_2")
except Exception as e:
logger.warning("rank: %s, failed saving EMA_2 separately: %s", rank,
str(e))
# Save generator model weights (consolidated) for inference
cpu_state = gather_state_dict_on_cpu_rank0(generator_transformer,
device=None)
@@ -511,45 +454,46 @@ def save_distillation_checkpoint(
logger.info("--> distillation checkpoint saved at step %s to %s", step,
weight_path)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
inference_save_dir_2 = os.path.join(
save_dir, "generator_2_inference_transformer")
cpu_state_2 = gather_state_dict_on_cpu_rank0(generator_transformer_2,
device=None)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
inference_save_dir_2 = os.path.join(
save_dir, "generator_2_inference_transformer")
cpu_state_2 = gather_state_dict_on_cpu_rank0(
generator_transformer_2, device=None)
if rank == 0:
os.makedirs(inference_save_dir_2, exist_ok=True)
weight_path_2 = os.path.join(inference_save_dir_2,
"diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
rank,
weight_path_2,
local_main_process_only=False)
if rank == 0:
os.makedirs(inference_save_dir_2, exist_ok=True)
weight_path_2 = os.path.join(
inference_save_dir_2, "diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Convert training format to diffusers format and save
diffusers_state_dict_2 = custom_to_hf_state_dict(
cpu_state_2,
generator_transformer_2.reverse_param_names_mapping)
save_file(diffusers_state_dict_2, weight_path_2)
# Convert training format to diffusers format and save
diffusers_state_dict_2 = custom_to_hf_state_dict(
cpu_state_2,
generator_transformer_2.reverse_param_names_mapping)
save_file(diffusers_state_dict_2, weight_path_2)
logger.info(
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
rank,
weight_path_2,
local_main_process_only=False)
logger.info(
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Save model config
config_dict_2 = generator_transformer_2.hf_config
if "dtype" in config_dict_2:
del config_dict_2["dtype"] # TODO
config_path_2 = os.path.join(inference_save_dir_2, "config.json")
with open(config_path_2, "w") as f:
json.dump(config_dict_2, f, indent=4)
logger.info(
"--> generator_2 distillation checkpoint saved at step %s to %s",
step, weight_path_2)
# Save model config
config_dict_2 = generator_transformer_2.hf_config
if "dtype" in config_dict_2:
del config_dict_2["dtype"] # TODO
config_path_2 = os.path.join(inference_save_dir_2,
"config.json")
with open(config_path_2, "w") as f:
json.dump(config_dict_2, f, indent=4)
logger.info(
"--> generator_2 distillation checkpoint saved at step %s to %s",
step, weight_path_2)
def load_checkpoint(transformer,
@@ -700,37 +644,18 @@ def load_distillation_checkpoint(
end_time - begin_time,
local_main_process_only=False)
# Load EMA separately if saved in rank0_full mode
# Load EMA state if available and generator_ema is provided
if generator_ema is not None:
try:
if getattr(generator_ema, "mode", None) == "rank0_full":
ema_path = os.path.join(checkpoint_path, "ema",
"generator_ema.pt")
if rank == 0 and os.path.exists(ema_path):
ema_state = torch.load(ema_path, map_location="cpu")
generator_ema.load_state_dict(ema_state)
logger.info(
"rank: %s, generator EMA (rank0_full) loaded from %s",
rank, ema_path)
elif rank == 0:
logger.info(
"rank: %s, generator EMA file not found at %s; skipping",
rank, ema_path)
elif getattr(generator_ema, "mode", None) == "local_shard":
ema_path = os.path.join(checkpoint_path, "ema_local_shard",
f"generator_ema_rank{rank}.pt")
if os.path.exists(ema_path):
ema_state = torch.load(ema_path, map_location="cpu")
generator_ema.load_state_dict(ema_state)
logger.info(
"rank: %s, generator EMA shard (local_shard) loaded from %s",
rank, ema_path)
else:
logger.info(
"rank: %s, generator EMA shard file not found at %s; skipping",
rank, ema_path)
ema_state = generator_states.get("ema")
if ema_state is not None:
generator_ema.load_state_dict(ema_state)
logger.info("rank: %s, generator EMA state loaded successfully",
rank)
else:
logger.info("rank: %s, no EMA state found in checkpoint", rank)
except Exception as e:
logger.warning("rank: %s, failed to load generator EMA: %s", rank,
logger.warning("rank: %s, failed to load EMA state: %s", rank,
str(e))
# Load generator_2 distributed checkpoint (MoE support)
@@ -925,7 +850,7 @@ def load_distillation_checkpoint(
def normalize_dit_input(model_type, latents, vae) -> torch.Tensor:
if model_type == "hunyuan_hf" or model_type == "hunyuan":
return latents * vae.config.scaling_factor
return latents * 0.476986
elif model_type == "wan":
latents_mean = torch.tensor(vae.latents_mean)
latents_std = 1.0 / torch.tensor(vae.latents_std)
@@ -1244,71 +1169,6 @@ def custom_to_hf_state_dict(
return new_state_dict
def _save_full_ema_safetensors_from_state(
state_dict: dict[str, Any],
reverse_param_names_mapping: dict[str, tuple[str, int, int]],
output_path: str,
) -> None:
"""
Convert a training-format state_dict to HF format and save as safetensors.
"""
diffusers_state_dict = custom_to_hf_state_dict(state_dict,
reverse_param_names_mapping)
save_file(diffusers_state_dict, output_path)
def _save_rank0_full_ema_safetensors(
ema: "EMA_FSDP",
module,
rank: int,
save_dir: str,
base_name: str,
) -> None:
if rank != 0:
return
ema_dir = os.path.join(save_dir, "ema")
os.makedirs(ema_dir, exist_ok=True)
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
ema_state = ema.state_dict()
_save_full_ema_safetensors_from_state(ema_state,
module.reverse_param_names_mapping,
output_path)
logger.info("rank: %s, saved %s as consolidated EMA safetensors to %s",
rank,
base_name,
output_path,
local_main_process_only=False)
def _consolidate_local_shard_ema_and_save_safetensors(
ema: "EMA_FSDP",
module,
rank: int,
save_dir: str,
base_name: str,
) -> None:
try:
# Temporarily apply EMA to the live (sharded) module and gather full CPU state on rank 0
with ema.apply_to_model(module):
cpu_state_full = gather_state_dict_on_cpu_rank0(module, device=None)
if rank == 0:
ema_dir = os.path.join(save_dir, "ema")
os.makedirs(ema_dir, exist_ok=True)
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
_save_full_ema_safetensors_from_state(
cpu_state_full, module.reverse_param_names_mapping, output_path)
logger.info(
"rank: %s, saved consolidated %s EMA (from local_shard) as safetensors to %s",
rank,
base_name,
output_path,
local_main_process_only=False)
except Exception as ce:
logger.warning(
"rank: %s, failed consolidating %s EMA (local_shard): %s", rank,
base_name, str(ce))
def shift_timestep(timestep: torch.Tensor, shift: float,
num_train_timestep: float) -> torch.Tensor:
if shift == 1:
@@ -1935,5 +1795,5 @@ class EMA_FSDP:
self.saved.clear()
return False
def apply_to_model(self, module: torch.nn.Module) -> _ApplyEMACtx:
def apply_to_model(self, module):
return EMA_FSDP._ApplyEMACtx(self, module)
@@ -10,10 +10,7 @@ from fastvideo.pipelines.basic.wan.wan_dmd_pipeline import WanDMDPipeline
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.utils import is_vsa_available
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
@@ -19,10 +19,7 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
@@ -18,10 +18,7 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
@@ -10,10 +10,7 @@ from fastvideo.training.self_forcing_distillation_pipeline import (
SelfForcingDistillationPipeline)
from fastvideo.utils import is_vsa_available
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
+1 -4
View File
@@ -10,10 +10,7 @@ from fastvideo.pipelines.basic.wan.wan_pipeline import WanPipeline
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
+1 -3
View File
@@ -650,10 +650,8 @@ class WorkerMultiprocProc:
logging_info = None
if envs.FASTVIDEO_STAGE_LOGGING:
logging_info = output_batch.logging_info
# result tensor shared by CUDA IPC to avoid serialization overhead
result = output_batch.output
self.pipe.send({
"output_batch": result,
"output_batch": output_batch.output.cpu(),
"logging_info": logging_info,
"extra": output_batch.extra,
})
@@ -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()
+7 -2
View File
@@ -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)
+29 -5
View File
@@ -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()
+1
View File
@@ -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))