Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bdaf2e2e45 | ||
|
|
f3a7a37312 | ||
|
|
b1abb5c251 | ||
|
|
a80824255e |
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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"
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()),
|
||||
])
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user