Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8fa6ba6178 | ||
|
|
2e5fef787b | ||
|
|
43d87816bd | ||
|
|
2ace7dc6f4 | ||
|
|
0ca75db738 | ||
|
|
375ffd3fd5 | ||
|
|
2615ba4291 | ||
|
|
474dd71f28 | ||
|
|
98ad2d2db6 |
@@ -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
|
||||
@@ -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}"
|
||||
|
||||
@@ -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
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]:
|
||||
"""
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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, )
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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 []),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user