Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a8dddfaa16 | ||
|
|
7210c68f1b | ||
|
|
5602dc1bad | ||
|
|
ac4bc4ab84 | ||
|
|
bee27f9f74 | ||
|
|
dff0ea401a | ||
|
|
becd379f58 | ||
|
|
1ed7d7e1b0 | ||
|
|
d9fabcc5ef | ||
|
|
9db48498de | ||
|
|
1cd7038315 |
@@ -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
|
||||
|
||||
@@ -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/
|
||||
|
||||
@@ -10,7 +10,7 @@ exclude: |
|
||||
demo/.*|
|
||||
predict\.py|
|
||||
scripts/.*|
|
||||
assets/prompts/.*|
|
||||
prompts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/models/.*|
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.**
|
||||
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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,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"]
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -60,61 +60,6 @@ class PatchEmbed(nn.Module):
|
||||
return x
|
||||
|
||||
|
||||
class WanCamControlPatchEmbedding(nn.Module):
|
||||
"""Lingbot World Patch embedding for Plucker features."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=(1, 2, 2),
|
||||
in_chans=384, # 6 * 64
|
||||
embed_dim=2048,
|
||||
bias=True,
|
||||
dtype=None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
# must be 3-tuple
|
||||
if isinstance(patch_size, list | tuple):
|
||||
if len(patch_size) != 3:
|
||||
raise ValueError(
|
||||
f"patch_size must have length 3, got {len(patch_size)}")
|
||||
else:
|
||||
raise ValueError(f"Unsupported patch_size type: {type(patch_size)}")
|
||||
|
||||
self.patch_size = patch_size
|
||||
pt, ph, pw = self.patch_size
|
||||
self.in_features = in_chans * pt * ph * pw
|
||||
self.proj = nn.Linear(self.in_features,
|
||||
embed_dim,
|
||||
bias=bias,
|
||||
dtype=dtype)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if x.dim() != 5:
|
||||
raise ValueError(
|
||||
f"Expected camera embedding shape [B, C, F, H, W], got {x.shape}"
|
||||
)
|
||||
bsz, channels, frames, height, width = x.shape
|
||||
pt, ph, pw = self.patch_size
|
||||
if (frames % pt) != 0 or (height % ph) != 0 or (width % pw) != 0:
|
||||
raise ValueError(
|
||||
f"Input shape {x.shape} must be divisible by patch_size {self.patch_size}"
|
||||
)
|
||||
|
||||
# '1 c (f c1) (h c2) (w c3) -> 1 (f h w) (c c1 c2 c3)',
|
||||
x = x.view(
|
||||
bsz,
|
||||
channels,
|
||||
frames // pt,
|
||||
pt,
|
||||
height // ph,
|
||||
ph,
|
||||
width // pw,
|
||||
pw,
|
||||
)
|
||||
x = x.permute(0, 2, 4, 6, 1, 3, 5, 7).reshape(bsz, -1, self.in_features)
|
||||
return self.proj(x)
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
@@ -307,4 +252,4 @@ class Timesteps(nn.Module):
|
||||
downscale_freq_shift=self.downscale_freq_shift,
|
||||
scale=self.scale,
|
||||
)
|
||||
return t_emb
|
||||
return t_emb
|
||||
@@ -1,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"]
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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__(
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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}")
|
||||
@@ -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:
|
||||
|
||||
@@ -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 |