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
201 changed files with 3654 additions and 12428 deletions
@@ -67,9 +67,6 @@ jobs:
- torch-version: '2.9.1'
cuda-version: '12.8.0'
torch-cuda-short: 'cu128'
# - torch-version: '2.10.0'
# cuda-version: '12.8.0'
# torch-cuda-short: 'cu128'
steps:
- name: Free up disk space
@@ -223,6 +220,3 @@ jobs:
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: fastvideo-kernel/dist/
# PyPI does not allow replacing an existing file with the same name.
# This makes re-runs idempotent by skipping files already uploaded.
skip-existing: true
-11
View File
@@ -18,7 +18,6 @@ venv/
.venv/
runs/
samples/
Miniconda3-latest-Linux-x86_64.sh
*validation/
data/
outputs/
@@ -33,11 +32,6 @@ env
**.txt
*.log
weights/
official_weights/
converted_weights/
# SSIM test outputs
fastvideo/tests/ssim/generated_videos/
# Distribution / packaging
build/
@@ -75,11 +69,6 @@ docs/distillation/examples/
!docs/assets/images/**/*.png
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
!assets/images/**/*.png
!assets/images/**/*.jpg
!assets/images/**/*.jpeg
!assets/images/**/*.gif
!assets/videos/**/*.mp4
dmd_t2v_output/
preprocess_output_text/
+1 -1
View File
@@ -10,7 +10,7 @@ exclude: |
demo/.*|
predict\.py|
scripts/.*|
assets/prompts/.*|
prompts/.*|
fastvideo/data_preprocess/.*|
fastvideo/dataset/.*|
fastvideo/models/.*|
-42
View File
@@ -1,42 +0,0 @@
# Repository Guidelines
## Project Structure & Module Organization
- Core Python package: `fastvideo/` (models, pipelines, training, distributed runtime, CLI entrypoints).
- CUDA/custom kernels: `fastvideo-kernel/` (separate build/test flow).
- Tests:
- `fastvideo/tests/` for package-level tests (dataset, encoders, inference, training, SSIM, workflow).
- `tests/local_tests/` for additional local/component checks.
- Docs and guides: `docs/` (MkDocs source), with contributor docs in `docs/contributing/`.
- Runnable examples and scripts: `examples/` and `scripts/`.
- Static assets: `assets/` (including `assets/images/`, `assets/videos/`, and `assets/prompts/`) and `comfyui/assets/`.
## Build, Test, and Development Commands
- `uv pip install -e .[dev]`: editable install with lint/test extras.
- `pre-commit install --hook-type pre-commit --hook-type commit-msg`: enable local hooks.
- `pre-commit run --all-files`: run formatter/lint/type/spelling checks.
- `pytest tests/`: run top-level test suite.
- `pytest fastvideo/tests/ -v`: run package tests.
- `pytest fastvideo/tests/ssim/ -vs`: run SSIM regression tests (GPU-heavy).
- `cd fastvideo-kernel && ./build.sh`: build kernel extensions.
## Coding Style & Naming Conventions
- Python 3.10+; 4-space indentation; keep code and imports readable and explicit.
- Style tools are configured in `pyproject.toml` and `.pre-commit-config.yaml`:
- `yapf` (format), `ruff` (lint, auto-fix), `mypy` (typing), `codespell`.
- Target line length is 80.
- Naming: `snake_case` for functions/files, `PascalCase` for classes, `UPPER_SNAKE_CASE` for constants.
## Testing Guidelines
- Use `pytest` and place tests near relevant domains (e.g., `fastvideo/tests/encoders/`).
- Prefer descriptive names like `test_<feature>_<expected_behavior>.py`.
- For new pipelines/backends, include at least one regression-oriented test; add SSIM coverage when output quality must be preserved.
- Document GPU assumptions in tests that require specific hardware.
## Commit & Pull Request Guidelines
- Follow existing commit style: short subject with optional tag prefix, e.g. `[bugfix]: ...`, `[feat]: ...`, `[misc]: ...`, and include PR reference like `(#1234)` when applicable.
- Keep commits focused by concern (feature, refactor, fix).
- PRs should include:
- clear problem/solution summary,
- test evidence (`pytest`/SSIM outputs or rationale if skipped),
- linked issue/PR context,
- screenshots or sample outputs for UI/demo/docs changes.
-1
View File
@@ -1 +0,0 @@
@AGENTS.md
+1 -6
View File
@@ -1,10 +1,5 @@
<div align="center">
<img src=assets/logos/logo.svg width="30%"/>
</div>
<p align="center">
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://github.com/hao-ai-lab/FastVideo/discussions/1097" target="_blank"> <b> WeChat </b> </a> |
</p>
| **[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.**
+195
View File
@@ -0,0 +1,195 @@
import argparse
import os
import tempfile
import gradio as gr
import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import export_to_video
from fastvideo.distill.solver import PCMFMScheduler
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
def init_args():
parser = argparse.ArgumentParser()
parser.add_argument("--prompts", nargs="+", default=[])
parser.add_argument("--num_frames", type=int, default=25)
parser.add_argument("--height", type=int, default=480)
parser.add_argument("--width", type=int, default=848)
parser.add_argument("--num_inference_steps", type=int, default=8)
parser.add_argument("--guidance_scale", type=float, default=4.5)
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--seed", type=int, default=12345)
parser.add_argument("--transformer_path", type=str, default=None)
parser.add_argument("--scheduler_type", type=str, default="pcm_linear_quadratic")
parser.add_argument("--lora_checkpoint_dir", type=str, default=None)
parser.add_argument("--shift", type=float, default=8.0)
parser.add_argument("--num_euler_timesteps", type=int, default=50)
parser.add_argument("--linear_threshold", type=float, default=0.1)
parser.add_argument("--linear_range", type=float, default=0.75)
parser.add_argument("--cpu_offload", action="store_true")
return parser.parse_args()
def load_model(args):
if args.scheduler_type == "euler":
scheduler = FlowMatchEulerDiscreteScheduler()
else:
linear_quadratic = True if "linear_quadratic" in args.scheduler_type else False
scheduler = PCMFMScheduler(
1000,
args.shift,
args.num_euler_timesteps,
linear_quadratic,
args.linear_threshold,
args.linear_range,
)
if args.transformer_path:
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder="transformer/")
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
pipe.enable_vae_tiling()
# pipe.to(device)
# if args.cpu_offload:
pipe.enable_sequential_cpu_offload()
return pipe
def generate_video(
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed=False,
):
if randomize_seed:
seed = torch.randint(0, 1000000, (1, )).item()
generator = torch.Generator(device="cuda").manual_seed(seed)
if not use_negative_prompt:
negative_prompt = None
with torch.autocast("cuda", dtype=torch.bfloat16):
output = pipe(
prompt=[prompt],
negative_prompt=negative_prompt,
height=height,
width=width,
num_frames=num_frames,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=generator,
).frames[0]
output_path = os.path.join(tempfile.mkdtemp(), "output.mp4")
export_to_video(output, output_path, fps=30)
return output_path, seed
examples = [
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
]
args = init_args()
pipe = load_model(args)
print("load model successfully")
with gr.Blocks() as demo:
gr.Markdown("# Fastvideo Mochi Video Generation Demo")
with gr.Group():
with gr.Row():
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=1,
placeholder="Enter your prompt",
container=False,
)
run_button = gr.Button("Run", scale=0)
result = gr.Video(label="Result", show_label=False)
with gr.Accordion("Advanced options", open=False):
with gr.Group():
with gr.Row():
height = gr.Slider(
label="Height",
minimum=256,
maximum=1024,
step=32,
value=args.height,
)
width = gr.Slider(label="Width", minimum=256, maximum=1024, step=32, value=args.width)
with gr.Row():
num_frames = gr.Slider(
label="Number of Frames",
minimum=21,
maximum=163,
value=args.num_frames,
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=args.guidance_scale,
)
num_inference_steps = gr.Slider(
label="Inference Steps",
minimum=4,
maximum=100,
value=args.num_inference_steps,
)
with gr.Row():
use_negative_prompt = gr.Checkbox(label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=1,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed)
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
seed_output = gr.Number(label="Used Seed")
gr.Examples(examples=examples, inputs=prompt)
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=negative_prompt,
)
run_button.click(
fn=generate_video,
inputs=[
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed,
],
outputs=[result, seed_output],
)
if __name__ == "__main__":
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
+15
View File
@@ -0,0 +1,15 @@
Fast-Hunyuan comparison with original Hunyuan, achieving an 8X diffusion speed boost with the FastVideo framework.
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
Fast-Mochi comparison with original Mochi, achieving an 8X diffusion speed boost with the FastVideo framework.
https://github.com/user-attachments/assets/5fbc4596-56d6-43aa-98e0-da472cf8e26c
Comparison between OpenAI Sora, original Hunyuan and FastHunyuan
https://github.com/user-attachments/assets/d323b712-3f68-42b2-952b-94f6a49c4836
Comparison between original FastHunyuan, LLM-INT8 quantized FastHunyuan and NF4 quantized FastHunyuan
https://github.com/user-attachments/assets/cf89efb5-5f68-4949-a085-f41c1ef26c94
@@ -20,7 +20,7 @@ def main():
sampling_param = SamplingParam.from_pretrained(model_path)
# image2world example from official repo
image_path = "assets/images/bus_terminal.jpg"
image_path = "images/bus_terminal.jpg"
prompt = (
"A nighttime city bus terminal gradually shifts from stillness to subtle movement. "
@@ -48,3 +48,4 @@ def main():
if __name__ == "__main__":
main()
@@ -20,7 +20,7 @@ def main():
sampling_param = SamplingParam.from_pretrained(model_path)
# video2world example from official repo
video_path = "assets/videos/robot_pouring.mp4"
video_path = "videos/robot_pouring.mp4"
prompt = (
"A robotic arm, primarily white with black joints and cables, is shown in a clean, modern indoor setting with a white tabletop. "
@@ -51,3 +51,4 @@ def main():
if __name__ == "__main__":
main()
-119
View File
@@ -1,119 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Basic inference script for HunyuanGameCraft video generation.
HunyuanGameCraft generates game-like videos with camera/action control.
It takes an optional image input and generates video with camera motion
based on simple action commands (forward, left, right, backward, rotations).
Available actions:
- forward (w): Move camera forward
- backward (s): Move camera backward
- left (a): Move camera left (strafe)
- right (d): Move camera right (strafe)
- left_rot: Rotate camera left (pan)
- right_rot: Rotate camera right (pan)
- up_rot: Rotate camera up (tilt)
- down_rot: Rotate camera down (tilt)
T2V vs I2V:
- Default: I2V (uses a default reference image). Set GAMECRAFT_I2V_IMAGE to a
URL or path to use a different image.
- T2V only (no reference image): run with GAMECRAFT_I2V_IMAGE= (empty).
"""
import os
import torch
from fastvideo import VideoGenerator
from fastvideo.models.camera import create_camera_trajectory
# Model configuration (use GAMECRAFT_MODEL_PATH for local weights)
MODEL_PATH = os.environ.get("GAMECRAFT_MODEL_PATH", "FastVideo/HunyuanGameCraft-Diffusers")
# Default prompts for demo
DEFAULT_PROMPTS = {
"village": "A charming medieval village with cobblestone streets, thatched-roof houses, and vibrant flower gardens under a bright blue sky.",
"temple": "A majestic ancient temple stands under a clear blue sky, its grandeur highlighted by towering Doric columns and intricate architectural details.",
"forest": "A lush green forest with tall trees, dappled sunlight filtering through the leaves, and a winding dirt path.",
"beach": "A tropical beach with crystal clear turquoise water, white sand, and palm trees swaying in the breeze.",
}
# I2V: default reference image (URL). Can override with a local path.
DEFAULT_I2V_IMAGE_URL = (
"https://huggingface.co/datasets/huggingface/documentation-images/"
"resolve/main/diffusers/astronaut.jpg"
)
DEFAULT_I2V_PROMPT = (
"An astronaut hatching from an egg, on the surface of the moon, "
"the darkness and depth of space realised in the background."
)
OUTPUT_PATH = "video_samples_gamecraft"
def main():
# Initialize generator
# FastVideo will automatically download weights from HuggingFace
generator = VideoGenerator.from_pretrained(
MODEL_PATH,
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
)
# Video parameters
height = 704
width = 1280
num_frames = 33
action = "forward"
action_speed = 0.2
# Create camera trajectory (Plücker coordinates)
camera_states = create_camera_trajectory(
action=action,
height=height,
width=width,
num_frames=num_frames,
action_speed=action_speed,
dtype=torch.bfloat16,
)
print(f"Camera states shape: {camera_states.shape}")
# I2V vs T2V: unset GAMECRAFT_I2V_IMAGE -> I2V (default image). Set to "" -> T2V.
env_image = os.environ.get("GAMECRAFT_I2V_IMAGE")
if env_image is None:
image_path = DEFAULT_I2V_IMAGE_URL # default: I2V
elif env_image.strip() == "":
image_path = None # T2V
else:
image_path = env_image.strip() # I2V with given URL/path
is_i2v = image_path is not None
prompt = DEFAULT_I2V_PROMPT if is_i2v else DEFAULT_PROMPTS["temple"]
print(f"Mode: {'I2V' if is_i2v else 'T2V'}, prompt: {prompt[:60]}...")
gen_kw = dict(
prompt=prompt,
negative_prompt="",
camera_states=camera_states,
height=height,
width=width,
num_frames=num_frames,
num_inference_steps=50,
guidance_scale=6.0,
seed=42,
fps=24,
output_path=OUTPUT_PATH,
save_video=True,
)
if is_i2v:
gen_kw["image_path"] = image_path
generator.generate_video(**gen_kw)
if __name__ == "__main__":
main()
@@ -1,49 +0,0 @@
from fastvideo import VideoGenerator
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_lingbotworld"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"FastVideo/LingBot-World-Base-Cam-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=False, # set to True if GPU is out of memory
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
num_frames = 81
prompt = "The video presents a soaring journey through a fantasy jungle. The wind whips past the rider's blue hands gripping the reins, causing the leather straps to vibrate. The ancient gothic castle approaches steadily, its stone details becoming clearer against the backdrop of floating islands and distant waterfalls."
image_path = "https://raw.githubusercontent.com/Robbyant/lingbot-world/main/examples/00/image.jpg"
action_path = "examples/inference/basic/lingbotworld_examples/00"
c2ws_plucker_emb, num_frames = prepare_camera_embedding(
action_path=action_path,
num_frames=num_frames,
height=480,
width=832,
spatial_scale=8,
)
generator.generate_video(
prompt,
image_path=image_path,
output_path=OUTPUT_PATH,
save_video=True,
num_frames=num_frames,
height=480,
width=832,
c2ws_plucker_emb=c2ws_plucker_emb,
)
if __name__ == "__main__":
main()
+3 -14
View File
@@ -1,6 +1,5 @@
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 "
@@ -17,26 +16,16 @@ PROMPT = (
def main() -> None:
# Uses FastVideo default sampling settings for LTX2 base.
generator = VideoGenerator.from_pretrained(
"Davids048/LTX2-Base-Diffusers",
num_gpus=8,
"FastVideo/LTX2-Distilled-Diffusers",
num_gpus=1,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.4.mp4"
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
num_frames=121,
height=1088,
width=1920,
# LTX2 uses these parameters for multi-modal CFG instead of guidance_scale
# ltx2_cfg_scale_video=3.0,
# ltx2_cfg_scale_audio=7.0,
# ltx2_modality_scale_video=3.0,
# ltx2_modality_scale_audio=3.0,
# ltx2_rescale_scale=0.7,
)
generator.shutdown()
@@ -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()
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo import VideoGenerator
PROMPT = (
@@ -14,17 +15,33 @@ PROMPT = (
"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(
"FastVideo/LTX2-Distilled-Diffusers",
num_gpus=1,
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,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
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()
-137
View File
@@ -1,137 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import argparse
import os
import re
from typing import List
DEFAULT_PROMPTS = [
"a photo of a cat",
"a cinematic photo of a red panda wearing a tiny backpack, standing on a rainy neon-lit street at night, shallow depth of field, sharp focus, 35mm, bokeh",
]
def _safe_filename(text: str, max_len: int = 100) -> str:
"""
Make a stable, filesystem-friendly filename base.
VideoGenerator uses prompt[:100].strip() internally, so we mirror that,
but also remove path separators and other problematic characters.
"""
s = text[:max_len].strip()
s = s.replace(os.sep, "_")
if os.altsep:
s = s.replace(os.altsep, "_")
s = re.sub(r"\s+", " ", s)
s = re.sub(r"[^A-Za-z0-9 .,_-]", "_", s)
s = s.strip(" .")
return s or "prompt"
def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
"""
Ensure deterministic naming by deleting any existing mp4s that would
cause VideoGenerator to append suffixes like _1, _2, etc.
"""
if not os.path.isdir(out_dir):
return
pattern = re.compile(rf"^{re.escape(filename_base)}(_\d+)?\.mp4$")
for fn in os.listdir(out_dir):
if pattern.match(fn):
try:
os.remove(os.path.join(out_dir, fn))
except FileNotFoundError:
pass
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Run SD3.5 Medium text-to-image with FastVideo VideoGenerator.")
p.add_argument("--model-path", default="stabilityai/stable-diffusion-3.5-medium", help="Path to local diffusers-format SD3.5 weights directory.")
p.add_argument(
"--out-dir",
"--outdir",
default="outputs/sd35/samples",
help="Output directory for generated mp4 files.",
)
p.add_argument(
"--prompt",
action="append",
default=None,
help="Prompt text. Repeat --prompt multiple times to generate multiple samples.",
)
p.add_argument("--negative", default="lowres, blurry, jpeg artifacts, watermark, text", help="Negative prompt.")
p.add_argument(
"--backend",
default=None,
help="Set FASTVIDEO_ATTENTION_BACKEND (e.g. TORCH_SDPA). If omitted, respects the existing env var.",
)
p.add_argument("--seed", type=int, default=42, help="Base seed. Each prompt uses seed + prompt_idx.")
p.add_argument("--height", type=int, default=768, help="Output height.")
p.add_argument("--width", type=int, default=768, help="Output width.")
p.add_argument("--steps", type=int, default=28, help="Number of inference steps.")
p.add_argument("--guidance", type=float, default=6.0, help="Guidance scale.")
p.add_argument("--num-gpus", type=int, default=1, help="Number of GPUs to use.")
return p.parse_args()
def main() -> None:
args = parse_args()
prompts: List[str] = args.prompt if args.prompt else DEFAULT_PROMPTS
if args.backend:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
from fastvideo import VideoGenerator
os.makedirs(args.out_dir, exist_ok=True)
init_kwargs = {
"num_gpus": args.num_gpus,
"workload_type": "t2i",
"sp_size": 1,
"tp_size": 1,
"dit_cpu_offload": False,
"dit_layerwise_offload": False,
"text_encoder_cpu_offload": False,
"vae_cpu_offload": False,
"image_encoder_cpu_offload": False,
"pin_cpu_memory": False,
"use_fsdp_inference": False,
}
generator = VideoGenerator.from_pretrained(model_path=args.model_path, **init_kwargs)
try:
for i, prompt in enumerate(prompts):
seed = args.seed + i
filename_base = f"sd35_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
_remove_existing_outputs(args.out_dir, filename_base)
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
print(f"[sd35] prompt_idx={i} seed={seed} output_path={output_path}")
generation_kwargs = {
"output_path": output_path,
"height": args.height,
"width": args.width,
"num_frames": 1,
"fps": 1,
"num_inference_steps": args.steps,
"guidance_scale": args.guidance,
"seed": seed,
"negative_prompt": args.negative,
"save_video": True,
}
generator.generate_video(prompt, **generation_kwargs)
print(f"[sd35] done. outputs written to: {args.out_dir}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -31,7 +31,7 @@ def main():
sampling_param.height = 480
sampling_param.seed = 1000
with open("assets/prompts/mixkit_i2v.jsonl", "r") as f:
with open("prompts/mixkit_i2v.jsonl", "r") as f:
prompt_image_pairs = json.load(f)
for prompt_image_pair in prompt_image_pairs:
@@ -188,9 +188,8 @@ def load_example_prompts():
prompt_to_image = {}
# Try to find the JSON file relative to project root
possible_json_paths = [
Path("assets/prompts/mixkit_i2v.jsonl"),
Path(__file__).resolve().parents[4] / "assets" / "prompts" /
"mixkit_i2v.jsonl",
Path("prompts/mixkit_i2v.jsonl"),
Path(__file__).parent.parent.parent.parent / "prompts" / "mixkit_i2v.jsonl",
]
json_path = None
for path in possible_json_paths:
@@ -202,8 +201,8 @@ def load_example_prompts():
try:
with open(json_path, "r", encoding='utf-8') as f:
data = json.load(f)
# Resolve paths relative to repository root.
project_root = Path(__file__).resolve().parents[4]
# Get the project root directory (parent of prompts directory)
project_root = json_path.parent.parent
for item in data:
prompt_text = item.get("prompt", "").strip()
image_path = item.get("image_path", "")
@@ -737,8 +736,8 @@ def main():
allowed_paths=[
os.path.abspath("outputs"),
os.path.abspath("fastvideo-logos"),
os.path.abspath("assets/prompts"),
os.path.abspath("assets/images"),
os.path.abspath("prompts"),
os.path.abspath("images"),
os.path.abspath(tempfile.gettempdir()),
os.path.abspath(os.path.join(tempfile.gettempdir(), "gradio")),
]
@@ -748,4 +747,4 @@ def main():
if __name__ == "__main__":
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,23 +0,0 @@
# LTX-2 Crush-Smol Example
# TODO: Update this doc.
These are e2e example scripts for finetuning LTX-2 on the crush-smol dataset.
## Execute the following commands from `FastVideo/` to run training:
### Download crush-smol dataset:
`bash examples/training/finetune/ltx2/overfit/download_dataset.sh`
### Preprocess the videos and captions into latents:
`bash examples/training/finetune/ltx2/overfit/preprocess_ltx2_data_t2v_new.sh`
### Edit the following file and run finetuning:
`bash examples/training/finetune/ltx2/overfit/finetune_t2v.sh`
Notes:
- Update `DATASET_PATH` in the preprocess script to point to your merged dataset root (`videos/` + `videos2caption.json`).
- `MODEL_PATH` should point to a local LTX-2 diffusers-style directory that contains `model_index.json` and `text_encoder/gemma`.
@@ -1,3 +0,0 @@
# #!/bin/bash
#
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -1,95 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="Davids048/LTX2-Base-Diffusers"
# Also can use simple 1 video for overfitting experiments.
# DATA_DIR="/home/hal-jundas/codes/FastVideo/data/crush-smol"
DATA_DIR="<PATH_TO_PROCESSED_DATASET>"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
echo VALIDATION_DATASET_FILE: $VALIDATION_DATASET_FILE
NUM_GPUS=4
OVERFIT_HEIGHT=480
OVERFIT_WIDTH=832
OVERFIT_FRAMES=73
training_args=(
--tracker_project_name "ltx2_t2v_finetune"
--output_dir "checkpoints/ltx2_t2v_finetune"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 10
--num_height $OVERFIT_HEIGHT
--num_width $OVERFIT_WIDTH
--num_frames $OVERFIT_FRAMES
--ltx2-first-frame-conditioning-p 0.1
--enable_gradient_checkpointing_type "full"
--mode "finetuning"
)
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
)
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
dataset_args=(
--data_path $DATA_DIR
--dataloader_num_workers 1
)
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "3.0"
)
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
--lr_scheduler "linear"
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--dit_precision "fp32"
--dit_cpu_offload False
--dit_layerwise_offload False
--text_encoder_cpu_offload False
--image_encoder_cpu_offload False
--vae_cpu_offload False
)
# NOTE: Setting this environment variable to TORCH_SDPA to avoid the issue of stacking that failed in flash attn.
export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ltx2_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,80 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="/path/to/LTX-2"
DATA_DIR="data/crush-smol"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
training_args=(
--tracker_project_name "ltx2_t2v_lora_finetune"
--output_dir "checkpoints/ltx2_t2v_lora_finetune"
--max_train_steps 2000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 8
--num_latent_t 10
--num_height 480
--num_width 832
--num_frames 77
--ltx2-first-frame-conditioning-p 0.1
--enable_gradient_checkpointing_type "full"
)
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
)
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
dataset_args=(
--data_path $DATA_DIR
--dataloader_num_workers 1
)
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 200
--validation_sampling_steps "50"
--validation_guidance_scale "3.0"
)
optimizer_args=(
--learning_rate 2e-4
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
--lora_training True
--lora_rank 16
--lora_alpha 16
)
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ltx2_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,35 +0,0 @@
#!/bin/bash
GPU_NUM=1
MODEL_PATH="Davids048/LTX2-Base-Diffusers"
# DATASET_PATH="data/overfit"
DATASET_PATH="data/crush-smol"
OUTPUT_DIR="$DATASET_PATH"
WITH_AUDIO=true
# Convert one-file overfit metadata into merged format if needed.
if [ ! -f "$DATASET_PATH/videos2caption.json" ] && [ -f "$DATASET_PATH/overfit.json" ]; then
python scripts/dataset_preparation/convert_to_merged_dataset.py \
--items-json "$DATASET_PATH/overfit.json" \
--output-dir "$DATASET_PATH"
fi
torchrun --nproc_per_node=$GPU_NUM \
--master_port=29513 \
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
--model_path $MODEL_PATH \
--mode preprocess \
--workload_type t2v \
--preprocess.video_loader_type torchvision \
--preprocess.dataset_type merged \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
--preprocess.with_audio $WITH_AUDIO \
--preprocess.preprocess_video_batch_size 1 \
--preprocess.dataloader_num_workers 0 \
--preprocess.max_height 480 \
--preprocess.max_width 832 \
--preprocess.num_frames 73 \
--preprocess.train_fps 16 \
--preprocess.video_length_tolerance_range 5
@@ -1,13 +0,0 @@
{
"data": [
{
"caption": "The camera opens in a calm, sunlit frog yoga studio. Warm morning light washes over the wooden floor as incense smoke drifts lazily in the air. The senior frog instructor sits cross-legged at the center, eyes closed, voice deep and calm. “We are one with the pond.” All the frogs answer softly: “Ommm...” “We are one with the mud.” “Ommm...” He smiles faintly. “We are one with the flies.” A quiet pause. The camera slowly pans to the side — one frog twitches, eyes darting. Suddenly — *thwip!* — its tongue snaps out, catching a fly mid-air and pulling it into its mouth. The master exhales slowly, still serene.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 1088,
"width": 1920,
"num_frames": 121
}
]
}
@@ -1,31 +0,0 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
+1 -1
View File
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
[project]
name = "fastvideo-kernel"
version = "0.2.6"
version = "0.2.5"
description = "Unified CUDA kernels for FastVideo"
readme = "README.md"
requires-python = ">=3.10"
@@ -287,8 +287,9 @@ def block_sparse_attn(
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
if (not _force_triton()) and _is_sm90() and (block_sparse_fwd is not None) and (block_sparse_bwd is not None):
return block_sparse_attn_sm90(q, k, v, block_map, variable_block_sizes)
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
# to a multiple of the block size (64 tokens).
# Triton path: generally assumes q/k/v share the same padded length
if q.shape[2] != k.shape[2] or q.shape[2] != v.shape[2]:
raise RuntimeError("Triton fallback requires q/k/v to have the same padded length.")
return block_sparse_attn_triton(q, k, v, block_map, variable_block_sizes)
@@ -141,6 +141,12 @@ def video_sparse_attn(
# Use autograd-enabled wrapper so backward works (and still uses SM90 kernel when available)
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
else:
if q_seq_len != kv_seq_len:
raise RuntimeError(
"q/k have different lengths, but the compiled CUDA kernel (block_sparse_fwd) "
"is not available. The Triton fallback currently requires q and k/v to have "
"the same padded length."
)
# Triton-only forward (kept for environments without the wrapper deps)
out_s, _ = triton_block_sparse_attn_forward(q, k, v, idx, num, variable_block_sizes)
@@ -29,7 +29,7 @@ configs = [
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
@triton.autotune(configs, key=["N_CTX_Q", "HEAD_DIM"])
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
@triton.jit
def _attn_fwd_sparse(
Q,
@@ -60,8 +60,7 @@ def _attn_fwd_sparse(
stride_on,
Z,
H,
N_CTX_Q, #
N_CTX_KV, #
N_CTX, #
HEAD_DIM: tl.constexpr, #
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
@@ -76,29 +75,24 @@ def _attn_fwd_sparse(
off_hz = tl.program_id(1) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX_Q // BLOCK_M
q_tiles = N_CTX // BLOCK_M
meta_base = ((b * H + h) * q_tiles + q_blk)
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
# ----- base pointers -----
# Note: when q and kv have different sequence lengths, their per-(batch,head)
# strides differ, so we must compute separate base offsets.
q_off = (b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh)
k_off = (b.to(tl.int64) * stride_kz + h.to(tl.int64) * stride_kh)
v_off = (b.to(tl.int64) * stride_vz + h.to(tl.int64) * stride_vh)
o_off = (b.to(tl.int64) * stride_oz + h.to(tl.int64) * stride_oh)
qvk_off = (b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh)
Q_ptr = tl.make_block_ptr(base=Q + q_off,
shape=(N_CTX_Q, HEAD_DIM),
Q_ptr = tl.make_block_ptr(base=Q + qvk_off,
shape=(N_CTX, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
order=(1, 0))
K_base = tl.make_block_ptr(base=K + k_off,
shape=(HEAD_DIM, N_CTX_KV),
K_base = tl.make_block_ptr(base=K + qvk_off,
shape=(HEAD_DIM, N_CTX),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N),
@@ -106,15 +100,15 @@ def _attn_fwd_sparse(
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1,
0)
V_base = tl.make_block_ptr(base=V + v_off,
shape=(N_CTX_KV, HEAD_DIM),
V_base = tl.make_block_ptr(base=V + qvk_off,
shape=(N_CTX, HEAD_DIM),
strides=(stride_vk, stride_vn),
offsets=(0, 0),
block_shape=(BLOCK_N, HEAD_DIM),
order=v_order)
O_ptr = tl.make_block_ptr(base=Out + o_off,
shape=(N_CTX_Q, HEAD_DIM),
O_ptr = tl.make_block_ptr(base=Out + qvk_off,
shape=(N_CTX, HEAD_DIM),
strides=(stride_om, stride_on),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM),
@@ -156,7 +150,7 @@ def _attn_fwd_sparse(
# ----- epilogue -----
m_i += tl.math.log2(l_i)
acc = acc / l_i[:, None]
tl.store(M + off_hz * N_CTX_Q + offs_m, m_i)
tl.store(M + off_hz * N_CTX + offs_m, m_i)
tl.store(O_ptr, acc.to(Out.type.element_ty))
@@ -207,7 +201,7 @@ def _attn_bwd_dkdv(
stride_tok,
stride_d, #
H,
N_CTX_KV,
N_CTX,
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
HEAD_DIM: tl.constexpr, #
@@ -227,8 +221,8 @@ def _attn_bwd_dkdv(
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
kv_tiles = N_CTX_KV // BLOCK_N1
meta_base = ((b * H + h) * kv_tiles + kv_blk)
q_tiles = N_CTX // BLOCK_N1
meta_base = ((b * H + h) * q_tiles + kv_blk)
q_blocks = tl.load(k2q_num + meta_base) # int32
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
@@ -308,21 +302,16 @@ def _attn_bwd_dq(
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
block_size = tl.load(variable_block_sizes + q_blk)
for blk_idx in range(kv_blocks * 2):
kv_idx = tl.load(kv_ptr + blk_idx // 2).to(tl.int32)
# variable_block_sizes is defined per KV block (tile). Mask must therefore
# use kv_idx (not q_blk). Also, because we split each 64-token block into
# two 32-token halves, the mask must account for the half-block offset.
block_size = tl.load(variable_block_sizes + kv_idx).to(tl.int32)
half = (blk_idx % 2).to(tl.int32)
block_sparse_offset = (kv_idx * 2 + half) * step_n * stride_tok
block_sparse_offset = (tl.load(kv_ptr + blk_idx // 2).to(tl.int32) * 2 +
blk_idx % 2) * step_n * stride_tok
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT)
p = tl.math.exp2(qk - m)
offs_in_block = half * step_n + tl.arange(0, BLOCK_N2)
mask = offs_in_block < block_size
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
p = tl.where(mask[None, :], p, 0.0)
# Compute dP and dS.
dp = tl.dot(do, vT).to(tl.float32)
@@ -478,235 +467,19 @@ def _attn_bwd(
tl.store(dq_ptrs, dq)
@triton.jit
def _attn_bwd_dkdv_kernel(
Q,
K,
V,
sm_scale, #
DO, #
DK,
DV, #
M,
D,
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
# shared token/dim strides (assumed contiguous along token and dim)
stride_tok,
stride_d, #
# batch/head strides (may differ between Q and KV)
stride_qz,
stride_qh,
stride_kz,
stride_kh,
stride_vz,
stride_vh,
stride_doz,
stride_doh,
stride_dkz,
stride_dkh,
stride_dvz,
stride_dvh,
H,
N_CTX_Q,
N_CTX_KV,
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
HEAD_DIM: tl.constexpr):
"""
Backward kernel that computes dK and dV for each KV block (64 tokens).
Grid:
pid0: kv_blk in [0, N_CTX_KV/BLOCK_N1)
pid2: fused (batch, head) in [0, B*H)
"""
bhid = tl.program_id(2)
b = bhid // H
h = bhid % H
kv_blk = tl.program_id(0)
q_adj = (b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh)
kv_adj_k = (b.to(tl.int64) * stride_kz + h.to(tl.int64) * stride_kh)
kv_adj_v = (b.to(tl.int64) * stride_vz + h.to(tl.int64) * stride_vh)
do_adj = (b.to(tl.int64) * stride_doz + h.to(tl.int64) * stride_doh)
dk_adj = (b.to(tl.int64) * stride_dkz + h.to(tl.int64) * stride_dkh)
dv_adj = (b.to(tl.int64) * stride_dvz + h.to(tl.int64) * stride_dvh)
Q = Q + q_adj
K = K + kv_adj_k
V = V + kv_adj_v
DO = DO + do_adj
DK = DK + dk_adj
DV = DV + dv_adj
# M and D (delta) are always sized by Q length.
M = M + (bhid * N_CTX_Q).to(tl.int64)
D = D + (bhid * N_CTX_Q).to(tl.int64)
offs_k = tl.arange(0, HEAD_DIM)
start_n = kv_blk * BLOCK_N1
offs_n = start_n + tl.arange(0, BLOCK_N1)
# load K and V: they stay in SRAM throughout the inner loop.
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
dv_acc = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
dk_acc = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
num_steps = N_CTX_Q // BLOCK_M1
dk_acc, dv_acc = _attn_bwd_dkdv(
dk_acc,
dv_acc,
Q,
k,
v,
sm_scale,
DO,
M,
D,
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
stride_tok,
stride_d,
H,
N_CTX_KV,
BLOCK_M1=BLOCK_M1,
BLOCK_N1=BLOCK_N1,
HEAD_DIM=HEAD_DIM,
start_n=start_n,
start_m=0,
num_steps=num_steps,
)
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dv_ptrs, dv_acc)
dk_acc *= sm_scale
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dk_ptrs, dk_acc)
@triton.jit
def _attn_bwd_dq_kernel(
Q,
K,
V,
DO, #
DQ,
M,
D,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
# shared token/dim strides (assumed contiguous along token and dim)
stride_tok,
stride_d, #
# batch/head strides (may differ between Q and KV)
stride_qz,
stride_qh,
stride_kz,
stride_kh,
stride_vz,
stride_vh,
stride_doz,
stride_doh,
stride_dqz,
stride_dqh,
H,
N_CTX_Q,
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr):
"""
Backward kernel that computes dQ for each Q block (64 tokens).
Grid:
pid0: q_blk in [0, N_CTX_Q/BLOCK_M2)
pid2: fused (batch, head) in [0, B*H)
"""
LN2 = 0.6931471824645996 # = ln(2)
bhid = tl.program_id(2)
b = bhid // H
h = bhid % H
q_blk = tl.program_id(0)
q_adj = (b.to(tl.int64) * stride_qz + h.to(tl.int64) * stride_qh)
kv_adj_k = (b.to(tl.int64) * stride_kz + h.to(tl.int64) * stride_kh)
kv_adj_v = (b.to(tl.int64) * stride_vz + h.to(tl.int64) * stride_vh)
do_adj = (b.to(tl.int64) * stride_doz + h.to(tl.int64) * stride_doh)
dq_adj = (b.to(tl.int64) * stride_dqz + h.to(tl.int64) * stride_dqh)
Q = Q + q_adj
K = K + kv_adj_k
V = V + kv_adj_v
DO = DO + do_adj
DQ = DQ + dq_adj
M = M + (bhid * N_CTX_Q).to(tl.int64)
D = D + (bhid * N_CTX_Q).to(tl.int64)
offs_k = tl.arange(0, HEAD_DIM)
start_m = q_blk * BLOCK_M2
offs_m = start_m + tl.arange(0, BLOCK_M2)
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
m = tl.load(M + offs_m)[:, None]
dq_acc = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
num_steps = 0 # unused in _attn_bwd_dq
dq_acc = _attn_bwd_dq(
dq_acc,
q,
K,
V,
do,
m,
D,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
stride_tok,
stride_d,
H,
N_CTX_Q,
BLOCK_M2=BLOCK_M2,
BLOCK_N2=BLOCK_N2,
HEAD_DIM=HEAD_DIM,
start_m=start_m,
start_n=0,
num_steps=num_steps,
)
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq_acc *= LN2
tl.store(dq_ptrs, dq_acc)
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num,
variable_block_sizes):
B, H, Tq, D = q.shape
Tkv = k.shape[2]
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
max_kv_blks = q2k_index.shape[-1]
assert Tq % 64 == 0, f"q length must be a multiple of 64, but got {Tq}"
assert Tkv % 64 == 0, f"kv length must be a multiple of 64, but got {Tkv}"
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
assert q2k_num.shape[
-1] == Tq // 64, f"shape mismatch, Tq // 64 = {Tq // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
assert variable_block_sizes.numel() == Tkv // 64, (
f"shape mismatch, variable_block_sizes must have length {Tkv // 64}, "
f"got {variable_block_sizes.numel()}"
)
-1] == T // 64, f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
o = torch.empty_like(q)
M = torch.empty((B, H, Tq), dtype=torch.float32, device=q.device)
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
grid = lambda _: (triton.cdiv(Tq, 64), B * H, 1)
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
_attn_fwd_sparse[grid](q,
k,
v,
@@ -735,8 +508,7 @@ def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num,
o.stride(3),
B,
H,
Tq,
Tkv,
T,
HEAD_DIM=D,
STAGE=3)
@@ -746,21 +518,21 @@ def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num,
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num,
k2q_index, k2q_num, variable_block_sizes):
assert do.is_contiguous()
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
B, H, Tq, D = q.shape
Tkv = k.shape[2]
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
BATCH, N_HEAD = q.shape[:2]
BATCH, N_HEAD, N_CTX = q.shape[:3]
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
arg_k = k
arg_k = arg_k * (sm_scale * RCP_LN2)
PRE_BLOCK = 64
assert Tq % PRE_BLOCK == 0
pre_grid = (Tq // PRE_BLOCK, BATCH * N_HEAD)
assert N_CTX % PRE_BLOCK == 0
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
delta = torch.empty_like(M)
_attn_bwd_preprocess[pre_grid](
o,
@@ -768,7 +540,7 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num,
delta, #
BATCH,
N_HEAD,
Tq, #
N_CTX, #
BLOCK_M=PRE_BLOCK,
HEAD_DIM=D #
)
@@ -776,75 +548,36 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num,
max_q_blks = k2q_index.shape[-1]
max_kv_blks = q2k_index.shape[-1]
# dK/dV kernel: grid over KV blocks
grid_kv = (Tkv // BLOCK_N1, 1, BATCH * N_HEAD)
_attn_bwd_dkdv_kernel[grid_kv](
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
_attn_bwd[grid](
q,
arg_k,
v,
sm_scale,
do,
dq,
dk,
dv,
dv, #
M,
delta,
delta, #
q2k_index,
q2k_num,
max_kv_blks,
k2q_index,
k2q_num,
max_q_blks,
variable_block_sizes,
q.stride(2),
q.stride(3),
q.stride(0),
q.stride(1),
arg_k.stride(0),
arg_k.stride(1),
v.stride(0),
v.stride(1),
do.stride(0),
do.stride(1),
dk.stride(0),
dk.stride(1),
dv.stride(0),
dv.stride(1),
q.stride(2),
q.stride(3), #
N_HEAD,
Tq,
Tkv,
N_CTX, #
BLOCK_M1=BLOCK_M1,
BLOCK_N1=BLOCK_N1,
HEAD_DIM=D,
)
# dQ kernel: grid over Q blocks
grid_q = (Tq // BLOCK_M2, 1, BATCH * N_HEAD)
_attn_bwd_dq_kernel[grid_q](
q,
arg_k,
v,
do,
dq,
M,
delta,
q2k_index,
q2k_num,
max_kv_blks,
variable_block_sizes,
q.stride(2),
q.stride(3),
q.stride(0),
q.stride(1),
arg_k.stride(0),
arg_k.stride(1),
v.stride(0),
v.stride(1),
do.stride(0),
do.stride(1),
dq.stride(0),
dq.stride(1),
N_HEAD,
Tq,
BLOCK_N1=BLOCK_N1, #
BLOCK_M2=BLOCK_M2,
BLOCK_N2=BLOCK_N2,
HEAD_DIM=D,
BLOCK_N2=BLOCK_N2, #
HEAD_DIM=D #
)
return dq, dk, dv
@@ -1 +1 @@
__version__ = "0.2.6"
__version__ = "0.2.5"
-5
View File
@@ -86,7 +86,6 @@ class PreprocessConfig:
# Model configuration
training_cfg_rate: float = 0.0
with_audio: bool = False
# framework configuration
seed: int = 42
@@ -191,10 +190,6 @@ class PreprocessConfig:
type=float,
default=PreprocessConfig.training_cfg_rate,
help="Training CFG rate")
preprocess_args.add_argument(f"--{prefix_with_dot}with-audio",
action=StoreBoolean,
default=PreprocessConfig.with_audio,
help="Whether to extract and encode audio")
preprocess_args.add_argument(f"--{prefix_with_dot}seed",
type=int,
default=PreprocessConfig.seed,
+3 -5
View File
@@ -1,6 +1,5 @@
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
@@ -10,8 +9,7 @@ from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig",
"WanVideoConfig", "StepVideoConfig", "CosmosVideoConfig",
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig",
"HYWorldConfig"
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig"
]
@@ -1,163 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Configuration for HunyuanGameCraft transformer model.
HunyuanGameCraft extends HunyuanVideo with:
1. CameraNet for camera/action conditioning
2. 33 input channels (16 latent + 16 gt_latent + 1 mask)
3. Mask-based conditioning for autoregressive generation
"""
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double" in n and str.isdigit(n.split(".")[-1])
def is_single_block(n: str, m) -> bool:
return "single" in n and str.isdigit(n.split(".")[-1])
def is_refiner_block(n: str, m) -> bool:
return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
return n.split(".")[-1] == "txt_in"
def is_camera_net(n: str, m) -> bool:
return "camera_net" in n
@dataclass
class HunyuanGameCraftArchConfig(DiTArchConfig):
"""Architecture config for HunyuanGameCraft transformer."""
# Version field for compatibility with saved config.json
_fastvideo_version: str = "0.1.0"
# Camera net flag (for config.json compatibility)
camera_net: bool = True
_fsdp_shard_conditions: list = field(
default_factory=lambda:
[is_double_block, is_single_block, is_refiner_block, is_camera_net])
_compile_conditions: list = field(
default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
# Parameter names mapping from official checkpoint to FastVideo naming
# GameCraft weights are already close to FastVideo format with minor adjustments
param_names_mapping: dict = field(
default_factory=lambda: {
# MLP naming: fc1 -> fc_in, fc2 -> fc_out
r"^(.*)\.img_mlp\.fc1\.(.*)$":
r"\1.img_mlp.fc_in.\2",
r"^(.*)\.img_mlp\.fc2\.(.*)$":
r"\1.img_mlp.fc_out.\2",
r"^(.*)\.txt_mlp\.fc1\.(.*)$":
r"\1.txt_mlp.fc_in.\2",
r"^(.*)\.txt_mlp\.fc2\.(.*)$":
r"\1.txt_mlp.fc_out.\2",
# Single block MLP naming
r"^single_blocks\.(\d+)\.mlp\.fc1\.(.*)$":
r"single_blocks.\1.mlp.fc_in.\2",
r"^single_blocks\.(\d+)\.mlp\.fc2\.(.*)$":
r"single_blocks.\1.mlp.fc_out.\2",
# Token refiner naming
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.(.*)$":
r"txt_in.refiner_blocks.\1.\2",
# Vector in naming
r"^vector_in\.in_layer\.(.*)$":
r"vector_in.fc_in.\1",
r"^vector_in\.out_layer\.(.*)$":
r"vector_in.fc_out.\1",
# Time embedder naming
r"^time_in\.mlp\.0\.(.*)$":
r"time_in.mlp.fc_in.\1",
r"^time_in\.mlp\.2\.(.*)$":
r"time_in.mlp.fc_out.\1",
# Guidance embedder naming (if present)
r"^guidance_in\.mlp\.0\.(.*)$":
r"guidance_in.mlp.fc_in.\1",
r"^guidance_in\.mlp\.2\.(.*)$":
r"guidance_in.mlp.fc_out.\1",
# Final layer adaLN modulation
r"^final_layer\.adaLN_modulation\.1\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
# Refiner block MLP naming
r"^txt_in\.refiner_blocks\.(\d+)\.mlp\.fc1\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^txt_in\.refiner_blocks\.(\d+)\.mlp\.fc2\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
# Camera net weights are already correctly named
})
# Reverse mapping for saving checkpoints
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Model architecture parameters
# patch_size can be int or tuple - if tuple, it's [T, H, W]
patch_size: int | tuple[int, int, int] = 2
patch_size_t: int = 1
in_channels: int = 33 # 16 latent + 16 gt_latent + 1 mask
out_channels: int = 16
num_attention_heads: int = 24
attention_head_dim: int = 128
mlp_ratio: float = 4.0
num_layers: int = 20 # Double stream blocks
num_single_layers: int = 40 # Single stream blocks
num_refiner_layers: int = 2
rope_axes_dim: tuple[int, int, int] = (16, 56, 56)
guidance_embeds: bool = False # GameCraft doesn't use guidance
dtype: torch.dtype | None = None
text_embed_dim: int = 4096 # LLaMA-3 hidden size
pooled_projection_dim: int = 768 # CLIP pooled output dim
rope_theta: int = 256
qk_norm: str = "rms_norm"
# Camera net parameters
camera_in_channels: int = 6 # Plücker coordinates
camera_downscale_coef: int = 8
camera_out_channels: int = 16
# Layers to exclude from LoRA
exclude_lora_layers: list[str] = field(
default_factory=lambda:
["img_in", "txt_in", "time_in", "vector_in", "camera_net"])
def __post_init__(self):
super().__post_init__()
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
self.num_channels_latents: int = 16 # Output is 16 channels
# Convert patch_size list to tuple if needed (from JSON deserialization)
if isinstance(self.patch_size, list):
self.patch_size = tuple(self.patch_size)
# Convert rope_axes_dim list to tuple if needed
if isinstance(self.rope_axes_dim, list):
self.rope_axes_dim = tuple(self.rope_axes_dim)
@dataclass
class HunyuanGameCraftConfig(DiTConfig):
"""Full config for HunyuanGameCraft model."""
arch_config: DiTArchConfig = field(
default_factory=HunyuanGameCraftArchConfig)
prefix: str = "HunyuanGameCraft"
@@ -1,110 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class LingBotWorldArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1",
r"^patch_embedding_wancamctrl\.(.*)$":
r"patch_embedding_wancamctrl.proj.\1",
r"^c2ws_hidden_states_layer1\.(.*)$": r"c2ws_mlp.fc_in.\1",
r"^c2ws_hidden_states_layer2\.(.*)$": r"c2ws_mlp.fc_out.\1",
r"^text_embedding\.0\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^text_embedding\.2\.(.*)$":
r"condition_embedder.text_embedder.fc_out.\1",
r"^time_embedding\.0\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^time_embedding\.2\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^time_projection\.1\.(.*)$":
r"condition_embedder.time_modulation.linear.\1",
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.scale_shift_table",
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.norm3\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$":
r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$":
r"blocks.\1.attn2.norm_q.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$":
r"blocks.\1.attn2.norm_k.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.cam_injector_layer1\.(.*)$":
r"blocks.\1.cam_conditioner.cam_injector.fc_in.\2",
r"^blocks\.(\d+)\.cam_injector_layer2\.(.*)$":
r"blocks.\1.cam_conditioner.cam_injector.fc_out.\2",
r"^blocks\.(\d+)\.cam_scale_layer\.(.*)$":
r"blocks.\1.cam_conditioner.cam_scale_layer.\2",
r"^blocks\.(\d+)\.cam_shift_layer\.(.*)$":
r"blocks.\1.cam_conditioner.cam_shift_layer.\2",
r"^head\.modulation$": r"scale_shift_table",
r"^head\.head\.(.*)$": r"proj_out.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Some LoRA adapters use the original official layer names instead of hf layer names,
# so apply this before the param_names_mapping
lora_param_names_mapping: dict = field(default_factory=lambda: {})
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len: int = 512
num_attention_heads: int = 40
attention_head_dim: int = 128
in_channels: int = 16
out_channels: int = 16
text_dim: int = 4096
freq_dim: int = 256
ffn_dim: int = 13824
num_layers: int = 40
cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6
image_dim: int | None = None
added_kv_proj_dim: int | None = None
rope_max_seq_len: int = 1024
pos_embed_seq_len: int | None = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
# Wan MoE
boundary_ratio: float | None = None
# Causal Wan
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
num_frames_per_block: int = 3
sliding_window_num_frames: int = 21
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.out_channels
@dataclass
class LingBotWorldVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=LingBotWorldArchConfig)
prefix: str = "Wan"
+2 -4
View File
@@ -7,12 +7,10 @@ from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
import re
def is_ltx2_blocks(name: str, _module) -> bool:
res = re.search(r"(?:^|\.)transformer_blocks\.\d+$", name) is not None
return res
"""FSDP shard condition for LTX-2 transformer blocks."""
return "transformer_blocks" in name
@dataclass
@@ -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])
-31
View File
@@ -1,31 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
@dataclass
class SD3Transformer2DArchConfig(DiTArchConfig):
# Diffusers SD3Transformer2DModel config fields.
sample_size: int = 128
patch_size: int = 2
num_layers: int = 24
attention_head_dim: int = 64
joint_attention_dim: int = 4096
caption_projection_dim: int = 1536
pooled_projection_dim: int = 2048
pos_embed_max_size: int = 384
dual_attention_layers: list[int] = field(
default_factory=lambda: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12])
qk_norm: str = "rms_norm"
in_channels: int = 16
out_channels: int = 16
num_attention_heads = 24
@dataclass
class SD3DiTConfig(DiTConfig):
arch_config: DiTArchConfig = field(
default_factory=SD3Transformer2DArchConfig)
prefix: str = "sd3"
@@ -1,6 +1,5 @@
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
from fastvideo.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
@@ -8,7 +7,6 @@ from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
__all__ = [
"GameCraftVAEConfig",
"HunyuanVAEConfig",
"WanVAEConfig",
"StepVideoVAEConfig",
@@ -1,39 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class AutoencoderKLArchConfig(VAEArchConfig):
_name_or_path: str = ""
act_fn: str = "silu"
block_out_channels: tuple[int, ...] | list[int] = field(
default_factory=list)
down_block_types: tuple[str, ...] | list[str] = field(default_factory=list)
up_block_types: tuple[str, ...] | list[str] = field(default_factory=list)
force_upcast: bool = True
in_channels: int = 3
latent_channels: int = 4
latents_mean: tuple[float, ...] | list[float] | None = None
latents_std: tuple[float, ...] | list[float] | None = None
layers_per_block: int = 1
mid_block_add_attention: bool = True
norm_num_groups: int = 32
out_channels: int = 3
sample_size: int = 32
scaling_factor: float | torch.Tensor = 0.18215
shift_factor: float | None = None
use_post_quant_conv: bool = True
use_quant_conv: bool = True
temporal_compression_ratio: int = 1
spatial_compression_ratio: int = 8
@dataclass
class AutoencoderKLVAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=AutoencoderKLArchConfig)
@@ -1,50 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
GameCraft VAE config - matches official config.json from Hunyuan-GameCraft-1.0.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class GameCraftVAEArchConfig(VAEArchConfig):
"""Architecture config matching official AutoencoderKLCausal3D config.json."""
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 16
down_block_types: tuple[str, ...] = (
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
)
up_block_types: tuple[str, ...] = (
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
)
block_out_channels: tuple[int, ...] = (128, 256, 512, 512)
layers_per_block: int = 2
act_fn: str = "silu"
norm_num_groups: int = 32
scaling_factor: float = 0.476986
spatial_compression_ratio: int = 8
temporal_compression_ratio: int = 4
time_compression_ratio: int = 4 # alias for DecoderCausal3D
mid_block_add_attention: bool = True
mid_block_causal_attn: bool = True
sample_size: int = 256 # from config.json
sample_tsize: int = 64 # from config.json
def __post_init__(self):
self.spatial_compression_ratio = 2**(len(self.block_out_channels) - 1)
@dataclass
class GameCraftVAEConfig(VAEConfig):
"""Full config for GameCraft VAE."""
arch_config: VAEArchConfig = field(default_factory=GameCraftVAEArchConfig)
+6 -7
View File
@@ -4,7 +4,6 @@ from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.registry import get_pipeline_config_cls_from_name
@@ -14,10 +13,10 @@ from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
WanT2V480PConfig, WanT2V720PConfig)
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig",
"PipelineConfig", "Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"SelfForcingWanT2V480PConfig", "CosmosConfig", "Cosmos25Config",
"LTX2T2VConfig", "HYWorldConfig", "get_pipeline_config_cls_from_name"
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig", "HYWorldConfig",
"get_pipeline_config_cls_from_name"
]
@@ -1,122 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Pipeline configuration for HunyuanGameCraft.
HunyuanGameCraft extends HunyuanVideo with:
1. CameraNet for camera/action conditioning (Plücker coordinates)
2. Mask-based conditioning for autoregressive generation
3. 33 input channels (16 latent + 16 gt_latent + 1 mask)
Text encoders are the same as HunyuanVideo:
- LLaVA-LLaMA-3-8B for primary text encoding (4096 dim)
- CLIP ViT-L/14 for secondary pooled embeddings (768 dim)
"""
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import TypedDict
import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits import HunyuanGameCraftConfig
from fastvideo.configs.models.encoders import (
BaseEncoderOutput,
CLIPTextConfig,
LlamaConfig,
)
from fastvideo.configs.models.vaes import GameCraftVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
# GameCraft uses the same prompt template as HunyuanVideo
PROMPT_TEMPLATE_ENCODE_VIDEO = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
class PromptTemplate(TypedDict):
template: str
crop_start: int
prompt_template_video: PromptTemplate = {
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
"crop_start": 95,
}
def llama_preprocess_text(prompt: str) -> str:
"""Apply prompt template for LLaMA encoder."""
return prompt_template_video["template"].format(prompt)
def llama_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
"""Extract hidden states from LLaMA output, skipping instruction tokens."""
hidden_state_skip_layer = 2
assert outputs.hidden_states is not None
hidden_states: tuple[torch.Tensor, ...] = outputs.hidden_states
last_hidden_state: torch.Tensor = hidden_states[-(hidden_state_skip_layer +
1)]
crop_start = prompt_template_video.get("crop_start", -1)
last_hidden_state = last_hidden_state[:, crop_start:]
return last_hidden_state
def clip_preprocess_text(prompt: str) -> str:
"""No preprocessing for CLIP encoder."""
return prompt
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
"""Extract pooled output from CLIP encoder."""
pooler_output: torch.Tensor = outputs.pooler_output
return pooler_output
@dataclass
class HunyuanGameCraftPipelineConfig(PipelineConfig):
"""Configuration for HunyuanGameCraft pipeline.
Inherits text encoding from HunyuanVideo but uses:
- GameCraft DiT with CameraNet
- Same VAE (HunyuanVAE)
- Same text encoders (LLaMA + CLIP)
"""
# DiT config - uses GameCraft config (33 input channels)
dit_config: DiTConfig = field(default_factory=HunyuanGameCraftConfig)
# VAE config - GameCraft VAE (mid_block_causal_attn=True, etc.)
vae_config: VAEConfig = field(default_factory=GameCraftVAEConfig)
# Denoising parameters
# Official GameCraft does NOT use embedded guidance (passes guidance=None)
# It uses standard CFG with guidance_scale=6.0 instead
embedded_cfg_scale = None
flow_shift: int = 5 # Official GameCraft uses flow_shift=5.0
# Text encoding stage - same as HunyuanVideo
# Uses LLaMA-3-8B (via LLaVA) + CLIP
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LlamaConfig(), CLIPTextConfig()))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (llama_preprocess_text, clip_preprocess_text))
postprocess_text_funcs: tuple[
Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(llama_postprocess_text, clip_postprocess_text))
# Precision for each component
dit_precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp16"))
def __post_init__(self):
# VAE only needs decoder for inference
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@@ -1,13 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models import DiTConfig
from fastvideo.configs.pipelines.wan import Wan2_2_I2V_A14B_Config
from fastvideo.configs.models.dits.lingbotworld import LingBotWorldVideoConfig
@dataclass
class LingBotWorldI2V480PConfig(Wan2_2_I2V_A14B_Config):
dit_config: DiTConfig = field(default_factory=LingBotWorldVideoConfig)
flow_shift: float | None = 10.0
boundary_ratio: float | None = 0.947
-65
View File
@@ -1,65 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models import EncoderConfig
from fastvideo.configs.models.encoders import (
BaseEncoderOutput,
CLIPTextConfig,
T5Config,
)
from fastvideo.configs.models.dits.sd3 import SD3DiTConfig
from fastvideo.configs.models.vaes.autoencoder_kl import AutoencoderKLVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
def _sd35_text_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
assert outputs.last_hidden_state is not None
return outputs.last_hidden_state
@dataclass
class SD35Config(PipelineConfig):
scheduler_arch: str = "FlowMatchEulerDiscreteScheduler"
transformer_arch: str = "SD3Transformer2DModel"
vae_arch: str = "AutoencoderKL"
text_encoder_archs: tuple[str, ...] = (
"CLIPTextModelWithProjection",
"CLIPTextModelWithProjection",
"T5EncoderModel",
)
tokenizer_archs: tuple[str, ...] = (
"CLIPTokenizer",
"CLIPTokenizer",
"T5TokenizerFast",
)
dit_config: SD3DiTConfig = field(default_factory=SD3DiTConfig)
vae_config: AutoencoderKLVAEConfig = field(
default_factory=AutoencoderKLVAEConfig)
embedded_cfg_scale: float = 0.0
flow_shift: float | None = None
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda:
(CLIPTextConfig(), CLIPTextConfig(), T5Config()))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda:
(preprocess_text, preprocess_text, preprocess_text))
postprocess_text_funcs: tuple[
Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(_sd35_text_postprocess, _sd35_text_postprocess,
_sd35_text_postprocess))
dit_precision: str = "bf16"
vae_precision: str = "fp32"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp32", "fp32", "bf16"))
+1 -11
View File
@@ -1,13 +1,3 @@
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.configs.sample.hunyuangamecraft import (
HunyuanGameCraftSamplingParam,
HunyuanGameCraft65FrameSamplingParam,
HunyuanGameCraft129FrameSamplingParam,
)
__all__ = [
"SamplingParam",
"HunyuanGameCraftSamplingParam",
"HunyuanGameCraft65FrameSamplingParam",
"HunyuanGameCraft129FrameSamplingParam",
]
__all__ = ["SamplingParam"]
-3
View File
@@ -31,9 +31,6 @@ class SamplingParam:
# Camera control inputs (HYWorld)
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
# Camera control inputs (LingBotWorld)
c2ws_plucker_emb: Any | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
# Refine inputs (LongCat 480p->720p upscaling)
# Path-based refine (load stage1 video from disk, e.g. MP4)
refine_from: str | None = None # Path to stage1 video (480p output from distill)
@@ -1,104 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Sampling parameters for HunyuanGameCraft video generation.
GameCraft generates game-like videos with camera/action control.
Default parameters are based on the official implementation.
"""
from dataclasses import dataclass, field
from typing import Any
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.configs.sample.teacache import TeaCacheParams
@dataclass
class HunyuanGameCraftSamplingParam(SamplingParam):
"""Sampling parameters for HunyuanGameCraft video generation.
Supports camera/action conditioning via:
- camera_trajectory: Plücker coordinates for camera motion
- action_list: List of actions (e.g., ["forward", "left", "right"])
- action_speed_list: Speed multipliers for each action
Default resolution is 704x1280 (same as HunyuanVideo).
Default frame count is 33 video frames -> 9 latent frames.
"""
# Number of denoising steps
num_inference_steps: int = 50
# Video dimensions
# 33 video frames -> 9 latent frames (4x temporal compression)
num_frames: int = 33
height: int = 704
width: int = 1280
fps: int = 24
# Guidance scale - official GameCraft uses CFG with guidance_scale=6.0
guidance_scale: float = 6.0
# Negative prompt for CFG (empty string = unconditional)
negative_prompt: str = ""
# Camera/Action conditioning
# Camera states as Plücker coordinates [B, T_video, 6, H, W]
camera_states: Any | None = None
# Camera trajectory file/identifier (alternative to camera_states)
camera_trajectory: str | None = None
# Action list for camera motion (e.g., ["forward", "left"])
action_list: list[str] | None = None
# Speed multipliers for each action
action_speed_list: list[float] | None = None
# History frame conditioning (for autoregressive generation)
# Ground truth latents for conditioning [B, 16, T, H, W]
gt_latents: Any | None = None
# Mask for conditioning (1=use gt, 0=generate) [B, 1, T, H, W]
conditioning_mask: Any | None = None
# Number of conditioning frames (for autoregressive) - maps to num_cond_frames
num_cond_frames: int = 0
# TeaCache parameters (if enabled)
teacache_params: TeaCacheParams = field(
default_factory=lambda: TeaCacheParams(
teacache_thresh=0.15,
coefficients=[
7.33226126e+02, -4.01131952e+02, 6.75869174e+01,
-3.14987800e+00, 9.61237896e-02
],
))
def __post_init__(self) -> None:
super().__post_init__()
# Validate action lists
if (self.action_list is not None and self.action_speed_list is not None
and len(self.action_list) != len(self.action_speed_list)):
raise ValueError(
f"action_list length ({len(self.action_list)}) must match "
f"action_speed_list length ({len(self.action_speed_list)})")
@dataclass
class HunyuanGameCraft65FrameSamplingParam(HunyuanGameCraftSamplingParam):
"""Sampling parameters for 65-frame GameCraft generation.
65 video frames -> 17 latent frames (with first frame as key frame).
This is useful for longer video generation.
"""
num_frames: int = 65
@dataclass
class HunyuanGameCraft129FrameSamplingParam(HunyuanGameCraftSamplingParam):
"""Sampling parameters for 129-frame GameCraft generation.
129 video frames -> 33 latent frames.
This is the maximum supported by the official implementation.
"""
num_frames: int = 129
-21
View File
@@ -1,21 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.wan import Wan2_2_I2V_A14B_SamplingParam
@dataclass
class LingBotWorld_SamplingParam(Wan2_2_I2V_A14B_SamplingParam):
guidance_scale: float = 5.0 # high_noise
guidance_scale_2: float = 5.0 # low_noise
num_inference_steps: int = 70
boundary_ratio: float | None = 0.947
negative_prompt: str | None = (
"画面突变,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,"
"最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,"
"畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走,"
"镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,"
"PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,"
"皮肤,肢体,面部特征,汽车,电线")
fps: int = 16
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
+16 -48
View File
@@ -5,51 +5,10 @@ from fastvideo.configs.sample.base import SamplingParam
@dataclass
class LTX2BaseSamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 base one-stage T2V.
Values follow the official LTX-2 one-stage defaults.
Multi-modal CFG params are read by ``LTX2DenoisingStage``.
class LTX2SamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled T2V.
"""
seed: int = 10
num_frames: int = 121
height: int = 512
width: int = 768
fps: int = 24
num_inference_steps: int = 40
guidance_scale: float = 3.0
# Copied/following official LTX-2 DEFAULT_NEGATIVE_PROMPT.
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.")
# Official LTX-2 multi-modal CFG defaults.
ltx2_cfg_scale_video: float = 3.0
ltx2_cfg_scale_audio: float = 7.0
ltx2_modality_scale_video: float = 3.0
ltx2_modality_scale_audio: float = 3.0
ltx2_rescale_scale: float = 0.7
@dataclass
class LTX2DistilledSamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled one-stage T2V."""
seed: int = 10
num_frames: int = 121
height: int = 1024
@@ -57,9 +16,18 @@ class LTX2DistilledSamplingParam(SamplingParam):
fps: int = 24
num_inference_steps: int = 8
guidance_scale: float = 1.0
# No default negative_prompt for distilled models
negative_prompt: str = ""
# Backward compatibility alias.
LTX2SamplingParam = LTX2DistilledSamplingParam
# 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."
)
-25
View File
@@ -1,25 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class SD35SamplingParam(SamplingParam):
prompt: str | None = "a photo of a cat"
negative_prompt: str = ""
num_videos_per_prompt: int = 1
seed: int = 0
num_frames: int = 1
height: int = 512
width: int = 512
fps: int = 1
num_inference_steps: int = 28
guidance_scale: float = 6.0
+2 -8
View File
@@ -4,8 +4,6 @@ from torchvision.transforms import Lambda
from fastvideo.dataset.parquet_dataset_map_style import (
build_parquet_map_style_dataloader)
from fastvideo.dataset.ltx2_precomputed_dataset import (
build_ltx2_precomputed_dataloader, LTX2PrecomputedDataset)
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset, TextDataset
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
@@ -48,10 +46,6 @@ def gettextdataset(args) -> TextDataset:
__all__ = [
"build_parquet_map_style_dataloader",
"build_ltx2_precomputed_dataloader",
"LTX2PrecomputedDataset",
"ValidationDataset",
"VideoCaptionMergedDataset",
"TextDataset",
"build_parquet_map_style_dataloader", "ValidationDataset",
"VideoCaptionMergedDataset", "TextDataset"
]
@@ -1,210 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Dataset utilities for loading LTX2 precomputed training artifacts.
#
# Usage:
# - Input root can be either `<data_root>/` or `<data_root>/.precomputed/`.
# - Required sources are `latents/` and `conditions/` with matching `.pt` files.
# - Optional source `audio_latents/` is loaded when provided in `data_sources`.
# - `build_ltx2_precomputed_dataloader(...)` is the intended entrypoint used by
# `fastvideo/training/ltx2_training_pipeline.py`.
from __future__ import annotations
from pathlib import Path
from typing import Any
import torch
from einops import rearrange
from torch.utils.data import Dataset
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.dataset.parquet_dataset_map_style import DP_SP_BatchSampler
from fastvideo.distributed import get_sp_world_size, get_world_rank, get_world_size
from fastvideo.logger import init_logger
logger = init_logger(__name__)
PRECOMPUTED_DIR_NAME = ".precomputed"
class LTX2PrecomputedDataset(Dataset):
"""Dataset for LTX-2 precomputed latents and conditions.
Expected directory structure (data_root):
.precomputed/
latents/*.pt
conditions/*.pt
audio_latents/*.pt (optional)
"""
def __init__(
self,
data_root: str,
data_sources: dict[str, str] | list[str] | None = None,
) -> None:
super().__init__()
self.data_root = self._setup_data_root(data_root)
self.data_sources = self._normalize_data_sources(data_sources)
self.source_paths = self._setup_source_paths()
self.sample_files = self._discover_samples()
self._validate_setup()
@staticmethod
def _setup_data_root(data_root: str) -> Path:
data_root_path = Path(data_root).expanduser().resolve()
if not data_root_path.exists():
raise FileNotFoundError(
f"Data root directory does not exist: {data_root_path}")
if (data_root_path / PRECOMPUTED_DIR_NAME).exists():
data_root_path = data_root_path / PRECOMPUTED_DIR_NAME
return data_root_path
@staticmethod
def _normalize_data_sources(
data_sources: dict[str, str] | list[str] | None,
) -> dict[str, str]:
if data_sources is None:
return {"latents": "latents", "conditions": "conditions"}
if isinstance(data_sources, list):
return {source: source for source in data_sources}
if isinstance(data_sources, dict):
return data_sources.copy()
raise TypeError(
f"data_sources must be dict, list, or None, got {type(data_sources)}")
def _setup_source_paths(self) -> dict[str, Path]:
source_paths: dict[str, Path] = {}
for dir_name in self.data_sources:
source_path = self.data_root / dir_name
if not source_path.exists():
raise FileNotFoundError(
f"Required {dir_name} directory does not exist: {source_path}")
source_paths[dir_name] = source_path
return source_paths
def _discover_samples(self) -> dict[str, list[Path]]:
data_key = ("latents"
if "latents" in self.data_sources else next(iter(
self.data_sources.keys())))
data_path = self.source_paths[data_key]
data_files = list(data_path.glob("**/*.pt"))
if not data_files:
raise ValueError(f"No data files found in {data_path}")
sample_files = {output_key: [] for output_key in self.data_sources.values()}
for data_file in data_files:
rel_path = data_file.relative_to(data_path)
if self._all_source_files_exist(data_file, rel_path):
self._fill_sample_data_files(data_file, rel_path, sample_files)
return sample_files
def _all_source_files_exist(self, data_file: Path, rel_path: Path) -> bool:
for dir_name in self.data_sources:
expected_path = self._get_expected_file_path(dir_name, data_file,
rel_path)
if not expected_path.exists():
logger.warning(
"No matching %s file found for: %s (expected in: %s)",
dir_name,
data_file.name,
expected_path,
)
return False
return True
def _get_expected_file_path(self, dir_name: str, data_file: Path,
rel_path: Path) -> Path:
source_path = self.source_paths[dir_name]
if dir_name == "conditions" and data_file.name.startswith("latent_"):
return source_path / f"condition_{data_file.stem[7:]}.pt"
return source_path / rel_path
def _fill_sample_data_files(self, data_file: Path, rel_path: Path,
sample_files: dict[str, list[Path]]) -> None:
for dir_name, output_key in self.data_sources.items():
expected_path = self._get_expected_file_path(dir_name, data_file,
rel_path)
sample_files[output_key].append(
expected_path.relative_to(self.source_paths[dir_name]))
def _validate_setup(self) -> None:
if not self.sample_files:
raise ValueError(
"No valid samples found - all data sources must have matching files"
)
sample_counts = {
key: len(files)
for key, files in self.sample_files.items()
}
if len(set(sample_counts.values())) > 1:
raise ValueError(
f"Mismatched sample counts across sources: {sample_counts}")
def __len__(self) -> int:
first_key = next(iter(self.sample_files.keys()))
return len(self.sample_files[first_key])
def __getitem__(self, index: int) -> dict[str, torch.Tensor]:
result: dict[str, Any] = {}
for dir_name, output_key in self.data_sources.items():
source_path = self.source_paths[dir_name]
file_rel_path = self.sample_files[output_key][index]
file_path = source_path / file_rel_path
try:
data = torch.load(file_path, map_location="cpu", weights_only=True)
if "latent" in dir_name.lower():
data = self._normalize_video_latents(data)
result[output_key] = data
except Exception as e:
raise RuntimeError(
f"Failed to load {output_key} from {file_path}: {e}") from e
result["idx"] = index
return result
@staticmethod
def _normalize_video_latents(data: dict) -> dict:
latents = data["latents"]
if latents.dim() == 2:
num_frames = data["num_frames"]
height = data["height"]
width = data["width"]
latents = rearrange(
latents,
"(f h w) c -> c f h w",
f=num_frames,
h=height,
w=width,
)
data = data.copy()
data["latents"] = latents
return data
def build_ltx2_precomputed_dataloader(
path: str,
batch_size: int,
num_data_workers: int,
data_sources: dict[str, str] | list[str] | None = None,
drop_last: bool = True,
seed: int = 42,
) -> tuple[LTX2PrecomputedDataset, StatefulDataLoader]:
dataset = LTX2PrecomputedDataset(path, data_sources=data_sources)
sampler = DP_SP_BatchSampler(
batch_size=batch_size,
dataset_size=len(dataset),
num_sp_groups=get_world_size() // get_sp_world_size(),
sp_world_size=get_sp_world_size(),
global_rank=get_world_rank(),
drop_last=drop_last,
drop_first_row=False,
seed=seed,
)
loader = StatefulDataLoader(
dataset,
batch_sampler=sampler,
collate_fn=None,
num_workers=num_data_workers,
pin_memory=True,
persistent_workers=num_data_workers > 0,
)
return dataset, loader
+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"]
+25 -92
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.
@@ -223,82 +207,64 @@ class VideoGenerator:
sampling_param=sampling_param,
**kwargs)
def _is_image_workload(self) -> bool:
"""Return True when the workload produces a single image (t2i, i2i …)."""
args = getattr(self, "fastvideo_args", None)
if args is None:
return False
return args.workload_type.value.endswith("2i")
def _prepare_output_path(
self,
output_path: str,
prompt: str,
) -> str:
"""Build a unique, sanitized output file path.
"""Build a unique, sanitized .mp4 output file path.
The file extension is chosen automatically based on the workload type:
``.png`` for image workloads (``t2i``, ``i2i``, …) and ``.mp4`` for
video workloads.
- If ``output_path`` already carries the correct extension, treat it
as a file path.
- Otherwise, treat ``output_path`` as a directory and derive the
filename from the prompt.
- If `output_path` ends with .mp4 (case-insensitive), treat it as a file path.
- Otherwise, treat `output_path` as a directory and derive the filename
from the prompt.
- Invalid filename characters are removed; if the name changes, a
warning is logged.
- If the target path already exists, a numeric suffix is appended.
"""
target_ext = ".png" if self._is_image_workload() else ".mp4"
def _sanitize_filename_component(name: str) -> str:
# Remove characters invalid on common filesystems, strip spaces/dots
sanitized = re.sub(r'[\\/:*?"<>|]', '', name)
sanitized = sanitized.strip().strip('.')
sanitized = re.sub(r'\s+', ' ', sanitized)
return sanitized or "output"
return sanitized or "video"
base_path, extension = os.path.splitext(output_path)
extension_lower = extension.lower()
if extension_lower == target_ext:
if extension_lower == ".mp4":
output_dir = os.path.dirname(output_path)
base_name = os.path.basename(
base_path) # filename without extension
sanitized_base = _sanitize_filename_component(base_name)
if sanitized_base != base_name:
logger.warning(
"The output name '%s' contained invalid characters. "
"It has been renamed to '%s%s'",
"The video name '%s' contained invalid characters. It has been renamed to '%s.mp4'",
os.path.basename(output_path),
sanitized_base,
target_ext,
)
out_name = f"{sanitized_base}{target_ext}"
video_name = f"{sanitized_base}.mp4"
else:
# Treat as directory; inform if an unexpected extension was
# provided.
# Treat as directory; inform if an unexpected extension was provided.
if extension:
logger.info(
"Output path '%s' has extension '%s' which does not "
"match the target '%s'; treating it as a directory",
"Output path '%s' has non-mp4 extension '%s'; treating it as a directory and using a .mp4 filename derived from the prompt",
output_path,
extension,
target_ext,
)
output_dir = output_path
prompt_component = _sanitize_filename_component(prompt[:100])
out_name = f"{prompt_component}{target_ext}"
video_name = f"{prompt_component}.mp4"
if output_dir:
os.makedirs(output_dir, exist_ok=True)
new_output_path = os.path.join(output_dir, out_name)
new_output_path = os.path.join(output_dir, video_name)
counter = 1
while os.path.exists(new_output_path):
name_part, ext_part = os.path.splitext(out_name)
new_name = f"{name_part}_{counter}{ext_part}"
new_output_path = os.path.join(output_dir, new_name)
name_part, ext_part = os.path.splitext(video_name)
new_video_name = f"{name_part}_{counter}{ext_part}"
new_output_path = os.path.join(output_dir, new_video_name)
counter += 1
return new_output_path
@@ -406,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
@@ -444,25 +387,15 @@ class VideoGenerator:
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
# Save output if requested
# Save video if requested
if batch.save_video:
if self._is_image_workload():
# Image workloads (t2i, i2i, …): save the first frame as PNG.
imageio.imwrite(output_path, frames[0])
logger.info("Saved image to %s", output_path)
else:
imageio.mimsave(output_path,
frames,
fps=batch.fps,
format="mp4")
logger.info("Saved video to %s", output_path)
audio = output_batch.extra.get("audio")
audio_sample_rate = output_batch.extra.get("audio_sample_rate")
if (audio is not None and audio_sample_rate is not None
and not self._mux_audio(output_path, audio,
audio_sample_rate)):
logger.warning(
"Audio mux failed; saved video without audio.")
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
audio = output_batch.extra.get("audio")
audio_sample_rate = output_batch.extra.get("audio_sample_rate")
if (audio is not None and audio_sample_rate is not None and
not self._mux_audio(output_path, audio, audio_sample_rate)):
logger.warning("Audio mux failed; saved video without audio.")
if batch.return_frames:
return frames
+185 -12
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(
@@ -903,7 +1088,6 @@ class TrainingArgs(FastVideoArgs):
lora_rank: int | None = None
lora_alpha: int | None = None
lora_training: bool = False
ltx2_first_frame_conditioning_p: float = 0.1
# distillation args
generator_update_interval: int = 5
@@ -917,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
@@ -1081,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")
@@ -1258,13 +1438,6 @@ class TrainingArgs(FastVideoArgs):
help="Whether to use LoRA training")
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
parser.add_argument("--lora-alpha", type=int, help="LoRA alpha")
parser.add_argument(
"--ltx2-first-frame-conditioning-p",
type=float,
default=TrainingArgs.ltx2_first_frame_conditioning_p,
help=
"Probability of conditioning on the first frame during LTX-2 training",
)
# V-MoBA parameters
parser.add_argument(
+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
+22 -61
View File
@@ -109,45 +109,26 @@ def _apply_rotary_emb(
"""
Args:
x: [num_tokens, num_heads, head_size]
cos: [num_tokens, head_size] or [num_tokens, head_size // 2]
sin: [num_tokens, head_size] or [num_tokens, head_size // 2]
cos: [num_tokens, head_size // 2]
sin: [num_tokens, head_size // 2]
is_neox_style: Whether to use the Neox-style or GPT-J-style rotary
positional embeddings.
The function auto-detects whether cos/sin are full or half head_size:
- If cos/sin have head_size: use rotate_half style (for HunyuanVideo/GameCraft)
- If cos/sin have head_size // 2: use Neox/GPT-J style
"""
head_size = x.shape[-1]
rope_dim = cos.shape[-1]
# Check if cos/sin are full head_dim (rotate_half style) or half (traditional style)
if rope_dim == head_size:
# Full head_dim - use rotate_half style (HunyuanVideo, GameCraft)
# x * cos + rotate_half(x) * sin
cos = cos.unsqueeze(-2) # [num_tokens, 1, head_size]
sin = sin.unsqueeze(-2) # [num_tokens, 1, head_size]
# rotate_half: split into pairs, negate and swap
x_real, x_imag = x.float().reshape(*x.shape[:-1], -1,
2).unbind(-1) # [B, H, D//2] each
x_rotated = torch.stack([-x_imag, x_real],
dim=-1).flatten(-2) # [B, H, D]
return (x.float() * cos + x_rotated * sin).type_as(x)
# cos = cos.unsqueeze(-2).to(x.dtype)
# sin = sin.unsqueeze(-2).to(x.dtype)
cos = cos.unsqueeze(-2)
sin = sin.unsqueeze(-2)
if is_neox_style:
x1, x2 = torch.chunk(x, 2, dim=-1)
else:
# Half head_dim - use traditional Neox/GPT-J style
cos = cos.unsqueeze(-2)
sin = sin.unsqueeze(-2)
if is_neox_style:
x1, x2 = torch.chunk(x, 2, dim=-1)
else:
x1 = x[..., ::2]
x2 = x[..., 1::2]
o1 = (x1.float() * cos - x2.float() * sin).type_as(x)
o2 = (x2.float() * cos + x1.float() * sin).type_as(x)
if is_neox_style:
return torch.cat((o1, o2), dim=-1)
else:
return torch.stack((o1, o2), dim=-1).flatten(-2)
x1 = x[..., ::2]
x2 = x[..., 1::2]
o1 = (x1.float() * cos - x2.float() * sin).type_as(x)
o2 = (x2.float() * cos + x1.float() * sin).type_as(x)
if is_neox_style:
return torch.cat((o1, o2), dim=-1)
else:
return torch.stack((o1, o2), dim=-1).flatten(-2)
@CustomOp.register("rotary_embedding")
@@ -297,7 +278,6 @@ def get_1d_rotary_pos_embed(
theta_rescale_factor: float = 1.0,
interpolation_factor: float = 1.0,
dtype: torch.dtype = torch.float32,
use_real: bool = True,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
@@ -312,12 +292,9 @@ def get_1d_rotary_pos_embed(
theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.
theta_rescale_factor (float, optional): Rescale factor for theta. Defaults to 1.0.
interpolation_factor (float, optional): Factor to scale positions. Defaults to 1.0.
use_real (bool, optional): If True, output full head_dim with repeated cos/sin for
rotate_half style RoPE. If False, output half head_dim for complex style. Defaults to True.
Returns:
freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately.
Shape is [S, D] if use_real=True, [S, D/2] if use_real=False.
freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately. [S, D]
"""
if isinstance(pos, int):
pos = torch.arange(pos).float()
@@ -332,16 +309,6 @@ def get_1d_rotary_pos_embed(
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
freqs_cos = freqs.cos() # [S, D/2]
freqs_sin = freqs.sin() # [S, D/2]
if use_real:
# For rotate_half style RoPE (used by HunyuanVideo, GameCraft),
# we need to expand cos/sin to full head_dim using repeat_interleave.
# The rotate_half operation works on consecutive PAIRS: (x0,x1), (x2,x3)...
# so cos/sin must be interleaved: [c0,c0,c1,c1,...] to match the pairing.
# Using torch.cat would produce [c0,c1,...,c0,c1,...] which is WRONG.
freqs_cos = freqs_cos.repeat_interleave(2, dim=-1) # [S, D]
freqs_sin = freqs_sin.repeat_interleave(2, dim=-1) # [S, D]
return freqs_cos, freqs_sin
@@ -357,7 +324,6 @@ def get_nd_rotary_pos_embed(
sp_world_size: int = 1,
dtype: torch.dtype = torch.float32,
start_frame: int = 0,
use_real: bool = True,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
@@ -375,10 +341,9 @@ def get_nd_rotary_pos_embed(
shard_dim (int): Which dimension to shard for sequence parallelism. Defaults to 0.
sp_rank (int): Rank in the sequence parallel group. Defaults to 0.
sp_world_size (int): World size of the sequence parallel group. Defaults to 1.
use_real (bool): If True, output full head_dim for rotate_half style. Defaults to True.
Returns:
Tuple[torch.Tensor, torch.Tensor]: (cos, sin) tensors of shape [HW, D] if use_real, [HW, D/2] otherwise
Tuple[torch.Tensor, torch.Tensor]: (cos, sin) tensors of shape [HW, D/2]
"""
# Get the full grid
full_grid = get_meshgrid_nd(
@@ -447,12 +412,11 @@ def get_nd_rotary_pos_embed(
theta_rescale_factor=theta_rescale_factor[i],
interpolation_factor=interpolation_factor[i],
dtype=dtype,
use_real=use_real,
) # 2 x [WHD, rope_dim_list[i]] or 2 x [WHD, rope_dim_list[i]*2] if use_real
) # 2 x [WHD, rope_dim_list[i]]
embs.append(emb)
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D) or (WHD, D/2)
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D) or (WHD, D/2)
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D/2)
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D/2)
return cos, sin
@@ -468,7 +432,6 @@ def get_rotary_pos_embed(
do_sp_sharding: bool = False,
dtype: torch.dtype = torch.float32,
start_frame: int = 0,
use_real: bool = True,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Generate rotary positional embeddings for the given sizes.
@@ -483,10 +446,9 @@ def get_rotary_pos_embed(
interpolation_factor: Factor to scale positions. Defaults to 1.0
shard_dim: Which dimension to shard for sequence parallelism. Defaults to 0.
do_sp_sharding: Whether to shard the positional embeddings for sequence parallelism. Defaults to False.
use_real: If True, output full head_dim for rotate_half style RoPE. Defaults to True.
Returns:
Tuple of (cos, sin) tensors for rotary embeddings. Shape [S, D] if use_real, [S, D/2] otherwise.
Tuple of (cos, sin) tensors for rotary embeddings
"""
target_ndim = 3
@@ -519,7 +481,6 @@ def get_rotary_pos_embed(
sp_world_size=sp_world_size,
dtype=dtype,
start_frame=start_frame,
use_real=use_real,
)
return freqs_cos, freqs_sin
+1 -56
View File
@@ -60,61 +60,6 @@ class PatchEmbed(nn.Module):
return x
class WanCamControlPatchEmbedding(nn.Module):
"""Lingbot World Patch embedding for Plucker features."""
def __init__(
self,
patch_size=(1, 2, 2),
in_chans=384, # 6 * 64
embed_dim=2048,
bias=True,
dtype=None,
prefix: str = ""):
super().__init__()
# must be 3-tuple
if isinstance(patch_size, list | tuple):
if len(patch_size) != 3:
raise ValueError(
f"patch_size must have length 3, got {len(patch_size)}")
else:
raise ValueError(f"Unsupported patch_size type: {type(patch_size)}")
self.patch_size = patch_size
pt, ph, pw = self.patch_size
self.in_features = in_chans * pt * ph * pw
self.proj = nn.Linear(self.in_features,
embed_dim,
bias=bias,
dtype=dtype)
def forward(self, x: torch.Tensor) -> torch.Tensor:
if x.dim() != 5:
raise ValueError(
f"Expected camera embedding shape [B, C, F, H, W], got {x.shape}"
)
bsz, channels, frames, height, width = x.shape
pt, ph, pw = self.patch_size
if (frames % pt) != 0 or (height % ph) != 0 or (width % pw) != 0:
raise ValueError(
f"Input shape {x.shape} must be divisible by patch_size {self.patch_size}"
)
# '1 c (f c1) (h c2) (w c3) -> 1 (f h w) (c c1 c2 c3)',
x = x.view(
bsz,
channels,
frames // pt,
pt,
height // ph,
ph,
width // pw,
pw,
)
x = x.permute(0, 2, 4, 6, 1, 3, 5, 7).reshape(bsz, -1, self.in_features)
return self.proj(x)
class TimestepEmbedder(nn.Module):
"""
Embeds scalar timesteps into vector representations.
@@ -307,4 +252,4 @@ class Timesteps(nn.Module):
downscale_freq_shift=self.downscale_freq_shift,
scale=self.scale,
)
return t_emb
return t_emb
@@ -1,61 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Audio preprocessing helpers for LTX-2 training."""
from __future__ import annotations
import torch
import torchaudio
from torch import nn
class AudioProcessor(nn.Module):
"""Converts audio waveforms to log-mel spectrograms with resampling."""
def __init__(
self,
sample_rate: int,
mel_bins: int,
mel_hop_length: int,
n_fft: int,
) -> None:
super().__init__()
self.sample_rate = sample_rate
self.mel_transform = torchaudio.transforms.MelSpectrogram(
sample_rate=sample_rate,
n_fft=n_fft,
win_length=n_fft,
hop_length=mel_hop_length,
f_min=0.0,
f_max=sample_rate / 2.0,
n_mels=mel_bins,
window_fn=torch.hann_window,
center=True,
pad_mode="reflect",
power=1.0,
mel_scale="slaney",
norm="slaney",
)
def resample_waveform(
self,
waveform: torch.Tensor,
source_rate: int,
target_rate: int,
) -> torch.Tensor:
if source_rate == target_rate:
return waveform
resampled = torchaudio.functional.resample(
waveform, source_rate, target_rate)
return resampled.to(device=waveform.device, dtype=waveform.dtype)
def waveform_to_mel(
self,
waveform: torch.Tensor,
waveform_sample_rate: int,
) -> torch.Tensor:
waveform = self.resample_waveform(
waveform, waveform_sample_rate, self.sample_rate)
mel = self.mel_transform(waveform)
mel = torch.log(torch.clamp(mel, min=1e-5))
mel = mel.to(device=waveform.device, dtype=waveform.dtype)
return mel.permute(0, 1, 3, 2).contiguous()
-6
View File
@@ -1,6 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Camera trajectory utilities for video generation models."""
from fastvideo.models.camera.trajectory import create_camera_trajectory
__all__ = ["create_camera_trajectory"]
-395
View File
@@ -1,395 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Camera trajectory generation for HunyuanGameCraft.
Generates Plücker coordinate embeddings from simple action commands
(forward, backward, left, right, rotations) for camera conditioning.
This is a self-contained implementation that does not depend on the
official Hunyuan-GameCraft-1.0 repository.
"""
import math
import numpy as np
import torch
from packaging import version as pver
# Action name -> motion type mapping
ACTION_DICT = {
"w": "forward",
"a": "left",
"d": "right",
"s": "backward",
"forward": "forward",
"backward": "backward",
"left": "left",
"right": "right",
"left_rot": "left_rot",
"right_rot": "right_rot",
"up_rot": "up_rot",
"down_rot": "down_rot",
}
def _custom_meshgrid(*args):
"""Torch meshgrid with consistent indexing."""
if pver.parse(torch.__version__) < pver.parse("1.10"):
return torch.meshgrid(*args)
else:
return torch.meshgrid(*args, indexing="ij")
def _generate_motion_segment(
current_pose: dict,
motion_type: str,
value: float,
duration: int = 30,
) -> tuple[list, list, dict]:
"""
Generate camera motion segment.
Args:
current_pose: Dict with 'position' (xyz) and 'rotation' (pitch, yaw, roll).
motion_type: One of 'forward', 'backward', 'left', 'right',
'left_rot', 'right_rot', 'up_rot', 'down_rot'.
value: Translation (meters) or rotation (degrees).
duration: Number of frames.
Returns:
positions: List of position arrays.
rotations: List of rotation arrays.
current_pose: Updated pose dict.
"""
positions = []
rotations = []
if motion_type in ["forward", "backward"]:
yaw_rad = np.radians(current_pose["rotation"][1])
pitch_rad = np.radians(current_pose["rotation"][0])
forward_vec = np.array([
-math.sin(yaw_rad) * math.cos(pitch_rad),
math.sin(pitch_rad),
-math.cos(yaw_rad) * math.cos(pitch_rad),
])
direction = 1 if motion_type == "forward" else -1
total_move = forward_vec * value * direction
step = total_move / duration
for i in range(1, duration + 1):
new_pos = current_pose["position"] + step * i
positions.append(new_pos.copy())
rotations.append(current_pose["rotation"].copy())
current_pose["position"] = positions[-1]
elif motion_type in ["left", "right"]:
yaw_rad = np.radians(current_pose["rotation"][1])
right_vec = np.array([math.cos(yaw_rad), 0, -math.sin(yaw_rad)])
direction = -1 if motion_type == "right" else 1
total_move = right_vec * value * direction
step = total_move / duration
for i in range(1, duration + 1):
new_pos = current_pose["position"] + step * i
positions.append(new_pos.copy())
rotations.append(current_pose["rotation"].copy())
current_pose["position"] = positions[-1]
elif motion_type.endswith("rot"):
axis = motion_type.split("_")[0]
total_rotation = np.zeros(3)
if axis == "left":
total_rotation[0] = value
elif axis == "right":
total_rotation[0] = -value
elif axis == "up":
total_rotation[2] = -value
elif axis == "down":
total_rotation[2] = value
step = total_rotation / duration
for i in range(1, duration + 1):
positions.append(current_pose["position"].copy())
new_rot = current_pose["rotation"] + step * i
rotations.append(new_rot.copy())
current_pose["rotation"] = rotations[-1]
return positions, rotations, current_pose
def _euler_to_quaternion(angles: np.ndarray) -> list[float]:
"""Convert Euler angles (pitch, yaw, roll in degrees) to quaternion."""
pitch, yaw, roll = np.radians(angles)
cy = math.cos(yaw * 0.5)
sy = math.sin(yaw * 0.5)
cp = math.cos(pitch * 0.5)
sp = math.sin(pitch * 0.5)
cr = math.cos(roll * 0.5)
sr = math.sin(roll * 0.5)
qw = cy * cp * cr + sy * sp * sr
qx = cy * cp * sr - sy * sp * cr
qy = sy * cp * sr + cy * sp * cr
qz = sy * cp * cr - cy * sp * sr
return [qw, qx, qy, qz]
def _quaternion_to_rotation_matrix(q: list[float]) -> np.ndarray:
"""Convert quaternion to 3x3 rotation matrix."""
qw, qx, qy, qz = q
return np.array([
[1 - 2 * (qy**2 + qz**2), 2 * (qx * qy - qw * qz), 2 * (qx * qz + qw * qy)],
[2 * (qx * qy + qw * qz), 1 - 2 * (qx**2 + qz**2), 2 * (qy * qz - qw * qx)],
[2 * (qx * qz - qw * qy), 2 * (qy * qz + qw * qx), 1 - 2 * (qx**2 + qy**2)],
])
def _action_to_pose_list(action_id: str, value: float = 0.2, duration: int = 33) -> list[str]:
"""
Convert an action ID to a list of pose strings.
Args:
action_id: Action identifier (e.g., 'w', 'forward', 'left_rot').
value: Motion magnitude (translation in meters, rotation in degrees).
duration: Number of frames.
Returns:
List of pose strings in the official GameCraft format.
"""
all_positions = []
all_rotations = []
current_pose = {
"position": np.array([0.0, 0.0, 0.0]),
"rotation": np.array([0.0, 0.0, 0.0]),
}
intrinsic = [0.50505, 0.8979, 0.5, 0.5]
motion_type = ACTION_DICT.get(action_id, action_id)
positions, rotations, current_pose = _generate_motion_segment(
current_pose, motion_type, value, duration
)
all_positions.extend(positions)
all_rotations.extend(rotations)
pose_list = []
# First frame: identity pose
row = [0] + intrinsic + [0, 0] + [1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0]
first_row = " ".join(map(str, row))
pose_list.append(first_row)
for i, (pos, rot) in enumerate(zip(all_positions, all_rotations)):
quat = _euler_to_quaternion(rot)
R = _quaternion_to_rotation_matrix(quat)
extrinsic = np.hstack([R, pos.reshape(3, 1)])
row = [i] + intrinsic + [0, 0] + extrinsic.flatten().tolist()
pose_list.append(" ".join(map(str, row)))
return pose_list
class _Camera:
"""Camera parameters from a pose string."""
def __init__(self, entry: list[float]):
fx, fy, cx, cy = entry[1:5]
self.fx = fx
self.fy = fy
self.cx = cx
self.cy = cy
w2c_mat = np.array(entry[7:]).reshape(3, 4)
w2c_mat_4x4 = np.eye(4)
w2c_mat_4x4[:3, :] = w2c_mat
self.w2c_mat = w2c_mat_4x4
self.c2w_mat = np.linalg.inv(w2c_mat_4x4)
def _get_relative_pose(cam_params: list[_Camera]) -> np.ndarray:
"""Convert camera parameters to relative poses."""
abs_w2cs = [cam_param.w2c_mat for cam_param in cam_params]
abs_c2ws = [cam_param.c2w_mat for cam_param in cam_params]
target_cam_c2w = np.array([
[1, 0, 0, 0],
[0, 1, 0, 0],
[0, 0, 1, 0],
[0, 0, 0, 1],
])
abs2rel = target_cam_c2w @ abs_w2cs[0]
ret_poses = [target_cam_c2w] + [abs2rel @ abs_c2w for abs_c2w in abs_c2ws[1:]]
for pose in ret_poses:
pose[:3, -1:] *= 10
ret_poses = np.array(ret_poses, dtype=np.float32)
return ret_poses
def _get_c2w(w2cs: list[np.ndarray], transform_matrix: np.ndarray) -> np.ndarray:
"""Convert w2c matrices to c2w with relative transform."""
target_cam_c2w = np.array([
[1, 0, 0, 0],
[0, 1, 0, 0],
[0, 0, 1, 0],
[0, 0, 0, 1],
])
abs2rel = target_cam_c2w @ w2cs[0]
ret_poses = [target_cam_c2w] + [abs2rel @ np.linalg.inv(w2c) for w2c in w2cs[1:]]
for pose in ret_poses:
pose[:3, -1:] *= 2
ret_poses = [transform_matrix @ x for x in ret_poses]
return np.array(ret_poses, dtype=np.float32)
def _ray_condition(
K: torch.Tensor,
c2w: torch.Tensor,
H: int,
W: int,
device: str | torch.device = "cpu",
flip_flag: torch.Tensor | None = None,
) -> torch.Tensor:
"""
Compute Plücker coordinates from camera intrinsics and extrinsics.
Args:
K: Intrinsics [B, V, 4] (fx, fy, cx, cy).
c2w: Camera-to-world matrices [B, V, 4, 4].
H: Image height.
W: Image width.
device: Torch device.
flip_flag: Optional flip flags [V].
Returns:
Plücker coordinates [B, V, H, W, 6].
"""
B, V = K.shape[:2]
j, i = _custom_meshgrid(
torch.linspace(0, H - 1, H, device=device, dtype=c2w.dtype),
torch.linspace(0, W - 1, W, device=device, dtype=c2w.dtype),
)
i = i.reshape([1, 1, H * W]).expand([B, V, H * W]) + 0.5
j = j.reshape([1, 1, H * W]).expand([B, V, H * W]) + 0.5
n_flip = torch.sum(flip_flag).item() if flip_flag is not None else 0
if n_flip > 0:
j_flip, i_flip = _custom_meshgrid(
torch.linspace(0, H - 1, H, device=device, dtype=c2w.dtype),
torch.linspace(W - 1, 0, W, device=device, dtype=c2w.dtype),
)
i_flip = i_flip.reshape([1, 1, H * W]).expand(B, 1, H * W) + 0.5
j_flip = j_flip.reshape([1, 1, H * W]).expand(B, 1, H * W) + 0.5
i[:, flip_flag, ...] = i_flip
j[:, flip_flag, ...] = j_flip
fx, fy, cx, cy = K.chunk(4, dim=-1)
zs = torch.ones_like(i)
xs = (i - cx) / fx * zs
ys = (j - cy) / fy * zs
zs = zs.expand_as(ys)
directions = torch.stack((xs, ys, zs), dim=-1)
directions = directions / directions.norm(dim=-1, keepdim=True)
rays_d = directions @ c2w[..., :3, :3].transpose(-1, -2)
rays_o = c2w[..., :3, 3]
rays_o = rays_o[:, :, None].expand_as(rays_d)
rays_dxo = torch.linalg.cross(rays_o, rays_d)
plucker = torch.cat([rays_dxo, rays_d], dim=-1)
plucker = plucker.reshape(B, c2w.shape[1], H, W, 6)
return plucker
def create_camera_trajectory(
action: str,
height: int,
width: int,
num_frames: int,
action_speed: float = 0.2,
device: torch.device | str = "cpu",
dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
"""
Create Plücker coordinate embeddings from an action command.
Args:
action: One of 'forward', 'backward', 'left', 'right',
'left_rot', 'right_rot', 'up_rot', 'down_rot'
(or shorthand 'w', 'a', 's', 'd').
height: Video height in pixels.
width: Video width in pixels.
num_frames: Number of video frames.
action_speed: Speed of motion (default 0.2).
device: Torch device for output tensor.
dtype: Torch dtype for output tensor.
Returns:
camera_states: [1, num_frames, 6, height, width] Plücker embeddings.
"""
# Generate pose list from action
poses = _action_to_pose_list(action, value=action_speed, duration=num_frames)
# Parse poses
poses_parsed = [pose.split(" ") for pose in poses]
start_idx = 0
sample_id = [start_idx + i for i in range(num_frames)]
poses_parsed = [poses_parsed[i] for i in sample_id]
# Convert to w2c matrices
w2cs = [np.asarray([float(p) for p in pose[7:]]).reshape(3, 4) for pose in poses_parsed]
transform_matrix = np.asarray(
[[1, 0, 0, 0], [0, 0, 1, 0], [0, -1, 0, 0], [0, 0, 0, 1]]
).reshape(4, 4)
last_row = np.zeros((1, 4))
last_row[0, -1] = 1.0
w2cs = [np.concatenate((w2c, last_row), axis=0) for w2c in w2cs]
c2ws = _get_c2w(w2cs, transform_matrix)
# Parse camera parameters
cam_params = [[float(x) for x in pose] for pose in poses_parsed]
assert len(cam_params) == num_frames
cam_params = [_Camera(cam_param) for cam_param in cam_params]
# Compute scaled intrinsics
monst3r_w = cam_params[0].cx * 2
monst3r_h = cam_params[0].cy * 2
ratio_w, ratio_h = width / monst3r_w, height / monst3r_h
intrinsics = np.asarray(
[
[
cam_param.fx * ratio_w,
cam_param.fy * ratio_h,
cam_param.cx * ratio_w,
cam_param.cy * ratio_h,
]
for cam_param in cam_params
],
dtype=np.float32,
)
intrinsics = torch.as_tensor(intrinsics)[None] # [1, n_frame, 4]
# Get relative poses
c2w_poses = _get_relative_pose(cam_params)
c2w = torch.as_tensor(c2w_poses)[None] # [1, n_frame, 4, 4]
# Compute Plücker embeddings
flip_flag = torch.zeros(num_frames, dtype=torch.bool, device="cpu")
plucker_embedding = _ray_condition(intrinsics, c2w, height, width, device="cpu", flip_flag=flip_flag)
# [1, n_frame, H, W, 6] -> [1, n_frame, 6, H, W]
plucker_embedding = plucker_embedding[0].permute(0, 3, 1, 2).contiguous()
# Add batch dim and convert to target dtype/device
# Shape: [n_frame, 6, H, W] -> [1, n_frame, 6, H, W]
camera_states = plucker_embedding.unsqueeze(0).to(device=device, dtype=dtype)
return camera_states
+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):
-363
View File
@@ -1,363 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
HunyuanGameCraft Transformer model for FastVideo.
Ported from official Hunyuan-GameCraft-1.0 implementation.
"""
from typing import Any, List, Optional
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from fastvideo.configs.models.dits.hunyuangamecraft import (
HunyuanGameCraftArchConfig,
HunyuanGameCraftConfig,
)
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
from fastvideo.layers.visual_embedding import ModulateProjection, PatchEmbed, TimestepEmbedder, unpatchify
from fastvideo.models.dits.base import CachableDiT
from fastvideo.models.dits.hunyuanvideo import (
MMDoubleStreamBlock,
MMSingleStreamBlock,
SingleTokenRefiner,
)
class GameCraftFinalLayer(nn.Module):
"""
GameCraft-specific FinalLayer with correct shift/scale order.
The official GameCraft implementation uses shift, scale order (not scale, shift).
This differs from the HunyuanVideo FinalLayer.
"""
def __init__(self,
hidden_size,
patch_size,
out_channels,
dtype=None,
prefix: str = "") -> None:
super().__init__()
self.norm_final = nn.LayerNorm(hidden_size,
eps=1e-6,
elementwise_affine=False,
dtype=dtype)
output_dim = patch_size[0] * patch_size[1] * patch_size[2] * out_channels
self.linear = ReplicatedLinear(hidden_size,
output_dim,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear")
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
def forward(self, x, c):
# GameCraft uses shift, scale order (verified against official implementation)
shift, scale = self.adaLN_modulation(c).chunk(2, dim=-1)
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
x, _ = self.linear(x)
return x
class CameraNet(nn.Module):
"""
Camera state encoding network - ported from official GameCraft.
Processes camera parameters (Plücker coordinates) into feature embeddings.
"""
def __init__(
self,
in_channels: int = 6,
downscale_coef: int = 8,
out_channels: int = 16,
patch_size: List[int] = [1, 2, 2],
hidden_size: int = 3072,
dtype: Optional[torch.dtype] = None,
prefix: str = "",
):
super().__init__()
_ = prefix # Unused
start_channels = in_channels * (downscale_coef ** 2)
input_channels = [start_channels, start_channels // 2, start_channels // 4]
self.input_channels = input_channels
self.unshuffle = nn.PixelUnshuffle(downscale_coef)
self.encode_first = nn.Sequential(
nn.Conv2d(input_channels[0], input_channels[1], kernel_size=1, stride=1, padding=0),
nn.GroupNorm(2, input_channels[1]),
nn.ReLU(),
)
self._initialize_weights(self.encode_first)
self.encode_second = nn.Sequential(
nn.Conv2d(input_channels[1], input_channels[2], kernel_size=1, stride=1, padding=0),
nn.GroupNorm(2, input_channels[2]),
nn.ReLU(),
)
self._initialize_weights(self.encode_second)
self.final_proj = nn.Conv2d(input_channels[2], out_channels, kernel_size=1)
self._zeros_init_linear(self.final_proj)
self.scale = nn.Parameter(torch.ones(1))
self.camera_in = PatchEmbed(
patch_size=patch_size,
in_chans=out_channels,
embed_dim=hidden_size,
)
def _zeros_init_linear(self, linear):
if hasattr(linear, "weight"):
nn.init.zeros_(linear.weight)
if hasattr(linear, "bias") and linear.bias is not None:
nn.init.zeros_(linear.bias)
def _initialize_weights(self, block):
for m in block:
if isinstance(m, nn.Conv2d):
n = m.kernel_size[0] * m.kernel_size[1] * m.in_channels
nn.init.normal_(m.weight, mean=0.0, std=np.sqrt(2.0 / n))
if m.bias is not None:
nn.init.zeros_(m.bias)
def compress_time(self, x: torch.Tensor, num_frames: int) -> torch.Tensor:
x = rearrange(x, '(b f) c h w -> b f c h w', f=num_frames)
batch_size, frames, channels, height, width = x.shape
x = rearrange(x, 'b f c h w -> (b h w) c f')
if x.shape[-1] == 66 or x.shape[-1] == 34:
x_len = x.shape[-1]
x_clip1 = x[..., :x_len // 2]
x_clip1_first = x_clip1[..., 0].unsqueeze(-1)
x_clip1_rest = F.avg_pool1d(x_clip1[..., 1:], kernel_size=2, stride=2)
x_clip2 = x[..., x_len // 2:]
x_clip2_first = x_clip2[..., 0].unsqueeze(-1)
x_clip2_rest = F.avg_pool1d(x_clip2[..., 1:], kernel_size=2, stride=2)
x = torch.cat([x_clip1_first, x_clip1_rest, x_clip2_first, x_clip2_rest], dim=-1)
elif x.shape[-1] % 2 == 1:
x_first = x[..., 0]
x_rest = x[..., 1:]
if x_rest.shape[-1] > 0:
x_rest = F.avg_pool1d(x_rest, kernel_size=2, stride=2)
x = torch.cat([x_first[..., None], x_rest], dim=-1)
else:
x = F.avg_pool1d(x, kernel_size=2, stride=2)
x = rearrange(x, '(b h w) c f -> (b f) c h w', b=batch_size, h=height, w=width)
return x
def forward(self, camera_states: torch.Tensor) -> torch.Tensor:
batch_size, num_frames, channels, height, width = camera_states.shape
camera_states = rearrange(camera_states, 'b f c h w -> (b f) c h w')
camera_states = self.unshuffle(camera_states)
camera_states = self.encode_first(camera_states)
camera_states = self.compress_time(camera_states, num_frames=num_frames)
num_frames = camera_states.shape[0] // batch_size
camera_states = self.encode_second(camera_states)
camera_states = self.compress_time(camera_states, num_frames=num_frames)
camera_states = self.final_proj(camera_states)
camera_states = rearrange(camera_states, "(b f) c h w -> b c f h w", b=batch_size)
camera_states = self.camera_in(camera_states)
return camera_states * self.scale
class HunyuanGameCraftTransformer3DModel(CachableDiT):
"""
HunyuanGameCraft Transformer - ported from official implementation.
"""
_fsdp_shard_conditions = HunyuanGameCraftArchConfig()._fsdp_shard_conditions
_compile_conditions = HunyuanGameCraftArchConfig()._compile_conditions
_supported_attention_backends = HunyuanGameCraftArchConfig()._supported_attention_backends
param_names_mapping = HunyuanGameCraftConfig().param_names_mapping
reverse_param_names_mapping = HunyuanGameCraftConfig().reverse_param_names_mapping
def __init__(self, config: HunyuanGameCraftConfig, hf_config: dict[str, Any]):
super().__init__(config=config, hf_config=hf_config)
arch = config.arch_config
if isinstance(arch.patch_size, (list, tuple)):
self.patch_size = list(arch.patch_size)
else:
self.patch_size = [arch.patch_size_t, arch.patch_size, arch.patch_size]
self.in_channels = arch.in_channels
self.out_channels = arch.out_channels
self.unpatchify_channels = self.out_channels
self.num_channels_latents = self.out_channels # Alias for latent_preparation stage
self.hidden_size = arch.hidden_size
self.num_heads = arch.num_attention_heads
self.num_attention_heads = arch.num_attention_heads # Alias for compatibility
self.guidance_embeds = arch.guidance_embeds
self.rope_dim_list = list(arch.rope_axes_dim)
self.rope_theta = arch.rope_theta
self.text_states_dim = arch.text_embed_dim
self.text_states_dim_2 = arch.pooled_projection_dim
self.dtype = arch.dtype
pe_dim = self.hidden_size // self.num_heads
if sum(self.rope_dim_list) != pe_dim:
raise ValueError(f"rope_axes_dim sum {sum(self.rope_dim_list)} != {pe_dim}")
factory_kwargs = {'dtype': self.dtype}
self.img_in = PatchEmbed(
patch_size=self.patch_size,
in_chans=self.in_channels,
embed_dim=self.hidden_size,
**factory_kwargs,
)
self.txt_in = SingleTokenRefiner(
self.text_states_dim,
self.hidden_size,
self.num_heads,
depth=arch.num_refiner_layers,
**factory_kwargs,
)
self.time_in = TimestepEmbedder(self.hidden_size, **factory_kwargs)
self.vector_in = MLP(
self.text_states_dim_2,
self.hidden_size,
self.hidden_size,
act_type="silu",
**factory_kwargs,
)
self.guidance_in = (
TimestepEmbedder(self.hidden_size, **factory_kwargs)
if self.guidance_embeds else None
)
self.double_blocks = nn.ModuleList([
MMDoubleStreamBlock(
hidden_size=self.hidden_size,
num_attention_heads=self.num_heads,
mlp_ratio=arch.mlp_ratio,
supported_attention_backends=self._supported_attention_backends,
**factory_kwargs,
)
for _ in range(arch.num_layers)
])
self.single_blocks = nn.ModuleList([
MMSingleStreamBlock(
hidden_size=self.hidden_size,
num_attention_heads=self.num_heads,
mlp_ratio=arch.mlp_ratio,
supported_attention_backends=self._supported_attention_backends,
**factory_kwargs,
)
for _ in range(arch.num_single_layers)
])
self.final_layer = GameCraftFinalLayer(
self.hidden_size,
self.patch_size,
self.out_channels,
**factory_kwargs,
)
self.camera_net = CameraNet(
in_channels=arch.camera_in_channels,
out_channels=16,
downscale_coef=arch.camera_downscale_coef,
patch_size=self.patch_size,
hidden_size=self.hidden_size,
)
def forward(
self,
x: torch.Tensor,
encoder_hidden_states: List[torch.Tensor],
timestep: torch.Tensor,
camera_states: Optional[torch.Tensor] = None,
encoder_attention_mask: Optional[List[torch.Tensor]] = None,
guidance: Optional[torch.Tensor] = None,
return_dict: bool = False,
) -> torch.Tensor:
img = x
_, _, ot, oh, ow = x.shape
tt = ot // self.patch_size[0]
th = oh // self.patch_size[1]
tw = ow // self.patch_size[2]
text_states = encoder_hidden_states[0]
text_states_2 = encoder_hidden_states[1] if len(encoder_hidden_states) > 1 else None
text_mask = encoder_attention_mask[0] if encoder_attention_mask else None
vec = self.time_in(timestep)
if text_states_2 is not None:
vec = vec + self.vector_in(text_states_2)
if self.guidance_in is not None and guidance is not None:
vec = vec + self.guidance_in(guidance)
img = self.img_in(img)
if camera_states is not None:
latent_len = ot
if latent_len == 18:
camera_latents = torch.cat([
self.camera_net(torch.zeros_like(camera_states)),
self.camera_net(camera_states)
], dim=1)
elif latent_len == 9:
camera_latents = self.camera_net(camera_states)
elif latent_len == 10:
camera_latents = torch.cat([
self.camera_net(torch.zeros_like(camera_states[:, 0:4, :, :, :])),
self.camera_net(camera_states)
], dim=1)
else:
camera_latents = self.camera_net(camera_states)
img = img + camera_latents
txt = self.txt_in(text_states, timestep)
txt_seq_len = txt.shape[1]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(tt, th, tw),
self.hidden_size,
self.num_heads,
self.rope_dim_list,
self.rope_theta,
)
freqs_cos = freqs_cos.to(device=img.device, dtype=img.dtype)
freqs_sin = freqs_sin.to(device=img.device, dtype=img.dtype)
freqs_cis = (freqs_cos, freqs_sin)
for block in self.double_blocks:
img, txt = block(img, txt, vec, freqs_cis)
x = torch.cat([img, txt], dim=1)
for block in self.single_blocks:
x = block(x, vec, txt_seq_len, freqs_cis)
img = x[:, :-txt_seq_len, ...]
img = self.final_layer(img, vec)
img = unpatchify(img, tt, th, tw, self.patch_size, self.out_channels)
return img
@@ -1,8 +0,0 @@
from .model import LingBotWorldTransformer3DModel
__all__ = [
"LingBotWorldTransformer3DModel",
]
# Entry point for model registry
EntryClass = [LingBotWorldTransformer3DModel]
@@ -1,203 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from LingBot World: https://github.com/Robbyant/lingbot-world/blob/main/wan/utils/cam_utils.py
import numpy as np
import os
import torch
from scipy.interpolate import interp1d
from scipy.spatial.transform import Rotation, Slerp
# --- Official Code (Leave Unchanged) ---
def interpolate_camera_poses(
src_indices: np.ndarray,
src_rot_mat: np.ndarray,
src_trans_vec: np.ndarray,
tgt_indices: np.ndarray,
) -> torch.Tensor:
# interpolate translation
interp_func_trans = interp1d(
src_indices,
src_trans_vec,
axis=0,
kind='linear',
bounds_error=False,
fill_value="extrapolate",
)
interpolated_trans_vec = interp_func_trans(tgt_indices)
# interpolate rotation
src_quat_vec = Rotation.from_matrix(src_rot_mat)
# ensure there is no sudden change in qw
quats = src_quat_vec.as_quat().copy() # [N, 4]
for i in range(1, len(quats)):
if np.dot(quats[i], quats[i-1]) < 0:
quats[i] = -quats[i]
src_quat_vec = Rotation.from_quat(quats)
slerp_func_rot = Slerp(src_indices, src_quat_vec)
interpolated_rot_quat = slerp_func_rot(tgt_indices)
interpolated_rot_mat = interpolated_rot_quat.as_matrix()
poses = np.zeros((len(tgt_indices), 4, 4))
poses[:, :3, :3] = interpolated_rot_mat
poses[:, :3, 3] = interpolated_trans_vec
poses[:, 3, 3] = 1.0
return torch.from_numpy(poses).float()
def SE3_inverse(T: torch.Tensor) -> torch.Tensor:
Rot = T[:, :3, :3] # [B,3,3]
trans = T[:, :3, 3:] # [B,3,1]
R_inv = Rot.transpose(-1, -2)
t_inv = -torch.bmm(R_inv, trans)
T_inv = torch.eye(4, device=T.device, dtype=T.dtype)[None, :, :].repeat(T.shape[0], 1, 1)
T_inv[:, :3, :3] = R_inv
T_inv[:, :3, 3:] = t_inv
return T_inv
def compute_relative_poses(
c2ws_mat: torch.Tensor,
framewise: bool = False,
normalize_trans: bool = True,
) -> torch.Tensor:
ref_w2cs = SE3_inverse(c2ws_mat[0:1])
relative_poses = torch.matmul(ref_w2cs, c2ws_mat)
# ensure identity matrix for 1st frame
relative_poses[0] = torch.eye(4, device=c2ws_mat.device, dtype=c2ws_mat.dtype)
if framewise:
# compute pose between i and i+1
relative_poses_framewise = torch.bmm(SE3_inverse(relative_poses[:-1]), relative_poses[1:])
relative_poses[1:] = relative_poses_framewise
if normalize_trans: # note refer to camctrl2: "we scale the coordinate inputs to roughly 1 standard deviation to simplify model learning."
translations = relative_poses[:, :3, 3] # [f, 3]
max_norm = torch.norm(translations, dim=-1).max()
# only normlaize when moving
if max_norm > 0:
relative_poses[:, :3, 3] = translations / max_norm
return relative_poses
@torch.no_grad()
def create_meshgrid(n_frames: int, height: int, width: int, bias: float = 0.5, device='cuda', dtype=torch.float32) -> torch.Tensor:
x_range = torch.arange(width, device=device, dtype=dtype)
y_range = torch.arange(height, device=device, dtype=dtype)
grid_y, grid_x = torch.meshgrid(y_range, x_range, indexing='ij')
grid_xy = torch.stack([grid_x, grid_y], dim=-1).view([-1, 2]) + bias # [h*w, 2]
grid_xy = grid_xy[None, ...].repeat(n_frames, 1, 1) # [f, h*w, 2]
return grid_xy
def get_plucker_embeddings(
c2ws_mat: torch.Tensor,
Ks: torch.Tensor,
height: int,
width: int,
):
n_frames = c2ws_mat.shape[0]
grid_xy = create_meshgrid(n_frames, height, width, device=c2ws_mat.device, dtype=c2ws_mat.dtype) # [f, h*w, 2]
fx, fy, cx, cy = Ks.chunk(4, dim=-1) # [f, 1]
i = grid_xy[..., 0] # [f, h*w]
j = grid_xy[..., 1] # [f, h*w]
zs = torch.ones_like(i) # [f, h*w]
xs = (i - cx) / fx * zs
ys = (j - cy) / fy * zs
directions = torch.stack([xs, ys, zs], dim=-1) # [f, h*w, 3]
directions = directions / directions.norm(dim=-1, keepdim=True) # [f, h*w, 3]
rays_d = directions @ c2ws_mat[:, :3, :3].transpose(-1, -2) # [f, h*w, 3]
rays_o = c2ws_mat[:, :3, 3] # [f, 3]
rays_o = rays_o[:, None, :].expand_as(rays_d) # [f, h*w, 3]
# rays_dxo = torch.cross(rays_o, rays_d, dim=-1) # [f, h*w, 3]
# note refer to: apt2
plucker_embeddings = torch.cat([rays_o, rays_d], dim=-1) # [f, h*w, 6]
plucker_embeddings = plucker_embeddings.view([n_frames, height, width, 6]) # [f*h*w, 6]
return plucker_embeddings
def get_Ks_transformed(
Ks: torch.Tensor,
height_org: int,
width_org: int,
height_resize: int,
width_resize: int,
height_final: int,
width_final: int,
):
fx, fy, cx, cy = Ks.chunk(4, dim=-1) # [f, 1]
scale_x = width_resize / width_org
scale_y = height_resize / height_org
fx_resize = fx * scale_x
fy_resize = fy * scale_y
cx_resize = cx * scale_x
cy_resize = cy * scale_y
crop_offset_x = (width_resize - width_final) / 2
crop_offset_y = (height_resize - height_final) / 2
cx_final = cx_resize - crop_offset_x
cy_final = cy_resize - crop_offset_y
Ks_transformed = torch.zeros_like(Ks)
Ks_transformed[:, 0:1] = fx_resize
Ks_transformed[:, 1:2] = fy_resize
Ks_transformed[:, 2:3] = cx_final
Ks_transformed[:, 3:4] = cy_final
return Ks_transformed
# --- Custom ---
def prepare_camera_embedding(
action_path: str,
num_frames: int,
height: int,
width: int,
spatial_scale: int = 8,
) -> tuple[torch.Tensor, int]:
c2ws = np.load(os.path.join(action_path, "poses.npy"))
len_c2ws = ((len(c2ws) - 1) // 4) * 4 + 1
num_frames = min(num_frames, len_c2ws)
c2ws = c2ws[:num_frames]
Ks = torch.from_numpy(
np.load(os.path.join(action_path, "intrinsics.npy"))
).float()
Ks = get_Ks_transformed(
Ks,
height_org=480,
width_org=832,
height_resize=height,
width_resize=width,
height_final=height,
width_final=width,
)
Ks = Ks[0] # use first frame
len_c2ws = len(c2ws)
num_latent_frames = (len_c2ws - 1) // 4 + 1
c2ws_infer = interpolate_camera_poses(
src_indices=np.linspace(0, len_c2ws - 1, len_c2ws),
src_rot_mat=c2ws[:, :3, :3],
src_trans_vec=c2ws[:, :3, 3],
tgt_indices=np.linspace(0, len_c2ws - 1, num_latent_frames),
)
c2ws_infer = compute_relative_poses(c2ws_infer, framewise=True)
Ks = Ks.repeat(num_latent_frames, 1)
plucker = get_plucker_embeddings(c2ws_infer, Ks, height, width) # [F, H, W, 6]
# reshpae
latent_height = height // spatial_scale
latent_width = width // spatial_scale
plucker = plucker.view(num_latent_frames, latent_height, spatial_scale, latent_width, spatial_scale, 6)
plucker = plucker.permute(0, 1, 3, 5, 2, 4).contiguous()
plucker = plucker.view(num_latent_frames, latent_height, latent_width, 6 * spatial_scale * spatial_scale)
c2ws_plucker_emb = plucker.permute(3, 0, 1, 2).contiguous().unsqueeze(0)
return c2ws_plucker_emb, num_frames
-569
View File
@@ -1,569 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import math
from contextlib import nullcontext
from typing import Any
import numpy as np
import torch
import torch.nn as nn
from fastvideo.attention import DistributedAttention
from fastvideo.configs.models.dits.lingbotworld import LingBotWorldVideoConfig
from fastvideo.configs.sample.wan import WanTeaCacheParams
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather_with_unpad,
sequence_model_parallel_shard)
from fastvideo.forward_context import get_forward_context
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
# from torch.nn import RMSNorm
# TODO: RMSNorm ....
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
from fastvideo.layers.visual_embedding import (PatchEmbed, WanCamControlPatchEmbedding)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import CachableDiT
from fastvideo.models.dits.wanvideo import (
WanI2VCrossAttention,
WanT2VCrossAttention,
WanTimeTextImageEmbedding,
)
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.distributed.utils import create_attention_mask_for_padding
logger = init_logger(__name__)
class LingBotWorldCamConditioner(nn.Module):
def __init__(self, dim: int) -> None:
super().__init__()
self.cam_injector = MLP(dim, dim, dim, bias=True, act_type="silu")
self.cam_scale_layer = nn.Linear(dim, dim)
self.cam_shift_layer = nn.Linear(dim, dim)
def forward(
self,
hidden_states: torch.Tensor,
c2ws_plucker_emb: torch.Tensor | None,
) -> torch.Tensor:
if c2ws_plucker_emb is None:
return hidden_states
assert c2ws_plucker_emb.shape == hidden_states.shape, (
f"c2ws_plucker_emb shape must match hidden_states shape, got "
f"{tuple(c2ws_plucker_emb.shape)} vs {tuple(hidden_states.shape)}"
)
c2ws_hidden_states = self.cam_injector(c2ws_plucker_emb)
c2ws_hidden_states = c2ws_hidden_states + c2ws_plucker_emb
cam_scale = self.cam_scale_layer(c2ws_hidden_states)
cam_shift = self.cam_shift_layer(c2ws_hidden_states)
return (1.0 + cam_scale) * hidden_states + cam_shift
class LingBotWorldTransformerBlock(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = DistributedAttention(
num_heads=num_heads,
head_size=dim // num_heads,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1")
self.hidden_dim = dim
self.num_attention_heads = num_heads
dim_head = dim // num_heads
if qk_norm == "rms_norm":
self.norm_q = RMSNorm(dim_head, eps=eps)
self.norm_k = RMSNorm(dim_head, eps=eps)
elif qk_norm == "rms_norm_across_heads":
# LTX applies qk norm across all heads
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
else:
raise NotImplementedError(
f"QK Norm type '{qk_norm}' not supported")
assert cross_attn_norm is True
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
# I2V
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
else:
# T2V
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
self.mlp_residual = ScaleResidual()
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
self.cam_conditioner = LingBotWorldCamConditioner(dim)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
attention_mask: torch.Tensor | None = None,
c2ws_plucker_emb: torch.Tensor | None = None,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
if temb.dim() == 4:
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table.unsqueeze(0) + temb.float()
).chunk(6, dim=2)
# batch_size, seq_len, 1, inner_dim
shift_msa = shift_msa.squeeze(2)
scale_msa = scale_msa.squeeze(2)
gate_msa = gate_msa.squeeze(2)
c_shift_msa = c_shift_msa.squeeze(2)
c_scale_msa = c_scale_msa.squeeze(2)
c_gate_msa = c_gate_msa.squeeze(2)
else:
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
attn_output, _ = self.attn1(query, key, value, freqs_cis=freqs_cis, attention_mask=attention_mask)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
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)
# Inject camera condition
# must be applied after the self-attention residual update.
hidden_states = self.cam_conditioner(hidden_states, c2ws_plucker_emb)
norm_hidden_states = self.self_attn_residual_norm.norm(hidden_states)
norm_hidden_states = norm_hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class LingBotWorldTransformer3DModel(CachableDiT):
_fsdp_shard_conditions = LingBotWorldVideoConfig()._fsdp_shard_conditions
_compile_conditions = LingBotWorldVideoConfig()._compile_conditions
_supported_attention_backends = LingBotWorldVideoConfig(
)._supported_attention_backends
param_names_mapping = LingBotWorldVideoConfig().param_names_mapping
reverse_param_names_mapping = LingBotWorldVideoConfig().reverse_param_names_mapping
lora_param_names_mapping = LingBotWorldVideoConfig().lora_param_names_mapping
def __init__(self, config: LingBotWorldVideoConfig, hf_config: dict[str,
Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size
self.text_len = config.text_len
assert config.num_attention_heads % get_sp_world_size() == 0, f"The number of attention heads ({config.num_attention_heads}) must be divisible by the sequence parallel size ({get_sp_world_size()})"
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False)
self.patch_embedding_wancamctrl = WanCamControlPatchEmbedding(in_chans=6 * 64,
embed_dim=inner_dim,
patch_size=config.patch_size)
self.c2ws_mlp = MLP(inner_dim, inner_dim, inner_dim, bias=True, act_type="silu")
# 2. Condition embeddings
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=config.freq_dim,
text_embed_dim=config.text_dim,
image_embed_dim=config.image_dim,
)
# 3. Transformer blocks
transformer_block = LingBotWorldTransformerBlock
self.blocks = nn.ModuleList([
transformer_block(inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{config.prefix}.blocks.{i}")
for i in range(config.num_layers)
])
# 4. Output norm & projection
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
self._logged_attention_mask = False
# For type checking
self.previous_e0_even = None
self.previous_e0_odd = None
self.previous_residual_even = None
self.previous_residual_odd = None
self.is_even = True
self.should_calc_even = True
self.should_calc_odd = True
self.accumulated_rel_l1_distance_even = 0
self.accumulated_rel_l1_distance_odd = 0
self.cnt = 0
self.__post_init__()
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance=None,
c2ws_plucker_emb: torch.Tensor | None = None,
**kwargs) -> torch.Tensor:
forward_batch = get_forward_context().forward_batch
enable_teacache = forward_batch is not None and forward_batch.enable_teacache
orig_dtype = hidden_states.dtype
if encoder_hidden_states is not None and not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# Get rotary embeddings
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames, post_patch_height,
post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000)
freqs_cis = (freqs_cos.to(hidden_states.device).float(),
freqs_sin.to(hidden_states.device).float())
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
c2ws_hidden_states = None
if c2ws_plucker_emb is not None:
c2ws_plucker_emb = self.patch_embedding_wancamctrl(
c2ws_plucker_emb.to(device=hidden_states.device, dtype=hidden_states.dtype)
)
c2ws_hidden_states = self.c2ws_mlp(c2ws_plucker_emb)
c2ws_plucker_emb = c2ws_plucker_emb + c2ws_hidden_states
# Shard with padding support - returns (sharded_tensor, original_seq_len)
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
# Shard c2ws_plucker_emb
if c2ws_plucker_emb is not None:
c2ws_plucker_emb, _ = sequence_model_parallel_shard(c2ws_plucker_emb, dim=1)
# Create attention mask for padded tokens if padding was applied
current_seq_len = hidden_states.shape[1]
sp_world_size = get_sp_world_size()
padded_seq_len = current_seq_len * sp_world_size
if padded_seq_len > original_seq_len:
if not self._logged_attention_mask:
logger.info(f"Padding applied, original seq len: {original_seq_len}, padded seq len: {padded_seq_len}")
self._logged_attention_mask = True
attention_mask = create_attention_mask_for_padding(
seq_len=original_seq_len,
padded_seq_len=padded_seq_len,
batch_size=batch_size,
device=hidden_states.device,
)
else:
if not self._logged_attention_mask:
logger.info(f"Padding not applied")
self._logged_attention_mask = True
attention_mask = None
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
timestep = timestep.flatten() # batch_size * seq_len
else:
ts_seq_len = None
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
if ts_seq_len is not None:
# batch_size, seq_len, 6, inner_dim
timestep_proj = timestep_proj.unflatten(2, (6, -1))
else:
# batch_size, 6, inner_dim
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
if encoder_hidden_states is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
else:
encoder_hidden_states = encoder_hidden_states_image
if current_platform.is_mps() or current_platform.is_npu():
encoder_hidden_states = encoder_hidden_states.to(orig_dtype)
else:
encoder_hidden_states = encoder_hidden_states # cast to orig_dtype for MPS & NPU
assert encoder_hidden_states.dtype == orig_dtype
# 4. Transformer blocks
# if caching is enabled, we might be able to skip the forward pass
should_skip_forward = self.should_skip_forward_for_cached_states(
timestep_proj=timestep_proj, temb=temb)
if should_skip_forward:
print("skipping forward, cached")
hidden_states = self.retrieve_cached_states(hidden_states)
else:
# if teacache is enabled, we need to cache the original hidden states
if enable_teacache:
original_hidden_states = hidden_states.clone()
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, attention_mask, c2ws_plucker_emb)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, attention_mask, c2ws_plucker_emb)
# if teacache is enabled, we need to cache the original hidden states
if enable_teacache:
self.maybe_cache_states(hidden_states, original_hidden_states)
# 5. Output norm, projection & unpatchify
if temb.dim() == 3:
# batch_size, seq_len, inner_dim (wan 2.2 ti2v)
shift, scale = (self.scale_shift_table.unsqueeze(0) + temb.unsqueeze(2)).chunk(2, dim=2)
shift = shift.squeeze(2)
scale = scale.squeeze(2)
else:
# batch_size, inner_dim
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
# Gather and unpad in one operation
hidden_states = sequence_model_parallel_all_gather_with_unpad(
hidden_states, original_seq_len, dim=1)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
def maybe_cache_states(self, hidden_states: torch.Tensor,
original_hidden_states: torch.Tensor) -> None:
if self.is_even:
self.previous_residual_even = hidden_states.squeeze(
0) - original_hidden_states
else:
self.previous_residual_odd = hidden_states.squeeze(
0) - original_hidden_states
def should_skip_forward_for_cached_states(self, **kwargs) -> bool:
forward_context = get_forward_context()
forward_batch = forward_context.forward_batch
if forward_batch is None or not forward_batch.enable_teacache:
return False
teacache_params = forward_batch.teacache_params
assert teacache_params is not None, "teacache_params is not initialized"
assert isinstance(
teacache_params,
WanTeaCacheParams), "teacache_params is not a WanTeaCacheParams"
current_timestep = forward_context.current_timestep
num_inference_steps = forward_batch.num_inference_steps
# initialize the coefficients, cutoff_steps, and ret_steps
coefficients = teacache_params.coefficients
use_ret_steps = teacache_params.use_ret_steps
cutoff_steps = teacache_params.get_cutoff_steps(num_inference_steps)
ret_steps = teacache_params.ret_steps
teacache_thresh = teacache_params.teacache_thresh
if current_timestep == 0:
self.cnt = 0
timestep_proj = kwargs["timestep_proj"]
temb = kwargs["temb"]
modulated_inp = timestep_proj if use_ret_steps else temb
if self.cnt % 2 == 0: # even -> condition
self.is_even = True
if self.cnt < ret_steps or self.cnt >= cutoff_steps:
self.should_calc_even = True
self.accumulated_rel_l1_distance_even = 0
else:
assert self.previous_e0_even is not None, "previous_e0_even is not initialized"
assert self.accumulated_rel_l1_distance_even is not None, "accumulated_rel_l1_distance_even is not initialized"
rescale_func = np.poly1d(coefficients)
self.accumulated_rel_l1_distance_even += rescale_func(
((modulated_inp - self.previous_e0_even).abs().mean() /
self.previous_e0_even.abs().mean()).cpu().item())
if self.accumulated_rel_l1_distance_even < teacache_thresh:
self.should_calc_even = False
else:
self.should_calc_even = True
self.accumulated_rel_l1_distance_even = 0
self.previous_e0_even = modulated_inp.clone()
else: # odd -> unconditon
self.is_even = False
if self.cnt < ret_steps or self.cnt >= cutoff_steps:
self.should_calc_odd = True
self.accumulated_rel_l1_distance_odd = 0
else:
assert self.previous_e0_odd is not None, "previous_e0_odd is not initialized"
assert self.accumulated_rel_l1_distance_odd is not None, "accumulated_rel_l1_distance_odd is not initialized"
rescale_func = np.poly1d(coefficients)
self.accumulated_rel_l1_distance_odd += rescale_func(
((modulated_inp - self.previous_e0_odd).abs().mean() /
self.previous_e0_odd.abs().mean()).cpu().item())
if self.accumulated_rel_l1_distance_odd < teacache_thresh:
self.should_calc_odd = False
else:
self.should_calc_odd = True
self.accumulated_rel_l1_distance_odd = 0
self.previous_e0_odd = modulated_inp.clone()
self.cnt += 1
should_skip_forward = False
if self.is_even:
if not self.should_calc_even:
should_skip_forward = True
else:
if not self.should_calc_odd:
should_skip_forward = True
return should_skip_forward
def retrieve_cached_states(self,
hidden_states: torch.Tensor) -> torch.Tensor:
if self.is_even:
return hidden_states + self.previous_residual_even
else:
return hidden_states + self.previous_residual_odd
# Entry point for model registry
EntryClass = LingBotWorldTransformer3DModel
+1 -40
View File
@@ -814,8 +814,6 @@ class TransformerArgsPreprocessor:
batch_size = x.shape[0]
if context.device != x.device:
context = context.to(x.device)
if context.dtype != x.dtype:
context = context.to(x.dtype)
if attention_mask is not None and attention_mask.device != x.device:
attention_mask = attention_mask.to(x.device)
context = self.caption_projection(context)
@@ -1478,26 +1476,6 @@ class BasicAVTransformerBlock(torch.nn.Module):
self.norm_eps = norm_eps
def _register_fsdp_backward_hooks_on_output(self, vx, ax):
"""Register backward hooks on output tensors to trigger FSDP2 unshard.
FSDP2's module-level backward hooks don't fire when the module returns
dataclass outputs. We must register hooks directly on the output tensors.
"""
if not hasattr(self, 'unshard'):
return # Not wrapped by FSDP2
def make_unshard_hook():
def hook(grad):
self.unshard()
return grad
return hook
if vx is not None and vx.requires_grad:
vx.register_hook(make_unshard_hook())
if ax is not None and ax.requires_grad:
ax.register_hook(make_unshard_hook())
def get_ada_values(
self, scale_shift_table: torch.Tensor, batch_size: int, timestep: torch.Tensor, indices: slice
) -> tuple[torch.Tensor, ...]:
@@ -1534,7 +1512,6 @@ class BasicAVTransformerBlock(torch.nn.Module):
audio: TransformerArgs | None,
video_attention_mask: torch.Tensor | None = None,
audio_attention_mask: torch.Tensor | None = None,
skip_cross_modal_attn: bool = False,
) -> tuple[TransformerArgs | None, TransformerArgs | None]:
"""Forward pass for transformer block.
@@ -1543,8 +1520,6 @@ class BasicAVTransformerBlock(torch.nn.Module):
audio: Audio transformer args
video_attention_mask: SP padding attention mask for video [B, padded_seq_len]
audio_attention_mask: SP padding attention mask for audio [B, padded_seq_len]
skip_cross_modal_attn: If True, skip A2V and V2A cross-modal
attention (used for the modality-isolated CFG pass).
"""
vx = video.x if video is not None else None
ax = audio.x if audio is not None else None
@@ -1588,7 +1563,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
mask=audio.context_mask,
)
if (run_a2v or run_v2a) and not skip_cross_modal_attn:
if run_a2v or run_v2a:
vx_norm3 = torch.nn.functional.rms_norm(vx, (vx.shape[-1],), eps=self.norm_eps)
ax_norm3 = torch.nn.functional.rms_norm(ax, (ax.shape[-1],), eps=self.norm_eps)
@@ -1729,10 +1704,6 @@ class BasicAVTransformerBlock(torch.nn.Module):
f"audio_sum={audio_sum:.6f}"
)
# Register FSDP2 backward hooks on output tensors (module-level hooks don't
# fire for dataclass outputs, so we must hook the tensors directly)
self._register_fsdp_backward_hooks_on_output(vx, ax)
return (
replace(video, x=vx) if video is not None else None,
replace(audio, x=ax) if audio is not None else None,
@@ -2009,7 +1980,6 @@ class LTXModel(torch.nn.Module):
audio: TransformerArgs | None,
video_attention_mask: torch.Tensor | None = None,
audio_attention_mask: torch.Tensor | None = None,
skip_cross_modal_attn: bool = False,
) -> tuple[TransformerArgs | None, TransformerArgs | None]:
for block in self.transformer_blocks:
video, audio = block(
@@ -2017,7 +1987,6 @@ class LTXModel(torch.nn.Module):
audio=audio,
video_attention_mask=video_attention_mask,
audio_attention_mask=audio_attention_mask,
skip_cross_modal_attn=skip_cross_modal_attn,
)
return video, audio
@@ -2044,7 +2013,6 @@ class LTXModel(torch.nn.Module):
audio: Modality | None,
video_attention_mask: torch.Tensor | None = None,
audio_attention_mask: torch.Tensor | None = None,
skip_cross_modal_attn: bool = False,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
"""Forward pass through the LTX model.
@@ -2053,9 +2021,6 @@ class LTXModel(torch.nn.Module):
audio: Audio modality input
video_attention_mask: SP padding attention mask for video [B, padded_seq_len]
audio_attention_mask: SP padding attention mask for audio [B, padded_seq_len]
skip_cross_modal_attn: If True, skip A2V and V2A cross-modal
attention in all transformer blocks (modality-isolated
CFG pass).
"""
if os.getenv("LTX2_PIPELINE_DEBUG_LOG", "0") == "1":
_debug_block_log_line(
@@ -2079,7 +2044,6 @@ class LTXModel(torch.nn.Module):
audio_args,
video_attention_mask=video_attention_mask,
audio_attention_mask=audio_attention_mask,
skip_cross_modal_attn=skip_cross_modal_attn,
)
vx = (
@@ -2111,7 +2075,6 @@ class LTX2Transformer3DModel(CachableDiT):
param_names_mapping = LTX2VideoConfig().param_names_mapping
reverse_param_names_mapping = LTX2VideoConfig().reverse_param_names_mapping
lora_param_names_mapping = LTX2VideoConfig().lora_param_names_mapping
_fsdp_shard_conditions = LTX2VideoConfig()._fsdp_shard_conditions
def __init__(self, config: LTX2VideoConfig, hf_config: dict[str, Any]):
super().__init__(config=config, hf_config=hf_config)
@@ -2248,7 +2211,6 @@ class LTX2Transformer3DModel(CachableDiT):
audio_encoder_hidden_states: torch.Tensor | None = None,
audio_timestep: torch.Tensor | None = None,
audio_encoder_attention_mask: torch.Tensor | None = None,
skip_cross_modal_attn: bool = False,
**kwargs,
) -> torch.Tensor:
if isinstance(encoder_hidden_states, list):
@@ -2405,7 +2367,6 @@ class LTX2Transformer3DModel(CachableDiT):
audio=audio_modality,
video_attention_mask=video_attention_mask,
audio_attention_mask=audio_attention_mask,
skip_cross_modal_attn=skip_cross_modal_attn,
)
# Denoised prediction
-44
View File
@@ -1,44 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from typing import Any
from diffusers import SD3Transformer2DModel as _DiffusersSD3Transformer2DModel
from fastvideo.configs.models import DiTConfig
class SD3Transformer2DModel(_DiffusersSD3Transformer2DModel):
_fsdp_shard_conditions: list = []
_compile_conditions: list = []
param_names_mapping: dict[str, Any] = {}
reverse_param_names_mapping: dict[str, Any] = {}
lora_param_names_mapping: dict[str, Any] = {}
def __init__(self, config: DiTConfig, hf_config: dict[str, Any], **kwargs):
self.fastvideo_config = config
self.hf_config = hf_config
arch = config.arch_config
dual_layers = getattr(arch, "dual_attention_layers", ())
if isinstance(dual_layers, list):
dual_layers = tuple(dual_layers)
super().__init__(
sample_size=arch.sample_size,
patch_size=arch.patch_size,
in_channels=arch.in_channels,
num_layers=arch.num_layers,
attention_head_dim=arch.attention_head_dim,
num_attention_heads=arch.num_attention_heads,
joint_attention_dim=arch.joint_attention_dim,
caption_projection_dim=arch.caption_projection_dim,
pooled_projection_dim=arch.pooled_projection_dim,
out_channels=arch.out_channels,
pos_embed_max_size=arch.pos_embed_max_size,
dual_attention_layers=dual_layers,
qk_norm=arch.qk_norm,
**kwargs,
)
-37
View File
@@ -491,43 +491,6 @@ class CLIPTextModel(TextEncoder):
return loaded_params
class CLIPTextModelWithProjection(CLIPTextModel):
def __init__(self, config: CLIPTextConfig) -> None:
super().__init__(config)
self.text_projection = nn.Linear(
config.hidden_size, config.projection_dim, bias=False
)
def forward(
self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
outputs = super().forward(
input_ids=input_ids,
position_ids=position_ids,
attention_mask=attention_mask,
inputs_embeds=inputs_embeds,
output_hidden_states=output_hidden_states,
**kwargs,
)
pooled = outputs.pooler_output
if pooled is not None:
pooled = self.text_projection(pooled)
return BaseEncoderOutput(
last_hidden_state=outputs.last_hidden_state,
pooler_output=pooled,
hidden_states=outputs.hidden_states,
attentions=outputs.attentions,
attention_mask=outputs.attention_mask,
)
class CLIPVisionTransformer(nn.Module):
def __init__(
+2 -56
View File
@@ -3,11 +3,11 @@ from __future__ import annotations
from dataclasses import dataclass
import os
from typing import Any, Iterable
from typing import Iterable
import torch
from torch import nn
from transformers import AutoTokenizer, Gemma3ForConditionalGeneration
from transformers import Gemma3ForConditionalGeneration
from fastvideo.configs.models.encoders import BaseEncoderOutput, TextEncoderConfig
from fastvideo.models.encoders.base import TextEncoder
@@ -447,60 +447,6 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
return encoded, encoded_for_audio, attention_mask.squeeze(-1)
@torch.no_grad()
def preprocess_text_embeddings(
self,
prompts: str | list[str],
tokenizer: AutoTokenizer,
tokenizer_kwargs: dict[str, Any] | None = None,
padding_side: str | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Compute pre-connector text embeddings for LTX-2 training preprocessing."""
if isinstance(prompts, str):
prompts = [prompts]
model = self.gemma_model
kwargs: dict[str, Any] = {
"padding": "max_length",
"truncation": True,
"return_tensors": "pt",
}
if tokenizer_kwargs is not None:
kwargs.update(tokenizer_kwargs)
if "max_length" not in kwargs:
kwargs["max_length"] = self.config.arch_config.text_len
original_padding_side = tokenizer.padding_side
target_padding_side = padding_side or self.padding_side
tokenizer.padding_side = target_padding_side
try:
text_inputs = tokenizer(prompts, **kwargs)
finally:
tokenizer.padding_side = original_padding_side
input_ids = text_inputs["input_ids"].to(device=model.device)
attention_mask = text_inputs["attention_mask"].to(device=model.device)
outputs = model(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
return_dict=True,
)
prompt_embeds = self._run_feature_extractor(
outputs.hidden_states,
attention_mask,
padding_side=target_padding_side,
)
return prompt_embeds, attention_mask
def run_connectors(
self,
encoded_input: torch.Tensor,
attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Apply embedding connectors to precomputed Gemma features."""
return self._run_connectors(encoded_input, attention_mask)
def forward(
self,
input_ids: torch.Tensor | None,
-75
View File
@@ -1,75 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from typing import Any
import torch
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
from fastvideo.models.encoders.base import TextEncoder
class T5EncoderModel(TextEncoder):
supports_hf_from_pretrained: bool = True
def __init__(self, config: T5Config, hf_model: Any | None = None) -> None:
super().__init__(config)
self.hf_model = hf_model
@classmethod
def from_pretrained_local(
cls,
model_path: str,
config: T5Config,
*,
dtype: torch.dtype,
device: torch.device,
) -> "T5EncoderModel":
from transformers import T5EncoderModel as HFT5EncoderModel
kwargs = {
"local_files_only": True,
"low_cpu_mem_usage": True,
}
try:
hf = HFT5EncoderModel.from_pretrained(model_path, dtype=dtype, **kwargs)
except TypeError:
# Backward-compatible fallback for older Transformers versions.
hf = HFT5EncoderModel.from_pretrained(
model_path, torch_dtype=dtype, **kwargs
)
hf = hf.eval().to(device=device, dtype=dtype)
return cls(config=config, hf_model=hf).eval()
def forward(
self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
if self.hf_model is None:
raise RuntimeError(
"T5EncoderModel(HF) is not initialized. Use "
"`from_pretrained_local(...)` to construct a loaded instance."
)
out = self.hf_model(
input_ids=input_ids,
attention_mask=attention_mask,
inputs_embeds=inputs_embeds,
output_hidden_states=bool(output_hidden_states)
if output_hidden_states is not None
else False,
return_dict=True,
)
return BaseEncoderOutput(
last_hidden_state=out.last_hidden_state,
hidden_states=out.hidden_states if output_hidden_states else None,
attention_mask=attention_mask,
)
+60 -44
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,12 +85,12 @@ 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"),
"text_encoder_3": (TextEncoderLoader, "transformers"),
"tokenizer": (TokenizerLoader, "transformers"),
"tokenizer_2": (TokenizerLoader, "transformers"),
"tokenizer_3": (TokenizerLoader, "transformers"),
"image_processor": (ImageProcessorLoader, "transformers"),
"feature_extractor": (ImageProcessorLoader, "transformers"),
"image_encoder": (ImageEncoderLoader, "transformers"),
@@ -294,26 +295,23 @@ class TextEncoderLoader(ComponentLoader):
pass
logger.info("HF Model config: %s", model_config)
base = os.path.basename(os.path.normpath(model_path))
idx = 0
if base.startswith("text_encoder_"):
try:
idx = int(base.split("_")[-1]) - 1
except Exception:
idx = 0
encoder_configs = fastvideo_args.pipeline_config.text_encoder_configs
encoder_precisions = fastvideo_args.pipeline_config.text_encoder_precisions
if idx < 0 or idx >= len(encoder_configs):
raise IndexError(
f"text encoder index {idx} out of range for text_encoder_configs (len={len(encoder_configs)}), model_path={model_path}"
# @TODO(Wei): Better way to handle this?
try:
encoder_config = (
fastvideo_args.pipeline_config.text_encoder_configs[0]
)
encoder_config = encoder_configs[idx]
encoder_config.update_model_arch(model_config)
if idx < 0 or idx >= len(encoder_precisions):
raise IndexError(
f"text encoder index {idx} out of range for text_encoder_precisions (len={len(encoder_precisions)}), model_path={model_path}"
encoder_config.update_model_arch(model_config)
encoder_precision = (
fastvideo_args.pipeline_config.text_encoder_precisions[0]
)
except Exception:
encoder_config = (
fastvideo_args.pipeline_config.text_encoder_configs[1]
)
encoder_config.update_model_arch(model_config)
encoder_precision = (
fastvideo_args.pipeline_config.text_encoder_precisions[1]
)
encoder_precision = encoder_precisions[idx]
target_device = get_local_torch_device()
# TODO(will): add support for other dtypes
@@ -367,16 +365,7 @@ class TextEncoderLoader(ComponentLoader):
with target_device:
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
if getattr(model_cls, "supports_hf_from_pretrained", False):
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
model_path,
model_config, # type: ignore[arg-type]
dtype=PRECISION_TO_TYPE[dtype],
device=target_device,
)
return model.eval()
model = model_cls(model_config) # type: ignore
model: TextEncoder = model_cls(model_config) # type: ignore
weights_to_load = {name for name, _ in model.named_parameters()}
if (
@@ -673,10 +662,7 @@ class VAELoader(ComponentLoader):
break
loaded = remapped
# Diffusers-format AutoencoderKL checkpoints should match exactly; load
# strictly so missing/unexpected keys are surfaced early.
strict_load = class_name == "AutoencoderKL"
vae.load_state_dict(loaded, strict=strict_load)
vae.load_state_dict(loaded, strict=False)
return vae.eval()
@@ -744,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."""
@@ -913,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()
@@ -931,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()
+7 -10
View File
@@ -25,8 +25,6 @@ logger = init_logger(__name__)
_TEXT_TO_VIDEO_DIT_MODELS = {
"HunyuanVideoTransformer3DModel":
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"HunyuanGameCraftTransformer3DModel":
("dits", "hunyuangamecraft", "HunyuanGameCraftTransformer3DModel"),
"HunyuanVideo15Transformer3DModel":
("dits", "hunyuanvideo15", "HunyuanVideo15Transformer3DModel"),
"HYWorldTransformer3DModel":
@@ -39,8 +37,6 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
"LTX2Transformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
"SD3Transformer2DModel": ("dits", "sd3", "SD3Transformer2DModel"),
"LingBotWorldTransformer3DModel": ("dits", "lingbotworld", "LingBotWorldTransformer3DModel"),
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
@@ -53,11 +49,9 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
_TEXT_ENCODER_MODELS = {
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
"CLIPTextModelWithProjection":
("encoders", "clip", "CLIPTextModelWithProjection"),
"LlamaModel": ("encoders", "llama", "LlamaModel"),
"UMT5EncoderModel": ("encoders", "t5", "UMT5EncoderModel"),
"T5EncoderModel": ("encoders", "t5_hf", "T5EncoderModel"),
"T5EncoderModel": ("encoders", "t5", "T5EncoderModel"),
"STEP1TextEncoder": ("encoders", "stepllm", "STEP1TextEncoder"),
"BertModel": ("encoders", "clip", "CLIPTextModel"),
"Qwen2_5_VLTextModel": ("encoders", "qwen2_5", "Qwen2_5_VLTextModel"),
@@ -77,12 +71,10 @@ _IMAGE_ENCODER_MODELS: dict[str, tuple] = {
_VAE_MODELS = {
"AutoencoderKLHunyuanVideo":
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
"AutoencoderKLCausal3D": ("vaes", "gamecraftvae", "GameCraftVAE"),
"AutoencoderKLHYWorld": ("vaes", "hyworldvae", "AutoencoderKLHYWorld"),
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo"),
"AutoencoderKL": ("vaes", "autoencoder_kl", "AutoencoderKL"),
"CausalVideoAutoencoder": ("vaes", "ltx2vae", "LTX2CausalVideoAutoencoder"),
}
@@ -92,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",
@@ -119,6 +115,7 @@ _LEGACY_FAST_VIDEO_MODELS = {
**_IMAGE_ENCODER_MODELS,
**_VAE_MODELS,
**_AUDIO_MODELS,
**_UPSAMPLER_MODELS,
**_SCHEDULERS,
**_UPSAMPLERS,
}
@@ -454,4 +451,4 @@ ModelRegistry = _ModelRegistry({
)
for model_arch, (component_name, mod_relname,
cls_name) in _FAST_VIDEO_MODELS.items()
})
})
+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",
]
-55
View File
@@ -1,55 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from typing import Any
from diffusers import AutoencoderKL as _DiffusersAutoencoderKL
from fastvideo.configs.models import VAEConfig
class AutoencoderKL(_DiffusersAutoencoderKL):
def __init__(self, config: VAEConfig, **kwargs: Any) -> None:
self.fastvideo_config = config
arch = config.arch_config
down_block_types = arch.down_block_types
if isinstance(down_block_types, list):
down_block_types = tuple(down_block_types)
up_block_types = arch.up_block_types
if isinstance(up_block_types, list):
up_block_types = tuple(up_block_types)
block_out_channels = arch.block_out_channels
if isinstance(block_out_channels, list):
block_out_channels = tuple(block_out_channels)
latents_mean = arch.latents_mean
if isinstance(latents_mean, list):
latents_mean = tuple(latents_mean)
latents_std = arch.latents_std
if isinstance(latents_std, list):
latents_std = tuple(latents_std)
super().__init__(
in_channels=arch.in_channels,
out_channels=arch.out_channels,
down_block_types=down_block_types,
up_block_types=up_block_types,
block_out_channels=block_out_channels,
layers_per_block=arch.layers_per_block,
act_fn=arch.act_fn,
latent_channels=arch.latent_channels,
norm_num_groups=arch.norm_num_groups,
sample_size=arch.sample_size,
scaling_factor=arch.scaling_factor,
shift_factor=arch.shift_factor,
latents_mean=latents_mean,
latents_std=latents_std,
force_upcast=arch.force_upcast,
use_quant_conv=arch.use_quant_conv,
use_post_quant_conv=arch.use_post_quant_conv,
mid_block_add_attention=arch.mid_block_add_attention,
**kwargs,
)
-462
View File
@@ -1,462 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
GameCraft VAE - ported from official Hunyuan-GameCraft-1.0/hymm_sp/vae/.
Matches the official AutoencoderKLCausal3D structure exactly for weight loading.
"""
from dataclasses import dataclass
from typing import Optional, Tuple
import numpy as np
import torch
import torch.nn as nn
from fastvideo.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
from fastvideo.models.vaes.common import DiagonalGaussianDistribution
from fastvideo.models.vaes.gamecraftvae_blocks import (
CausalConv3d,
DownEncoderBlockCausal3D,
UpDecoderBlockCausal3D,
UNetMidBlockCausal3D,
)
@dataclass
class AutoencoderKLOutput:
"""Matches official AutoencoderKLOutput interface."""
latent_dist: DiagonalGaussianDistribution
@dataclass
class DecoderOutput:
"""Matches official DecoderOutput interface."""
sample: torch.Tensor
class EncoderCausal3D(nn.Module):
"""Encoder - ported from official vae.py. Structure matches for weight loading."""
def __init__(
self,
in_channels: int = 3,
out_channels: int = 16,
down_block_types: Tuple[str, ...] = ("DownEncoderBlockCausal3D",),
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
double_z: bool = True,
mid_block_add_attention: bool = True,
time_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
disable_causal: bool = False,
mid_block_causal_attn: bool = False,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(
in_channels, block_out_channels[0], kernel_size=3, stride=1, disable_causal=disable_causal
)
self.down_blocks = nn.ModuleList([])
output_channel = block_out_channels[0]
for i, _ in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial = int(np.log2(spatial_compression_ratio))
num_time = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial = bool(i < num_spatial)
add_time = bool(
i >= (len(block_out_channels) - 1 - num_time) and not is_final_block
)
elif time_compression_ratio == 8:
add_spatial = bool(i < num_spatial)
add_time = bool(i < num_time)
else:
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}")
downsample_stride_HW = (2, 2) if add_spatial else (1, 1)
downsample_stride_T = (2,) if add_time else (1,)
downsample_stride = tuple(downsample_stride_T + downsample_stride_HW)
down_block = DownEncoderBlockCausal3D(
in_channels=input_channel,
out_channels=output_channel,
num_layers=layers_per_block,
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
add_downsample=bool(add_spatial or add_time),
downsample_stride=downsample_stride,
downsample_padding=0,
disable_causal=disable_causal,
)
self.down_blocks.append(down_block)
self.mid_block = UNetMidBlockCausal3D(
in_channels=block_out_channels[-1],
temb_channels=None,
num_layers=1,
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
add_attention=mid_block_add_attention,
attention_head_dim=block_out_channels[-1],
disable_causal=disable_causal,
causal_attention=mid_block_causal_attn,
)
self.conv_norm_out = nn.GroupNorm(
num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6
)
self.conv_act = nn.SiLU()
conv_out_channels = 2 * out_channels if double_z else out_channels
self.conv_out = CausalConv3d(
block_out_channels[-1],
conv_out_channels,
kernel_size=3,
disable_causal=disable_causal,
)
def forward(self, sample: torch.Tensor) -> torch.Tensor:
sample = self.conv_in(sample)
for down_block in self.down_blocks:
sample = down_block(sample, scale=1.0)
sample = self.mid_block(sample, temb=None)
sample = self.conv_norm_out(sample)
sample = self.conv_act(sample)
sample = self.conv_out(sample)
return sample
class DecoderCausal3D(nn.Module):
"""Decoder - ported from official vae.py. Structure matches for weight loading."""
def __init__(
self,
in_channels: int = 16,
out_channels: int = 3,
up_block_types: Tuple[str, ...] = ("UpDecoderBlockCausal3D",),
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
mid_block_add_attention: bool = True,
time_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
disable_causal: bool = False,
mid_block_causal_attn: bool = False,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(
in_channels,
block_out_channels[-1],
kernel_size=3,
stride=1,
disable_causal=disable_causal,
)
self.mid_block = UNetMidBlockCausal3D(
in_channels=block_out_channels[-1],
temb_channels=None,
num_layers=1,
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
add_attention=mid_block_add_attention,
attention_head_dim=block_out_channels[-1],
disable_causal=disable_causal,
causal_attention=mid_block_causal_attn,
)
self.up_blocks = nn.ModuleList([])
reversed_channels = list(reversed(block_out_channels))
output_channel = reversed_channels[0]
for i, _ in enumerate(up_block_types):
prev_output_channel = output_channel
output_channel = reversed_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial = int(np.log2(spatial_compression_ratio))
num_time = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial = bool(i < num_spatial)
add_time = bool(
i >= len(block_out_channels) - 1 - num_time and not is_final_block
)
elif time_compression_ratio == 8:
add_spatial = bool(i >= len(block_out_channels) - num_spatial)
add_time = bool(i >= len(block_out_channels) - num_time)
else:
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}")
upsample_HW = (2, 2) if add_spatial else (1, 1)
upsample_T = (2,) if add_time else (1,)
upsample_scale_factor = tuple(upsample_T + upsample_HW)
up_block = UpDecoderBlockCausal3D(
in_channels=prev_output_channel,
out_channels=output_channel,
num_layers=layers_per_block + 1,
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
add_upsample=bool(add_spatial or add_time),
upsample_scale_factor=upsample_scale_factor,
temb_channels=None,
disable_causal=disable_causal,
)
self.up_blocks.append(up_block)
output_channel = output_channel
self.conv_norm_out = nn.GroupNorm(
num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6
)
self.conv_act = nn.SiLU()
self.conv_out = CausalConv3d(
block_out_channels[0], out_channels, kernel_size=3, disable_causal=disable_causal
)
def forward(
self,
sample: torch.Tensor,
latent_embeds: Optional[torch.Tensor] = None,
) -> torch.Tensor:
sample = self.conv_in(sample)
sample = self.mid_block(sample, temb=latent_embeds)
for up_block in self.up_blocks:
sample = up_block(sample, temb=latent_embeds, scale=1.0)
sample = self.conv_norm_out(sample)
sample = self.conv_act(sample)
sample = self.conv_out(sample)
return sample
class GameCraftVAE(nn.Module):
"""
GameCraft VAE - ported from official AutoencoderKLCausal3D.
Structure matches exactly for loading official weights.
"""
def __init__(self, config: GameCraftVAEConfig):
super().__init__()
self.config = config
arch = config.arch_config
time_ratio = getattr(arch, "time_compression_ratio", arch.temporal_compression_ratio)
self.encoder = EncoderCausal3D(
in_channels=arch.in_channels,
out_channels=arch.latent_channels,
down_block_types=tuple(arch.down_block_types),
block_out_channels=tuple(arch.block_out_channels),
layers_per_block=arch.layers_per_block,
norm_num_groups=arch.norm_num_groups,
act_fn=arch.act_fn,
double_z=True,
time_compression_ratio=time_ratio,
spatial_compression_ratio=arch.spatial_compression_ratio,
disable_causal=getattr(arch, "disable_causal_conv", False),
mid_block_add_attention=arch.mid_block_add_attention,
mid_block_causal_attn=getattr(arch, "mid_block_causal_attn", False),
)
self.decoder = DecoderCausal3D(
in_channels=arch.latent_channels,
out_channels=arch.out_channels,
up_block_types=tuple(arch.up_block_types),
block_out_channels=tuple(arch.block_out_channels),
layers_per_block=arch.layers_per_block,
norm_num_groups=arch.norm_num_groups,
act_fn=arch.act_fn,
time_compression_ratio=time_ratio,
spatial_compression_ratio=arch.spatial_compression_ratio,
disable_causal=getattr(arch, "disable_causal_conv", False),
mid_block_add_attention=arch.mid_block_add_attention,
mid_block_causal_attn=getattr(arch, "mid_block_causal_attn", False),
)
self.quant_conv = nn.Conv3d(
2 * arch.latent_channels, 2 * arch.latent_channels, kernel_size=1
)
self.post_quant_conv = nn.Conv3d(
arch.latent_channels, arch.latent_channels, kernel_size=1
)
# Scaling factor for latent normalization (required for decoding stage)
self.scaling_factor = arch.scaling_factor
# Tiling support - matches official GameCraft VAE settings
self._tiling_enabled = False
self.tile_overlap_factor = 0.25
# Temporal tiling params (for >64 output frames)
self.tile_sample_min_tsize = 64 # Minimum sample temporal size (video frames)
self.tile_latent_min_tsize = 16 # = 64 // 4 (time_compression_ratio)
# Spatial tiling params - use small tiles to reduce memory
self.tile_sample_min_size = 256 # Minimum spatial tile size in pixel space
self.tile_latent_min_size = 32 # = 256 // 8 (spatial_compression_ratio)
def encode(self, x: torch.Tensor) -> AutoencoderKLOutput:
"""Encode to latent distribution."""
h = self.encoder(x)
moments = self.quant_conv(h)
posterior = DiagonalGaussianDistribution(moments)
return AutoencoderKLOutput(latent_dist=posterior)
def decode(self, z: torch.Tensor) -> torch.Tensor:
"""Decode from latents.
Args:
z: Latent tensor [B, C, T, H, W]
Returns:
Decoded tensor [B, C, T_out, H_out, W_out]
"""
# Use tiled decode for memory efficiency when enabled
if self._tiling_enabled:
# Check if temporal tiling needed (>64 output frames)
if z.shape[2] > self.tile_latent_min_tsize:
return self._temporal_tiled_decode(z)
# Check if spatial tiling needed (large H or W)
if z.shape[-1] > self.tile_latent_min_size or z.shape[-2] > self.tile_latent_min_size:
return self._spatial_tiled_decode(z)
z = self.post_quant_conv(z)
dec = self.decoder(z, latent_embeds=None)
return dec
def _temporal_tiled_decode(self, z: torch.Tensor) -> torch.Tensor:
"""Decode latents in temporal tiles with overlapping and blending.
Based on official GameCraft temporal_tiled_decode implementation.
Only used when T > tile_latent_min_tsize (16).
"""
B, C, T, H, W = z.shape
# Use the pre-configured tiling parameters
overlap_size = int(self.tile_latent_min_tsize * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_tsize * self.tile_overlap_factor)
t_limit = self.tile_sample_min_tsize - blend_extent
row = []
for i in range(0, T, overlap_size):
tile = z[:, :, i : i + self.tile_latent_min_tsize + 1, :, :]
tile = self.post_quant_conv(tile)
decoded = self.decoder(tile, latent_embeds=None)
if i > 0:
decoded = decoded[:, :, 1:, :, :] # Skip first frame for non-first tiles
row.append(decoded)
# Blend overlapping regions
result_row = []
for i, tile in enumerate(row):
if i > 0:
tile = self._blend_t(row[i - 1], tile, blend_extent)
result_row.append(tile[:, :, :t_limit, :, :])
else:
result_row.append(tile[:, :, :t_limit+1, :, :])
return torch.cat(result_row, dim=2)
def _spatial_tiled_decode(self, z: torch.Tensor) -> torch.Tensor:
"""Decode latents in spatial tiles with overlapping and blending.
Based on official GameCraft spatial_tiled_decode implementation.
"""
overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
row_limit = self.tile_sample_min_size - blend_extent
# Split z into overlapping tiles and decode them separately
rows = []
for i in range(0, z.shape[-2], overlap_size):
row = []
for j in range(0, z.shape[-1], overlap_size):
tile = z[:, :, :, i: i + self.tile_latent_min_size, j: j + self.tile_latent_min_size]
tile = self.post_quant_conv(tile)
decoded = self.decoder(tile, latent_embeds=None)
row.append(decoded)
rows.append(row)
# Blend overlapping regions
result_rows = []
for i, row in enumerate(rows):
result_row = []
for j, tile in enumerate(row):
# Blend with above tile and left tile
if i > 0:
tile = self._blend_v(rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self._blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :, :row_limit, :row_limit])
result_rows.append(torch.cat(result_row, dim=-1))
return torch.cat(result_rows, dim=-2)
def _blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
"""Blend two tensors along temporal dimension."""
blend_extent = min(a.shape[-3], b.shape[-3], blend_extent)
if blend_extent == 0:
return b
a_region = a[..., -blend_extent:, :, :]
b_region = b[..., :blend_extent, :, :]
weights = torch.arange(blend_extent, device=a.device, dtype=a.dtype) / blend_extent
weights = weights.view(1, 1, blend_extent, 1, 1)
blended = a_region * (1 - weights) + b_region * weights
b[..., :blend_extent, :, :] = blended
return b
def _blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
"""Blend two tensors along vertical (height) dimension."""
blend_extent = min(a.shape[-2], b.shape[-2], blend_extent)
if blend_extent == 0:
return b
a_region = a[..., -blend_extent:, :]
b_region = b[..., :blend_extent, :]
weights = torch.arange(blend_extent, device=a.device, dtype=a.dtype) / blend_extent
weights = weights.view(1, 1, 1, blend_extent, 1)
blended = a_region * (1 - weights) + b_region * weights
b[..., :blend_extent, :] = blended
return b
def _blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
"""Blend two tensors along horizontal (width) dimension."""
blend_extent = min(a.shape[-1], b.shape[-1], blend_extent)
if blend_extent == 0:
return b
a_region = a[..., -blend_extent:]
b_region = b[..., :blend_extent]
weights = torch.arange(blend_extent, device=a.device, dtype=a.dtype) / blend_extent
weights = weights.view(1, 1, 1, 1, blend_extent)
blended = a_region * (1 - weights) + b_region * weights
b[..., :blend_extent] = blended
return b
def enable_tiling(self) -> None:
"""Enable tiling for large inputs."""
self._tiling_enabled = True
def disable_tiling(self) -> None:
"""Disable tiling."""
self._tiling_enabled = False
EntryClass = GameCraftVAE
@@ -1,446 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
GameCraft VAE building blocks - ported from official Hunyuan-GameCraft-1.0/hymm_sp/vae/unet_causal_3d_blocks.py.
Matches the official structure exactly for weight loading.
"""
from typing import Optional, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
def prepare_causal_attention_mask(
n_frame: int, n_hw: int, dtype, device, batch_size: Optional[int] = None
):
seq_len = n_frame * n_hw
mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device)
for i in range(seq_len):
i_frame = i // n_hw
mask[i, : (i_frame + 1) * n_hw] = 0
if batch_size is not None:
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
return mask
class CausalConv3d(nn.Module):
"""Causal 3D convolution - matches official structure (has .conv)."""
def __init__(
self,
chan_in: int,
chan_out: int,
kernel_size: Union[int, Tuple[int, int, int]] = 3,
stride: Union[int, Tuple[int, int, int]] = 1,
dilation: Union[int, Tuple[int, int, int]] = 1,
pad_mode: str = "replicate",
disable_causal: bool = False,
**kwargs,
):
super().__init__()
self.pad_mode = pad_mode
if isinstance(kernel_size, int):
k = kernel_size
else:
k = kernel_size[0]
if disable_causal:
padding = (k // 2, k // 2, k // 2, k // 2, k // 2, k // 2)
else:
padding = (k // 2, k // 2, k // 2, k // 2, k - 1, 0)
self.time_causal_padding = padding
self.conv = nn.Conv3d(
chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
return self.conv(x)
class DownsampleCausal3D(nn.Module):
"""Causal 3D downsampling - matches official (has .conv)."""
def __init__(
self,
channels: int,
out_channels: Optional[int] = None,
padding: int = 1,
stride: Union[int, Tuple[int, int, int]] = 2,
kernel_size: int = 3,
bias: bool = True,
disable_causal: bool = False,
):
super().__init__()
self.out_channels = out_channels or channels
self.conv = CausalConv3d(
channels,
self.out_channels,
kernel_size=kernel_size,
stride=stride,
padding=0,
disable_causal=disable_causal,
bias=bias,
)
def forward(self, hidden_states: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
return self.conv(hidden_states)
class UpsampleCausal3D(nn.Module):
"""Causal 3D upsampling - matches official (has .conv when use_conv=True)."""
def __init__(
self,
channels: int,
out_channels: Optional[int] = None,
kernel_size: int = 3,
upsample_factor: Tuple[int, int, int] = (2, 2, 2),
disable_causal: bool = False,
bias: bool = True,
):
super().__init__()
self.out_channels = out_channels or channels
self.upsample_factor = upsample_factor
self.disable_causal = disable_causal
self.conv = CausalConv3d(
channels,
self.out_channels,
kernel_size=kernel_size,
stride=1,
disable_causal=disable_causal,
bias=bias,
)
def forward(
self,
hidden_states: torch.Tensor,
output_size: Optional[int] = None,
scale: float = 1.0,
) -> torch.Tensor:
B, C, T, H, W = hidden_states.shape
dtype = hidden_states.dtype
if dtype == torch.bfloat16:
hidden_states = hidden_states.to(torch.float32)
if not self.disable_causal and T > 1:
first_h, other_h = hidden_states.split((1, T - 1), dim=2)
other_h = F.interpolate(
other_h, scale_factor=self.upsample_factor, mode="nearest"
)
first_h = F.interpolate(
first_h.squeeze(2), scale_factor=self.upsample_factor[1:], mode="nearest"
).unsqueeze(2)
hidden_states = torch.cat((first_h, other_h), dim=2)
else:
hidden_states = F.interpolate(
hidden_states, scale_factor=self.upsample_factor, mode="nearest"
)
if dtype == torch.bfloat16:
hidden_states = hidden_states.to(dtype)
return self.conv(hidden_states)
class GameCraftVAEAttention(nn.Module):
"""Attention block matching official diffusers Attention structure (group_norm, to_q, to_k, to_v, to_out)."""
def __init__(
self,
in_channels: int,
heads: int,
dim_head: int,
eps: float = 1e-6,
norm_num_groups: Optional[int] = 32,
bias: bool = True,
):
super().__init__()
self.heads = heads
self.dim_head = dim_head
inner_dim = heads * dim_head
self.group_norm = nn.GroupNorm(
norm_num_groups or in_channels, in_channels, eps=eps
)
self.to_q = nn.Linear(in_channels, inner_dim, bias=bias)
self.to_k = nn.Linear(in_channels, inner_dim, bias=bias)
self.to_v = nn.Linear(in_channels, inner_dim, bias=bias)
self.to_out = nn.Sequential(nn.Linear(inner_dim, in_channels, bias=bias))
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
residual = hidden_states
batch_size, seq_len, _ = hidden_states.shape
hidden_states = self.group_norm(
hidden_states.permute(0, 2, 1)
).permute(0, 2, 1)
q = self.to_q(hidden_states)
k = self.to_k(hidden_states)
v = self.to_v(hidden_states)
q = q.view(batch_size, seq_len, self.heads, self.dim_head).transpose(1, 2)
k = k.view(batch_size, seq_len, self.heads, self.dim_head).transpose(1, 2)
v = v.view(batch_size, seq_len, self.heads, self.dim_head).transpose(1, 2)
scale = self.dim_head**-0.5
attn = torch.matmul(q, k.transpose(-2, -1)) * scale
if attention_mask is not None:
attn = attn + attention_mask
# Official uses upcast_softmax=True: compute softmax in fp32 for numerical stability
attn = F.softmax(attn.float(), dim=-1).to(q.dtype)
hidden_states = torch.matmul(attn, v)
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, seq_len, -1)
hidden_states = self.to_out(hidden_states) + residual
return hidden_states
class ResnetBlockCausal3D(nn.Module):
"""ResNet block - matches official structure (conv1.conv, conv2.conv, norm1, norm2)."""
def __init__(
self,
in_channels: int,
out_channels: Optional[int] = None,
temb_channels: Optional[int] = None,
eps: float = 1e-6,
groups: int = 32,
dropout: float = 0.0,
non_linearity: str = "swish",
disable_causal: bool = False,
):
super().__init__()
out_channels = out_channels or in_channels
self.norm1 = nn.GroupNorm(groups, in_channels, eps=eps)
self.conv1 = CausalConv3d(in_channels, out_channels, 3, 1, disable_causal=disable_causal)
self.norm2 = nn.GroupNorm(groups, out_channels, eps=eps)
self.conv2 = CausalConv3d(out_channels, out_channels, 3, 1, disable_causal=disable_causal)
self.dropout = nn.Dropout(dropout)
self.conv_shortcut = (
CausalConv3d(in_channels, out_channels, 1, 1, disable_causal=disable_causal)
if in_channels != out_channels
else None
)
self.nonlinearity = getattr(F, non_linearity, F.silu)
def forward(
self,
x: torch.Tensor,
temb: Optional[torch.Tensor] = None,
scale: float = 1.0,
) -> torch.Tensor:
h = self.norm1(x)
h = self.nonlinearity(h)
h = self.conv1(h)
h = self.norm2(h)
h = self.nonlinearity(h)
h = self.dropout(h)
h = self.conv2(h)
if self.conv_shortcut is not None:
x = self.conv_shortcut(x)
return (x + h) / 1.0
class UNetMidBlockCausal3D(nn.Module):
"""Mid block with resnets and optional attention - matches official structure."""
def __init__(
self,
in_channels: int,
temb_channels: Optional[int] = None,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_act_fn: str = "swish",
resnet_groups: int = 32,
add_attention: bool = True,
attention_head_dim: int = 1,
disable_causal: bool = False,
causal_attention: bool = False,
):
super().__init__()
self.add_attention = add_attention
self.causal_attention = causal_attention
self.resnets = nn.ModuleList()
self.attentions = nn.ModuleList()
self.resnets.append(
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=in_channels,
temb_channels=temb_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=0.0,
non_linearity=resnet_act_fn,
disable_causal=disable_causal,
)
)
for _ in range(num_layers):
if add_attention:
self.attentions.append(
GameCraftVAEAttention(
in_channels=in_channels,
heads=in_channels // attention_head_dim,
dim_head=attention_head_dim,
eps=resnet_eps,
norm_num_groups=resnet_groups,
bias=True,
)
)
else:
self.attentions.append(None)
self.resnets.append(
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=in_channels,
temb_channels=temb_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=0.0,
non_linearity=resnet_act_fn,
disable_causal=disable_causal,
)
)
def forward(
self,
hidden_states: torch.Tensor,
temb: Optional[torch.Tensor] = None,
) -> torch.Tensor:
hidden_states = self.resnets[0](hidden_states, temb)
for attn, resnet in zip(self.attentions, self.resnets[1:]):
if attn is not None:
B, C, T, H, W = hidden_states.shape
hidden_states = rearrange(hidden_states, "b c f h w -> b (f h w) c")
if self.causal_attention:
mask = prepare_causal_attention_mask(
T, H * W, hidden_states.dtype, hidden_states.device, batch_size=B
)
else:
mask = None
hidden_states = attn(hidden_states, attention_mask=mask)
hidden_states = rearrange(
hidden_states, "b (f h w) c -> b c f h w", f=T, h=H, w=W
)
hidden_states = resnet(hidden_states, temb)
return hidden_states
class DownEncoderBlockCausal3D(nn.Module):
"""Encoder down block - matches official structure."""
def __init__(
self,
in_channels: int,
out_channels: int,
num_layers: int = 2,
resnet_eps: float = 1e-6,
resnet_act_fn: str = "swish",
resnet_groups: int = 32,
add_downsample: bool = True,
downsample_stride: Union[int, Tuple[int, int, int]] = 2,
downsample_padding: int = 0,
disable_causal: bool = False,
):
super().__init__()
self.resnets = nn.ModuleList()
for i in range(num_layers):
inc = in_channels if i == 0 else out_channels
self.resnets.append(
ResnetBlockCausal3D(
in_channels=inc,
out_channels=out_channels,
temb_channels=None,
eps=resnet_eps,
groups=resnet_groups,
dropout=0.0,
non_linearity=resnet_act_fn,
disable_causal=disable_causal,
)
)
self.downsamplers = None
if add_downsample:
self.downsamplers = nn.ModuleList([
DownsampleCausal3D(
out_channels,
out_channels=out_channels,
padding=downsample_padding,
stride=downsample_stride,
disable_causal=disable_causal,
)
])
def forward(
self, hidden_states: torch.Tensor, scale: float = 1.0
) -> torch.Tensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, temb=None, scale=scale)
if self.downsamplers is not None:
for ds in self.downsamplers:
hidden_states = ds(hidden_states, scale)
return hidden_states
class UpDecoderBlockCausal3D(nn.Module):
"""Decoder up block - matches official structure."""
def __init__(
self,
in_channels: int,
out_channels: int,
num_layers: int = 3,
resnet_eps: float = 1e-6,
resnet_act_fn: str = "swish",
resnet_groups: int = 32,
add_upsample: bool = True,
upsample_scale_factor: Tuple[int, int, int] = (2, 2, 2),
temb_channels: Optional[int] = None,
disable_causal: bool = False,
):
super().__init__()
self.resnets = nn.ModuleList()
for i in range(num_layers):
inc = in_channels if i == 0 else out_channels
self.resnets.append(
ResnetBlockCausal3D(
in_channels=inc,
out_channels=out_channels,
temb_channels=temb_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=0.0,
non_linearity=resnet_act_fn,
disable_causal=disable_causal,
)
)
self.upsamplers = None
if add_upsample:
self.upsamplers = nn.ModuleList([
UpsampleCausal3D(
out_channels,
out_channels=out_channels,
upsample_factor=upsample_scale_factor,
disable_causal=disable_causal,
)
])
def forward(
self,
hidden_states: torch.Tensor,
temb: Optional[torch.Tensor] = None,
scale: float = 1.0,
) -> torch.Tensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, temb=temb, scale=scale)
if self.upsamplers is not None:
for us in self.upsamplers:
hidden_states = us(hidden_states)
return hidden_states
+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(
@@ -1,2 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""HunyuanGameCraft pipeline implementations."""
@@ -1,101 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
HunyuanGameCraft video diffusion pipeline implementation.
This module implements the HunyuanGameCraft pipeline for camera/action-conditioned
video generation with the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (
ConditioningStage,
DecodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
)
from fastvideo.pipelines.stages.gamecraft_denoising import GameCraftDenoisingStage
logger = init_logger(__name__)
class HunyuanGameCraftPipeline(ComposedPipelineBase):
"""
Pipeline for HunyuanGameCraft video generation.
This pipeline supports:
- Text-to-video generation with camera/action conditioning
- Autoregressive generation with history frames
- 33-channel input (16 latent + 16 gt_latent + 1 mask)
- CameraNet for encoding Plücker coordinates
"""
_required_config_modules = [
"text_encoder",
"text_encoder_2",
"tokenizer",
"tokenizer_2",
"vae",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(
stage_name="input_validation_stage",
stage=InputValidationStage(),
)
self.add_stage(
stage_name="prompt_encoding_stage_primary",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2"),
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2"),
],
),
)
self.add_stage(
stage_name="conditioning_stage",
stage=ConditioningStage(),
)
self.add_stage(
stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")),
)
self.add_stage(
stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
),
)
self.add_stage(
stage_name="denoising_stage",
stage=GameCraftDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
self.add_stage(
stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")),
)
EntryClass = HunyuanGameCraftPipeline
@@ -1 +0,0 @@
@@ -1,16 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan video diffusion pipeline implementation.
This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.pipelines.basic.wan.wan_i2v_pipeline import WanImageToVideoPipeline
class LingBotWorldImageToVideoPipeline(WanImageToVideoPipeline):
pass
EntryClass = LingBotWorldImageToVideoPipeline
+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
@@ -1,118 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
from fastvideo.pipelines.stages.timestep_preparation import (
TimestepPreparationStage, )
from fastvideo.pipelines.stages.sd35_conditioning import (
SD35ConditioningStage,
SD35DecodingStage,
SD35DenoisingStage,
SD35LatentPreparationStage,
)
logger = init_logger(__name__)
class SD35Pipeline(ComposedPipelineBase):
"""Minimal SD3.5 Medium text-to-image pipeline (treat as num_frames=1)."""
_required_config_modules = [
"scheduler",
"transformer",
"vae",
"text_encoder",
"text_encoder_2",
"text_encoder_3",
"tokenizer",
"tokenizer_2",
"tokenizer_3",
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
te_cfgs = list(fastvideo_args.pipeline_config.text_encoder_configs)
if len(te_cfgs) >= 2:
for i in (0, 1):
te_cfgs[i].tokenizer_kwargs.setdefault("padding", "max_length")
te_cfgs[i].tokenizer_kwargs.setdefault("max_length", 77)
te_cfgs[i].tokenizer_kwargs.setdefault("truncation", True)
te_cfgs[i].tokenizer_kwargs.setdefault("return_tensors", "pt")
if len(te_cfgs) >= 3:
te_cfgs[2].tokenizer_kwargs["max_length"] = min(
int(te_cfgs[2].tokenizer_kwargs.get("max_length", 256)), 256)
te_cfgs[2].tokenizer_kwargs.setdefault("padding", "max_length")
te_cfgs[2].tokenizer_kwargs.setdefault("truncation", True)
te_cfgs[2].tokenizer_kwargs.setdefault("return_tensors", "pt")
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(
stage_name="text_encoding_stage",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2"),
self.get_module("text_encoder_3"),
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2"),
self.get_module("tokenizer_3"),
],
),
)
self.add_stage(
stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")),
)
self.add_stage(
stage_name="latent_preparation_stage",
stage=SD35LatentPreparationStage(
scheduler=self.get_module("scheduler"), ),
)
self.add_stage(
stage_name="sd35_conditioning_stage",
stage=SD35ConditioningStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2"),
self.get_module("text_encoder_3"),
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2"),
self.get_module("tokenizer_3"),
],
),
)
self.add_stage(
stage_name="denoising_stage",
stage=SD35DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
self.add_stage(
stage_name="decoding_stage",
stage=SD35DecodingStage(vae=self.get_module("vae"), ),
)
class StableDiffusion3Pipeline(SD35Pipeline):
"""Alias name to match SD3.5 diffusers `model_index.json` _class_name."""
EntryClass = [SD35Pipeline, StableDiffusion3Pipeline]
+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
@@ -136,17 +136,6 @@ class ForwardBatch:
# Camera control inputs (HYWorld)
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
# Camera/action control inputs (GameCraft)
camera_states: torch.Tensor | None = None # Plücker coordinates [B, T, 6, H, W]
gt_latents: torch.Tensor | None = None # Ground truth latents for conditioning [B, 16, T, H, W]
conditioning_mask: torch.Tensor | None = None # Mask for conditioning [B, 1, T, H, W]
camera_trajectory: str | None = None # Camera trajectory file/identifier
action_list: list[
str] | None = None # List of actions (e.g., ['forward', 'left'])
action_speed_list: list[float] | None = None # Speed for each action
# Camera control inputs (LingBotWorld)
c2ws_plucker_emb: torch.Tensor | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
# Latent dimensions
height_latents: list[int] | int | None = None
width_latents: list[int] | int | None = None
@@ -175,13 +164,6 @@ class ForwardBatch:
eta: float = 0.0
sigmas: list[float] | None = None
# LTX-2 multi-modal CFG parameters
ltx2_cfg_scale_video: float = 1.0
ltx2_cfg_scale_audio: float = 1.0
ltx2_modality_scale_video: float = 1.0
ltx2_modality_scale_audio: float = 1.0
ltx2_rescale_scale: float = 0.0
n_tokens: int | None = None
# Other parameters that may be needed by specific schedulers
@@ -248,14 +230,6 @@ class TrainingBatch:
noise_latents: torch.Tensor | None = None
encoder_hidden_states: torch.Tensor | None = None
encoder_attention_mask: torch.Tensor | None = None
# LTX related audio inputs
audio_latents: torch.Tensor | None = None
audio_noisy_model_input: torch.Tensor | None = None
audio_timesteps: torch.Tensor | None = None
audio_noise: torch.Tensor | None = None
audio_encoder_hidden_states: torch.Tensor | None = None
audio_encoder_attention_mask: torch.Tensor | None = None
conditioning_mask: torch.Tensor | None = None
# i2v
preprocessed_image: torch.Tensor | None = None
image_embeds: torch.Tensor | None = None
@@ -1,290 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""LTX-2 preprocessing pipeline for native FastVideo training data generation.
This module defines the LTX-2 preprocess pipeline used by FastVideo workflows
to build precomputed training artifacts from raw text/video datasets.
Usage:
- Entry is through preprocess workflows that register `PreprocessPipelineT2V`.
- Input samples should provide prompt text plus video metadata/loader fields
consumed by `TextTransformStage` and `VideoTransformStage`.
- Output artifacts are written by the shared preprocessing workflow into
`.precomputed/` (latents, conditions, and optional audio_latents).
Optional audio path:
- When audio preprocessing is enabled, this pipeline loads the native
LTX-2 audio encoder and stores per-sample audio latents in
`batch.extra["ltx2_audio_latents"]`.
"""
from __future__ import annotations
import glob
import os
from typing import Any, cast
import torch
import torchaudio
from safetensors.torch import load_file as safetensors_load_file
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.audio.ltx2_audio_processing import AudioProcessor
from fastvideo.models.audio.ltx2_audio_vae import LTX2AudioEncoder
from fastvideo.models.hf_transformer_utils import get_diffusers_config
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, PreprocessBatch
from fastvideo.pipelines.preprocess.preprocess_stages import (
TextTransformStage, VideoTransformStage)
from fastvideo.pipelines.stages import EncodingStage, PipelineStage
logger = init_logger(__name__)
class LTX2TextPrecomputeStage(PipelineStage):
"""Compute pre-connector Gemma embeddings for LTX-2 training."""
def __init__(
self,
text_encoder: torch.nn.Module,
tokenizer: Any,
preprocess_text_fn,
tokenizer_kwargs: dict[str, Any],
padding_side: str,
) -> None:
self.text_encoder = text_encoder
self.tokenizer = tokenizer
self.preprocess_text_fn = preprocess_text_fn
self.tokenizer_kwargs = tokenizer_kwargs
self.padding_side = padding_side
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
batch = cast(PreprocessBatch, batch)
assert isinstance(batch.prompt, list)
prompts = []
for prompt in batch.prompt:
if not isinstance(prompt, str):
prompt = str(prompt)
processed_prompt = self.preprocess_text_fn(prompt)
prompts.append(
processed_prompt if processed_prompt is not None else "")
prompt_embeds, prompt_attention_mask = (
self.text_encoder.preprocess_text_embeddings(
prompts=prompts,
tokenizer=self.tokenizer,
tokenizer_kwargs=self.tokenizer_kwargs,
padding_side=self.padding_side,
))
batch.prompt_embeds = [prompt_embeds]
batch.prompt_attention_mask = [prompt_attention_mask]
return batch
class LTX2AudioEncodingStage(PipelineStage):
"""Extract audio from input videos and encode into LTX-2 audio latents."""
def __init__(
self,
audio_encoder: torch.nn.Module,
audio_processor: AudioProcessor,
fallback_fps: int,
) -> None:
self.audio_encoder = audio_encoder.eval()
self.audio_processor = audio_processor
self.fallback_fps = fallback_fps
self.audio_dtype = next(audio_encoder.parameters()).dtype
self.audio_device = next(audio_encoder.parameters()).device
@staticmethod
def _extract_audio(
video_path: str,
target_duration: float,
) -> tuple[torch.Tensor, int] | None:
try:
waveform, sample_rate = torchaudio.load(video_path)
except Exception as e:
logger.error("Failed to load audio from %s: %s", video_path, e)
raise e
target_samples = int(target_duration * sample_rate)
if target_samples <= 0:
return None
current_samples = waveform.shape[-1]
if current_samples > target_samples:
waveform = waveform[..., :target_samples]
elif current_samples < target_samples:
padding = target_samples - current_samples
waveform = torch.nn.functional.pad(waveform, (0, padding))
return waveform, sample_rate
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
batch = cast(PreprocessBatch, batch)
assert isinstance(batch.video_loader, list)
assert isinstance(batch.num_frames, list)
assert isinstance(batch.fps, list)
audio_latents: list[torch.Tensor | None] = []
for idx, video_input in enumerate(batch.video_loader):
if not isinstance(video_input, str):
logger.warning(
"Skipping audio for sample %s: video loader is not a path string",
idx,
)
audio_latents.append(None)
continue
fps = float(batch.fps[idx]) if batch.fps[idx] else float(
self.fallback_fps)
if fps <= 0:
fps = float(self.fallback_fps)
target_duration = float(batch.num_frames[idx]) / fps
audio_data = self._extract_audio(video_input, target_duration)
if audio_data is None:
audio_latents.append(None)
continue
waveform, sample_rate = audio_data
waveform = waveform.unsqueeze(0).to(device=self.audio_device,
dtype=self.audio_dtype)
mel = self.audio_processor.waveform_to_mel(
waveform,
waveform_sample_rate=sample_rate).to(device=self.audio_device,
dtype=self.audio_dtype)
latents = self.audio_encoder(mel).squeeze(0).detach().cpu()
audio_latents.append(latents)
batch.extra["ltx2_audio_latents"] = audio_latents
return batch
class PreprocessPipelineT2V(ComposedPipelineBase):
"""Native LTX-2 preprocessing pipeline (text/video with optional audio)."""
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
tokenizer = self.get_module("tokenizer")
if tokenizer is not None:
tokenizer.padding_side = "left"
if tokenizer.pad_token is None and tokenizer.eos_token is not None:
tokenizer.pad_token = tokenizer.eos_token
def _load_ltx2_audio_encoder(
self) -> tuple[torch.nn.Module, AudioProcessor]:
audio_vae_path = os.path.join(self.model_path, "audio_vae")
if not os.path.isdir(audio_vae_path):
raise FileNotFoundError(
f"Expected audio_vae directory for LTX-2 audio preprocessing: {audio_vae_path}"
)
config = get_diffusers_config(model=audio_vae_path)
audio_encoder = LTX2AudioEncoder(config).to(
device=get_local_torch_device(),
dtype=torch.float32,
)
safetensors_list = glob.glob(
os.path.join(audio_vae_path, "*.safetensors"))
loaded: dict[str, torch.Tensor] = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
encoder_state = {}
for name, tensor in loaded.items():
if name.startswith("encoder."):
encoder_state[name.replace("encoder.", "")] = tensor
elif name.startswith("per_channel_statistics."):
encoder_state[name] = tensor
target_module = getattr(audio_encoder, "model", audio_encoder)
missing, unexpected = target_module.load_state_dict(encoder_state,
strict=False)
if missing:
logger.warning("Missing LTX-2 audio encoder keys: %s", missing[:8])
if unexpected:
logger.warning("Unexpected LTX-2 audio encoder keys: %s",
unexpected[:8])
target_module.eval()
audio_processor = AudioProcessor(
sample_rate=target_module.sample_rate,
mel_bins=target_module.mel_bins,
mel_hop_length=target_module.mel_hop_length,
n_fft=target_module.n_fft,
).to(next(target_module.parameters()).device)
return target_module, audio_processor
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
assert fastvideo_args.preprocess_config is not None
preprocess_cfg = fastvideo_args.preprocess_config
self.add_stage(
stage_name="text_transform_stage",
stage=TextTransformStage(
cfg_uncondition_drop_rate=preprocess_cfg.training_cfg_rate,
seed=preprocess_cfg.seed,
),
)
text_encoder = self.get_module("text_encoder")
tokenizer = self.get_module("tokenizer")
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[0]
tokenizer_kwargs = dict(encoder_config.tokenizer_kwargs)
if "max_length" not in tokenizer_kwargs:
tokenizer_kwargs["max_length"] = encoder_config.arch_config.text_len
self.add_stage(
stage_name="prompt_precompute_stage",
stage=LTX2TextPrecomputeStage(
text_encoder=text_encoder,
tokenizer=tokenizer,
preprocess_text_fn=fastvideo_args.pipeline_config.
preprocess_text_funcs[0],
tokenizer_kwargs=tokenizer_kwargs,
padding_side=encoder_config.arch_config.padding_side,
),
)
self.add_stage(
stage_name="video_transform_stage",
stage=VideoTransformStage(
train_fps=preprocess_cfg.train_fps,
num_frames=preprocess_cfg.num_frames,
max_height=preprocess_cfg.max_height,
max_width=preprocess_cfg.max_width,
do_temporal_sample=preprocess_cfg.do_temporal_sample,
),
)
if preprocess_cfg.with_audio:
audio_encoder, audio_processor = self._load_ltx2_audio_encoder()
self.add_stage(
stage_name="audio_encoding_stage",
stage=LTX2AudioEncodingStage(
audio_encoder=audio_encoder,
audio_processor=audio_processor,
fallback_fps=preprocess_cfg.train_fps,
),
)
self.add_stage(
stage_name="video_encoding_stage",
stage=EncodingStage(vae=self.get_module("vae")),
)
EntryClass = PreprocessPipelineT2V

Some files were not shown because too many files have changed in this diff Show More