Compare commits

..
Author SHA1 Message Date
SolitaryThinker a8dddfaa16 missing file 2026-02-10 08:33:00 +00:00
Will Lin 7210c68f1b lint 2026-02-10 00:30:03 -08:00
SolitaryThinker 5602dc1bad revert 2026-02-10 08:22:05 +00:00
SolitaryThinker ac4bc4ab84 uipdate 2026-02-10 07:58:43 +00:00
Matthew Noto bee27f9f74 Merge branch 'main' into ltx-base 2026-02-09 17:51:37 -08:00
Will Lin dff0ea401a update 2026-02-08 00:22:21 -08:00
Davids048 becd379f58 Split LTX2 mappings and add registry coverage.
Detail:

- Merged LTX2 sampling behavior into the global registry and removed the ambiguous local converted-path default.
- Replaced single LTX2 sampling mapping with explicit model-ID mappings:
    - Lightricks/LTX-2 -> LTX2BaseSamplingParam
    - FastVideo/LTX2-base -> LTX2BaseSamplingParam
    - FastVideo/LTX2-Distilled-Diffusers -> LTX2DistilledSamplingParam
- Kept LTX2T2VConfig as the pipeline config for all explicitly mapped LTX2 IDs.
- Removed implicit mapping for converted/ltx2_diffusers to avoid guessing base vs distilled for user-local
  conversions.
- Added focused local tests at tests/local_tests/test_ltx2_registry.py for:
    - exact base/distilled sampling resolution,
    - pipeline config resolution,
    - no fallback behavior for ambiguous local converted paths.

Assumptions:

- Canonical rename/ID intent:
    - “base” names map to LTX2BaseSamplingParam.
    - “Distilled” names map to LTX2DistilledSamplingParam.
- converted/ltx2_diffusers is intentionally ambiguous across users and must not be auto-assigned.
- Unknown/non-canonical names containing “LTX”/“LTX2” (but not matching explicit registered IDs) should not auto-
  resolve to base or distilled.
    - Sampling resolver returns None (caller falls back to generic defaults/user overrides).
    - Pipeline config lookup raises a “No match found” error.

Notes:

- This change prioritizes explicitness over convenience: only predetermined, canonical model IDs get LTX2-specific
  defaults; everything else requires user intent.
2026-02-05 20:28:46 -08:00
Davids048 1ed7d7e1b0 Add gemma tokenizer to LTX2 conversion script.
- Also clean up for PR.
2026-02-05 18:53:22 -08:00
Davids048 d9fabcc5ef Add some annotations. 2026-02-05 18:53:22 -08:00
Davids048 9db48498de Update quality test script. 2026-02-05 18:53:22 -08:00
Davids048 1cd7038315 Add LTX2 base model. 2026-02-05 18:53:12 -08:00
140 changed files with 442 additions and 9882 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/.*|
+1 -1
View File
@@ -8,7 +8,7 @@
- `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/`.
- Static assets: `assets/`, `images/`, `videos/`, and `comfyui/assets/`.
## Build, Test, and Development Commands
- `uv pip install -e .[dev]`: editable install with lint/test extras.
+1 -1
View File
@@ -3,7 +3,7 @@
</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> |
| <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://ibb.co/sv3MMKyv" target="_blank"> <b> WeChat </b> </a> |
</p>
**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()
+2 -8
View File
@@ -20,10 +20,10 @@ def main() -> None:
# Uses FastVideo default sampling settings for LTX2 base.
generator = VideoGenerator.from_pretrained(
"Davids048/LTX2-Base-Diffusers",
num_gpus=8,
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_base_t2v_1088_1920_1.1.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
@@ -31,12 +31,6 @@ def main() -> None:
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()
-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()
+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"
+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"
-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
-7
View File
@@ -9,7 +9,6 @@ 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``.
"""
seed: int = 10
@@ -38,12 +37,6 @@ class LTX2BaseSamplingParam(SamplingParam):
"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
-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
+23 -51
View File
@@ -223,82 +223,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
@@ -444,25 +426,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
+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
-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
-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 -13
View File
@@ -1534,7 +1534,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 +1542,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 +1585,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)
@@ -2009,7 +2006,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 +2013,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 +2039,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 +2047,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 +2070,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 = (
@@ -2248,7 +2238,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 +2394,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__(
-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,
)
+18 -35
View File
@@ -86,10 +86,8 @@ class ComponentLoader(ABC):
"vocoder": (VocoderLoader, "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 +292,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 +362,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 +659,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()
@@ -1022,4 +1005,4 @@ class PipelineComponentLoader:
)
# Load the module
return loader.load(component_model_path, fastvideo_args)
return loader.load(component_model_path, fastvideo_args)
+2 -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"),
}
@@ -454,4 +446,4 @@ ModelRegistry = _ModelRegistry({
)
for model_arch, (component_name, mod_relname,
cls_name) in _FAST_VIDEO_MODELS.items()
})
})
-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
@@ -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
@@ -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]
@@ -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
-5
View File
@@ -20,8 +20,6 @@ from fastvideo.pipelines.stages.image_encoding import (
ImageEncodingStage, MatrixGameImageEncodingStage, RefImageEncodingStage,
ImageVAEEncodingStage, VideoVAEEncodingStage, Hy15ImageEncodingStage,
HYWorldImageEncodingStage)
from fastvideo.pipelines.stages.gamecraft_image_encoding import (
GameCraftImageVAEEncodingStage)
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.latent_preparation import (
Cosmos25LatentPreparationStage, CosmosLatentPreparationStage,
@@ -35,7 +33,6 @@ from fastvideo.pipelines.stages.ltx2_text_encoding import LTX2TextEncodingStage
from fastvideo.pipelines.stages.matrixgame_denoising import (
MatrixGameCausalDenoisingStage)
from fastvideo.pipelines.stages.hyworld_denoising import HYWorldDenoisingStage
from fastvideo.pipelines.stages.gamecraft_denoising import GameCraftDenoisingStage
from fastvideo.pipelines.stages.stepvideo_encoding import (
StepvideoPromptEncodingStage)
from fastvideo.pipelines.stages.text_encoding import (Cosmos25TextEncodingStage,
@@ -67,7 +64,6 @@ __all__ = [
"CausalDMDDenosingStage",
"MatrixGameCausalDenoisingStage",
"HYWorldDenoisingStage",
"GameCraftDenoisingStage",
"CosmosDenoisingStage",
"Cosmos25DenoisingStage",
"Cosmos25T2WDenoisingStage",
@@ -85,7 +81,6 @@ __all__ = [
"RefImageEncodingStage",
"ImageVAEEncodingStage",
"VideoVAEEncodingStage",
"GameCraftImageVAEEncodingStage",
"TextEncodingStage",
"Cosmos25TextEncodingStage",
"StepvideoPromptEncodingStage",
-10
View File
@@ -173,14 +173,6 @@ class DenoisingStage(PipelineStage):
{
"mouse_cond": batch.mouse_cond,
"keyboard_cond": batch.keyboard_cond,
"c2ws_plucker_emb": batch.c2ws_plucker_emb,
},
)
camera_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"camera_states": batch.camera_states,
},
)
@@ -435,7 +427,6 @@ class DenoisingStage(PipelineStage):
**image_kwargs,
**pos_cond_kwargs,
**action_kwargs,
**camera_kwargs,
**timesteps_r_kwarg,
)
@@ -454,7 +445,6 @@ class DenoisingStage(PipelineStage):
**image_kwargs,
**neg_cond_kwargs,
**action_kwargs,
**camera_kwargs,
**timesteps_r_kwarg,
)
@@ -1,352 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
GameCraft denoising stage for camera/action-conditioned video generation.
This stage implements the denoising loop for HunyuanGameCraft, which generates
game-like videos with camera and action conditioning via:
1. CameraNet - Encodes Plücker coordinates into features added to image embeddings
2. Concatenated input - 33 channels (16 latent + 16 gt_latent + 1 mask)
3. Mask-based conditioning for autoregressive generation
"""
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
class GameCraftDenoisingStage(DenoisingStage):
"""
Denoising stage for HunyuanGameCraft with camera/action conditioning.
This stage handles:
- Camera state encoding via CameraNet (Plücker coordinates)
- Concatenation of latents with gt_latents and mask (33 channels)
- Flow matching denoising with camera conditioning
- Support for autoregressive generation with history frames
"""
def __init__(
self,
transformer,
scheduler,
pipeline=None,
transformer_2=None,
vae=None,
) -> None:
super().__init__(transformer, scheduler, pipeline, transformer_2, vae)
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Run the denoising loop with camera/action conditioning.
Args:
batch: The current batch information. Must contain:
- latents: Noise latents [B, 16, T, H, W]
- camera_states: Plücker coordinates [B, T_video, 6, H_video, W_video]
- gt_latents (optional): Ground truth latents for conditioning [B, 16, T, H, W]
- conditioning_mask (optional): Mask for conditioning [B, 1, T, H, W]
fastvideo_args: The inference arguments.
Returns:
The batch with denoised latents.
"""
pipeline = self.pipeline() if self.pipeline else None
if not fastvideo_args.model_loaded["transformer"]:
loader = TransformerLoader()
self.transformer = loader.load(
fastvideo_args.model_paths["transformer"], fastvideo_args)
if pipeline:
pipeline.add_module("transformer", self.transformer)
fastvideo_args.model_loaded["transformer"] = True
# Extract GameCraft-specific parameters
camera_states = getattr(batch, "camera_states", None)
if camera_states is None:
camera_states = batch.extra.get("camera_states", None)
gt_latents = getattr(batch, "gt_latents", None)
if gt_latents is None:
gt_latents = batch.extra.get("gt_latents", None)
conditioning_mask = getattr(batch, "conditioning_mask", None)
if conditioning_mask is None:
conditioning_mask = batch.extra.get("conditioning_mask", None)
# Prepare extra step kwargs for scheduler
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta,
},
)
# Setup precision and autocast settings
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
# Get timesteps
timesteps = batch.timesteps
if timesteps is None:
raise ValueError("Timesteps must be provided")
num_inference_steps = batch.num_inference_steps
num_warmup_steps = len(
timesteps) - num_inference_steps * self.scheduler.order
# Prepare image embeddings for I2V generation (if any)
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert not torch.isnan(
image_embeds[0]).any(), "image_embeds contains nan"
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
image_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{"encoder_hidden_states_image": image_embeds},
)
pos_cond_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_2": batch.clip_embedding_pos,
"encoder_attention_mask": batch.prompt_attention_mask,
},
)
neg_cond_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"encoder_hidden_states_2": batch.clip_embedding_neg,
"encoder_attention_mask": batch.negative_attention_mask,
},
)
# Get latents and embeddings
latents = batch.latents
prompt_embeds = batch.prompt_embeds
assert not torch.isnan(
prompt_embeds[0]).any(), "prompt_embeds contains nan"
if batch.do_classifier_free_guidance:
neg_prompt_embeds = batch.negative_prompt_embeds
assert neg_prompt_embeds is not None
assert not torch.isnan(
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
# Prepare gt_latents and mask for concatenation
# If not provided, use zeros (for unconditional generation)
if gt_latents is None:
gt_latents = torch.zeros_like(latents)
else:
gt_latents = gt_latents.to(target_dtype)
if conditioning_mask is None:
# Default mask: all zeros (generate everything)
conditioning_mask = torch.zeros(
latents.shape[0],
1,
*latents.shape[2:],
device=latents.device,
dtype=target_dtype,
)
else:
conditioning_mask = conditioning_mask.to(target_dtype)
# Move camera states to device if provided
if camera_states is not None:
camera_states = camera_states.to(device=latents.device,
dtype=target_dtype)
# Debug logging
logger.debug("[GameCraft DEBUG] latents shape: %s, min/max: %.4f/%.4f",
latents.shape, latents.min(), latents.max())
logger.debug("[GameCraft DEBUG] camera_states: %s",
camera_states.shape if camera_states is not None else None)
logger.debug("[GameCraft DEBUG] prompt_embeds[0] shape: %s",
prompt_embeds[0].shape)
# Initialize lists for trajectory
trajectory_timesteps: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = []
# For I2V: extract clean reference latents to inject at every step.
# Official GameCraft replaces conditioned frames of the denoising latents
# with the clean reference at EVERY timestep (not just the first).
# This keeps the conditioned frames noise-free so the model has a strong
# reference signal throughout the denoising process.
#
# ref_latent_for_injection: [B, 16, 1, H, W] or None
ref_latent_for_injection = getattr(batch, "_ref_latent_for_injection",
None)
if (ref_latent_for_injection is None and gt_latents is not None
and conditioning_mask.sum() > 0
and gt_latents[:, :, 0].abs().sum() > 0):
# Extract the clean reference from gt_latents first frame
# gt_latents[:, :, 0] should have the VAE-encoded reference image
ref_latent_for_injection = gt_latents[:, :, 0:1].clone(
) # [B, 16, 1, H, W]
logger.info(
"[GameCraft I2V] Will inject ref latent at conditioned frames each step. "
"ref mean=%.4f",
ref_latent_for_injection.abs().mean())
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
# Skip if interrupted
if hasattr(self, "interrupt") and self.interrupt:
continue
current_model = self.transformer
current_guidance_scale = batch.guidance_scale
# I2V: replace conditioned frames with clean reference latent
# (matches official: latents[:,:,0,:,:] = last_latents[:,:,-1,:,:])
if ref_latent_for_injection is not None:
# conditioning_mask shape: [B, 1, T, H, W]
# Where mask == 1, replace latents with the clean ref
cond_frames = conditioning_mask[0, 0, :, 0, 0] > 0.5 # [T]
for t_idx in range(cond_frames.shape[0]):
if cond_frames[t_idx]:
latents[:, :,
t_idx, :, :] = ref_latent_for_injection[:, :,
0, :, :]
# Prepare model input: concatenate latents, gt_latents, and mask
# [B, 33, T, H, W] = [B, 16, T, H, W] + [B, 16, T, H, W] + [B, 1, T, H, W]
latent_model_input = torch.cat(
[latents.to(target_dtype), gt_latents, conditioning_mask],
dim=1,
)
assert not torch.isnan(
latent_model_input).any(), "latent_model_input contains nan"
t_expand = t.repeat(latent_model_input.shape[0])
latent_model_input = self.scheduler.scale_model_input(
latent_model_input, t)
# Official GameCraft does NOT use embedded guidance (guidance=None)
# It uses standard CFG with guidance_scale instead
guidance_expand = None
# Run transformer with camera conditioning
with torch.autocast(
device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled,
):
batch.is_cfg_negative = False
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
):
noise_pred = current_model(
latent_model_input,
prompt_embeds,
t_expand,
camera_states=camera_states,
guidance=guidance_expand,
**image_kwargs,
**pos_cond_kwargs,
)
# Debug: log first step output
if i == 0:
logger.info(
"[GameCraft DEBUG] Step 0 noise_pred: shape=%s, min/max=%.4f/%.4f",
noise_pred.shape, noise_pred.min(),
noise_pred.max())
# Classifier-free guidance
if batch.do_classifier_free_guidance:
batch.is_cfg_negative = True
with set_forward_context(
current_timestep=i,
attn_metadata=None,
forward_batch=batch,
):
noise_pred_uncond = current_model(
latent_model_input,
neg_prompt_embeds,
t_expand,
camera_states=camera_states,
guidance=guidance_expand,
**image_kwargs,
**neg_cond_kwargs,
)
noise_pred = noise_pred_uncond + current_guidance_scale * (
noise_pred - noise_pred_uncond)
# Compute the previous noisy sample
latents = self.scheduler.step(
noise_pred,
t,
latents,
**extra_step_kwargs,
return_dict=False,
)[0]
# Store trajectory if requested
if batch.return_trajectory_latents:
trajectory_timesteps.append(t.clone())
trajectory_latents.append(latents.clone())
# Update progress bar
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
(i + 1) % self.scheduler.order == 0):
progress_bar.update()
# Debug: log final latents
logger.info(
"[GameCraft DEBUG] Final latents: shape=%s, min/max=%.4f/%.4f",
latents.shape, latents.min(), latents.max())
# Store final latents and trajectory
batch.latents = latents
if batch.return_trajectory_latents:
batch.trajectory_timesteps = trajectory_timesteps
batch.trajectory_latents = torch.stack(trajectory_latents, dim=0)
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify that required inputs are present."""
result = VerificationResult()
result.add_check("timesteps", batch.timesteps,
[V.is_tensor, V.min_dims(1)])
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
result.add_check("num_inference_steps", batch.num_inference_steps,
V.positive_int)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify that outputs are properly set."""
result = VerificationResult()
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
return result
@@ -1,183 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
GameCraft image-to-video encoding stage.
Encodes a reference image into gt_latents and conditioning_mask for
HunyuanGameCraft I2V generation. For T2V this stage is a no-op.
"""
import torch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.utils import PRECISION_TO_TYPE
logger = init_logger(__name__)
class GameCraftImageVAEEncodingStage(PipelineStage):
"""
Stage for encoding a reference image into gt_latents and conditioning_mask
for HunyuanGameCraft image-to-video generation.
Official GameCraft I2V flow:
1. VAE-encode the reference image -> [B, 16, 1, H_lat, W_lat]
2. Scale by VAE scaling_factor (0.476986)
3. Repeat to all temporal frames
4. Zero out non-conditioned frames (first frame only for short videos,
first half for longer autoregressive generation)
5. Build a binary mask (1 = conditioned, 0 = generate)
6. Store gt_latents and conditioning_mask on the batch for the denoising stage
If no image is provided (T2V mode), this stage is a no-op; the denoising
stage already falls back to zero gt_latents and zero mask.
"""
def __init__(self, vae) -> None:
super().__init__()
self.vae = vae
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""Encode reference image for I2V, or skip for T2V."""
if batch.pil_image is None:
# T2V mode: nothing to do; denoising stage handles the fallback
return batch
device = get_local_torch_device()
image = batch.pil_image
# ------------------------------------------------------------------
# 1. Preprocess image to tensor [B, 3, 1, H, W] (values in [-1, 1])
# ------------------------------------------------------------------
from torchvision import transforms
target_height, target_width = batch.height, batch.width
if isinstance(image, torch.Tensor):
# Already a tensor (e.g. from InputValidationStage causal path)
if image.dim() == 5:
# [B, C, F, H, W] -> take first frame
ref_pixel = image[:, :, :1]
elif image.dim() == 4:
ref_pixel = image.unsqueeze(
2) # [B, C, H, W] -> [B, C, 1, H, W]
else:
raise ValueError(f"Unexpected image tensor dims: {image.dim()}")
ref_pixel = ref_pixel.to(device=device, dtype=torch.float32)
else:
# PIL Image – resize, center-crop, normalize to [-1, 1]
from PIL import Image as PILImage
if not isinstance(image, PILImage.Image):
import numpy as np
if isinstance(image, np.ndarray):
image = PILImage.fromarray(image)
original_w, original_h = image.size
scale = max(target_width / original_w, target_height / original_h)
resize_w = int(round(original_w * scale))
resize_h = int(round(original_h * scale))
ref_transform = transforms.Compose([
transforms.Resize(
(resize_h, resize_w),
interpolation=transforms.InterpolationMode.LANCZOS,
),
transforms.CenterCrop((target_height, target_width)),
transforms.ToTensor(),
transforms.Normalize([0.5], [0.5]),
])
ref_pixel = ref_transform(image) # [3, H, W]
ref_pixel = ref_pixel.unsqueeze(0).unsqueeze(2) # [1, 3, 1, H, W]
ref_pixel = ref_pixel.to(device=device, dtype=torch.float32)
# ------------------------------------------------------------------
# 2. VAE-encode
# ------------------------------------------------------------------
self.vae = self.vae.to(device)
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
if not vae_autocast_enabled:
ref_pixel = ref_pixel.to(vae_dtype)
encoder_output = self.vae.encode(ref_pixel)
# Sample from the distribution (official uses .sample())
ref_latents = encoder_output.latent_dist.sample().to(
dtype=torch.float32)
# Scale by VAE scaling factor (0.476986 for GameCraft)
ref_latents.mul_(self.vae.config.scaling_factor)
# ref_latents: [B, 16, 1, H_lat, W_lat]
# ------------------------------------------------------------------
# 3. Build gt_latents + conditioning_mask
# ------------------------------------------------------------------
# Get latent temporal dimension from batch latents (set by LatentPreparationStage)
latent_frames = batch.latents.shape[
2] if batch.latents is not None else ((batch.num_frames - 1) // 4 +
1)
# Repeat to all frames
gt_latents = ref_latents.repeat(1, 1, latent_frames, 1, 1)
# [B, 16, T_lat, H_lat, W_lat]
# Mask construction following official GameCraft logic:
# - Short videos (latent_frames <= 10): first frame conditioned
# - Longer videos: first half conditioned (autoregressive)
mask = torch.ones(
gt_latents.shape[0],
1,
gt_latents.shape[2],
gt_latents.shape[3],
gt_latents.shape[4],
device=gt_latents.device,
dtype=gt_latents.dtype,
)
if latent_frames <= 10:
# I2V: only first frame is conditioned
gt_latents[:, :, 1:, :, :] = 0.0
mask[:, :, 1:, :, :] = 0.0
else:
# Autoregressive: first half conditioned
half = latent_frames // 2
gt_latents[:, :, half:, :, :] = 0.0
mask[:, :, half:, :, :] = 0.0
batch.gt_latents = gt_latents.to(device=device)
batch.conditioning_mask = mask.to(device=device)
# Offload
if fastvideo_args.vae_cpu_offload:
self.vae.to("cpu")
return batch
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
# Stage is a no-op when pil_image is None, so nothing required
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
# gt_latents and conditioning_mask are only set for I2V
return result
+9 -64
View File
@@ -141,7 +141,6 @@ class LTX2DenoisingStage(PipelineStage):
video_shape)
else:
token_count = 1
timestep_template = torch.ones(
(latents.shape[0], token_count),
device=latents.device,
@@ -210,27 +209,12 @@ class LTX2DenoisingStage(PipelineStage):
device=latents.device,
dtype=torch.float32,
)
# Multi-modal CFG parameters (per-stream scales).
cfg_scale_video = batch.ltx2_cfg_scale_video
cfg_scale_audio = batch.ltx2_cfg_scale_audio
modality_scale_video = batch.ltx2_modality_scale_video
modality_scale_audio = batch.ltx2_modality_scale_audio
rescale_scale = batch.ltx2_rescale_scale
do_cfg = batch.do_classifier_free_guidance
logger.info(
"[LTX2] Denoising start: steps=%d dtype=%s "
"cfg_video=%.1f cfg_audio=%.1f mod_video=%.1f "
"mod_audio=%.1f rescale=%.2f "
"[LTX2] Denoising start: steps=%d dtype=%s guidance=%s "
"sigmas_shape=%s latents_shape=%s",
batch.num_inference_steps,
target_dtype,
cfg_scale_video,
cfg_scale_audio,
modality_scale_video,
modality_scale_audio,
rescale_scale,
batch.guidance_scale,
tuple(sigmas.shape),
tuple(latents.shape),
)
@@ -251,7 +235,6 @@ class LTX2DenoisingStage(PipelineStage):
attn_metadata=None,
forward_batch=batch,
):
# Pass 1: Full conditioning (text + cross-modal)
pos_outputs = self.transformer(
hidden_states=latents.to(target_dtype),
encoder_hidden_states=prompt_embeds,
@@ -267,8 +250,8 @@ class LTX2DenoisingStage(PipelineStage):
pos_denoised = pos_outputs
pos_audio = None
if do_cfg:
# Pass 2: Unconditioned text (negative prompt)
# Only run negative pass if CFG is enabled
if batch.do_classifier_free_guidance:
neg_outputs = self.transformer(
hidden_states=latents.to(target_dtype),
encoder_hidden_states=neg_prompt_embeds,
@@ -283,48 +266,11 @@ class LTX2DenoisingStage(PipelineStage):
else:
neg_denoised = neg_outputs
neg_audio = None
# Pass 3: Modality-isolated (skip cross-modal attn)
mod_outputs = self.transformer(
hidden_states=latents.to(target_dtype),
encoder_hidden_states=prompt_embeds,
encoder_attention_mask=prompt_mask,
timestep=timestep,
audio_hidden_states=audio_latents,
audio_encoder_hidden_states=audio_context_p,
audio_timestep=audio_timestep,
skip_cross_modal_attn=True,
)
if isinstance(mod_outputs, tuple):
mod_denoised, mod_audio = mod_outputs
else:
mod_denoised = mod_outputs
mod_audio = None
# Multi-modal guidance formula per stream.
vid = (pos_denoised + (cfg_scale_video - 1) *
(pos_denoised - neg_denoised) +
(modality_scale_video - 1) *
(pos_denoised - mod_denoised))
aud = None
if pos_audio is not None:
aud = (pos_audio + (cfg_scale_audio - 1) *
(pos_audio - neg_audio) +
(modality_scale_audio - 1) *
(pos_audio - mod_audio))
# Guidance rescaling (prevents saturation).
if rescale_scale > 0:
f_v = pos_denoised.std() / vid.std()
f_v = rescale_scale * f_v + (1 - rescale_scale)
vid = vid * f_v
if aud is not None:
f_a = pos_audio.std() / aud.std()
f_a = rescale_scale * f_a + (1 - rescale_scale)
aud = aud * f_a
pos_denoised = vid
pos_audio = aud
pos_denoised = pos_denoised + (batch.guidance_scale - 1) * (
pos_denoised - neg_denoised)
if pos_audio is not None and neg_audio is not None:
pos_audio = pos_audio + (batch.guidance_scale -
1) * (pos_audio - neg_audio)
sigma_value = sigma.to(torch.float32) if isinstance(
sigma, torch.Tensor) else torch.tensor(
@@ -335,7 +281,6 @@ class LTX2DenoisingStage(PipelineStage):
dt = sigma_next - sigma
velocity = ((latents.float() - pos_denoised.float()) /
sigma_value).to(latents.dtype)
latents = (latents.float() + velocity.float() * dt).to(
latents.dtype)
if pos_audio is not None and audio_latents is not None:
@@ -1,374 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import inspect
from typing import Any
import torch
import torch.nn.functional as F
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.utils import PRECISION_TO_TYPE
class SD35LatentPreparationStage(PipelineStage):
def __init__(self, scheduler) -> None:
super().__init__()
self.scheduler = scheduler
@torch.no_grad()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.height is None or batch.width is None:
raise ValueError(
"height/width must be set for SD35LatentPreparationStage")
if isinstance(batch.prompt, list):
batch_size = len(batch.prompt)
elif batch.prompt is not None:
batch_size = 1
else:
if not batch.prompt_embeds:
raise ValueError("prompt or prompt_embeds must be provided")
batch_size = batch.prompt_embeds[0].shape[0]
batch_size *= batch.num_videos_per_prompt
dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
device = get_local_torch_device()
if isinstance(batch.generator, list) and len(
batch.generator) != batch_size:
raise ValueError(
f"generator list length {len(batch.generator)} does not match batch_size {batch_size}"
)
in_channels = fastvideo_args.pipeline_config.dit_config.arch_config.in_channels
spatial_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
h_lat = batch.height // spatial_ratio
w_lat = batch.width // spatial_ratio
shape = (batch_size, in_channels, 1, h_lat, w_lat)
latents = batch.latents
if latents is None:
latents = randn_tensor(
shape,
generator=batch.generator,
device=device,
dtype=dtype,
)
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
else:
latents = latents.to(device=device, dtype=dtype)
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
batch.latents = latents
batch.raw_latent_shape = shape
return batch
class SD35ConditioningStage(PipelineStage):
def __init__(self, text_encoders, tokenizers) -> None:
super().__init__()
self.text_encoders = text_encoders
self.tokenizers = tokenizers
@staticmethod
def _tokenize(
tokenizer: Any,
prompts: str | list[str],
tok_kwargs: dict[str, Any],
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
texts = [prompts] if isinstance(prompts, str) else prompts
enc = tokenizer(texts, **tok_kwargs)
input_ids = enc["input_ids"].to(device)
attention_mask = enc["attention_mask"].to(device)
return input_ids, attention_mask
@torch.no_grad()
def _clip_pooled(
self,
text_encoder,
tokenizer,
prompts: str | list[str],
tok_kwargs: dict[str, Any],
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor:
input_ids, attention_mask = self._tokenize(tokenizer, prompts,
tok_kwargs, device)
with set_forward_context(current_timestep=0, attn_metadata=None):
out = text_encoder(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=False,
)
pooled = getattr(out, "pooler_output", None)
if pooled is None:
raise RuntimeError(
"CLIP pooled output is required for SD3.5 conditioning")
return pooled.to(dtype=dtype)
@torch.no_grad()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
if len(batch.prompt_embeds) < 3:
raise ValueError(
f"SD35ConditioningStage expects 3 prompt_embeds entries (2x CLIP + 1x T5), got {len(batch.prompt_embeds)}"
)
device = get_local_torch_device()
target_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.dit_precision]
clip_1 = batch.prompt_embeds[0].to(device=device, dtype=target_dtype)
clip_2 = batch.prompt_embeds[1].to(device=device, dtype=target_dtype)
t5 = batch.prompt_embeds[2].to(device=device, dtype=target_dtype)
clip_prompt = torch.cat([clip_1, clip_2], dim=-1)
if clip_prompt.shape[-1] > t5.shape[-1]:
raise ValueError(
f"CLIP prompt dim {clip_prompt.shape[-1]} exceeds T5 dim {t5.shape[-1]}"
)
clip_prompt = F.pad(clip_prompt,
(0, t5.shape[-1] - clip_prompt.shape[-1]))
prompt_embeds = torch.cat([clip_prompt, t5], dim=-2)
te_cfgs = fastvideo_args.pipeline_config.text_encoder_configs
clip_tok_kwargs_1 = dict(getattr(te_cfgs[0], "tokenizer_kwargs", {}))
clip_tok_kwargs_2 = dict(getattr(te_cfgs[1], "tokenizer_kwargs", {}))
clip_tok_kwargs_1.setdefault("padding", "max_length")
clip_tok_kwargs_2.setdefault("padding", "max_length")
clip_tok_kwargs_1.setdefault("max_length", 77)
clip_tok_kwargs_2.setdefault("max_length", 77)
clip_tok_kwargs_1.setdefault("truncation", True)
clip_tok_kwargs_2.setdefault("truncation", True)
clip_tok_kwargs_1.setdefault("return_tensors", "pt")
clip_tok_kwargs_2.setdefault("return_tensors", "pt")
pooled_1 = self._clip_pooled(self.text_encoders[0], self.tokenizers[0],
batch.prompt, clip_tok_kwargs_1, device,
target_dtype)
pooled_2 = self._clip_pooled(self.text_encoders[1], self.tokenizers[1],
batch.prompt, clip_tok_kwargs_2, device,
target_dtype)
pooled = torch.cat([pooled_1, pooled_2], dim=-1)
batch.extra["sd35_encoder_hidden_states"] = prompt_embeds
batch.extra["sd35_pooled_projections"] = pooled
if batch.do_classifier_free_guidance:
if batch.negative_prompt_embeds is None or len(
batch.negative_prompt_embeds) < 3:
raise ValueError(
"negative_prompt_embeds must contain 3 entries when CFG is enabled"
)
neg_clip_1 = batch.negative_prompt_embeds[0].to(device=device,
dtype=target_dtype)
neg_clip_2 = batch.negative_prompt_embeds[1].to(device=device,
dtype=target_dtype)
neg_t5 = batch.negative_prompt_embeds[2].to(device=device,
dtype=target_dtype)
neg_clip_prompt = torch.cat([neg_clip_1, neg_clip_2], dim=-1)
neg_clip_prompt = F.pad(
neg_clip_prompt,
(0, neg_t5.shape[-1] - neg_clip_prompt.shape[-1]))
neg_prompt_embeds = torch.cat([neg_clip_prompt, neg_t5], dim=-2)
negative_pooled_1 = self._clip_pooled(
self.text_encoders[0],
self.tokenizers[0],
batch.negative_prompt if isinstance(batch.negative_prompt, str
| list) else "",
clip_tok_kwargs_1,
device,
target_dtype,
)
negative_pooled_2 = self._clip_pooled(
self.text_encoders[1],
self.tokenizers[1],
batch.negative_prompt if isinstance(batch.negative_prompt, str
| list) else "",
clip_tok_kwargs_2,
device,
target_dtype,
)
neg_pooled = torch.cat([negative_pooled_1, negative_pooled_2],
dim=-1)
batch.extra[
"sd35_negative_encoder_hidden_states"] = neg_prompt_embeds
batch.extra["sd35_negative_pooled_projections"] = neg_pooled
return batch
class SD35DenoisingStage(PipelineStage):
"""Denoising loop for SD3.5 (2D transformer + FlowMatch scheduler)."""
def __init__(self, transformer, scheduler) -> None:
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
@staticmethod
def _prepare_extra_func_kwargs(func, kwargs) -> dict[str, Any]:
extra_kwargs: dict[str, Any] = {}
sig = inspect.signature(func)
for k, v in kwargs.items():
if k in sig.parameters:
extra_kwargs[k] = v
return extra_kwargs
@torch.no_grad()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.timesteps is None:
raise ValueError("timesteps must be set before SD35DenoisingStage")
if batch.latents is None:
raise ValueError("latents must be set before SD35DenoisingStage")
prompt_embeds: torch.Tensor = batch.extra["sd35_encoder_hidden_states"]
pooled: torch.Tensor = batch.extra["sd35_pooled_projections"]
neg_prompt_embeds: torch.Tensor | None = batch.extra.get(
"sd35_negative_encoder_hidden_states")
neg_pooled: torch.Tensor | None = batch.extra.get(
"sd35_negative_pooled_projections")
timesteps = batch.timesteps
latents = batch.latents
guidance_scale = float(batch.guidance_scale)
target_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.dit_precision]
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
extra_step_kwargs = self._prepare_extra_func_kwargs(
self.scheduler.step, {
"generator":
batch.generator[0]
if isinstance(batch.generator, list) else batch.generator
})
for t in timesteps:
latents_4d = latents.squeeze(2)
if batch.do_classifier_free_guidance:
if neg_prompt_embeds is None or neg_pooled is None:
raise ValueError(
"Missing negative conditioning tensors for CFG")
latent_model_input = torch.cat([latents_4d] * 2, dim=0)
cond_embeds = torch.cat([neg_prompt_embeds, prompt_embeds],
dim=0)
cond_pooled = torch.cat([neg_pooled, pooled], dim=0)
else:
latent_model_input = latents_4d
cond_embeds = prompt_embeds
cond_pooled = pooled
latent_model_input = self.scheduler.scale_model_input(
latent_model_input, t)
timestep = t.expand(latent_model_input.shape[0])
with torch.autocast(
device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled
and (get_local_torch_device().type == "cuda"),
):
noise_pred = self.transformer(
hidden_states=latent_model_input,
timestep=timestep,
encoder_hidden_states=cond_embeds,
pooled_projections=cond_pooled,
return_dict=False,
)[0]
if batch.do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + guidance_scale * (
noise_pred_text - noise_pred_uncond)
noise_pred_5d = noise_pred.unsqueeze(2)
latents = self.scheduler.step(noise_pred_5d,
t,
latents,
return_dict=False,
**extra_step_kwargs)[0]
batch.latents = latents
return batch
class SD35DecodingStage(PipelineStage):
def __init__(self, vae) -> None:
super().__init__()
self.vae = vae
@staticmethod
def _denormalize_latents(latents: torch.Tensor, vae) -> torch.Tensor:
# Prefer config fields to avoid diffusers deprecation warnings for direct
# attribute access (vae.scaling_factor / vae.shift_factor).
cfg = getattr(vae, "config", None)
sf = getattr(cfg, "scaling_factor", None) if cfg is not None else None
sh = getattr(cfg, "shift_factor", None) if cfg is not None else None
if sf is None and hasattr(vae, "scaling_factor"):
sf = vae.scaling_factor
if sh is None and hasattr(vae, "shift_factor"):
sh = vae.shift_factor
if sf is not None:
latents = latents / (sf.to(latents.device, latents.dtype)
if isinstance(sf, torch.Tensor) else sf)
if sh is not None:
latents = latents + (sh.to(latents.device, latents.dtype)
if isinstance(sh, torch.Tensor) else sh)
return latents
@torch.no_grad()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.latents is None:
raise ValueError("latents must be set before SD35DecodingStage")
device = get_local_torch_device()
latents_5d = batch.latents.to(device)
latents_4d = latents_5d.squeeze(2)
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
latents_4d = self._denormalize_latents(latents_4d, self.vae)
with torch.autocast(
device_type="cuda",
dtype=vae_dtype,
enabled=autocast_enabled and (device.type == "cuda"),
):
if not autocast_enabled:
latents_4d = latents_4d.to(dtype=vae_dtype)
dec = self.vae.decode(latents_4d)
image = dec.sample if hasattr(dec, "sample") else dec[0]
image = (image / 2 + 0.5).clamp(0, 1)
batch.output = image.unsqueeze(2).detach().float().cpu()
return batch
+1 -7
View File
@@ -253,17 +253,11 @@ class TextEncodingStage(PipelineStage):
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
want_hidden_states = bool(
getattr(getattr(encoder_config, "arch_config", None),
"output_hidden_states", False))
if is_ltx2:
want_hidden_states = True
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs = text_encoder(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=want_hidden_states,
output_hidden_states=True,
)
try:
+3 -53
View File
@@ -18,12 +18,10 @@ from fastvideo.configs.pipelines.base import PipelineConfig
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.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.hunyuan15 import (
Hunyuan15T2V480PConfig, Hunyuan15I2V480PStepDistilledConfig,
Hunyuan15T2V720PConfig, Hunyuan15I2V720PConfig, Hunyuan15SR1080PConfig)
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.lingbotworld import LingBotWorldI2V480PConfig
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
@@ -47,7 +45,6 @@ from fastvideo.configs.pipelines.wan import (
WanT2V480PConfig,
WanT2V720PConfig,
)
from fastvideo.configs.pipelines.sd35 import SD35Config
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.configs.sample.cosmos import (
Cosmos_Predict2_2B_Video2World_SamplingParam, )
@@ -60,8 +57,6 @@ from fastvideo.configs.sample.hunyuan15 import (
Hunyuan15_720P_SamplingParam, Hunyuan15_720P_Distilled_I2V_SamplingParam,
Hunyuan15_SR_1080P_SamplingParam)
from fastvideo.configs.sample.hyworld import HYWorld_SamplingParam
from fastvideo.configs.sample.hunyuangamecraft import HunyuanGameCraftSamplingParam
from fastvideo.configs.sample.lingbotworld import LingBotWorld_SamplingParam
from fastvideo.configs.sample.ltx2 import (LTX2BaseSamplingParam,
LTX2DistilledSamplingParam)
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
@@ -85,8 +80,6 @@ from fastvideo.configs.sample.wan import (
WanT2V_14B_SamplingParam,
WanT2V_1_3B_SamplingParam,
)
from fastvideo.configs.sample.sd35 import SD35SamplingParam
from fastvideo.fastvideo_args import WorkloadType
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -254,7 +247,6 @@ def _register_configs() -> None:
hf_model_paths=[
"Lightricks/LTX-2",
"FastVideo/LTX2-base",
"FastVideo/LTX2-Diffusers",
],
model_detectors=[
lambda path: ("ltx2" in path.lower() or "ltx-2" in path.lower()) and
@@ -320,18 +312,14 @@ def _register_configs() -> None:
],
)
# Hunyuan (excludes gamecraft, hyworld, and versioned models)
# Hunyuan
register_configs(
sampling_param_cls=HunyuanSamplingParam,
pipeline_config_cls=HunyuanConfig,
hf_model_paths=[
"hunyuanvideo-community/HunyuanVideo",
],
model_detectors=[
lambda path: "hunyuan" in path.lower()
and "gamecraft" not in path.lower() and "hyworld" not in path.lower(
) and "1.5" not in path.lower() and "1-5" not in path.lower()
],
model_detectors=[lambda path: "hunyuan" in path.lower()],
)
register_configs(
sampling_param_cls=FastHunyuanSamplingParam,
@@ -351,28 +339,6 @@ def _register_configs() -> None:
model_detectors=[lambda path: "hyworld" in path.lower()],
)
# HunyuanGameCraft
register_configs(
sampling_param_cls=HunyuanGameCraftSamplingParam,
pipeline_config_cls=HunyuanGameCraftPipelineConfig,
hf_model_paths=[
"FastVideo/HunyuanGameCraft-Diffusers",
],
model_detectors=[lambda path: "gamecraft" in path.lower()],
)
# LingBotWorld
register_configs(
sampling_param_cls=LingBotWorld_SamplingParam,
pipeline_config_cls=LingBotWorldI2V480PConfig,
hf_model_paths=[
"FastVideo/LingBot-World-Base-Cam-Diffusers",
],
model_detectors=[
lambda path:
("lingbotworld" in path.lower() or "lingbot-world" in path.lower())
],
)
# LongCat
register_configs(
sampling_param_cls=None,
@@ -571,22 +537,6 @@ def _register_configs() -> None:
],
)
# SD3.5
register_configs(
sampling_param_cls=SD35SamplingParam,
pipeline_config_cls=SD35Config,
hf_model_paths=[
"stabilityai/stable-diffusion-3.5-medium",
],
model_detectors=[
lambda path: any(token in path.lower() for token in (
"sd35",
"stablediffusion3",
"stabilityai__stable-diffusion-3.5-medium",
)),
],
)
# --- Part 3: Main Resolver ---
@@ -679,4 +629,4 @@ __all__ = [
"get_pipeline_config_cls_from_name",
"get_sampling_param_cls_for_name",
"get_pipeline_config_classes",
]
]
@@ -6,7 +6,7 @@ import pytest
import torch
from transformers import AutoConfig, AutoTokenizer, CLIPTextModel
import gc
from fastvideo.configs.pipelines import HunyuanConfig, PipelineConfig
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
@@ -40,7 +40,7 @@ def test_clip_encoder():
- Produce nearly identical outputs for the same input prompts
"""
args = FastVideoArgs(model_path="openai/clip-vit-large-patch14",
pipeline_config=HunyuanConfig())
pipeline_config=PipelineConfig(text_encoder_configs=(CLIPTextConfig(),), text_encoder_precisions=("fp16",)))
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
logger.info("Loading models from %s", args.model_path)
@@ -8,7 +8,7 @@ from torch.distributed.tensor import DTensor
from torch.testing import assert_close
from transformers import AutoConfig, AutoTokenizer, T5EncoderModel
from fastvideo.configs.pipelines import Hunyuan15T2V480PConfig, PipelineConfig
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TextEncoderLoader
@@ -23,18 +23,18 @@ os.environ["MASTER_PORT"] = "29503"
@pytest.fixture
def t5_model_paths_and_config():
def t5_model_paths():
base_model_path = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
model_path = maybe_download_model(base_model_path)
text_encoder_path = os.path.join(model_path, "text_encoder_2")
tokenizer_path = os.path.join(model_path, "tokenizer_2")
return text_encoder_path, tokenizer_path, Hunyuan15T2V480PConfig()
return text_encoder_path, tokenizer_path
@pytest.mark.usefixtures("distributed_setup")
def test_t5_encoder(t5_model_paths_and_config):
def test_t5_encoder(t5_model_paths):
# Initialize the two model implementations
text_encoder_path, tokenizer_path, pipeline_config = t5_model_paths_and_config
text_encoder_path, tokenizer_path = t5_model_paths
hf_config = AutoConfig.from_pretrained(text_encoder_path)
print(hf_config)
@@ -47,7 +47,8 @@ def test_t5_encoder(t5_model_paths_and_config):
args = FastVideoArgs(model_path=text_encoder_path,
pipeline_config=pipeline_config,
pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),),
text_encoder_precisions=(precision_str,)),
pin_cpu_memory=False)
loader = TextEncoderLoader()
model2 = loader.load(text_encoder_path, args)
@@ -6,7 +6,7 @@ import pytest
import torch
from transformers import AutoConfig, AutoTokenizer, LlamaModel
import gc
from fastvideo.configs.pipelines import HunyuanConfig
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
@@ -40,7 +40,7 @@ def test_llama_encoder():
- Produce nearly identical outputs for the same input prompts
"""
args = FastVideoArgs(model_path="meta-llama/Llama-2-7b-hf",
pipeline_config=HunyuanConfig())
pipeline_config=PipelineConfig(text_encoder_configs=(LlamaConfig(),), text_encoder_precisions=("fp16",)))
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
@@ -6,7 +6,7 @@ from torch.distributed.tensor import DTensor
from torch.testing import assert_close
from transformers import AutoConfig, AutoTokenizer, Qwen2_5_VLTextModel
from fastvideo.configs.pipelines import Hunyuan15T2V480PConfig, PipelineConfig
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TextEncoderLoader
@@ -20,16 +20,16 @@ os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29505"
@pytest.fixture
def qwen_model_path_and_config():
def qwen_model_path():
base_model_path = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
model_path = maybe_download_model(base_model_path)
text_encoder_path = os.path.join(model_path, "text_encoder")
tokenizer_path = os.path.join(model_path, "tokenizer")
return text_encoder_path, tokenizer_path, Hunyuan15T2V480PConfig()
return text_encoder_path, tokenizer_path
@pytest.mark.usefixtures("distributed_setup")
def test_qwen2_5_encoder(qwen_model_path_and_config):
text_encoder_path, tokenizer_path, pipeline_config = qwen_model_path_and_config
def test_qwen2_5_encoder(qwen_model_path):
text_encoder_path, tokenizer_path = qwen_model_path
hf_config = AutoConfig.from_pretrained(text_encoder_path)
print(hf_config)
@@ -47,7 +47,8 @@ def test_qwen2_5_encoder(qwen_model_path_and_config):
# Load FastVideo model
args = FastVideoArgs(model_path=text_encoder_path,
pipeline_config=pipeline_config,
pipeline_config=PipelineConfig(text_encoder_configs=(Qwen2_5_VLConfig(),),
text_encoder_precisions=(precision_str,)),
pin_cpu_memory=False)
loader = TextEncoderLoader()
+13 -11
View File
@@ -8,7 +8,7 @@ from torch.distributed.tensor import DTensor
from torch.testing import assert_close
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel, T5EncoderModel
from fastvideo.configs.pipelines import CosmosConfig, PipelineConfig, WanT2V480PConfig
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TextEncoderLoader
@@ -23,31 +23,31 @@ os.environ["MASTER_PORT"] = "29503"
@pytest.fixture
def t5_model_paths_and_config():
def t5_model_paths():
base_model_path = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
model_path = maybe_download_model(base_model_path,
local_dir=os.path.join(
'data', base_model_path))
text_encoder_path = os.path.join(model_path, "text_encoder")
tokenizer_path = os.path.join(model_path, "tokenizer")
return text_encoder_path, tokenizer_path, WanT2V480PConfig()
return text_encoder_path, tokenizer_path
@pytest.fixture
def t5_large_model_paths_and_config():
def t5_large_model_paths():
base_model_path = "nvidia/Cosmos-Predict2-2B-Video2World"
model_path = maybe_download_model(base_model_path,
local_dir=os.path.join(
'data', base_model_path))
text_encoder_path = os.path.join(model_path, "text_encoder")
tokenizer_path = os.path.join(model_path, "tokenizer")
return text_encoder_path, tokenizer_path, CosmosConfig()
return text_encoder_path, tokenizer_path
@pytest.mark.usefixtures("distributed_setup")
def test_t5_encoder(t5_model_paths_and_config):
def test_t5_encoder(t5_model_paths):
# Initialize the two model implementations
text_encoder_path, tokenizer_path, pipeline_config = t5_model_paths_and_config
text_encoder_path, tokenizer_path = t5_model_paths
hf_config = AutoConfig.from_pretrained(text_encoder_path)
print(hf_config)
@@ -60,7 +60,8 @@ def test_t5_encoder(t5_model_paths_and_config):
args = FastVideoArgs(model_path=text_encoder_path,
pipeline_config=pipeline_config,
pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),),
text_encoder_precisions=(precision_str,)),
pin_cpu_memory=False)
loader = TextEncoderLoader()
model2 = loader.load(text_encoder_path, args)
@@ -136,9 +137,9 @@ def test_t5_encoder(t5_model_paths_and_config):
@pytest.mark.usefixtures("distributed_setup")
def test_t5_large_encoder(t5_large_model_paths_and_config):
def test_t5_large_encoder(t5_large_model_paths):
# Initialize the two model implementations
text_encoder_path, tokenizer_path, pipeline_config = t5_large_model_paths_and_config
text_encoder_path, tokenizer_path = t5_large_model_paths
hf_config = AutoConfig.from_pretrained(text_encoder_path)
print(hf_config)
@@ -150,7 +151,8 @@ def test_t5_large_encoder(t5_large_model_paths_and_config):
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
args = FastVideoArgs(model_path=text_encoder_path,
pipeline_config=pipeline_config,
pipeline_config=PipelineConfig(text_encoder_configs=(T5LargeConfig(),),
text_encoder_precisions=(precision_str,)),
pin_cpu_memory=False)
loader = TextEncoderLoader()
model2 = loader.load(text_encoder_path, args)
@@ -74,5 +74,5 @@ def test_prepare_output_path_empty_prompt_fallback(tmp_path):
result = vg._prepare_output_path(str(out_dir), prompt=bad_prompt)
assert os.path.dirname(result) == str(out_dir)
assert os.path.basename(result) == "output.mp4"
assert os.path.basename(result) == "video.mp4"
@@ -19,10 +19,6 @@ if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
elif "H100" in device_name:
device_reference_folder = "H100" + device_reference_folder_suffix
elif "H200" in device_name:
device_reference_folder = "H200" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
logger.warning(f"Unsupported device for ssim tests: {device_name}")
@@ -1,355 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
SSIM regression test for HunyuanGameCraft (T2V and I2V).
Generates a video with deterministic seed and camera trajectory,
then compares against a device-specific reference video via MS-SSIM.
Reference videos must be pre-generated and stored under:
<device>_reference_videos/HunyuanGameCraft/<ATTENTION_BACKEND>/
To create initial reference videos, run this test once and copy the
generated videos into the appropriate reference folder.
"""
import os
import torch
import pytest
from fastvideo import VideoGenerator
from fastvideo.logger import init_logger
from fastvideo.tests.utils import (
compute_video_ssim_torchvision,
write_ssim_results,
)
from fastvideo.worker.multiproc_executor import MultiprocExecutor
logger = init_logger(__name__)
# ---------------------------------------------------------------------------
# Device-dependent reference folder
# ---------------------------------------------------------------------------
device_name = torch.cuda.get_device_name()
device_reference_folder_suffix = "_reference_videos"
if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
elif "A100" in device_name:
device_reference_folder = "A100" + device_reference_folder_suffix
elif "H100" in device_name:
device_reference_folder = "H100" + device_reference_folder_suffix
else:
device_reference_folder = "Unknown" + device_reference_folder_suffix
logger.warning(f"Unsupported device for ssim tests: {device_name}")
# ---------------------------------------------------------------------------
# Helpers – camera trajectory (self-contained, no official repo dependency)
# ---------------------------------------------------------------------------
from fastvideo.models.camera import create_camera_trajectory as _create_camera_trajectory
def _shutdown_executor(generator: VideoGenerator | None) -> None:
if generator is None:
return
if isinstance(generator.executor, MultiprocExecutor):
generator.executor.shutdown()
# ---------------------------------------------------------------------------
# Model parameters
# ---------------------------------------------------------------------------
# Same as basic_gamecraft.py: default HF path; set GAMECRAFT_MODEL_PATH for local weights.
_GAMECRAFT_MODEL_PATH = os.environ.get(
"GAMECRAFT_MODEL_PATH",
"FastVideo/HunyuanGameCraft-Diffusers",
)
GAMECRAFT_T2V_PARAMS = {
"num_gpus": 1,
"model_path": _GAMECRAFT_MODEL_PATH,
"height": 480,
"width": 832,
"num_frames": 33,
"num_inference_steps": 20,
"guidance_scale": 6.0,
"seed": 1024,
"action": "forward",
"action_speed": 0.2,
"negative_prompt": "",
}
GAMECRAFT_I2V_PARAMS = {
**GAMECRAFT_T2V_PARAMS,
"image_path": (
"https://huggingface.co/datasets/huggingface/documentation-images/"
"resolve/main/diffusers/astronaut.jpg"
),
}
MODEL_TO_PARAMS = {
"HunyuanGameCraft-T2V": GAMECRAFT_T2V_PARAMS,
}
I2V_MODEL_TO_PARAMS = {
"HunyuanGameCraft-I2V": GAMECRAFT_I2V_PARAMS,
}
TEST_PROMPTS = [
"A majestic ancient temple stands under a clear blue sky, its grandeur highlighted by towering Doric columns and intricate architectural details.",
]
I2V_TEST_PROMPTS = [
"An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background.",
]
# ---------------------------------------------------------------------------
# T2V SSIM test
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
def test_gamecraft_t2v_similarity(prompt, ATTENTION_BACKEND, model_id):
"""Generate a T2V video with GameCraft and compare to reference via SSIM."""
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
script_dir = os.path.dirname(os.path.abspath(__file__))
base_output_dir = os.path.join(script_dir, "generated_videos", model_id)
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
BASE_PARAMS = MODEL_TO_PARAMS[model_id]
num_inference_steps = BASE_PARAMS["num_inference_steps"]
# Build camera trajectory
camera_states = _create_camera_trajectory(
action=BASE_PARAMS["action"],
height=BASE_PARAMS["height"],
width=BASE_PARAMS["width"],
num_frames=BASE_PARAMS["num_frames"],
action_speed=BASE_PARAMS["action_speed"],
dtype=torch.bfloat16,
)
init_kwargs = {
"num_gpus": BASE_PARAMS["num_gpus"],
"use_fsdp_inference": True,
"dit_cpu_offload": True,
"vae_cpu_offload": True,
"text_encoder_cpu_offload": True,
"pin_cpu_memory": True,
}
generation_kwargs = {
"num_inference_steps": num_inference_steps,
"output_path": output_dir,
"height": BASE_PARAMS["height"],
"width": BASE_PARAMS["width"],
"num_frames": BASE_PARAMS["num_frames"],
"guidance_scale": BASE_PARAMS["guidance_scale"],
"seed": BASE_PARAMS["seed"],
"fps": 24,
"camera_states": camera_states,
"negative_prompt": BASE_PARAMS.get("negative_prompt", ""),
"save_video": True,
}
generator: VideoGenerator | None = None
try:
generator = VideoGenerator.from_pretrained(
model_path=BASE_PARAMS["model_path"], **init_kwargs
)
generator.generate_video(prompt, **generation_kwargs)
finally:
_shutdown_executor(generator)
assert os.path.exists(output_dir), (
f"Output video was not generated at {output_dir}"
)
# Compare to reference
reference_folder = os.path.join(
script_dir, device_reference_folder, model_id, ATTENTION_BACKEND
)
if not os.path.exists(reference_folder):
logger.error("Reference folder missing")
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}"
)
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith(".mp4") and prompt[:100].strip() in filename:
reference_video_name = filename
break
if not reference_video_name:
logger.error(
f"Reference video not found for prompt: {prompt} "
f"with backend: {ATTENTION_BACKEND}"
)
raise FileNotFoundError("Reference video missing")
reference_video_path = os.path.join(reference_folder, reference_video_name)
generated_video_path = os.path.join(output_dir, output_video_name)
logger.info(
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
)
ssim_values = compute_video_ssim_torchvision(
reference_video_path, generated_video_path, use_ms_ssim=True
)
mean_ssim = ssim_values[0]
logger.info(f"SSIM mean value: {mean_ssim}")
logger.info(f"Writing SSIM results to directory: {output_dir}")
success = write_ssim_results(
output_dir,
ssim_values,
reference_video_path,
generated_video_path,
num_inference_steps,
prompt,
)
if not success:
logger.error("Failed to write SSIM results to file")
min_acceptable_ssim = 0.93
assert mean_ssim >= min_acceptable_ssim, (
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
f"for {model_id} with backend {ATTENTION_BACKEND}"
)
# ---------------------------------------------------------------------------
# I2V SSIM test
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("prompt", I2V_TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
@pytest.mark.parametrize("model_id", list(I2V_MODEL_TO_PARAMS.keys()))
def test_gamecraft_i2v_similarity(prompt, ATTENTION_BACKEND, model_id):
"""Generate an I2V video with GameCraft and compare to reference via SSIM."""
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
script_dir = os.path.dirname(os.path.abspath(__file__))
base_output_dir = os.path.join(script_dir, "generated_videos", model_id)
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
BASE_PARAMS = I2V_MODEL_TO_PARAMS[model_id]
num_inference_steps = BASE_PARAMS["num_inference_steps"]
# Build camera trajectory
camera_states = _create_camera_trajectory(
action=BASE_PARAMS["action"],
height=BASE_PARAMS["height"],
width=BASE_PARAMS["width"],
num_frames=BASE_PARAMS["num_frames"],
action_speed=BASE_PARAMS["action_speed"],
dtype=torch.bfloat16,
)
init_kwargs = {
"num_gpus": BASE_PARAMS["num_gpus"],
"use_fsdp_inference": True,
"dit_cpu_offload": True,
"vae_cpu_offload": True,
"text_encoder_cpu_offload": True,
"pin_cpu_memory": True,
}
generation_kwargs = {
"num_inference_steps": num_inference_steps,
"output_path": output_dir,
"image_path": BASE_PARAMS["image_path"],
"height": BASE_PARAMS["height"],
"width": BASE_PARAMS["width"],
"num_frames": BASE_PARAMS["num_frames"],
"guidance_scale": BASE_PARAMS["guidance_scale"],
"seed": BASE_PARAMS["seed"],
"fps": 24,
"camera_states": camera_states,
"negative_prompt": BASE_PARAMS.get("negative_prompt", ""),
"save_video": True,
}
generator: VideoGenerator | None = None
try:
generator = VideoGenerator.from_pretrained(
model_path=BASE_PARAMS["model_path"], **init_kwargs
)
generator.generate_video(prompt, **generation_kwargs)
finally:
_shutdown_executor(generator)
assert os.path.exists(output_dir), (
f"Output video was not generated at {output_dir}"
)
# Compare to reference
reference_folder = os.path.join(
script_dir, device_reference_folder, model_id, ATTENTION_BACKEND
)
if not os.path.exists(reference_folder):
logger.error("Reference folder missing")
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}"
)
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith(".mp4") and prompt[:100].strip() in filename:
reference_video_name = filename
break
if not reference_video_name:
logger.error(
f"Reference video not found for prompt: {prompt} "
f"with backend: {ATTENTION_BACKEND}"
)
raise FileNotFoundError("Reference video missing")
reference_video_path = os.path.join(reference_folder, reference_video_name)
generated_video_path = os.path.join(output_dir, output_video_name)
logger.info(
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
)
ssim_values = compute_video_ssim_torchvision(
reference_video_path, generated_video_path, use_ms_ssim=True
)
mean_ssim = ssim_values[0]
logger.info(f"SSIM mean value: {mean_ssim}")
logger.info(f"Writing SSIM results to directory: {output_dir}")
success = write_ssim_results(
output_dir,
ssim_values,
reference_video_path,
generated_video_path,
num_inference_steps,
prompt,
)
if not success:
logger.error("Failed to write SSIM results to file")
min_acceptable_ssim = 0.93
assert mean_ssim >= min_acceptable_ssim, (
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
f"for {model_id} with backend {ATTENTION_BACKEND}"
)
@@ -41,10 +41,6 @@ if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
elif "H100" in device_name:
device_reference_folder = "H100" + device_reference_folder_suffix
elif "H200" in device_name:
device_reference_folder = "H200" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
logger.warning(f"Unsupported device for ssim tests: {device_name}")
@@ -1,206 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
SSIM-based similarity test for LingBotWorld I2V with camera control.
Camera trajectory is loaded from LingBot example npy files
(poses.npy/intrinsics.npy), matching the official script workflow.
Note: num_inference_steps is reduced to 4 for faster CI.
"""
import os
import pytest
import torch
from fastvideo import VideoGenerator
from fastvideo.logger import init_logger
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
logger = init_logger(__name__)
def _find_lingbotworld_examples_root() -> str | None:
script_dir = os.path.dirname(os.path.abspath(__file__))
repo_root = os.path.abspath(os.path.join(script_dir, "..", "..", ".."))
candidates = [
os.path.join(repo_root, "examples", "inference", "basic",
"lingbotworld_examples"),
os.path.join(repo_root, "..", "FastVideo", "examples", "inference",
"basic", "lingbotworld_examples"),
]
for candidate in candidates:
if (os.path.exists(os.path.join(candidate, "00", "poses.npy"))
and os.path.exists(
os.path.join(candidate, "00", "intrinsics.npy"))):
return os.path.abspath(candidate)
return None
device_name = torch.cuda.get_device_name()
device_reference_folder_suffix = "_reference_videos"
if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
elif "H100" in device_name:
device_reference_folder = "H100" + device_reference_folder_suffix
elif "H200" in device_name:
device_reference_folder = "H200" + device_reference_folder_suffix
else:
device_reference_folder = None
logger.warning("Unsupported device for ssim tests: %s", device_name)
LINGBOT_PARAMS = {
"model_path": "FastVideo/LingBot-World-Base-Cam-Diffusers",
"num_gpus": 2,
"height": 256,
"width": 448,
"num_frames": 45, # must be 4k+1
"num_inference_steps": 4,
"guidance_scale": 5.0,
"guidance_scale_2": 5.0,
"embedded_cfg_scale": 6,
"flow_shift": 10.0,
"boundary_ratio": 0.947,
"seed": 42,
"fps": 16,
"spatial_scale": 8,
"example_case": "00",
"image_path": (
"https://raw.githubusercontent.com/Robbyant/lingbot-world/main/"
"examples/00/image.jpg"
),
"negative_prompt": (
"画面突变,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,"
"最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,"
"畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走,"
"镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,"
"PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,"
"皮肤,肢体,面部特征,汽车,电线"
),
}
TEST_PROMPTS = [
"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.",
]
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
def test_lingbot_i2v_similarity(prompt: str, ATTENTION_BACKEND: str):
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
if device_reference_folder is None:
pytest.skip(f"Unsupported device for LingBot SSIM test: {device_name}")
if torch.cuda.device_count() < LINGBOT_PARAMS["num_gpus"]:
pytest.skip(
f"LingBot SSIM test requires {LINGBOT_PARAMS['num_gpus']} GPUs, "
f"but only {torch.cuda.device_count()} detected."
)
examples_root = _find_lingbotworld_examples_root()
if examples_root is None:
pytest.skip(
"lingbotworld_examples not found under examples/inference/basic.")
action_path = os.path.join(examples_root, LINGBOT_PARAMS["example_case"])
if not (os.path.exists(os.path.join(action_path, "poses.npy"))
and os.path.exists(os.path.join(action_path, "intrinsics.npy"))):
pytest.skip(f"Missing camera npy files under {action_path}")
c2ws_plucker_emb, aligned_num_frames = prepare_camera_embedding(
action_path=action_path,
num_frames=LINGBOT_PARAMS["num_frames"],
height=LINGBOT_PARAMS["height"],
width=LINGBOT_PARAMS["width"],
spatial_scale=LINGBOT_PARAMS["spatial_scale"],
)
script_dir = os.path.dirname(os.path.abspath(__file__))
model_id = "LingBot-World-Base-Cam-Diffusers"
output_dir = os.path.join(script_dir, "generated_videos", model_id,
ATTENTION_BACKEND)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
init_kwargs = {
"num_gpus": LINGBOT_PARAMS["num_gpus"],
"flow_shift": LINGBOT_PARAMS["flow_shift"],
"boundary_ratio": LINGBOT_PARAMS["boundary_ratio"],
"use_fsdp_inference": True,
"dit_cpu_offload": True,
"dit_layerwise_offload": False,
"text_encoder_cpu_offload": True,
"vae_cpu_offload": False,
"pin_cpu_memory": True,
}
generation_kwargs = {
"output_path": output_dir,
"image_path": LINGBOT_PARAMS["image_path"],
"height": LINGBOT_PARAMS["height"],
"width": LINGBOT_PARAMS["width"],
"num_frames": aligned_num_frames,
"num_inference_steps": LINGBOT_PARAMS["num_inference_steps"],
"guidance_scale": LINGBOT_PARAMS["guidance_scale"],
"guidance_scale_2": LINGBOT_PARAMS["guidance_scale_2"],
"embedded_cfg_scale": LINGBOT_PARAMS["embedded_cfg_scale"],
"seed": LINGBOT_PARAMS["seed"],
"fps": LINGBOT_PARAMS["fps"],
"negative_prompt": LINGBOT_PARAMS["negative_prompt"],
"c2ws_plucker_emb": c2ws_plucker_emb,
}
generator: VideoGenerator | None = None
try:
generator = VideoGenerator.from_pretrained(
model_path=LINGBOT_PARAMS["model_path"], **init_kwargs)
generator.generate_video(prompt, **generation_kwargs)
finally:
if generator is not None:
generator.shutdown()
generated_video_path = os.path.join(output_dir, output_video_name)
assert os.path.exists(generated_video_path), (
f"Output video was not generated at {generated_video_path}")
reference_folder = os.path.join(script_dir, device_reference_folder, model_id,
ATTENTION_BACKEND)
if not os.path.exists(reference_folder):
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}")
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith(".mp4") and prompt[:100].strip() in filename:
reference_video_name = filename
break
if not reference_video_name:
raise FileNotFoundError(
f"Reference video missing for prompt/backend under {reference_folder}"
)
reference_video_path = os.path.join(reference_folder, reference_video_name)
logger.info("Computing SSIM between %s and %s", reference_video_path,
generated_video_path)
ssim_values = compute_video_ssim_torchvision(reference_video_path,
generated_video_path,
use_ms_ssim=True)
mean_ssim = ssim_values[0]
logger.info("SSIM mean value: %s", mean_ssim)
write_ssim_results(output_dir, ssim_values, reference_video_path,
generated_video_path,
LINGBOT_PARAMS["num_inference_steps"], prompt)
min_acceptable_ssim = 0.90
assert mean_ssim >= min_acceptable_ssim, (
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
f"for {model_id} with backend {ATTENTION_BACKEND}")
@@ -35,8 +35,6 @@ elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
elif "H100" in device_name:
device_reference_folder = "H100" + device_reference_folder_suffix
elif "H200" in device_name:
device_reference_folder = "H200" + device_reference_folder_suffix
else:
logger.warning(f"Unsupported device for ssim tests: {device_name}")
@@ -22,10 +22,6 @@ if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
elif "H100" in device_name:
device_reference_folder = "H100" + device_reference_folder_suffix
elif "H200" in device_name:
device_reference_folder = "H200" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
logger.warning(f"Unsupported device for ssim tests: {device_name}")
@@ -1,195 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import os
import shlex
import logging
import pytest
import torch
logger = logging.getLogger(__name__)
def _device_reference_folder() -> str:
"""Pick a reference folder name based on the current CUDA device."""
suffix = "_reference_videos"
device_name = torch.cuda.get_device_name(0)
if "A40" in device_name:
return "A40" + suffix
if "L40S" in device_name:
return "L40S" + suffix
if "H100" in device_name:
return "H100" + suffix
if "H200" in device_name:
return "H200" + suffix
if "RTX 4090" in device_name or "4090" in device_name:
return "RTX4090" + suffix
logger.warning(
"Unsupported device for ssim tests: %s; using L40S references", device_name
)
return "L40S" + suffix
SD35_MODEL_PATH = os.getenv(
"SD35_MODEL_DIR",
"/FastVideo/official_weights/stabilityai__stable-diffusion-3.5-medium",
)
MODEL_ID = "stabilityai__stable-diffusion-3.5-medium"
TEST_PROMPTS = [
"a photo of a cat",
]
@pytest.mark.skipif(not torch.cuda.is_available(), reason="SD3.5 SSIM test requires CUDA")
@pytest.mark.parametrize("ATTENTION_BACKEND", ["TORCH_SDPA"])
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
def test_sd35_similarity(prompt: str, ATTENTION_BACKEND: str) -> None:
from fastvideo import VideoGenerator
from fastvideo.tests.utils import (
compute_video_ssim_torchvision,
write_ssim_results,
)
if not os.path.isdir(SD35_MODEL_PATH):
pytest.skip(
f"SD3.5 weights not found at {SD35_MODEL_PATH} (set SD35_MODEL_DIR to override)"
)
old_backend = os.environ.get("FASTVIDEO_ATTENTION_BACKEND")
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
try:
script_dir = os.path.dirname(os.path.abspath(__file__))
output_dir = os.path.join(
script_dir, "generated_videos", MODEL_ID, ATTENTION_BACKEND
)
os.makedirs(output_dir, exist_ok=True)
prompt_prefix = prompt[:100].strip()
output_video_name = f"{prompt_prefix}.mp4"
expected_video_path = os.path.join(output_dir, output_video_name)
for filename in os.listdir(output_dir):
if filename.endswith(".mp4") and filename.startswith(prompt_prefix):
try:
os.remove(os.path.join(output_dir, filename))
except FileNotFoundError:
pass
init_kwargs = {
"num_gpus": 1,
"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,
}
num_inference_steps = 8
generation_kwargs = {
"output_path": output_dir,
"height": 256,
"width": 256,
"num_frames": 1,
"fps": 1,
"num_inference_steps": num_inference_steps,
"guidance_scale": 6.0,
"seed": 0,
"negative_prompt": "lowres, blurry",
"save_video": True,
}
generator = VideoGenerator.from_pretrained(model_path=SD35_MODEL_PATH, **init_kwargs)
try:
generator.generate_video(prompt, **generation_kwargs)
finally:
generator.shutdown()
generated_video_path = None
if os.path.exists(expected_video_path):
generated_video_path = expected_video_path
else:
candidates = [
os.path.join(output_dir, f)
for f in os.listdir(output_dir)
if f.endswith(".mp4") and f.startswith(prompt_prefix)
]
if candidates:
generated_video_path = max(candidates, key=os.path.getmtime)
assert generated_video_path is not None and os.path.exists(generated_video_path), (
f"Output video was not generated under {output_dir} for prompt '{prompt}'"
)
device_reference_folder = _device_reference_folder()
reference_folder = os.path.join(script_dir, device_reference_folder, MODEL_ID, ATTENTION_BACKEND)
if not os.path.exists(reference_folder):
bless_cmd = (
f"mkdir -p {shlex.quote(reference_folder)} && "
f"cp {shlex.quote(generated_video_path)} {shlex.quote(reference_folder)}/"
)
pytest.fail(
f"Reference folder does not exist: {reference_folder}\n"
f"Generated video saved at: {generated_video_path}\n"
"To bless references, copy the generated mp4 into the reference folder with a matching name.\n"
f"Example:\n {bless_cmd}"
)
reference_video_path = os.path.join(reference_folder, output_video_name)
if not os.path.exists(reference_video_path):
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith(".mp4") and filename.startswith(prompt_prefix):
reference_video_name = filename
break
if not reference_video_name:
bless_cmd = (
f"cp {shlex.quote(generated_video_path)} {shlex.quote(reference_folder)}/"
)
pytest.fail(
f"Reference video not found for prompt '{prompt}' under: {reference_folder}\n"
f"Expected name: {output_video_name}\n"
f"Generated video saved at: {generated_video_path}\n"
"To bless references, copy the generated mp4 into the reference folder.\n"
f"Example:\n {bless_cmd}"
)
reference_video_path = os.path.join(reference_folder, reference_video_name)
logger.info("Computing SSIM between %s and %s", reference_video_path, generated_video_path)
ssim_values = compute_video_ssim_torchvision(
reference_video_path, generated_video_path, use_ms_ssim=True
)
mean_ssim = ssim_values[0]
logger.info("SSIM mean value: %s", mean_ssim)
write_ssim_results(
output_dir,
ssim_values,
reference_video_path,
generated_video_path,
num_inference_steps,
prompt,
)
min_acceptable_ssim = 0.98
assert mean_ssim >= min_acceptable_ssim, (
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} "
f"for {MODEL_ID} with backend {ATTENTION_BACKEND}"
)
finally:
if old_backend is None:
os.environ.pop("FASTVIDEO_ATTENTION_BACKEND", None)
else:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = old_backend
@@ -121,9 +121,9 @@ def test_distributed_training():
wandb_summary = json.load(open(summary_file))
fields_and_thresholds = {
'avg_step_time': 15,
'avg_step_time': 5,
'grad_norm': 0.5,
'step_time': 15,
'step_time': 5,
'train_loss': 0.04
}
@@ -136,9 +136,9 @@ def test_distributed_training():
wandb_summary = json.load(open(summary_file))
fields_and_thresholds = {
'avg_step_time': 15.0,
'avg_step_time': 6.0,
'grad_norm': 0.3,
'step_time': 15.0,
'step_time': 6.0,
'train_loss': 0.0025
}
@@ -1,141 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Regression test for HunyuanGameCraft transformer.
Tests that the FastVideo HunyuanGameCraft implementation produces consistent outputs.
"""
import os
import pytest
import torch
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.distributed.parallel_state import (
get_sp_parallel_rank,
get_sp_world_size,
)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29507"
os.environ.setdefault("TORCHDYNAMO_DISABLE", "1")
os.environ.setdefault("DISABLE_SP", "1")
# Path to converted weights (local path, not HuggingFace)
TRANSFORMER_PATH = "official_weights/hunyuan-gamecraft/transformer"
# Reference latent computed from FastVideo HunyuanGameCraft model with:
# - seed=42, batch=1, frames=9, H=44, W=80, text_seq=32
# - 33 camera frames at 352x640 resolution
# - timestep=500
REFERENCE_LATENT = 42351.12903189659
@pytest.mark.usefixtures("distributed_setup")
def test_hunyuangamecraft_transformer():
"""Test HunyuanGameCraft transformer regression."""
if not os.path.exists(TRANSFORMER_PATH):
pytest.skip(f"Weights not found at {TRANSFORMER_PATH}")
sp_rank = get_sp_parallel_rank()
sp_world_size = get_sp_world_size()
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(
model_path=TRANSFORMER_PATH,
dit_cpu_offload=False,
use_fsdp_inference=False,
pipeline_config=PipelineConfig(
dit_config=HunyuanGameCraftConfig(), dit_precision=precision_str
),
)
args.device = device
loader = TransformerLoader()
model = loader.load(TRANSFORMER_PATH, args).to(device, dtype=precision)
model.eval()
total_params = sum(p.numel() for p in model.parameters())
weight_sum = sum(p.to(torch.float64).sum().item() for p in model.parameters())
weight_mean = weight_sum / total_params
logger.info("Total parameters: %s", total_params)
logger.info("Weight sum: %s", weight_sum)
logger.info("Weight mean: %s", weight_mean)
torch.manual_seed(42)
batch_size = 1
latent_frames = 9 # GameCraft uses 9 latent frames
latent_height = 44
latent_width = 80
text_seq_len = 32
# Input latents [B, 33, T, H, W] - 16 latent + 16 gt_latent + 1 mask
hidden_states = torch.randn(
batch_size, 33, latent_frames, latent_height, latent_width,
device=device, dtype=precision
)
if sp_world_size > 1:
chunk_per_rank = hidden_states.shape[2] // sp_world_size
hidden_states = hidden_states[:, :, sp_rank * chunk_per_rank:(sp_rank + 1) * chunk_per_rank]
# Text embeddings (LLaMA)
text_states = torch.randn(batch_size, text_seq_len, 4096, device=device, dtype=precision)
# CLIP pooled embeddings
text_states_2 = torch.randn(batch_size, 768, device=device, dtype=precision)
# Text mask
text_mask = torch.ones(batch_size, text_seq_len, device=device, dtype=torch.long)
# Timestep
timestep = torch.tensor([500], device=device, dtype=precision)
# Camera states [B, num_frames, 6, H, W] - 33 video frames at full resolution
# For 9 latent frames, we use 33 camera frames (matching official model)
camera_states = torch.randn(
batch_size, 33, 6, 352, 640,
device=device, dtype=precision
)
encoder_hidden_states = [text_states, text_states_2]
forward_batch = ForwardBatch(data_type="video", enable_teacache=False)
with torch.no_grad():
with torch.amp.autocast('cuda', dtype=precision):
with set_forward_context(current_timestep=0, attn_metadata=None, forward_batch=forward_batch):
output = model(
hidden_states,
encoder_hidden_states,
timestep,
camera_states=camera_states,
encoder_attention_mask=[text_mask],
)
latent = output.double().sum().item()
logger.info(f"Output shape: {output.shape}")
logger.info(f"Current latent: {latent}")
if REFERENCE_LATENT is not None:
diff = abs(REFERENCE_LATENT - latent)
relative_diff = diff / abs(REFERENCE_LATENT)
logger.info(f"Reference latent: {REFERENCE_LATENT}")
logger.info(f"Absolute diff: {diff}, Relative diff: {relative_diff * 100:.4f}%")
assert relative_diff < 0.005, \
f"Output latents differ significantly: relative diff = {relative_diff * 100:.4f}% (max allowed: 0.5%)"
else:
logger.info(f"No reference latent set. To enable regression testing, set REFERENCE_LATENT = {latent}")
+1 -1
View File
@@ -153,4 +153,4 @@ def test_hyworld_transformer():
logger.info(f"Absolute diff: {diff}, Relative diff: {relative_diff * 100:.4f}%")
# Allow 0.5% relative difference
assert relative_diff < 0.006, f"Output latents differ significantly: relative diff = {relative_diff * 100:.4f}% (max allowed: 0.6%)"
assert relative_diff < 0.005, f"Output latents differ significantly: relative diff = {relative_diff * 100:.4f}% (max allowed: 0.5%)"
@@ -402,7 +402,6 @@ class LTX2TrainingPipeline(TrainingPipeline):
with torch.autocast("cuda", dtype=training_batch.latents.dtype
), torch.autograd.set_detect_anomaly(True):
outputs = self.transformer(**input_kwargs)
if isinstance(outputs, tuple):
video_denoised, audio_denoised = outputs
else:
-2
View File
@@ -19,7 +19,6 @@ from einops import rearrange
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
from diffusers import FlowMatchEulerDiscreteScheduler
import fastvideo.envs as envs
try:
@@ -606,7 +605,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
device="cpu").manual_seed(self.seed + self.global_rank)
logger.info("Initialized random seeds with seed: %s",
self.seed + self.global_rank)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()

Before

Width:  |  Height:  |  Size: 113 KiB

After

Width:  |  Height:  |  Size: 113 KiB

Before

Width:  |  Height:  |  Size: 1.2 MiB

After

Width:  |  Height:  |  Size: 1.2 MiB

Before

Width:  |  Height:  |  Size: 229 KiB

After

Width:  |  Height:  |  Size: 229 KiB

Before

Width:  |  Height:  |  Size: 168 KiB

After

Width:  |  Height:  |  Size: 168 KiB

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