Compare commits

...
41 changed files with 2627 additions and 11 deletions
@@ -0,0 +1,18 @@
Total Files: 16
00. Hold [W] + Static
01. Hold [S] + Static
02. Hold [A] + Static
03. Hold [D] + Static
04. Hold [WA] + Static
05. Hold [WD] + Static
06. Hold [SA] + Static
07. Hold [SD] + Static
08. No Key + Hold [up]
09. No Key + Hold [down]
10. No Key + Hold [left]
11. No Key + Hold [right]
12. No Key + Hold [up_right]
13. No Key + Hold [up_left]
14. No Key + Hold [down_right]
15. No Key + Hold [down_left]
@@ -0,0 +1,93 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export TOKENIZERS_PARALLELISM=false
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
MODEL_PATH="weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers"
DATA_DIR="mc_wasd_10/preprocessed/combined_parquet_dataset"
VALIDATION_DATASET_FILE="mc_wasd_10/validation.json"
NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=0,1,2,3
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wangame_1.3b_overfit"
--output_dir "wangame_1.3b_overfit"
--max_train_steps 1500
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 20
--num_height 352
--num_width 640
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wangame_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,120 @@
#!/bin/bash
#SBATCH --job-name=wangame_1.3b_overfit
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=wangame_1.3b_overfit_output/wangame_1.3b_overfit_%j.out
#SBATCH --error=wangame_1.3b_overfit_output/wangame_1.3b_overfit_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export NCCL_DEBUG_SUBSYS=INIT,NET
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export TOKENIZERS_PARALLELISM=false
# export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
# export WANDB_API_KEY='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
source ~/conda/miniconda/bin/activate
conda activate wei-fv-distill
export HOME="/mnt/weka/home/hao.zhang/wei"
MODEL_PATH="weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers"
DATA_DIR="mc_wasd_10/preprocessed/combined_parquet_dataset"
VALIDATION_DATASET_FILE="examples/training/finetune/WanGame2.1_1.3b_i2v/validation.json"
# Configs
NUM_GPUS=8
# Training arguments
training_args=(
--tracker_project_name "wangame_1.3b_overfit"
--output_dir "wangame_1.3b_overfit"
--max_train_steps 15000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 20
--num_height 352
--num_width 640
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000000
--training_state_checkpointing_steps 10000000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wangame_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,193 @@
import os
import numpy as np
# Configuration
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
BASE_OUTPUT_DIR = os.path.join(SCRIPT_DIR, 'action')
VIDEO_OUTPUT_DIR = BASE_OUTPUT_DIR
os.makedirs(VIDEO_OUTPUT_DIR, exist_ok=True)
FRAME_COUNT = 81
CAM_VALUE = 0.1
# Action Mapping
KEY_TO_INDEX = {
'W': 0, 'S': 1, 'A': 2, 'D': 3,
}
VIEW_ACTION_TO_MOUSE = {
"stop": [0.0, 0.0],
"up": [CAM_VALUE, 0.0],
"down": [-CAM_VALUE, 0.0],
"left": [0.0, -CAM_VALUE],
"right": [0.0, CAM_VALUE],
"up_right": [CAM_VALUE, CAM_VALUE],
"up_left": [CAM_VALUE, -CAM_VALUE],
"down_right": [-CAM_VALUE, CAM_VALUE],
"down_left": [-CAM_VALUE, -CAM_VALUE],
}
def get_multihot_vector(keys_str):
"""Convert string like 'WA' to [1, 0, 1, 0, 0, 0]"""
vector = [0.0] * 6
if not keys_str:
return vector
for char in keys_str.upper():
if char in KEY_TO_INDEX:
vector[KEY_TO_INDEX[char]] = 1.0
return vector
def get_mouse_vector(view_str):
"""Convert view string to [x, y]"""
return VIEW_ACTION_TO_MOUSE.get(view_str.lower(), [0.0, 0.0])
def generate_sequence(key_seq, mouse_seq):
"""
Generates action arrays based on sequences.
"""
keyboard_arr = np.zeros((FRAME_COUNT, 6), dtype=np.float32)
mouse_arr = np.zeros((FRAME_COUNT, 2), dtype=np.float32)
mid_point = FRAME_COUNT // 2
# First Half
k_vec1 = get_multihot_vector(key_seq[0])
m_vec1 = get_mouse_vector(mouse_seq[0])
keyboard_arr[:mid_point] = k_vec1
mouse_arr[:mid_point] = m_vec1
# Second Half
k_vec2 = get_multihot_vector(key_seq[1])
m_vec2 = get_mouse_vector(mouse_seq[1])
keyboard_arr[mid_point:] = k_vec2
mouse_arr[mid_point:] = m_vec2
return keyboard_arr, mouse_arr
def save_action(index, keyboard_arr, mouse_arr):
filename = f"{index:06d}_action.npy"
filepath = os.path.join(VIDEO_OUTPUT_DIR, filename)
action_dict = {
'keyboard': keyboard_arr,
'mouse': mouse_arr
}
np.save(filepath, action_dict)
return filename
def generate_description(key_seq, mouse_seq):
"""Generates a human-readable string for the combination."""
k1, k2 = key_seq
m1, m2 = mouse_seq
# Format Keyboard Description
if not k1 and not k2:
k_desc = "No Key"
elif k1 == k2:
k_desc = f"Hold [{k1}]"
else:
k_desc = f"Switch [{k1}]->[{k2}]"
# Format Mouse Description
if m1 == "stop" and m2 == "stop":
m_desc = "Static"
elif m1 == m2:
m_desc = f"Hold [{m1}]"
else:
m_desc = f"Switch [{m1}]->[{m2}]"
return f"{k_desc} + {m_desc}"
# ==========================================
# Main Generation Logic
# ==========================================
configs = []
readme_content = []
# Group 1: Constant Keyboard, No Mouse (0-7)
keys_basic = ['W', 'S', 'A', 'D', 'WA', 'WD', 'SA', 'SD']
for k in keys_basic:
configs.append(((k, k), ("stop", "stop")))
# Group 2: No Keyboard, Constant Mouse (8-15)
mouse_basic = ['up', 'down', 'left', 'right', 'up_right', 'up_left', 'down_right', 'down_left']
for m in mouse_basic:
configs.append((("", ""), (m, m)))
# Group 3: Split Keyboard, No Mouse (16-23)
split_keys = [
('W', 'S'), ('S', 'W'),
('A', 'D'), ('D', 'A'),
('W', 'A'), ('W', 'D'),
('S', 'A'), ('S', 'D')
]
for k1, k2 in split_keys:
configs.append(((k1, k2), ("stop", "stop")))
# Group 4: No Keyboard, Split Mouse (24-31)
split_mouse = [
('left', 'right'), ('right', 'left'),
('up', 'down'), ('down', 'up'),
('up_left', 'up_right'), ('up_right', 'up_left'),
('left', 'up'), ('right', 'down')
]
for m1, m2 in split_mouse:
configs.append((("", ""), (m1, m2)))
# Group 5: Constant Keyboard + Constant Mouse (32-47)
combo_keys = ['W', 'S', 'W', 'S', 'A', 'D', 'WA', 'WD', 'W', 'S', 'W', 'S', 'A', 'D', 'WA', 'WD']
combo_mice = ['left', 'left', 'right', 'right', 'up', 'up', 'down', 'down', 'up_left', 'up_left', 'up_right', 'up_right', 'down_left', 'down_right', 'right', 'left']
for i in range(16):
configs.append(((combo_keys[i], combo_keys[i]), (combo_mice[i], combo_mice[i])))
# Group 6: Constant Keyboard, Split Mouse (48-55)
complex_1_keys = ['W'] * 8
complex_1_mice = [
('left', 'right'), ('right', 'left'),
('up', 'down'), ('down', 'up'),
('left', 'up'), ('right', 'up'),
('left', 'down'), ('right', 'down')
]
for i in range(8):
configs.append(((complex_1_keys[i], complex_1_keys[i]), complex_1_mice[i]))
# Group 7: Split Keyboard, Constant Mouse (56-63)
complex_2_keys = [
('W', 'S'), ('S', 'W'),
('A', 'D'), ('D', 'A'),
('W', 'A'), ('W', 'D'),
('S', 'A'), ('S', 'D')
]
complex_2_mouse = 'up'
for k1, k2 in complex_2_keys:
configs.append(((k1, k2), (complex_2_mouse, complex_2_mouse)))
# Execution
print(f"Preparing to generate {len(configs)} action files...")
for i, (key_seq, mouse_seq) in enumerate(configs):
if i >= 16: break
# Generate Data
kb_arr, ms_arr = generate_sequence(key_seq, mouse_seq)
filename = save_action(i, kb_arr, ms_arr)
# Generate Description for README
description = generate_description(key_seq, mouse_seq)
readme_entry = f"{i:02d}. {description}"
readme_content.append(readme_entry)
print(f"Generated {filename} -> {description}")
# Write README
readme_path = os.path.join(BASE_OUTPUT_DIR, 'README.md')
with open(readme_path, 'w', encoding='utf-8') as f:
f.write(f"Total Files: {len(readme_content)}\n\n")
for line in readme_content:
f.write(line + '\n')
print(f"\nProcessing complete.")
print(f"64 .npy files generated in {VIDEO_OUTPUT_DIR}")
print(f"Manifest saved to {readme_path}")
@@ -0,0 +1,27 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
DATA_MERGE_PATH="mc_wasd_10/merge.txt"
OUTPUT_DIR="mc_wasd_10/preprocessed/"
# export CUDA_VISIBLE_DEVICES=0
export MASTER_ADDR=localhost
export MASTER_PORT=29500
export RANK=0
export WORLD_SIZE=1
python fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 10 \
--seed 42 \
--max_height 352 \
--max_width 640 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--samples_per_file 10 \
--train_fps 25 \
--flush_frequency 10 \
--preprocess_task wangame
@@ -0,0 +1,404 @@
{
"data": [
{
"caption": "0",
"image_path": "../../../../mc_wasd_10/validate/000000.jpg",
"action_path": "../../../../mc_wasd_10/videos/000000_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "1",
"image_path": "../../../../mc_wasd_10/validate/000001.jpg",
"action_path": "../../../../mc_wasd_10/videos/000001_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "2",
"image_path": "../../../../mc_wasd_10/validate/000002.jpg",
"action_path": "../../../../mc_wasd_10/videos/000002_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "3",
"image_path": "../../../../mc_wasd_10/validate/000003.jpg",
"action_path": "../../../../mc_wasd_10/videos/000003_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "4",
"image_path": "../../../../mc_wasd_10/validate/000004.jpg",
"action_path": "../../../../mc_wasd_10/videos/000004_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "5",
"image_path": "../../../../mc_wasd_10/validate/000005.jpg",
"action_path": "../../../../mc_wasd_10/videos/000005_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "6",
"image_path": "../../../../mc_wasd_10/validate/000006.jpg",
"action_path": "../../../../mc_wasd_10/videos/000006_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "7",
"image_path": "../../../../mc_wasd_10/validate/000007.jpg",
"action_path": "../../../../mc_wasd_10/videos/000007_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "00. Hold [W] + Static",
"image_path": "../../../../mc_wasd_10/validate/000000.jpg",
"action_path": "action/000000_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "01. Hold [S] + Static",
"image_path": "../../../../mc_wasd_10/validate/000001.jpg",
"action_path": "action/000001_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "02. Hold [A] + Static",
"image_path": "../../../../mc_wasd_10/validate/000002.jpg",
"action_path": "action/000002_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "03. Hold [D] + Static",
"image_path": "../../../../mc_wasd_10/validate/000003.jpg",
"action_path": "action/000003_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "04. Hold [WA] + Static",
"image_path": "../../../../mc_wasd_10/validate/000004.jpg",
"action_path": "action/000004_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "05. Hold [WD] + Static",
"image_path": "../../../../mc_wasd_10/validate/000005.jpg",
"action_path": "action/000005_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "06. Hold [SA] + Static",
"image_path": "../../../../mc_wasd_10/validate/000006.jpg",
"action_path": "action/000006_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "07. Hold [SD] + Static",
"image_path": "../../../../mc_wasd_10/validate/000007.jpg",
"action_path": "action/000007_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "08. No Key + Hold [up]",
"image_path": "../../../../mc_wasd_10/validate/000000.jpg",
"action_path": "action/000008_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "09. No Key + Hold [down]",
"image_path": "../../../../mc_wasd_10/validate/000001.jpg",
"action_path": "action/000009_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "10. No Key + Hold [left]",
"image_path": "../../../../mc_wasd_10/validate/000002.jpg",
"action_path": "action/000010_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "11. No Key + Hold [right]",
"image_path": "../../../../mc_wasd_10/validate/000003.jpg",
"action_path": "action/000011_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "12. No Key + Hold [up_right]",
"image_path": "../../../../mc_wasd_10/validate/000004.jpg",
"action_path": "action/000012_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "13. No Key + Hold [up_left]",
"image_path": "../../../../mc_wasd_10/validate/000005.jpg",
"action_path": "action/000013_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "14. No Key + Hold [down_right]",
"image_path": "../../../../mc_wasd_10/validate/000006.jpg",
"action_path": "action/000014_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "15. No Key + Hold [down_left]",
"image_path": "../../../../mc_wasd_10/validate/000007.jpg",
"action_path": "action/000015_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "00. Hold [W] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000000.jpg",
"action_path": "action/000000_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "01. Hold [S] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000001.jpg",
"action_path": "action/000001_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "02. Hold [A] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000002.jpg",
"action_path": "action/000002_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "03. Hold [D] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000003.jpg",
"action_path": "action/000003_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "04. Hold [WA] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000004.jpg",
"action_path": "action/000004_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "05. Hold [WD] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000005.jpg",
"action_path": "action/000005_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "06. Hold [SA] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000006.jpg",
"action_path": "action/000006_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "07. Hold [SD] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000007.jpg",
"action_path": "action/000007_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "08. No Key + Hold [up]",
"image_path": "../../../../mc_wasd_10/validate/gen_000000.jpg",
"action_path": "action/000008_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "09. No Key + Hold [down]",
"image_path": "../../../../mc_wasd_10/validate/gen_000001.jpg",
"action_path": "action/000009_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "10. No Key + Hold [left]",
"image_path": "../../../../mc_wasd_10/validate/gen_000002.jpg",
"action_path": "action/000010_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "11. No Key + Hold [right]",
"image_path": "../../../../mc_wasd_10/validate/gen_000003.jpg",
"action_path": "action/000011_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "12. No Key + Hold [up_right]",
"image_path": "../../../../mc_wasd_10/validate/gen_000004.jpg",
"action_path": "action/000012_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "13. No Key + Hold [up_left]",
"image_path": "../../../../mc_wasd_10/validate/gen_000005.jpg",
"action_path": "action/000013_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "14. No Key + Hold [down_right]",
"image_path": "../../../../mc_wasd_10/validate/gen_000006.jpg",
"action_path": "action/000014_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "15. No Key + Hold [down_left]",
"image_path": "../../../../mc_wasd_10/validate/gen_000007.jpg",
"action_path": "action/000015_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
}
]
}
@@ -0,0 +1,115 @@
# 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 WanGameVideoArchConfig(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"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$":
r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_out.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.norm2\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
})
# 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: {
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$":
r"blocks.\1.attn1.to_out.0.\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.0.\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",
})
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len = 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 WanGameVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=WanGameVideoArchConfig)
prefix: str = "WanGame"
+1
View File
@@ -42,6 +42,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers": HYWorldConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
"weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers": WANV2VConfig,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
+2
View File
@@ -58,6 +58,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
Wan2_1_Fun_1_3B_Control_SamplingParam,
+6 -6
View File
@@ -113,13 +113,13 @@ class FastWanT2V480P_SamplingParam(WanT2V_1_3B_SamplingParam):
@dataclass
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
height: int = 352
width: int = 640
num_frames: int = 77
fps: int = 25
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale: float = 6.0
num_inference_steps: int = 50
guidance_scale: float = 1.0
num_inference_steps: int = 40
@dataclass
+40
View File
@@ -157,3 +157,43 @@ pyarrow_schema_matrixgame = pa.schema([
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_wangame = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
pa.field("first_frame_latent_bytes", pa.binary()),
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
pa.field("first_frame_latent_dtype", pa.string()),
# --- Action ---
pa.field("mouse_cond_bytes", pa.binary()),
pa.field("mouse_cond_shape", pa.list_(pa.int64())), # [T, 2]
pa.field("mouse_cond_dtype", pa.string()),
pa.field("keyboard_cond_bytes", pa.binary()),
pa.field("keyboard_cond_shape", pa.list_(pa.int64())), # [T, 4]
pa.field("keyboard_cond_dtype", pa.string()),
# I2V Validation
pa.field("pil_image_bytes", pa.binary()),
pa.field("pil_image_shape", pa.list_(pa.int64())),
pa.field("pil_image_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
+21
View File
@@ -4,6 +4,7 @@ import os
import pathlib
import datasets
import numpy as np
from torch.utils.data import IterableDataset
from fastvideo.distributed import (get_sp_world_size, get_world_rank,
@@ -160,5 +161,25 @@ class ValidationDataset(IterableDataset):
else:
sample["control_video"] = load_video(control_video_path)
if sample.get("action_path", None) is not None:
action_path = sample["action_path"]
action_path = os.path.join(self.dir, action_path)
sample["action_path"] = action_path
if not pathlib.Path(action_path).is_file():
logger.warning("Action file %s does not exist.", action_path)
else:
try:
action_data = np.load(action_path, allow_pickle=True)
num_frames = sample["num_frames"]
if action_data.dtype == object: action_data = action_data.item()
if isinstance(action_data, dict):
sample["keyboard_cond"] = action_data["keyboard"][:num_frames]
sample["mouse_cond"] = action_data["mouse"][:num_frames]
else:
sample["keyboard_cond"] = action_data[:num_frames]
except Exception as e:
logger.error("Error loading action file %s: %s",
action_path, e)
sample = {k: v for k, v in sample.items() if v is not None}
yield sample
+267 -1
View File
@@ -15,7 +15,7 @@ import torch
from scipy.spatial.transform import Rotation as R
from typing import Union, Optional
from .trajectory import generate_camera_trajectory_local
from fastvideo.models.dits.hyworld.trajectory import generate_camera_trajectory_local
# Mapping from one-hot action encoding to single label
@@ -411,3 +411,269 @@ def compute_num_frames(latent_num: int) -> int:
Number of video frames
"""
return (latent_num - 1) * 4 + 1
def reformat_keyboard_and_mouse_tensors(keyboard_tensor, mouse_tensor):
"""
Reformat the keyboard and mouse tensors to the format compatible with HyWorld.
"""
num_frames = keyboard_tensor.shape[0]
assert (num_frames - 1) % 4 == 0, "num_frames must be a multiple of 4"
assert mouse_tensor.shape[0] == num_frames, "mouse_tensor must have the same number of frames as keyboard_tensor"
keyboard_tensor = keyboard_tensor[1:, :]
mouse_tensor = mouse_tensor[1:, :]
groups = keyboard_tensor.view(-1, 4, keyboard_tensor.shape[1])
assert (groups == groups[:, 0:1]).all(dim=1).all(), "keyboard_tensor must have the same value for each group"
groups = mouse_tensor.view(-1, 4, mouse_tensor.shape[1])
assert (groups == groups[:, 0:1]).all(dim=1).all(), "mouse_tensor must have the same value for each group"
return keyboard_tensor[::4], mouse_tensor[::4]
def process_custom_actions(keyboard_tensor, mouse_tensor, forward_speed=DEFAULT_FORWARD_SPEED):
"""
Process custom keyboard and mouse tensors into model inputs (viewmats, intrinsics, action_labels).
Assumes inputs correspond to each LATENT frame.
"""
if keyboard_tensor.ndim == 3:
keyboard_tensor = keyboard_tensor.squeeze(0)
if mouse_tensor.ndim == 3:
mouse_tensor = mouse_tensor.squeeze(0)
keyboard_tensor, mouse_tensor = reformat_keyboard_and_mouse_tensors(keyboard_tensor, mouse_tensor)
motions = []
# 1. Translate tensors to motions for trajectory generation
for t in range(keyboard_tensor.shape[0]):
frame_motion = {}
# --- Translation ---
# MatrixGame convention: 0:W, 1:S, 2:A, 3:D
fwd = 0.0
if keyboard_tensor[t, 0] > 0.5: fwd += forward_speed # W
if keyboard_tensor[t, 1] > 0.5: fwd -= forward_speed # S
if fwd != 0: frame_motion["forward"] = fwd
rgt = 0.0
if keyboard_tensor[t, 2] > 0.5: rgt -= forward_speed # A (Left is negative Right)
if keyboard_tensor[t, 3] > 0.5: rgt += forward_speed # D (Right)
if rgt != 0: frame_motion["right"] = rgt
# --- Rotation ---
# MatrixGame convention: mouse is [Pitch, Yaw] (or Y, X)
# Apply scaling (e.g. to match HyWorld distribution)
pitch = mouse_tensor[t, 0].item()
yaw = mouse_tensor[t, 1].item()
if abs(pitch) > 1e-4: frame_motion["pitch"] = pitch
if abs(yaw) > 1e-4: frame_motion["yaw"] = yaw
motions.append(frame_motion)
# 2. Generate Camera Trajectory
# generate_camera_trajectory_local returns T+1 poses (starting at Identity)
# We take the first T poses to match the latent count.
# Pose 0 is Identity. Pose 1 is Identity + Motion[0].
poses = generate_camera_trajectory_local(motions)
# poses = np.array(poses[:T])
# 3. Compute Viewmats (w2c) and Intrinsics
w2c_list = []
intrinsic_list = []
# Setup default intrinsic (normalized)
K = np.array(DEFAULT_INTRINSIC)
K[0, 0] /= K[0, 2] * 2
K[1, 1] /= K[1, 2] * 2
K[0, 2] = 0.5
K[1, 2] = 0.5
for i in range(len(poses)):
c2w = np.array(poses[i])
w2c = np.linalg.inv(c2w)
w2c_list.append(w2c)
intrinsic_list.append(K)
viewmats = torch.as_tensor(np.array(w2c_list))
intrinsics = torch.as_tensor(np.array(intrinsic_list))
# 4. Generate Action Labels by analyzing the generated trajectory
# This ensures consistency with complex simultaneous movements, exactly as pose_to_input does.
# Calculate relative camera-to-world transforms
# c2ws = inverse(viewmats)
c2ws = np.linalg.inv(np.array(w2c_list))
# Calculate relative movement between frames
# relative_c2w[i] = inv(c2ws[i-1]) @ c2ws[i]
C_inv = np.linalg.inv(c2ws[:-1])
relative_c2w = np.zeros_like(c2ws)
relative_c2w[0, ...] = c2ws[0, ...] # First is anchor
relative_c2w[1:, ...] = C_inv @ c2ws[1:, ...]
# Initialize one-hot action encodings
trans_one_hot = np.zeros((relative_c2w.shape[0], 4), dtype=np.int32)
rotate_one_hot = np.zeros((relative_c2w.shape[0], 4), dtype=np.int32)
move_norm_valid = 0.0001
# Skip index 0 (anchor/identity)
for i in range(1, relative_c2w.shape[0]):
move_dirs = relative_c2w[i, :3, 3] # direction vector
move_norms = np.linalg.norm(move_dirs)
if move_norms > move_norm_valid: # threshold for movement
move_norm_dirs = move_dirs / move_norms
angles_rad = np.arccos(move_norm_dirs.clip(-1.0, 1.0))
trans_angles_deg = angles_rad * (180.0 / np.pi) # convert to degrees
else:
trans_angles_deg = np.zeros(3)
R_rel = relative_c2w[i, :3, :3]
r = R.from_matrix(R_rel)
rot_angles_deg = r.as_euler("xyz", degrees=True)
# Determine movement actions based on trajectory
# Note: HyWorld logic checks if rotation is small before assigning translation labels
# to avoid ambiguity in TPS mode, but here we generally want to capture the dominant movement.
tps = False # Default assumption, can be made an arg if needed
if move_norms > move_norm_valid:
if (not tps) or (
tps and abs(rot_angles_deg[1]) < 5e-2 and abs(rot_angles_deg[0]) < 5e-2
):
# Z-axis (Forward/Back)
if trans_angles_deg[2] < 60:
trans_one_hot[i, 0] = 1 # forward
elif trans_angles_deg[2] > 120:
trans_one_hot[i, 1] = 1 # backward
# X-axis (Right/Left)
if trans_angles_deg[0] < 60:
trans_one_hot[i, 2] = 1 # right
elif trans_angles_deg[0] > 120:
trans_one_hot[i, 3] = 1 # left
# Determine rotation actions
# Y-axis (Yaw)
if rot_angles_deg[1] > 5e-2:
rotate_one_hot[i, 0] = 1 # right
elif rot_angles_deg[1] < -5e-2:
rotate_one_hot[i, 1] = 1 # left
# X-axis (Pitch)
if rot_angles_deg[0] > 5e-2:
rotate_one_hot[i, 2] = 1 # up
elif rot_angles_deg[0] < -5e-2:
rotate_one_hot[i, 3] = 1 # down
trans_one_hot = torch.tensor(trans_one_hot)
rotate_one_hot = torch.tensor(rotate_one_hot)
# Convert to single labels
trans_label = one_hot_to_one_dimension(trans_one_hot)
rotate_label = one_hot_to_one_dimension(rotate_one_hot)
action_labels = trans_label * 9 + rotate_label
return viewmats, intrinsics, action_labels
if __name__ == "__main__":
print("Running comparison test between process_custom_actions and pose_to_input...")
def test_process_custom_actions(pose_string: str, keyboard: torch.Tensor, mouse: torch.Tensor, latent_num: int):
# Run process_custom_actions
# Note: We need to pass float tensors
print("Running process_custom_actions...")
viewmats_1, intrinsics_1, labels_1 = process_custom_actions(
keyboard, mouse
)
print(f"Running pose_to_input with string: '{pose_string}'...")
viewmats_2, intrinsics_2, labels_2 = pose_to_input(
pose_string, latent_num=latent_num
)
# print(f"Viewmats: {viewmats_1} vs \n {viewmats_2}")
# print(f"Intrinsics: {intrinsics_1} vs \n {intrinsics_2}")
# print(f"Labels: {labels_1} vs \n {labels_2}")
# 3. Compare Results
print("\nComparison Results:")
# Check Shapes
print(f"Shapes (Viewmats): {viewmats_1.shape} vs {viewmats_2.shape}")
assert viewmats_1.shape == viewmats_2.shape, "Shape mismatch for viewmats"
# Check Values
# Viewmats
diff_viewmats = (viewmats_1 - viewmats_2).abs().max().item()
print(f"Max difference in Viewmats: {diff_viewmats}")
if diff_viewmats < 1e-5:
print("✅ Viewmats match.")
else:
print("❌ Viewmats mismatch.")
# Check intrinsics
diff_intrinsics = (intrinsics_1 - intrinsics_2).abs().max().item()
print(f"Max difference in Intrinsics: {diff_intrinsics}")
if diff_intrinsics < 1e-5:
print("✅ Intrinsics match.")
else:
print("❌ Intrinsics mismatch.")
# Check labels
diff_labels = (labels_1 - labels_2).abs().max().item()
print(f"Max difference in Labels: {diff_labels}")
if diff_labels < 1e-5:
print("✅ Labels match.")
else:
print("❌ Labels mismatch.")
print("All checks passed.")
# Define shared parameters
latent_num = 13
pose_string = "w-2, a-3, s-1, d-6"
num_frames = 4 * (latent_num - 1) + 1
keyboard = torch.zeros((num_frames, 6))
mouse = torch.zeros((num_frames, 2))
# Frame 0 is ignored/start
# Frames 1-8: Press W (index 0)
keyboard[1:9, 0] = 1.0
# Frames 9-20: Press A (index 2)
keyboard[9:21, 2] = 1.0
# Frames 21-24: Press S (index 1)
keyboard[21:25, 1] = 1.0
# Frames 25-48: Press D (index 3)
keyboard[25:49, 3] = 1.0
test_process_custom_actions(pose_string, keyboard, mouse, latent_num)
# Test keyboard AND mouse
latent_num = 25
pose_string = "w-2, up-2, a-3, down-4, s-1, left-2, d-6, right-4"
num_frames = 4 * (latent_num - 1) + 1
keyboard = torch.zeros((num_frames, 6))
mouse = torch.zeros((num_frames, 2))
# Frame 0 is ignored/start
# Frames 1-8: Press W (index 0)
keyboard[1:9, 0] = 1.0
# Frames 17-28: Press A (index 2)
keyboard[17:29, 2] = 1.0
# Frames 45-48: Press S (index 1)
keyboard[45:49, 1] = 1.0
# Frames 57-80: Press D (index 3)
keyboard[57:81, 3] = 1.0
# Frames 9-16: Press Up (index 4)
mouse[9:17, 0] = DEFAULT_PITCH_SPEED
# Frames 25-32: Press Down (index 5)
mouse[29:45, 0] = -DEFAULT_PITCH_SPEED
# Frames 41-48: Press Left (index 6)
mouse[49:57, 1] = -DEFAULT_YAW_SPEED
# Frames 57-64: Press Right (index 7)
mouse[81:, 1] = DEFAULT_YAW_SPEED
test_process_custom_actions(pose_string, keyboard, mouse, latent_num)
@@ -0,0 +1,5 @@
from .model import WanGameActionTransformer3DModel
__all__ = [
"WanGameActionTransformer3DModel",
]
@@ -0,0 +1,231 @@
import torch
import torch.nn as nn
from fastvideo.layers.visual_embedding import TimestepEmbedder, ModulateProjection, timestep_embedding
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.attention import DistributedAttention
from fastvideo.forward_context import set_forward_context
from fastvideo.models.dits.wanvideo import WanImageEmbedding
from fastvideo.models.dits.hyworld.camera_rope import prope_qkv
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
from fastvideo.layers.mlp import MLP
class WanGameActionTimeImageEmbedding(nn.Module):
def __init__(
self,
dim: int,
time_freq_dim: int,
image_embed_dim: int | None = None,
):
super().__init__()
self.time_freq_dim = time_freq_dim
self.time_embedder = TimestepEmbedder(
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
self.time_modulation = ModulateProjection(dim,
factor=6,
act_layer="silu")
self.image_embedder = None
if image_embed_dim is not None:
self.image_embedder = WanImageEmbedding(image_embed_dim, dim)
self.action_embedder = MLP(
time_freq_dim,
dim,
dim,
bias=True,
act_type="silu"
)
# Initialize with zeros for residual-like behavior
nn.init.zeros_(self.action_embedder.fc_out.weight)
if self.action_embedder.fc_out.bias is not None:
nn.init.zeros_(self.action_embedder.fc_out.bias)
def forward(
self,
timestep: torch.Tensor,
action: torch.Tensor,
encoder_hidden_states: torch.Tensor, # Kept for interface compatibility
encoder_hidden_states_image: torch.Tensor | None = None,
timestep_seq_len: int | None = None,
):
temb = self.time_embedder(timestep, timestep_seq_len)
action_emb = timestep_embedding(action.flatten(), self.time_freq_dim)
action_embedder_dtype = next(iter(self.action_embedder.parameters())).dtype
if (
action_emb.dtype != action_embedder_dtype
and action_embedder_dtype != torch.int8
):
action_emb = action_emb.to(action_embedder_dtype)
action_emb = self.action_embedder(action_emb).type_as(temb)
temb = temb + action_emb
timestep_proj = self.time_modulation(temb)
# MatrixGame does not use text embeddings, so we ignore encoder_hidden_states
if encoder_hidden_states_image is not None:
assert self.image_embedder is not None
encoder_hidden_states_image = self.image_embedder(
encoder_hidden_states_image)
encoder_hidden_states = torch.zeros((timestep.shape[0], 0, temb.shape[-1]),
device=temb.device,
dtype=temb.dtype)
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
class WanGameActionSelfAttention(nn.Module):
"""
Self-attention module with support for:
- Standard RoPE-based attention
- Camera PRoPE-based attention (when viewmats and Ks are provided)
- KV caching for autoregressive generation
"""
def __init__(self,
dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm=True,
eps=1e-6) -> None:
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
self.qk_norm = qk_norm
self.eps = eps
self.max_attention_size = 32760 if local_attn_size == -1 else local_attn_size * 1560
# Scaled dot product attention (using DistributedAttention for SP support)
self.attn = DistributedAttention(
num_heads=num_heads,
head_size=self.head_dim,
softmax_scale=None,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
def forward(self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
kv_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None,
viewmats: torch.Tensor | None = None,
Ks: torch.Tensor | None = None,
is_cache: bool = False,
attention_mask: torch.Tensor | None = None):
"""
Forward pass with camera PRoPE attention combining standard RoPE and projective positional encoding.
Args:
q, k, v: Query, key, value tensors [B, L, num_heads, head_dim]
freqs_cis: RoPE frequency cos/sin tensors
kv_cache: KV cache dict (may have None values for training)
current_start: Current position for KV cache
cache_start: Cache start position
viewmats: Camera view matrices for PRoPE [B, cameras, 4, 4]
Ks: Camera intrinsics for PRoPE [B, cameras, 3, 3]
is_cache: Whether to store to KV cache (for inference)
attention_mask: Attention mask [B, L] (1 = attend, 0 = mask)
"""
if cache_start is None:
cache_start = current_start
# Apply RoPE manually
cos, sin = freqs_cis
query_rope = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v)
key_rope = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
value_rope = v
# Get PRoPE transformed q, k, v
query_prope, key_prope, value_prope, apply_fn_o = prope_qkv(
q.transpose(1, 2), # [B, num_heads, L, head_dim]
k.transpose(1, 2),
v.transpose(1, 2),
viewmats=viewmats,
Ks=Ks,
patches_x=40, # hardcoded for now
patches_y=22,
)
# PRoPE returns [B, num_heads, L, head_dim], convert to [B, L, num_heads, head_dim]
query_prope = query_prope.transpose(1, 2)
key_prope = key_prope.transpose(1, 2)
value_prope = value_prope.transpose(1, 2)
# KV cache handling
if kv_cache is not None:
cache_key = kv_cache.get("k", None)
cache_value = kv_cache.get("v", None)
if cache_value is not None and not is_cache:
cache_key_rope, cache_key_prope = cache_key.chunk(2, dim=-1)
cache_value_rope, cache_value_prope = cache_value.chunk(2, dim=-1)
key_rope = torch.cat([cache_key_rope, key_rope], dim=1)
value_rope = torch.cat([cache_value_rope, value_rope], dim=1)
key_prope = torch.cat([cache_key_prope, key_prope], dim=1)
value_prope = torch.cat([cache_value_prope, value_prope], dim=1)
if is_cache:
# Store to cache (update input dict directly)
kv_cache["k"] = torch.cat([key_rope, key_prope], dim=-1)
kv_cache["v"] = torch.cat([value_rope, value_prope], dim=-1)
# Concatenate rope and prope paths (matching original)
query_all = torch.cat([query_rope, query_prope], dim=0)
key_all = torch.cat([key_rope, key_prope], dim=0)
value_all = torch.cat([value_rope, value_prope], dim=0)
# Check if Q and KV have different sequence lengths (KV cache mode)
# In this case, use LocalAttention (supports different Q/KV lengths)
if query_all.shape[1] != key_all.shape[1]:
raise ValueError("Q and KV have different sequence lengths")
# KV cache mode: Q has new tokens only, KV has cached + new tokens
# Use LocalAttention which supports different Q/KV lengths
# LocalAttention will use the appropriate backend (SageAttn, FlashAttn, etc.)
if not hasattr(self, '_kv_cache_attn'):
from fastvideo.attention import LocalAttention
self._kv_cache_attn = LocalAttention(
num_heads=self.num_heads,
head_size=self.head_dim,
causal=False,
supported_attention_backends=(AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)
)
hidden_states_all = self._kv_cache_attn(query_all, key_all, value_all)
else:
# Same sequence length: use DistributedAttention (supports SP)
# Create default attention mask if not provided
if attention_mask is None:
batch_size, seq_len = q.shape[0], q.shape[1]
attention_mask = torch.ones(batch_size, seq_len, device=q.device, dtype=q.dtype)
if q.dtype == torch.float32:
from fastvideo.attention.backends.sdpa import SDPAMetadataBuilder
attn_metadata_builder = SDPAMetadataBuilder
else:
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata_builder = FlashAttnMetadataBuilder
attn_metadata = attn_metadata_builder().build(
current_timestep=0,
attn_mask=attention_mask,
)
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
hidden_states_all, _ = self.attn(query_all, key_all, value_all, attention_mask=attention_mask)
hidden_states_rope, hidden_states_prope = hidden_states_all.chunk(2, dim=0)
hidden_states_prope = apply_fn_o(hidden_states_prope.transpose(1, 2)).transpose(1, 2)
return hidden_states_rope, hidden_states_prope
+424
View File
@@ -0,0 +1,424 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import torch
import torch.nn as nn
from fastvideo.configs.models.dits.wangamevideo import WanGameVideoConfig
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.visual_embedding import PatchEmbed
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanI2VCrossAttention
from fastvideo.platforms import AttentionBackendEnum, current_platform
# Import ActionModule
from fastvideo.models.dits.wangame.hyworld_action_module import WanGameActionTimeImageEmbedding, WanGameActionSelfAttention
logger = init_logger(__name__)
class WanGameCrossAttention(WanI2VCrossAttention):
def forward(self, x, context, context_lens=None):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
context_lens(Tensor): Shape [B]
"""
context_img = context
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
b, -1, n, d)
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
img_x = self.attn(q, k_img, v_img)
# output
x = img_x.flatten(2)
x, _ = self.to_out(x)
return x
class WanGameActionTransformerBlock(nn.Module):
"""
Transformer block for WAN Action model with support for:
- Self-attention with RoPE and camera PRoPE
- Cross-attention with text/image context
- Feed-forward network with AdaLN modulation
"""
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
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 = WanGameActionSelfAttention(
dim,
num_heads,
local_attn_size=local_attn_size,
sink_size=sink_size,
qk_norm=qk_norm,
eps=eps)
self.hidden_dim = dim
self.num_attention_heads = num_heads
self.local_attn_size = local_attn_size
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":
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
else:
raise ValueError(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,
compute_dtype=torch.float32)
# 2. Cross-attention (I2V only for now)
self.attn2 = WanGameCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
# norm3 for FFN input
self.norm3 = LayerNormScaleShift(dim, norm_type="layer", eps=eps,
elementwise_affine=False)
# 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)
# PRoPE output projection (initialized via add_discrete_action_parameters on the model)
self.to_out_prope = nn.ModuleList([
nn.Linear(dim, dim, bias=True),
])
nn.init.zeros_(self.to_out_prope[0].weight)
if self.to_out_prope[0].bias is not None:
nn.init.zeros_(self.to_out_prope[0].bias)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
kv_cache: dict | None = None,
crossattn_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None,
viewmats: torch.Tensor | None = None,
Ks: torch.Tensor | None = None,
is_cache: bool = False,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
num_frames = temb.shape[1]
frame_seqlen = hidden_states.shape[1] // num_frames
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# Cast temb to float32 for scale/shift computation
e = self.scale_shift_table + temb.float()
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(6, dim=2)
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).to(orig_dtype).flatten(1, 2)
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.forward_native(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(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))
# Self-attention with optional camera PRoPE
attn_output_rope, attn_output_prope = self.attn1(
query, key, value, freqs_cis,
kv_cache, current_start, cache_start, viewmats, Ks,
is_cache=is_cache
)
# Combine rope and prope outputs
attn_output_rope = attn_output_rope.flatten(2)
attn_output_rope, _ = self.to_out(attn_output_rope)
attn_output_prope = attn_output_prope.flatten(2)
attn_output_prope = self.to_out_prope[0](attn_output_prope)
attn_output = attn_output_rope.squeeze(1) + attn_output_prope.squeeze(1)
# Self-attention residual + norm in float32
null_shift = null_scale = torch.zeros(1, device=hidden_states.device, dtype=torch.float32)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states.float(), attn_output.float(), gate_msa, null_shift, null_scale)
hidden_states = hidden_states.type_as(attn_output)
norm_hidden_states = norm_hidden_states.type_as(attn_output)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states.to(orig_dtype),
context=encoder_hidden_states,
context_lens=None)
# Cross-attention residual in bfloat16
hidden_states = hidden_states + attn_output
# norm3 for FFN input in float32
norm_hidden_states = self.norm3(
hidden_states.float(), c_shift_msa, c_scale_msa
).type_as(hidden_states)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states.to(orig_dtype))
hidden_states = self.mlp_residual(hidden_states.float(), ff_output.float(), c_gate_msa)
hidden_states = hidden_states.to(orig_dtype) # Cast back to original dtype
return hidden_states
class WanGameActionTransformer3DModel(BaseDiT):
"""
WAN Action Transformer 3D Model for video generation with action conditioning.
Extends the base WAN video model with:
- Action embedding support for controllable generation
- camera PRoPE attention for 3D-aware generation
- KV caching for autoregressive inference
"""
_fsdp_shard_conditions = WanGameVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanGameVideoConfig()._compile_conditions
_supported_attention_backends = WanGameVideoConfig()._supported_attention_backends
param_names_mapping = WanGameVideoConfig().param_names_mapping
reverse_param_names_mapping = WanGameVideoConfig().reverse_param_names_mapping
lora_param_names_mapping = WanGameVideoConfig().lora_param_names_mapping
def __init__(self, config: WanGameVideoConfig, 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.attention_head_dim = config.attention_head_dim
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.local_attn_size = config.local_attn_size
self.inner_dim = inner_dim
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False)
# 2. Condition embeddings (with action support)
self.condition_embedder = WanGameActionTimeImageEmbedding(
dim=inner_dim,
time_freq_dim=config.freq_dim,
image_embed_dim=config.image_dim,
)
# 3. Transformer blocks
self.blocks = nn.ModuleList([
WanGameActionTransformerBlock(
inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.local_attn_size,
config.sink_size,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
supported_attention_backends=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)
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
# Causal-specific
self.num_frame_per_block = config.arch_config.num_frames_per_block
assert self.num_frame_per_block <= 3
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],
guidance=None,
action: torch.Tensor | None = None,
viewmats: torch.Tensor | None = None,
Ks: torch.Tensor | None = None,
kv_cache: list[dict] | None = None,
crossattn_cache: list[dict] | None = None,
current_start: int = 0,
cache_start: int = 0,
start_frame: int = 0,
is_cache: bool = False,
**kwargs
) -> torch.Tensor:
"""
Forward pass for both training and inference with KV caching.
Args:
hidden_states: Video latents [B, C, T, H, W]
encoder_hidden_states: Text embeddings [B, L, D]
timestep: Timestep tensor
encoder_hidden_states_image: Optional image embeddings
action: Action tensor [B, T] for per-frame conditioning
viewmats: Camera view matrices for PRoPE [B, T, 4, 4]
Ks: Camera intrinsics for PRoPE [B, T, 3, 3]
kv_cache: KV cache for autoregressive inference (list of dicts per layer)
crossattn_cache: Cross-attention cache for inference
current_start: Current position for KV cache
cache_start: Cache start position
start_frame: RoPE offset for new frames in autoregressive mode
is_cache: If True, populate KV cache and return early (cache-only mode)
"""
orig_dtype = hidden_states.dtype
# if 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 * get_sp_world_size(), 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,
start_frame=start_frame
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
if timestep.dim() == 2:
timestep = timestep.flatten()
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, action, encoder_hidden_states, encoder_hidden_states_image=encoder_hidden_states_image)
# Reshape timestep_proj: [T, 6*dim] -> [B, T, 6, dim]
# For training: batch_size=1, T=num_frames (diffusion forcing)
# For inference: batch_size can vary
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size))
if timestep_proj.shape[0] == post_patch_num_frames and batch_size == 1:
# Training mode: timestep_proj is [T, 6, dim], add batch dim -> [1, T, 6, dim]
timestep_proj = timestep_proj.unsqueeze(0)
else:
# Inference mode: reshape based on timestep shape
timestep_proj = timestep_proj.unflatten(dim=0, sizes=timestep.shape)
encoder_hidden_states = encoder_hidden_states_image
# Transformer blocks
for block_idx, block in enumerate(self.blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states, timestep_proj, freqs_cis,
kv_cache[block_idx] if kv_cache else None,
crossattn_cache[block_idx] if crossattn_cache else None,
current_start, cache_start,
viewmats, Ks, is_cache)
else:
hidden_states = block(
hidden_states, encoder_hidden_states, timestep_proj, freqs_cis,
kv_cache[block_idx] if kv_cache else None,
crossattn_cache[block_idx] if crossattn_cache else None,
current_start, cache_start,
viewmats, Ks, is_cache)
# If cache-only mode, return early
if is_cache:
return kv_cache
# Output norm, projection & unpatchify
# Reshape temb to match timestep_proj shape: [T, dim] -> [B, T, 1, dim]
if temb.shape[0] == post_patch_num_frames and batch_size == 1:
# Training mode: temb is [T, dim] -> [1, T, 1, dim]
temb = temb.unsqueeze(0).unsqueeze(2)
else:
# Inference mode: reshape based on timestep shape
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2, dim=2)
hidden_states = self.norm_out(hidden_states, shift, scale)
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
@@ -806,6 +806,10 @@ class TransformerLoader(ComponentLoader):
cls_name.startswith("Cosmos25")
or cls_name == "Cosmos25Transformer3DModel"
or getattr(fastvideo_args.pipeline_config, "prefix", "") == "Cosmos25"
) and not (
cls_name.startswith("WanGame")
or cls_name == "WanGameActionTransformer3DModel"
or getattr(fastvideo_args.pipeline_config, "prefix", "") == "WanGame"
)
model = maybe_load_fsdp_model(
model_cls=model_cls,
+5 -2
View File
@@ -138,7 +138,7 @@ def maybe_load_fsdp_model(
weight_iterator = safetensors_weights_iterator(weight_dir_list)
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
load_model_from_full_model_state_dict(
incompatible_keys, unexpected_keys = load_model_from_full_model_state_dict(
model,
weight_iterator,
device,
@@ -147,6 +147,9 @@ def maybe_load_fsdp_model(
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
)
if incompatible_keys or unexpected_keys:
logger.warning("Incompatible keys: %s", incompatible_keys)
logger.warning("Unexpected keys: %s", unexpected_keys)
for n, p in chain(model.named_parameters(), model.named_buffers()):
if p.is_meta:
raise RuntimeError(
@@ -340,7 +343,7 @@ def load_model_from_full_model_state_dict(
unused_keys)
# List of allowed parameter name patterns
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l"] # Can be extended as needed
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l", "to_out_prope", "action_embedder"] # Can be extended as needed
for new_param_name in unused_keys:
if not any(pattern in new_param_name
for pattern in ALLOWED_NEW_PARAM_PATTERNS):
+1
View File
@@ -42,6 +42,7 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"WanGameActionTransformer3DModel": ("dits", "wangame", "WanGameActionTransformer3DModel"),
"MatrixGameWanModel": ("dits", "matrixgame", "MatrixGameWanModel"),
"CausalMatrixGameWanModel": ("dits", "matrixgame", "CausalMatrixGameWanModel"),
}
@@ -0,0 +1,74 @@
# 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.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (
ImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
logger = init_logger(__name__)
class WanGameActionImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"vae", "transformer", "scheduler", \
"image_encoder", "image_processor"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
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="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
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="image_latent_preparation_stage",
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
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 = WanGameActionImageToVideoPipeline
+1
View File
@@ -21,6 +21,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanPipeline": "wan",
"WanDMDPipeline": "wan",
"WanImageToVideoPipeline": "wan",
"WanGameActionImageToVideoPipeline": "wan",
"WanVideoToVideoPipeline": "wan",
"WanCausalDMDPipeline": "wan",
"TurboDiffusionPipeline": "turbodiffusion",
@@ -18,6 +18,8 @@ from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
PreprocessPipeline_Text)
from fastvideo.pipelines.preprocess.matrixgame.matrixgame_preprocess_pipeline import (
PreprocessPipeline_MatrixGame)
from fastvideo.pipelines.preprocess.wangame.wangame_preprocess_pipeline import (
PreprocessPipeline_WanGame)
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -64,10 +66,12 @@ def main(args) -> None:
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
elif args.preprocess_task == "matrixgame":
PreprocessPipeline = PreprocessPipeline_MatrixGame
elif args.preprocess_task == "wangame":
PreprocessPipeline = PreprocessPipeline_WanGame
else:
raise ValueError(
f"Invalid preprocess task: {args.preprocess_task}. "
f"Valid options: t2v, i2v, ode_trajectory, text_only, matrixgame")
f"Valid options: t2v, i2v, ode_trajectory, text_only, matrixgame, wangame")
logger.info("Preprocess task: %s using %s", args.preprocess_task,
PreprocessPipeline.__name__)
@@ -115,7 +119,7 @@ if __name__ == "__main__":
"--preprocess_task",
type=str,
default="t2v",
choices=["t2v", "i2v", "text_only", "ode_trajectory", "matrixgame"],
choices=["t2v", "i2v", "text_only", "ode_trajectory", "matrixgame", "wangame"],
help="Type of preprocessing task to run")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
@@ -0,0 +1,303 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Any
import numpy as np
import torch
from PIL import Image
from fastvideo.dataset.dataloader.schema import pyarrow_schema_wangame
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.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import ImageEncodingStage
class PreprocessPipeline_WanGame(BasePreprocessPipeline):
"""I2V preprocessing pipeline implementation."""
_required_config_modules = ["vae", "image_encoder", "image_processor"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
def get_pyarrow_schema(self):
"""Return the PyArrow schema for I2V pipeline."""
return pyarrow_schema_wangame
def get_extra_features(self, valid_data: dict[str, Any],
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
# TODO(will): move these to cpu at some point
self.get_module("image_encoder").to(get_local_torch_device())
self.get_module("vae").to(get_local_torch_device())
features = {}
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
_, _, num_frames, height, width = valid_data["pixel_values"].shape
# latent_height = height // self.get_module(
# "vae").spatial_compression_ratio
# latent_width = width // self.get_module("vae").spatial_compression_ratio
processed_images = []
# Frame has values between -1 and 1
for frame in first_frame:
frame = (frame + 1) * 127.5
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
processed_img = self.get_module("image_processor")(
images=frame_pil, return_tensors="pt")
processed_images.append(processed_img)
# Get CLIP features
pixel_values = torch.cat(
[img['pixel_values'] for img in processed_images],
dim=0).to(get_local_torch_device())
with torch.no_grad():
image_inputs = {'pixel_values': pixel_values}
with set_forward_context(current_timestep=0, attn_metadata=None):
clip_features = self.get_module("image_encoder")(**image_inputs)
clip_features = clip_features.last_hidden_state
features["clip_feature"] = clip_features
"""Get VAE features from the first frame of each video"""
video_conditions = []
for frame in first_frame:
processed_img = frame.to(device="cpu", dtype=torch.float32)
processed_img = processed_img.unsqueeze(0).permute(0, 3, 1,
2).unsqueeze(2)
# (B, H, W, C) -> (B, C, 1, H, W)
video_condition = torch.cat([
processed_img,
processed_img.new_zeros(processed_img.shape[0],
processed_img.shape[1], num_frames - 1,
height, width)
],
dim=2)
video_condition = video_condition.to(
device=get_local_torch_device(), dtype=torch.float32)
video_conditions.append(video_condition)
video_conditions = torch.cat(video_conditions, dim=0)
with torch.autocast(device_type="cuda",
dtype=torch.float32,
enabled=True):
encoder_outputs = self.get_module("vae").encode(video_conditions)
latent_condition = encoder_outputs.mean
if (hasattr(self.get_module("vae"), "shift_factor")
and self.get_module("vae").shift_factor is not None):
if isinstance(self.get_module("vae").shift_factor, torch.Tensor):
latent_condition -= self.get_module("vae").shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.get_module("vae").shift_factor
if isinstance(self.get_module("vae").scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor.to(latent_condition.device,
latent_condition.dtype)
else:
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor
# mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
# latent_width)
# mask_lat_size[:, :, list(range(1, num_frames))] = 0
# first_frame_mask = mask_lat_size[:, :, 0:1]
# first_frame_mask = torch.repeat_interleave(
# first_frame_mask,
# dim=2,
# repeats=self.get_module("vae").temporal_compression_ratio)
# mask_lat_size = torch.concat(
# [first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
# mask_lat_size = mask_lat_size.view(
# batch_size, -1,
# self.get_module("vae").temporal_compression_ratio, latent_height,
# latent_width)
# mask_lat_size = mask_lat_size.transpose(1, 2)
# mask_lat_size = mask_lat_size.to(latent_condition.device)
# image_latent = torch.concat([mask_lat_size, latent_condition], dim=1)
features["first_frame_latent"] = latent_condition
if "action_path" in valid_data and valid_data["action_path"]:
keyboard_cond_list = []
mouse_cond_list = []
num_bits = 6
for action_path in valid_data["action_path"]:
if action_path:
action_data = np.load(action_path, allow_pickle=True)
if isinstance(
action_data,
np.ndarray) and action_data.dtype == np.dtype('O'):
action_dict = action_data.item()
if "keyboard" in action_dict:
keyboard_raw = action_dict["keyboard"]
# Convert 1D bit-flag values to 2D multi-hot encoding
if isinstance(keyboard_raw, np.ndarray):
if keyboard_raw.ndim == 1:
# [T] -> [T, num_bits]
T = len(keyboard_raw)
multi_hot = np.zeros((T, num_bits),
dtype=np.float32)
action_values = keyboard_raw.astype(int)
for bit_idx in range(num_bits):
target_idx = (
2 -
(bit_idx % 3)) + 3 * (bit_idx // 3)
if target_idx < num_bits:
multi_hot[:, target_idx] = (
(action_values >> bit_idx)
& 1).astype(np.float32)
keyboard_cond_list.append(multi_hot)
else:
# If already 2D, pad to num_bits if necessary
k_data = keyboard_raw.astype(np.float32)
if k_data.ndim == 2 and k_data.shape[
-1] < num_bits:
padding = np.zeros(
(k_data.shape[0],
num_bits - k_data.shape[-1]),
dtype=np.float32)
k_data = np.concatenate(
[k_data, padding], axis=-1)
keyboard_cond_list.append(k_data)
else:
keyboard_cond_list.append(keyboard_raw)
if "mouse" in action_dict:
mouse_cond_list.append(action_dict["mouse"])
else:
if isinstance(action_data,
np.ndarray) and action_data.ndim == 1:
T = len(action_data)
multi_hot = np.zeros((T, num_bits),
dtype=np.float32)
action_values = action_data.astype(int)
for bit_idx in range(num_bits):
target_idx = (
2 - (bit_idx % 3)) + 3 * (bit_idx // 3)
if target_idx < num_bits:
multi_hot[:, target_idx] = (
(action_values >> bit_idx) & 1).astype(
np.float32)
keyboard_cond_list.append(multi_hot)
else:
# If already 2D, pad to num_bits if necessary
k_data = action_data.astype(np.float32)
if k_data.ndim == 2 and k_data.shape[-1] < num_bits:
padding = np.zeros(
(k_data.shape[0],
num_bits - k_data.shape[-1]),
dtype=np.float32)
k_data = np.concatenate([k_data, padding],
axis=-1)
keyboard_cond_list.append(k_data)
if keyboard_cond_list:
features["keyboard_cond"] = keyboard_cond_list
if mouse_cond_list:
features["mouse_cond"] = mouse_cond_list
return features
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int,
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
"""Create a record for the Parquet dataset with CLIP features."""
record = super().create_record(video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "clip_feature" in extra_features:
clip_feature = extra_features["clip_feature"]
record.update({
"clip_feature_bytes": clip_feature.tobytes(),
"clip_feature_shape": list(clip_feature.shape),
"clip_feature_dtype": str(clip_feature.dtype),
})
else:
record.update({
"clip_feature_bytes": b"",
"clip_feature_shape": [],
"clip_feature_dtype": "",
})
if extra_features and "first_frame_latent" in extra_features:
first_frame_latent = extra_features["first_frame_latent"]
record.update({
"first_frame_latent_bytes":
first_frame_latent.tobytes(),
"first_frame_latent_shape":
list(first_frame_latent.shape),
"first_frame_latent_dtype":
str(first_frame_latent.dtype),
})
else:
record.update({
"first_frame_latent_bytes": b"",
"first_frame_latent_shape": [],
"first_frame_latent_dtype": "",
})
if extra_features and "pil_image" in extra_features:
pil_image = extra_features["pil_image"]
record.update({
"pil_image_bytes": pil_image.tobytes(),
"pil_image_shape": list(pil_image.shape),
"pil_image_dtype": str(pil_image.dtype),
})
else:
record.update({
"pil_image_bytes": b"",
"pil_image_shape": [],
"pil_image_dtype": "",
})
if extra_features and "keyboard_cond" in extra_features:
keyboard_cond = extra_features["keyboard_cond"]
record.update({
"keyboard_cond_bytes": keyboard_cond.tobytes(),
"keyboard_cond_shape": list(keyboard_cond.shape),
"keyboard_cond_dtype": str(keyboard_cond.dtype),
})
else:
record.update({
"keyboard_cond_bytes": b"",
"keyboard_cond_shape": [],
"keyboard_cond_dtype": "",
})
if extra_features and "mouse_cond" in extra_features:
mouse_cond = extra_features["mouse_cond"]
record.update({
"mouse_cond_bytes": mouse_cond.tobytes(),
"mouse_cond_shape": list(mouse_cond.shape),
"mouse_cond_dtype": str(mouse_cond.dtype),
})
else:
record.update({
"mouse_cond_bytes": b"",
"mouse_cond_shape": [],
"mouse_cond_dtype": "",
})
return record
EntryClass = PreprocessPipeline_WanGame
+16
View File
@@ -168,6 +168,20 @@ class DenoisingStage(PipelineStage):
},
)
if batch.mouse_cond is not None and batch.keyboard_cond is not None:
from fastvideo.models.dits.hyworld.pose import process_custom_actions
viewmats, intrinsics, action_labels = process_custom_actions(batch.keyboard_cond, batch.mouse_cond)
camera_action_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"viewmats": viewmats.unsqueeze(0).to(get_local_torch_device(), dtype=target_dtype),
"Ks": intrinsics.unsqueeze(0).to(get_local_torch_device(), dtype=target_dtype),
"action": action_labels.unsqueeze(0).to(get_local_torch_device(), dtype=target_dtype),
},
)
else:
camera_action_kwargs = {}
action_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
@@ -406,6 +420,7 @@ class DenoisingStage(PipelineStage):
**image_kwargs,
**pos_cond_kwargs,
**action_kwargs,
**camera_action_kwargs,
)
if batch.do_classifier_free_guidance:
@@ -423,6 +438,7 @@ class DenoisingStage(PipelineStage):
**image_kwargs,
**neg_cond_kwargs,
**action_kwargs,
**camera_action_kwargs,
)
noise_pred_text = noise_pred
@@ -0,0 +1,250 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import Any
import torch
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset.dataloader.schema import pyarrow_schema_wangame
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.pipelines.basic.wan.wangame_i2v_pipeline import WanGameActionImageToVideoPipeline
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanGameTrainingPipeline(TrainingPipeline):
"""
A training pipeline for WanGame-2.1-Fun-1.3B-InP.
"""
_required_config_modules = ["scheduler", "transformer", "vae"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_wangame
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
# args_copy.pipeline_config.vae_config.load_encoder = False
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
self.validation_pipeline = WanGameActionImageToVideoPipeline.from_pretrained(
training_args.model_path,
args=None,
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
dit_cpu_offload=False)
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
latents = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
# encoder_hidden_states = batch['text_embedding']
# encoder_attention_mask = batch['text_attention_mask']
clip_features = batch['clip_feature']
image_latents = batch['first_frame_latent']
image_latents = image_latents[:, :, :self.training_args.num_latent_t]
pil_image = batch['pil_image']
infos = batch['info_list']
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = None
training_batch.encoder_attention_mask = None
# MatrixGame doesn't use text encoder
training_batch.preprocessed_image = pil_image.to(
get_local_torch_device())
training_batch.image_embeds = clip_features.to(get_local_torch_device())
training_batch.image_latents = image_latents.to(
get_local_torch_device())
training_batch.infos = infos
# Action conditioning
if 'mouse_cond' in batch and batch['mouse_cond'].numel() > 0:
training_batch.mouse_cond = batch['mouse_cond'].to(
get_local_torch_device(), dtype=torch.bfloat16)
else:
training_batch.mouse_cond = None
if 'keyboard_cond' in batch and batch['keyboard_cond'].numel() > 0:
training_batch.keyboard_cond = batch['keyboard_cond'].to(
get_local_torch_device(), dtype=torch.bfloat16)
else:
training_batch.keyboard_cond = None
return training_batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
# First, call parent method to prepare noise, timesteps, etc. for video latents
training_batch = super()._prepare_dit_inputs(training_batch)
assert isinstance(training_batch.image_latents, torch.Tensor)
image_latents = training_batch.image_latents.to(
get_local_torch_device(), dtype=torch.bfloat16)
temporal_compression_ratio = self.training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (self.training_args.num_latent_t -
1) * temporal_compression_ratio + 1
batch_size, num_channels, _, latent_height, latent_width = image_latents.shape
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
latent_width)
mask_lat_size[:, :, 1:] = 0
first_frame_mask = mask_lat_size[:, :, :1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask, dim=2, repeats=temporal_compression_ratio)
mask_lat_size = torch.cat([first_frame_mask, mask_lat_size[:, :, 1:]],
dim=2)
mask_lat_size = mask_lat_size.view(batch_size, -1,
temporal_compression_ratio,
latent_height, latent_width)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(
image_latents.device).to(dtype=torch.bfloat16)
training_batch.noisy_model_input = torch.cat(
[training_batch.noisy_model_input, mask_lat_size, image_latents],
dim=1)
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
# Image Embeds for conditioning
image_embeds = training_batch.image_embeds
assert torch.isnan(image_embeds).sum() == 0
image_embeds = image_embeds.to(get_local_torch_device(),
dtype=torch.bfloat16)
encoder_hidden_states_image = image_embeds
from fastvideo.models.dits.hyworld.pose import process_custom_actions
viewmats, intrinsics, action_labels = process_custom_actions(training_batch.keyboard_cond, training_batch.mouse_cond)
viewmats = viewmats.unsqueeze(0).to(get_local_torch_device(), dtype=torch.bfloat16)
intrinsics = intrinsics.unsqueeze(0).to(get_local_torch_device(), dtype=torch.bfloat16)
action_labels = action_labels.unsqueeze(0).to(get_local_torch_device(), dtype=torch.bfloat16)
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
training_batch.encoder_hidden_states, # None for MatrixGame
"timestep":
training_batch.timesteps.to(get_local_torch_device(),
dtype=torch.bfloat16),
# "encoder_attention_mask":
# training_batch.encoder_attention_mask,
"encoder_hidden_states_image":
encoder_hidden_states_image,
# Action conditioning
"viewmats": viewmats,
"Ks": intrinsics,
"action": action_labels,
"return_dict":
False,
}
return training_batch
def _prepare_validation_batch(self, sampling_param: SamplingParam,
training_args: TrainingArgs,
validation_batch: dict[str, Any],
num_inference_steps: int) -> ForwardBatch:
sampling_param.prompt = validation_batch['prompt']
sampling_param.height = training_args.num_height
sampling_param.width = training_args.num_width
sampling_param.image_path = validation_batch.get(
'image_path') or validation_batch.get('video_path')
sampling_param.num_inference_steps = num_inference_steps
sampling_param.data_type = "video"
assert self.seed is not None
sampling_param.seed = self.seed
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
sampling_param.num_frames = num_frames
batch = ForwardBatch(
**shallow_asdict(sampling_param),
latents=None,
generator=torch.Generator(device="cpu").manual_seed(self.seed),
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
if "image" in validation_batch and validation_batch["image"] is not None:
batch.pil_image = validation_batch["image"]
if "keyboard_cond" in validation_batch and validation_batch[
"keyboard_cond"] is not None:
keyboard_cond = validation_batch["keyboard_cond"]
keyboard_cond = torch.tensor(keyboard_cond, dtype=torch.bfloat16)
keyboard_cond = keyboard_cond.unsqueeze(0)
batch.keyboard_cond = keyboard_cond
if "mouse_cond" in validation_batch and validation_batch[
"mouse_cond"] is not None:
mouse_cond = validation_batch["mouse_cond"]
mouse_cond = torch.tensor(mouse_cond, dtype=torch.bfloat16)
mouse_cond = mouse_cond.unsqueeze(0)
batch.mouse_cond = mouse_cond
return batch
def main(args) -> None:
logger.info("Starting training pipeline...")
pipeline = WanGameTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.dit_cpu_offload = False
main(args)