Compare commits

...
Author SHA1 Message Date
mignonjia 8fa6ba6178 mc dfsft 2026-04-01 03:53:20 +00:00
H1yori233 2e5fef787b fix tf scheduler 2026-03-17 18:19:27 -07:00
H1yori233 43d87816bd add logger 2026-03-16 17:27:59 -07:00
H1yori233 2ace7dc6f4 update df scheduler 2026-03-16 15:49:09 -07:00
H1yori233 0ca75db738 fix train / val step mismatch 2026-03-15 20:55:53 -07:00
H1yori233 375ffd3fd5 make visualization in 1 panel 2026-03-15 16:59:34 -07:00
H1yori233 2615ba4291 upload more validation to wandb 2026-03-15 16:43:59 -07:00
RandNMR73 474dd71f28 config 2026-03-15 22:10:50 +00:00
RandNMR73 98ad2d2db6 wangame 2026-03-15 22:01:33 +00:00
49 changed files with 7926 additions and 302 deletions
@@ -0,0 +1,95 @@
# V3 config: WanGame causal Diffusion-Forcing SFT (DFSFT).
#
# Uses _target_-based instantiation — each model role is an independent
# class instance; the method class is resolved directly from the YAML.
models:
student:
_target_: fastvideo.train.models.wangame.WanGameCausalModel
init_from: /mnt/weka/home/hao.zhang/kaiqin/wg_models/WanGame-2.1-0306-61500steps
trainable: true
# transformer_override_safetensor: /mnt/weka/home/hao.zhang/mhuo/FastVideo-hyw/outputs/wangame_dfsft_causal_v3/checkpoint-best-step-36500/transformer/model.safetensors
method:
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
attn_kind: dense
# use_ema: true
chunk_size: 3
min_timestep_ratio: 0.02
max_timestep_ratio: 0.98
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: >-
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_2130/preprocessed:1,
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_1600/preprocessed:0,
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/0_static_plus_w_only/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/1_wasd_only/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/wasdonly_alpha1/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/camera/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/camera4hold_alpha1/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/preprocessed:3
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 352
num_width: 640
num_frames: 69
apply_bot_died_filter: true
optimizer:
learning_rate: 1e-4
betas: [0.9, 0.95]
weight_decay: 1e-5
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 60000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wangame_dfsft_causal_v3
# resume_from_checkpoint: /mnt/weka/home/hao.zhang/mhuo/FastVideo-hyw/outputs/wangame_dfsft_causal_v3/checkpoint-best-step-36500
training_state_checkpointing_steps: 1000000
weight_only_checkpointing_steps: 1000000
checkpoints_total_limit: 0
best_checkpoint_start_step: 1000000
best_checkpoint_top_k: 0
tracker:
project_name: distillation_wangame_r
run_name: wangame_dfsft_causal_v3
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
# ema:
# _target_: fastvideo.train.callbacks.ema.EMACallback
# beta: 0.9999
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.wan.wangame_causal_dmd_pipeline.WangameCausalOdeDMDPipeline
dataset_file: examples/training/finetune/WanGame2.1_1.3b_i2v/validation_8.json
every_steps: 500
sampling_steps: [40]
scheduler_target: fastvideo.models.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler
guidance_scale: 1.0
num_frames: 69
evaluate_ptlflow: false
pipeline:
flow_shift: 3
@@ -0,0 +1,102 @@
models:
student:
_target_: fastvideo.train.models.wangame.WanGameCausalModel
init_from: /mnt/weka/home/hao.zhang/kaiqin/wg_models/SFWanGame-2.1-0308-10000steps
trainable: true
teacher:
_target_: fastvideo.train.models.wangame.WanGameModel
init_from: /mnt/weka/home/hao.zhang/kaiqin/wg_models/WanGame-2.1-0306-61500steps
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.wangame.WanGameModel
init_from: /mnt/weka/home/hao.zhang/kaiqin/wg_models/WanGame-2.1-0306-61500steps
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
rollout_mode: simulate
generator_update_interval: 5
real_score_guidance_scale: 1.0
dmd_denoising_steps: [1000, 750, 500, 250]
warp_denoising_step: true
chunk_size: 3
student_sample_type: sde
context_noise: 0.0
enable_gradient_in_rollout: true
start_gradient_frame: 0
# Critic optimizer
fake_score_learning_rate: 8.0e-6
fake_score_betas: [0.0, 0.999]
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: >-
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_2130/preprocessed:1,
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_1600/preprocessed:0,
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/0_static_plus_w_only/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/1_wasd_only/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/wasdonly_alpha1/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/camera/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/camera4hold_alpha1/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/preprocessed:3
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 352
num_width: 640
num_frames: 69
apply_bot_died_filter: true
optimizer:
learning_rate: 1e-5
betas: [0.0, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 6
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wangame_1.3b_self_forcing
training_state_checkpointing_steps: 500
checkpoints_total_limit: 5
tracker:
project_name: wangame_sf
run_name: wangame_1.3b_self_forcing
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wangame_causal_dmd_pipeline.WangameCausalSdeDMDPipeline
scheduler_target: fastvideo.models.schedulers.scheduling_self_forcing_flow_match.SelfForcingFlowMatchScheduler
dataset_file: examples/training/finetune/WanGame2.1_1.3b_i2v/validation_4.json
every_steps: 5
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250]
num_frames: 69
guidance_scale: 1.0
evaluate_ptlflow: false
pipeline:
flow_shift: 3
@@ -0,0 +1,94 @@
# V3 config: WanGame causal Diffusion-Forcing SFT (DFSFT).
#
# Uses _target_-based instantiation — each model role is an independent
# class instance; the method class is resolved directly from the YAML.
models:
student:
_target_: fastvideo.train.models.wangame.WanGameCausalModel
init_from: /mnt/weka/home/hao.zhang/kaiqin/wg_models/WanGame-2.1-0306-61500steps
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
attn_kind: dense
# use_ema: true
chunk_size: 3
min_timestep_ratio: 0.02
max_timestep_ratio: 0.98
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: >-
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_2130/preprocessed:1,
/mnt/weka/home/hao.zhang/mhuo/traindata_0204_1600/preprocessed:0,
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/0_static_plus_w_only/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0205_1330/data/1_wasd_only/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/wasdonly_alpha1/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0206_1200/data/camera/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/camera4hold_alpha1/preprocessed:3,
/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/preprocessed:3
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 352
num_width: 640
num_frames: 69
apply_bot_died_filter: true
optimizer:
learning_rate: 1e-5
betas: [0.0, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 10000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wangame_tfsft_causal
# resume_from_checkpoint: /mnt/weka/home/hao.zhang/mhuo/FastVideo-hyw/outputs/wangame_dfsft_causal_v3/checkpoint-best-step-36500
training_state_checkpointing_steps: 5000
weight_only_checkpointing_steps: 5000
checkpoints_total_limit: 10
best_checkpoint_start_step: 2000
best_checkpoint_top_k: 5
tracker:
project_name: distillation_wangame_r
run_name: wangame_tfsft_causal
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
# ema:
# _target_: fastvideo.train.callbacks.ema.EMACallback
# beta: 0.9999
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.wan.wangame_causal_dmd_pipeline.WangameCausalOdeDMDPipeline
dataset_file: examples/training/finetune/WanGame2.1_1.3b_i2v/validation_4.json
every_steps: 500
sampling_steps: [40]
scheduler_target: fastvideo.models.schedulers.scheduling_flow_match_euler_discrete.FlowMatchEulerDiscreteScheduler
guidance_scale: 1.0
num_frames: 69
evaluate_ptlflow: false
pipeline:
flow_shift: 3
+9 -2
View File
@@ -22,16 +22,17 @@ shift
# ── GPU / node settings ──────────────────────────────────────────
NUM_GPUS="${NUM_GPUS:-$(nvidia-smi -L 2>/dev/null | wc -l)}"
NUM_GPUS="${NUM_GPUS:-1}"
NUM_GPUS="${NUM_GPUS:-8}"
NNODES="${NNODES:-1}"
NODE_RANK="${NODE_RANK:-0}"
MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}"
MASTER_PORT="${MASTER_PORT:-29501}"
export TOKENIZERS_PARALLELISM=false
# ── W&B ──────────────────────────────────────────────────────────
export WANDB_API_KEY="${WANDB_API_KEY:-}"
export WANDB_API_KEY="${WANDB_API_KEY:-7ff8b6e8356924f7a6dd51a0342dd1a422ea9352}"
export WANDB_MODE="${WANDB_MODE:-online}"
# ── Log file ─────────────────────────────────────────────────────
CONFIG_NAME="$(basename "${CONFIG}" .yaml)"
TIMESTAMP="$(date +%Y%m%d_%H%M%S)"
@@ -39,6 +40,12 @@ LOG_DIR="${LOG_DIR:-examples/train}"
mkdir -p "${LOG_DIR}"
LOG_FILE="${LOG_DIR}/${CONFIG_NAME}_${TIMESTAMP}.log"
set +u
source ~/conda/miniconda/bin/activate
conda activate mhuo-fv
set -u
export PYTHONPATH="/mnt/weka/home/hao.zhang/mhuo/FastVideo-refactor:${PYTHONPATH:-}"
echo "=== Train Training ==="
echo "Config: ${CONFIG}"
echo "Num GPUs: ${NUM_GPUS}"
+80
View File
@@ -0,0 +1,80 @@
#!/bin/bash
#SBATCH --job-name=wg-sf
#SBATCH --partition=main
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=examples/train/slurm_%j.out
#SBATCH --error=examples/train/slurm_%j.err
#SBATCH --exclusive
set -euo pipefail
CONFIG="${1:?Usage: sbatch examples/train/run.slurm <config.yaml> [extra flags...]}"
shift
EXTRA_ARGS=("$@")
set --
cd /mnt/weka/home/hao.zhang/mhuo/FastVideo-refactor
get_num_gpus() {
if [[ -n "${NUM_GPUS:-}" ]]; then
echo "${NUM_GPUS}"
return
fi
if [[ -n "${SLURM_GPUS_ON_NODE:-}" ]]; then
echo "${SLURM_GPUS_ON_NODE%%(*}"
return
fi
if [[ -n "${SLURM_GPUS_PER_NODE:-}" ]]; then
echo "${SLURM_GPUS_PER_NODE%%(*}"
return
fi
if command -v nvidia-smi >/dev/null 2>&1; then
nvidia-smi -L 2>/dev/null | wc -l | tr -d " "
else
echo 8
fi
}
export NNODES="${NNODES:-${SLURM_JOB_NUM_NODES:-1}}"
export NUM_GPUS="$(get_num_gpus)"
export MASTER_PORT="${MASTER_PORT:-29501}"
if [[ -z "${MASTER_ADDR:-}" ]]; then
nodes=( $(scontrol show hostnames "${SLURM_JOB_NODELIST}") )
export MASTER_ADDR="${nodes[0]}"
fi
export NCCL_P2P_DISABLE="${NCCL_P2P_DISABLE:-1}"
export TORCH_NCCL_ENABLE_MONITORING="${TORCH_NCCL_ENABLE_MONITORING:-0}"
export NCCL_DEBUG_SUBSYS="${NCCL_DEBUG_SUBSYS:-INIT,NET}"
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="${WANDB_API_KEY:-7ff8b6e8356924f7a6dd51a0342dd1a422ea9352}"
export WANDB_MODE="${WANDB_MODE:-online}"
set +u
source ~/conda/miniconda/bin/activate
conda activate mhuo-fv
set -u
export PYTHONPATH="/mnt/weka/home/hao.zhang/mhuo/FastVideo-refactor:${PYTHONPATH:-}"
echo "=== Distillation Training (Slurm) ==="
echo "Config: ${CONFIG}"
echo "Num GPUs: ${NUM_GPUS}"
echo "Num Nodes: ${NNODES}"
echo "Master: ${MASTER_ADDR}:${MASTER_PORT}"
echo "Extra args: ${EXTRA_ARGS[*]:-}"
echo "====================================="
srun torchrun \
--nnodes "${NNODES}" \
--nproc_per_node "${NUM_GPUS}" \
--rdzv_backend c10d \
--rdzv_endpoint "${MASTER_ADDR}:${MASTER_PORT}" \
--node_rank "${SLURM_PROCID}" \
fastvideo/train/entrypoint/train.py \
--config "${CONFIG}" \
"${EXTRA_ARGS[@]}"
@@ -0,0 +1,44 @@
{
"data": [
{
"caption": "00 Val-00: W",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000002.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/W.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "01 Val-01: S",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000003.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/S.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "02 Val-02: A",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000004.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/A.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "03 Val-03: D",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000005.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/D.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
}
]
}
@@ -0,0 +1,84 @@
{
"data": [
{
"caption": "00 Val-00: W",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000002.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/W.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "01 Val-01: S",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000003.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/S.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "02 Val-02: A",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000004.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/A.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "03 Val-03: D",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000005.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/D.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "04 Val-04: u",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000000.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/u.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "05 Val-05: d",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000001.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/d.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "06 Val-06: l",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000006.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/l.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "07 Val-07: r",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000007.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/r.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
}
]
}
@@ -0,0 +1,324 @@
{
"data": [
{
"caption": "00 Val-00: W",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000002.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/W.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "01 Val-01: S",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000003.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/S.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "02 Val-02: A",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000004.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/A.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "03 Val-03: D",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000005.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/D.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "04 Val-04: u",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000000.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/u.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "05 Val-05: d",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000001.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/d.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "06 Val-06: l",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000006.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/l.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "07 Val-07: r",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000007.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/r.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "08 Val-00: key rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000002.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_1_action_rand_1.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "09 Val-01: key rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000003.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_1_action_rand_2.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "10 Val-02: camera rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000004.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/camera_1_action_rand_1.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "11 Val-03: camera rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000005.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/camera_1_action_rand_2.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "12 Val-00: key+camera excl rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000002.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_1_action_rand_1.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "13 Val-01: key+camera excl rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000003.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_1_action_rand_2.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "14 Val-02: key+camera rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000004.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_1_action_rand_1.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "15 Val-03: key+camera rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000005.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_1_action_rand_2.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "16 Val-04: (simultaneous) key rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000000.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_2_action_rand_1.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "17 Val-05: (simultaneous) camera rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000001.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/camera_2_action_rand_1.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "18 Val-06: (simultaneous) key+camera excl rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000006.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_2_action_rand_1.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "19 Val-07: (simultaneous) key+camera rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000007.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_2_action_rand_1.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "20 Val-08: W+A",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/humanplay/000005.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/WA.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "21 Val-09: S+u",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/humanplay/000013.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/S_u.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "22 Val-08: Still",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/humanplay/000005.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/still.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "23 Val-09: Still",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/humanplay/000013.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/still.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "24 Val-06: key+camera excl rand Frame 4",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000006.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_1_action_rand_1_f4.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "25 Val-07: key+camera excl rand Frame 4",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/mc_wasd_10/validate/000007.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_1_action_rand_2_f4.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "26 Train-00",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/first_frame/000500.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/videos/000500_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "27 Train-01",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/first_frame/001000.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/traindata_0208_2000/data/wasd4holdrandview_simple_1key1mouse1/videos/001000_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "28 Doom-00: W",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/doom/000000.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/W.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "29 Doom-01: key rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/doom/000001.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_1_action_rand_1.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "30 Doom-02: camera rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/doom/000002.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/camera_1_action_rand_1.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "31 Doom-03: key+camera excl rand",
"image_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/doom/000003.jpg",
"action_path": "/mnt/weka/home/hao.zhang/mhuo/FastVideo/examples/training/finetune/WanGame2.1_1.3b_i2v/actions/key_camera_excl_1_action_rand_1.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
}
]
}
+3 -1
View File
@@ -7,9 +7,11 @@ from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
from fastvideo.configs.models.dits.wangamevideo import WanGameVideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig",
"WanVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig"
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig",
"WanGameVideoConfig"
]
@@ -0,0 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.wanvideo import (
WanVideoArchConfig,
WanVideoConfig,
)
@dataclass
class WanGameVideoArchConfig(WanVideoArchConfig):
"""Wangame keeps WanVideo architecture defaults and checkpoint mappings."""
@dataclass
class WanGameVideoConfig(WanVideoConfig):
arch_config: WanGameVideoArchConfig = field(
default_factory=WanGameVideoArchConfig
)
prefix: str = "WanGame"
+10
View File
@@ -69,6 +69,8 @@ class PipelineConfig:
# DMD parameters
dmd_denoising_steps: list[int] | None = field(default=None)
ode_solver: str = "unipc"
# Wan2.2 TI2V parameters
ti2v_task: bool = False
boundary_ratio: float | None = None
@@ -175,6 +177,14 @@ class PipelineConfig:
help=
"Comma-separated list of denoising steps (e.g., '1000,757,522')",
)
parser.add_argument(
f"--{prefix_with_dot}ode-solver",
type=str,
choices=["unipc", "euler"],
dest=f"{prefix_with_dot.replace('-', '_')}ode_solver",
default=PipelineConfig.ode_solver,
help="ODE solver selection for ode sampling.",
)
# Add VAE configuration arguments
from fastvideo.configs.models.vaes.base import VAEConfig
+2
View File
@@ -157,3 +157,5 @@ pyarrow_schema_matrixgame = pa.schema([
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_wangame = pyarrow_schema_matrixgame
+219 -60
View File
@@ -1,4 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
import hashlib
import os
import pickle
import random
@@ -36,7 +37,9 @@ class DP_SP_BatchSampler(Sampler[list[int]]):
global_rank: int,
drop_last: bool = True,
drop_first_row: bool = False,
reshuffle_each_epoch: bool = True,
seed: int = 0,
candidate_indices: list[int] | None = None,
):
self.batch_size = batch_size
self.dataset_size = dataset_size
@@ -45,38 +48,43 @@ class DP_SP_BatchSampler(Sampler[list[int]]):
self.num_sp_groups = num_sp_groups
self.global_rank = global_rank
self.sp_world_size = sp_world_size
self.drop_first_row = drop_first_row
self.reshuffle_each_epoch = reshuffle_each_epoch
self.candidate_indices = (
torch.as_tensor(candidate_indices, dtype=torch.long)
if candidate_indices is not None else None
)
# ── epoch-level RNG ────────────────────────────────────────────────
rng = torch.Generator().manual_seed(self.seed)
# Create a random permutation of all indices
global_indices = torch.randperm(self.dataset_size, generator=rng)
self._build_indices(0)
if drop_first_row:
# drop 0 in global_indices
def _build_indices(self, epoch: int) -> None:
rng = torch.Generator().manual_seed(self.seed + epoch)
if self.candidate_indices is None:
global_indices = torch.randperm(self.dataset_size, generator=rng)
else:
perm = torch.randperm(len(self.candidate_indices), generator=rng)
global_indices = self.candidate_indices[perm]
dataset_size = len(global_indices)
if self.drop_first_row:
global_indices = global_indices[global_indices != 0]
self.dataset_size = self.dataset_size - 1
dataset_size = len(global_indices)
if self.drop_last:
# For drop_last=True, we:
# 1. Ensure total samples is divisible by (batch_size * num_sp_groups)
# 2. This guarantees each SP group gets same number of complete batches
# 3. Prevents uneven batch sizes across SP groups at end of epoch
num_batches = self.dataset_size // self.batch_size
num_batches = dataset_size // self.batch_size
num_global_batches = num_batches // self.num_sp_groups
global_indices = global_indices[:num_global_batches *
self.num_sp_groups *
self.batch_size]
else:
if self.dataset_size % (self.num_sp_groups * self.batch_size) != 0:
# add more indices to make it divisible by (batch_size * num_sp_groups)
if dataset_size % (self.num_sp_groups * self.batch_size) != 0:
padding_size = self.num_sp_groups * self.batch_size - (
self.dataset_size % (self.num_sp_groups * self.batch_size))
dataset_size % (self.num_sp_groups * self.batch_size))
logger.info("Padding the dataset from %d to %d",
self.dataset_size, self.dataset_size + padding_size)
dataset_size, dataset_size + padding_size)
global_indices = torch.cat(
[global_indices, global_indices[:padding_size]])
# shard the indices to each sp group
ith_sp_group = self.global_rank // self.sp_world_size
sp_group_local_indices = global_indices[ith_sp_group::self.
num_sp_groups]
@@ -84,6 +92,22 @@ class DP_SP_BatchSampler(Sampler[list[int]]):
logger.info("Dataset size for each sp group: %d",
len(sp_group_local_indices))
def set_candidate_indices(
self,
candidate_indices: list[int] | None,
epoch: int = 0,
) -> None:
self.candidate_indices = (
torch.as_tensor(candidate_indices, dtype=torch.long)
if candidate_indices is not None else None
)
self._build_indices(epoch)
def set_epoch(self, epoch: int) -> None:
if not self.reshuffle_each_epoch:
return
self._build_indices(epoch)
def __iter__(self):
indices = self.sp_group_local_indices
for i in range(0, len(indices), self.batch_size):
@@ -94,19 +118,89 @@ class DP_SP_BatchSampler(Sampler[list[int]]):
return len(self.sp_group_local_indices) // self.batch_size
def _parse_data_path_specs(path: str) -> list[tuple[str, int]]:
"""
Parse data_path into a list of (directory, repeat_count).
Syntax: comma-separated entries; each entry is "path" (default 1) or "path:N" (N = repeat count).
N=0 means skip that path (convenience to disable without removing). If no ":" present, default is 1.
Example: "/dir1:2,/dir2,/dir3:0" -> dir1 2x, dir2 1x, dir3 skipped.
"""
specs: list[tuple[str, int]] = []
for part in path.split(","):
part = part.strip()
if not part:
continue
if ":" in part:
p, _, count_str = part.rpartition(":")
p = p.strip()
try:
count = int(count_str.strip())
except ValueError:
raise ValueError(
f"data_path repeat count must be an integer, got {count_str!r}"
) from None
if count < 0:
raise ValueError(
f"data_path repeat count must be >= 0, got {count}"
)
specs.append((p, count))
else:
specs.append((part, 1))
return specs
def _scan_parquet_files_for_path(p: str) -> tuple[list[str], list[int]]:
"""Return (file_paths, row_lengths) for a single directory."""
file_names: list[str] = []
for root, _, files in os.walk(p):
for file in sorted(files):
if file.endswith(".parquet"):
file_names.append(os.path.join(root, file))
lengths = []
for file_path in tqdm.tqdm(
file_names, desc="Reading parquet files to get lengths"):
lengths.append(pq.ParquetFile(file_path).metadata.num_rows)
logger.info("Found %d parquet files with %d total rows", len(file_names),
sum(lengths))
return file_names, lengths
def get_parquet_files_and_length(path: str):
dataset_root = os.path.realpath(os.path.expanduser(path))
# Check if cached info exists
cache_dir = os.path.join(dataset_root, "map_style_cache")
cache_file = os.path.join(cache_dir, "file_info.pkl")
"""
Collect parquet file paths and row lengths from one or more directories.
path: single directory, or comma-separated "path" or "path:N" (N = repeat count).
E.g. "/dir1:2,/dir2:1" -> dir1's files appear 2x (oversampled), dir2 once.
"""
path_specs = _parse_data_path_specs(path)
if not path_specs:
raise ValueError(
"data_path must be a non-empty path or comma-separated path specs"
)
first_path = next((p for p, c in path_specs if c > 0), path_specs[0][0])
is_single_no_repeat = len(path_specs) == 1 and path_specs[0][1] == 1
effective_path = path.strip()
if is_single_no_repeat:
cache_dir = os.path.join(first_path, "map_style_cache")
cache_suffix = "file_info.pkl"
else:
neutral_root = os.environ.get(
"FASTVIDEO_MAP_STYLE_CACHE_DIR",
os.path.join(os.path.expanduser("~"), ".cache", "fastvideo",
"map_style_cache"),
)
cache_dir = neutral_root
cache_suffix = ("file_info_" +
hashlib.md5(effective_path.encode()).hexdigest()[:16] +
".pkl")
cache_file = os.path.join(cache_dir, cache_suffix)
# Only rank 0 checks for cache and scans files if needed
if get_world_rank() == 0:
cache_loaded = False
file_names_sorted = None
lengths_sorted = None
# First try to load existing cache
if os.path.exists(cache_file):
logger.info("Loading cached file info from %s", cache_file)
try:
@@ -117,24 +211,11 @@ def get_parquet_files_and_length(path: str):
os.path.join(os.getcwd(), p)
if not os.path.isabs(p) else p)
for p in file_names_sorted)
files_outside_dataset_root = [
file_path for file_path in file_names_sorted
if os.path.commonpath([dataset_root, file_path
]) != dataset_root
]
missing_files = [
file_path for file_path in file_names_sorted
if not os.path.exists(file_path)
]
if files_outside_dataset_root:
logger.warning(
"Cached parquet file list points outside dataset root "
"(%s). Cache will be rebuilt. First out-of-root file: %s",
dataset_root,
files_outside_dataset_root[0],
)
cache_loaded = False
elif missing_files:
if missing_files:
logger.warning(
"Cached parquet file list contains %d missing files. "
"Cache will be rebuilt. First missing file: %s",
@@ -150,43 +231,43 @@ def get_parquet_files_and_length(path: str):
logger.info("Falling back to scanning files")
cache_loaded = False
# If cache not loaded (either doesn't exist or failed to load), scan files
if not cache_loaded:
logger.info("Scanning parquet files to get lengths")
lengths = []
file_names = []
for root, _, files in os.walk(dataset_root):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.realpath(os.path.join(root, file))
file_names.append(file_path)
if len(file_names) == 0:
logger.info(
"Scanning parquet files (path specs: %s)",
[(p, c) for p, c in path_specs],
)
combined: list[tuple[str, int, int]] = []
sort_index = 0
for p, count in path_specs:
if count == 0:
continue
fnames, lens = _scan_parquet_files_for_path(p)
if not fnames:
logger.warning("No parquet files found under path spec %s", p)
continue
for _ in range(count):
for f, ln in zip(fnames, lens, strict=True):
combined.append((f, ln, sort_index))
sort_index += 1
if len(combined) == 0:
raise FileNotFoundError(
"No parquet files found under dataset path: "
f"{path}. "
"Please verify this path points to preprocessed parquet "
"data.")
for file_path in tqdm.tqdm(
file_names, desc="Reading parquet files to get lengths"):
num_rows = pq.ParquetFile(file_path).metadata.num_rows
lengths.append(num_rows)
# sort according to file name to ensure all rank has the same order
file_names_sorted, lengths_sorted = zip(*sorted(zip(file_names,
lengths,
strict=True),
key=lambda x: x[0]),
strict=True)
# Save the cache
file_names_sorted, lengths_sorted, _ = zip(
*sorted(combined, key=lambda x: (x[0], x[2])), strict=True)
os.makedirs(cache_dir, exist_ok=True)
with open(cache_file, "wb") as f:
pickle.dump((file_names_sorted, lengths_sorted), f)
logger.info("Saved file info to %s", cache_file)
# Wait for rank 0 to finish creating/loading cache
world_group = get_world_group()
world_group.barrier()
# Now all ranks load the cache (it should exist and be valid now)
logger.info("Loading cached file info from %s after barrier", cache_file)
with open(cache_file, "rb") as f:
file_names_sorted, lengths_sorted = pickle.load(f)
@@ -366,6 +447,84 @@ def passthrough(batch):
return batch
def build_bot_died_excluded_indices(
data_path: str,
parquet_files: list[str],
lengths: list[int],
) -> set[int]:
"""Build global row indices to exclude based on per-dir bot_died.json."""
import json
path_specs = [
(os.path.realpath(os.path.expanduser(p)), count)
for p, count in _parse_data_path_specs(data_path)
]
bot_died_per_dir: dict[str, set[int]] = {}
for p, count in path_specs:
if count == 0:
continue
candidates = [
os.path.join(os.path.dirname(p), "filter", "blue_water_random_half.json"),
os.path.join(
os.path.dirname(os.path.abspath(os.path.expanduser(p))),
"filter",
"blue_water_random_half.json",
),
]
filter_path = next((fp for fp in candidates if os.path.exists(fp)), None)
if filter_path is None:
continue
with open(filter_path, "r", encoding="utf-8") as f:
bot_died_per_dir[p] = set(json.load(f))
logger.info(
"Loaded bot_died filter from %s: %d entries to exclude",
filter_path,
len(bot_died_per_dir[p]),
)
if not bot_died_per_dir:
return set()
excluded: set[int] = set()
matched_file_count: dict[str, int] = {k: 0 for k in bot_died_per_dir}
global_offset = 0
for file_path, length in zip(parquet_files, lengths, strict=True):
file_path_real = os.path.realpath(file_path)
matching_dir = None
for dir_path in bot_died_per_dir:
if (file_path.startswith(dir_path)
or file_path_real.startswith(dir_path)):
matching_dir = dir_path
break
if matching_dir is not None:
matched_file_count[matching_dir] += 1
bot_died_set = bot_died_per_dir[matching_dir]
try:
table = pq.read_table(file_path, columns=["file_name"])
file_names = table.column("file_name").to_pylist()
for local_idx, fn in enumerate(file_names):
if int(str(fn).strip()) in bot_died_set:
excluded.add(global_offset + local_idx)
except Exception as e:
logger.warning(
"Failed to read file_name from %s for bot_died filter: %s",
file_path,
e,
)
global_offset += length
for dir_path, count in matched_file_count.items():
logger.info(
"bot_died matching: dir=%s matched_parquet_files=%d",
dir_path,
count,
)
return excluded
def build_parquet_map_style_dataloader(
path,
batch_size,
+29 -1
View File
@@ -4,6 +4,7 @@ import os
import pathlib
import datasets
import numpy as np
from torch.utils.data import IterableDataset
from fastvideo.distributed import (get_sp_world_size, get_world_rank,
@@ -16,8 +17,9 @@ logger = init_logger(__name__)
class ValidationDataset(IterableDataset):
def __init__(self, filename: str):
def __init__(self, filename: str, num_samples: int | None = None):
super().__init__()
self.num_samples = num_samples
self.filename = pathlib.Path(filename)
# get directory of filename
@@ -58,6 +60,12 @@ class ValidationDataset(IterableDataset):
# Convert to list to get total samples
self.all_samples = list(data)
# Limit number of samples if specified
if self.num_samples is not None and self.num_samples < len(self.all_samples):
self.all_samples = self.all_samples[:self.num_samples]
logger.info("Limiting validation samples to %s", self.num_samples)
self.original_total_samples = len(self.all_samples)
# Extend samples to be a multiple of DP degree (num_sp_groups)
@@ -160,5 +168,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
+150
View File
@@ -10,6 +10,7 @@ Adapted from HY-WorldPlay: https://github.com/Tencent-Hunyuan/HY-WorldPlay
"""
import json
import logging
import numpy as np
import torch
from scipy.spatial.transform import Rotation as R
@@ -18,6 +19,9 @@ from typing import Union, Optional
from .trajectory import generate_camera_trajectory_local
logger = logging.getLogger(__name__)
# Mapping from one-hot action encoding to single label
mapping = {
(0, 0, 0, 0): 0,
@@ -305,6 +309,152 @@ def camera_center_normalization(w2c: np.ndarray) -> np.ndarray:
return np.linalg.inv(c2w_aligned)
def reformat_keyboard_and_mouse_tensors(
keyboard_tensor: torch.Tensor,
mouse_tensor: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Collapse frame-level keyboard/mouse controls to latent-level controls.
The first frame is the anchor frame. The remaining frames are grouped into
chunks of 4 and each chunk is expected to be constant.
"""
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])
if not (groups == groups[:, 0:1]).all(dim=1).all():
logger.warning("keyboard_tensor has different values per 4-frame group")
groups = mouse_tensor.view(-1, 4, mouse_tensor.shape[1])
if not (groups == groups[:, 0:1]).all(dim=1).all():
logger.warning("mouse_tensor has different values per 4-frame group")
return keyboard_tensor[::4], mouse_tensor[::4]
def process_custom_actions(
keyboard_tensor: torch.Tensor,
mouse_tensor: torch.Tensor,
forward_speed: float = DEFAULT_FORWARD_SPEED,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""
Convert custom keyboard/mouse controls into viewmats, intrinsics, and labels.
"""
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: list[dict[str, float]] = []
for t in range(keyboard_tensor.shape[0]):
frame_motion: dict[str, float] = {}
fwd = 0.0
if keyboard_tensor[t, 0] > 0.5:
fwd += forward_speed
if keyboard_tensor[t, 1] > 0.5:
fwd -= forward_speed
if fwd != 0.0:
frame_motion["forward"] = fwd
rgt = 0.0
if keyboard_tensor[t, 2] > 0.5:
rgt -= forward_speed
if keyboard_tensor[t, 3] > 0.5:
rgt += forward_speed
if rgt != 0.0:
frame_motion["right"] = rgt
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)
poses = generate_camera_trajectory_local(motions)
w2c_list = []
intrinsic_list = []
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 pose in poses:
c2w = np.array(pose)
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))
c2ws = np.linalg.inv(np.array(w2c_list))
c_inv = np.linalg.inv(c2ws[:-1])
relative_c2w = np.zeros_like(c2ws)
relative_c2w[0, ...] = c2ws[0, ...]
relative_c2w[1:, ...] = c_inv @ c2ws[1:, ...]
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
for i in range(1, relative_c2w.shape[0]):
move_dirs = relative_c2w[i, :3, 3]
move_norms = np.linalg.norm(move_dirs)
if move_norms > move_norm_valid:
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)
else:
trans_angles_deg = np.zeros(3)
r_rel = relative_c2w[i, :3, :3]
rot_angles_deg = R.from_matrix(r_rel).as_euler("xyz", degrees=True)
if move_norms > move_norm_valid:
if trans_angles_deg[2] < 60:
trans_one_hot[i, 0] = 1
elif trans_angles_deg[2] > 120:
trans_one_hot[i, 1] = 1
if trans_angles_deg[0] < 60:
trans_one_hot[i, 2] = 1
elif trans_angles_deg[0] > 120:
trans_one_hot[i, 3] = 1
if rot_angles_deg[1] > 5e-2:
rotate_one_hot[i, 0] = 1
elif rot_angles_deg[1] < -5e-2:
rotate_one_hot[i, 1] = 1
if rot_angles_deg[0] > 5e-2:
rotate_one_hot[i, 2] = 1
elif rot_angles_deg[0] < -5e-2:
rotate_one_hot[i, 3] = 1
trans_label = one_hot_to_one_dimension(torch.tensor(trans_one_hot))
rotate_label = one_hot_to_one_dimension(torch.tensor(rotate_one_hot))
action_labels = trans_label * 9 + rotate_label
return viewmats, intrinsics, action_labels
def parse_pose_string_to_actions(pose_string: str, fps: int = 24) -> list[dict]:
"""
+203
View File
@@ -299,6 +299,209 @@ def parse_config(config, mode="universal"):
)
return key_data, mouse_data
def _get_cv2():
import cv2
return cv2
def draw_rounded_rectangle(
image,
top_left,
bottom_right,
color,
radius=10,
alpha=0.5,
):
cv2 = _get_cv2()
overlay = image.copy()
x1, y1 = top_left
x2, y2 = bottom_right
cv2.rectangle(overlay, (x1 + radius, y1), (x2 - radius, y2), color, -1)
cv2.rectangle(overlay, (x1, y1 + radius), (x2, y2 - radius), color, -1)
cv2.ellipse(
overlay,
(x1 + radius, y1 + radius),
(radius, radius),
180,
0,
90,
color,
-1,
)
cv2.ellipse(
overlay,
(x2 - radius, y1 + radius),
(radius, radius),
270,
0,
90,
color,
-1,
)
cv2.ellipse(
overlay,
(x1 + radius, y2 - radius),
(radius, radius),
90,
0,
90,
color,
-1,
)
cv2.ellipse(
overlay,
(x2 - radius, y2 - radius),
(radius, radius),
0,
0,
90,
color,
-1,
)
cv2.addWeighted(overlay, alpha, image, 1 - alpha, 0, image)
def draw_keys_on_frame(
frame,
keys,
key_size=(30, 30),
spacing=5,
top_margin=15,
mode="universal",
):
"""Draw keyboard action badges on the top-left of the frame."""
cv2 = _get_cv2()
del spacing # Preserved for compatibility with the original helper.
left_margin = 15
gap = 3
key_positions = {
"W": (left_margin + key_size[0] + gap, top_margin),
"A": (left_margin, top_margin + key_size[1] + gap),
"S": (left_margin + key_size[0] + gap, top_margin + key_size[1] + gap),
"D": (
left_margin + (key_size[0] + gap) * 2,
top_margin + key_size[1] + gap,
),
}
key_icon = {
"W": "W",
"A": "A",
"S": "S",
"D": "D",
"left": "L",
"right": "R",
}
if mode == "templerun":
key_positions.update(
{
"left": (
left_margin + (key_size[0] + gap) * 3 + 10,
top_margin + key_size[1] + gap,
),
"right": (
left_margin + (key_size[0] + gap) * 4 + 15,
top_margin + key_size[1] + gap,
),
}
)
for key, (x, y) in key_positions.items():
is_pressed = keys.get(key, False)
top_left = (x, y)
bottom_right = (x + key_size[0], y + key_size[1])
color = (0, 255, 0) if is_pressed else (200, 200, 200)
alpha = 0.8 if is_pressed else 0.5
draw_rounded_rectangle(
frame,
top_left,
bottom_right,
color,
radius=5,
alpha=alpha,
)
text = key_icon[key]
text_size = cv2.getTextSize(
text,
cv2.FONT_HERSHEY_SIMPLEX,
0.5,
1,
)[0]
text_x = x + (key_size[0] - text_size[0]) // 2
text_y = y + (key_size[1] + text_size[1]) // 2
cv2.putText(
frame,
text,
(text_x, text_y),
cv2.FONT_HERSHEY_SIMPLEX,
0.5,
(0, 0, 0),
1,
)
def draw_mouse_on_frame(frame, pitch, yaw, top_margin=15):
"""Draw a mouse-look crosshair with a direction arrow."""
cv2 = _get_cv2()
h, w, _ = frame.shape
right_margin = 15
crosshair_radius = 25
crosshair_x = w - right_margin - crosshair_radius
crosshair_y = top_margin + crosshair_radius
dx = int(yaw * crosshair_radius * 8)
dy = int(-pitch * crosshair_radius * 8)
max_arrow = crosshair_radius - 5
dx = max(-max_arrow, min(max_arrow, dx))
dy = max(-max_arrow, min(max_arrow, dy))
cv2.circle(
frame,
(crosshair_x, crosshair_y),
crosshair_radius,
(50, 50, 50),
-1,
)
cv2.circle(
frame,
(crosshair_x, crosshair_y),
crosshair_radius,
(200, 200, 200),
1,
)
cv2.line(
frame,
(crosshair_x - crosshair_radius + 5, crosshair_y),
(crosshair_x + crosshair_radius - 5, crosshair_y),
(100, 100, 100),
1,
)
cv2.line(
frame,
(crosshair_x, crosshair_y - crosshair_radius + 5),
(crosshair_x, crosshair_y + crosshair_radius - 5),
(100, 100, 100),
1,
)
if abs(dx) > 1 or abs(dy) > 1:
cv2.arrowedLine(
frame,
(crosshair_x, crosshair_y),
(crosshair_x + dx, crosshair_y + dy),
(0, 255, 0),
2,
tipLength=0.3,
)
# NOTE: drawing functions are commented out to avoid cv2/libGL dependency.
#
# def draw_rounded_rectangle(image, top_left, bottom_right, color, radius=10, alpha=0.5):
+16
View File
@@ -0,0 +1,16 @@
from .model import WanGameActionTransformer3DModel
from .causal_model import CausalWanGameActionTransformer3DModel
from .hyworld_action_module import WanGameActionTimeImageEmbedding, WanGameActionSelfAttention
__all__ = [
"WanGameActionTransformer3DModel",
"CausalWanGameActionTransformer3DModel",
"WanGameActionTimeImageEmbedding",
"WanGameActionSelfAttention",
]
# Entry point for model registry
EntryClass = [
WanGameActionTransformer3DModel,
CausalWanGameActionTransformer3DModel,
]
@@ -0,0 +1,863 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import torch
import torch.nn as nn
from torch.nn.attention.flex_attention import create_block_mask, flex_attention
from torch.nn.attention.flex_attention import BlockMask
# wan 1.3B model has a weird channel / head configurations and require max-autotune to work with flexattention
# see https://github.com/pytorch/pytorch/issues/133254
# change to default for other models
flex_attention = torch.compile(
flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
import torch.distributed as dist
from fastvideo.attention import LocalAttention
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
from .hyworld_action_module import (
WanGameActionTimeImageEmbedding,
WanGameActionSelfAttention,
)
from fastvideo.models.dits.hyworld.camera_rope import prope_qkv
logger = init_logger(__name__)
_DEFAULT_WANGAME_CONFIG = WanGameVideoConfig()
class CausalWanGameCrossAttention(WanI2VCrossAttention):
"""Cross-attention for WanGame causal model"""
def forward(self, x, context, context_lens=None, crossattn_cache=None):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
context_lens(Tensor): Shape [B]
crossattn_cache: Optional cache dict for inference
"""
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)
if crossattn_cache is not None:
if not crossattn_cache["is_init"]:
crossattn_cache["is_init"] = True
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)
crossattn_cache["k"] = k_img
crossattn_cache["v"] = v_img
else:
k_img = crossattn_cache["k"]
v_img = crossattn_cache["v"]
else:
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 CausalWanGameActionSelfAttention(WanGameActionSelfAttention):
def __init__(self,
dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm=True,
eps=1e-6) -> None:
super().__init__(
dim=dim,
num_heads=num_heads,
local_attn_size=local_attn_size,
sink_size=sink_size,
qk_norm=qk_norm,
eps=eps,
)
self.max_attention_size = 32760 if local_attn_size == -1 else local_attn_size * 1560
# Local attention for KV-cache inference
self.local_attn = LocalAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
@staticmethod
def _masked_flex_attn(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
block_mask: BlockMask,
) -> torch.Tensor:
padded_length = math.ceil(query.shape[1] / 128) * 128 - query.shape[1]
if padded_length > 0:
query = torch.cat(
[
query,
torch.zeros(
[query.shape[0], padded_length, query.shape[2], query.shape[3]],
device=query.device,
dtype=value.dtype,
),
],
dim=1,
)
key = torch.cat(
[
key,
torch.zeros(
[key.shape[0], padded_length, key.shape[2], key.shape[3]],
device=key.device,
dtype=value.dtype,
),
],
dim=1,
)
value = torch.cat(
[
value,
torch.zeros(
[value.shape[0], padded_length, value.shape[2], value.shape[3]],
device=value.device,
dtype=value.dtype,
),
],
dim=1,
)
out = flex_attention(
query=query.transpose(2, 1),
key=key.transpose(2, 1),
value=value.transpose(2, 1),
block_mask=block_mask,
).transpose(2, 1)
if padded_length > 0:
out = out[:, :-padded_length]
return out
def forward(self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
block_mask: BlockMask | None = None,
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):
"""
Forward pass with causal attention.
"""
if cache_start is None:
cache_start = current_start
if kv_cache is None:
if block_mask is None:
raise ValueError(
"block_mask must be provided for causal training attention")
if viewmats is None or Ks is None:
raise ValueError(
"viewmats and Ks must be provided for WanGame causal attention")
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)
rope_output = self._masked_flex_attn(
query_rope, key_rope, v, block_mask)
# PRoPE path with the same causal mask.
query_prope, key_prope, value_prope, apply_fn_o = prope_qkv(
q.transpose(1, 2),
k.transpose(1, 2),
v.transpose(1, 2),
viewmats=viewmats,
Ks=Ks,
patches_x=40,
patches_y=22,
)
query_prope = query_prope.transpose(1, 2)
key_prope = key_prope.transpose(1, 2)
value_prope = value_prope.transpose(1, 2)
prope_output = self._masked_flex_attn(
query_prope, key_prope, value_prope, block_mask)
prope_output = apply_fn_o(
prope_output.transpose(1, 2)).transpose(1, 2)
return rope_output, prope_output
else:
# Inference mode with KV cache
if viewmats is None or Ks is None:
raise ValueError(
"viewmats and Ks must be provided for WanGame causal attention")
cos, sin = freqs_cis
roped_query = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v)
roped_key = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
query_prope, key_prope, value_prope, apply_fn_o = prope_qkv(
q.transpose(1, 2),
k.transpose(1, 2),
v.transpose(1, 2),
viewmats=viewmats,
Ks=Ks,
patches_x=40,
patches_y=22,
)
query_prope = query_prope.transpose(1, 2).type_as(v)
key_prope = key_prope.transpose(1, 2).type_as(v)
value_prope = value_prope.transpose(1, 2).type_as(v)
frame_seqlen = q.shape[1]
current_end = current_start + roped_query.shape[1]
sink_tokens = self.sink_size * frame_seqlen
# If we are using local attention and the current KV cache size is larger than the local attention size, we need to truncate the KV cache
kv_cache_size = kv_cache["k"].shape[1]
num_new_tokens = roped_query.shape[1]
# rope+prope
cache_head_dim = kv_cache["k"].shape[-1]
local_end_index = kv_cache["local_end_index"].item()
# read cache but never mutate it.
if not is_cache:
if cache_head_dim not in (self.head_dim, self.head_dim * 2):
raise ValueError(
f"Unexpected kv_cache head dim: {cache_head_dim}, "
f"expected {self.head_dim} or {self.head_dim * 2}")
cache_k_rope = kv_cache["k"][..., :self.head_dim]
cache_v_rope = kv_cache["v"][..., :self.head_dim]
rope_k = torch.cat(
[cache_k_rope[:, :local_end_index], roped_key], dim=1)
rope_v = torch.cat(
[cache_v_rope[:, :local_end_index], v], dim=1)
rope_k = rope_k[:, -self.max_attention_size:]
rope_v = rope_v[:, -self.max_attention_size:]
rope_x = self.local_attn(roped_query, rope_k, rope_v)
if cache_head_dim == self.head_dim * 2:
cache_k_prope = kv_cache["k"][..., self.head_dim:]
cache_v_prope = kv_cache["v"][..., self.head_dim:]
prope_k = torch.cat(
[cache_k_prope[:, :local_end_index], key_prope], dim=1)
prope_v = torch.cat(
[cache_v_prope[:, :local_end_index], value_prope], dim=1)
prope_k = prope_k[:, -self.max_attention_size:]
prope_v = prope_v[:, -self.max_attention_size:]
prope_x = self.local_attn(query_prope, prope_k, prope_v)
else:
prope_x = self.local_attn(
query_prope, key_prope, value_prope)
prope_x = apply_fn_o(prope_x.transpose(1, 2)).transpose(1, 2)
return rope_x, prope_x
# update cache.
if cache_head_dim == self.head_dim:
kv_cache["k"] = torch.cat(
[kv_cache["k"], torch.zeros_like(kv_cache["k"])], dim=-1)
kv_cache["v"] = torch.cat(
[kv_cache["v"], torch.zeros_like(kv_cache["v"])], dim=-1)
elif cache_head_dim != self.head_dim * 2:
raise ValueError(
f"Unexpected kv_cache head dim: {cache_head_dim}, "
f"expected {self.head_dim} or {self.head_dim * 2}")
cache_k_rope = kv_cache["k"][..., :self.head_dim]
cache_k_prope = kv_cache["k"][..., self.head_dim:]
cache_v_rope = kv_cache["v"][..., :self.head_dim]
cache_v_prope = kv_cache["v"][..., self.head_dim:]
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
# Calculate the number of new tokens added in this step
# Shift existing cache content left to discard oldest tokens
# Clone the source slice to avoid overlapping memory error
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
cache_k_rope[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
cache_k_rope[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
cache_v_rope[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
cache_v_rope[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
cache_k_prope[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
cache_k_prope[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
cache_v_prope[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
cache_v_prope[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache["local_end_index"].item() + current_end - \
kv_cache["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
cache_k_rope[:, local_start_index:local_end_index] = roped_key
cache_v_rope[:, local_start_index:local_end_index] = v
cache_k_prope[:, local_start_index:local_end_index] = key_prope
cache_v_prope[:, local_start_index:local_end_index] = value_prope
else:
# Assign new keys/values directly up to current_end
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache["k"] = kv_cache["k"].detach()
kv_cache["v"] = kv_cache["v"].detach()
cache_k_rope = kv_cache["k"][..., :self.head_dim]
cache_k_prope = kv_cache["k"][..., self.head_dim:]
cache_v_rope = kv_cache["v"][..., :self.head_dim]
cache_v_prope = kv_cache["v"][..., self.head_dim:]
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
cache_k_rope[:, local_start_index:local_end_index] = roped_key
cache_v_rope[:, local_start_index:local_end_index] = v
cache_k_prope[:, local_start_index:local_end_index] = key_prope
cache_v_prope[:, local_start_index:local_end_index] = value_prope
rope_x = self.local_attn(
roped_query,
cache_k_rope[:, max(0, local_end_index - self.max_attention_size):local_end_index],
cache_v_rope[:, max(0, local_end_index - self.max_attention_size):local_end_index]
)
prope_x = self.local_attn(
query_prope,
cache_k_prope[:, max(0, local_end_index - self.max_attention_size):local_end_index],
cache_v_prope[:, max(0, local_end_index - self.max_attention_size):local_end_index]
)
prope_x = apply_fn_o(prope_x.transpose(1, 2)).transpose(1, 2)
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
return rope_x, prope_x
class CausalWanGameActionTransformerBlock(nn.Module):
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 = CausalWanGameActionSelfAttention(
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)
self.attn2 = CausalWanGameCrossAttention(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 = ReplicatedLinear(dim, dim, bias=True)
nn.init.zeros_(self.to_out_prope.weight)
if self.to_out_prope.bias is not None:
nn.init.zeros_(self.to_out_prope.bias)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
block_mask: BlockMask | None = None,
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 camera PRoPE
attn_output_rope, attn_output_prope = self.attn1(
query, key, value, freqs_cis,
block_mask, 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(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,
crossattn_cache=crossattn_cache)
# 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)
return hidden_states
class CausalWanGameActionTransformer3DModel(BaseDiT):
supports_action_input = True
_fsdp_shard_conditions = _DEFAULT_WANGAME_CONFIG._fsdp_shard_conditions
_compile_conditions = _DEFAULT_WANGAME_CONFIG._compile_conditions
_supported_attention_backends = (
_DEFAULT_WANGAME_CONFIG._supported_attention_backends
)
param_names_mapping = _DEFAULT_WANGAME_CONFIG.param_names_mapping
reverse_param_names_mapping = (
_DEFAULT_WANGAME_CONFIG.reverse_param_names_mapping
)
lora_param_names_mapping = _DEFAULT_WANGAME_CONFIG.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
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([
CausalWanGameActionTransformerBlock(
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.block_mask = None
self.num_frame_per_block = config.arch_config.num_frames_per_block
assert self.num_frame_per_block <= 3
self.__post_init__()
@staticmethod
def _prepare_blockwise_causal_attn_mask(
device: torch.device | str, num_frames: int = 21,
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1
) -> BlockMask:
"""
we will divide the token sequence into the following format
[1 latent frame] [1 latent frame] ... [1 latent frame]
We use flexattention to construct the attention mask
"""
total_length = num_frames * frame_seqlen
# we do right padding to get to a multiple of 128
padded_length = math.ceil(total_length / 128) * 128 - total_length
ends = torch.zeros(total_length + padded_length,
device=device, dtype=torch.long)
# Block-wise causal mask will attend to all elements that are before the end of the current chunk
frame_indices = torch.arange(
start=0,
end=total_length,
step=frame_seqlen * num_frame_per_block,
device=device
)
for tmp in frame_indices:
ends[tmp:tmp + frame_seqlen * num_frame_per_block] = tmp + \
frame_seqlen * num_frame_per_block
def attention_mask(b, h, q_idx, kv_idx):
if local_attn_size == -1:
return (kv_idx < ends[q_idx]) | (q_idx == kv_idx)
else:
return ((kv_idx < ends[q_idx]) & (kv_idx >= (ends[q_idx] - local_attn_size * frame_seqlen))) | (q_idx == kv_idx)
block_mask = create_block_mask(attention_mask, B=None, H=None, Q_LEN=total_length + padded_length,
KV_LEN=total_length + padded_length, _compile=False, device=device)
if not dist.is_initialized() or dist.get_rank() == 0:
print(
f" cache a block wise causal mask with block size of {num_frame_per_block} frames")
print(block_mask)
return block_mask
def _forward_inference(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
guidance=None,
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:
r"""
Run the diffusion model with kv caching.
See Algorithm 2 of CausVid paper https://arxiv.org/abs/2412.07772 for details.
This function will be run for num_frame times.
Process the latent frames one by one (1560 tokens each)
"""
orig_dtype = hidden_states.dtype
if isinstance(encoder_hidden_states_image, list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
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)
# condition_embedder returns:
# - temb: [B*T, dim] where T = post_patch_num_frames
# - timestep_proj: [B*T, 6*dim]
# Reshape to [B, T, 6, dim] for transformer blocks
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)) # [B*T, 6, dim]
timestep_proj = timestep_proj.view(batch_size, post_patch_num_frames, 6, self.hidden_size) # [B, T, 6, dim]
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,
self.block_mask,
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,
block_mask=self.block_mask,
kv_cache=kv_cache[block_idx] if kv_cache else None,
crossattn_cache=crossattn_cache[block_idx] if crossattn_cache else None,
current_start=current_start, cache_start=cache_start,
viewmats=viewmats, Ks=Ks, is_cache=is_cache)
# If cache-only mode, return early
if is_cache:
return kv_cache
# Output norm, projection & unpatchify
# temb is [B*T, dim], reshape to [B, T, 1, dim]
temb = temb.view(batch_size, post_patch_num_frames, -1).unsqueeze(2) # [B, T, 1, dim]
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
def _forward_train(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
guidance=None,
action: torch.Tensor | None = None,
viewmats: torch.Tensor | None = None,
Ks: torch.Tensor | None = None,
start_frame: int = 0,
**kwargs
) -> torch.Tensor:
orig_dtype = hidden_states.dtype
if isinstance(encoder_hidden_states_image, list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
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
# Construct blockwise causal attn mask
if self.block_mask is None:
self.block_mask = self._prepare_blockwise_causal_attn_mask(
device=hidden_states.device,
num_frames=num_frames,
frame_seqlen=post_patch_height * post_patch_width,
num_frame_per_block=self.num_frame_per_block,
local_attn_size=self.local_attn_size
)
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)
# condition_embedder returns:
# - temb: [B*T, dim] where T = post_patch_num_frames
# - timestep_proj: [B*T, 6*dim]
# Reshape to [B, T, 6, dim] for transformer blocks
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)) # [B*T, 6, dim]
timestep_proj = timestep_proj.view(batch_size, post_patch_num_frames, 6, self.hidden_size) # [B, T, 6, dim]
encoder_hidden_states = encoder_hidden_states_image
# Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
self.block_mask,
None, None, # kv_cache, crossattn_cache
0, None, # current_start, cache_start
viewmats, Ks, False) # viewmats, Ks, is_cache
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
block_mask=self.block_mask,
viewmats=viewmats, Ks=Ks)
# Output norm, projection & unpatchify
# temb is [B*T, dim], reshape to [B, T, 1, dim]
temb = temb.view(batch_size, post_patch_num_frames, -1).unsqueeze(2) # [B, T, 1, dim]
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
def forward(
self,
*args,
**kwargs
):
if kwargs.get('kv_cache', None) is not None:
return self._forward_inference(*args, **kwargs)
else:
return self._forward_train(*args, **kwargs)
@@ -0,0 +1,284 @@
import math
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 fc_in with kaiming_uniform (same as nn.Linear default)
nn.init.kaiming_uniform_(self.action_embedder.fc_in.weight, a=math.sqrt(5))
# Initialize fc_out 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,
):
"""
Args:
timestep: [B] diffusion timesteps (one per batch sample)
action: [B, T] action labels (one per frame per batch sample)
Returns:
temb: [B*T, dim] combined timestep + action embedding
timestep_proj: [B*T, 6*dim] modulation projection
"""
# timestep may be [B] (one per sample) or [B*T] (one per frame, from causal training)
temb = self.time_embedder(timestep, timestep_seq_len)
# Handle action embedding for batch > 1
# action shape: [B, T] where B=batch_size, T=num_frames
batch_size = action.shape[0]
num_frames = action.shape[1]
# Compute action embeddings: [B, T] -> [B*T] -> [B*T, dim]
action_flat = action.flatten() # [B*T]
action_emb = timestep_embedding(action_flat, 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) # [B*T, dim]
# temb is [B*T, dim] when timestep was already per-frame (causal training),
# or [B, dim] when timestep is per-sample (inference).
# Only expand if temb is per-sample [B, dim].
if temb.shape[0] == batch_size and num_frames > 1:
# Expand temb: [B, dim] -> [B, T, dim] -> [B*T, dim]
temb_expanded = temb.unsqueeze(1).expand(-1, num_frames, -1) # [B, T, dim]
temb_expanded = temb_expanded.reshape(batch_size * num_frames, -1) # [B*T, dim]
else:
# temb is already [B*T, dim] (per-frame timesteps)
temb_expanded = temb
# Add action embedding to expanded temb
temb = temb_expanded + action_emb # [B*T, dim]
timestep_proj = self.time_modulation(temb) # [B*T, 6*dim]
# 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((batch_size, 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
# # DEBUG: Check camera matrices
# if self.training and torch.distributed.get_rank() == 0:
# vm_info = f"viewmats={viewmats.shape if viewmats is not None else None}"
# ks_info = f"Ks={Ks.shape if Ks is not None else None}"
# vm_nonzero = (viewmats != 0).sum().item() if viewmats is not None else 0
# ks_nonzero = (Ks != 0).sum().item() if Ks is not None else 0
# print(f"[DEBUG] PRoPE input: {vm_info} nonzero={vm_nonzero}, {ks_info} nonzero={ks_nonzero}", flush=True)
# 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)
# # DEBUG: Check prope_qkv output
# if self.training and torch.distributed.get_rank() == 0:
# q_nz = (query_prope != 0).sum().item()
# k_nz = (key_prope != 0).sum().item()
# v_nz = (value_prope != 0).sum().item()
# print(f"[DEBUG] prope_qkv output: q_nonzero={q_nz}, k_nonzero={k_nz}, v_nonzero={v_nz}", flush=True)
# 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")
else:
# Same sequence length: use DistributedAttention (supports SP)
# Create default attention mask if not provided
# NOTE: query_all has shape [2*B, L, ...] (rope+prope concatenated), so mask needs 2*B
if attention_mask is None:
batch_size, seq_len = q.shape[0], q.shape[1]
attention_mask = torch.ones(batch_size * 2, 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,
)
hidden_states_rope, hidden_states_prope = hidden_states_all.chunk(2, dim=0)
# # DEBUG: Check attention output and apply_fn_o
# if self.training and torch.distributed.get_rank() == 0:
# attn_all_nz = (hidden_states_all != 0).sum().item()
# rope_nz = (hidden_states_rope != 0).sum().item()
# prope_before = (hidden_states_prope != 0).sum().item()
# print(f"[DEBUG] attn output: all_nonzero={attn_all_nz}, rope_nonzero={rope_nz}, prope_before_apply={prope_before}", flush=True)
hidden_states_prope = apply_fn_o(hidden_states_prope.transpose(1, 2)).transpose(1, 2)
# # DEBUG: Check after apply_fn_o
# if self.training and torch.distributed.get_rank() == 0:
# prope_after = (hidden_states_prope != 0).sum().item()
# print(f"[DEBUG] prope_after_apply_fn_o={prope_after}", flush=True)
return hidden_states_rope, hidden_states_prope
+433
View File
@@ -0,0 +1,433 @@
# 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
from .hyworld_action_module import (
WanGameActionSelfAttention,
WanGameActionTimeImageEmbedding,
)
logger = init_logger(__name__)
_DEFAULT_WANGAME_CONFIG = WanGameVideoConfig()
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 = ReplicatedLinear(dim, dim, bias=True)
nn.init.zeros_(self.to_out_prope.weight)
if self.to_out_prope.bias is not None:
nn.init.zeros_(self.to_out_prope.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 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)
# # DEBUG: Check if prope input is zero
# if self.training and torch.distributed.get_rank() == 0:
# prope_nonzero = (attn_output_prope != 0).sum().item()
# prope_total = attn_output_prope.numel()
# if prope_nonzero == 0:
# print(f"[DEBUG] to_out_prope INPUT is ALL ZEROS! shape={attn_output_prope.shape}", flush=True)
attn_output_prope, _ = self.to_out_prope(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
"""
supports_action_input = True
_fsdp_shard_conditions = _DEFAULT_WANGAME_CONFIG._fsdp_shard_conditions
_compile_conditions = _DEFAULT_WANGAME_CONFIG._compile_conditions
_supported_attention_backends = (
_DEFAULT_WANGAME_CONFIG._supported_attention_backends
)
param_names_mapping = _DEFAULT_WANGAME_CONFIG.param_names_mapping
reverse_param_names_mapping = (
_DEFAULT_WANGAME_CONFIG.reverse_param_names_mapping
)
lora_param_names_mapping = _DEFAULT_WANGAME_CONFIG.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)
# condition_embedder returns:
# - temb: [B*T, dim] where T = post_patch_num_frames
# - timestep_proj: [B*T, 6*dim]
# Reshape to [B, T, 6, dim] for transformer blocks
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)) # [B*T, 6, dim]
timestep_proj = timestep_proj.view(batch_size, post_patch_num_frames, 6, self.hidden_size) # [B, T, 6, dim]
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
# temb is [B*T, dim], reshape to [B, T, 1, dim]
temb = temb.view(batch_size, post_patch_num_frames, -1).unsqueeze(2) # [B, T, 1, dim]
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
@@ -844,6 +844,15 @@ 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 cls_name.startswith("CausalWan")
or getattr(fastvideo_args.pipeline_config, "prefix", "") == "WanGame"
or cls_name.startswith("WanLingBot")
or cls_name == "WanLingBotTransformer3DModel"
or getattr(fastvideo_args.pipeline_config, "prefix", "") == "WanLingBot"
or cls_name.startswith("CausalWanGameActionTransformer3DModel")
)
model = maybe_load_fsdp_model(
model_cls=model_cls,
+25 -2
View File
@@ -290,8 +290,31 @@ def load_model_from_full_model_state_dict(
"""
meta_sd = model.state_dict()
sharded_sd = {}
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
full_sd_iterator, param_names_mapping) # type: ignore
if param_names_mapping is None:
custom_param_sd = dict(full_sd_iterator)
reverse_param_names_mapping = {
name: (name, None, None)
for name in custom_param_sd
}
else:
def _mapping_with_passthrough(
source_param_name: str,
) -> tuple[str, Any, Any]:
target_param_name, merge_index, num_params_to_merge = (
param_names_mapping(source_param_name)
)
# Custom override safetensors may already use FastVideo-internal
# names. In that case, keep the source name instead of remapping it
# a second time through the HF -> custom regex rules.
if (target_param_name not in meta_sd
and source_param_name in meta_sd):
return source_param_name, None, None
return target_param_name, merge_index, num_params_to_merge
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
full_sd_iterator,
_mapping_with_passthrough,
)
for target_param_name, full_tensor in custom_param_sd.items():
meta_sharded_param = meta_sd.get(target_param_name)
if meta_sharded_param is None:
+10 -1
View File
@@ -46,6 +46,12 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"CausalWanGameTransformer3DModel":
("dits", "wangame", "CausalWanGameActionTransformer3DModel"),
"CausalWanGameActionTransformer3DModel":
("dits", "wangame", "CausalWanGameActionTransformer3DModel"),
"WanGameActionTransformer3DModel":
("dits", "wangame", "WanGameActionTransformer3DModel"),
"MatrixGameWanModel": ("dits", "matrixgame", "MatrixGameWanModel"),
"CausalMatrixGameWanModel": ("dits", "matrixgame", "CausalMatrixGameWanModel"),
}
@@ -93,6 +99,9 @@ _SCHEDULERS = {
"FlowMatchEulerDiscreteScheduler":
("schedulers", "scheduling_flow_match_euler_discrete",
"FlowMatchEulerDiscreteScheduler"),
"DiffusionForcingScheduler":
("schedulers", "scheduling_diffusion_forcing",
"DiffusionForcingScheduler"),
"UniPCMultistepScheduler":
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
"FlowUniPCMultistepScheduler":
@@ -451,4 +460,4 @@ ModelRegistry = _ModelRegistry({
)
for model_arch, (component_name, mod_relname,
cls_name) in _FAST_VIDEO_MODELS.items()
})
})
@@ -0,0 +1,205 @@
# SPDX-License-Identifier: Apache-2.0
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput
import torch
from fastvideo.models.schedulers.base import BaseScheduler
class DiffusionForcingSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
class DiffusionForcingScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
config_name = "scheduler_config.json"
order = 1
@register_to_config
def __init__(
self,
num_inference_steps: int = 100,
num_train_timesteps: int = 1000,
shift: float = 5.0,
sigma_max: float = 1.0,
sigma_min: float = 0.0,
extra_one_step: bool = True,
training: bool = False,
):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
self.sigma_max = sigma_max
self.sigma_min = sigma_min
self.extra_one_step = extra_one_step
self.set_timesteps(num_inference_steps, training=training)
def sigma_from_timestep(self, timestep: torch.Tensor) -> torch.Tensor:
if not torch.is_tensor(timestep):
timestep = torch.as_tensor(timestep, dtype=torch.float32)
timestep = self._flatten_timestep(timestep)
device = timestep.device
self.sigmas = self.sigmas.to(device)
timestep_id = self._lookup_timestep_indices(
timestep=timestep,
device=device,
)
return self.sigmas[timestep_id]
def _flatten_timestep(self, timestep: torch.Tensor) -> torch.Tensor:
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
elif timestep.ndim == 0:
timestep = timestep.unsqueeze(0)
elif timestep.ndim != 1:
raise ValueError("timestep must be scalar, [B], [B, T], "
"or [B*T]")
return timestep.to(torch.float32)
def _lookup_timestep_indices(
self,
*,
timestep: torch.Tensor,
device: torch.device,
) -> torch.Tensor:
self.timesteps = self.timesteps.to(device)
return torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(),
dim=1,
)
def set_timesteps(
self,
num_inference_steps: int = 100,
denoising_strength: float = 1.0,
training: bool = False,
return_dict: bool = False,
**kwargs,
):
sigma_start = self.sigma_min + (
self.sigma_max - self.sigma_min
) * denoising_strength
if self.extra_one_step:
sigmas = torch.linspace(
sigma_start, self.sigma_min, num_inference_steps + 1
)[:-1]
else:
sigmas = torch.linspace(
sigma_start, self.sigma_min, num_inference_steps
)
if self.shift != 1.0:
sigmas = self.shift * sigmas / (1 + (self.shift - 1) * sigmas)
self.sigmas = sigmas.to(torch.float32)
self.timesteps = (
self.sigmas * float(self.num_train_timesteps)
).to(torch.float32)
if training:
x = self.timesteps
y = torch.exp(
-2 * ((x - num_inference_steps / 2) / num_inference_steps) ** 2
)
y_shifted = y - y.min()
self.linear_timesteps_weights = (
y_shifted * (num_inference_steps / y_shifted.sum())
)
def step(
self,
model_output: torch.FloatTensor,
timestep: torch.FloatTensor,
sample: torch.FloatTensor,
to_final: bool = False,
return_dict: bool = False,
**kwargs,
):
timestep = self._flatten_timestep(timestep).to(
model_output.device,
dtype=torch.float32,
)
self.sigmas = self.sigmas.to(model_output.device)
timestep_id = self._lookup_timestep_indices(
timestep=timestep,
device=model_output.device,
)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
if to_final:
sigma_next = torch.zeros_like(sigma)
else:
sigma_next = torch.zeros_like(sigma)
valid = timestep_id + 1 < len(self.timesteps)
if valid.any():
sigma_next[valid] = self.sigmas[
timestep_id[valid] + 1
].reshape(-1, 1, 1, 1)
prev_sample = sample + model_output * (sigma_next - sigma)
if isinstance(prev_sample, (torch.Tensor, float)) and not return_dict:
return (prev_sample,)
return DiffusionForcingSchedulerOutput(prev_sample=prev_sample)
@staticmethod
def calculate_alpha_beta_high(sigma, sigma_bound):
alpha = (1 - sigma) / (1 - sigma_bound)
beta = torch.sqrt(sigma**2 - (alpha * sigma_bound) ** 2)
return alpha, beta
def add_noise(self, original_samples, noise, timestep):
timestep = self._flatten_timestep(timestep).to(noise.device)
self.sigmas = self.sigmas.to(noise.device)
timestep_id = self._lookup_timestep_indices(
timestep=timestep,
device=noise.device,
)
sigma = self.sigmas[timestep_id].reshape(
-1, 1, 1, 1,
)
sample = (1 - sigma) * original_samples + sigma * noise
return sample.type_as(noise)
def add_noise_high(
self, original_samples, noise, timestep, boundary_timestep
):
timestep = self._flatten_timestep(timestep).to(noise.device)
boundary_timestep = self._flatten_timestep(boundary_timestep).to(
noise.device
)
self.sigmas = self.sigmas.to(noise.device)
timestep_id = self._lookup_timestep_indices(
timestep=timestep,
device=noise.device,
)
boundary_timestep_id = self._lookup_timestep_indices(
timestep=boundary_timestep,
device=noise.device,
)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
sigma_boundary = self.sigmas[boundary_timestep_id].reshape(
-1, 1, 1, 1
)
alpha, beta = self.calculate_alpha_beta_high(sigma, sigma_boundary)
sample = alpha * original_samples + beta * noise
return sample.type_as(noise)
def training_target(self, sample, noise, timestep):
return noise - sample
def training_weight(self, timestep):
timestep = self._flatten_timestep(timestep)
device = timestep.device
self.linear_timesteps_weights = self.linear_timesteps_weights.to(device)
timestep_id = self._lookup_timestep_indices(
timestep=timestep,
device=device,
)
return self.linear_timesteps_weights[timestep_id]
def scale_model_input(
self, sample: torch.Tensor, timestep: int | None = None
) -> torch.Tensor:
return sample
def set_shift(self, shift: float) -> None:
self.shift = shift
EntryClass = DiffusionForcingScheduler
@@ -0,0 +1,3 @@
from fastvideo.pipelines.basic.wan.wangame_causal_dmd_pipeline import (
WangameCausalOdeDMDPipeline as WangameCausalOdeDMDPipeline,
WangameCausalSdeDMDPipeline as WangameCausalSdeDMDPipeline, )
@@ -0,0 +1,94 @@
# SPDX-License-Identifier: Apache-2.0
"""Wangame causal DMD pipeline implementations."""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (
ConditioningStage, DecodingStage, MatrixGameCausalDenoisingStage,
MatrixGameCausalOdeDenoisingStage, MatrixGameImageEncodingStage,
InputValidationStage, LatentPreparationStage, TextEncodingStage,
TimestepPreparationStage)
from fastvideo.pipelines.stages.image_encoding import (
MatrixGameImageVAEEncodingStage)
logger = init_logger(__name__)
class _WangameCausalDMDPipelineBase(LoRAPipeline, ComposedPipelineBase):
requires_timestep_preparation: bool
requires_dmd_denoising_steps: bool
denoising_stage_cls: type[MatrixGameCausalDenoisingStage]
_required_config_modules = [
"vae", "transformer", "scheduler", "image_encoder", "image_processor"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
del fastvideo_args
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
if (self.get_module("text_encoder", None) is not None
and self.get_module("tokenizer", None) is not None):
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
if (self.get_module("image_encoder", None) is not None
and self.get_module("image_processor", None) is not None):
self.add_stage(
stage_name="image_encoding_stage",
stage=MatrixGameImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
if self.requires_timestep_preparation:
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", None)))
self.add_stage(
stage_name="image_latent_preparation_stage",
stage=MatrixGameImageVAEEncodingStage(vae=self.get_module("vae")))
denoising_stage = self.denoising_stage_cls(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
pipeline=self,
vae=self.get_module("vae"),
)
self.add_stage(stage_name="denoising_stage", stage=denoising_stage)
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
logger.info("%s initialized with action support",
self.__class__.__name__)
class WangameCausalOdeDMDPipeline(_WangameCausalDMDPipelineBase):
requires_timestep_preparation = True
requires_dmd_denoising_steps = False
denoising_stage_cls = MatrixGameCausalOdeDenoisingStage
class WangameCausalSdeDMDPipeline(_WangameCausalDMDPipelineBase):
requires_timestep_preparation = False
requires_dmd_denoising_steps = True
denoising_stage_cls = MatrixGameCausalDenoisingStage
EntryClass = [WangameCausalOdeDMDPipeline, WangameCausalSdeDMDPipeline]
@@ -160,6 +160,10 @@ class ForwardBatch:
# Timesteps
timesteps: torch.Tensor | None = None
# Optional explicit denoising-loop timesteps (sampler-specific).
# When set, some samplers iterate this list instead of `timesteps`
# produced by the timestep preparation stage.
sampling_timesteps: torch.Tensor | None = None
timestep: torch.Tensor | float | int | None = None
step_index: int | None = None
boundary_ratio: float | None = None
+2 -1
View File
@@ -34,7 +34,7 @@ from fastvideo.pipelines.stages.ltx2_latent_preparation import (
LTX2LatentPreparationStage)
from fastvideo.pipelines.stages.ltx2_text_encoding import LTX2TextEncodingStage
from fastvideo.pipelines.stages.matrixgame_denoising import (
MatrixGameCausalDenoisingStage)
MatrixGameCausalDenoisingStage, MatrixGameCausalOdeDenoisingStage)
from fastvideo.pipelines.stages.hyworld_denoising import HYWorldDenoisingStage
from fastvideo.pipelines.stages.gamecraft_denoising import GameCraftDenoisingStage
from fastvideo.pipelines.stages.text_encoding import (Cosmos25TextEncodingStage,
@@ -66,6 +66,7 @@ __all__ = [
"CausalDMDDenosingStage",
"CausalDenoisingStage",
"MatrixGameCausalDenoisingStage",
"MatrixGameCausalOdeDenoisingStage",
"HYWorldDenoisingStage",
"GameCraftDenoisingStage",
"CosmosDenoisingStage",
+57 -9
View File
@@ -1237,6 +1237,37 @@ class DmdDenoisingStage(DenoisingStage):
},
)
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,
{
"mouse_cond": batch.mouse_cond,
"keyboard_cond": batch.keyboard_cond,
"c2ws_plucker_emb": batch.c2ws_plucker_emb,
},
)
# Get latents and embeddings
assert batch.latents is not None, "latents must be provided"
latents = batch.latents
@@ -1245,14 +1276,29 @@ class DmdDenoisingStage(DenoisingStage):
prompt_embeds = batch.prompt_embeds
assert not torch.isnan(
prompt_embeds[0]).any(), "prompt_embeds contains nan"
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
loop_timesteps = batch.sampling_timesteps
if loop_timesteps is None:
legacy = getattr(fastvideo_args.pipeline_config,
"dmd_denoising_steps", None)
if legacy is not None:
loop_timesteps = torch.tensor(legacy, dtype=torch.long)
else:
loop_timesteps = batch.timesteps
if loop_timesteps is None:
raise ValueError(
"SDE sampling requires `batch.sampling_timesteps` "
"(preferred) or `pipeline_config.dmd_denoising_steps`.")
if not isinstance(loop_timesteps, torch.Tensor):
loop_timesteps = torch.tensor(loop_timesteps, dtype=torch.long)
if loop_timesteps.ndim != 1:
raise ValueError("Expected 1D `sampling_timesteps`, got shape "
f"{tuple(loop_timesteps.shape)}")
loop_timesteps = loop_timesteps.to(get_local_torch_device())
# Run denoising loop
with self.progress_bar(total=len(timesteps)) as progress_bar:
for i, t in enumerate(timesteps):
with self.progress_bar(total=len(loop_timesteps)) as progress_bar:
for i, t in enumerate(loop_timesteps):
# Skip if interrupted
if hasattr(self, 'interrupt') and self.interrupt:
continue
@@ -1326,6 +1372,8 @@ class DmdDenoisingStage(DenoisingStage):
guidance=guidance_expand,
**image_kwargs,
**pos_cond_kwargs,
**action_kwargs,
**camera_action_kwargs,
).permute(0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
@@ -1335,8 +1383,8 @@ class DmdDenoisingStage(DenoisingStage):
scheduler=self.scheduler).unflatten(
0, pred_noise.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
if i < len(loop_timesteps) - 1:
next_timestep = loop_timesteps[i + 1] * torch.ones(
[1], dtype=torch.long, device=pred_video.device)
noise = torch.randn(video_raw_latent_shape,
dtype=pred_video.dtype,
@@ -1349,7 +1397,7 @@ class DmdDenoisingStage(DenoisingStage):
latents = pred_video
# Update progress bar
if i == len(timesteps) - 1 or (
if i == len(loop_timesteps) - 1 or (
(i + 1) > num_warmup_steps and
(i + 1) % self.scheduler.order == 0
and progress_bar is not None):
@@ -1,6 +1,8 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
import hashlib
import os
from typing import Any
import torch # type: ignore
@@ -26,6 +28,42 @@ except ImportError:
logger = init_logger(__name__)
def _shape_repr(value: Any) -> str:
if isinstance(value, torch.Tensor):
return "x".join(str(int(dim)) for dim in value.shape)
if value is None:
return "None"
return type(value).__name__
def _preview_tensor(value: Any, *, limit: int = 8) -> list[Any] | None:
if not isinstance(value, torch.Tensor):
return None
flat = value.detach().flatten().cpu()
if flat.numel() == 0:
return []
flat = flat[:limit]
if torch.is_floating_point(flat):
return [float(x.item()) for x in flat]
return [int(x.item()) for x in flat]
def _should_debug_timesteps(prompt: object) -> bool:
if not os.environ.get("FASTVIDEO_DEBUG_TIMESTEPS"):
return False
if isinstance(prompt, str):
return prompt.startswith("00 Val-00:")
if isinstance(prompt, list) and prompt:
first = prompt[0]
return isinstance(first, str) and first.startswith("00 Val-00:")
return False
def _tensor_md5(tensor: torch.Tensor) -> str:
array = tensor.detach().cpu().to(torch.int64).contiguous().numpy()
return hashlib.md5(array.tobytes()).hexdigest()
@dataclass
class BlockProcessingContext:
"""Dataclass contains for block processing."""
@@ -54,6 +92,9 @@ class BlockProcessingContext:
image_kwargs: dict[str, Any]
pos_cond_kwargs: dict[str, Any]
viewmats_full: torch.Tensor | None = None
intrinsics_full: torch.Tensor | None = None
action_full: torch.Tensor | None = None
def get_kv_cache(self, timestep_val: float) -> list[dict[Any, Any]]:
if self.boundary_timestep is not None:
@@ -97,10 +138,12 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
-1)
except Exception:
self.local_attn_size = -1
try:
self.local_attn_size = getattr(self.transformer.model,
"local_attn_size", -1)
except Exception:
self.local_attn_size = -1
assert self.local_attn_size != -1, (
f"local_attn_size must be set for Matrix-Game causal inference, "
f"got {self.local_attn_size}. Check MatrixGameWanVideoArchConfig.")
assert self.num_frame_per_block > 0, (
f"num_frame_per_block must be positive, got {self.num_frame_per_block}"
)
@@ -115,6 +158,7 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
self._streaming_initialized: bool = False
self._streaming_ctx: BlockProcessingContext | None = None
self._logged_forward_setup = False
def forward(
self,
@@ -126,7 +170,10 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
) and not fastvideo_args.disable_autocast
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
patch_size = self.transformer.patch_size
if hasattr(self.transformer, "patch_size"):
patch_size = self.transformer.patch_size
else:
patch_size = self.transformer.config.arch_config.patch_size
patch_ratio = patch_size[-1] * patch_size[-2]
self.frame_seq_length = latent_seq_length // patch_ratio
@@ -166,6 +213,31 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
viewmats_full = None
intrinsics_full = None
action_full = None
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_list = []
intrinsics_list = []
action_list = []
for bi in range(b):
vm, ks, action = process_custom_actions(batch.keyboard_cond[bi],
batch.mouse_cond[bi])
viewmats_list.append(vm)
intrinsics_list.append(ks)
action_list.append(action)
viewmats_full = torch.stack(viewmats_list,
dim=0).to(device=latents.device,
dtype=target_dtype)
intrinsics_full = torch.stack(intrinsics_list,
dim=0).to(device=latents.device,
dtype=target_dtype)
action_full = torch.stack(action_list,
dim=0).to(device=latents.device,
dtype=target_dtype)
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
@@ -200,6 +272,8 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
if boundary_timestep is not None:
block_sizes[0] = 1
self._logged_forward_setup = True
# NOTE: MatrixGame does NOT process the first frame separately.
# The first frame information is already encoded in batch.image_latent (cond_concat)
# and will be used by the model via channel concatenation: torch.cat([x, cond_concat], dim=1)
@@ -225,6 +299,9 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
"context_noise", 0),
image_kwargs=image_kwargs,
pos_cond_kwargs=pos_cond_kwargs,
viewmats_full=viewmats_full,
intrinsics_full=intrinsics_full,
action_full=action_full,
)
context_noise = getattr(fastvideo_args.pipeline_config, "context_noise",
@@ -240,6 +317,8 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
action_kwargs = self._prepare_action_kwargs(
batch, start_index, current_num_frames)
camera_action_kwargs = self._prepare_camera_action_kwargs(
ctx, start_index, current_num_frames)
current_latents = self._process_single_block(
current_latents=current_latents,
@@ -249,6 +328,7 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
timesteps=timesteps,
ctx=ctx,
action_kwargs=action_kwargs,
camera_action_kwargs=camera_action_kwargs,
progress_bar=progress_bar,
)
@@ -263,6 +343,7 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
current_num_frames=current_num_frames,
ctx=ctx,
action_kwargs=action_kwargs,
camera_action_kwargs=camera_action_kwargs,
context_noise=context_noise,
)
@@ -324,9 +405,9 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
torch.zeros((), dtype=torch.long, device=device),
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
torch.zeros((), dtype=torch.long, device=device),
})
return kv_cache
@@ -362,9 +443,9 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
torch.zeros((), dtype=torch.long, device=device),
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
torch.zeros((), dtype=torch.long, device=device),
})
kv_cache_mouse.append({
"k":
@@ -382,9 +463,9 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
torch.zeros((), dtype=torch.long, device=device),
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
torch.zeros((), dtype=torch.long, device=device),
})
return kv_cache_mouse, kv_cache_keyboard
@@ -418,6 +499,19 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
})
return crossattn_cache
def _prepare_camera_action_kwargs(
self, ctx: BlockProcessingContext, start_index: int,
current_num_frames: int) -> dict[str, Any]:
if ctx.action_full is None or ctx.viewmats_full is None or ctx.intrinsics_full is None:
return {}
end_index = start_index + current_num_frames
result = {
"viewmats": ctx.viewmats_full[:, start_index:end_index],
"Ks": ctx.intrinsics_full[:, start_index:end_index],
"action": ctx.action_full[:, start_index:end_index],
}
return result
def _process_single_block(
self,
current_latents: torch.Tensor,
@@ -427,6 +521,7 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
timesteps: torch.Tensor,
ctx: BlockProcessingContext,
action_kwargs: dict[str, Any],
camera_action_kwargs: dict[str, Any],
noise_generator: Callable[[tuple, torch.dtype, int], torch.Tensor]
| None = None,
progress_bar: Any | None = None,
@@ -445,7 +540,16 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
independent_first_frame = getattr(self.transformer,
'independent_first_frame', False)
if batch.image_latent is not None and independent_first_frame and start_index == 0:
if batch.image_latent is not None and not independent_first_frame:
image_latent_chunk = batch.image_latent[:, :, start_index:
start_index +
current_num_frames, :, :]
latent_model_input = torch.cat([
latent_model_input,
image_latent_chunk.to(ctx.target_dtype)
],
dim=1)
elif batch.image_latent is not None and independent_first_frame and start_index == 0:
latent_model_input = torch.cat([
latent_model_input,
batch.image_latent.to(ctx.target_dtype)
@@ -495,6 +599,7 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
"crossattn_cache": ctx.crossattn_cache,
"current_start": start_index * self.frame_seq_length,
"start_frame": start_index,
"is_cache": False,
}
if self.use_action_module and current_model == self.transformer:
@@ -510,6 +615,7 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
latent_model_input,
prompt_embeds,
t_expanded_noise,
**camera_action_kwargs,
**ctx.image_kwargs,
**ctx.pos_cond_kwargs,
**model_kwargs,
@@ -582,6 +688,7 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
current_num_frames: int,
ctx: BlockProcessingContext,
action_kwargs: dict[str, Any],
camera_action_kwargs: dict[str, Any],
context_noise: float,
) -> None:
prompt_embeds = batch.prompt_embeds
@@ -592,6 +699,17 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
device=latents_device,
dtype=torch.long) * int(context_noise)
context_bcthw = current_latents.to(ctx.target_dtype)
context_input = context_bcthw
independent_first_frame = getattr(self.transformer,
"independent_first_frame", False)
if batch.image_latent is not None and not independent_first_frame:
image_context_chunk = batch.image_latent[:, :,
start_index:start_index +
current_num_frames, :, :]
context_input = torch.cat(
[context_input,
image_context_chunk.to(ctx.target_dtype)],
dim=1)
with torch.autocast(device_type="cuda",
dtype=ctx.target_dtype,
@@ -605,6 +723,7 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
"crossattn_cache": ctx.crossattn_cache,
"current_start": start_index * self.frame_seq_length,
"start_frame": start_index,
"is_cache": True,
}
if self.use_action_module:
@@ -617,26 +736,409 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
context_model_kwargs.update(action_kwargs)
if ctx.boundary_timestep is not None and self.transformer_2 is not None:
self.transformer_2(
context_bcthw,
cache_update_ret_2 = self.transformer_2(
context_input,
prompt_embeds,
t_context,
kv_cache=ctx.kv_cache2,
crossattn_cache=ctx.crossattn_cache,
current_start=start_index * self.frame_seq_length,
start_frame=start_index,
is_cache=True,
**camera_action_kwargs,
**ctx.image_kwargs,
**ctx.pos_cond_kwargs,
)
if isinstance(cache_update_ret_2,
list) and len(cache_update_ret_2) > 0:
ctx.kv_cache2 = cache_update_ret_2
self.transformer(
context_bcthw,
cache_update_ret = self.transformer(
context_input,
prompt_embeds,
t_context,
**camera_action_kwargs,
**ctx.image_kwargs,
**ctx.pos_cond_kwargs,
**context_model_kwargs,
)
if isinstance(cache_update_ret, list) and len(cache_update_ret) > 0:
ctx.kv_cache1 = cache_update_ret
class MatrixGameCausalOdeDenoisingStage(MatrixGameCausalDenoisingStage):
"""Causal ODE denoising for WanGame/MatrixGame.
This is the deterministic counterpart of `MatrixGameCausalDenoisingStage`.
It performs block-by-block causal rollout, but uses the scheduler's ODE-style
`step()` update (no re-noising between steps).
"""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
timesteps = batch.timesteps
if timesteps is None:
raise ValueError(
"MatrixGameCausalOdeDenoisingStage requires batch.timesteps. "
"Make sure TimestepPreparationStage runs before this stage.")
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
if hasattr(self.transformer, "patch_size"):
patch_size = self.transformer.patch_size
else:
patch_size = self.transformer.config.arch_config.patch_size
patch_ratio = patch_size[-1] * patch_size[-2]
self.frame_seq_length = latent_seq_length // patch_ratio
timesteps = timesteps.to(get_local_torch_device())
if _should_debug_timesteps(batch.prompt):
logger.info(
"DEBUG_TIMESTEPS stage=ode_denoising rank=%s prompt=%r "
"len=%s md5=%s head=%s tail=%s",
os.environ.get("RANK", "?"),
batch.prompt,
int(timesteps.numel()),
_tensor_md5(timesteps),
timesteps[:10].detach().cpu().tolist(),
timesteps[-10:].detach().cpu().tolist(),
)
boundary_ratio = getattr(fastvideo_args.pipeline_config.dit_config,
"boundary_ratio", None)
if boundary_ratio is not None:
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
else:
boundary_timestep = None
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert torch.isnan(image_embeds[0]).sum() == 0
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
# directly set the kwarg.
image_kwargs = {"encoder_hidden_states_image": image_embeds}
pos_cond_kwargs: dict[str, Any] = {}
assert batch.latents is not None, "latents must be provided"
latents = batch.latents
b, c, t, h, w = latents.shape
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
viewmats_full = None
intrinsics_full = None
action_full = None
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_list = []
intrinsics_list = []
action_list = []
for bi in range(b):
vm, ks, action = process_custom_actions(batch.keyboard_cond[bi],
batch.mouse_cond[bi])
viewmats_list.append(vm)
intrinsics_list.append(ks)
action_list.append(action)
viewmats_full = torch.stack(viewmats_list,
dim=0).to(device=latents.device,
dtype=target_dtype)
intrinsics_full = torch.stack(intrinsics_list,
dim=0).to(device=latents.device,
dtype=target_dtype)
action_full = torch.stack(action_list,
dim=0).to(device=latents.device,
dtype=target_dtype)
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache2 = None
if boundary_timestep is not None:
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache_mouse = None
kv_cache_keyboard = None
if self.use_action_module:
kv_cache_mouse, kv_cache_keyboard = self._initialize_action_kv_cache(
batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=257, # 1 CLS + 256 patch tokens
dtype=target_dtype,
device=latents.device)
if t % self.num_frame_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frame_per_block for causal denoising"
)
num_blocks = t // self.num_frame_per_block
block_sizes = [self.num_frame_per_block] * num_blocks
start_index = 0
if boundary_timestep is not None:
block_sizes[0] = 1
ctx = BlockProcessingContext(
batch=batch,
block_idx=0,
start_index=0,
kv_cache1=kv_cache1,
kv_cache2=kv_cache2,
kv_cache_mouse=kv_cache_mouse,
kv_cache_keyboard=kv_cache_keyboard,
crossattn_cache=crossattn_cache,
timesteps=timesteps,
block_sizes=block_sizes,
noise_pool=None,
fastvideo_args=fastvideo_args,
target_dtype=target_dtype,
autocast_enabled=autocast_enabled,
boundary_timestep=boundary_timestep,
high_noise_timesteps=None,
context_noise=getattr(fastvideo_args.pipeline_config,
"context_noise", 0),
image_kwargs=image_kwargs,
pos_cond_kwargs=pos_cond_kwargs,
viewmats_full=viewmats_full,
intrinsics_full=intrinsics_full,
action_full=action_full,
)
context_noise = getattr(fastvideo_args.pipeline_config, "context_noise",
0)
with self.progress_bar(total=len(block_sizes) *
len(timesteps)) as progress_bar:
for block_idx, current_num_frames in enumerate(block_sizes):
ctx.block_idx = block_idx
ctx.start_index = start_index
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
# The scheduler maintains an internal `step_index` (and potentially
# additional multistep state, e.g. UniPC). Since causal streaming runs
# a full denoising trajectory *per block*, reset that state before
# each block rollout.
self._reset_scheduler_state_for_new_rollout()
action_kwargs = self._prepare_action_kwargs(
batch, start_index, current_num_frames)
camera_action_kwargs = self._prepare_camera_action_kwargs(
ctx, start_index, current_num_frames)
current_latents = self._process_single_block_ode(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
timesteps=timesteps,
ctx=ctx,
action_kwargs=action_kwargs,
camera_action_kwargs=camera_action_kwargs,
progress_bar=progress_bar,
)
latents[:, :, start_index:start_index +
current_num_frames, :, :] = current_latents
# Update KV caches with clean context
self._update_context_cache(
current_latents=current_latents,
batch=batch,
start_index=start_index,
current_num_frames=current_num_frames,
ctx=ctx,
action_kwargs=action_kwargs,
camera_action_kwargs=camera_action_kwargs,
context_noise=context_noise,
)
start_index += current_num_frames
if boundary_timestep is not None:
num_frames_to_remove = self.num_frame_per_block - 1
if num_frames_to_remove > 0:
latents = latents[:, :, :-num_frames_to_remove, :, :]
batch.latents = latents
return batch
def _reset_scheduler_state_for_new_rollout(self) -> None:
scheduler = self.scheduler
# Common diffusers-like state.
if hasattr(scheduler, "_step_index"):
scheduler._step_index = None # type: ignore[attr-defined]
if hasattr(scheduler, "_begin_index"):
scheduler._begin_index = None # type: ignore[attr-defined]
# UniPC multistep state (FlowUniPCMultistepScheduler) needs additional reset
# between independent trajectories.
if hasattr(scheduler, "model_outputs") and hasattr(scheduler, "config"):
try:
solver_order = int(
getattr(scheduler.config, "solver_order", 0) or 0)
except Exception:
solver_order = 0
if solver_order > 0:
scheduler.model_outputs = [
None
] * solver_order # type: ignore[attr-defined]
if hasattr(scheduler, "timestep_list") and hasattr(scheduler, "config"):
try:
solver_order = int(
getattr(scheduler.config, "solver_order", 0) or 0)
except Exception:
solver_order = 0
if solver_order > 0:
scheduler.timestep_list = [
None
] * solver_order # type: ignore[attr-defined]
if hasattr(scheduler, "lower_order_nums"):
scheduler.lower_order_nums = 0 # type: ignore[attr-defined]
if hasattr(scheduler, "last_sample"):
scheduler.last_sample = None # type: ignore[attr-defined]
def _process_single_block_ode(
self,
*,
current_latents: torch.Tensor,
batch: ForwardBatch,
start_index: int,
current_num_frames: int,
timesteps: torch.Tensor,
ctx: BlockProcessingContext,
action_kwargs: dict[str, Any],
camera_action_kwargs: dict[str, Any],
progress_bar: Any | None = None,
) -> torch.Tensor:
prompt_embeds = batch.prompt_embeds
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step,
{
"generator": batch.generator,
"eta": batch.eta,
},
)
for i, t_cur in enumerate(timesteps):
if ctx.boundary_timestep is not None and t_cur < ctx.boundary_timestep:
current_model = self.transformer_2 if self.transformer_2 is not None else self.transformer
else:
current_model = self.transformer
latent_model_input = current_latents.to(ctx.target_dtype)
independent_first_frame = getattr(self.transformer,
"independent_first_frame", False)
if batch.image_latent is not None and not independent_first_frame:
image_latent_chunk = batch.image_latent[:, :, start_index:
start_index +
current_num_frames, :, :]
latent_model_input = torch.cat([
latent_model_input,
image_latent_chunk.to(ctx.target_dtype)
],
dim=1)
elif (batch.image_latent is not None and independent_first_frame
and start_index == 0):
latent_model_input = torch.cat([
latent_model_input,
batch.image_latent.to(ctx.target_dtype)
],
dim=2)
latent_model_input = self.scheduler.scale_model_input(
latent_model_input, t_cur)
# Build attention metadata if VSA is available
if vsa_available and self.attn_backend == VideoSparseAttentionBackend:
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls(
)
h, w = current_latents.shape[-2:]
attn_metadata = self.attn_metadata_builder.build(
current_timestep=i,
raw_latent_shape=(current_num_frames, h, w),
patch_size=ctx.fastvideo_args.pipeline_config.
dit_config.patch_size,
VSA_sparsity=ctx.fastvideo_args.VSA_sparsity,
device=get_local_torch_device(),
)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
with torch.autocast(device_type="cuda",
dtype=ctx.target_dtype,
enabled=ctx.autocast_enabled), \
set_forward_context(current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded = t_cur * torch.ones(
(latent_model_input.shape[0], current_num_frames),
device=latent_model_input.device,
dtype=t_cur.dtype)
model_kwargs = {
"kv_cache": ctx.get_kv_cache(t_cur),
"crossattn_cache": ctx.crossattn_cache,
"current_start": start_index * self.frame_seq_length,
"start_frame": start_index,
"is_cache": False,
}
if self.use_action_module and current_model == self.transformer:
model_kwargs.update({
"kv_cache_mouse":
ctx.kv_cache_mouse,
"kv_cache_keyboard":
ctx.kv_cache_keyboard,
})
model_kwargs.update(action_kwargs)
noise_pred = current_model(
latent_model_input,
prompt_embeds,
t_expanded,
**camera_action_kwargs,
**ctx.image_kwargs,
**ctx.pos_cond_kwargs,
**model_kwargs,
)
current_latents = self.scheduler.step(
noise_pred,
t_cur,
current_latents,
**extra_step_kwargs,
return_dict=False,
)[0]
if progress_bar is not None:
progress_bar.update()
return current_latents
def streaming_reset(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
@@ -645,7 +1147,10 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
) and not fastvideo_args.disable_autocast
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
patch_size = self.transformer.patch_size
if hasattr(self.transformer, "patch_size"):
patch_size = self.transformer.patch_size
else:
patch_size = self.transformer.config.arch_config.patch_size
patch_ratio = patch_size[-1] * patch_size[-2]
self.frame_seq_length = latent_seq_length // patch_ratio
@@ -821,6 +1326,7 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
timesteps=ctx.timesteps,
ctx=ctx,
action_kwargs=action_kwargs,
camera_action_kwargs={},
noise_generator=streaming_noise_generator,
)
@@ -835,6 +1341,7 @@ class MatrixGameCausalDenoisingStage(DenoisingStage):
current_num_frames=current_num_frames,
ctx=ctx,
action_kwargs=action_kwargs,
camera_action_kwargs={},
context_noise=ctx.context_noise,
)
+3
View File
@@ -173,3 +173,6 @@ class CallbackDict:
fn(*args, **kwargs)
return _dispatch
def get_callback(self, name: str) -> Callback | None:
return self._callbacks.get(name)
File diff suppressed because it is too large Load Diff
+68
View File
@@ -33,6 +33,10 @@ def run_training_from_config(
config_path: str,
*,
dry_run: bool = False,
resume_from_checkpoint: str | None = None,
override_output_dir: str | None = None,
best_checkpoint_start_step: int | None = None,
best_checkpoint_top_k: int | None = None,
overrides: list[str] | None = None,
) -> None:
"""YAML-only training entrypoint (schema v2)."""
@@ -54,6 +58,15 @@ def run_training_from_config(
cfg = load_run_config(config_path, overrides=overrides)
tc = cfg.training
if resume_from_checkpoint is not None:
tc.checkpoint.resume_from_checkpoint = str(resume_from_checkpoint)
if override_output_dir is not None:
tc.checkpoint.output_dir = str(override_output_dir)
if best_checkpoint_start_step is not None:
tc.checkpoint.best_checkpoint_start_step = int(best_checkpoint_start_step)
if best_checkpoint_top_k is not None:
tc.checkpoint.best_checkpoint_top_k = max(1, int(best_checkpoint_top_k))
# Auto-set attention backend for VSA when sparsity is configured.
if tc.vsa_sparsity > 0.0:
os.environ.setdefault(
@@ -97,6 +110,7 @@ def run_training_from_config(
output_dir=tc.checkpoint.output_dir,
config=ckpt_config,
callbacks=trainer.callbacks,
tracker=trainer.tracker,
raw_config=cfg.raw,
)
@@ -115,6 +129,18 @@ def main(
) -> None:
config_path = str(args.config)
dry_run = bool(args.dry_run)
resume_from_checkpoint = getattr(
args, "resume_from_checkpoint", None
)
override_output_dir = getattr(
args, "override_output_dir", None
)
best_checkpoint_start_step = getattr(
args, "best_checkpoint_start_step", None
)
best_checkpoint_top_k = getattr(
args, "best_checkpoint_top_k", None
)
logger.info(
"Starting training from config=%s",
config_path,
@@ -122,6 +148,10 @@ def main(
run_training_from_config(
config_path,
dry_run=dry_run,
resume_from_checkpoint=resume_from_checkpoint,
override_output_dir=override_output_dir,
best_checkpoint_start_step=best_checkpoint_start_step,
best_checkpoint_top_k=best_checkpoint_top_k,
overrides=overrides,
)
logger.info("Training completed")
@@ -142,5 +172,43 @@ if __name__ == "__main__":
help=("Parse config and build runtime, "
"but do not start training."),
)
parser.add_argument(
"--resume-from-checkpoint",
type=str,
default=None,
help=(
"Path to a checkpoint directory "
"(checkpoint-<step>), its 'dcp/' subdir, "
"or an output_dir containing checkpoints "
"(auto-picks latest)."
),
)
parser.add_argument(
"--override-output-dir",
type=str,
default=None,
help=(
"Override training.output_dir from YAML "
"(useful for repeated runs)."
),
)
parser.add_argument(
"--best-checkpoint-start-step",
type=int,
default=None,
help=(
"Override training.checkpoint.best_checkpoint_start_step "
"(0 disables best-checkpoint saving)."
),
)
parser.add_argument(
"--best-checkpoint-top-k",
type=int,
default=None,
help=(
"Override training.checkpoint.best_checkpoint_top_k "
"(minimum 1)."
),
)
args, unknown = parser.parse_known_args(argv[1:])
main(args, overrides=unknown if unknown else None)
+4
View File
@@ -8,6 +8,7 @@ __all__ = [
"FineTuneMethod",
"SelfForcingMethod",
"DiffusionForcingSFTMethod",
"TeacherForcingSFTMethod",
]
@@ -24,4 +25,7 @@ def __getattr__(name: str) -> object:
if name == "DiffusionForcingSFTMethod":
from fastvideo.train.methods.fine_tuning.dfsft import DiffusionForcingSFTMethod
return DiffusionForcingSFTMethod
if name == "TeacherForcingSFTMethod":
from fastvideo.train.methods.fine_tuning.tfsft import TeacherForcingSFTMethod
return TeacherForcingSFTMethod
raise AttributeError(name)
+4
View File
@@ -73,6 +73,10 @@ class TrainingMethod(torch.nn.Module, ABC):
def set_tracker(self, tracker: Any) -> None:
self.tracker = tracker
@property
def transformer_inference(self) -> torch.nn.Module:
return self.student.transformer
@abstractmethod
def single_train_step(
self,
@@ -126,6 +126,12 @@ class DMD2Method(TrainingMethod):
}
outputs: dict[str, Any] = dict(critic_outputs)
if training_batch.dmd_latent_vis_dict:
outputs["dmd_latent_vis_dict"] = training_batch.dmd_latent_vis_dict
if training_batch.fake_score_latent_vis_dict:
outputs["fake_score_latent_vis_dict"] = (
training_batch.fake_score_latent_vis_dict
)
outputs["_fv_backward"] = {
"update_student": update_student,
"student_ctx": student_ctx,
@@ -623,6 +629,12 @@ class DMD2Method(TrainingMethod):
grad = (faker_x0 - real_cfg_x0) / denom
grad = torch.nan_to_num(grad)
batch.dmd_latent_vis_dict.update({
"generator_pred_video": generator_pred_x0.detach(),
"real_score_pred_video": real_cfg_x0.detach(),
"faker_score_pred_video": faker_x0.detach(),
"dmd_timestep": timestep.float().detach(),
})
loss = 0.5 * F.mse_loss(
generator_pred_x0.float(),
(generator_pred_x0.float() - grad.float()).detach(),
@@ -8,6 +8,7 @@ from typing import Any, Literal, TYPE_CHECKING
import torch
import torch.distributed as dist
from fastvideo.logger import init_logger
from fastvideo.train.models.base import (
CausalModelBase,
ModelBase,
@@ -25,6 +26,8 @@ from fastvideo.models.utils import pred_noise_to_pred_video
if TYPE_CHECKING:
from fastvideo.pipelines import TrainingBatch
logger = init_logger(__name__)
def _require_bool(raw: Any, *, where: str) -> bool:
if isinstance(raw, bool):
@@ -38,6 +41,26 @@ def _require_str(raw: Any, *, where: str) -> str:
return raw
def _shape_repr(value: Any) -> str:
if isinstance(value, torch.Tensor):
return "x".join(str(int(dim)) for dim in value.shape)
if value is None:
return "None"
return type(value).__name__
def _preview_tensor(value: Any, *, limit: int = 8) -> list[Any] | None:
if not isinstance(value, torch.Tensor):
return None
flat = value.detach().flatten().cpu()
if flat.numel() == 0:
return []
flat = flat[:limit]
if torch.is_floating_point(flat):
return [float(x.item()) for x in flat]
return [int(x.item()) for x in flat]
class SelfForcingMethod(DMD2Method):
"""Self-Forcing DMD2 (distribution matching) method.
@@ -154,6 +177,38 @@ class SelfForcingMethod(DMD2Method):
)
self._sf_denoising_step_list: torch.Tensor | None = None
self._logged_denoising_steps = False
self._rollout_debug_calls = 0
self._fake_score_debug_calls = 0
self._real_score_debug_calls = 0
scheduler_timesteps = self._sf_scheduler.timesteps
logger.warning(
"SF_DEBUG init sample_type=%s chunk_size=%s "
"same_step_across_blocks=%s last_step_only=%s "
"enable_gradient_in_rollout=%s start_gradient_frame=%s "
"context_noise=%s raw_denoising_steps=%s "
"warp_denoising_step=%s flow_shift=%s cfg_uncond=%s "
"scheduler=%s scheduler_num_train_timesteps=%s "
"scheduler_timesteps_len=%s scheduler_timesteps_head=%s "
"scheduler_timesteps_tail=%s",
self._student_sample_type,
self._chunk_size,
self._same_step_across_blocks,
self._last_step_only,
self._enable_gradient_in_rollout,
self._start_gradient_frame,
self._context_noise,
self.method_config.get("dmd_denoising_steps"),
bool(self.method_config.get("warp_denoising_step", False)),
shift,
self._cfg_uncond,
type(self._sf_scheduler).__name__,
int(self.student.num_train_timesteps),
int(scheduler_timesteps.numel()),
_preview_tensor(scheduler_timesteps[:6]),
_preview_tensor(scheduler_timesteps[-6:]),
)
def _get_denoising_step_list(self, device: torch.device) -> torch.Tensor:
if (self._sf_denoising_step_list is not None and self._sf_denoising_step_list.device == device):
@@ -179,6 +234,8 @@ class SelfForcingMethod(DMD2Method):
)).to(device)
steps = timesteps[int(self.student.num_train_timesteps) - steps]
self._logged_denoising_steps = True
self._sf_denoising_step_list = steps
return steps
@@ -321,6 +378,32 @@ class SelfForcingMethod(DMD2Method):
device=device,
)
rollout_call = self._rollout_debug_calls
if rollout_call < 3:
logger.warning(
"SF_DEBUG rollout call=%s with_grad=%s latents=%s "
"noise_full=%s num_frames=%s chunk_size=%s remaining=%s "
"num_blocks=%s denoising_steps=%s exit_indices=%s "
"sample_type=%s same_step_across_blocks=%s "
"last_step_only=%s start_gradient_frame=%s "
"context_noise=%s",
rollout_call,
with_grad,
_shape_repr(latents),
_shape_repr(noise_full),
num_frames,
chunk,
remaining,
num_blocks,
_preview_tensor(denoising_steps, limit=16),
exit_indices,
self._student_sample_type,
self._same_step_across_blocks,
self._last_step_only,
self._start_gradient_frame,
self._context_noise,
)
denoised_blocks: list[torch.Tensor] = []
cache_tag = "pos"
@@ -340,6 +423,20 @@ class SelfForcingMethod(DMD2Method):
noisy_block = noise_full[:, start:end]
exit_idx = int(exit_indices[block_idx])
if rollout_call < 2:
logger.warning(
"SF_DEBUG rollout_block call=%s block_idx=%s "
"frame_range=[%s,%s) block_frames=%s exit_idx=%s "
"exit_timestep=%s noisy_block=%s",
rollout_call,
block_idx,
start,
end,
end - start,
exit_idx,
float(denoising_steps[exit_idx].item()),
_shape_repr(noisy_block),
)
for step_idx, current_timestep in enumerate(denoising_steps):
exit_flag = step_idx == exit_idx
@@ -462,6 +559,7 @@ class SelfForcingMethod(DMD2Method):
raise RuntimeError("Self-forcing rollout produced no blocks")
self.student.clear_caches(cache_tag=cache_tag)
self._rollout_debug_calls += 1
return torch.cat(denoised_blocks, dim=1)
def _critic_flow_matching_loss(self, batch: Any) -> tuple[torch.Tensor, Any, dict[str, Any]]:
@@ -498,6 +596,20 @@ class SelfForcingMethod(DMD2Method):
target = noise - generator_pred_x0
flow_matching_loss = torch.mean((pred_noise - target)**2)
if self._fake_score_debug_calls < 5:
logger.warning(
"SF_DEBUG fake_score call=%s timestep=%s "
"generator_pred_x0=%s noisy_x0=%s pred_noise=%s target=%s "
"critic_attn_kind=dense conditional=True",
self._fake_score_debug_calls,
int(fake_score_timestep.item()),
_shape_repr(generator_pred_x0),
_shape_repr(noisy_x0),
_shape_repr(pred_noise),
_shape_repr(target),
)
self._fake_score_debug_calls += 1
batch.fake_score_latent_vis_dict = {
"generator_pred_video": generator_pred_x0,
"fake_score_timestep": fake_score_timestep,
@@ -572,5 +684,30 @@ class SelfForcingMethod(DMD2Method):
grad = (faker_x0 - real_cfg_x0) / denom
grad = torch.nan_to_num(grad)
if self._real_score_debug_calls < 5:
logger.warning(
"SF_DEBUG real_score call=%s timestep=%s "
"guidance_scale=%s generator_pred_x0=%s noisy_latents=%s "
"faker_x0=%s real_cond_x0=%s real_uncond_x0=%s "
"real_cfg_x0=%s denom=%s",
self._real_score_debug_calls,
int(timestep.item()),
float(guidance_scale),
_shape_repr(generator_pred_x0),
_shape_repr(noisy_latents),
_shape_repr(faker_x0),
_shape_repr(real_cond_x0),
_shape_repr(real_uncond_x0),
_shape_repr(real_cfg_x0),
float(denom.detach().item()),
)
self._real_score_debug_calls += 1
batch.dmd_latent_vis_dict.update({
"generator_pred_video": generator_pred_x0.detach(),
"real_score_pred_video": real_cfg_x0.detach(),
"faker_score_pred_video": faker_x0.detach(),
"dmd_timestep": timestep.float().detach(),
})
loss = 0.5 * torch.mean((generator_pred_x0.float() - (generator_pred_x0.float() - grad.float()).detach())**2)
return loss
@@ -7,10 +7,12 @@ from typing import TYPE_CHECKING
if TYPE_CHECKING:
from fastvideo.train.methods.fine_tuning.dfsft import DiffusionForcingSFTMethod
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
from fastvideo.train.methods.fine_tuning.tfsft import TeacherForcingSFTMethod
__all__ = [
"DiffusionForcingSFTMethod",
"FineTuneMethod",
"TeacherForcingSFTMethod",
]
@@ -25,4 +27,9 @@ def __getattr__(name: str) -> object:
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
return FineTuneMethod
if name == "TeacherForcingSFTMethod":
from fastvideo.train.methods.fine_tuning.tfsft import (
TeacherForcingSFTMethod, )
return TeacherForcingSFTMethod
raise AttributeError(name)
+74 -33
View File
@@ -3,11 +3,14 @@
from __future__ import annotations
import copy
from typing import Any, Literal
import torch
import torch.nn.functional as F
from fastvideo.models.schedulers.scheduling_diffusion_forcing import (
DiffusionForcingScheduler, )
from fastvideo.train.methods.base import TrainingMethod, LogScalar
from fastvideo.train.models.base import ModelBase
from fastvideo.train.utils.optimizer import (
@@ -31,9 +34,15 @@ class DiffusionForcingSFTMethod(TrainingMethod):
raise ValueError("DFSFT requires role 'student'")
if not self.student._trainable:
raise ValueError("DFSFT requires student to be trainable")
if self.training_config.model.precondition_outputs:
raise ValueError(
"DFSFT only supports official diffusion-forcing loss; "
"set training.model.precondition_outputs=false"
)
self._attn_kind: Literal["dense", "vsa"] = (self._infer_attn_kind())
self._chunk_size = self._parse_chunk_size(self.method_config.get("chunk_size", None))
self._dfsft_scheduler = self._build_dfsft_scheduler()
self._timestep_index_range = (self._parse_timestep_index_range())
# Initialize preprocessors on student.
@@ -103,36 +112,30 @@ class DiffusionForcingSFTMethod(TrainingMethod):
if (sp_size > 1 and sp_group is not None and hasattr(sp_group, "broadcast")):
sp_group.broadcast(timestep_indices, src=0)
scheduler = self.student.noise_scheduler
if scheduler is None:
raise ValueError("DFSFT requires student.noise_scheduler")
schedule_timesteps = scheduler.timesteps.to(device=clean_latents.device, dtype=torch.float32)
schedule_sigmas = scheduler.sigmas.to(
scheduler = self._dfsft_scheduler
t_inhom = scheduler.timesteps.to(
device=clean_latents.device,
dtype=clean_latents.dtype,
)
t_inhom = schedule_timesteps[timestep_indices]
dtype=torch.float32,
)[timestep_indices]
# Override the homogeneous timesteps from prepare_batch
# so that set_forward_context (in predict_noise and
# backward) receives the correct per-chunk timesteps.
training_batch.timesteps = t_inhom
training_batch = self._refresh_attention_metadata(training_batch)
noise = getattr(training_batch, "noise", None)
if noise is None:
noise = torch.randn_like(clean_latents)
else:
if not torch.is_tensor(noise):
raise TypeError("TrainingBatch.noise must be a "
"torch.Tensor when set")
noise = noise.permute(0, 2, 1, 3, 4).to(dtype=clean_latents.dtype)
noisy_latents = self.student.add_noise(
clean_latents,
noise,
t_inhom.flatten(),
noise = torch.randn(
clean_latents.shape,
generator=self.cuda_generator,
device=clean_latents.device,
dtype=clean_latents.dtype,
)
noisy_latents = scheduler.add_noise(
clean_latents.flatten(0, 1),
noise.flatten(0, 1),
t_inhom.flatten(),
).unflatten(0, clean_latents.shape[:2])
training_batch.noise = noise
pred = self.student.predict_noise(
noisy_latents,
@@ -142,14 +145,17 @@ class DiffusionForcingSFTMethod(TrainingMethod):
attn_kind=self._attn_kind,
)
if bool(self.training_config.model.precondition_outputs):
sigmas = schedule_sigmas[timestep_indices]
sigmas = sigmas.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1)
pred_x0 = noisy_latents - pred * sigmas
loss = F.mse_loss(pred_x0.float(), clean_latents.float())
else:
target = noise - clean_latents
loss = F.mse_loss(pred.float(), target.float())
target = scheduler.training_target(clean_latents, noise, t_inhom)
per_frame_loss = F.mse_loss(
pred.float(),
target.float(),
reduction="none",
).mean(dim=(2, 3, 4))
weight = scheduler.training_weight(t_inhom).reshape(
batch_size,
num_latents,
)
loss = (per_frame_loss * weight.float()).mean()
if self._attn_kind == "vsa":
attn_metadata = training_batch.attn_metadata_vsa
@@ -245,9 +251,7 @@ class DiffusionForcingSFTMethod(TrainingMethod):
f"got {type(raw).__name__}")
def _parse_timestep_index_range(self, ) -> tuple[int, int]:
scheduler = self.student.noise_scheduler
if scheduler is None:
raise ValueError("DFSFT requires student.noise_scheduler")
scheduler = self._dfsft_scheduler
num_steps = int(getattr(scheduler, "config", scheduler).num_train_timesteps)
min_ratio = self._parse_ratio(
@@ -320,3 +324,40 @@ class DiffusionForcingSFTMethod(TrainingMethod):
)
expanded = chunk_indices.repeat_interleave(chunk_size, dim=1)
return expanded[:, :num_latents]
def _build_dfsft_scheduler(self) -> DiffusionForcingScheduler:
student_scheduler = getattr(self.student, "noise_scheduler", None)
if student_scheduler is None:
raise ValueError("DFSFT requires student.noise_scheduler")
num_steps = int(
getattr(student_scheduler, "config", student_scheduler).num_train_timesteps
)
pipeline_config = self.training_config.pipeline_config
if pipeline_config is None:
raise ValueError("DFSFT requires training_config.pipeline_config")
shift = float(
getattr(
pipeline_config,
"flow_shift",
getattr(self.student, "timestep_shift", 1.0),
)
)
scheduler = DiffusionForcingScheduler(
num_inference_steps=num_steps,
num_train_timesteps=num_steps,
shift=shift,
sigma_min=0.0,
extra_one_step=True,
training=True,
)
scheduler.set_timesteps(num_steps, training=True)
return scheduler
def _refresh_attention_metadata(self, training_batch: Any) -> Any:
training_batch = self.student._build_attention_metadata(training_batch)
training_batch.attn_metadata_vsa = copy.deepcopy(
training_batch.attn_metadata
)
if training_batch.attn_metadata is not None:
training_batch.attn_metadata.VSA_sparsity = 0.0
return training_batch
@@ -0,0 +1,531 @@
# SPDX-License-Identifier: Apache-2.0
"""Teacher-forcing SFT method (TFSFT; algorithm layer)."""
from __future__ import annotations
import copy
import math
from typing import Any, Literal
import torch
import torch.nn.functional as F
from torch.nn.attention.flex_attention import create_block_mask
from fastvideo.models.schedulers.scheduling_diffusion_forcing import (
DiffusionForcingScheduler, )
from fastvideo.train.methods.base import TrainingMethod, LogScalar
from fastvideo.train.models.base import ModelBase
from fastvideo.train.utils.optimizer import (
build_optimizer_and_scheduler, )
class TeacherForcingSFTMethod(TrainingMethod):
"""Teacher-forcing SFT (TFSFT) for causal students.
Training uses a concatenated ``[clean | noisy]`` latent sequence plus a
custom block mask so noisy chunks can attend to a clean prefix while the
denoising loss is applied only on the noisy half.
"""
def __init__(
self,
*,
cfg: Any,
role_models: dict[str, ModelBase],
) -> None:
super().__init__(cfg=cfg, role_models=role_models)
if "student" not in role_models:
raise ValueError("TFSFT requires role 'student'")
if not self.student._trainable:
raise ValueError("TFSFT requires student to be trainable")
if self.training_config.model.precondition_outputs:
raise ValueError(
"TFSFT only supports official diffusion-forcing loss; "
"set training.model.precondition_outputs=false"
)
self._attn_kind: Literal["dense", "vsa"] = (self._infer_attn_kind())
self._chunk_size = self._parse_chunk_size(
self.method_config.get("chunk_size", None))
self._tfsft_scheduler = self._build_tfsft_scheduler()
self._timestep_index_range = self._parse_timestep_index_range()
# Initialize preprocessors on student.
self.student.init_preprocessors(self.training_config)
self._block_mask_cache: dict[tuple[str, int, int, int], Any] = {}
self._init_optimizers_and_schedulers()
@property
def _optimizer_dict(self) -> dict[str, Any]:
return {"student": self._student_optimizer}
@property
def _lr_scheduler_dict(self) -> dict[str, Any]:
return {"student": self._student_lr_scheduler}
def single_train_step(
self,
batch: dict[str, Any],
iteration: int,
) -> tuple[
dict[str, torch.Tensor],
dict[str, Any],
dict[str, LogScalar],
]:
del iteration
training_batch = self.student.prepare_batch(
batch,
generator=self.cuda_generator,
latents_source="data",
)
if training_batch.latents is None:
raise RuntimeError("prepare_batch() must set TrainingBatch.latents")
clean_latents = training_batch.latents
if not torch.is_tensor(clean_latents):
raise TypeError("TrainingBatch.latents must be a torch.Tensor")
if clean_latents.ndim != 5:
raise ValueError("TrainingBatch.latents must be "
"[B, T, C, H, W], got "
f"shape={tuple(clean_latents.shape)}")
batch_size, num_latents = (
int(clean_latents.shape[0]),
int(clean_latents.shape[1]),
)
expected_chunk = getattr(
self.student.transformer,
"num_frame_per_block",
None,
)
if (expected_chunk is not None
and int(expected_chunk) != int(self._chunk_size)):
raise ValueError("TFSFT chunk_size must match "
"transformer.num_frame_per_block for "
f"causal training (got {self._chunk_size}, "
f"expected {expected_chunk}).")
timestep_indices = self._sample_t_inhom_indices(
batch_size=batch_size,
num_latents=num_latents,
device=clean_latents.device,
)
sp_size = int(self.training_config.distributed.sp_size)
sp_group = getattr(self.student, "sp_group", None)
if (sp_size > 1 and sp_group is not None
and hasattr(sp_group, "broadcast")):
sp_group.broadcast(timestep_indices, src=0)
scheduler = self._tfsft_scheduler
schedule_timesteps = scheduler.timesteps.to(
device=clean_latents.device,
dtype=torch.float32,
)
t_inhom = schedule_timesteps[timestep_indices]
noise = torch.randn(
clean_latents.shape,
generator=self.cuda_generator,
device=clean_latents.device,
dtype=clean_latents.dtype,
)
noisy_latents = scheduler.add_noise(
clean_latents.flatten(0, 1),
noise.flatten(0, 1),
t_inhom.flatten(),
).unflatten(0, clean_latents.shape[:2])
training_batch.noise = noise
concat_latents = torch.cat([clean_latents, noisy_latents], dim=1)
concat_timesteps = torch.cat([torch.zeros_like(t_inhom), t_inhom], dim=1)
saved_batch_state = self._expand_batch_for_concat(
training_batch,
num_concat_latents=int(concat_latents.shape[1]),
)
training_batch.timesteps = concat_timesteps
self._refresh_attention_metadata(training_batch)
transformer = self.student.transformer
custom_block_mask = self._get_teacher_forcing_block_mask(
transformer=transformer,
latents=concat_latents,
)
prev_block_mask = getattr(transformer, "block_mask", None)
try:
transformer.block_mask = custom_block_mask
pred = self.student.predict_noise(
concat_latents,
concat_timesteps,
training_batch,
conditional=True,
attn_kind=self._attn_kind,
)
finally:
transformer.block_mask = prev_block_mask
self._restore_batch_after_concat(training_batch, saved_batch_state)
pred_noisy = pred[:, num_latents:]
target = scheduler.training_target(clean_latents, noise, t_inhom)
per_frame_loss = F.mse_loss(
pred_noisy.float(),
target.float(),
reduction="none",
).mean(dim=(2, 3, 4))
weight = scheduler.training_weight(t_inhom).reshape(
batch_size,
num_latents,
)
total_loss = (per_frame_loss * weight.float()).mean()
if self._attn_kind == "vsa":
attn_metadata = training_batch.attn_metadata_vsa
else:
attn_metadata = training_batch.attn_metadata
loss_map = {"total_loss": total_loss, "tfsft_loss": total_loss}
outputs: dict[str, Any] = {
"_fv_backward": (
concat_timesteps,
attn_metadata,
)
}
metrics: dict[str, LogScalar] = {}
return loss_map, outputs, metrics
def backward(
self,
loss_map: dict[str, torch.Tensor],
outputs: dict[str, Any],
*,
grad_accum_rounds: int = 1,
) -> None:
grad_accum_rounds = max(1, int(grad_accum_rounds))
ctx = outputs.get("_fv_backward")
if ctx is None:
super().backward(
loss_map,
outputs,
grad_accum_rounds=grad_accum_rounds,
)
return
self.student.backward(
loss_map["total_loss"],
ctx,
grad_accum_rounds=grad_accum_rounds,
)
def get_optimizers(
self,
iteration: int,
) -> list[torch.optim.Optimizer]:
del iteration
return [self._student_optimizer]
def get_lr_schedulers(
self,
iteration: int,
) -> list[Any]:
del iteration
return [self._student_lr_scheduler]
def _refresh_attention_metadata(self, training_batch: Any) -> None:
build_fn = getattr(self.student, "_build_attention_metadata", None)
if not callable(build_fn):
return
build_fn(training_batch)
training_batch.attn_metadata_vsa = copy.deepcopy(
training_batch.attn_metadata)
if training_batch.attn_metadata is not None:
training_batch.attn_metadata.VSA_sparsity = 0.0 # type: ignore[attr-defined]
def _expand_batch_for_concat(
self,
training_batch: Any,
*,
num_concat_latents: int,
) -> dict[str, Any]:
saved: dict[str, Any] = {}
raw_latent_shape = getattr(training_batch, "raw_latent_shape", None)
if raw_latent_shape is not None and len(raw_latent_shape) == 5:
saved["raw_latent_shape"] = raw_latent_shape
batch_size, channels, _num_frames, height, width = raw_latent_shape
training_batch.raw_latent_shape = (
batch_size,
channels,
num_concat_latents,
height,
width,
)
temporal_dims = {
"image_latents": 2,
"mask_lat_size": 2,
"viewmats": 1,
"Ks": 1,
"action": 1,
"mouse_cond": 1,
"keyboard_cond": 1,
}
for name, dim in temporal_dims.items():
value = getattr(training_batch, name, None)
if isinstance(value, torch.Tensor):
saved[name] = value
setattr(training_batch, name, torch.cat([value, value], dim=dim))
return saved
def _restore_batch_after_concat(
self,
training_batch: Any,
saved: dict[str, Any],
) -> None:
for name, value in saved.items():
setattr(training_batch, name, value)
def _get_teacher_forcing_block_mask(
self,
*,
transformer: Any,
latents: torch.Tensor,
) -> Any:
num_frames = int(latents.shape[1])
if num_frames % 2 != 0:
raise ValueError("TFSFT concat-and-mask requires an even "
f"number of frames, got {num_frames}")
patch_size = getattr(transformer, "patch_size", (1, 1, 1))
if len(patch_size) != 3:
raise ValueError("Unexpected transformer.patch_size: "
f"{patch_size!r}")
patch_t, patch_h, patch_w = [int(v) for v in patch_size]
post_patch_num_frames = num_frames // patch_t
if post_patch_num_frames * patch_t != num_frames:
raise ValueError("TFSFT requires num_frames divisible by "
f"patch temporal size {patch_t}, got {num_frames}")
frame_seqlen = (int(latents.shape[-2]) // patch_h) * (
int(latents.shape[-1]) // patch_w)
clean_frames = post_patch_num_frames // 2
if clean_frames <= 0:
raise ValueError("TFSFT requires at least one clean frame")
device_key = str(latents.device)
cache_key = (
device_key,
post_patch_num_frames,
frame_seqlen,
self._chunk_size,
)
cached = self._block_mask_cache.get(cache_key)
if cached is not None:
return cached
total_length = post_patch_num_frames * frame_seqlen
padded_length = math.ceil(total_length / 128) * 128 - total_length
total_tokens = total_length + padded_length
prefix_starts = torch.zeros(total_tokens, device=latents.device, dtype=torch.long)
prefix_ends = torch.zeros_like(prefix_starts)
noisy_starts = torch.zeros_like(prefix_starts)
noisy_ends = torch.zeros_like(prefix_starts)
for frame_idx in range(post_patch_num_frames):
token_start = frame_idx * frame_seqlen
token_end = min(total_length, token_start + frame_seqlen)
if frame_idx < clean_frames:
chunk_end = min(clean_frames, ((frame_idx // self._chunk_size) + 1) * self._chunk_size)
prefix_start = 0
prefix_end = chunk_end * frame_seqlen
noisy_start = 0
noisy_end = 0
else:
noisy_frame = frame_idx - clean_frames
chunk_start = (noisy_frame // self._chunk_size) * self._chunk_size
chunk_end = min(clean_frames, chunk_start + self._chunk_size)
prefix_start = 0
prefix_end = chunk_start * frame_seqlen
noisy_start = (clean_frames + chunk_start) * frame_seqlen
noisy_end = (clean_frames + chunk_end) * frame_seqlen
prefix_starts[token_start:token_end] = prefix_start
prefix_ends[token_start:token_end] = prefix_end
noisy_starts[token_start:token_end] = noisy_start
noisy_ends[token_start:token_end] = noisy_end
def attention_mask(b, h, q_idx, kv_idx):
valid = (q_idx < total_length) & (kv_idx < total_length)
prefix_ok = ((kv_idx >= prefix_starts[q_idx])
& (kv_idx < prefix_ends[q_idx]))
noisy_ok = ((kv_idx >= noisy_starts[q_idx])
& (kv_idx < noisy_ends[q_idx]))
return valid & (prefix_ok | noisy_ok)
block_mask = create_block_mask(
attention_mask,
B=None,
H=None,
Q_LEN=total_tokens,
KV_LEN=total_tokens,
_compile=False,
device=latents.device,
)
self._block_mask_cache[cache_key] = block_mask
return block_mask
def _parse_chunk_size(self, raw: Any) -> int:
if raw in (None, ""):
return 3
if isinstance(raw, bool):
raise ValueError("method_config.chunk_size must be an int, "
"got bool")
if isinstance(raw, float) and not raw.is_integer():
raise ValueError("method_config.chunk_size must be an int, "
"got float")
if isinstance(raw, str) and not raw.strip():
raise ValueError("method_config.chunk_size must be an int, "
"got empty string")
try:
value = int(raw)
except (TypeError, ValueError) as e:
raise ValueError("method_config.chunk_size must be an int, "
f"got {type(raw).__name__}") from e
if value <= 0:
raise ValueError("method_config.chunk_size must be > 0")
return value
def _parse_ratio(
self,
raw: Any,
*,
where: str,
default: float,
) -> float:
if raw in (None, ""):
return float(default)
if isinstance(raw, bool):
raise ValueError(f"{where} must be a number/string, got bool")
if isinstance(raw, int | float):
return float(raw)
if isinstance(raw, str) and raw.strip():
return float(raw)
raise ValueError(f"{where} must be a number/string, "
f"got {type(raw).__name__}")
def _parse_timestep_index_range(self) -> tuple[int, int]:
scheduler = self._tfsft_scheduler
num_steps = int(
getattr(scheduler, "config", scheduler).num_train_timesteps)
min_ratio = self._parse_ratio(
self.method_config.get("min_timestep_ratio", None),
where="method.min_timestep_ratio",
default=0.0,
)
max_ratio = self._parse_ratio(
self.method_config.get("max_timestep_ratio", None),
where="method.max_timestep_ratio",
default=1.0,
)
if not (0.0 <= min_ratio <= 1.0 and 0.0 <= max_ratio <= 1.0):
raise ValueError("TFSFT timestep ratios must be in [0,1], "
f"got min={min_ratio}, max={max_ratio}")
if max_ratio < min_ratio:
raise ValueError("method_config.max_timestep_ratio must be "
">= min_timestep_ratio")
min_index = int(min_ratio * num_steps)
max_index = int(max_ratio * num_steps)
min_index = max(0, min(min_index, num_steps - 1))
max_index = max(0, min(max_index, num_steps - 1))
if max_index <= min_index:
max_index = min(num_steps - 1, min_index + 1)
return min_index, max_index + 1
def _init_optimizers_and_schedulers(self) -> None:
tc = self.training_config
student_lr = float(tc.optimizer.learning_rate)
if student_lr <= 0.0:
raise ValueError("training.learning_rate must be > 0 "
"for tfsft")
student_betas = tc.optimizer.betas
student_sched = str(tc.optimizer.lr_scheduler)
student_params = [
p for p in self.student.transformer.parameters() if p.requires_grad
]
(
self._student_optimizer,
self._student_lr_scheduler,
) = build_optimizer_and_scheduler(
params=student_params,
optimizer_config=tc.optimizer,
loop_config=tc.loop,
learning_rate=student_lr,
betas=student_betas,
scheduler_name=student_sched,
)
def _sample_t_inhom_indices(
self,
*,
batch_size: int,
num_latents: int,
device: torch.device,
) -> torch.Tensor:
chunk_size = self._chunk_size
num_chunks = ((num_latents + chunk_size - 1) // chunk_size)
low, high = self._timestep_index_range
chunk_indices = torch.randint(
low=low,
high=high,
size=(batch_size, num_chunks),
device=device,
dtype=torch.long,
generator=self.cuda_generator,
)
expanded = chunk_indices.repeat_interleave(chunk_size, dim=1)
return expanded[:, :num_latents]
def _build_tfsft_scheduler(self) -> DiffusionForcingScheduler:
student_scheduler = getattr(self.student, "noise_scheduler", None)
if student_scheduler is None:
raise ValueError("TFSFT requires student.noise_scheduler")
num_steps = int(
getattr(
student_scheduler,
"config",
student_scheduler,
).num_train_timesteps
)
pipeline_config = self.training_config.pipeline_config
if pipeline_config is None:
raise ValueError("TFSFT requires training_config.pipeline_config")
shift = float(
getattr(
pipeline_config,
"flow_shift",
getattr(self.student, "timestep_shift", 1.0),
)
)
scheduler = DiffusionForcingScheduler(
num_inference_steps=num_steps,
num_train_timesteps=num_steps,
shift=shift,
sigma_min=0.0,
extra_one_step=True,
training=True,
)
scheduler.set_timesteps(num_steps, training=True)
return scheduler
@@ -0,0 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
"""WanGame model plugin package."""
from fastvideo.train.models.wangame.wangame import (
WanGameModel as WanGameModel, )
from fastvideo.train.models.wangame.wangame_causal import (
WanGameCausalModel as WanGameCausalModel, )
+862
View File
@@ -0,0 +1,862 @@
# SPDX-License-Identifier: Apache-2.0
"""WanGame bidirectional model plugin (per-role instance)."""
from __future__ import annotations
import copy
from typing import Any, Literal, TYPE_CHECKING
import torch
import fastvideo.envs as envs
from fastvideo.distributed import (
get_local_torch_device,
get_sp_group,
get_world_group,
)
from fastvideo.forward_context import set_forward_context
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.pipelines import TrainingBatch
from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing, )
from fastvideo.training.training_utils import (
compute_density_for_timestep_sampling,
get_sigmas,
normalize_dit_input,
shift_timestep,
)
from fastvideo.utils import (
is_vmoba_available,
is_vsa_available,
set_random_seed,
)
from fastvideo.train.models.base import ModelBase
from fastvideo.train.utils.module_state import (
apply_trainable, )
from fastvideo.train.utils.moduleloader import (
load_module_from_path, )
if TYPE_CHECKING:
from fastvideo.train.utils.training_config import (
TrainingConfig, )
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder, )
from fastvideo.attention.backends.vmoba import (
VideoMobaAttentionMetadataBuilder, )
except Exception:
VideoSparseAttentionMetadataBuilder = None # type: ignore[assignment]
VideoMobaAttentionMetadataBuilder = None # type: ignore[assignment]
class WanGameModel(ModelBase):
"""WanGame per-role model: owns transformer + noise_scheduler."""
_transformer_cls_name: str = ("WanGameActionTransformer3DModel")
def __init__(
self,
*,
init_from: str,
training_config: TrainingConfig,
trainable: bool = True,
disable_custom_init_weights: bool = False,
flow_shift: float = 3.0,
enable_gradient_checkpointing_type: str | None = None,
transformer_override_safetensor: str | None = None,
) -> None:
self._init_from = str(init_from)
self._trainable = bool(trainable)
self.transformer = self._load_transformer(
init_from=self._init_from,
trainable=self._trainable,
disable_custom_init_weights=(disable_custom_init_weights),
enable_gradient_checkpointing_type=(enable_gradient_checkpointing_type),
training_config=training_config,
transformer_override_safetensor=(transformer_override_safetensor),
)
self.noise_scheduler = (FlowMatchEulerDiscreteScheduler(shift=float(flow_shift)))
# Filled by init_preprocessors (student only).
self.vae: Any = None
self.training_config: TrainingConfig = training_config
self.dataloader: Any = None
self.validator: Any = None
self.start_step: int = 0
self.world_group: Any = None
self.sp_group: Any = None
self.device: Any = get_local_torch_device()
self.noise_random_generator: (torch.Generator | None) = None
self.noise_gen_cuda: torch.Generator | None = None
self.timestep_shift: float = float(flow_shift)
self.num_train_timestep: int = int(self.noise_scheduler.num_train_timesteps)
self.min_timestep: int = 0
self.max_timestep: int = self.num_train_timestep
def _load_transformer(
self,
*,
init_from: str,
trainable: bool,
disable_custom_init_weights: bool,
enable_gradient_checkpointing_type: str | None,
training_config: TrainingConfig,
transformer_override_safetensor: str | None = None,
) -> torch.nn.Module:
transformer = load_module_from_path(
model_path=init_from,
module_type="transformer",
training_config=training_config,
disable_custom_init_weights=(disable_custom_init_weights),
override_transformer_cls_name=(self._transformer_cls_name),
transformer_override_safetensor=(transformer_override_safetensor),
)
transformer = apply_trainable(transformer, trainable=trainable)
ckpt_type = (enable_gradient_checkpointing_type or getattr(
getattr(training_config, "model", None),
"enable_gradient_checkpointing_type",
None,
))
if trainable and ckpt_type:
transformer = apply_activation_checkpointing(
transformer,
checkpointing_type=ckpt_type,
)
return transformer
# ------------------------------------------------------------------
# Lifecycle
# ------------------------------------------------------------------
def init_preprocessors(self, training_config: TrainingConfig) -> None:
"""Load VAE, build dataloader, seed RNGs."""
self.vae = load_module_from_path(
model_path=str(training_config.model_path),
module_type="vae",
training_config=training_config,
)
self.world_group = get_world_group()
self.sp_group = get_sp_group()
self._init_timestep_mechanics()
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_wangame, )
from fastvideo.train.utils.dataloader import (
build_parquet_wangame_train_dataloader, )
self.dataloader = (build_parquet_wangame_train_dataloader(
training_config.data,
parquet_schema=pyarrow_schema_wangame,
))
self.start_step = 0
# ------------------------------------------------------------------
# ModelBase overrides: timestep helpers
# ------------------------------------------------------------------
@property
def num_train_timesteps(self) -> int:
return int(self.num_train_timestep)
def shift_and_clamp_timestep(self, timestep: torch.Tensor) -> torch.Tensor:
timestep = shift_timestep(
timestep,
self.timestep_shift,
self.num_train_timestep,
)
return timestep.clamp(self.min_timestep, self.max_timestep)
# ------------------------------------------------------------------
# ModelBase overrides: lifecycle hooks
# ------------------------------------------------------------------
def on_train_start(self) -> None:
assert self.training_config is not None
tc = self.training_config
seed = tc.data.seed
if seed is None:
raise ValueError("training.data.seed must be set "
"for training")
global_rank = int(getattr(self.world_group, "rank", 0))
sp_world_size = int(tc.distributed.sp_size or 1)
if sp_world_size > 1:
sp_group_seed = int(seed) + (global_rank // sp_world_size)
set_random_seed(sp_group_seed)
else:
set_random_seed(int(seed) + global_rank)
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(int(seed))
self.noise_gen_cuda = torch.Generator(device=self.device).manual_seed(int(seed))
def get_rng_generators(self, ) -> dict[str, torch.Generator]:
generators: dict[str, torch.Generator] = {}
if self.noise_random_generator is not None:
generators["noise_cpu"] = (self.noise_random_generator)
if self.noise_gen_cuda is not None:
generators["noise_cuda"] = self.noise_gen_cuda
return generators
# ------------------------------------------------------------------
# ModelBase overrides: runtime primitives
# ------------------------------------------------------------------
def _trim_temporal_prefix(
self,
value: torch.Tensor,
*,
target_length: int,
where: str,
dim: int,
) -> torch.Tensor:
current_length = int(value.shape[dim])
if current_length < target_length:
raise ValueError(
f"{where} temporal dim mismatch: got {current_length}, "
f"expected at least {target_length}"
)
if current_length == target_length:
return value
index = [slice(None)] * value.ndim
index[dim] = slice(0, target_length)
return value[tuple(index)]
def prepare_batch(
self,
raw_batch: dict[str, Any],
*,
generator: torch.Generator | None = None,
current_vsa_sparsity: float = 0.0,
latents_source: Literal["data", "zeros"] = "data",
) -> TrainingBatch:
del generator
assert self.training_config is not None
tc = self.training_config
dtype = self._get_training_dtype()
device = self.device
training_batch = TrainingBatch(current_vsa_sparsity=current_vsa_sparsity)
infos = raw_batch.get("info_list")
if latents_source == "zeros":
clip_feature = raw_batch["clip_feature"]
batch_size = int(clip_feature.shape[0])
vae_config = (
tc.pipeline_config.vae_config.arch_config # type: ignore[union-attr]
)
num_channels = int(vae_config.z_dim)
spatial_compression_ratio = int(vae_config.spatial_compression_ratio)
latent_height = (int(tc.data.num_height) // spatial_compression_ratio)
latent_width = (int(tc.data.num_width) // spatial_compression_ratio)
latents = torch.zeros(
batch_size,
num_channels,
int(tc.data.num_latent_t),
latent_height,
latent_width,
device=device,
dtype=dtype,
)
elif latents_source == "data":
if "vae_latent" not in raw_batch:
raise ValueError("vae_latent not found in batch "
"and latents_source='data'")
latents = raw_batch["vae_latent"]
latents = self._trim_temporal_prefix(
latents,
target_length=int(tc.data.num_latent_t),
where="vae_latent",
dim=2,
)
latents = latents.to(device, dtype=dtype)
else:
raise ValueError(f"Unknown latents_source: "
f"{latents_source!r}")
if "clip_feature" not in raw_batch:
raise ValueError("clip_feature must be present for WanGame")
image_embeds = raw_batch["clip_feature"].to(device, dtype=dtype)
if "first_frame_latent" not in raw_batch:
raise ValueError("first_frame_latent must be present "
"for WanGame")
image_latents = raw_batch["first_frame_latent"]
image_latents = self._trim_temporal_prefix(
image_latents,
target_length=int(tc.data.num_latent_t),
where="first_frame_latent",
dim=2,
)
image_latents = image_latents.to(device, dtype=dtype)
pil_image = raw_batch.get("pil_image")
if isinstance(pil_image, torch.Tensor):
training_batch.preprocessed_image = (pil_image.to(device=device))
else:
training_batch.preprocessed_image = pil_image
keyboard_cond = raw_batch.get("keyboard_cond")
if (isinstance(keyboard_cond, torch.Tensor) and keyboard_cond.numel() > 0):
training_batch.keyboard_cond = keyboard_cond.to(device, dtype=dtype)
else:
training_batch.keyboard_cond = None
mouse_cond = raw_batch.get("mouse_cond")
if (isinstance(mouse_cond, torch.Tensor) and mouse_cond.numel() > 0):
training_batch.mouse_cond = mouse_cond.to(device, dtype=dtype)
else:
training_batch.mouse_cond = None
temporal_compression_ratio = (
tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio # type: ignore[union-attr]
)
expected_num_frames = ((tc.data.num_latent_t - 1) * temporal_compression_ratio + 1)
if training_batch.keyboard_cond is not None:
training_batch.keyboard_cond = self._trim_temporal_prefix(
training_batch.keyboard_cond,
target_length=int(expected_num_frames),
where="keyboard_cond",
dim=1,
)
if training_batch.mouse_cond is not None:
training_batch.mouse_cond = self._trim_temporal_prefix(
training_batch.mouse_cond,
target_length=int(expected_num_frames),
where="mouse_cond",
dim=1,
)
training_batch.latents = latents
training_batch.encoder_hidden_states = None
training_batch.encoder_attention_mask = None
training_batch.image_embeds = image_embeds
training_batch.image_latents = image_latents
training_batch.infos = infos
training_batch.latents = normalize_dit_input("wan", training_batch.latents, self.vae)
training_batch = self._prepare_dit_inputs(training_batch)
training_batch = self._build_attention_metadata(training_batch)
training_batch.attn_metadata_vsa = copy.deepcopy(training_batch.attn_metadata)
if training_batch.attn_metadata is not None:
training_batch.attn_metadata.VSA_sparsity = 0.0 # type: ignore[attr-defined]
training_batch.mask_lat_size = (self._build_i2v_mask_latents(image_latents))
viewmats, intrinsics, action_labels = (self._process_actions(training_batch))
training_batch.viewmats = viewmats
training_batch.Ks = intrinsics
training_batch.action = action_labels
return training_batch
def add_noise(
self,
clean_latents: torch.Tensor,
noise: torch.Tensor,
timestep: torch.Tensor,
) -> torch.Tensor:
b, t = clean_latents.shape[:2]
noisy = self.noise_scheduler.add_noise(
clean_latents.flatten(0, 1),
noise.flatten(0, 1),
timestep,
).unflatten(0, (b, t))
return noisy
def predict_x0(
self,
noisy_latents: torch.Tensor,
timestep: torch.Tensor,
batch: TrainingBatch,
*,
conditional: bool,
cfg_uncond: dict[str, Any] | None = None,
attn_kind: Literal["dense", "vsa"] = "dense",
) -> torch.Tensor:
device_type = self.device.type
dtype = noisy_latents.dtype
if attn_kind == "dense":
attn_metadata = batch.attn_metadata
elif attn_kind == "vsa":
attn_metadata = batch.attn_metadata_vsa
else:
raise ValueError(f"Unknown attn_kind: {attn_kind!r}")
with torch.autocast(device_type, dtype=dtype), set_forward_context(
current_timestep=batch.timesteps,
attn_metadata=attn_metadata,
):
cond_inputs = (self._select_cfg_condition_inputs(
batch,
conditional=conditional,
cfg_uncond=cfg_uncond,
))
input_kwargs = (self._build_distill_input_kwargs(
noisy_latents,
timestep,
image_embeds=cond_inputs["image_embeds"],
image_latents=cond_inputs["image_latents"],
mask_lat_size=cond_inputs["mask_lat_size"],
viewmats=cond_inputs["viewmats"],
Ks=cond_inputs["Ks"],
action=cond_inputs["action"],
mouse_cond=cond_inputs["mouse_cond"],
keyboard_cond=cond_inputs["keyboard_cond"],
))
transformer = self._get_transformer(timestep)
pred_noise = transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
pred_x0 = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=noisy_latents.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler,
).unflatten(0, pred_noise.shape[:2])
return pred_x0
def predict_noise(
self,
noisy_latents: torch.Tensor,
timestep: torch.Tensor,
batch: TrainingBatch,
*,
conditional: bool,
cfg_uncond: dict[str, Any] | None = None,
attn_kind: Literal["dense", "vsa"] = "dense",
) -> torch.Tensor:
device_type = self.device.type
dtype = noisy_latents.dtype
if attn_kind == "dense":
attn_metadata = batch.attn_metadata
elif attn_kind == "vsa":
attn_metadata = batch.attn_metadata_vsa
else:
raise ValueError(f"Unknown attn_kind: {attn_kind!r}")
with torch.autocast(device_type, dtype=dtype), set_forward_context(
current_timestep=batch.timesteps,
attn_metadata=attn_metadata,
):
cond_inputs = (self._select_cfg_condition_inputs(
batch,
conditional=conditional,
cfg_uncond=cfg_uncond,
))
input_kwargs = (self._build_distill_input_kwargs(
noisy_latents,
timestep,
image_embeds=cond_inputs["image_embeds"],
image_latents=cond_inputs["image_latents"],
mask_lat_size=cond_inputs["mask_lat_size"],
viewmats=cond_inputs["viewmats"],
Ks=cond_inputs["Ks"],
action=cond_inputs["action"],
mouse_cond=cond_inputs["mouse_cond"],
keyboard_cond=cond_inputs["keyboard_cond"],
))
transformer = self._get_transformer(timestep)
pred_noise = transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
return pred_noise
def backward(
self,
loss: torch.Tensor,
ctx: Any,
*,
grad_accum_rounds: int,
) -> None:
timesteps, attn_metadata = ctx
with set_forward_context(
current_timestep=timesteps,
attn_metadata=attn_metadata,
):
(loss / max(1, int(grad_accum_rounds))).backward()
# ------------------------------------------------------------------
# Internal helpers
# ------------------------------------------------------------------
def _get_training_dtype(self) -> torch.dtype:
return torch.bfloat16
def _init_timestep_mechanics(self) -> None:
assert self.training_config is not None
tc = self.training_config
self.timestep_shift = float(tc.pipeline_config.flow_shift # type: ignore[union-attr]
)
self.num_train_timestep = int(self.noise_scheduler.num_train_timesteps)
self.min_timestep = 0
self.max_timestep = self.num_train_timestep
def _sample_timesteps(self, batch_size: int, device: torch.device) -> torch.Tensor:
if self.noise_random_generator is None:
raise RuntimeError("on_train_start() must be called "
"before prepare_batch()")
assert self.training_config is not None
tc = self.training_config
u = compute_density_for_timestep_sampling(
weighting_scheme=tc.model.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=tc.model.logit_mean,
logit_std=tc.model.logit_std,
mode_scale=tc.model.mode_scale,
)
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
return self.noise_scheduler.timesteps[indices].to(device=device)
def _build_attention_metadata(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_config is not None
tc = self.training_config
latents_shape = training_batch.raw_latent_shape
patch_size = (
tc.pipeline_config.dit_config.patch_size # type: ignore[union-attr]
)
current_vsa_sparsity = (training_batch.current_vsa_sparsity)
assert latents_shape is not None
assert training_batch.timesteps is not None
if (envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN"):
if (not is_vsa_available() or VideoSparseAttentionMetadataBuilder is None):
raise ImportError("FASTVIDEO_ATTENTION_BACKEND is "
"VIDEO_SPARSE_ATTN, but "
"fastvideo_kernel is not correctly "
"installed or detected.")
training_batch.attn_metadata = VideoSparseAttentionMetadataBuilder().build( # type: ignore[misc]
raw_latent_shape=latents_shape[2:5],
current_timestep=(training_batch.timesteps),
patch_size=patch_size,
VSA_sparsity=current_vsa_sparsity,
device=self.device,
)
elif (envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN"):
if (not is_vmoba_available() or VideoMobaAttentionMetadataBuilder is None):
raise ImportError("FASTVIDEO_ATTENTION_BACKEND is "
"VMOBA_ATTN, but fastvideo_kernel "
"(or flash_attn>=2.7.4) is not "
"correctly installed.")
moba_params = tc.model.moba_config.copy()
moba_params.update({
"current_timestep": (training_batch.timesteps),
"raw_latent_shape": (training_batch.raw_latent_shape[2:5]),
"patch_size": patch_size,
"device": self.device,
})
training_batch.attn_metadata = VideoMobaAttentionMetadataBuilder().build(**
moba_params) # type: ignore[misc]
else:
training_batch.attn_metadata = None
return training_batch
def _prepare_dit_inputs(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_config is not None
tc = self.training_config
latents = training_batch.latents
assert isinstance(latents, torch.Tensor)
batch_size = latents.shape[0]
if self.noise_gen_cuda is None:
raise RuntimeError("on_train_start() must be called "
"before prepare_batch()")
noise = torch.randn(
latents.shape,
generator=self.noise_gen_cuda,
device=latents.device,
dtype=latents.dtype,
)
timesteps = self._sample_timesteps(batch_size, latents.device)
if int(tc.distributed.sp_size or 1) > 1:
self.sp_group.broadcast(timesteps, src=0)
sigmas = get_sigmas(
self.noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = ((1.0 - sigmas) * latents + sigmas * noise)
training_batch.noisy_model_input = (noisy_model_input)
training_batch.timesteps = timesteps
training_batch.sigmas = sigmas
training_batch.noise = noise
training_batch.raw_latent_shape = latents.shape
training_batch.latents = (training_batch.latents.permute(0, 2, 1, 3, 4))
return training_batch
def _build_i2v_mask_latents(self, image_latents: torch.Tensor) -> torch.Tensor:
assert self.training_config is not None
tc = self.training_config
temporal_compression_ratio = (
tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio # type: ignore[union-attr]
)
num_frames = ((tc.data.num_latent_t - 1) * temporal_compression_ratio + 1)
(
batch_size,
_num_channels,
_t,
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)
return mask_lat_size.to(
device=image_latents.device,
dtype=image_latents.dtype,
)
def _process_actions(self, training_batch: TrainingBatch) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
keyboard_cond = getattr(training_batch, "keyboard_cond", None)
mouse_cond = getattr(training_batch, "mouse_cond", None)
if keyboard_cond is None or mouse_cond is None:
raise ValueError("WanGame batch must provide "
"keyboard_cond and mouse_cond")
from fastvideo.models.dits.hyworld.pose import (
process_custom_actions, )
batch_size = int(training_batch.noisy_model_input.shape[0] # type: ignore[union-attr]
)
viewmats_list: list[torch.Tensor] = []
intrinsics_list: list[torch.Tensor] = []
action_labels_list: list[torch.Tensor] = []
for b in range(batch_size):
v, i, a = process_custom_actions(keyboard_cond[b], mouse_cond[b])
viewmats_list.append(v)
intrinsics_list.append(i)
action_labels_list.append(a)
viewmats = torch.stack(viewmats_list, dim=0).to(device=self.device, dtype=torch.bfloat16)
intrinsics = torch.stack(intrinsics_list, dim=0).to(device=self.device, dtype=torch.bfloat16)
action_labels = torch.stack(action_labels_list, dim=0).to(device=self.device, dtype=torch.bfloat16)
num_latent_t = int(training_batch.noisy_model_input.shape[2] # type: ignore[union-attr]
)
if int(action_labels.shape[1]) != num_latent_t:
raise ValueError("Action conditioning temporal dim "
"mismatch: "
f"action={tuple(action_labels.shape)} "
f"vs latent_t={num_latent_t}")
if int(viewmats.shape[1]) != num_latent_t:
raise ValueError("Viewmats temporal dim mismatch: "
f"viewmats={tuple(viewmats.shape)} "
f"vs latent_t={num_latent_t}")
return viewmats, intrinsics, action_labels
def _build_distill_input_kwargs(
self,
noisy_video_latents: torch.Tensor,
timestep: torch.Tensor,
*,
image_embeds: torch.Tensor,
image_latents: torch.Tensor,
mask_lat_size: torch.Tensor,
viewmats: torch.Tensor | None,
Ks: torch.Tensor | None,
action: torch.Tensor | None,
mouse_cond: torch.Tensor | None,
keyboard_cond: torch.Tensor | None,
) -> dict[str, Any]:
hidden_states = torch.cat(
[
noisy_video_latents.permute(0, 2, 1, 3, 4),
mask_lat_size,
image_latents,
],
dim=1,
)
return {
"hidden_states": hidden_states,
"encoder_hidden_states": None,
"timestep": timestep.to(device=self.device, dtype=torch.bfloat16),
"encoder_hidden_states_image": image_embeds,
"viewmats": viewmats,
"Ks": Ks,
"action": action,
"mouse_cond": mouse_cond,
"keyboard_cond": keyboard_cond,
"return_dict": False,
}
def _select_cfg_condition_inputs(
self,
batch: TrainingBatch,
*,
conditional: bool,
cfg_uncond: dict[str, Any] | None,
) -> dict[str, Any]:
image_embeds = batch.image_embeds
image_latents = batch.image_latents
mask_lat_size = batch.mask_lat_size
if image_embeds is None:
raise RuntimeError("WanGameModel requires "
"TrainingBatch.image_embeds")
if image_latents is None:
raise RuntimeError("WanGameModel requires "
"TrainingBatch.image_latents")
if mask_lat_size is None:
raise RuntimeError("WanGameModel requires "
"TrainingBatch.mask_lat_size")
viewmats = getattr(batch, "viewmats", None)
Ks = getattr(batch, "Ks", None)
action = getattr(batch, "action", None)
mouse_cond = getattr(batch, "mouse_cond", None)
keyboard_cond = getattr(batch, "keyboard_cond", None)
if conditional or cfg_uncond is None:
return {
"image_embeds": image_embeds,
"image_latents": image_latents,
"mask_lat_size": mask_lat_size,
"viewmats": viewmats,
"Ks": Ks,
"action": action,
"mouse_cond": mouse_cond,
"keyboard_cond": keyboard_cond,
}
on_missing_raw = cfg_uncond.get("on_missing", "error")
if not isinstance(on_missing_raw, str):
raise ValueError("method_config.cfg_uncond.on_missing "
"must be a string, got "
f"{type(on_missing_raw).__name__}")
on_missing = on_missing_raw.strip().lower()
if on_missing not in {"error", "ignore"}:
raise ValueError("method_config.cfg_uncond.on_missing "
"must be one of {error, ignore}, got "
f"{on_missing_raw!r}")
supported_channels = {"image", "action"}
for channel, policy_raw in cfg_uncond.items():
if channel in {"on_missing"}:
continue
if channel in supported_channels:
continue
if policy_raw is None:
continue
if not isinstance(policy_raw, str):
raise ValueError("method_config.cfg_uncond values "
"must be strings, got "
f"{channel}="
f"{type(policy_raw).__name__}")
policy = policy_raw.strip().lower()
if policy == "keep":
continue
if on_missing == "ignore":
continue
raise ValueError("WanGameModel does not support "
"cfg_uncond channel "
f"{channel!r} (policy={policy!r}). "
"Set cfg_uncond.on_missing=ignore or "
"remove the channel.")
def _get_policy(channel: str) -> str:
raw = cfg_uncond.get(channel, "keep")
if raw is None:
return "keep"
if not isinstance(raw, str):
raise ValueError("method_config.cfg_uncond values "
"must be strings, got "
f"{channel}={type(raw).__name__}")
policy = raw.strip().lower()
if policy not in {"keep", "zero", "drop"}:
raise ValueError("method_config.cfg_uncond values "
"must be one of "
"{keep, zero, drop}, got "
f"{channel}={raw!r}")
return policy
image_policy = _get_policy("image")
if image_policy == "zero":
image_embeds = torch.zeros_like(image_embeds)
image_latents = torch.zeros_like(image_latents)
mask_lat_size = torch.zeros_like(mask_lat_size)
elif image_policy == "drop":
raise ValueError("cfg_uncond.image=drop is not supported "
"for WanGame I2V; use "
"cfg_uncond.image=zero or keep.")
action_policy = _get_policy("action")
if action_policy == "zero":
if (viewmats is None or Ks is None or action is None):
if on_missing == "ignore":
pass
else:
raise ValueError("cfg_uncond.action=zero requires "
"action conditioning tensors, "
"but TrainingBatch is missing "
"{viewmats, Ks, action}.")
else:
viewmats = torch.zeros_like(viewmats)
Ks = torch.zeros_like(Ks)
action = torch.zeros_like(action)
if mouse_cond is not None:
mouse_cond = torch.zeros_like(mouse_cond)
if keyboard_cond is not None:
keyboard_cond = torch.zeros_like(keyboard_cond)
elif action_policy == "drop":
viewmats = None
Ks = None
action = None
mouse_cond = None
keyboard_cond = None
return {
"image_embeds": image_embeds,
"image_latents": image_latents,
"mask_lat_size": mask_lat_size,
"viewmats": viewmats,
"Ks": Ks,
"action": action,
"mouse_cond": mouse_cond,
"keyboard_cond": keyboard_cond,
}
def _get_transformer(self, timestep: torch.Tensor) -> torch.nn.Module:
return self.transformer
@@ -0,0 +1,505 @@
# SPDX-License-Identifier: Apache-2.0
"""WanGame causal model plugin (per-role instance, streaming/cache)."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Literal, TYPE_CHECKING
import torch
from fastvideo.forward_context import set_forward_context
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.train.models.base import CausalModelBase
from fastvideo.train.models.wangame.wangame import WanGameModel
if TYPE_CHECKING:
from fastvideo.train.utils.training_config import (
TrainingConfig, )
@dataclass(slots=True)
class _StreamingCaches:
kv_cache: list[dict[str, Any]]
crossattn_cache: list[dict[str, Any]] | None
frame_seq_length: int
local_attn_size: int
sliding_window_num_frames: int
batch_size: int
dtype: torch.dtype
device: torch.device
class WanGameCausalModel(WanGameModel, CausalModelBase):
"""WanGame per-role model with causal/streaming primitives."""
_transformer_cls_name: str = ("CausalWanGameActionTransformer3DModel")
def __init__(
self,
*,
init_from: str,
training_config: TrainingConfig,
trainable: bool = True,
disable_custom_init_weights: bool = False,
flow_shift: float = 3.0,
enable_gradient_checkpointing_type: str | None = None,
transformer_override_safetensor: str | None = None,
) -> None:
super().__init__(
init_from=init_from,
trainable=trainable,
disable_custom_init_weights=disable_custom_init_weights,
flow_shift=flow_shift,
enable_gradient_checkpointing_type=(enable_gradient_checkpointing_type),
training_config=training_config,
transformer_override_safetensor=(transformer_override_safetensor),
)
self._streaming_caches: dict[tuple[int, str], _StreamingCaches] = {}
# --- CausalModelBase override: clear_caches ---
def clear_caches(self, *, cache_tag: str = "pos") -> None:
self._streaming_caches.pop((id(self), str(cache_tag)), None)
# --- CausalModelBase override: predict_noise_streaming ---
def predict_noise_streaming(
self,
noisy_latents: torch.Tensor,
timestep: torch.Tensor,
batch: Any,
*,
conditional: bool,
cache_tag: str = "pos",
store_kv: bool = False,
cur_start_frame: int = 0,
cfg_uncond: dict[str, Any] | None = None,
attn_kind: Literal["dense", "vsa"] = "dense",
) -> torch.Tensor | None:
if attn_kind == "dense":
attn_metadata = batch.attn_metadata
elif attn_kind == "vsa":
attn_metadata = batch.attn_metadata_vsa
else:
raise ValueError(f"Unknown attn_kind: {attn_kind!r}")
cache_tag = str(cache_tag)
cur_start_frame = int(cur_start_frame)
if cur_start_frame < 0:
raise ValueError("cur_start_frame must be >= 0")
batch_size = int(noisy_latents.shape[0])
num_frames = int(noisy_latents.shape[1])
timestep_full = self._ensure_per_frame_timestep(
timestep=timestep,
batch_size=batch_size,
num_frames=num_frames,
device=noisy_latents.device,
)
transformer = self._get_transformer(timestep_full)
caches = self._get_or_init_streaming_caches(
cache_tag=cache_tag,
transformer=transformer,
noisy_latents=noisy_latents,
)
frame_seq_length = int(caches.frame_seq_length)
kv_cache = caches.kv_cache
crossattn_cache = caches.crossattn_cache
if (self._should_snapshot_streaming_cache() and torch.is_grad_enabled()):
kv_cache = self._snapshot_kv_cache_indices(kv_cache)
model_kwargs: dict[str, Any] = {
"kv_cache": kv_cache,
"crossattn_cache": crossattn_cache,
"current_start": cur_start_frame * frame_seq_length,
"start_frame": cur_start_frame,
"is_cache": bool(store_kv),
}
device_type = self.device.type
dtype = noisy_latents.dtype
with torch.autocast(device_type, dtype=dtype), set_forward_context(
current_timestep=batch.timesteps,
attn_metadata=attn_metadata,
):
cond_inputs = self._select_cfg_condition_inputs(
batch,
conditional=conditional,
cfg_uncond=cfg_uncond,
)
cond_inputs = self._slice_cond_inputs_for_streaming(
cond_inputs=cond_inputs,
cur_start_frame=cur_start_frame,
num_frames=num_frames,
)
input_kwargs = self._build_distill_input_kwargs(
noisy_latents,
timestep_full,
image_embeds=cond_inputs["image_embeds"],
image_latents=cond_inputs["image_latents"],
mask_lat_size=cond_inputs["mask_lat_size"],
viewmats=cond_inputs["viewmats"],
Ks=cond_inputs["Ks"],
action=cond_inputs["action"],
mouse_cond=cond_inputs["mouse_cond"],
keyboard_cond=cond_inputs["keyboard_cond"],
)
input_kwargs["timestep"] = timestep_full.to(device=self.device, dtype=torch.long)
input_kwargs.update(model_kwargs)
if store_kv:
with torch.no_grad():
_ = transformer(**input_kwargs)
return None
pred_noise = transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
return pred_noise
# --- CausalModelBase override: predict_x0_streaming ---
def predict_x0_streaming(
self,
noisy_latents: torch.Tensor,
timestep: torch.Tensor,
batch: Any,
*,
conditional: bool,
cache_tag: str = "pos",
store_kv: bool = False,
cur_start_frame: int = 0,
cfg_uncond: dict[str, Any] | None = None,
attn_kind: Literal["dense", "vsa"] = "dense",
) -> torch.Tensor | None:
pred_noise = self.predict_noise_streaming(
noisy_latents,
timestep,
batch,
conditional=conditional,
cache_tag=cache_tag,
store_kv=store_kv,
cur_start_frame=cur_start_frame,
cfg_uncond=cfg_uncond,
attn_kind=attn_kind,
)
if pred_noise is None:
return None
pred_x0 = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=noisy_latents.flatten(0, 1),
timestep=self.shift_and_clamp_timestep(
self._ensure_per_frame_timestep(
timestep=timestep,
batch_size=int(noisy_latents.shape[0]),
num_frames=int(noisy_latents.shape[1]),
device=noisy_latents.device,
).flatten()),
scheduler=self.noise_scheduler,
).unflatten(0, pred_noise.shape[:2])
return pred_x0
# --- internal helpers ---
def _ensure_per_frame_timestep(
self,
*,
timestep: torch.Tensor,
batch_size: int,
num_frames: int,
device: torch.device,
) -> torch.Tensor:
if timestep.ndim == 0:
return (timestep.view(1, 1).expand(batch_size, num_frames).to(device=device))
if timestep.ndim == 1:
if int(timestep.shape[0]) == batch_size:
return (timestep.view(batch_size, 1).expand(batch_size, num_frames).to(device=device))
raise ValueError("streaming timestep must be scalar, [B], or "
f"[B, T]; got shape={tuple(timestep.shape)}")
if timestep.ndim == 2:
return timestep.to(device=device)
raise ValueError("streaming timestep must be scalar, [B], or [B, T]; "
f"got ndim={int(timestep.ndim)}")
def _slice_cond_inputs_for_streaming(
self,
*,
cond_inputs: dict[str, Any],
cur_start_frame: int,
num_frames: int,
) -> dict[str, Any]:
start = int(cur_start_frame)
num_frames = int(num_frames)
if num_frames <= 0:
raise ValueError("num_frames must be positive for streaming")
if start < 0:
raise ValueError("cur_start_frame must be >= 0 for streaming")
end = start + num_frames
sliced: dict[str, Any] = dict(cond_inputs)
image_latents = cond_inputs.get("image_latents")
if isinstance(image_latents, torch.Tensor):
sliced["image_latents"] = image_latents[:, :, start:end]
mask_lat_size = cond_inputs.get("mask_lat_size")
if isinstance(mask_lat_size, torch.Tensor):
sliced["mask_lat_size"] = mask_lat_size[:, :, start:end]
viewmats = cond_inputs.get("viewmats")
if isinstance(viewmats, torch.Tensor):
sliced["viewmats"] = viewmats[:, start:end]
Ks = cond_inputs.get("Ks")
if isinstance(Ks, torch.Tensor):
sliced["Ks"] = Ks[:, start:end]
action = cond_inputs.get("action")
if isinstance(action, torch.Tensor):
sliced["action"] = action[:, start:end]
temporal_compression_ratio = int(
self.training_config.pipeline_config.vae_config.arch_config.temporal_compression_ratio)
raw_end_frame_idx = (1 + temporal_compression_ratio * max(0, end - 1))
mouse_cond = cond_inputs.get("mouse_cond")
if isinstance(mouse_cond, torch.Tensor):
sliced["mouse_cond"] = mouse_cond[:, :raw_end_frame_idx]
keyboard_cond = cond_inputs.get("keyboard_cond")
if isinstance(keyboard_cond, torch.Tensor):
sliced["keyboard_cond"] = keyboard_cond[:, :raw_end_frame_idx]
return sliced
def _get_or_init_streaming_caches(
self,
*,
cache_tag: str,
transformer: torch.nn.Module,
noisy_latents: torch.Tensor,
) -> _StreamingCaches:
key = (id(self), cache_tag)
cached = self._streaming_caches.get(key)
batch_size = int(noisy_latents.shape[0])
dtype = noisy_latents.dtype
device = noisy_latents.device
frame_seq_length = self._compute_frame_seq_length(transformer, noisy_latents)
local_attn_size = self._get_local_attn_size(transformer)
sliding_window_num_frames = (self._get_sliding_window_num_frames(transformer))
meta = (
frame_seq_length,
local_attn_size,
sliding_window_num_frames,
batch_size,
dtype,
device,
)
if cached is not None:
cached_meta = (
cached.frame_seq_length,
cached.local_attn_size,
cached.sliding_window_num_frames,
cached.batch_size,
cached.dtype,
cached.device,
)
if cached_meta == meta:
return cached
kv_cache = self._initialize_kv_cache(
transformer=transformer,
batch_size=batch_size,
dtype=dtype,
device=device,
frame_seq_length=frame_seq_length,
local_attn_size=local_attn_size,
sliding_window_num_frames=sliding_window_num_frames,
checkpoint_safe=(self._should_use_checkpoint_safe_kv_cache()),
)
crossattn_cache = self._initialize_crossattn_cache(transformer=transformer, device=device)
caches = _StreamingCaches(
kv_cache=kv_cache,
crossattn_cache=crossattn_cache,
frame_seq_length=frame_seq_length,
local_attn_size=local_attn_size,
sliding_window_num_frames=sliding_window_num_frames,
batch_size=batch_size,
dtype=dtype,
device=device,
)
self._streaming_caches[key] = caches
return caches
def _compute_frame_seq_length(
self,
transformer: torch.nn.Module,
noisy_latents: torch.Tensor,
) -> int:
latent_seq_length = int(noisy_latents.shape[-1]) * int(noisy_latents.shape[-2])
patch_size = getattr(transformer, "patch_size", None)
if patch_size is None:
patch_size = getattr(
getattr(transformer, "config", None),
"arch_config",
None,
)
patch_size = getattr(patch_size, "patch_size", None)
if patch_size is None:
raise ValueError("Unable to determine transformer.patch_size "
"for causal streaming")
patch_ratio = int(patch_size[-1]) * int(patch_size[-2])
if patch_ratio <= 0:
raise ValueError("Invalid patch_size for causal streaming")
return latent_seq_length // patch_ratio
def _get_sliding_window_num_frames(self, transformer: torch.nn.Module) -> int:
cfg = getattr(transformer, "config", None)
arch_cfg = getattr(cfg, "arch_config", None)
value = (getattr(arch_cfg, "sliding_window_num_frames", None) if arch_cfg is not None else None)
if value is None:
return 15
return int(value)
def _get_local_attn_size(self, transformer: torch.nn.Module) -> int:
try:
value = getattr(transformer, "local_attn_size", -1)
except Exception:
value = -1
if value is None:
return -1
return int(value)
def _initialize_kv_cache(
self,
*,
transformer: torch.nn.Module,
batch_size: int,
dtype: torch.dtype,
device: torch.device,
frame_seq_length: int,
local_attn_size: int,
sliding_window_num_frames: int,
checkpoint_safe: bool,
) -> list[dict[str, Any]]:
num_blocks = len(getattr(transformer, "blocks", []))
if num_blocks <= 0:
raise ValueError("Unexpected transformer.blocks for causal "
"streaming")
try:
num_attention_heads = int(transformer.num_attention_heads # type: ignore[attr-defined]
)
except AttributeError as e:
raise ValueError("Transformer is missing num_attention_heads") from e
try:
attention_head_dim = int(transformer.attention_head_dim # type: ignore[attr-defined]
)
except AttributeError:
try:
hidden_size = int(transformer.hidden_size # type: ignore[attr-defined]
)
except AttributeError as e:
raise ValueError("Transformer is missing attention_head_dim "
"and hidden_size") from e
attention_head_dim = hidden_size // max(1, num_attention_heads)
if local_attn_size != -1:
kv_cache_size = (int(local_attn_size) * int(frame_seq_length))
else:
kv_cache_size = int(frame_seq_length) * int(sliding_window_num_frames)
if checkpoint_safe:
tc = getattr(self, "training_config", None)
total_frames = int(tc.data.num_frames if tc is not None else 0)
if total_frames <= 0:
raise ValueError("training.num_frames must be set to enable "
"checkpoint-safe streaming KV cache; "
f"got {total_frames}")
kv_cache_size = max(
kv_cache_size,
int(frame_seq_length) * total_frames,
)
kv_cache: list[dict[str, Any]] = []
for _ in range(num_blocks):
kv_cache.append({
"k":
torch.zeros(
[
batch_size,
kv_cache_size,
num_attention_heads,
attention_head_dim,
],
dtype=dtype,
device=device,
),
"v":
torch.zeros(
[
batch_size,
kv_cache_size,
num_attention_heads,
attention_head_dim,
],
dtype=dtype,
device=device,
),
"global_end_index":
torch.zeros((), dtype=torch.long, device=device),
"local_end_index":
torch.zeros((), dtype=torch.long, device=device),
})
return kv_cache
def _should_use_checkpoint_safe_kv_cache(self) -> bool:
tc = getattr(self, "training_config", None)
if tc is not None:
checkpointing_type = tc.model.enable_gradient_checkpointing_type
else:
checkpointing_type = None
return bool(checkpointing_type) and bool(self._trainable)
def _should_snapshot_streaming_cache(self) -> bool:
return self._should_use_checkpoint_safe_kv_cache()
def _snapshot_kv_cache_indices(self, kv_cache: list[dict[str, Any]]) -> list[dict[str, Any]]:
snapshot: list[dict[str, Any]] = []
for block_cache in kv_cache:
global_end_index = block_cache.get("global_end_index")
local_end_index = block_cache.get("local_end_index")
if not isinstance(global_end_index, torch.Tensor) or not isinstance(local_end_index, torch.Tensor):
raise ValueError("Unexpected kv_cache index tensors; expected "
"tensors at kv_cache[*].{global_end_index, "
"local_end_index}")
copied = dict(block_cache)
copied["global_end_index"] = (global_end_index.detach().clone())
copied["local_end_index"] = (local_end_index.detach().clone())
snapshot.append(copied)
return snapshot
def _initialize_crossattn_cache(
self,
*,
transformer: torch.nn.Module,
device: torch.device,
) -> list[dict[str, Any]] | None:
num_blocks = len(getattr(transformer, "blocks", []))
if num_blocks <= 0:
return None
return [{
"is_init": False,
"k": torch.empty(0, device=device),
"v": torch.empty(0, device=device),
} for _ in range(num_blocks)]
+232 -14
View File
@@ -2,15 +2,18 @@
from __future__ import annotations
import os
import time
from collections.abc import Iterator
from dataclasses import dataclass
from typing import Any, TYPE_CHECKING
import numpy as np
import torch
from tqdm.auto import tqdm
from fastvideo.distributed import get_sp_group, get_world_group
from fastvideo.logger import init_logger
from fastvideo.train.callbacks.callback import CallbackDict
from fastvideo.train.methods.base import TrainingMethod
from fastvideo.train.utils.tracking import build_tracker
@@ -19,6 +22,8 @@ if TYPE_CHECKING:
from fastvideo.train.utils.training_config import (
TrainingConfig, )
logger = init_logger(__name__)
def _coerce_log_scalar(value: Any, *, where: str) -> float:
if isinstance(value, torch.Tensor):
@@ -32,6 +37,101 @@ def _coerce_log_scalar(value: Any, *, where: str) -> float:
f"{where}, got {type(value).__name__}")
def _maybe_log_resume_fingerprint(
method: TrainingMethod,
*,
global_rank: int,
step: int,
) -> None:
if os.getenv("FASTVIDEO_DEBUG_RESUME_HASH", "").lower() not in {"1", "true", "yes"}:
return
transformer = method.transformer_inference
fingerprints: list[str] = []
for idx, (name, param) in enumerate(transformer.named_parameters()):
if idx >= 3:
break
data = param.detach().reshape(-1)
if data.numel() == 0:
fingerprints.append(f"{name}:empty")
continue
sample = data[:16].float()
fingerprints.append(
f"{name}:shape={tuple(param.shape)} "
f"sample_sum={sample.sum().item():.8f} "
f"sample_mean={sample.mean().item():.8f} "
f"first={sample[0].item():.8f}"
)
print(
"DEBUG_RESUME_FINGERPRINT "
f"rank={global_rank} step={step} " + " | ".join(fingerprints),
flush=True,
)
def _decode_latent_video(
vae: Any,
latents: torch.Tensor,
) -> np.ndarray:
with torch.no_grad():
latents = latents.detach()
latents = latents.permute(0, 2, 1, 3, 4)
scaling_factor = getattr(vae, "scaling_factor", None)
if scaling_factor is not None:
if isinstance(scaling_factor, torch.Tensor):
latents = latents / scaling_factor.to(
latents.device,
latents.dtype,
)
else:
latents = latents / float(scaling_factor)
shift_factor = getattr(vae, "shift_factor", None)
if shift_factor is not None:
if isinstance(shift_factor, torch.Tensor):
latents = latents + shift_factor.to(
latents.device,
latents.dtype,
)
else:
latents = latents + float(shift_factor)
with torch.autocast(
device_type="cuda",
dtype=torch.bfloat16,
enabled=latents.is_cuda,
):
video = vae.decode(latents)
video = (video / 2 + 0.5).clamp(0, 1)
video = video.detach().cpu().float().permute(0, 2, 1, 3, 4)
return (video * 255).numpy().astype(np.uint8)
def _maybe_add_video_artifact(
tracker: Any,
videos: list[Any],
*,
vae: Any,
source_dict: dict[str, Any],
latent_key: str,
caption: str,
) -> None:
latents = source_dict.get(latent_key)
if not isinstance(latents, torch.Tensor):
return
video = _decode_latent_video(vae, latents)
artifact = tracker.video(
video,
caption=caption,
fps=24,
format="mp4",
)
if artifact is not None:
videos.append(artifact)
@dataclass(slots=True)
class TrainLoopState:
step: int
@@ -63,6 +163,92 @@ class Trainer:
training_config,
)
def _should_log_train_artifacts(
self,
step: int,
) -> bool:
validation_cb = self.callbacks.get_callback("validation")
every_steps = getattr(validation_cb, "every_steps", None)
if every_steps is None:
return False
every_steps = int(every_steps)
if every_steps <= 0:
return False
return step % every_steps == 0
def _log_train_artifacts(
self,
method: TrainingMethod,
outputs: dict[str, Any],
*,
step: int,
) -> None:
if self.global_rank != 0:
return
if not self._should_log_train_artifacts(step):
return
vae = getattr(method.student, "vae", None)
if vae is None:
return
artifacts: dict[str, Any] = {}
video_artifacts: list[Any] = []
dmd_latent_dict = outputs.get("dmd_latent_vis_dict")
if isinstance(dmd_latent_dict, dict) and dmd_latent_dict:
_maybe_add_video_artifact(
self.tracker,
video_artifacts,
latent_key="generator_pred_video",
caption="generator",
vae=vae,
source_dict=dmd_latent_dict,
)
_maybe_add_video_artifact(
self.tracker,
video_artifacts,
latent_key="real_score_pred_video",
caption="real_score",
vae=vae,
source_dict=dmd_latent_dict,
)
_maybe_add_video_artifact(
self.tracker,
video_artifacts,
latent_key="faker_score_pred_video",
caption="fake_score",
vae=vae,
source_dict=dmd_latent_dict,
)
for scalar_key in ("generator_timestep", "dmd_timestep"):
value = dmd_latent_dict.get(scalar_key)
if isinstance(value, torch.Tensor) and value.numel() == 1:
artifacts[scalar_key] = float(value.detach().item())
fake_score_latent_dict = outputs.get("fake_score_latent_vis_dict")
if isinstance(fake_score_latent_dict, dict) and fake_score_latent_dict:
_maybe_add_video_artifact(
self.tracker,
video_artifacts,
latent_key="generator_pred_video",
caption="critic_generator",
vae=vae,
source_dict=fake_score_latent_dict,
)
value = fake_score_latent_dict.get("fake_score_timestep")
if isinstance(value, torch.Tensor) and value.numel() == 1:
artifacts["fake_score_timestep"] = float(
value.detach().item()
)
if video_artifacts:
artifacts["train_visualization"] = video_artifacts
if artifacts:
self.tracker.log_artifacts(artifacts, step)
def _iter_dataloader(self, dataloader: Any) -> Iterator[dict[str, Any]]:
data_iter = iter(dataloader)
while True:
@@ -89,18 +275,21 @@ class Trainer:
method.set_tracker(self.tracker)
method.on_train_start()
resume_from_checkpoint = (tc.checkpoint.resume_from_checkpoint or "")
if checkpoint_manager is not None:
resumed_step = (checkpoint_manager.maybe_resume(resume_from_checkpoint=(resume_from_checkpoint)))
if resumed_step is not None:
start_step = int(resumed_step)
_maybe_log_resume_fingerprint(
method,
global_rank=self.global_rank,
step=start_step,
)
self.callbacks.on_train_start(
method,
iteration=start_step,
)
resume_from_checkpoint = (tc.checkpoint.resume_from_checkpoint or "")
if checkpoint_manager is not None:
if resume_from_checkpoint:
method.seed_optimizer_state_for_resume()
resumed_step = (checkpoint_manager.maybe_resume(resume_from_checkpoint=(resume_from_checkpoint)))
if resumed_step is not None:
start_step = int(resumed_step)
self.callbacks.on_validation_begin(
method,
iteration=start_step,
@@ -108,14 +297,9 @@ class Trainer:
method.optimizers_zero_grad(start_step)
data_stream = self._iter_dataloader(dataloader)
# Restore the RNG snapshot LAST — after dcp.load,
# after iter(dataloader), after everything that may
# have advanced the RNG as a side-effect.
if (checkpoint_manager is not None and resume_from_checkpoint):
checkpoint_manager.load_rng_snapshot(resume_from_checkpoint, )
progress = tqdm(
range(start_step + 1, max_steps + 1),
total=max_steps,
initial=start_step,
desc="Steps",
disable=self.local_rank > 0,
@@ -125,12 +309,14 @@ class Trainer:
loss_sums: dict[str, float] = {}
metric_sums: dict[str, float] = {}
last_outputs: dict[str, Any] = {}
for accum_iter in range(grad_accum):
batch = next(data_stream)
loss_map, outputs, step_metrics = method.single_train_step(
batch,
step,
)
last_outputs = outputs
method.backward(
loss_map,
@@ -166,6 +352,11 @@ class Trainer:
metrics["vsa_sparsity"] = float(tc.vsa_sparsity)
if self.global_rank == 0 and metrics:
self.tracker.log(metrics, step)
self._log_train_artifacts(
method,
last_outputs,
step=step,
)
self.callbacks.on_training_step_end(
method,
@@ -184,6 +375,33 @@ class Trainer:
method,
iteration=step,
)
if checkpoint_manager is not None:
validation_cb = self.callbacks.get_callback("validation")
latest_mf_metric: float | None = None
get_latest_metric = getattr(
validation_cb,
"get_latest_metric",
None,
)
if callable(get_latest_metric):
latest_mf_metric = get_latest_metric(
"mf_angle_err_mean",
step=step,
)
checkpoint_manager.maybe_save_best(
step=step,
metric_value=latest_mf_metric,
metric_name="mf_angle_err_mean",
start_step=int(
tc.checkpoint.best_checkpoint_start_step
or 0
),
top_k=int(
tc.checkpoint.best_checkpoint_top_k
or 1
),
)
self.callbacks.on_train_end(
method,
+417 -88
View File
@@ -3,6 +3,7 @@
from __future__ import annotations
import json
import math
import os
import random
import re
@@ -20,6 +21,9 @@ from fastvideo.logger import init_logger
logger = init_logger(__name__)
_CHECKPOINT_DIR_RE = re.compile(r"^checkpoint-(\d+)$")
_BEST_CHECKPOINT_DIR_RE = re.compile(
r"^checkpoint-best-step-(\d+)$"
)
def _is_stateful(obj: Any) -> bool:
@@ -37,11 +41,38 @@ def _barrier() -> None:
dist.barrier()
def _broadcast_int(value: int, *, src: int = 0) -> int:
if not (dist.is_available() and dist.is_initialized()):
return int(value)
backend = dist.get_backend()
use_cuda = (
backend == "nccl"
and torch.cuda.is_available()
)
device = (
torch.device("cuda")
if use_cuda
else torch.device("cpu")
)
tensor = torch.tensor(
[int(value)],
dtype=torch.int32,
device=device,
)
dist.broadcast(tensor, src=src)
return int(tensor.item())
def _parse_step_from_dir(checkpoint_dir: Path) -> int:
match = _CHECKPOINT_DIR_RE.match(checkpoint_dir.name)
match = (
_CHECKPOINT_DIR_RE.match(checkpoint_dir.name)
or _BEST_CHECKPOINT_DIR_RE.match(checkpoint_dir.name)
)
if not match:
raise ValueError(f"Invalid checkpoint directory name {checkpoint_dir.name!r}; "
"expected 'checkpoint-<step>'")
raise ValueError(
f"Invalid checkpoint directory name {checkpoint_dir.name!r}; "
"expected 'checkpoint-<step>' or 'checkpoint-best-step-<step>'"
)
return int(match.group(1))
@@ -53,7 +84,10 @@ def _find_latest_checkpoint(output_dir: Path) -> Path | None:
for child in output_dir.iterdir():
if not child.is_dir():
continue
if not _CHECKPOINT_DIR_RE.match(child.name):
if not (
_CHECKPOINT_DIR_RE.match(child.name)
or _BEST_CHECKPOINT_DIR_RE.match(child.name)
):
continue
if not (child / "dcp").is_dir():
continue
@@ -100,7 +134,10 @@ def _resolve_resume_checkpoint(resume_from_checkpoint: str, *, output_dir: str)
if path.is_dir() and path.name == "dcp":
path = path.parent
if path.is_dir() and _CHECKPOINT_DIR_RE.match(path.name):
if path.is_dir() and (
_CHECKPOINT_DIR_RE.match(path.name)
or _BEST_CHECKPOINT_DIR_RE.match(path.name)
):
if not (path / "dcp").is_dir():
raise FileNotFoundError(f"Missing dcp dir under checkpoint: {path / 'dcp'}")
return path
@@ -113,8 +150,9 @@ def _resolve_resume_checkpoint(resume_from_checkpoint: str, *, output_dir: str)
# Give a clearer error message.
out = Path(os.path.expanduser(str(output_dir))).resolve()
raise ValueError("Could not resolve resume checkpoint. Expected a checkpoint directory "
f"named 'checkpoint-<step>' (with 'dcp/' inside), or an output_dir "
f"containing such checkpoints. Got: {path} (output_dir={out}).")
"named 'checkpoint-<step>' or 'checkpoint-best-step-<step>' "
f"(with 'dcp/' inside), or an output_dir containing such "
f"checkpoints. Got: {path} (output_dir={out}).")
class _RoleModuleContainer(torch.nn.Module):
@@ -141,8 +179,7 @@ class _CallbackStateWrapper:
return self._callbacks.state_dict()
def load_state_dict(
self,
state_dict: dict[str, Any],
self, state_dict: dict[str, Any],
) -> None:
self._callbacks.load_state_dict(state_dict)
@@ -168,6 +205,7 @@ class CheckpointManager:
output_dir: str,
config: CheckpointConfig,
callbacks: Any | None = None,
tracker: Any | None = None,
raw_config: dict[str, Any] | None = None,
) -> None:
self.method = method
@@ -175,8 +213,10 @@ class CheckpointManager:
self.output_dir = str(output_dir)
self.config = config
self._callbacks = callbacks
self._tracker = tracker
self._raw_config = raw_config
self._last_saved_step: int | None = None
self._last_best_saved_step: int | None = None
def _build_states(self) -> dict[str, Any]:
states: dict[str, Any] = self.method.checkpoint_state()
@@ -187,7 +227,9 @@ class CheckpointManager:
# Callback state (e.g. EMA shadow weights, validation RNG).
if self._callbacks is not None and _is_stateful(self._callbacks):
states["callbacks"] = _CallbackStateWrapper(self._callbacks, )
states["callbacks"] = _CallbackStateWrapper(
self._callbacks,
)
return states
@@ -197,6 +239,15 @@ class CheckpointManager:
def _dcp_dir(self, step: int) -> Path:
return self._checkpoint_dir(step) / "dcp"
def _best_checkpoint_dir(self, step: int) -> Path:
return (
Path(self.output_dir)
/ f"checkpoint-best-step-{step}"
)
def _checkpoint_best_alias_path(self) -> Path:
return Path(self.output_dir) / "checkpoint-best"
def maybe_save(self, step: int) -> None:
save_steps = int(self.config.save_steps or 0)
if save_steps <= 0:
@@ -215,120 +266,298 @@ class CheckpointManager:
def save(self, step: int) -> None:
checkpoint_dir = self._checkpoint_dir(step)
dcp_dir = self._dcp_dir(step)
self._save_checkpoint_dir(
checkpoint_dir,
step=step,
log_prefix="Saving checkpoint",
)
self._last_saved_step = step
self._cleanup_old_checkpoints()
def _save_checkpoint_dir(
self,
checkpoint_dir: Path,
*,
step: int,
log_prefix: str,
metadata_extra: dict[str, Any] | None = None,
) -> None:
dcp_dir = checkpoint_dir / "dcp"
os.makedirs(dcp_dir, exist_ok=True)
states = self._build_states()
if _rank() == 0:
logger.info(
"Saving checkpoint to %s",
logger.info("%s to %s", log_prefix, checkpoint_dir)
self._write_metadata(
checkpoint_dir,
step,
extra=metadata_extra,
)
self._write_metadata(checkpoint_dir, step)
dcp.save(states, checkpoint_id=str(dcp_dir))
_barrier()
# Save RNG state AFTER dcp.save so it captures the
# exact state the continuous run continues with.
# dcp.save triggers FSDP all-gather ops that can
# advance the RNG between when DCP captures it and
# when the save completes.
self._save_rng_snapshot(checkpoint_dir)
_barrier()
self._last_saved_step = step
self._cleanup_old_checkpoints()
def _write_metadata(
self,
checkpoint_dir: Path,
step: int,
*,
extra: dict[str, Any] | None = None,
) -> None:
metadata: dict[str, Any] = {"step": step}
if self._raw_config is not None:
metadata["config"] = self._raw_config
if extra:
metadata.update(extra)
meta_path = checkpoint_dir / "metadata.json"
with open(meta_path, "w", encoding="utf-8") as f:
json.dump(metadata, f, indent=2)
@staticmethod
def load_metadata(checkpoint_dir: str | Path, ) -> dict[str, Any]:
def load_metadata(
checkpoint_dir: str | Path,
) -> dict[str, Any]:
"""Read ``metadata.json`` from a checkpoint dir."""
meta_path = Path(checkpoint_dir) / "metadata.json"
if not meta_path.is_file():
raise FileNotFoundError(f"No metadata.json in {checkpoint_dir}")
raise FileNotFoundError(
f"No metadata.json in {checkpoint_dir}"
)
with open(meta_path, encoding="utf-8") as f:
return json.load(f) # type: ignore[no-any-return]
def _save_rng_snapshot(self, checkpoint_dir: Path) -> None:
"""Save per-rank RNG state to a separate file.
Called AFTER ``dcp.save`` so the snapshot reflects
the exact state the continuous run continues with.
Each rank saves its own file because CUDA RNG and
custom generators differ across ranks.
"""
rank = _rank()
rng: dict[str, Any] = {
"torch_rng": torch.get_rng_state(),
"python_rng": random.getstate(),
"numpy_rng": np.random.get_state(),
}
rng["cuda_rng"] = torch.cuda.get_rng_state()
rng["gen_cuda"] = self.method.cuda_generator.get_state()
torch.save(
rng,
checkpoint_dir / f"rng_state_rank{rank}.pt",
)
def load_rng_snapshot(
def _list_best_checkpoint_entries(
self,
checkpoint_path: str,
) -> None:
"""Restore per-rank RNG state from the snapshot file.
metric_name: str,
) -> list[dict[str, Any]]:
output_dir = Path(self.output_dir)
if not output_dir.is_dir():
return []
Must be called AFTER ``dcp.load`` **and** after
``iter(dataloader)`` so no later operation can
clobber the restored state.
"""
resolved = _resolve_resume_checkpoint(
checkpoint_path,
output_dir=self.output_dir,
)
if resolved is None:
return
rank = _rank()
rng_path = resolved / f"rng_state_rank{rank}.pt"
if not rng_path.is_file():
# Fall back to legacy single-file snapshot.
rng_path = resolved / "rng_state.pt"
if not rng_path.is_file():
logger.warning(
"No rng_state in %s; skipping "
"RNG snapshot restore.",
resolved,
entries: list[dict[str, Any]] = []
for child in output_dir.iterdir():
if (
not child.is_dir()
or not _BEST_CHECKPOINT_DIR_RE.match(
child.name
)
):
continue
metric_path = child / "best_metric.json"
if not metric_path.is_file():
logger.warning(
"Skipping %s: best_metric.json "
"is missing",
child,
)
continue
try:
with open(
metric_path, encoding="utf-8"
) as f:
meta = json.load(f)
metric_raw = meta.get(metric_name)
if metric_raw is None:
metric_raw = meta.get(
"mf_angle_err_mean"
)
if metric_raw is None:
metric_raw = meta.get(
"metric_value"
)
step_raw = meta.get("step")
if step_raw is None:
match = (
_BEST_CHECKPOINT_DIR_RE.match(
child.name
)
)
if match is None:
continue
step_raw = int(match.group(1))
metric_val = float(metric_raw)
step_val = int(step_raw)
if not math.isfinite(metric_val):
raise ValueError(
"metric is non-finite"
)
except Exception as e:
logger.warning(
"Skipping %s: invalid "
"best metric metadata (%s)",
child,
e,
)
continue
entries.append(
{
"path": child,
"step": step_val,
"metric": metric_val,
}
)
entries.sort(
key=lambda x: (
float(x["metric"]),
int(x["step"]),
)
)
return entries
def _update_best_checkpoint_alias(
self,
best_checkpoint_path: Path,
) -> None:
alias_path = self._checkpoint_best_alias_path()
try:
if alias_path.is_symlink() or alias_path.is_file():
alias_path.unlink()
elif alias_path.is_dir():
shutil.rmtree(alias_path)
os.symlink(
os.path.basename(str(best_checkpoint_path)),
str(alias_path),
)
except OSError as e:
logger.warning(
"Failed to update checkpoint-best "
"alias: %s",
e,
)
def maybe_save_best(
self,
*,
step: int,
metric_value: float | None,
metric_name: str = "mf_angle_err_mean",
start_step: int = 0,
top_k: int = 1,
) -> None:
start_step = int(start_step or 0)
if start_step <= 0:
return
if int(step) < start_step:
return
if metric_value is None:
return
metric = float(metric_value)
if not math.isfinite(metric):
return
if self._last_best_saved_step == int(step):
return
rng = torch.load(
rng_path,
map_location="cpu",
weights_only=False,
)
if "torch_rng" in rng:
torch.set_rng_state(rng["torch_rng"])
if "python_rng" in rng:
random.setstate(rng["python_rng"])
if "numpy_rng" in rng:
np.random.set_state(rng["numpy_rng"])
top_k = max(1, int(top_k or 1))
metric_name = str(metric_name)
torch.cuda.set_rng_state(rng["cuda_rng"])
self.method.cuda_generator.set_state(rng["gen_cuda"])
logger.info(
"Restored RNG snapshot from %s",
rng_path,
should_save = 0
if _rank() == 0:
entries = self._list_best_checkpoint_entries(
metric_name
)
if len(entries) < top_k:
should_save = 1
else:
worst = entries[-1]
if metric < float(
worst["metric"]
):
should_save = 1
should_save = _broadcast_int(
should_save, src=0
)
if should_save == 0:
return
checkpoint_dir = self._best_checkpoint_dir(
int(step)
)
if _rank() == 0 and checkpoint_dir.exists():
shutil.rmtree(checkpoint_dir, ignore_errors=True)
_barrier()
logger.info(
"%s=%.6f at step %s entered top-%s "
"best checkpoints; saving.",
metric_name,
metric,
step,
top_k,
)
self._save_checkpoint_dir(
checkpoint_dir,
step=int(step),
log_prefix="Saving best checkpoint",
metadata_extra={
"kind": "best",
"metric_name": metric_name,
"metric_value": metric,
},
)
self._last_best_saved_step = int(step)
if _rank() == 0:
metric_meta = {
"step": int(step),
metric_name: metric,
"metric_name": metric_name,
"metric_value": metric,
}
if metric_name == "mf_angle_err_mean":
metric_meta[
"mf_angle_err_mean"
] = metric
metric_path = (
checkpoint_dir / "best_metric.json"
)
with open(
metric_path, "w", encoding="utf-8"
) as f:
json.dump(metric_meta, f, indent=2)
best_entries = self._list_best_checkpoint_entries(
metric_name
)
kept_entries = best_entries[:top_k]
for stale_entry in best_entries[top_k:]:
stale_path = Path(
stale_entry["path"]
)
logger.info(
"Removing non-top-k best "
"checkpoint: %s",
stale_path,
)
shutil.rmtree(
stale_path, ignore_errors=True
)
if kept_entries:
top1 = kept_entries[0]
self._update_best_checkpoint_alias(
Path(top1["path"])
)
if self._tracker is not None:
self._tracker.log(
{
f"best/{metric_name}": float(
top1["metric"]
),
"best/step": int(
top1["step"]
),
"best/topk_count": len(
kept_entries
),
},
int(step),
)
_barrier()
def maybe_resume(self, *, resume_from_checkpoint: str | None) -> int | None:
if not resume_from_checkpoint:
@@ -344,11 +573,111 @@ class CheckpointManager:
states = self._build_states()
logger.info("Loading Phase 2 checkpoint from %s", resolved)
dcp.load(states, checkpoint_id=str(resolved / "dcp"))
try:
dcp.load(states, checkpoint_id=str(resolved / "dcp"))
except BaseException as exc:
if not isinstance(exc, dcp.CheckpointException):
raise
msg = str(exc)
fallback_prefixes = (
"optimizers.",
"schedulers.",
"dataloader",
"callbacks.",
"random_state",
)
can_fallback = (
"Missing key in checkpoint state_dict:" in msg
and any(
f"Missing key in checkpoint state_dict: {prefix}"
in msg
for prefix in fallback_prefixes
)
)
if not can_fallback:
raise
model_only_states = {
key: value
for key, value in states.items()
if key.startswith("roles.")
}
if not model_only_states:
raise
logger.warning(
"Resume checkpoint is missing non-model state "
"(optimizer/scheduler/dataloader/callback/RNG). "
"Falling back to model-only restore from %s "
"with %d role states. Optimizer/scheduler/etc. "
"will be reinitialized.",
resolved,
len(model_only_states),
)
dcp.load(
model_only_states,
checkpoint_id=str(resolved / "dcp"),
)
_barrier()
logger.info("Checkpoint loaded; resuming from step=%s", step)
return step
def _save_rng_snapshot(self, checkpoint_dir: Path) -> None:
"""Save per-rank RNG state after DCP save completes."""
rank = _rank()
rng: dict[str, Any] = {
"torch_rng": torch.get_rng_state(),
"python_rng": random.getstate(),
"numpy_rng": np.random.get_state(),
}
if torch.cuda.is_available():
rng["cuda_rng"] = torch.cuda.get_rng_state()
cuda_generator = getattr(self.method, "cuda_generator", None)
if cuda_generator is not None:
rng["gen_cuda"] = cuda_generator.get_state()
torch.save(
rng,
checkpoint_dir / f"rng_state_rank{rank}.pt",
)
def load_rng_snapshot(
self,
checkpoint_path: str,
) -> None:
resolved = _resolve_resume_checkpoint(
checkpoint_path,
output_dir=self.output_dir,
)
if resolved is None:
return
rank = _rank()
rng_path = resolved / f"rng_state_rank{rank}.pt"
if not rng_path.is_file():
rng_path = resolved / "rng_state.pt"
if not rng_path.is_file():
logger.warning(
"No rng_state in %s; skipping RNG snapshot restore.",
resolved,
)
return
rng = torch.load(
rng_path,
map_location="cpu",
weights_only=False,
)
if "torch_rng" in rng:
torch.set_rng_state(rng["torch_rng"])
if "python_rng" in rng:
random.setstate(rng["python_rng"])
if "numpy_rng" in rng:
np.random.set_state(rng["numpy_rng"])
if torch.cuda.is_available() and "cuda_rng" in rng:
torch.cuda.set_rng_state(rng["cuda_rng"])
cuda_generator = getattr(self.method, "cuda_generator", None)
if cuda_generator is not None and "gen_cuda" in rng:
cuda_generator.set_state(rng["gen_cuda"])
def _cleanup_old_checkpoints(self) -> None:
keep_last = int(self.config.keep_last or 0)
if keep_last <= 0:
+3
View File
@@ -321,6 +321,7 @@ def _build_training_config(
data_path=str(da.get("data_path", "") or ""),
train_batch_size=int(da.get("train_batch_size", 1) or 1),
dataloader_num_workers=int(da.get("dataloader_num_workers", 0) or 0),
apply_bot_died_filter=bool(da.get("apply_bot_died_filter", False)),
training_cfg_rate=float(da.get("training_cfg_rate", 0.0) or 0.0),
seed=int(da.get("seed", 0) or 0),
num_height=int(da.get("num_height", 0) or 0),
@@ -347,6 +348,8 @@ def _build_training_config(
resume_from_checkpoint=str(ck.get("resume_from_checkpoint", "") or ""),
training_state_checkpointing_steps=int(ck.get("training_state_checkpointing_steps", 0) or 0),
checkpoints_total_limit=int(ck.get("checkpoints_total_limit", 0) or 0),
best_checkpoint_start_step=int(ck.get("best_checkpoint_start_step", 0) or 0),
best_checkpoint_top_k=max(1, int(ck.get("best_checkpoint_top_k", 1) or 1)),
),
tracker=TrackerConfig(
trackers=list(tr.get("trackers", []) or []),
+89
View File
@@ -4,10 +4,68 @@ from __future__ import annotations
from typing import Any, TYPE_CHECKING
from fastvideo.logger import init_logger
if TYPE_CHECKING:
from fastvideo.train.utils.training_config import (
DataConfig, )
logger = init_logger(__name__)
def _maybe_apply_bot_died_filter(
*,
data_config: "DataConfig",
dataset: Any,
) -> None:
if not bool(getattr(data_config, "apply_bot_died_filter", False)):
return
from fastvideo.dataset.parquet_dataset_map_style import (
build_bot_died_excluded_indices,
)
from fastvideo.distributed import (
get_world_group,
get_world_rank,
)
world_group = get_world_group()
world_rank = int(get_world_rank())
total = int(sum(dataset.lengths))
bot_died_excluded: set[int] | None = None
if world_rank == 0:
bot_died_excluded = build_bot_died_excluded_indices(
data_path=str(data_config.data_path),
parquet_files=list(dataset.parquet_files),
lengths=list(dataset.lengths),
)
bot_died_excluded = world_group.broadcast_object(bot_died_excluded, src=0)
excluded_count = int(len(bot_died_excluded or set()))
remaining_count = int(total - excluded_count)
if bot_died_excluded:
valid_indices = [i for i in range(total) if i not in bot_died_excluded]
dataset.sampler.set_candidate_indices(valid_indices, epoch=0)
local_samples = int(len(getattr(dataset.sampler, "sp_group_local_indices", [])))
local_batches = int(len(dataset.sampler))
if world_rank == 0:
logger.info(
"BOT_DIED_FILTER_SUMMARY enabled=true total=%d excluded=%d remaining=%d",
total,
excluded_count,
remaining_count,
)
logger.info(
"BOT_DIED_FILTER_LOCAL rank=%d samples=%d batches=%d",
world_rank,
local_samples,
local_batches,
)
def build_parquet_t2v_train_dataloader(
data_config: DataConfig,
@@ -30,4 +88,35 @@ def build_parquet_t2v_train_dataloader(
text_padding_length=int(text_len),
seed=int(data_config.seed or 0),
))
_maybe_apply_bot_died_filter(
data_config=data_config,
dataset=_dataset,
)
return dataloader
def build_parquet_wangame_train_dataloader(
data_config: DataConfig,
*,
parquet_schema: Any,
) -> Any:
"""Build a parquet dataloader for WanGame datasets."""
from fastvideo.dataset import (
build_parquet_map_style_dataloader, )
_dataset, dataloader = (build_parquet_map_style_dataloader(
data_config.data_path,
data_config.train_batch_size,
num_data_workers=(data_config.dataloader_num_workers),
parquet_schema=parquet_schema,
cfg_rate=float(data_config.training_cfg_rate or 0.0),
drop_last=True,
text_padding_length=512,
seed=int(data_config.seed or 0),
))
_maybe_apply_bot_died_filter(
data_config=data_config,
dataset=_dataset,
)
return dataloader
+3
View File
@@ -25,6 +25,7 @@ class DataConfig:
data_path: str = ""
train_batch_size: int = 1
dataloader_num_workers: int = 0
apply_bot_died_filter: bool = False
training_cfg_rate: float = 0.0
seed: int = 0
num_height: int = 0
@@ -57,6 +58,8 @@ class CheckpointConfig:
resume_from_checkpoint: str = ""
training_state_checkpointing_steps: int = 0
checkpoints_total_limit: int = 0
best_checkpoint_start_step: int = 0
best_checkpoint_top_k: int = 1
@dataclass(slots=True)
+14 -4
View File
@@ -21,10 +21,20 @@ class ModelWrapper(torch.distributed.checkpoint.stateful.Stateful):
state_dict = get_model_state_dict(
self.model) # type: ignore[no-any-return]
# filter out non-trainable parameters
param_requires_grad = set([
k for k, v in dict(self.model.named_parameters()).items()
if v.requires_grad
])
param_requires_grad: set[str] = set()
for name, param in self.model.named_parameters():
if not bool(param.requires_grad):
continue
param_requires_grad.add(name)
# Activation checkpointing wraps modules with an internal attribute
# `_checkpoint_wrapped_module`, which changes the parameter names
# returned by `named_parameters()` but not the keys returned by
# `get_model_state_dict()`.
if "._checkpoint_wrapped_module." in name:
param_requires_grad.add(
name.replace("._checkpoint_wrapped_module.", ".")
)
state_dict = {
k: v
for k, v in state_dict.items() if k in param_requires_grad