Compare commits
18
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2e5302817f | ||
|
|
33376b7347 | ||
|
|
d91c8c5541 | ||
|
|
96683c2fa9 | ||
|
|
907616e412 | ||
|
|
7e2c8f8d49 | ||
|
|
421dcbb50b | ||
|
|
b9662bf882 | ||
|
|
140fb9f20e | ||
|
|
2078876b98 | ||
|
|
a953f46bd6 | ||
|
|
adae957008 | ||
|
|
fa40553afb | ||
|
|
b93ef4289d | ||
|
|
401bdbd316 | ||
|
|
1048d79cf8 | ||
|
|
1e8406162d | ||
|
|
03edd35c83 |
+12
-1
@@ -198,4 +198,15 @@ steps:
|
||||
env:
|
||||
- TEST_TYPE=inference_vmoba
|
||||
agents:
|
||||
queue: "default"
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Unit Tests"
|
||||
env:
|
||||
- TEST_TYPE=unit_test
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
@@ -118,6 +118,10 @@ case "$TEST_TYPE" in
|
||||
log "Running V-MoBA precision tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
|
||||
;;
|
||||
"unit_test")
|
||||
log "Running unit tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
|
||||
;;
|
||||
*)
|
||||
log "Error: Unknown test type: $TEST_TYPE"
|
||||
exit 1
|
||||
|
||||
@@ -62,8 +62,8 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_nightly_test:
|
||||
description: "Run nightly-test"
|
||||
run_unit_test:
|
||||
description: "Run unit-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
@@ -93,6 +93,7 @@ jobs:
|
||||
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
|
||||
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
|
||||
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
|
||||
unit-test: ${{ steps.filter.outputs.unit-test }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
@@ -102,6 +103,8 @@ jobs:
|
||||
# Define reusable path patterns
|
||||
common-paths: &common-paths
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.10'
|
||||
- 'docker/Dockerfile.python3.11'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
sta-kernel-paths: &sta-kernel-paths
|
||||
- 'csrc/attn/sliding_tile_attn/**'
|
||||
@@ -155,6 +158,9 @@ jobs:
|
||||
precision-test-VSA:
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
unit-test:
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -333,23 +339,42 @@ jobs:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
nightly-test:
|
||||
unit-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.unit-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_unit_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "nightly-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
job_id: "unit-test"
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
# nightly-test:
|
||||
# if: >-
|
||||
# (github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
# uses: ./.github/workflows/runpod-test.yml
|
||||
# with:
|
||||
# job_id: "nightly-test"
|
||||
# gpu_type: "NVIDIA A40"
|
||||
# gpu_count: 4
|
||||
# volume_size: 100
|
||||
# disk_size: 100
|
||||
# image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
# test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
# timeout_minutes: 30
|
||||
# secrets:
|
||||
# RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
# RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
# WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
runpod-cleanup:
|
||||
# Add other jobs to this list as you create them
|
||||
|
||||
+3
-1
@@ -64,4 +64,6 @@ docs/source/distillation/examples/
|
||||
!docs/source/_static/images/**/*.png
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
dmd_t2v_output/
|
||||
|
||||
dmd_t2v_output/
|
||||
preprocess_output_text/
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
# VidProm Dataset
|
||||
|
||||
From [Self-Forcing](https://github.com/gdhe17/Self-Forcing) repository.
|
||||
|
||||
## Download the dataset
|
||||
|
||||
```bash
|
||||
./download_dataset.sh
|
||||
```
|
||||
@@ -0,0 +1,3 @@
|
||||
#! /bin/bash
|
||||
|
||||
huggingface-cli download gdhe17/Self-Forcing vidprom_filtered_extended.txt --local-dir prompts
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,76 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-056.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-059.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/EJqsC21GSBY-Scene-059.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/GBSfpTcKegk-Scene-003.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,13 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -16,23 +16,23 @@ export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29503
|
||||
export MASTER_PORT=29501
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="2f25ad37933894dbf0966c838c0b8494987f9f2f"
|
||||
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_MODE=offline
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=1
|
||||
NUM_GPUS=4
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation:
|
||||
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
DATA_DIR="data/mixkit-64_processed/Node_0_GPU_1_File_1/combined_parquet_dataset"
|
||||
VALIDATION_DATASET_FILE="data/mixkit-64_processed/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
@@ -84,7 +84,7 @@ dataset_args=(
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
@@ -1,151 +0,0 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:1
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29503
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="2f25ad37933894dbf0966c838c0b8494987f9f2f"
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
# Model paths for Self-Forcing DMD distillation with Wan2.2:
|
||||
GENERATOR_MODEL_PATH="Wan-AI/Wan2.2-T2V-A14B-Diffusers" # Updated to Wan2.2
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
DATA_DIR="data/test-text-preprocessing/Node_0_GPU_1_File_1/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name SFwan2.2_t2v_distill_self_forcing_dmd # Updated for Wan2.2
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan2.2_t2v_finetune"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 16
|
||||
--num_height 448 # Updated to match Wan2.2 config
|
||||
--num_width 832 # Updated to match Wan2.2 config
|
||||
--num_frames 61 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--log_visualization
|
||||
--simulate_generator_forward
|
||||
--num_frame_per_block 4 # Frame generation block size for self-forcing
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 16
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS # 64
|
||||
--sp_size 4
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--generator_model_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
|
||||
)
|
||||
|
||||
# Self-forcing DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.0
|
||||
--fake_score_learning_rate 8e-6
|
||||
--fake_score_betas '0.0,0.999'
|
||||
--warp_denoising_step
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks True # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
--validate_cache_structure False # Set to True for debugging KV cache issues
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--master_port $MASTER_PORT \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
@@ -9,33 +9,38 @@ def main():
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=4,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
|
||||
ti2v_task=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param.image_path = "test.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open."
|
||||
"A girl is packing a suitcase when stuff suddently starts flying around the room."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object")
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
# prompt2 = (
|
||||
# "A majestic lion strides across the golden savanna, its powerful frame "
|
||||
# "glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
# "the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
# "embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
# "cinematic.")
|
||||
# video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
@@ -19,13 +19,15 @@ def main():
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.num_frames = 81
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
prompts = [
|
||||
"A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
|
||||
"A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
|
||||
]
|
||||
|
||||
for prompt in prompts:
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,99 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
# DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_5/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing-extended-t2v-1-3b/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_70k"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "fixed_wan_ode_init_70k_6e-6"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--max_train_steps 6000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--warp_denoising_step
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 6e-6
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,135 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=1e5B2_16kFV_warp_ode_vidprom
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=ode_vidprom16k_warp/Dode_vidprom8b16k_1e-5.out
|
||||
#SBATCH --error=ode_vidprom16k_warp/Dode_vidprom8b16k_1e-5.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv2
|
||||
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
|
||||
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing-16k-t2v-1-3b-81/"
|
||||
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "Dwarp_vidprom_8b16k_test_warp_1e-5"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "Dwarp_vidprom_8b16k_wan_ode_init_1e-5"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--warp_denoising_step
|
||||
--log_visualization
|
||||
--max_train_steps 6001
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--dmd_denoising_steps "1000,750,500,250"
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim $NUM_GPUS
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 500
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,131 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=ode_vidprom2k
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=ode_vidprom2k_output/ode_vidprom2k.out
|
||||
#SBATCH --error=ode_vidprom2k_output/ode_vidprom2k.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv2
|
||||
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
|
||||
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing/"
|
||||
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_vidprom2k"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "vidprom2k_wan_ode_init_5e-6"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--max_train_steps 6001
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 8
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-6
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 2000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,132 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=ode_crush
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=ode_crush_output/ode_crush.out
|
||||
#SBATCH --error=ode_crush_output/ode_crush.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv2
|
||||
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
|
||||
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_5/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
|
||||
NUM_GPUS=2
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_warp_2"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "2warp_fixed_wan_ode_init_5e-6"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
# --warp_denoising_step
|
||||
--max_train_steps 6001
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim $NUM_GPUS
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 20
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-6
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 2000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,98 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="/mnt/weka/home/hao.zhang/wl/FastVideo2/data/crush-smol_processed_t2v_1_3b_ode_init_single"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "wan_ode_init_crush_smol"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "overfitwan_ode_init_crush_smol"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--max_train_steps 2001
|
||||
# --warp_denoising_step
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 20
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 500
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,100 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
# DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_5/combined_parquet_dataset/"
|
||||
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing-16k-t2v-1-3b/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_ode_init"
|
||||
--output_dir "debug_ode_init"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "debug_ode_init"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--max_train_steps 1000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81
|
||||
--warp_denoising_step
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 10
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ode_causal_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,24 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol_single/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_single/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 1 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 81 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 1 \
|
||||
--flush_frequency 1 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "ode_trajectory"
|
||||
@@ -0,0 +1,76 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Elon Musk, dressed in a sleek white spacesuit with a reflective visor, walks confidently across the lunar surface. His posture is upright, and he moves steadily with purpose. The moon's rocky terrain and scattered boulders surround him, casting shadows under the dim sunlight. The background shows vast stretches of the moon's barren landscape with craters and dust clouds kicked up by his boots. The scene captures a wide shot, emphasizing the vastness and desolation of the lunar environment. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "In a dynamic action-packed sequence set in the Marvel multiverse, Spider-Man and Venom engage in an intense battle. Spider-Man, in his classic red and blue suit, swings and dodges venomous attacks from the black symbiote-covered Venom. Both characters display a range of acrobatic moves and powerful strikes. The environment is a chaotic urban landscape with crumbling buildings and neon lights, reflecting the multiversal theme. The camera captures the epic fight from various angles, including wide shots to show the scale of destruction and close-ups to highlight their fierce expressions and physical combat. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A warm, family-oriented scene depicting a father getting ready to leave the house to buy milk. The father, a middle-aged man with a kind face and a casual outfit, picks up a jacket from the coat rack. His posture is upright as he bends down slightly to put on his shoes. In the background, there are glimpses of a cozy living room with a family photograph on the wall. The camera focuses closely on the father, capturing his gentle smile and reassuring nod towards the camera before he opens the front door and steps outside. Static medium close-up shot. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Close-up shot of a man with a prosthetic hand that functions as a rocket launcher. He looks at his new hand with a mix of amazement and concern, his facial expression showing a blend of curiosity and apprehension. The prosthetic hand is sleek and metallic, with intricate details that resemble a high-tech weapon. The background is a dimly lit laboratory with various scientific equipment and monitors displaying data. The man stands in a relaxed posture, his other hand resting on his hip, as he inspects his new limb. The scene is rendered in a realistic sci-fi style, emphasizing the futuristic technology and the man's emotional response to his new appendage. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Realistic CCTV footage style, Kim Taehyung from the band BTS is involved in a drug deal, caught on camera. Kim Taehyung appears nervous and cautious, wearing casual clothing typical of a public space. He exchanges items discreetly with another person, who is partially obscured. Both individuals maintain a watchful demeanor, occasionally glancing around to ensure no one is watching them. The lighting is dim, with flickering fluorescent lights casting shadows on their faces. The background shows a typical urban setting with blurred figures moving in the distance. Static camera angle, medium close-up shot focusing on the interaction between Taehyung and the other individual. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "Photorealistic studio setup with professional lighting, showcasing detailed cubic dissections of experimental plastic and felt-like materials on a pristine white background. Each cube reveals intricate layers and textures of the materials, emphasizing their unique properties. The scene has a shallow depth of field initially, then slowly pulls out to reveal the full arrangement of cubes, maintaining a wide depth of field throughout the transition. ",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -6,8 +6,8 @@ export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_old"
|
||||
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
NUM_GPUS=4
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
|
||||
@@ -52,7 +52,7 @@ dataset_args=(
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 200
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
@@ -4,7 +4,7 @@ GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v_old/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=2 # 2,4,8
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATASET_PATH="data/crush-smol/"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
@@ -14,7 +14,7 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--preprocess.dataset_type merged \
|
||||
--preprocess.dataset_path $DATASET_PATH \
|
||||
--preprocess.dataset_output_dir $OUTPUT_DIR \
|
||||
--preprocess.preprocess_video_batch_size 2 \
|
||||
--preprocess.preprocess_video_batch_size 8 \
|
||||
--preprocess.dataloader_num_workers 0 \
|
||||
--preprocess.max_height 480 \
|
||||
--preprocess.max_width 832 \
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_old"
|
||||
VALIDATION_DATASET_FILE="examples/datasets/crush_smol/validation.json"
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_t2v_i2v_finetune"
|
||||
--output_dir "checkpoints/wan_t2v_i2v_finetune"
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 2
|
||||
--num_latent_t 20
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 4
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 2
|
||||
--hsdp_shard_dim 4
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path $DATA_DIR
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--t2v_as_i2v_task True
|
||||
# --resume_from_checkpoint "checkpoints/wan_t2v_finetune/checkpoint-2500"
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_t2v_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,24 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v_i2v_1_3b/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 2 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v_ode_trajectory"
|
||||
@@ -92,6 +92,9 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
pos_embed_seq_len: int | None = None
|
||||
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
|
||||
|
||||
# Wan MoE
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Causal Wan
|
||||
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
|
||||
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
|
||||
|
||||
@@ -45,6 +45,8 @@ class PipelineConfig:
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: float | None = None
|
||||
disable_autocast: bool = False
|
||||
ti2v_task: bool = False
|
||||
t2v_as_i2v_task: bool = False
|
||||
|
||||
# Model configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
@@ -85,9 +87,6 @@ class PipelineConfig:
|
||||
# DMD parameters
|
||||
dmd_denoising_steps: list[int] | None = field(default=None)
|
||||
|
||||
# Wan2.2 TI2V parameters
|
||||
ti2v_task: bool = False
|
||||
|
||||
# Compilation
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
@@ -214,6 +213,24 @@ class PipelineConfig:
|
||||
"Comma-separated list of denoising steps (e.g., '1000,757,522')",
|
||||
)
|
||||
|
||||
# TI2V task
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}ti2v-task",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}ti2v_task",
|
||||
default=PipelineConfig.ti2v_task,
|
||||
help="Enable TI2V",
|
||||
)
|
||||
|
||||
# T2V to I2V task
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}t2v-as-i2v-task",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}t2v_as_i2v_task",
|
||||
default=PipelineConfig.t2v_as_i2v_task,
|
||||
help="Enable T2V to I2V task",
|
||||
)
|
||||
|
||||
# Add VAE configuration arguments
|
||||
from fastvideo.configs.models.vaes.base import VAEConfig
|
||||
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
|
||||
@@ -245,7 +262,9 @@ class PipelineConfig:
|
||||
"""
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
logger.info("WTF model_path: %s", model_path)
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
logger.info("pipeline_config_cls: %s", pipeline_config_cls)
|
||||
|
||||
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
|
||||
|
||||
|
||||
@@ -13,7 +13,7 @@ from fastvideo.configs.pipelines.wan import (
|
||||
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
|
||||
SelfForcingWanT2V480PConfig, Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config,
|
||||
Wan2_2_TI2V_5B_Config, WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig,
|
||||
WanT2V720PConfig)
|
||||
WanT2V720PConfig, SelfForcingWanT2V480PConfig)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
@@ -49,6 +49,7 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
|
||||
"stepvideo": lambda id: "stepvideo" in id.lower(),
|
||||
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
@@ -60,7 +61,8 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
|
||||
"stepvideo": StepVideoT2VConfig
|
||||
"stepvideo": StepVideoT2VConfig,
|
||||
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
@@ -82,7 +82,7 @@ class WanI2V480PConfig(WanT2V480PConfig):
|
||||
default_factory=CLIPVisionConfig)
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
def __post_init__(self):
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
@@ -108,19 +108,17 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 757, 522])
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
|
||||
flow_shift: float | None = 5.0
|
||||
ti2v_task: bool = True
|
||||
expand_timesteps: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
self.dit_config.expand_timesteps = self.expand_timesteps
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -132,12 +130,21 @@ class FastWan2_2_TI2V_5B_Config(Wan2_2_TI2V_5B_Config):
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
|
||||
pass
|
||||
flow_shift: float | None = 12.0
|
||||
boundary_ratio: float | None = 0.875
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.dit_config.boundary_ratio = self.boundary_ratio
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
|
||||
pass
|
||||
class Wan2_2_I2V_A14B_Config(WanI2V480PConfig):
|
||||
flow_shift: float | None = 5.0
|
||||
boundary_ratio: float | None = 0.900
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.dit_config.boundary_ratio = self.boundary_ratio
|
||||
|
||||
|
||||
# =============================================
|
||||
|
||||
@@ -40,6 +40,7 @@ class SamplingParam:
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
@@ -47,6 +48,8 @@ class SamplingParam:
|
||||
# Misc
|
||||
save_video: bool = True
|
||||
return_frames: bool = False
|
||||
return_trajectory_latents: bool = False # returns all latents for each timestep
|
||||
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.data_type = "video" if self.num_frames > 1 else "image"
|
||||
@@ -167,6 +170,12 @@ class SamplingParam:
|
||||
default=SamplingParam.guidance_rescale,
|
||||
help="Guidance rescale factor",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--boundary-ratio",
|
||||
type=float,
|
||||
default=SamplingParam.boundary_ratio,
|
||||
help="Boundary timestep ratio",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save-video",
|
||||
action="store_true",
|
||||
@@ -198,6 +207,18 @@ class SamplingParam:
|
||||
help=
|
||||
"Path to a JSON file containing V-MoBA specific configurations.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--return-trajectory-latents",
|
||||
action="store_true",
|
||||
default=SamplingParam.return_trajectory_latents,
|
||||
help="Whether to return the trajectory",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--return-trajectory-decoded",
|
||||
action="store_true",
|
||||
default=SamplingParam.return_trajectory_decoded,
|
||||
help="Whether to return the decoded trajectory",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
@@ -144,18 +144,22 @@ class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
guidance_scale: float = 4.0
|
||||
guidance_scale_2: float = 3.0
|
||||
guidance_scale: float = 4.0 # high_noise
|
||||
guidance_scale_2: float = 3.0 # low_noise
|
||||
num_inference_steps: int = 40
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
|
||||
|
||||
@dataclass
|
||||
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
|
||||
guidance_scale: float = 3.5
|
||||
guidance_scale_2: float = 3.5
|
||||
guidance_scale: float = 3.5 # high_noise
|
||||
guidance_scale_2: float = 3.5 # low_noise
|
||||
num_inference_steps: int = 40
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
|
||||
|
||||
# =============================================
|
||||
|
||||
@@ -4,7 +4,7 @@ from torchvision.transforms import Lambda
|
||||
|
||||
from fastvideo.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset
|
||||
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset, TextDataset
|
||||
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
from fastvideo.dataset.validation_dataset import ValidationDataset
|
||||
@@ -39,7 +39,13 @@ def getdataset(args) -> VideoCaptionMergedDataset:
|
||||
seed=args.seed)
|
||||
|
||||
|
||||
def gettextdataset(args) -> TextDataset:
|
||||
return TextDataset(data_merge_path=args.data_merge_path,
|
||||
args=args,
|
||||
seed=args.seed)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_parquet_map_style_dataloader", "ValidationDataset",
|
||||
"VideoCaptionMergedDataset"
|
||||
"VideoCaptionMergedDataset", "TextDataset"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,264 @@
|
||||
"""
|
||||
Utilities for converting preprocessing records (dicts) into Arrow tables and
|
||||
writing Parquet datasets in fixed-size chunks.
|
||||
|
||||
This module centralizes table construction and Parquet file writing so
|
||||
pipelines only need to define their PyArrow schema and produce per-sample
|
||||
record dictionaries.
|
||||
|
||||
Key APIs:
|
||||
- records_to_table(records, schema): Safely convert a list of dictionaries into
|
||||
a pa.Table, casting to the provided schema.
|
||||
- ParquetDatasetWriter: Buffer tables and flush to a directory as multiple
|
||||
Parquet files with a fixed number of rows per file. Uses temporary files and
|
||||
atomic rename to avoid partially written outputs.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import multiprocessing
|
||||
import os
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from typing import Any
|
||||
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
|
||||
def records_to_table(records: list[dict[str, Any]], schema: pa.Schema) -> pa.Table:
|
||||
"""Build a PyArrow table from Python record dicts using an explicit schema.
|
||||
|
||||
Arrow will cast values to the target schema when possible (e.g., promoting
|
||||
Python ints/floats to pa.int64/pa.float64), eliminating hand-written per-
|
||||
field array construction.
|
||||
|
||||
Args:
|
||||
records: List of dictionaries, each representing one row. Keys must
|
||||
match schema field names.
|
||||
schema: Target PyArrow schema. Controls field names and types.
|
||||
|
||||
Returns:
|
||||
pa.Table: In-memory table matching the provided schema. If ``records``
|
||||
is empty, returns an empty table with the given schema.
|
||||
"""
|
||||
if not records:
|
||||
return pa.table({}, schema=schema)
|
||||
return pa.Table.from_pylist(records, schema=schema)
|
||||
|
||||
|
||||
class ParquetDatasetWriter:
|
||||
"""Accumulate tables and flush them to a Parquet directory in fixed-size chunks.
|
||||
|
||||
Behavior:
|
||||
- Writes files under worker-specific subdirectories for parallelism.
|
||||
- Uses temporary files and atomic rename to avoid partial files being left
|
||||
behind on failure.
|
||||
- Only full chunks of ``samples_per_file`` rows are written on each flush;
|
||||
any remainder rows are re-buffered for the next flush.
|
||||
|
||||
Note:
|
||||
- Instances are not meant to be shared across processes. Create one writer
|
||||
per process if using multiprocessing.
|
||||
"""
|
||||
|
||||
def __init__(self, out_dir: str, samples_per_file: int, compression: str = "zstd") -> None:
|
||||
"""Initialize the dataset writer.
|
||||
|
||||
Args:
|
||||
out_dir: Output directory where Parquet files will be written.
|
||||
samples_per_file: Fixed number of rows per Parquet file.
|
||||
compression: Compression codec passed to ``pyarrow.parquet.write_table``
|
||||
(e.g., ``"zstd"``, ``"snappy"``, ``"gzip"``).
|
||||
"""
|
||||
self.out_dir = out_dir
|
||||
self.samples_per_file = max(int(samples_per_file), 1)
|
||||
self.compression = compression
|
||||
os.makedirs(self.out_dir, exist_ok=True)
|
||||
self._tables: list[pa.Table] = []
|
||||
|
||||
def append_table(self, table: pa.Table) -> None:
|
||||
"""Append a non-empty table to the internal buffer.
|
||||
|
||||
Args:
|
||||
table: A ``pa.Table`` to buffer. Empty or ``None`` tables are ignored.
|
||||
"""
|
||||
if table is None or len(table) == 0:
|
||||
return
|
||||
self._tables.append(table)
|
||||
|
||||
def _combine(self) -> pa.Table | None:
|
||||
"""Combine all buffered tables into a single table, if any.
|
||||
|
||||
Returns:
|
||||
A concatenated table, a single table if only one was buffered, or
|
||||
``None`` if no tables are buffered.
|
||||
"""
|
||||
if not self._tables:
|
||||
return None
|
||||
if len(self._tables) == 1:
|
||||
return self._tables[0]
|
||||
return pa.concat_tables(self._tables, promote_options='none')
|
||||
|
||||
def flush(self, num_workers: int | None = None, write_remainder: bool = False) -> int:
|
||||
"""Write accumulated tables to disk and clear the written portion.
|
||||
|
||||
Only complete chunks of size ``samples_per_file`` are written. Any
|
||||
remainder rows are kept buffered for the next flush.
|
||||
|
||||
Args:
|
||||
num_workers: Optional override for the number of parallel workers
|
||||
used to write chunks. Defaults to ``min(cpu_count, chunks)``.
|
||||
write_remainder: If True, also write any leftover rows (< samples_per_file)
|
||||
as a final small Parquet file (useful for the last flush at the
|
||||
end of preprocessing).
|
||||
|
||||
Returns:
|
||||
int: Number of rows successfully written in this flush call.
|
||||
"""
|
||||
combined = self._combine()
|
||||
self._tables = []
|
||||
if combined is None or len(combined) == 0:
|
||||
return 0
|
||||
|
||||
num_samples = len(combined)
|
||||
total_chunks = num_samples // self.samples_per_file
|
||||
if total_chunks == 0:
|
||||
if not write_remainder:
|
||||
# Not enough to form a full chunk; keep buffered for next round
|
||||
# Re-buffer and return 0 written
|
||||
self._tables = [combined]
|
||||
return 0
|
||||
# Last flush: write the small remainder as a final file in worker_0
|
||||
worker_dir = os.path.join(self.out_dir, "worker_0")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
# Determine next index
|
||||
num_parquets = 0
|
||||
for _, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
chunk_path = os.path.join(worker_dir, f"data_chunk_{num_parquets}.parquet")
|
||||
temp_path = chunk_path + '.tmp'
|
||||
pq.write_table(combined, temp_path, compression=self.compression)
|
||||
if os.path.exists(chunk_path):
|
||||
os.remove(chunk_path)
|
||||
os.rename(temp_path, chunk_path)
|
||||
return num_samples
|
||||
|
||||
# Only write full chunks; keep remainder for next flush
|
||||
written_rows = total_chunks * self.samples_per_file
|
||||
remainder = num_samples - written_rows
|
||||
|
||||
table_to_write = combined.slice(0, written_rows)
|
||||
remainder_table = combined.slice(written_rows, remainder) if remainder > 0 else None
|
||||
if remainder_table is not None and len(remainder_table) > 0:
|
||||
if write_remainder:
|
||||
# Write the remainder as a final small file (worker_0)
|
||||
worker_dir = os.path.join(self.out_dir, "worker_0")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
num_parquets = 0
|
||||
for _, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
remainder_path = os.path.join(worker_dir,
|
||||
f"data_chunk_{num_parquets}.parquet")
|
||||
temp_path = remainder_path + '.tmp'
|
||||
pq.write_table(remainder_table,
|
||||
temp_path,
|
||||
compression=self.compression)
|
||||
if os.path.exists(remainder_path):
|
||||
os.remove(remainder_path)
|
||||
os.rename(temp_path, remainder_path)
|
||||
else:
|
||||
self._tables = [remainder_table]
|
||||
|
||||
# Parallel write by chunk ranges
|
||||
if num_workers is None:
|
||||
num_workers = min(multiprocessing.cpu_count(), max(total_chunks, 1))
|
||||
num_workers = max(int(num_workers), 1)
|
||||
chunks_per_worker = (total_chunks + num_workers - 1) // num_workers
|
||||
|
||||
work_ranges: list[tuple[int, int, pa.Table, int, str, int, str]] = []
|
||||
for worker_id in range(num_workers):
|
||||
start_chunk = worker_id * chunks_per_worker
|
||||
end_chunk = min((worker_id + 1) * chunks_per_worker, total_chunks)
|
||||
if start_chunk < end_chunk:
|
||||
work_ranges.append(
|
||||
(
|
||||
start_chunk,
|
||||
end_chunk,
|
||||
table_to_write,
|
||||
worker_id,
|
||||
self.out_dir,
|
||||
self.samples_per_file,
|
||||
self.compression,
|
||||
)
|
||||
)
|
||||
|
||||
written_total = 0
|
||||
if len(work_ranges) == 1:
|
||||
written_total += _process_chunk_range(work_ranges[0])
|
||||
return written_total
|
||||
|
||||
with ProcessPoolExecutor(max_workers=num_workers) as executor:
|
||||
futures = [executor.submit(_process_chunk_range, args) for args in work_ranges]
|
||||
for f in futures:
|
||||
written_total += f.result()
|
||||
return written_total + (len(remainder_table) if write_remainder and remainder_table is not None else 0)
|
||||
|
||||
|
||||
def _process_chunk_range(args: Any) -> int:
|
||||
"""Worker function to write a contiguous range of chunk files.
|
||||
|
||||
Args:
|
||||
args: Tuple containing
|
||||
- start_chunk (int): inclusive start chunk index
|
||||
- end_chunk (int): exclusive end chunk index
|
||||
- table (pa.Table): concatenated table containing all rows to write
|
||||
- worker_id (int): numeric worker identifier
|
||||
- output_dir (str): base output directory
|
||||
- samples_per_file (int): rows per chunk file
|
||||
- compression (str): compression codec for Parquet
|
||||
|
||||
Returns:
|
||||
int: Total number of rows written by this worker.
|
||||
"""
|
||||
start_chunk, end_chunk, table, worker_id, output_dir, samples_per_file, compression = args
|
||||
total_written = 0
|
||||
num_samples = len(table)
|
||||
|
||||
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
|
||||
# Offset to continue numbering if files exist
|
||||
num_parquets = 0
|
||||
for root, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
|
||||
for i in range(start_chunk, end_chunk):
|
||||
start_sample = i * samples_per_file
|
||||
end_sample = min((i + 1) * samples_per_file, num_samples)
|
||||
if end_sample <= start_sample:
|
||||
continue
|
||||
chunk = table.slice(start_sample, end_sample - start_sample)
|
||||
|
||||
chunk_path = os.path.join(worker_dir, f"data_chunk_{i + num_parquets}.parquet")
|
||||
temp_path = chunk_path + '.tmp'
|
||||
try:
|
||||
pq.write_table(chunk, temp_path, compression=compression)
|
||||
if os.path.exists(chunk_path):
|
||||
os.remove(chunk_path)
|
||||
os.rename(temp_path, chunk_path)
|
||||
total_written += len(chunk)
|
||||
except Exception:
|
||||
if os.path.exists(temp_path):
|
||||
os.remove(temp_path)
|
||||
raise
|
||||
|
||||
return total_written
|
||||
|
||||
|
||||
|
||||
+68
@@ -1,5 +1,7 @@
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
|
||||
|
||||
|
||||
@@ -120,3 +122,69 @@ def i2v_record_creator(batch: PreprocessBatch) -> list[dict[str, Any]]:
|
||||
})
|
||||
|
||||
return records
|
||||
|
||||
|
||||
def ode_text_only_record_creator(
|
||||
video_name: str, text_embedding: np.ndarray, caption: str,
|
||||
trajectory_latents: np.ndarray,
|
||||
trajectory_timesteps: np.ndarray) -> dict[str, Any]:
|
||||
"""Create a text-only ODE trajectory record matching pyarrow_schema_ode_trajectory_text_only.
|
||||
|
||||
Args:
|
||||
video_name: Base name/id for the sample (without extension).
|
||||
text_embedding: Text encoder output array [SeqLen, Dim].
|
||||
caption: Original text prompt.
|
||||
trajectory_latents: Collected trajectory latents array.
|
||||
trajectory_timesteps: Collected timesteps array.
|
||||
|
||||
Returns:
|
||||
dict suitable for records_to_table(…, pyarrow_schema_ode_trajectory_text_only)
|
||||
"""
|
||||
assert trajectory_latents is not None, "trajectory_latents is required"
|
||||
assert trajectory_timesteps is not None, "trajectory_timesteps is required"
|
||||
|
||||
record = {
|
||||
"id": f"text_{video_name}",
|
||||
"text_embedding_bytes": text_embedding.tobytes(),
|
||||
"text_embedding_shape": list(text_embedding.shape),
|
||||
"text_embedding_dtype": str(text_embedding.dtype),
|
||||
"file_name": video_name,
|
||||
"caption": caption,
|
||||
"media_type": "text",
|
||||
}
|
||||
|
||||
record.update({
|
||||
"trajectory_latents_bytes": trajectory_latents.tobytes(),
|
||||
"trajectory_latents_shape": list(trajectory_latents.shape),
|
||||
"trajectory_latents_dtype": str(trajectory_latents.dtype),
|
||||
})
|
||||
|
||||
record.update({
|
||||
"trajectory_timesteps_bytes": trajectory_timesteps.tobytes(),
|
||||
"trajectory_timesteps_shape": list(trajectory_timesteps.shape),
|
||||
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
|
||||
})
|
||||
|
||||
return record
|
||||
|
||||
|
||||
def text_only_record_creator(text_name: str, text_embedding: np.ndarray,
|
||||
caption: str) -> dict[str, Any]:
|
||||
"""Create a text-only record matching pyarrow_schema_text_only.
|
||||
|
||||
Args:
|
||||
text_name: Base id/name for the text sample.
|
||||
text_embedding: Text encoder output array [SeqLen, Dim].
|
||||
caption: Original text prompt.
|
||||
|
||||
Returns:
|
||||
dict suitable for records_to_table(…, pyarrow_schema_text_only)
|
||||
"""
|
||||
record = {
|
||||
"id": f"text_{text_name}",
|
||||
"text_embedding_bytes": text_embedding.tobytes(),
|
||||
"text_embedding_shape": list(text_embedding.shape),
|
||||
"text_embedding_dtype": str(text_embedding.dtype),
|
||||
"caption": caption,
|
||||
}
|
||||
return record
|
||||
@@ -50,6 +50,7 @@ pyarrow_schema_i2v = pa.schema([
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
|
||||
pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Image/Video VAE latents ---
|
||||
@@ -78,3 +79,40 @@ pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
|
||||
pyarrow_schema_ode_trajectory_text_only = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
# --- ODE Trajectory ---
|
||||
pa.field("trajectory_latents_bytes", pa.binary()),
|
||||
pa.field("trajectory_latents_shape", pa.list_(pa.int64())),
|
||||
pa.field("trajectory_latents_dtype", pa.string()),
|
||||
pa.field("trajectory_timesteps_bytes", pa.binary()),
|
||||
pa.field("trajectory_timesteps_shape", pa.list_(pa.int64())),
|
||||
pa.field("trajectory_timesteps_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # Always 'text' for text-only
|
||||
])
|
||||
|
||||
|
||||
pyarrow_schema_text_only = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("caption", pa.string()),
|
||||
])
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.dataset.lmdb_utils import get_array_shape_from_lmdb, retrieve_row_from_lmdb
|
||||
from torch.utils.data import Dataset
|
||||
import numpy as np
|
||||
import torch
|
||||
import lmdb
|
||||
|
||||
# from Self-Forcing: https://github.com/guandeh17/Self-Forcing/blob/main/utils/dataset.py
|
||||
class ODERegressionLMDBDataset(Dataset):
|
||||
def __init__(self, data_path: str, max_pair: int = int(1e8)):
|
||||
print(f"data_path: {data_path}")
|
||||
self.env = lmdb.open(data_path, readonly=True,
|
||||
lock=False, readahead=False, meminit=False)
|
||||
|
||||
self.latents_shape = get_array_shape_from_lmdb(self.env, 'latents')
|
||||
self.max_pair = max_pair
|
||||
|
||||
def __len__(self):
|
||||
return min(self.latents_shape[0], self.max_pair)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
"""
|
||||
Outputs:
|
||||
- prompts: List of Strings
|
||||
- latents: Tensor of shape (num_denoising_steps, num_frames, num_channels, height, width). It is ordered from pure noise to clean image.
|
||||
"""
|
||||
latents = retrieve_row_from_lmdb(
|
||||
self.env,
|
||||
"latents", np.float16, idx, shape=self.latents_shape[1:]
|
||||
)
|
||||
|
||||
if len(latents.shape) == 4:
|
||||
latents = latents[None, ...]
|
||||
|
||||
prompts = retrieve_row_from_lmdb(
|
||||
self.env,
|
||||
"prompts", str, idx
|
||||
)
|
||||
return {
|
||||
"prompts": prompts,
|
||||
"ode_latent": torch.tensor(latents, dtype=torch.float32)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# from Self-Forcing: https://github.com/guandeh17/Self-Forcing/blob/main/utils/lmdb.py
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def get_array_shape_from_lmdb(env, array_name):
|
||||
with env.begin() as txn:
|
||||
image_shape = txn.get(f"{array_name}_shape".encode()).decode()
|
||||
image_shape = tuple(map(int, image_shape.split()))
|
||||
return image_shape
|
||||
|
||||
|
||||
def store_arrays_to_lmdb(env, arrays_dict, start_index=0):
|
||||
"""
|
||||
Store rows of multiple numpy arrays in a single LMDB.
|
||||
Each row is stored separately with a naming convention.
|
||||
"""
|
||||
with env.begin(write=True) as txn:
|
||||
for array_name, array in arrays_dict.items():
|
||||
for i, row in enumerate(array):
|
||||
# Convert row to bytes
|
||||
if isinstance(row, str):
|
||||
row_bytes = row.encode()
|
||||
else:
|
||||
row_bytes = row.tobytes()
|
||||
|
||||
data_key = f'{array_name}_{start_index + i}_data'.encode()
|
||||
|
||||
txn.put(data_key, row_bytes)
|
||||
|
||||
|
||||
def process_data_dict(data_dict, seen_prompts):
|
||||
output_dict = {}
|
||||
|
||||
all_videos = []
|
||||
all_prompts = []
|
||||
for prompt, video in data_dict.items():
|
||||
if prompt in seen_prompts:
|
||||
continue
|
||||
else:
|
||||
seen_prompts.add(prompt)
|
||||
|
||||
video = video.half().numpy()
|
||||
all_videos.append(video)
|
||||
all_prompts.append(prompt)
|
||||
|
||||
if len(all_videos) == 0:
|
||||
return {"latents": np.array([]), "prompts": np.array([])}
|
||||
|
||||
all_videos = np.concatenate(all_videos, axis=0)
|
||||
|
||||
output_dict['latents'] = all_videos
|
||||
output_dict['prompts'] = np.array(all_prompts)
|
||||
|
||||
return output_dict
|
||||
|
||||
|
||||
def retrieve_row_from_lmdb(lmdb_env, array_name, dtype, row_index, shape=None):
|
||||
"""
|
||||
Retrieve a specific row from a specific array in the LMDB.
|
||||
"""
|
||||
data_key = f'{array_name}_{row_index}_data'.encode()
|
||||
|
||||
with lmdb_env.begin() as txn:
|
||||
row_bytes = txn.get(data_key)
|
||||
|
||||
if dtype == str:
|
||||
array = row_bytes.decode()
|
||||
else:
|
||||
array = np.frombuffer(row_bytes, dtype=dtype)
|
||||
|
||||
if shape is not None and len(shape) > 0:
|
||||
array = array.reshape(shape)
|
||||
return array
|
||||
@@ -628,3 +628,134 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
|
||||
"""Load state dict from checkpoint."""
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
|
||||
|
||||
class TextDataset(torch.utils.data.IterableDataset,
|
||||
torch.distributed.checkpoint.stateful.Stateful):
|
||||
"""
|
||||
Text-only dataset for processing prompts from a simple text file.
|
||||
|
||||
Assumes that data_merge_path is a text file with one prompt per line:
|
||||
A cat playing with a ball
|
||||
A dog running in the park
|
||||
A person cooking dinner
|
||||
...
|
||||
|
||||
This dataset processes text data through text encoding stages only.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
data_merge_path: str,
|
||||
args,
|
||||
start_idx: int = 0,
|
||||
seed: int = 42):
|
||||
self.data_merge_path = data_merge_path
|
||||
self.start_idx = start_idx
|
||||
self.args = args
|
||||
self.seed = seed
|
||||
|
||||
# Initialize tokenizer
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
|
||||
# Initialize text encoding stage
|
||||
self.text_encoding_stage = TextEncodingStage(
|
||||
tokenizer=tokenizer,
|
||||
text_max_length=args.text_max_length,
|
||||
cfg_rate=getattr(args, 'training_cfg_rate', 0.0),
|
||||
seed=self.seed)
|
||||
|
||||
# Process text data
|
||||
self.processed_batches = self._process_text_data()
|
||||
|
||||
def _load_text_data(self) -> list[str]:
|
||||
"""Load text prompts from file."""
|
||||
prompts = []
|
||||
with open(self.data_merge_path, 'r', encoding='utf-8') as f:
|
||||
for line in f:
|
||||
line = line.strip()
|
||||
if line: # Skip empty lines
|
||||
prompts.append(line)
|
||||
|
||||
logger.info(f"Loaded {len(prompts)} text prompts from {self.data_merge_path}")
|
||||
return prompts
|
||||
|
||||
def _process_text_data(self) -> list[PreprocessBatch]:
|
||||
"""Process the text prompts through text encoding stage."""
|
||||
raw_prompts = self._load_text_data()
|
||||
processed_batches = []
|
||||
|
||||
for idx, prompt in enumerate(raw_prompts):
|
||||
# Create a text-only batch with dummy path
|
||||
batch = PreprocessBatch(
|
||||
path=f"text_prompt_{idx}",
|
||||
cap=[prompt], # TextEncodingStage expects a list
|
||||
resolution=None,
|
||||
fps=None,
|
||||
duration=None,
|
||||
num_frames=0,
|
||||
sample_frame_index=None,
|
||||
sample_num_frames=0
|
||||
)
|
||||
|
||||
processed_batches.append(batch)
|
||||
|
||||
logger.info(f"Processed {len(processed_batches)} text batches")
|
||||
return processed_batches
|
||||
|
||||
def __iter__(self):
|
||||
"""Iterator for the dataset."""
|
||||
# Set up distributed sampling if needed
|
||||
if torch.distributed.is_available() and torch.distributed.is_initialized():
|
||||
rank = torch.distributed.get_rank()
|
||||
world_size = torch.distributed.get_world_size()
|
||||
else:
|
||||
rank = 0
|
||||
world_size = 1
|
||||
|
||||
# Calculate chunk for this rank
|
||||
total_items = len(self.processed_batches)
|
||||
items_per_rank = math.ceil(total_items / world_size)
|
||||
start_idx = rank * items_per_rank + self.start_idx
|
||||
end_idx = min(start_idx + items_per_rank, total_items)
|
||||
|
||||
# Yield items for this rank
|
||||
for idx in range(start_idx, end_idx):
|
||||
if idx < len(self.processed_batches):
|
||||
yield self._get_item(idx)
|
||||
|
||||
def _get_item(self, idx: int) -> dict:
|
||||
"""Get a single processed text item."""
|
||||
batch = self.processed_batches[idx]
|
||||
|
||||
# Apply text encoding stage
|
||||
batch = self.text_encoding_stage.process(batch)
|
||||
|
||||
# Build result dictionary for text-only processing with required schema fields
|
||||
result = {
|
||||
"text": batch.text,
|
||||
"input_ids": batch.input_ids,
|
||||
"cond_mask": batch.cond_mask,
|
||||
"path": batch.path,
|
||||
# Required schema fields for ODE trajectory processing
|
||||
"id": f"text_{idx}",
|
||||
"file_name": batch.path,
|
||||
"caption": batch.text,
|
||||
"media_type": "text",
|
||||
"width": 1,
|
||||
"height": 1,
|
||||
"num_frames": 0,
|
||||
"duration_sec": 0.0,
|
||||
"fps": 0.0,
|
||||
}
|
||||
|
||||
return result
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
"""Return state dict for checkpointing."""
|
||||
return {"processed_batches": self.processed_batches}
|
||||
|
||||
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
|
||||
"""Load state dict from checkpoint."""
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
|
||||
@@ -3,9 +3,12 @@ from typing import Any, cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
|
||||
def pad(t: torch.Tensor, padding_length: int) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Pad or crop an embedding [L, D] to exactly padding_length tokens.
|
||||
Return:
|
||||
|
||||
@@ -344,6 +344,9 @@ class VideoGenerator:
|
||||
"size": (target_height, target_width, batch.num_frames),
|
||||
"generation_time": gen_time,
|
||||
"logging_info": logging_info,
|
||||
"trajectory": output_batch.trajectory_latents,
|
||||
"trajectory_timesteps": output_batch.trajectory_timesteps,
|
||||
"trajectory_decoded": output_batch.trajectory_decoded,
|
||||
}
|
||||
|
||||
def set_lora_adapter(self,
|
||||
|
||||
@@ -158,6 +158,7 @@ class FastVideoArgs:
|
||||
"transformer": True,
|
||||
"vae": True,
|
||||
})
|
||||
override_transformer_cls_name: str | None = None
|
||||
|
||||
# # DMD parameters
|
||||
# dmd_denoising_steps: List[int] | None = field(default=None)
|
||||
@@ -396,6 +397,12 @@ class FastVideoArgs:
|
||||
default=FastVideoArgs.enable_stage_verification,
|
||||
help="Enable input/output verification for pipeline stages",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--override-transformer-cls-name",
|
||||
type=str,
|
||||
default=FastVideoArgs.override_transformer_cls_name,
|
||||
help="Override transformer cls name",
|
||||
)
|
||||
# Add pipeline configuration arguments
|
||||
PipelineConfig.add_cli_args(parser)
|
||||
|
||||
@@ -698,6 +705,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
# simulate generator forward to match inference
|
||||
simulate_generator_forward: bool = False
|
||||
warp_denoising_step: bool = False
|
||||
intermediate_latents_visualization: bool = False
|
||||
|
||||
# Self-forcing specific arguments
|
||||
num_frame_per_block: int = 3
|
||||
@@ -1133,11 +1141,10 @@ class TrainingArgs(FastVideoArgs):
|
||||
"--last-step-only",
|
||||
action=StoreBoolean,
|
||||
help="Whether to only use the last timestep for training")
|
||||
parser.add_argument(
|
||||
"--context-noise",
|
||||
type=int,
|
||||
default=TrainingArgs.context_noise,
|
||||
help="Context noise level for cache updates")
|
||||
parser.add_argument("--context-noise",
|
||||
type=int,
|
||||
default=TrainingArgs.context_noise,
|
||||
help="Context noise level for cache updates")
|
||||
|
||||
return parser
|
||||
|
||||
@@ -1145,4 +1152,4 @@ class TrainingArgs(FastVideoArgs):
|
||||
def parse_int_list(value: str) -> list[int]:
|
||||
if not value:
|
||||
return []
|
||||
return [int(x.strip()) for x in value.split(",")]
|
||||
return [int(x.strip()) for x in value.split(",")]
|
||||
|
||||
@@ -276,4 +276,4 @@ class LayerNormScaleShift(nn.Module):
|
||||
if self.compute_dtype == torch.float32:
|
||||
output = output.to(x.dtype)
|
||||
|
||||
return output
|
||||
return output
|
||||
|
||||
@@ -77,9 +77,11 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
lora_A = self.lora_A.to_local()
|
||||
|
||||
if not self.merged and not self.disable_lora:
|
||||
delta = x @ (
|
||||
self.slice_lora_b_weights(lora_B.to(x, non_blocking=True))
|
||||
@ self.slice_lora_a_weights(lora_A.to(x, non_blocking=True)))
|
||||
lora_A_sliced = self.slice_lora_a_weights(
|
||||
lora_A.to(x, non_blocking=True))
|
||||
lora_B_sliced = self.slice_lora_b_weights(
|
||||
lora_B.to(x, non_blocking=True))
|
||||
delta = x @ lora_A_sliced.T @ lora_B_sliced.T
|
||||
if self.lora_alpha != self.lora_rank:
|
||||
delta = delta * (
|
||||
self.lora_alpha / self.lora_rank # type: ignore
|
||||
|
||||
@@ -679,4 +679,4 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
return out
|
||||
return out
|
||||
|
||||
@@ -415,6 +415,10 @@ class TransformerLoader(ComponentLoader):
|
||||
raise ValueError(
|
||||
"Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported.")
|
||||
logger.info("transformer cls_name: %s", cls_name)
|
||||
if fastvideo_args.override_transformer_cls_name is not None:
|
||||
cls_name = fastvideo_args.override_transformer_cls_name
|
||||
logger.info("Overriding transformer cls_name to %s", cls_name)
|
||||
|
||||
fastvideo_args.model_paths["transformer"] = model_path
|
||||
|
||||
@@ -460,6 +464,7 @@ class TransformerLoader(ComponentLoader):
|
||||
device=get_local_torch_device(),
|
||||
hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim,
|
||||
hsdp_shard_dim=fastvideo_args.hsdp_shard_dim,
|
||||
default_dtype=default_dtype,
|
||||
cpu_offload=fastvideo_args.dit_cpu_offload,
|
||||
pin_cpu_memory=fastvideo_args.pin_cpu_memory,
|
||||
fsdp_inference=fastvideo_args.use_fsdp_inference,
|
||||
@@ -472,10 +477,11 @@ class TransformerLoader(ComponentLoader):
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||
for param in model.parameters():
|
||||
logger.info("Param dtype: %s", param.dtype)
|
||||
|
||||
dtypes = set(param.dtype for param in model.parameters())
|
||||
if len(dtypes) > 1:
|
||||
model = model.to(default_dtype)
|
||||
logger.info("Converting model to dtype: %s", default_dtype)
|
||||
model = model.to(default_dtype)
|
||||
model = model.eval()
|
||||
return model
|
||||
|
||||
|
||||
@@ -62,6 +62,7 @@ def maybe_load_fsdp_model(
|
||||
device: torch.device,
|
||||
hsdp_replicate_dim: int,
|
||||
hsdp_shard_dim: int,
|
||||
default_dtype: torch.dtype,
|
||||
param_dtype: torch.dtype,
|
||||
reduce_dtype: torch.dtype,
|
||||
cpu_offload: bool = False,
|
||||
@@ -87,7 +88,7 @@ def maybe_load_fsdp_model(
|
||||
mp_policy=mp_policy,
|
||||
)
|
||||
|
||||
with set_default_dtype(param_dtype), torch.device("meta"):
|
||||
with set_default_dtype(default_dtype), torch.device("meta"):
|
||||
model = model_cls(**init_params)
|
||||
|
||||
# Check if we should use FSDP
|
||||
@@ -125,7 +126,8 @@ def maybe_load_fsdp_model(
|
||||
model,
|
||||
weight_iterator,
|
||||
device,
|
||||
param_dtype,
|
||||
# param_dtype,
|
||||
default_dtype,
|
||||
strict=True,
|
||||
cpu_offload=cpu_offload,
|
||||
param_names_mapping=param_names_mapping_fn,
|
||||
@@ -137,6 +139,8 @@ def maybe_load_fsdp_model(
|
||||
# Avoid unintended computation graph accumulation during inference
|
||||
if isinstance(p, torch.nn.Parameter):
|
||||
p.requires_grad = False
|
||||
for param in model.parameters():
|
||||
assert param.dtype == torch.float32
|
||||
return model
|
||||
|
||||
|
||||
|
||||
@@ -171,10 +171,10 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
|
||||
# timestep shape should be [B]
|
||||
dtype = pred_noise.dtype
|
||||
device = pred_noise.device
|
||||
pred_noise = pred_noise.float().to(device)
|
||||
noise_input_latent = noise_input_latent.float().to(device)
|
||||
sigmas = scheduler.sigmas.float().to(device)
|
||||
timesteps = scheduler.timesteps.float().to(device)
|
||||
pred_noise = pred_noise.double().to(device)
|
||||
noise_input_latent = noise_input_latent.double().to(device)
|
||||
sigmas = scheduler.sigmas.double().to(device)
|
||||
timesteps = scheduler.timesteps.double().to(device)
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
@@ -121,14 +121,25 @@ class ComposedPipelineBase(ABC):
|
||||
model_path: str,
|
||||
device: str | None = None,
|
||||
torch_dtype: torch.dtype | None = None,
|
||||
pipeline_config: str | PipelineConfig | None = None,
|
||||
pipeline_config: PipelineConfig | None = None,
|
||||
args: argparse.Namespace | None = None,
|
||||
required_config_modules: list[str] | None = None,
|
||||
loaded_modules: dict[str, torch.nn.Module]
|
||||
| None = None,
|
||||
**kwargs) -> "ComposedPipelineBase":
|
||||
"""
|
||||
Load a pipeline from a pretrained model.
|
||||
Load a pipeline from a pretrained model.
|
||||
Few different patterns are supported:
|
||||
- Only provide model_path:
|
||||
- This will load the pipeline in inference mode.
|
||||
- The pipeline will be initialized with the default config.
|
||||
- The pipeline will be initialized with the default modules.
|
||||
- The pipeline will be initialized with the default stages.
|
||||
- The pipeline will be initialized with the default stages.
|
||||
- override the default config using pipeline_config or args or kwargs
|
||||
- override the default modules using loaded_modules
|
||||
- override the pipelineconfig
|
||||
|
||||
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
|
||||
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
|
||||
"""
|
||||
@@ -136,9 +147,18 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
kwargs['model_path'] = model_path
|
||||
fastvideo_args = FastVideoArgs.from_kwargs(**kwargs)
|
||||
if pipeline_config is not None:
|
||||
fastvideo_args.pipeline_config = pipeline_config
|
||||
if fastvideo_args.override_transformer_cls_name is not None:
|
||||
pipeline_config = PipelineConfig.from_pretrained("wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers")
|
||||
fastvideo_args.pipeline_config = pipeline_config
|
||||
else:
|
||||
assert args is not None, "args must be provided for training mode"
|
||||
fastvideo_args = TrainingArgs.from_cli_args(args)
|
||||
if fastvideo_args.override_transformer_cls_name is not None:
|
||||
pipeline_config = PipelineConfig.from_pretrained("wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers")
|
||||
fastvideo_args.pipeline_config = pipeline_config
|
||||
logger.info("in 2 Overriding transformer cls name to %s", fastvideo_args.override_transformer_cls_name)
|
||||
# TODO(will): fix this so that its not so ugly
|
||||
fastvideo_args.model_path = model_path
|
||||
for key, value in kwargs.items():
|
||||
@@ -149,7 +169,8 @@ class ComposedPipelineBase(ABC):
|
||||
# model is loaded with the correct precision. Subsequently we will
|
||||
# use FSDP2's MixedPrecisionPolicy to set the precision for the
|
||||
# fwd, bwd, and other operations' precision.
|
||||
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
|
||||
fastvideo_args.pipeline_config.dit_precision = 'fp32'
|
||||
# assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
|
||||
|
||||
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
|
||||
|
||||
@@ -237,20 +258,19 @@ class ComposedPipelineBase(ABC):
|
||||
# remove keys that are not pipeline modules
|
||||
model_index.pop("_class_name")
|
||||
model_index.pop("_diffusers_version")
|
||||
# @TODO(Wei): Temporary hack
|
||||
if "boundary_ratio" in model_index and model_index[
|
||||
"boundary_ratio"] is not None:
|
||||
logger.info(
|
||||
"MoE pipeline detected. Adding transformer_2 to self.required_config_modules..."
|
||||
)
|
||||
self.required_config_modules.append("transformer_2")
|
||||
if fastvideo_args.boundary_ratio is None:
|
||||
logger.info(
|
||||
"MoE pipeline detected. Setting boundary ratio to %s",
|
||||
model_index["boundary_ratio"])
|
||||
fastvideo_args.boundary_ratio = model_index["boundary_ratio"]
|
||||
logger.info("MoE pipeline detected. Setting boundary ratio to %s",
|
||||
model_index["boundary_ratio"])
|
||||
fastvideo_args.pipeline_config.dit_config.boundary_ratio = model_index[
|
||||
"boundary_ratio"]
|
||||
|
||||
model_index.pop("boundary_ratio", None)
|
||||
# used by Wan2.2 ti2v
|
||||
model_index.pop("expand_timesteps", None)
|
||||
|
||||
# some sanity checks
|
||||
@@ -283,8 +303,8 @@ class ComposedPipelineBase(ABC):
|
||||
architecture) in model_index.items():
|
||||
if transformers_or_diffusers is None:
|
||||
logger.warning(
|
||||
"Module in model_index.json has null value, removing from required_config_modules"
|
||||
)
|
||||
"Module %s in model_index.json has null value, removing from required_config_modules",
|
||||
module_name)
|
||||
if module_name in self.required_config_modules:
|
||||
self.required_config_modules.remove(module_name)
|
||||
continue
|
||||
|
||||
@@ -129,6 +129,7 @@ class ForwardBatch:
|
||||
timesteps: torch.Tensor | None = None
|
||||
timestep: torch.Tensor | float | int | None = None
|
||||
step_index: int | None = None
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Scheduler parameters
|
||||
num_inference_steps: int = 50
|
||||
@@ -147,7 +148,12 @@ class ForwardBatch:
|
||||
modules: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# Final output (after pipeline completion)
|
||||
output: Any = None
|
||||
output: torch.Tensor | None = None
|
||||
return_trajectory_latents: bool = False
|
||||
return_trajectory_decoded: bool = False
|
||||
trajectory_timesteps: list[int] | None = None
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_decoded: list[torch.Tensor] | None = None
|
||||
|
||||
# Extra parameters that might be needed by specific pipeline implementations
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
@@ -206,6 +212,10 @@ class TrainingBatch:
|
||||
infos: list[dict[str, Any]] | None = None
|
||||
mask_lat_size: torch.Tensor | None = None
|
||||
|
||||
# ODE trajectory supervision
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_timesteps: torch.Tensor | None = None
|
||||
|
||||
# Transformer inputs
|
||||
noisy_model_input: torch.Tensor | None = None
|
||||
timesteps: torch.Tensor | None = None
|
||||
@@ -236,6 +246,7 @@ class TrainingBatch:
|
||||
fake_score_loss: float = 0.0
|
||||
|
||||
dmd_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
latent_vis_dict: dict[str, torch.Tensor] = field(default_factory=dict)
|
||||
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import multiprocessing
|
||||
import os
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
@@ -12,6 +10,8 @@ from torch.utils.data import DataLoader
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset import getdataset
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.dataset.preprocessing_datasets import PreprocessBatch
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -54,10 +54,14 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"""Get additional features specific to the pipeline type. Override in subclasses."""
|
||||
return {}
|
||||
|
||||
def get_schema_fields(self) -> list[str]:
|
||||
"""Get the schema fields for the pipeline type. Override in subclasses."""
|
||||
def get_pyarrow_schema(self) -> pa.Schema:
|
||||
"""Return the PyArrow schema for this pipeline. Must be overridden."""
|
||||
raise NotImplementedError
|
||||
|
||||
def get_schema_fields(self) -> list[str]:
|
||||
"""Get the schema fields for the pipeline type."""
|
||||
return [f.name for f in self.get_pyarrow_schema()]
|
||||
|
||||
def create_record_for_schema(self,
|
||||
preprocess_batch: PreprocessBatch,
|
||||
schema: pa.Schema,
|
||||
@@ -400,166 +404,22 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
batch_data.append(record)
|
||||
|
||||
if batch_data:
|
||||
# Add progress bar for writing to Parquet dataset
|
||||
write_pbar = tqdm(total=1,
|
||||
desc="Writing to Parquet dataset",
|
||||
unit="batch")
|
||||
# Convert batch data to PyArrow arrays
|
||||
arrays = []
|
||||
for field in self.get_schema_fields():
|
||||
if field.endswith('_bytes'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.binary()))
|
||||
elif field.endswith('_shape'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.list_(pa.int32())))
|
||||
elif field in ['width', 'height', 'num_frames']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.int32()))
|
||||
elif field in ['duration_sec', 'fps']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.float32()))
|
||||
else:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data]))
|
||||
|
||||
table = pa.Table.from_arrays(arrays,
|
||||
names=self.get_schema_fields())
|
||||
table = records_to_table(batch_data, self.get_pyarrow_schema())
|
||||
write_pbar.update(1)
|
||||
write_pbar.close()
|
||||
|
||||
# Store the table in a list for later processing
|
||||
if not hasattr(self, 'all_tables'):
|
||||
self.all_tables = []
|
||||
self.all_tables.append(table)
|
||||
|
||||
if not hasattr(self, 'dataset_writer'):
|
||||
self.dataset_writer = ParquetDatasetWriter(
|
||||
out_dir=combined_parquet_dir,
|
||||
samples_per_file=args.samples_per_file,
|
||||
)
|
||||
self.dataset_writer.append_table(table)
|
||||
logger.info("Collected batch with %s samples", len(table))
|
||||
|
||||
if num_processed_samples >= args.flush_frequency:
|
||||
self._flush_tables(num_processed_samples, args,
|
||||
combined_parquet_dir)
|
||||
written = self.dataset_writer.flush()
|
||||
logger.info("Flushed %s samples to parquet", written)
|
||||
num_processed_samples = 0
|
||||
self.all_tables = []
|
||||
|
||||
def _flush_tables(self, num_processed_samples: int, args,
|
||||
combined_parquet_dir: str):
|
||||
"""Flush collected tables to disk."""
|
||||
assert hasattr(self, 'all_tables') and self.all_tables
|
||||
print(f"Combining {len(self.all_tables)} batches...")
|
||||
combined_table = pa.concat_tables(self.all_tables)
|
||||
assert len(combined_table) == num_processed_samples
|
||||
print(f"Total samples collected: {len(combined_table)}")
|
||||
|
||||
# Calculate total number of chunks needed, discarding remainder
|
||||
total_chunks = max(num_processed_samples // args.samples_per_file, 1)
|
||||
|
||||
print(f"Fixed samples per parquet file: {args.samples_per_file}")
|
||||
print(f"Total number of parquet files: {total_chunks}")
|
||||
print(
|
||||
f"Total samples to be processed: {total_chunks * args.samples_per_file} (discarding {num_processed_samples % args.samples_per_file} samples)"
|
||||
)
|
||||
|
||||
# Split work among processes
|
||||
num_workers = int(min(multiprocessing.cpu_count(), total_chunks))
|
||||
chunks_per_worker = (total_chunks + num_workers - 1) // num_workers
|
||||
|
||||
print(f"Using {num_workers} workers to process {total_chunks} chunks")
|
||||
logger.info("Chunks per worker: %s", chunks_per_worker)
|
||||
|
||||
# Prepare work ranges
|
||||
work_ranges = []
|
||||
for i in range(num_workers):
|
||||
start_idx = i * chunks_per_worker
|
||||
end_idx = min((i + 1) * chunks_per_worker, total_chunks)
|
||||
if start_idx < total_chunks:
|
||||
work_ranges.append(
|
||||
(start_idx, end_idx, combined_table, i,
|
||||
combined_parquet_dir, args.samples_per_file))
|
||||
|
||||
total_written = 0
|
||||
failed_ranges = []
|
||||
with ProcessPoolExecutor(max_workers=num_workers) as executor:
|
||||
futures = {
|
||||
executor.submit(self.process_chunk_range, work_range):
|
||||
work_range
|
||||
for work_range in work_ranges
|
||||
}
|
||||
for future in tqdm(futures, desc="Processing chunks"):
|
||||
try:
|
||||
written = future.result()
|
||||
total_written += written
|
||||
logger.info("Processed chunk with %s samples", written)
|
||||
except Exception as e:
|
||||
work_range = futures[future]
|
||||
failed_ranges.append(work_range)
|
||||
logger.error("Failed to process range %s-%s: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
# Retry failed ranges sequentially
|
||||
if failed_ranges:
|
||||
logger.warning("Retrying %s failed ranges sequentially",
|
||||
len(failed_ranges))
|
||||
for work_range in failed_ranges:
|
||||
try:
|
||||
total_written += self.process_chunk_range(work_range)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to process range %s-%s after retry: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
logger.info("Total samples written: %s", total_written)
|
||||
|
||||
@staticmethod
|
||||
def process_chunk_range(args: Any) -> int:
|
||||
start_idx, end_idx, table, worker_id, output_dir, samples_per_file = args
|
||||
try:
|
||||
total_written = 0
|
||||
num_samples = len(table)
|
||||
|
||||
# Create worker-specific subdirectory
|
||||
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
|
||||
os.makedirs(worker_dir, exist_ok=True)
|
||||
|
||||
# Check how many files there are already in the dir, and update i accordingly
|
||||
num_parquets = 0
|
||||
for root, _, files in os.walk(worker_dir):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
num_parquets += 1
|
||||
|
||||
for i in range(start_idx, end_idx):
|
||||
start_sample = i * samples_per_file
|
||||
end_sample = min((i + 1) * samples_per_file, num_samples)
|
||||
chunk = table.slice(start_sample, end_sample - start_sample)
|
||||
|
||||
# Create chunk file in worker's directory
|
||||
chunk_path = os.path.join(
|
||||
worker_dir, f"data_chunk_{i + num_parquets}.parquet")
|
||||
temp_path = chunk_path + '.tmp'
|
||||
|
||||
try:
|
||||
# Write to temporary file
|
||||
pq.write_table(chunk, temp_path, compression='zstd')
|
||||
|
||||
# Rename temporary file to final file
|
||||
if os.path.exists(chunk_path):
|
||||
os.remove(
|
||||
chunk_path) # Remove existing file if it exists
|
||||
os.rename(temp_path, chunk_path)
|
||||
|
||||
total_written += len(chunk)
|
||||
except Exception as e:
|
||||
# Clean up temporary file if it exists
|
||||
if os.path.exists(temp_path):
|
||||
os.remove(temp_path)
|
||||
raise e
|
||||
|
||||
return total_written
|
||||
except Exception as e:
|
||||
logger.error("Error processing chunks %s-%s for worker %s: %s",
|
||||
start_idx, end_idx, worker_id, str(e))
|
||||
raise
|
||||
|
||||
@@ -40,9 +40,9 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
image_processor=self.get_module("image_processor"),
|
||||
))
|
||||
|
||||
def get_schema_fields(self) -> list[str]:
|
||||
"""Get the schema fields for I2V pipeline."""
|
||||
return [f.name for f in pyarrow_schema_i2v]
|
||||
def get_pyarrow_schema(self):
|
||||
"""Return the PyArrow schema for I2V pipeline."""
|
||||
return pyarrow_schema_i2v
|
||||
|
||||
def get_extra_features(self, valid_data: dict[str, Any],
|
||||
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
|
||||
|
||||
@@ -0,0 +1,654 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
ODE Trajectory Data Preprocessing pipeline implementation.
|
||||
|
||||
This module contains an implementation of the ODE Trajectory Data Preprocessing pipeline
|
||||
using the modular pipeline architecture.
|
||||
|
||||
Sec 4.3 of CausVid paper: https://arxiv.org/pdf/2412.07772
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import torch
|
||||
from PIL import Image
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset import getdataset
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_ode_trajectory
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
|
||||
ImageVAEEncodingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage)
|
||||
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class FlowMatchScheduler:
|
||||
|
||||
order = 1
|
||||
|
||||
def __init__(self,
|
||||
num_inference_steps=100,
|
||||
num_train_timesteps=1000,
|
||||
shift=3.0,
|
||||
sigma_max=1.0,
|
||||
sigma_min=0.003 / 1.002,
|
||||
inverse_timesteps=False,
|
||||
extra_one_step=False,
|
||||
reverse_sigmas=False):
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.shift = shift
|
||||
self.sigma_max = sigma_max
|
||||
self.sigma_min = sigma_min
|
||||
self.inverse_timesteps = inverse_timesteps
|
||||
self.extra_one_step = extra_one_step
|
||||
self.reverse_sigmas = reverse_sigmas
|
||||
self.set_timesteps(num_inference_steps)
|
||||
|
||||
def set_timesteps(self,
|
||||
num_inference_steps=100,
|
||||
denoising_strength=1.0,
|
||||
training=False,
|
||||
device=None):
|
||||
sigma_start = self.sigma_min + \
|
||||
(self.sigma_max - self.sigma_min) * denoising_strength
|
||||
if self.extra_one_step:
|
||||
self.sigmas = torch.linspace(sigma_start, self.sigma_min,
|
||||
num_inference_steps + 1)[:-1]
|
||||
else:
|
||||
self.sigmas = torch.linspace(sigma_start, self.sigma_min,
|
||||
num_inference_steps)
|
||||
if self.inverse_timesteps:
|
||||
self.sigmas = torch.flip(self.sigmas, dims=[0])
|
||||
self.sigmas = self.shift * self.sigmas / \
|
||||
(1 + (self.shift - 1) * self.sigmas)
|
||||
if self.reverse_sigmas:
|
||||
self.sigmas = 1 - self.sigmas
|
||||
self.timesteps = self.sigmas * self.num_train_timesteps
|
||||
if training:
|
||||
x = self.timesteps
|
||||
y = torch.exp(
|
||||
-2 * ((x - num_inference_steps / 2) / num_inference_steps)**2)
|
||||
y_shifted = y - y.min()
|
||||
bsmntw_weighing = y_shifted * \
|
||||
(num_inference_steps / y_shifted.sum())
|
||||
self.linear_timesteps_weights = bsmntw_weighing
|
||||
|
||||
def step(self,
|
||||
model_output,
|
||||
timestep,
|
||||
sample,
|
||||
to_final=False,
|
||||
return_dict=False,
|
||||
**kwargs):
|
||||
assert return_dict is False
|
||||
assert kwargs == {}
|
||||
self.sigmas = self.sigmas.to(model_output.device)
|
||||
self.timesteps = self.timesteps.to(model_output.device)
|
||||
logger.info('step timestep: %s', timestep)
|
||||
logger.info('step timestep: %s', timestep.shape)
|
||||
# timestep is [num_frames]
|
||||
# timestep_id = torch.argmin(
|
||||
# (self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
# assert timestep.ndim == 1
|
||||
# assert timestep.shape[0] == 1
|
||||
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
if to_final or (timestep_id + 1 >= len(self.timesteps)).any():
|
||||
sigma_ = 1 if (self.inverse_timesteps or self.reverse_sigmas) else 0
|
||||
else:
|
||||
sigma_ = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
|
||||
prev_sample = sample + model_output * (sigma_ - sigma)
|
||||
return (prev_sample, )
|
||||
|
||||
def scale_model_input(self, sample: torch.Tensor, *args,
|
||||
**kwargs) -> torch.Tensor:
|
||||
"""
|
||||
Ensures interchangeability with schedulers that need to scale the denoising model input depending on the
|
||||
current timestep.
|
||||
|
||||
Args:
|
||||
sample (`torch.Tensor`):
|
||||
The input sample.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
A scaled input sample.
|
||||
"""
|
||||
return sample
|
||||
|
||||
def add_noise(self, original_samples, noise, timestep):
|
||||
"""
|
||||
Diffusion forward corruption process.
|
||||
Input:
|
||||
- clean_latent: the clean latent with shape [B, C, H, W]
|
||||
- noise: the noise with shape [B, C, H, W]
|
||||
- timestep: the timestep with shape [B]
|
||||
Output: the corrupted latent with shape [B, C, H, W]
|
||||
"""
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
sample = (1 - sigma) * original_samples + sigma * noise
|
||||
return sample.type_as(noise)
|
||||
|
||||
def training_target(self, sample, noise, timestep):
|
||||
target = noise - sample
|
||||
return target
|
||||
|
||||
def training_weight(self, timestep):
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps - timestep.to(self.timesteps.device)).abs())
|
||||
weights = self.linear_timesteps_weights[timestep_id]
|
||||
return weights
|
||||
|
||||
|
||||
class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
"""ODE Trajectory preprocessing pipeline implementation."""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
preprocess_dataloader: StatefulDataLoader
|
||||
preprocess_loader_iter: Iterator[dict[str, Any]]
|
||||
|
||||
def get_schema_fields(self):
|
||||
"""Get the schema fields for ODE Trajectory pipeline."""
|
||||
return [f.name for f in pyarrow_schema_ode_trajectory]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
fastvideo_args.pipeline_config.flow_shift = 5
|
||||
logger.info('WTF flow_shift: %s',
|
||||
fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
assert fastvideo_args.pipeline_config.flow_shift == 5
|
||||
# self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
|
||||
# shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
self.modules["scheduler"] = FlowMatchScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
sigma_min=0.0,
|
||||
extra_one_step=True)
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=48,
|
||||
denoising_strength=1.0)
|
||||
logger.info('WTF scheduler timesteps: %s',
|
||||
self.modules["scheduler"].timesteps)
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
self.add_stage(stage_name="vae_encoding_stage",
|
||||
stage=ImageVAEEncodingStage(
|
||||
vae=self.get_module("vae"), ))
|
||||
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="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
pipeline=self,
|
||||
))
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def preprocess_video_and_text_and_trajectory(self,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
args):
|
||||
|
||||
for batch_idx, data in enumerate(self.pbar):
|
||||
if data is None:
|
||||
continue
|
||||
|
||||
with torch.inference_mode():
|
||||
# Filter out invalid samples (those with all zeros)
|
||||
valid_indices = []
|
||||
for i, pixel_values in enumerate(data["pixel_values"]):
|
||||
if not torch.all(
|
||||
pixel_values == 0): # Check if all values are zero
|
||||
valid_indices.append(i)
|
||||
self.num_processed_samples += len(valid_indices)
|
||||
|
||||
if not valid_indices:
|
||||
continue
|
||||
|
||||
# Create new batch with only valid samples
|
||||
valid_data = {
|
||||
"pixel_values":
|
||||
torch.stack(
|
||||
[data["pixel_values"][i] for i in valid_indices]),
|
||||
"text": [data["text"][i] for i in valid_indices],
|
||||
"path": [data["path"][i] for i in valid_indices],
|
||||
"fps": [data["fps"][i] for i in valid_indices],
|
||||
"duration": [data["duration"][i] for i in valid_indices],
|
||||
}
|
||||
|
||||
# VAE
|
||||
with torch.autocast("cuda", dtype=torch.float32):
|
||||
latents = self.get_module("vae").encode(
|
||||
valid_data["pixel_values"].to(
|
||||
get_local_torch_device())).mean
|
||||
|
||||
# Get extra features if needed
|
||||
extra_features = self.get_extra_features(
|
||||
valid_data, fastvideo_args)
|
||||
|
||||
batch_captions = valid_data["text"]
|
||||
logger.info(f"===== batch_captions: {batch_captions}")
|
||||
# Encode text using the standalone TextEncodingStage API
|
||||
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
batch_captions,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
prompt_embeds = prompt_embeds_list[0]
|
||||
prompt_attention_masks = prompt_masks_list[0]
|
||||
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
|
||||
|
||||
# # Get sequence lengths from attention masks (number of 1s)
|
||||
# seq_lens = prompt_attention_mask.sum(dim=1)
|
||||
|
||||
# non_padded_embeds = []
|
||||
# non_padded_masks = []
|
||||
|
||||
# # Process each item in the batch
|
||||
# for i in range(prompt_embeds.size(0)):
|
||||
# seq_len = seq_lens[i].item()
|
||||
# # Slice the embeddings and masks to keep only non-padding parts
|
||||
# non_padded_embeds.append(prompt_embeds[i, :seq_len])
|
||||
# non_padded_masks.append(prompt_attention_mask[i, :seq_len])
|
||||
|
||||
# Update the tensors with non-padded versions
|
||||
# prompt_embeds = non_padded_embeds
|
||||
# prompt_attention_masks = non_padded_masks
|
||||
# prompt_embeds = prompt_embeds
|
||||
|
||||
# logger.info(f"===== prompt_embeds: {prompt_embeds[0].shape}")
|
||||
# logger.info(f"===== prompt_attention_masks: {prompt_attention_masks[0].shape}")
|
||||
|
||||
sampling_params = SamplingParam.from_pretrained(args.model_path)
|
||||
|
||||
# encode negative prompt for trajectory collection
|
||||
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
|
||||
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
sampling_params.negative_prompt,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
negative_prompt_embed = negative_prompt_embeds_list[0][0]
|
||||
negative_prompt_attention_mask = negative_prompt_masks_list[
|
||||
0][0]
|
||||
else:
|
||||
negative_prompt_embed = None
|
||||
negative_prompt_attention_mask = None
|
||||
|
||||
trajectory_latents = []
|
||||
trajectory_timesteps = []
|
||||
trajectory_decoded = []
|
||||
for i, (prompt_embed, prompt_attention_mask) in enumerate(
|
||||
zip(prompt_embeds, prompt_attention_masks, strict=False)):
|
||||
prompt_embed = prompt_embed.unsqueeze(0)
|
||||
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
|
||||
logger.info("what")
|
||||
logger.info(f"===== prompt_embed: {prompt_embed.shape}")
|
||||
logger.info(
|
||||
f"===== prompt_attention_mask: {prompt_attention_mask.shape}"
|
||||
)
|
||||
# Collect the trajectory data
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_params),
|
||||
# data_type="video",
|
||||
# seed=args.seed,
|
||||
# prompt=batch_captions[i],
|
||||
# prompt_embeds=[prompt_embed],
|
||||
# prompt_attention_mask=[prompt_attention_mask],
|
||||
# height=args.max_height,
|
||||
# width=args.max_width,
|
||||
# num_frames=81,
|
||||
# fps=args.train_fps,
|
||||
# return_trajectory_latents=True,
|
||||
# guidance_scale=3.0,
|
||||
# do_classifier_free_guidance=True,
|
||||
)
|
||||
batch.prompt_embeds = [prompt_embed]
|
||||
batch.prompt_attention_mask = [prompt_attention_mask]
|
||||
batch.negative_prompt_embeds = [negative_prompt_embed]
|
||||
batch.negative_attention_mask = [
|
||||
negative_prompt_attention_mask
|
||||
]
|
||||
batch.return_trajectory_latents = True
|
||||
batch.return_trajectory_decoded = False
|
||||
batch.height = args.max_height
|
||||
batch.width = args.max_width
|
||||
batch.num_inference_steps = 48
|
||||
# batch.num_frames = 81
|
||||
batch.fps = args.train_fps
|
||||
batch.guidance_scale = 6.0
|
||||
batch.do_classifier_free_guidance = True
|
||||
# fastvideo_args.pipeline_config.ti2v_task = True
|
||||
|
||||
result_batch = self.input_validation_stage(
|
||||
batch, fastvideo_args)
|
||||
# result_batch = self.prompt_encoding_stage(result_batch, fastvideo_args)
|
||||
# result_batch = self.vae_encoding_stage(result_batch, fastvideo_args)
|
||||
result_batch = self.timestep_preparation_stage(
|
||||
batch, fastvideo_args)
|
||||
result_batch = self.latent_preparation_stage(
|
||||
result_batch, fastvideo_args)
|
||||
result_batch = self.denoising_stage(result_batch,
|
||||
fastvideo_args)
|
||||
result_batch = self.decoding_stage(result_batch,
|
||||
fastvideo_args)
|
||||
# trajectory_latents = result_batch.trajectory_latents
|
||||
trajectory_latents.append(
|
||||
result_batch.trajectory_latents.cpu())
|
||||
trajectory_timesteps.append(
|
||||
result_batch.trajectory_timesteps.cpu())
|
||||
trajectory_decoded.append(result_batch.trajectory_decoded)
|
||||
|
||||
extra_features["trajectory_latents"] = trajectory_latents
|
||||
extra_features["trajectory_timesteps"] = trajectory_timesteps
|
||||
logger.info(
|
||||
f"===== trajectory_latents: {trajectory_latents[0].shape}")
|
||||
logger.info(
|
||||
f"===== trajectory_latents len: {len(trajectory_latents)}")
|
||||
logger.info(f"===== trajectory_timesteps: {trajectory_timesteps}")
|
||||
logger.info(
|
||||
f"===== trajectory_timesteps len: {len(trajectory_timesteps)}")
|
||||
|
||||
if batch.return_trajectory_decoded:
|
||||
logger.info("===== SAVING TRAJECTORY DECODED")
|
||||
for i, decoded_frames in enumerate(trajectory_decoded):
|
||||
for j, decoded_frame in enumerate(decoded_frames):
|
||||
logger.info(
|
||||
f"===== SAVING TRAJECTORY DECODED {i} for prompt {batch_captions[i]}"
|
||||
)
|
||||
save_decoded_latents_as_video(
|
||||
decoded_frame,
|
||||
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
|
||||
args.train_fps)
|
||||
# assert False
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data = []
|
||||
|
||||
# Add progress bar for saving outputs
|
||||
save_pbar = tqdm(enumerate(valid_data["path"]),
|
||||
desc="Saving outputs",
|
||||
unit="item",
|
||||
leave=False)
|
||||
for idx, video_path in save_pbar:
|
||||
# Get the corresponding latent and info using video name
|
||||
latent = latents[idx].cpu()
|
||||
video_name = os.path.basename(video_path).split(".")[0]
|
||||
|
||||
# Convert tensors to numpy arrays
|
||||
vae_latent = latent.cpu().numpy()
|
||||
text_embedding = prompt_embeds[idx].cpu().numpy()
|
||||
|
||||
# Get extra features for this sample if needed
|
||||
sample_extra_features = {}
|
||||
if extra_features:
|
||||
for key, value in extra_features.items():
|
||||
logger.info(f"===== key: {key}")
|
||||
if isinstance(value, torch.Tensor):
|
||||
logger.info(f"===== value: {value[idx].shape}")
|
||||
sample_extra_features[key] = value[idx].cpu().numpy(
|
||||
)
|
||||
else:
|
||||
assert isinstance(value, list)
|
||||
if isinstance(value[idx], torch.Tensor):
|
||||
logger.info(
|
||||
f"===== value in list: {value[idx].shape}")
|
||||
sample_extra_features[key] = value[idx].cpu(
|
||||
).float().numpy()
|
||||
else:
|
||||
logger.info("===== value in list: not tensor")
|
||||
sample_extra_features[key] = value[idx]
|
||||
# logger.info(f"===== value: not tensor")
|
||||
# sample_extra_features[key] = value[idx]
|
||||
|
||||
# Create record for Parquet dataset
|
||||
record = self.create_record(
|
||||
video_name=video_name,
|
||||
vae_latent=vae_latent,
|
||||
text_embedding=text_embedding,
|
||||
valid_data=valid_data,
|
||||
idx=idx,
|
||||
extra_features=sample_extra_features)
|
||||
batch_data.append(record)
|
||||
|
||||
if batch_data:
|
||||
# Add progress bar for writing to Parquet dataset
|
||||
write_pbar = tqdm(total=1,
|
||||
desc="Writing to Parquet dataset",
|
||||
unit="batch")
|
||||
# Convert batch data to PyArrow arrays
|
||||
arrays = []
|
||||
for field in self.get_schema_fields():
|
||||
if field.endswith('_bytes'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.binary()))
|
||||
elif field.endswith('_shape'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.list_(pa.int32())))
|
||||
elif field in ['width', 'height', 'num_frames']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.int32()))
|
||||
elif field in ['duration_sec', 'fps']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.float32()))
|
||||
else:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data]))
|
||||
|
||||
table = pa.Table.from_arrays(arrays,
|
||||
names=self.get_schema_fields())
|
||||
write_pbar.update(1)
|
||||
write_pbar.close()
|
||||
|
||||
# Store the table in a list for later processing
|
||||
if not hasattr(self, 'all_tables'):
|
||||
self.all_tables = []
|
||||
self.all_tables.append(table)
|
||||
|
||||
logger.info("Collected batch with %s samples", len(table))
|
||||
|
||||
if self.num_processed_samples >= args.flush_frequency:
|
||||
self._flush_tables(self.num_processed_samples, args,
|
||||
self.combined_parquet_dir)
|
||||
self.num_processed_samples = 0
|
||||
self.all_tables = []
|
||||
|
||||
def get_extra_features(self, valid_data: dict[str, Any],
|
||||
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
|
||||
|
||||
# TODO(will): move these to cpu at some point
|
||||
self.get_module("vae").to(get_local_torch_device())
|
||||
|
||||
# generator = torch.Generator(device=get_local_torch_device(), seed=42)
|
||||
generator = torch.Generator("cpu").manual_seed(42)
|
||||
|
||||
features = {}
|
||||
"""Get CLIP features from the first frame of each video."""
|
||||
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
|
||||
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
|
||||
_, _, num_frames, height, width = valid_data["pixel_values"].shape
|
||||
# latent_height = height // self.get_module(
|
||||
# "vae").spatial_compression_ratio
|
||||
# latent_width = width // self.get_module("vae").spatial_compression_ratio
|
||||
|
||||
unprocessed_images = []
|
||||
pil_images = []
|
||||
# Frame has values between -1 and 1
|
||||
for frame in first_frame:
|
||||
frame = (frame + 1) * 127.5
|
||||
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
|
||||
pil_images.append(frame_pil)
|
||||
# processed_img = self.get_module("image_processor")(
|
||||
# images=frame_pil, return_tensors="pt")
|
||||
unprocessed_images.append(frame_pil)
|
||||
"""Get VAE features from the first frame of each video"""
|
||||
video_conditions = []
|
||||
for frame in unprocessed_images:
|
||||
|
||||
latent = self.vae_encoding_stage.encode_image(
|
||||
frame, height, width, fastvideo_args, generator)
|
||||
video_conditions.append(latent)
|
||||
|
||||
features["image_condition_latents"] = video_conditions
|
||||
features["pil_images"] = pil_images
|
||||
return features
|
||||
|
||||
def create_record(
|
||||
self,
|
||||
video_name: str,
|
||||
vae_latent: np.ndarray,
|
||||
text_embedding: np.ndarray,
|
||||
valid_data: dict[str, Any],
|
||||
idx: int,
|
||||
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
"""Create a record for the Parquet dataset with CLIP features."""
|
||||
record = super().create_record(video_name=video_name,
|
||||
vae_latent=vae_latent,
|
||||
text_embedding=text_embedding,
|
||||
valid_data=valid_data,
|
||||
idx=idx,
|
||||
extra_features=extra_features)
|
||||
|
||||
if extra_features and "image_condition_latents" in extra_features:
|
||||
image_condition_latents = extra_features["image_condition_latents"]
|
||||
record.update({
|
||||
"image_condition_latents_bytes":
|
||||
image_condition_latents.tobytes(),
|
||||
"image_condition_latents_shape":
|
||||
list(image_condition_latents.shape),
|
||||
"image_condition_latents_dtype":
|
||||
str(image_condition_latents.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"image_condition_latents_bytes": b"",
|
||||
"image_condition_latents_shape": [],
|
||||
"image_condition_latents_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "trajectory_latents" in extra_features:
|
||||
trajectory_latents = extra_features["trajectory_latents"]
|
||||
record.update({
|
||||
"trajectory_latents_bytes":
|
||||
trajectory_latents.tobytes(),
|
||||
"trajectory_latents_shape":
|
||||
list(trajectory_latents.shape),
|
||||
"trajectory_latents_dtype":
|
||||
str(trajectory_latents.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"trajectory_latents_bytes": b"",
|
||||
"trajectory_latents_shape": [],
|
||||
"trajectory_latents_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "trajectory_timesteps" in extra_features:
|
||||
trajectory_timesteps = extra_features["trajectory_timesteps"]
|
||||
record.update({
|
||||
"trajectory_timesteps_bytes":
|
||||
trajectory_timesteps.tobytes(),
|
||||
"trajectory_timesteps_shape":
|
||||
list(trajectory_timesteps.shape),
|
||||
"trajectory_timesteps_dtype":
|
||||
str(trajectory_timesteps.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"trajectory_timesteps_bytes": b"",
|
||||
"trajectory_timesteps_shape": [],
|
||||
"trajectory_timesteps_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "pil_image" in extra_features:
|
||||
pil_image = extra_features["pil_image"]
|
||||
record.update({
|
||||
"pil_image_bytes": pil_image.tobytes(),
|
||||
"pil_image_shape": list(pil_image.shape),
|
||||
"pil_image_dtype": str(pil_image.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"pil_image_bytes": b"",
|
||||
"pil_image_shape": [],
|
||||
"pil_image_dtype": "",
|
||||
})
|
||||
|
||||
return record
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
|
||||
self.local_rank = int(os.getenv("RANK", 0))
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
# Create directory for combined data
|
||||
self.combined_parquet_dir = os.path.join(args.output_dir,
|
||||
"combined_parquet_dataset")
|
||||
os.makedirs(self.combined_parquet_dir, exist_ok=True)
|
||||
|
||||
# Loading dataset
|
||||
train_dataset = getdataset(args)
|
||||
|
||||
self.preprocess_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
batch_size=args.preprocess_video_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
|
||||
|
||||
self.num_processed_samples = 0
|
||||
# Add progress bar for video preprocessing
|
||||
self.pbar = tqdm(self.preprocess_loader_iter,
|
||||
desc="Processing videos",
|
||||
unit="batch",
|
||||
disable=self.local_rank != 0)
|
||||
|
||||
# Initialize class variables for data sharing
|
||||
self.video_data: dict[str, Any] = {} # Store video metadata and paths
|
||||
self.latent_data: dict[str, Any] = {} # Store latent tensors
|
||||
self.preprocess_video_and_text_and_trajectory(fastvideo_args, args)
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_ODE_Trajectory
|
||||
@@ -15,9 +15,9 @@ class PreprocessPipeline_T2V(BasePreprocessPipeline):
|
||||
|
||||
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
|
||||
|
||||
def get_schema_fields(self):
|
||||
"""Get the schema fields for T2V pipeline."""
|
||||
return [f.name for f in pyarrow_schema_t2v]
|
||||
def get_pyarrow_schema(self):
|
||||
"""Return the PyArrow schema for T2V pipeline."""
|
||||
return pyarrow_schema_t2v
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_T2V
|
||||
|
||||
@@ -0,0 +1,184 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Text-only Data Preprocessing pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Text-only Data Preprocessing pipeline
|
||||
using the modular pipeline architecture, based on the ODE Trajectory preprocessing.
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset import gettextdataset
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.dataset.dataloader.record_schema import text_only_record_creator
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_text_only
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.pipelines.stages import TextEncodingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PreprocessPipeline_Text(BasePreprocessPipeline):
|
||||
"""Text-only preprocessing pipeline implementation."""
|
||||
|
||||
_required_config_modules = ["text_encoder", "tokenizer"]
|
||||
preprocess_dataloader: StatefulDataLoader
|
||||
preprocess_loader_iter: Iterator[dict[str, Any]]
|
||||
pbar: Any
|
||||
num_processed_samples: int = 0
|
||||
|
||||
def get_pyarrow_schema(self):
|
||||
"""Return the PyArrow schema for text-only pipeline."""
|
||||
return pyarrow_schema_text_only
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
def preprocess_text_only(self, fastvideo_args: FastVideoArgs, args):
|
||||
"""Preprocess text-only data."""
|
||||
|
||||
for batch_idx, data in enumerate(self.pbar):
|
||||
if data is None:
|
||||
continue
|
||||
|
||||
with torch.inference_mode():
|
||||
# For text-only processing, we only need text data
|
||||
# Filter out samples without text
|
||||
valid_indices = []
|
||||
for i, text in enumerate(data["text"]):
|
||||
if text and text.strip(): # Check if text is not empty
|
||||
valid_indices.append(i)
|
||||
self.num_processed_samples += len(valid_indices)
|
||||
|
||||
if not valid_indices:
|
||||
continue
|
||||
|
||||
# Create new batch with only valid samples (text-only)
|
||||
valid_data = {
|
||||
"text": [data["text"][i] for i in valid_indices],
|
||||
"path": [data["path"][i] for i in valid_indices],
|
||||
}
|
||||
|
||||
batch_captions = valid_data["text"]
|
||||
# Encode text using the standalone TextEncodingStage API
|
||||
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
|
||||
batch_captions,
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
prompt_embeds = prompt_embeds_list[0]
|
||||
prompt_attention_masks = prompt_masks_list[0]
|
||||
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
|
||||
|
||||
logger.info("===== prompt_embeds: %s", prompt_embeds.shape)
|
||||
logger.info("===== prompt_attention_masks: %s",
|
||||
prompt_attention_masks.shape)
|
||||
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data = []
|
||||
|
||||
# Add progress bar for saving outputs
|
||||
save_pbar = tqdm(enumerate(valid_data["path"]),
|
||||
desc="Saving outputs",
|
||||
unit="item",
|
||||
leave=False)
|
||||
|
||||
for idx, text_path in save_pbar:
|
||||
text_name = os.path.basename(text_path).split(".")[0]
|
||||
|
||||
# Convert tensors to numpy arrays
|
||||
text_embedding = prompt_embeds[idx].cpu().numpy()
|
||||
|
||||
# Create record for Parquet dataset (text-only schema)
|
||||
record = text_only_record_creator(
|
||||
text_name=text_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=valid_data["text"][idx],
|
||||
)
|
||||
batch_data.append(record)
|
||||
|
||||
if batch_data:
|
||||
write_pbar = tqdm(total=1,
|
||||
desc="Writing to Parquet dataset",
|
||||
unit="batch")
|
||||
table = records_to_table(batch_data,
|
||||
pyarrow_schema_text_only)
|
||||
write_pbar.update(1)
|
||||
write_pbar.close()
|
||||
|
||||
if not hasattr(self, 'dataset_writer'):
|
||||
self.dataset_writer = ParquetDatasetWriter(
|
||||
out_dir=self.combined_parquet_dir,
|
||||
samples_per_file=args.samples_per_file,
|
||||
)
|
||||
self.dataset_writer.append_table(table)
|
||||
|
||||
logger.info("Collected batch with %s samples", len(table))
|
||||
|
||||
if self.num_processed_samples >= args.flush_frequency:
|
||||
written = self.dataset_writer.flush()
|
||||
logger.info("Flushed %s samples to parquet", written)
|
||||
self.num_processed_samples = 0
|
||||
|
||||
# Final flush for any remaining samples
|
||||
if hasattr(self, 'dataset_writer'):
|
||||
written = self.dataset_writer.flush(write_remainder=True)
|
||||
if written:
|
||||
logger.info("Final flush wrote %s samples", written)
|
||||
|
||||
# Text-only record creation moved to fastvideo.dataset.dataloader.record_schema
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
|
||||
self.local_rank = int(os.getenv("RANK", 0))
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
# Create directory for combined data
|
||||
self.combined_parquet_dir = os.path.join(args.output_dir,
|
||||
"combined_parquet_dataset")
|
||||
os.makedirs(self.combined_parquet_dir, exist_ok=True)
|
||||
|
||||
# Loading text dataset
|
||||
train_dataset = gettextdataset(args)
|
||||
|
||||
self.preprocess_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
batch_size=args.preprocess_video_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
|
||||
|
||||
self.num_processed_samples = 0
|
||||
# Add progress bar for text preprocessing
|
||||
self.pbar = tqdm(self.preprocess_loader_iter,
|
||||
desc="Processing text",
|
||||
unit="batch",
|
||||
disable=self.local_rank != 0)
|
||||
|
||||
# Initialize class variables for data sharing
|
||||
self.text_data: dict[str, Any] = {} # Store text metadata and paths
|
||||
|
||||
self.preprocess_text_only(fastvideo_args, args)
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_Text
|
||||
@@ -1,5 +1,6 @@
|
||||
import argparse
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
@@ -9,8 +10,12 @@ from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v import (
|
||||
PreprocessPipeline_I2V)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (
|
||||
PreprocessPipeline_ODE_Trajectory)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
|
||||
PreprocessPipeline_T2V)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
|
||||
PreprocessPipeline_Text)
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -21,12 +26,22 @@ def main(args) -> None:
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
num_gpus = int(os.environ["WORLD_SIZE"])
|
||||
assert num_gpus == 1, "Only support 1 GPU"
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
|
||||
kwargs = {
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
|
||||
}
|
||||
|
||||
kwargs: dict[str, Any] = {}
|
||||
if args.preprocess_task == "text_only":
|
||||
kwargs = {
|
||||
"text_encoder_cpu_offload": False,
|
||||
}
|
||||
else:
|
||||
# Full config for video/image processing
|
||||
kwargs = {
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
|
||||
}
|
||||
pipeline_config.update_config_from_dict(kwargs)
|
||||
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=args.model_path,
|
||||
num_gpus=get_world_size(),
|
||||
@@ -35,7 +50,19 @@ def main(args) -> None:
|
||||
text_encoder_cpu_offload=False,
|
||||
pipeline_config=pipeline_config,
|
||||
)
|
||||
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
|
||||
if args.preprocess_task == "t2v":
|
||||
PreprocessPipeline = PreprocessPipeline_T2V
|
||||
elif args.preprocess_task == "i2v":
|
||||
PreprocessPipeline = PreprocessPipeline_I2V
|
||||
elif args.preprocess_task == "text_only":
|
||||
PreprocessPipeline = PreprocessPipeline_Text
|
||||
else:
|
||||
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
|
||||
f"Valid options: t2v, i2v, ode_trajectory, text_only")
|
||||
|
||||
logger.info("Preprocess task: %s using %s", args.preprocess_task,
|
||||
PreprocessPipeline.__name__)
|
||||
|
||||
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
|
||||
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
|
||||
|
||||
@@ -74,7 +101,11 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--preprocess_task", type=str, default="t2v")
|
||||
parser.add_argument("--preprocess_task",
|
||||
type=str,
|
||||
default="t2v",
|
||||
choices=["t2v", "i2v", "text_only"],
|
||||
help="Type of preprocessing task to run")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
|
||||
@@ -78,6 +78,8 @@ class CausalDMDDenosingStage(DenoisingStage):
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32)))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
else:
|
||||
assert False, "warp_denoising_step must be true"
|
||||
timesteps = timesteps.to(get_local_torch_device())
|
||||
logger.info("Using timesteps: %s", timesteps)
|
||||
|
||||
|
||||
@@ -50,6 +50,50 @@ class DecodingStage(PipelineStage):
|
||||
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latents: torch.Tensor,
|
||||
fastvideo_args: FastVideoArgs) -> torch.Tensor:
|
||||
"""Decode latents into pixel space."""
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
latents = latents.to(get_local_torch_device())
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
# Decode latents
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
return image
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
@@ -66,6 +110,7 @@ class DecodingStage(PipelineStage):
|
||||
Returns:
|
||||
The batch with decoded outputs.
|
||||
"""
|
||||
# load vae if not already loaded (used for memory constrained devices)
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
if not fastvideo_args.model_loaded["vae"]:
|
||||
loader = VAELoader()
|
||||
@@ -75,58 +120,31 @@ class DecodingStage(PipelineStage):
|
||||
pipeline.add_module("vae", self.vae)
|
||||
fastvideo_args.model_loaded["vae"] = True
|
||||
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
latents = batch.latents
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
if latents is None:
|
||||
raise ValueError("Latents must be provided")
|
||||
|
||||
# Skip decoding if output type is latent
|
||||
if fastvideo_args.output_type == "latent":
|
||||
image = latents
|
||||
frames = batch.latents
|
||||
else:
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
frames = self.decode(batch.latents, fastvideo_args)
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
# Decode latents
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
# decode trajectory latents if needed
|
||||
if batch.return_trajectory_decoded:
|
||||
batch.trajectory_decoded = []
|
||||
logger.info(f"batch.trajectory_latents.shape: {batch.trajectory_latents.shape}")
|
||||
assert batch.trajectory_latents is not None, "batch should have trajectory latents"
|
||||
for idx in range(batch.trajectory_latents.shape[1]):
|
||||
# bathc.trajectory_latents is [batch_size, timesteps, channels, frames, height, width]
|
||||
cur_latent = batch.trajectory_latents[:, idx, :, :, :, :]
|
||||
logger.info(f"cur_latent.shape: {cur_latent.shape}")
|
||||
cur_timestep = batch.trajectory_timesteps[idx]
|
||||
logger.info(
|
||||
f"decoding trajectory latent for timestep: {cur_timestep}")
|
||||
decoded_frames = self.decode(cur_latent, fastvideo_args)
|
||||
batch.trajectory_decoded.append(decoded_frames.cpu().float())
|
||||
|
||||
# Convert to CPU float32 for compatibility
|
||||
image = image.cpu().float()
|
||||
frames = frames.cpu().float()
|
||||
|
||||
# Update batch with decoded image
|
||||
batch.output = image
|
||||
batch.output = frames
|
||||
|
||||
# Offload models if needed
|
||||
if hasattr(self, 'maybe_free_model_hooks'):
|
||||
|
||||
@@ -140,11 +140,12 @@ class DenoisingStage(PipelineStage):
|
||||
latents = latents[:, :, rank_in_sp_group, :, :, :]
|
||||
batch.latents = latents
|
||||
if batch.image_latent is not None:
|
||||
image_latent = rearrange(batch.image_latent,
|
||||
"b c (n t) h w -> b c n t h w",
|
||||
n=sp_world_size).contiguous()
|
||||
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
|
||||
batch.image_latent = image_latent
|
||||
if not fastvideo_args.pipeline_config.ti2v_task and not fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
image_latent = rearrange(batch.image_latent,
|
||||
"b c (n t) h w -> b c n t h w",
|
||||
n=sp_world_size).contiguous()
|
||||
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
|
||||
batch.image_latent = image_latent
|
||||
# Get timesteps and calculate warmup steps
|
||||
timesteps = batch.timesteps
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
@@ -204,8 +205,14 @@ class DenoisingStage(PipelineStage):
|
||||
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
|
||||
|
||||
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
|
||||
if fastvideo_args.boundary_ratio is not None:
|
||||
boundary_timestep = fastvideo_args.boundary_ratio * self.scheduler.num_train_timesteps
|
||||
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
|
||||
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio
|
||||
if batch.boundary_timestep is not None:
|
||||
logger.info("Overriding boundary timestep from %s to %s",
|
||||
boundary_timestep, batch.boundary_timestep)
|
||||
boundary_timestep = batch.boundary_timestep
|
||||
|
||||
boundary_timestep *= self.scheduler.num_train_timesteps
|
||||
else:
|
||||
boundary_timestep = None
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
@@ -247,6 +254,9 @@ class DenoisingStage(PipelineStage):
|
||||
patch_size[2])
|
||||
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
|
||||
|
||||
trajectory_timesteps: list[int] = []
|
||||
trajectory_latents: list[torch.Tensor] = []
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
@@ -274,14 +284,27 @@ class DenoisingStage(PipelineStage):
|
||||
|
||||
# Expand latents for I2V
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if batch.image_latent is not None:
|
||||
if batch.image_latent is not None and not fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.image_latent],
|
||||
dim=1).to(target_dtype)
|
||||
elif batch.image_latent is not None and fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
assert batch.image_latent is not None, "image latents should be provided for T2V to I2V task"
|
||||
if rank_in_sp_group == 0:
|
||||
logger.info("latent_model_input.shape: %s",
|
||||
latent_model_input.shape)
|
||||
latent_model_input = torch.cat([
|
||||
batch.image_latent,
|
||||
latent_model_input[:, :, 1:, :, :],
|
||||
],
|
||||
dim=2).to(target_dtype)
|
||||
logger.info("latent_model_input.shape: %s",
|
||||
latent_model_input.shape)
|
||||
|
||||
assert not torch.isnan(
|
||||
latent_model_input).any(), "latent_model_input contains nan"
|
||||
|
||||
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
|
||||
timestep = torch.stack([t]).to(get_local_torch_device())
|
||||
temp_ts = (mask2[0][0][:, ::2, ::2] * timestep).flatten()
|
||||
@@ -296,6 +319,13 @@ class DenoisingStage(PipelineStage):
|
||||
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
if rank_in_sp_group == 0:
|
||||
latent_model_input = torch.cat([
|
||||
batch.image_latent,
|
||||
latent_model_input[:, :, 1:, :, :],
|
||||
],
|
||||
dim=2).to(target_dtype)
|
||||
|
||||
# Prepare inputs for transformer
|
||||
guidance_expand = (
|
||||
@@ -427,6 +457,12 @@ class DenoisingStage(PipelineStage):
|
||||
latents = (1. - mask2[0]) * z + mask2[0] * latents
|
||||
# latents = latents.unsqueeze(0)
|
||||
|
||||
# save trajectory latents if needed
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_timesteps.append(t)
|
||||
# trajectory_latents.append(latents.cpu())
|
||||
trajectory_latents.append(latents)
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
(i + 1) > num_warmup_steps and
|
||||
@@ -435,8 +471,32 @@ class DenoisingStage(PipelineStage):
|
||||
progress_bar.update()
|
||||
|
||||
# Gather results if using sequence parallelism
|
||||
trajectory_tensor: torch.Tensor | None = None
|
||||
if trajectory_latents:
|
||||
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
|
||||
else:
|
||||
trajectory_tensor = None
|
||||
|
||||
if sp_group:
|
||||
latents = sequence_model_parallel_all_gather(latents, dim=2)
|
||||
if batch.return_trajectory_latents:
|
||||
# logger.info("before stack trajectory_latents.shape: %s", trajectory_latents[0].shape)
|
||||
logger.info("after stack trajectory_latents.shape: %s", trajectory_tensor.shape)
|
||||
trajectory_tensor = trajectory_tensor.to(
|
||||
get_local_torch_device())
|
||||
trajectory_tensor = sequence_model_parallel_all_gather(
|
||||
trajectory_tensor, dim=3)
|
||||
|
||||
if trajectory_tensor is not None:
|
||||
batch.trajectory_timesteps = torch.tensor(trajectory_timesteps).cpu()
|
||||
batch.trajectory_latents = trajectory_tensor.cpu()
|
||||
|
||||
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
latents = torch.cat([
|
||||
batch.image_latent,
|
||||
latents[:, :, 1:, :, :],
|
||||
],
|
||||
dim=2)
|
||||
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
@@ -105,6 +105,81 @@ class ImageVAEEncodingStage(PipelineStage):
|
||||
def __init__(self, vae: ParallelTiledVAE) -> None:
|
||||
self.vae: ParallelTiledVAE = vae
|
||||
|
||||
def encode_image(self,
|
||||
image: PIL.Image.Image,
|
||||
height: int,
|
||||
width: int,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
generator: torch.Generator | None = None) -> torch.Tensor:
|
||||
"""
|
||||
Encode image into latent space.
|
||||
"""
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=height,
|
||||
width=width).to(get_local_torch_device(), dtype=torch.float32)
|
||||
|
||||
# (B, C, H, W) -> (B, C, 1, H, W)
|
||||
print(f"image.shape: {image.shape}")
|
||||
image = image.unsqueeze(2)
|
||||
print(f"after unsqueeze image.shape: {image.shape}")
|
||||
return self.encode_tensor(image, fastvideo_args, generator)
|
||||
|
||||
def encode_tensor(self,
|
||||
video_condition: torch.Tensor,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
generator: torch.Generator | None = None) -> torch.Tensor:
|
||||
"""
|
||||
Encode frames into latent space.
|
||||
"""
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
video_condition = video_condition.to(device=get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Encode Image
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
video_condition = video_condition.to(vae_dtype)
|
||||
encoder_output = self.vae.encode(video_condition)
|
||||
|
||||
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
|
||||
latent_condition = encoder_output.mean
|
||||
else:
|
||||
generator = generator
|
||||
if generator is None:
|
||||
raise ValueError("Generator must be provided")
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator)
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latent_condition -= self.vae.shift_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition -= self.vae.shift_factor
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latent_condition = latent_condition * self.vae.scaling_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition = latent_condition * self.vae.scaling_factor
|
||||
|
||||
return latent_condition
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
@@ -157,58 +232,29 @@ class ImageVAEEncodingStage(PipelineStage):
|
||||
# (B, C, H, W) -> (B, C, 1, H, W)
|
||||
image = image.unsqueeze(2)
|
||||
|
||||
video_condition = torch.cat([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1], num_frames - 1,
|
||||
image.shape[3], image.shape[4])
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
# repeat the image self.vae.temporal_compression_ratio times
|
||||
video_condition = image.repeat(1, 1,
|
||||
self.vae.temporal_compression_ratio,
|
||||
1, 1)
|
||||
# video_condition = image
|
||||
logger.info("video_condition.shape: %s", video_condition.shape)
|
||||
else:
|
||||
video_condition = torch.cat([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1], num_frames - 1,
|
||||
image.shape[3], image.shape[4])
|
||||
],
|
||||
dim=2)
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Encode Image
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
video_condition = video_condition.to(vae_dtype)
|
||||
encoder_output = self.vae.encode(video_condition)
|
||||
|
||||
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
|
||||
latent_condition = encoder_output.mean
|
||||
else:
|
||||
generator = batch.generator
|
||||
if generator is None:
|
||||
raise ValueError("Generator must be provided")
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator)
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latent_condition -= self.vae.shift_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition -= self.vae.shift_factor
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latent_condition = latent_condition * self.vae.scaling_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition = latent_condition * self.vae.scaling_factor
|
||||
latent_condition = self.encode_tensor(video_condition, fastvideo_args,
|
||||
batch.generator)
|
||||
|
||||
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
|
||||
batch.image_latent = latent_condition
|
||||
elif fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
logger.info("latent_condition.shape: %s", latent_condition.shape)
|
||||
batch.image_latent = latent_condition
|
||||
else:
|
||||
mask_lat_size = torch.ones(1, 1, num_frames, latent_height,
|
||||
latent_width)
|
||||
|
||||
@@ -35,9 +35,15 @@ class InputValidationStage(PipelineStage):
|
||||
"""Generate seeds for the inference"""
|
||||
seed = batch.seed
|
||||
num_videos_per_prompt = batch.num_videos_per_prompt
|
||||
if isinstance(batch.prompt, list):
|
||||
num_prompts = len(batch.prompt)
|
||||
else:
|
||||
num_prompts = 1
|
||||
|
||||
total_num_videos = num_prompts * num_videos_per_prompt
|
||||
|
||||
assert seed is not None
|
||||
seeds = [seed + i for i in range(num_videos_per_prompt)]
|
||||
seeds = [seed + i for i in range(total_num_videos)]
|
||||
batch.seeds = seeds
|
||||
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
|
||||
batch.generator = [
|
||||
|
||||
@@ -82,8 +82,8 @@ def rocm_platform_plugin() -> str | None:
|
||||
logger.info("ROCm platform is available")
|
||||
finally:
|
||||
amdsmi.amdsmi_shut_down()
|
||||
except Exception as e:
|
||||
logger.info("ROCm platform is unavailable: %s", e)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return "fastvideo.platforms.rocm.RocmPlatform" if is_rocm else None
|
||||
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
from fastvideo.dataset.dataloader.parquet_io import (
|
||||
ParquetDatasetWriter,
|
||||
records_to_table,
|
||||
)
|
||||
|
||||
|
||||
def test_records_to_table_types():
|
||||
schema = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
pa.field("vae_latent_bytes", pa.binary()),
|
||||
pa.field("vae_latent_shape", pa.list_(pa.int64())),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("width", pa.int64()),
|
||||
])
|
||||
records = [{
|
||||
"id": "a",
|
||||
"vae_latent_bytes": b"\x00\x01",
|
||||
"vae_latent_shape": [1, 2, 3],
|
||||
"duration_sec": 1.5,
|
||||
"width": 640,
|
||||
}]
|
||||
|
||||
table = records_to_table(records, schema)
|
||||
assert table.schema == schema
|
||||
assert table.num_rows == 1
|
||||
cols = {name: table.column(name).to_pylist()[0] for name in schema.names}
|
||||
assert cols["id"] == "a"
|
||||
assert isinstance(cols["vae_latent_bytes"], (bytes, bytearray))
|
||||
assert cols["vae_latent_shape"] == [1, 2, 3]
|
||||
assert abs(cols["duration_sec"] - 1.5) < 1e-6
|
||||
assert cols["width"] == 640
|
||||
|
||||
|
||||
def test_writer_flush_and_remainder(tmp_path: Path):
|
||||
schema = pa.schema([pa.field("id", pa.string())])
|
||||
records = [{"id": str(i)} for i in range(25)]
|
||||
table = records_to_table(records, schema)
|
||||
|
||||
out_dir = tmp_path / "out"
|
||||
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
|
||||
writer.append_table(table)
|
||||
written = writer.flush(num_workers=1)
|
||||
assert written == 20
|
||||
|
||||
files = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files) == 2
|
||||
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
|
||||
assert total_rows == 20
|
||||
|
||||
# Append remainder to complete another chunk
|
||||
extra = records_to_table([{"id": str(i)} for i in range(5)], schema)
|
||||
writer.append_table(extra)
|
||||
written2 = writer.flush(num_workers=1)
|
||||
assert written2 == 10
|
||||
files2 = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files2) == 3
|
||||
total_rows2 = sum(pq.read_table(str(f)).num_rows for f in files2)
|
||||
assert total_rows2 == 30
|
||||
|
||||
|
||||
def test_writer_flush_write_remainder(tmp_path: Path):
|
||||
schema = pa.schema([pa.field("id", pa.string())])
|
||||
# 25 rows, 10 per file => 2 full files + 1 remainder(5)
|
||||
records = [{"id": str(i)} for i in range(25)]
|
||||
table = records_to_table(records, schema)
|
||||
|
||||
out_dir = tmp_path / "out_last"
|
||||
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
|
||||
writer.append_table(table)
|
||||
# First flush writes 20
|
||||
written1 = writer.flush(num_workers=1)
|
||||
assert written1 == 20
|
||||
# Final flush with remainder
|
||||
written2 = writer.flush(num_workers=1, write_remainder=True)
|
||||
assert written2 == 5
|
||||
files = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files) == 3
|
||||
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
|
||||
assert total_rows == 25
|
||||
|
||||
|
||||
def test_writer_parallel_workers(tmp_path: Path):
|
||||
schema = pa.schema([pa.field("id", pa.string())])
|
||||
# 40 rows, 10 per file => 4 files
|
||||
records = [{"id": str(i)} for i in range(40)]
|
||||
table = records_to_table(records, schema)
|
||||
|
||||
out_dir = tmp_path / "out_parallel"
|
||||
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
|
||||
writer.append_table(table)
|
||||
written = writer.flush(num_workers=2)
|
||||
assert written == 40
|
||||
|
||||
# Ensure files exist under worker subdirs
|
||||
worker_dirs = [p for p in out_dir.iterdir() if p.is_dir() and p.name.startswith("worker_")]
|
||||
assert len(worker_dirs) >= 1
|
||||
files = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files) == 4
|
||||
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
|
||||
assert total_rows == 40
|
||||
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.dataset.dataloader.record_schema import (
|
||||
basic_t2v_record_creator,
|
||||
i2v_record_creator,
|
||||
ode_text_only_record_creator,
|
||||
text_only_record_creator,
|
||||
)
|
||||
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
|
||||
|
||||
|
||||
def _mk_basic_batch(N: int) -> PreprocessBatch:
|
||||
batch = PreprocessBatch(data_type="video")
|
||||
batch.video_file_name = [f"vid_{i}" for i in range(N)]
|
||||
batch.prompt = [f"caption_{i}" for i in range(N)]
|
||||
batch.width = [640 for _ in range(N)]
|
||||
batch.height = [360 for _ in range(N)]
|
||||
batch.fps = [4 for _ in range(N)]
|
||||
batch.num_frames = [2 for _ in range(N)]
|
||||
# Latents: shape (N, C, T, H, W); per-record use latents[idx]
|
||||
batch.latents = np.zeros((N, 4, 2, 8, 8), dtype=np.float32)
|
||||
# Prompt embeds: list of per-record arrays [Seq, Dim]
|
||||
batch.prompt_embeds = [np.ones((6, 16), dtype=np.float32) for _ in range(N)]
|
||||
return batch
|
||||
|
||||
|
||||
def test_basic_t2v_record_creator_fields():
|
||||
N = 2
|
||||
batch = _mk_basic_batch(N)
|
||||
|
||||
records = basic_t2v_record_creator(batch)
|
||||
assert isinstance(records, list) and len(records) == N
|
||||
|
||||
for i, rec in enumerate(records):
|
||||
assert rec["id"] == batch.video_file_name[i]
|
||||
# Latents bytes/shape/dtype
|
||||
assert isinstance(rec["vae_latent_bytes"], (bytes, bytearray))
|
||||
assert rec["vae_latent_shape"] == list(batch.latents[i].shape)
|
||||
assert rec["vae_latent_dtype"] == str(batch.latents[i].dtype)
|
||||
# Text embedding
|
||||
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
|
||||
assert rec["text_embedding_shape"] == list(batch.prompt_embeds[i].shape)
|
||||
assert rec["text_embedding_dtype"] == str(batch.prompt_embeds[i].dtype)
|
||||
# Meta
|
||||
assert rec["caption"] == batch.prompt[i]
|
||||
assert rec["media_type"] == "video"
|
||||
assert rec["width"] == int(batch.width[i])
|
||||
assert rec["height"] == int(batch.height[i])
|
||||
assert rec["num_frames"] == batch.latents[i].shape[1]
|
||||
|
||||
|
||||
def test_i2v_record_creator_additional_fields():
|
||||
N = 3
|
||||
batch = _mk_basic_batch(N)
|
||||
# image_embeds is a list of length 1, with an array of shape [N, D]
|
||||
batch.image_embeds = [np.ones((N, 32), dtype=np.float32)]
|
||||
# first frame latent per record
|
||||
batch.image_latent = np.zeros((N, 4, 1, 8, 8), dtype=np.float32)
|
||||
# pil image per record
|
||||
batch.pil_image = np.zeros((N, 8, 8, 3), dtype=np.uint8)
|
||||
|
||||
records = i2v_record_creator(batch)
|
||||
assert isinstance(records, list) and len(records) == N
|
||||
|
||||
for i, rec in enumerate(records):
|
||||
# clip feature
|
||||
assert isinstance(rec["clip_feature_bytes"], (bytes, bytearray))
|
||||
assert rec["clip_feature_shape"] == list(batch.image_embeds[0][i].shape)
|
||||
assert rec["clip_feature_dtype"] == str(batch.image_embeds[0][i].dtype)
|
||||
# first frame latent
|
||||
assert isinstance(rec["first_frame_latent_bytes"], (bytes, bytearray))
|
||||
assert rec["first_frame_latent_shape"] == list(batch.image_latent[i].shape)
|
||||
assert rec["first_frame_latent_dtype"] == str(batch.image_latent[i].dtype)
|
||||
# pil image
|
||||
assert isinstance(rec["pil_image_bytes"], (bytes, bytearray))
|
||||
assert rec["pil_image_shape"] == list(batch.pil_image[i].shape)
|
||||
assert rec["pil_image_dtype"] == str(batch.pil_image[i].dtype)
|
||||
|
||||
|
||||
def test_ode_text_only_record_creator():
|
||||
video_name = "ex"
|
||||
caption = "a prompt"
|
||||
text_embedding = np.ones((6, 16), dtype=np.float32)
|
||||
traj = np.ones((5, 4, 2, 2), dtype=np.float32)
|
||||
tsteps = np.arange(5, dtype=np.float32)
|
||||
|
||||
rec = ode_text_only_record_creator(
|
||||
video_name=video_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=caption,
|
||||
trajectory_latents=traj,
|
||||
trajectory_timesteps=tsteps,
|
||||
)
|
||||
assert rec["id"] == f"text_{video_name}"
|
||||
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
|
||||
assert rec["text_embedding_shape"] == list(text_embedding.shape)
|
||||
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
|
||||
assert rec["file_name"] == video_name
|
||||
assert rec["caption"] == caption
|
||||
assert rec["media_type"] == "text"
|
||||
# Trajectory fields
|
||||
assert isinstance(rec["trajectory_latents_bytes"], (bytes, bytearray))
|
||||
assert rec["trajectory_latents_shape"] == list(traj.shape)
|
||||
assert rec["trajectory_latents_dtype"] == str(traj.dtype)
|
||||
assert isinstance(rec["trajectory_timesteps_bytes"], (bytes, bytearray))
|
||||
assert rec["trajectory_timesteps_shape"] == list(tsteps.shape)
|
||||
assert rec["trajectory_timesteps_dtype"] == str(tsteps.dtype)
|
||||
|
||||
|
||||
def test_text_only_record_creator():
|
||||
text_name = "note1"
|
||||
caption = "a prompt"
|
||||
text_embedding = np.ones((7, 16), dtype=np.float32)
|
||||
rec = text_only_record_creator(
|
||||
text_name=text_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=caption,
|
||||
)
|
||||
assert rec["id"] == f"text_{text_name}"
|
||||
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
|
||||
assert rec["text_embedding_shape"] == list(text_embedding.shape)
|
||||
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
|
||||
assert rec["caption"] == caption
|
||||
@@ -117,3 +117,7 @@ def run_inference_lora_tests():
|
||||
@app.function(gpu="L40S:2", image=image, timeout=900)
|
||||
def run_distill_dmd_tests():
|
||||
run_test("pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_unit_test():
|
||||
run_test("pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ -vs")
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
|
||||
from fastvideo.workflow.preprocess.components import ParquetDatasetSaver
|
||||
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
|
||||
|
||||
|
||||
def _simple_record_creator(batch: PreprocessBatch) -> list[dict]:
|
||||
# batch.latents will be converted to numpy by the saver before this call
|
||||
assert isinstance(batch.latents, np.ndarray)
|
||||
num = len(batch.video_file_name)
|
||||
records = []
|
||||
for i in range(num):
|
||||
arr = batch.latents[i]
|
||||
records.append({
|
||||
"id": batch.video_file_name[i],
|
||||
"data_bytes": arr.tobytes(),
|
||||
"data_shape": list(arr.shape),
|
||||
})
|
||||
return records
|
||||
|
||||
|
||||
def test_parquet_dataset_saver_flush_and_last(tmp_path: Path):
|
||||
# Schema for the simple record creator
|
||||
schema = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
pa.field("data_bytes", pa.binary()),
|
||||
pa.field("data_shape", pa.list_(pa.int64())),
|
||||
])
|
||||
|
||||
B = 5
|
||||
# Build a minimal PreprocessBatch
|
||||
batch = PreprocessBatch(
|
||||
data_type="video",
|
||||
latents=torch.randn(B, 2),
|
||||
prompt_embeds=[torch.randn(B, 1, 1)],
|
||||
# Attention mask should be integer dtype in real pipelines
|
||||
prompt_attention_mask=[torch.ones(B, 1, dtype=torch.int64)],
|
||||
)
|
||||
batch.video_file_name = [f"vid_{i}" for i in range(B)]
|
||||
|
||||
saver = ParquetDatasetSaver(
|
||||
flush_frequency=10, # higher than B to avoid auto-flush
|
||||
samples_per_file=3,
|
||||
schema=schema,
|
||||
record_creator=_simple_record_creator,
|
||||
)
|
||||
|
||||
out_dir = tmp_path / "saver_out"
|
||||
saver.save_and_write_parquet_batch(batch, str(out_dir))
|
||||
# First flush: should write one full file (3 rows), keep 2 in buffer
|
||||
saver.flush_tables()
|
||||
files = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files) == 1
|
||||
assert pq.read_table(str(files[0])).num_rows == 3
|
||||
|
||||
# Final flush: write remainder 2 rows
|
||||
saver.flush_tables(write_remainder=True)
|
||||
files2 = sorted(out_dir.rglob("*.parquet"))
|
||||
assert len(files2) == 2
|
||||
total = sum(pq.read_table(str(f)).num_rows for f in files2)
|
||||
assert total == 5
|
||||
|
||||
|
||||
@@ -40,7 +40,7 @@ from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.training.training_utils import (
|
||||
EMA_FSDP, clip_grad_norm_while_handling_failing_dtensor_cases,
|
||||
get_scheduler, load_distillation_checkpoint, save_distillation_checkpoint,
|
||||
shift_timestep, compute_density_for_timestep_sampling, get_sigmas)
|
||||
shift_timestep)
|
||||
from fastvideo.utils import (is_vsa_available, maybe_download_model,
|
||||
set_random_seed, verify_model_config_and_directory)
|
||||
|
||||
@@ -91,9 +91,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
|
||||
shift=self.timestep_shift)
|
||||
|
||||
self.transformer_2 = self.get_module("transformer_2", None)
|
||||
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
|
||||
|
||||
if training_args.real_score_model_path:
|
||||
logger.info(
|
||||
f"Loading real score transformer from: {training_args.real_score_model_path}"
|
||||
@@ -130,39 +127,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
self.real_score_transformer,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
if self.transformer_2 is not None:
|
||||
self.transformer_2 = apply_activation_checkpointing(
|
||||
self.transformer_2,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
|
||||
if self.transformer_2 is not None:
|
||||
self.transformer_2.train()
|
||||
self.transformer_2.requires_grad_(True)
|
||||
params_to_optimize_2 = self.transformer_2.parameters()
|
||||
params_to_optimize_2 = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize_2))
|
||||
|
||||
betas_str = training_args.betas
|
||||
betas = tuple(float(x.strip()) for x in betas_str.split(","))
|
||||
|
||||
self.optimizer_2 = torch.optim.AdamW(
|
||||
params_to_optimize_2,
|
||||
lr=training_args.learning_rate,
|
||||
betas=betas,
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
self.lr_scheduler_2 = get_scheduler(
|
||||
training_args.lr_scheduler,
|
||||
optimizer=self.optimizer_2,
|
||||
num_warmup_steps=training_args.lr_warmup_steps,
|
||||
num_training_steps=training_args.max_train_steps,
|
||||
num_cycles=training_args.lr_num_cycles,
|
||||
power=training_args.lr_power,
|
||||
min_lr_ratio=training_args.min_lr_ratio,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
# Initialize optimizers
|
||||
fake_score_params = list(
|
||||
@@ -317,9 +281,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"""Prepare training environment for distillation."""
|
||||
self.transformer.requires_grad_(True)
|
||||
self.transformer.train()
|
||||
if self.transformer_2 is not None:
|
||||
self.transformer_2.requires_grad_(True)
|
||||
self.transformer_2.train()
|
||||
self.fake_score_transformer.requires_grad_(True)
|
||||
self.fake_score_transformer.train()
|
||||
|
||||
@@ -475,9 +436,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_latent, timestep, training_batch.conditional_dict,
|
||||
training_batch)
|
||||
|
||||
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
|
||||
pred_noise = current_model(**training_batch.input_kwargs).permute(
|
||||
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
|
||||
0, 2, 1, 3, 4)
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise.flatten(0, 1),
|
||||
@@ -519,7 +478,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
max_target_idx = len(self.denoising_step_list) - 1
|
||||
noise_latents = []
|
||||
noise_latent_index = target_timestep_idx_int - 1
|
||||
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
|
||||
if max_target_idx > 0:
|
||||
# Run student model for all steps before the target timestep
|
||||
with torch.no_grad():
|
||||
@@ -531,7 +489,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
current_noise_latents, current_timestep_tensor,
|
||||
training_batch.conditional_dict, training_batch)
|
||||
pred_flow = current_model(
|
||||
pred_flow = self.transformer(
|
||||
**training_batch_temp.input_kwargs).permute(
|
||||
0, 2, 1, 3, 4)
|
||||
pred_clean = pred_noise_to_pred_video(
|
||||
@@ -574,7 +532,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
noisy_input, target_timestep, training_batch.conditional_dict,
|
||||
training_batch)
|
||||
pred_noise = current_model(**training_batch.input_kwargs).permute(
|
||||
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
|
||||
0, 2, 1, 3, 4)
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise.flatten(0, 1),
|
||||
@@ -772,12 +730,11 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"encoder_hidden_states": training_batch.encoder_hidden_states,
|
||||
"encoder_attention_mask": training_batch.encoder_attention_mask,
|
||||
}
|
||||
if getattr(self, "negative_prompt_embeds", None) is not None:
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
|
||||
training_batch.dmd_latent_vis_dict = {}
|
||||
training_batch.fake_score_latent_vis_dict = {}
|
||||
@@ -816,9 +773,6 @@ class DistillationPipeline(TrainingPipeline):
|
||||
batches.append(batch)
|
||||
|
||||
self.optimizer.zero_grad()
|
||||
# TODO: confirm this
|
||||
if self.transformer_2 is not None:
|
||||
self.optimizer_2.zero_grad()
|
||||
total_dmd_loss = 0.0
|
||||
dmd_latent_vis_dict = {}
|
||||
fake_score_latent_vis_dict = {}
|
||||
@@ -847,31 +801,15 @@ class DistillationPipeline(TrainingPipeline):
|
||||
attn_metadata=batch_gen.attn_metadata_vsa):
|
||||
(dmd_loss / gradient_accumulation_steps).backward()
|
||||
total_dmd_loss += dmd_loss.detach().item()
|
||||
|
||||
# Only clip gradients for the model that is currently training
|
||||
if self.train_transformer_2 and self.transformer_2 is not None:
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer_2)
|
||||
for param in self.transformer_2.parameters():
|
||||
# check if the gradient is not None and not zero
|
||||
assert param.grad is not None and param.grad.abs().sum() > 0
|
||||
self.optimizer_2.step()
|
||||
self.optimizer_2.zero_grad(set_to_none=True)
|
||||
else:
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer)
|
||||
for param in self.transformer.parameters():
|
||||
# check if the gradient is not None and not zero
|
||||
assert param.grad is not None and param.grad.abs().sum() > 0
|
||||
self.optimizer.step()
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
self._clip_model_grad_norm_(batch_gen, self.transformer)
|
||||
for param in self.transformer.parameters():
|
||||
# check if the gradient is not None and not zero
|
||||
assert param.grad is not None and param.grad.abs().sum() > 0
|
||||
self.optimizer.step()
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
if self.generator_ema is not None:
|
||||
# TODO: support EMA for transformer_2?
|
||||
if self.train_transformer_2 and self.transformer_2 is not None:
|
||||
# Note: EMA currently only supports the main transformer
|
||||
# Could be extended to support transformer_2 in the future
|
||||
pass
|
||||
else:
|
||||
self.generator_ema.update(self.transformer)
|
||||
self.generator_ema.update(self.transformer)
|
||||
|
||||
avg_dmd_loss = torch.tensor(total_dmd_loss /
|
||||
gradient_accumulation_steps,
|
||||
@@ -901,13 +839,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
assert param.grad is not None and param.grad.abs().sum() > 0
|
||||
self.fake_score_optimizer.step()
|
||||
self.fake_score_lr_scheduler.step()
|
||||
|
||||
# Step the appropriate scheduler
|
||||
if self.train_transformer_2 and self.transformer_2 is not None:
|
||||
self.lr_scheduler_2.step()
|
||||
else:
|
||||
self.lr_scheduler.step()
|
||||
|
||||
self.lr_scheduler.step()
|
||||
self.fake_score_optimizer.zero_grad(set_to_none=True)
|
||||
avg_fake_score_loss = torch.tensor(total_fake_score_loss /
|
||||
gradient_accumulation_steps,
|
||||
|
||||
@@ -0,0 +1,476 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from typing import cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
import wandb
|
||||
from fastvideo.dataset.dataloader.schema import (
|
||||
pyarrow_schema_ode_trajectory_text_only)
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
|
||||
SelfForcingFlowMatchScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
|
||||
WanCausalDMDPipeline)
|
||||
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.training.training_utils import (
|
||||
clip_grad_norm_while_handling_failing_dtensor_cases)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
Training pipeline for ODE-init using precomputed denoising trajectories.
|
||||
|
||||
Supervision: predict the next latent in the stored trajectory by
|
||||
- feeding current latent at timestep t into the transformer to predict noise
|
||||
- stepping the scheduler with the predicted noise
|
||||
- minimizing MSE to the stored next latent at timestep t_next
|
||||
"""
|
||||
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# Match the preprocess/generation scheduler for consistent stepping
|
||||
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift,
|
||||
sigma_min=0.0,
|
||||
extra_one_step=True)
|
||||
self.modules["scheduler"].set_timesteps(num_inference_steps=1000,
|
||||
training=True)
|
||||
|
||||
def set_schemas(self):
|
||||
self.train_dataset_schema = pyarrow_schema_ode_trajectory_text_only
|
||||
|
||||
def initialize_training_pipeline(self, training_args: TrainingArgs):
|
||||
super().initialize_training_pipeline(training_args)
|
||||
|
||||
self.noise_scheduler = self.get_module("scheduler")
|
||||
self.vae = self.get_module("vae")
|
||||
self.vae.requires_grad_(False)
|
||||
|
||||
self.timestep_shift = self.training_args.pipeline_config.flow_shift
|
||||
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
|
||||
self.noise_scheduler = SelfForcingFlowMatchScheduler(
|
||||
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
|
||||
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
|
||||
training=True)
|
||||
|
||||
# logger.info(f"ARG dmd_denoising_steps: {training_args.pipeline_config.dmd_denoising_steps}")
|
||||
logger.info(
|
||||
f"ARG dmd_denoising_steps: {self.training_args.pipeline_config.dmd_denoising_steps}"
|
||||
)
|
||||
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
# self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250], dtype=torch.long, device=get_local_torch_device())
|
||||
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
|
||||
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
|
||||
torch.tensor([0],
|
||||
dtype=torch.float32))).cuda()
|
||||
logger.info(f"timesteps: {timesteps}")
|
||||
self.dmd_denoising_steps = timesteps[1000 -
|
||||
self.dmd_denoising_steps]
|
||||
logger.info(
|
||||
f"warped self.dmd_denoising_steps: {self.dmd_denoising_steps}")
|
||||
# assert False, "warp_denoising_step must be false"
|
||||
else:
|
||||
assert False, "warp_denoising_step must be true"
|
||||
logger.info("not warped")
|
||||
self.dmd_denoising_steps = self.dmd_denoising_steps.to(
|
||||
get_local_torch_device())
|
||||
|
||||
logger.info(f"denoising_step_list: {self.dmd_denoising_steps}")
|
||||
|
||||
logger.info(
|
||||
"Initialized ODE-init training pipeline with %s denoising steps",
|
||||
len(self.dmd_denoising_steps))
|
||||
# Cache for nearest trajectory index per DMD step (computed lazily on first batch)
|
||||
self._cached_closest_idx_per_dmd = None
|
||||
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
|
||||
# self.min_timestep = int(self.training_args.min_timestep_ratio *
|
||||
# self.num_train_timestep)
|
||||
# self.max_timestep = int(self.training_args.max_timestep_ratio *
|
||||
# self.num_train_timestep)
|
||||
# self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
|
||||
self.manual_idx = 0
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
args_copy.inference_mode = True
|
||||
# Warm start validation with current transformer
|
||||
self.validation_pipeline = WanCausalDMDPipeline.from_pretrained(
|
||||
# training_args.model_path,
|
||||
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
|
||||
args=args_copy, # type: ignore
|
||||
inference_mode=True,
|
||||
loaded_modules={
|
||||
"transformer": self.get_module("transformer"),
|
||||
},
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
pin_cpu_memory=training_args.pin_cpu_memory,
|
||||
dit_cpu_offload=True)
|
||||
|
||||
def _get_next_batch(self, training_batch): # type: ignore[override]
|
||||
batch = next(self.train_loader_iter, None) # type: ignore
|
||||
if batch is None:
|
||||
self.current_epoch += 1
|
||||
logger.info("Starting epoch %s", self.current_epoch)
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
batch = next(self.train_loader_iter)
|
||||
|
||||
# Required fields from parquet (ODE trajectory schema)
|
||||
encoder_hidden_states = batch['text_embedding']
|
||||
encoder_attention_mask = batch['text_attention_mask']
|
||||
infos = batch['info_list']
|
||||
|
||||
# Trajectory tensors may include a leading singleton batch dim per row
|
||||
trajectory_latents = batch['trajectory_latents']
|
||||
if trajectory_latents.dim() == 7:
|
||||
# [B, 1, S, C, T, H, W] -> [B, S, C, T, H, W]
|
||||
trajectory_latents = trajectory_latents[:, 0]
|
||||
elif trajectory_latents.dim() == 6:
|
||||
# already [B, S, C, T, H, W]
|
||||
pass
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unexpected trajectory_latents dim: {trajectory_latents.dim()}"
|
||||
)
|
||||
|
||||
trajectory_timesteps = batch['trajectory_timesteps']
|
||||
if trajectory_timesteps.dim() == 3:
|
||||
# [B, 1, S] -> [B, S]
|
||||
trajectory_timesteps = trajectory_timesteps[:, 0]
|
||||
elif trajectory_timesteps.dim() == 2:
|
||||
# [B, S]
|
||||
pass
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unexpected trajectory_timesteps dim: {trajectory_timesteps.dim()}"
|
||||
)
|
||||
# [B, S, C, T, H, W] -> [B, S, T, C, H, W] to match self-forcing
|
||||
trajectory_latents = trajectory_latents.permute(0, 1, 3, 2, 4, 5)
|
||||
|
||||
# Move to device
|
||||
device = get_local_torch_device()
|
||||
# training_batch.encoder_hidden_states = encoder_hidden_states.to(
|
||||
# device, dtype=torch.bfloat16)
|
||||
# training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
# device, dtype=torch.bfloat16)
|
||||
# training_batch.infos = infos
|
||||
|
||||
# return training_batch, trajectory_latents.to(
|
||||
# device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
|
||||
|
||||
## TEMP
|
||||
self.manual_idx = self.manual_idx % 55
|
||||
path = f"/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_pt_vidprom_1000/{self.manual_idx:05d}.pt"
|
||||
logger.info(f"path: {path}")
|
||||
self.manual_idx += 1
|
||||
# path = "/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_single_full/00000.pt"
|
||||
b = torch.load(path)
|
||||
for k, v in b.items():
|
||||
logger.info(f"b[{k}]: {type(v)}")
|
||||
if isinstance(v, torch.Tensor):
|
||||
logger.info(f"b[{k}]: {v.shape}")
|
||||
else:
|
||||
logger.info(f"b[{k}]: {v}")
|
||||
training_batch.encoder_hidden_states = b["text_embedding"][0].unsqueeze(
|
||||
0).to(device, dtype=torch.bfloat16)
|
||||
trajectory_latents = b["ode_latent"].to(device, dtype=torch.bfloat16)
|
||||
logger.info(f"trajectory_latents: {trajectory_latents.shape}")
|
||||
logger.info(
|
||||
f"encoder_hidden_states: {training_batch.encoder_hidden_states.shape}"
|
||||
)
|
||||
return training_batch, trajectory_latents.to(
|
||||
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
|
||||
|
||||
def _get_timestep(self,
|
||||
min_timestep: int,
|
||||
max_timestep: int,
|
||||
batch_size: int,
|
||||
num_frame: int,
|
||||
num_frame_per_block: int,
|
||||
uniform_timestep: bool = False) -> torch.Tensor:
|
||||
if uniform_timestep:
|
||||
timestep = torch.randint(min_timestep,
|
||||
max_timestep, [batch_size, 1],
|
||||
device=self.device,
|
||||
dtype=torch.long).repeat(1, num_frame)
|
||||
return timestep
|
||||
else:
|
||||
timestep = torch.randint(min_timestep,
|
||||
max_timestep, [batch_size, num_frame],
|
||||
device=self.device,
|
||||
dtype=torch.long)
|
||||
# logger.info(f"individual timestep: {timestep}")
|
||||
# make the noise level the same within every block
|
||||
timestep = timestep.reshape(timestep.shape[0], -1,
|
||||
num_frame_per_block)
|
||||
timestep[:, :, 1:] = timestep[:, :, 0:1]
|
||||
timestep = timestep.reshape(timestep.shape[0], -1)
|
||||
return timestep
|
||||
|
||||
def _step_predict_next_latent(
|
||||
self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_attention_mask: torch.Tensor
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str,
|
||||
torch.Tensor]]:
|
||||
latent_vis_dict = {}
|
||||
device = get_local_torch_device()
|
||||
target_latent = traj_latents[:, -1]
|
||||
|
||||
# logger.info(f"traj_latents: {traj_latents.shape}")
|
||||
# logger.info(f"traj_timesteps: {traj_timesteps.shape}")
|
||||
|
||||
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
|
||||
B, S, num_frames, num_channels, height, width = traj_latents.shape
|
||||
|
||||
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
|
||||
if self._cached_closest_idx_per_dmd is None:
|
||||
# Use the first sample's trajectory timesteps; assumed identical across batches
|
||||
# s_steps = traj_timesteps[0].to(torch.long) # [S]
|
||||
# dmd = cast(torch.Tensor, self.dmd_denoising_steps).to(s_steps.device) # [K]
|
||||
# distances_ks: [K, S] = |s_steps - dmd|
|
||||
# distances_ks = (s_steps.unsqueeze(0) - dmd.unsqueeze(1)).abs()
|
||||
# self._cached_closest_idx_per_dmd = distances_ks.argmin(dim=1).to(torch.long).cpu() # [K]
|
||||
self._cached_closest_idx_per_dmd = torch.tensor(
|
||||
# [0, 12, 24, 36], dtype=torch.long).cpu()
|
||||
[0, 1, 2, 3], dtype=torch.long).cpu()
|
||||
logger.info(
|
||||
f"self._cached_closest_idx_per_dmd: {self._cached_closest_idx_per_dmd}"
|
||||
)
|
||||
logger.info(
|
||||
f"corresponding timesteps: {self.noise_scheduler.timesteps[self._cached_closest_idx_per_dmd]}"
|
||||
)
|
||||
|
||||
# logger.info(f"traj_latents: {traj_latents.shape}")
|
||||
# Select the K indexes from traj_latents using self._cached_closest_idx_per_dmd
|
||||
# traj_latents: [B, S, C, T, H, W], self._cached_closest_idx_per_dmd: [K]
|
||||
# Output: [B, K, C, T, H, W]
|
||||
relevant_traj_latents = torch.index_select(
|
||||
traj_latents,
|
||||
dim=1,
|
||||
index=self._cached_closest_idx_per_dmd.to(traj_latents.device))
|
||||
logger.info(f"relevant_traj_latents: {relevant_traj_latents.shape}")
|
||||
# assert relevant_traj_latents.shape[0] == 1
|
||||
|
||||
indexes = self._get_timestep( # [B, num_frames]
|
||||
0,
|
||||
len(self.dmd_denoising_steps),
|
||||
B,
|
||||
num_frames,
|
||||
3,
|
||||
uniform_timestep=False)
|
||||
logger.info(f"indexes: {indexes.shape}")
|
||||
logger.info(f"indexes: {indexes}")
|
||||
# noisy_input = relevant_traj_latents[indexes]
|
||||
noisy_input = torch.gather(
|
||||
relevant_traj_latents,
|
||||
dim=1,
|
||||
index=indexes.reshape(B, 1, num_frames, 1, 1,
|
||||
1).expand(-1, -1, -1, num_channels, height,
|
||||
width).to(self.device)).squeeze(1)
|
||||
# noisy_input = noisy_input.unsqueeze(0)
|
||||
|
||||
# # Sample a single DMD step for the whole batch and fetch its cached nearest S-index
|
||||
# K = len(self.dmd_denoising_steps)
|
||||
# dmd_idx = torch.randint(0, K, (1,), device=device)
|
||||
# logger.info(f"dmd_idx: {dmd_idx}")
|
||||
# assert self._cached_closest_idx_per_dmd is not None
|
||||
# nearest_s_idx = int(self._cached_closest_idx_per_dmd[int(dmd_idx.item())])
|
||||
# nearest_idx = torch.full((B,), nearest_s_idx, device=device, dtype=torch.long)
|
||||
|
||||
# batch_indices = torch.arange(B, device=device)
|
||||
# noisy_input = traj_latents[batch_indices, nearest_idx] # [B, C, T, H, W]
|
||||
# target_latent = traj_latents[batch_indices, -1] # [B, C, T, H, W]
|
||||
# t = traj_timesteps[batch_indices, nearest_idx] # [B]
|
||||
|
||||
# Scale model input as in inference for consistency with stored trajectories
|
||||
# noisy_input = self.modules["scheduler"].scale_model_input(noisy_input, t)
|
||||
# logger.info(f"indexes: {indexes.shape}")
|
||||
# logger.info(f"indexes: {indexes}")
|
||||
timestep = self.dmd_denoising_steps[indexes]
|
||||
# logger.info(f"timestep: {timestep.shape}")
|
||||
# logger.info(f"timestep: {timestep}")
|
||||
|
||||
# Prepare inputs for transformer
|
||||
latent_vis_dict["noisy_input"] = noisy_input.permute(
|
||||
0, 2, 1, 3, 4).detach().clone().cpu()
|
||||
latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3,
|
||||
4).detach().clone().cpu()
|
||||
|
||||
model_dtype = next(self.transformer.parameters()).dtype
|
||||
input_kwargs = {
|
||||
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timestep.to(device, dtype=model_dtype),
|
||||
"return_dict": False,
|
||||
}
|
||||
# Predict noise and step the scheduler to obtain next latent
|
||||
with set_forward_context(current_timestep=timestep,
|
||||
attn_metadata=None,
|
||||
forward_batch=None):
|
||||
noise_pred = self.transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
# logger.info(f"noise_pred: {noise_pred.shape}")
|
||||
if isinstance(noise_pred, (tuple, list)):
|
||||
noise_pred = noise_pred[0]
|
||||
|
||||
from fastvideo.models.utils import pred_noise_to_pred_video
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=noise_pred.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep.to(dtype=model_dtype).flatten(0, 1),
|
||||
scheduler=self.modules["scheduler"]).unflatten(
|
||||
0, noise_pred.shape[:2])
|
||||
latent_vis_dict["pred_video"] = pred_video.permute(
|
||||
0, 2, 1, 3, 4).detach().clone().cpu()
|
||||
|
||||
# noisy_input = pred_noise_to_pred_video(noise_pred, noisy_input, t, self.modules["scheduler"])
|
||||
# next_latent_pred = self.modules["scheduler"].step(
|
||||
# noise_pred, t, current_latents, return_dict=False)[0]
|
||||
return pred_video, target_latent, timestep, latent_vis_dict
|
||||
|
||||
def train_one_step(self, training_batch): # type: ignore[override]
|
||||
self.transformer.train()
|
||||
for name, param in self.transformer.named_parameters():
|
||||
assert param.requires_grad, "FUBAR"
|
||||
self.optimizer.zero_grad()
|
||||
training_batch.total_loss = 0.0
|
||||
args = cast(TrainingArgs, self.training_args)
|
||||
|
||||
# Using cached nearest index per DMD step; computation happens in _step_predict_next_latent
|
||||
|
||||
for _ in range(args.gradient_accumulation_steps):
|
||||
training_batch, traj_latents, traj_timesteps = self._get_next_batch(
|
||||
training_batch)
|
||||
text_embeds = training_batch.encoder_hidden_states
|
||||
text_attention_mask = training_batch.encoder_attention_mask
|
||||
assert traj_latents.shape[0] == 1
|
||||
|
||||
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
|
||||
B, S = traj_latents.shape[0], traj_latents.shape[1]
|
||||
if S < 2:
|
||||
raise ValueError("Trajectory must contain at least 2 steps")
|
||||
|
||||
# Sample per-sample current step i in [0, S-2]
|
||||
|
||||
# idx = torch.randint(low=0, high=S - 1, size=(B, ),
|
||||
# device=traj_latents.device)
|
||||
|
||||
# Gather current latents and next latents
|
||||
# batch_indices = torch.arange(B, device=traj_latents.device)
|
||||
# current_latents = traj_latents[batch_indices, idx] # [B, C, T,H,W]
|
||||
# current_latent = traj_timesteps[:, -1, :, :, :, :]
|
||||
# target_latents = traj_latents[:, -1, :, :, :, :]
|
||||
|
||||
# Corresponding timesteps t (long) -> cast per sample
|
||||
# t = traj_timesteps[:, -1, :, :, :, :]
|
||||
# if t.dtype != torch.long:
|
||||
# t = t.long()
|
||||
|
||||
# Forward to predict next latent by stepping scheduler with predicted noise
|
||||
noise_pred, target_latent, t, latent_vis_dict = self._step_predict_next_latent(
|
||||
traj_latents, traj_timesteps, text_embeds, text_attention_mask)
|
||||
|
||||
training_batch.latent_vis_dict.update(latent_vis_dict)
|
||||
|
||||
mask = t != 0
|
||||
|
||||
# Compute loss
|
||||
loss = F.mse_loss(noise_pred[mask],
|
||||
target_latent[mask],
|
||||
reduction="mean")
|
||||
loss = loss / args.gradient_accumulation_steps
|
||||
|
||||
with set_forward_context(current_timestep=t,
|
||||
attn_metadata=None,
|
||||
forward_batch=None):
|
||||
loss.backward()
|
||||
avg_loss = loss.detach().clone()
|
||||
training_batch.total_loss += avg_loss.item()
|
||||
|
||||
# Clip grad and step optimizers
|
||||
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
|
||||
[p for p in self.transformer.parameters() if p.requires_grad],
|
||||
args.max_grad_norm if args.max_grad_norm is not None else 0.0)
|
||||
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
if grad_norm is None:
|
||||
grad_value = 0.0
|
||||
else:
|
||||
try:
|
||||
if isinstance(grad_norm, torch.Tensor):
|
||||
grad_value = float(grad_norm.detach().float().item())
|
||||
else:
|
||||
grad_value = float(grad_norm)
|
||||
except Exception:
|
||||
grad_value = 0.0
|
||||
training_batch.grad_norm = grad_value
|
||||
return training_batch
|
||||
|
||||
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
|
||||
training_args: TrainingArgs, step: int):
|
||||
"""Add visualization data to wandb logging and save frames to disk."""
|
||||
wandb_loss_dict = {}
|
||||
latents_vis_dict = training_batch.latent_vis_dict
|
||||
latent_log_keys = ['noisy_input', 'x0', 'pred_video']
|
||||
for latent_key in latent_log_keys:
|
||||
assert latent_key in latents_vis_dict and latents_vis_dict[
|
||||
latent_key] is not None
|
||||
latent = latents_vis_dict[latent_key]
|
||||
pixel_latent = self.validation_pipeline.decoding_stage.decode(
|
||||
latent, training_args)
|
||||
|
||||
video = pixel_latent.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
video = (video * 255).numpy().astype(np.uint8)
|
||||
wandb_loss_dict[latent_key] = wandb.Video(
|
||||
video, fps=16, format="mp4") # change to 16 for Wan2.1
|
||||
# Clean up references
|
||||
del video, pixel_latent, latent
|
||||
|
||||
# Log to wandb
|
||||
if self.global_rank == 0:
|
||||
wandb.log(wandb_loss_dict, step=step)
|
||||
|
||||
# dmd_latents_vis_dict = training_batch.dmd_latent_vis_dict
|
||||
# fake_score_latents_vis_dict = training_batch.fake_score_latent_vis_dict
|
||||
# fake_score_log_keys = ['generator_pred_video']
|
||||
# dmd_log_keys = ['faker_score_pred_video', 'real_score_pred_video']
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting ODE-init training pipeline...")
|
||||
logger.info(f"ARG dmd_denoising_steps: {args.dmd_denoising_steps}")
|
||||
pipeline = ODEInitTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("ODE-init training pipeline done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.fastvideo_args import TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.dit_cpu_offload = False
|
||||
main(args)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,5 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import gc
|
||||
import dataclasses
|
||||
import math
|
||||
import os
|
||||
@@ -45,8 +44,7 @@ from fastvideo.training.training_utils import (
|
||||
shard_latents_across_sp)
|
||||
# from fastvideo.utils import (is_vmoba_available, is_vsa_available,
|
||||
# set_random_seed, shallow_asdict)
|
||||
from fastvideo.utils import (is_vsa_available,
|
||||
set_random_seed, shallow_asdict)
|
||||
from fastvideo.utils import is_vsa_available, set_random_seed, shallow_asdict
|
||||
|
||||
import wandb # isort: skip
|
||||
|
||||
@@ -70,7 +68,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
train_dataloader: StatefulDataLoader
|
||||
train_loader_iter: Iterator[dict[str, Any]]
|
||||
current_epoch: int = 0
|
||||
train_transformer_2: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -106,7 +103,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
self.sp_world_size = self.sp_group.world_size
|
||||
self.local_rank = world_group.local_rank
|
||||
self.transformer = self.get_module("transformer")
|
||||
self.transformer_2 = self.get_module("transformer_2", None)
|
||||
self.seed = training_args.seed
|
||||
self.set_schemas()
|
||||
|
||||
@@ -119,11 +115,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
self.transformer,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
if self.transformer_2 is not None:
|
||||
self.transformer_2 = apply_activation_checkpointing(
|
||||
self.transformer_2,
|
||||
checkpointing_type=training_args.
|
||||
enable_gradient_checkpointing_type)
|
||||
|
||||
noise_scheduler = self.modules["scheduler"]
|
||||
self.set_trainable()
|
||||
@@ -133,7 +124,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# Parse betas from string format "beta1,beta2"
|
||||
betas_str = training_args.betas
|
||||
betas = tuple(float(x.strip()) for x in betas_str.split(","))
|
||||
|
||||
|
||||
self.optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
lr=training_args.learning_rate,
|
||||
@@ -155,30 +146,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
min_lr_ratio=training_args.min_lr_ratio,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
if self.transformer_2 is not None:
|
||||
# Ensure transformer_2 has trainable parameters before creating optimizer
|
||||
self.transformer_2.train()
|
||||
self.transformer_2.requires_grad_(True)
|
||||
params_to_optimize_2 = self.transformer_2.parameters()
|
||||
params_to_optimize_2 = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize_2))
|
||||
self.optimizer_2 = torch.optim.AdamW(
|
||||
params_to_optimize_2,
|
||||
lr=training_args.learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=training_args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
self.lr_scheduler_2 = get_scheduler(
|
||||
training_args.lr_scheduler,
|
||||
optimizer=self.optimizer_2,
|
||||
num_warmup_steps=training_args.lr_warmup_steps,
|
||||
num_training_steps=training_args.max_train_steps,
|
||||
num_cycles=training_args.lr_num_cycles,
|
||||
power=training_args.lr_power,
|
||||
min_lr_ratio=training_args.min_lr_ratio,
|
||||
last_epoch=self.init_steps - 1,
|
||||
)
|
||||
|
||||
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
|
||||
training_args.data_path,
|
||||
@@ -193,7 +160,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
seed=self.seed)
|
||||
|
||||
self.noise_scheduler = noise_scheduler
|
||||
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
|
||||
|
||||
self.num_update_steps_per_epoch = math.ceil(
|
||||
len(self.train_dataloader) /
|
||||
training_args.gradient_accumulation_steps * training_args.sp_size /
|
||||
@@ -219,25 +186,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
def _prepare_training(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
self.transformer.train()
|
||||
self.optimizer.zero_grad()
|
||||
if self.transformer_2 is not None:
|
||||
self.transformer_2.train()
|
||||
self.optimizer_2.zero_grad()
|
||||
training_batch.total_loss = 0.0
|
||||
return training_batch
|
||||
|
||||
def _enable_training(self, model: torch.nn.Module, optimizer: torch.optim.Optimizer) -> None:
|
||||
"""Enable training mode and gradients for the specified model."""
|
||||
for param in model.parameters():
|
||||
param.requires_grad = True
|
||||
model.train()
|
||||
optimizer.zero_grad()
|
||||
|
||||
def _disable_training(self, model: torch.nn.Module, optimizer: torch.optim.Optimizer) -> None:
|
||||
"""Disable training mode and gradients for the specified model."""
|
||||
for param in model.parameters():
|
||||
param.requires_grad = False
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
batch = next(self.train_loader_iter, None) # type: ignore
|
||||
if batch is None:
|
||||
@@ -281,17 +232,17 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
generator=self.noise_gen_cuda,
|
||||
device=latents.device,
|
||||
dtype=latents.dtype)
|
||||
timesteps = self._sample_timesteps(batch_size, latents.device)
|
||||
|
||||
# Enable training for the model that will be trained next and disable the other
|
||||
if self.train_transformer_2:
|
||||
self._enable_training(self.transformer_2, self.optimizer_2)
|
||||
self._disable_training(self.transformer, self.optimizer)
|
||||
else:
|
||||
self._enable_training(self.transformer, self.optimizer)
|
||||
if self.transformer_2 is not None:
|
||||
self._disable_training(self.transformer_2, self.optimizer_2)
|
||||
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=self.training_args.weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=self.noise_random_generator,
|
||||
logit_mean=self.training_args.logit_mean,
|
||||
logit_std=self.training_args.logit_std,
|
||||
mode_scale=self.training_args.mode_scale,
|
||||
)
|
||||
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
|
||||
timesteps = self.noise_scheduler.timesteps[indices].to(
|
||||
device=latents.device)
|
||||
if self.training_args.sp_size > 1:
|
||||
# Make sure that the timesteps are the same across all sp processes.
|
||||
sp_group = get_sp_group()
|
||||
@@ -314,38 +265,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
return training_batch
|
||||
|
||||
def _sample_timesteps(self, batch_size, device):
|
||||
# Determine which model to train based on the boundary timestep
|
||||
if (self.transformer_2 is not None and self.boundary_timestep is not None and
|
||||
torch.rand(1, generator=self.noise_random_generator).item() <= self.training_args.boundary_ratio):
|
||||
self.train_transformer_2 = True
|
||||
else:
|
||||
self.train_transformer_2 = False
|
||||
|
||||
# Broadcast the decision to all processes
|
||||
decision = torch.tensor(1.0 if self.train_transformer_2 else 0.0, device=self.device)
|
||||
dist.broadcast(decision, src=0)
|
||||
self.train_transformer_2 = decision.item() == 1.0
|
||||
|
||||
# Sample u from the appropriate range
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=self.training_args.weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=self.noise_random_generator,
|
||||
logit_mean=self.training_args.logit_mean,
|
||||
logit_std=self.training_args.logit_std,
|
||||
mode_scale=self.training_args.mode_scale,
|
||||
)
|
||||
|
||||
boundary_ratio = self.training_args.boundary_ratio
|
||||
if self.train_transformer_2:
|
||||
u = (1 - boundary_ratio) + u * boundary_ratio # min: 1 - boundary_ratio, max: 1
|
||||
else:
|
||||
u = u * (1 - boundary_ratio) # min: 0, max: 1 - boundary_ratio
|
||||
|
||||
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:
|
||||
latents_shape = training_batch.raw_latent_shape
|
||||
@@ -361,6 +280,20 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
patch_size=patch_size,
|
||||
VSA_sparsity=current_vsa_sparsity,
|
||||
device=get_local_torch_device())
|
||||
# elif vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
|
||||
# moba_params = self.training_args.moba_config.copy()
|
||||
# moba_params.update({
|
||||
# "current_timestep":
|
||||
# training_batch.timesteps,
|
||||
# "raw_latent_shape":
|
||||
# training_batch.raw_latent_shape[2:5],
|
||||
# "patch_size":
|
||||
# self.training_args.pipeline_config.dit_config.patch_size,
|
||||
# "device":
|
||||
# get_local_torch_device(),
|
||||
# })
|
||||
# training_batch.attn_metadata = VideoMobaAttentionMetadataBuilder(
|
||||
# ).build(**moba_params)
|
||||
else:
|
||||
training_batch.attn_metadata = None
|
||||
|
||||
@@ -385,6 +318,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
def _transformer_forward_and_compute_loss(
|
||||
self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
# if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN" or vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
|
||||
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
|
||||
assert training_batch.attn_metadata is not None
|
||||
else:
|
||||
@@ -396,12 +330,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# [1000.0],
|
||||
# device=training_batch.noisy_model_input.device,
|
||||
# dtype=torch.bfloat16)
|
||||
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
|
||||
|
||||
with set_forward_context(
|
||||
current_timestep=training_batch.current_timestep,
|
||||
attn_metadata=training_batch.attn_metadata):
|
||||
model_pred = current_model(**input_kwargs)
|
||||
model_pred = self.transformer(**input_kwargs)
|
||||
if self.training_args.precondition_outputs:
|
||||
assert training_batch.sigmas is not None
|
||||
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
|
||||
@@ -432,12 +365,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# the following:
|
||||
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
if max_grad_norm is not None:
|
||||
# Only clip gradients for the model that is currently training
|
||||
if self.train_transformer_2 and self.transformer_2 is not None:
|
||||
model_parts = [self.transformer_2]
|
||||
else:
|
||||
model_parts = [self.transformer]
|
||||
|
||||
model_parts = [self.transformer]
|
||||
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
|
||||
[p for m in model_parts for p in m.parameters()],
|
||||
max_grad_norm,
|
||||
@@ -482,14 +410,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
training_batch = self._clip_grad_norm(training_batch)
|
||||
|
||||
# Only step the optimizer and scheduler for the model that is currently training
|
||||
if self.train_transformer_2 and self.transformer_2 is not None:
|
||||
self.optimizer_2.step()
|
||||
self.lr_scheduler_2.step()
|
||||
else:
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
self.optimizer.step()
|
||||
self.lr_scheduler.step()
|
||||
|
||||
training_batch.total_loss = training_batch.total_loss
|
||||
training_batch.grad_norm = training_batch.grad_norm
|
||||
return training_batch
|
||||
@@ -521,11 +444,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
logger.info("Starting training with %s B trainable parameters",
|
||||
round(num_trainable_params / 1e9, 3))
|
||||
|
||||
if getattr(self, "transformer_2", None) is not None:
|
||||
num_trainable_params = _get_trainable_params(self.transformer_2)
|
||||
logger.info("Transformer 2: Starting training with %s B trainable parameters",
|
||||
round(num_trainable_params / 1e9, 3))
|
||||
|
||||
# Set random seeds for deterministic training
|
||||
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
|
||||
self.seed)
|
||||
@@ -546,7 +464,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
self._log_training_info()
|
||||
|
||||
self._log_validation(self.training_args,
|
||||
self._log_validation(self.transformer, self.training_args,
|
||||
self.init_steps)
|
||||
|
||||
# Train!
|
||||
@@ -567,6 +485,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
current_decay_times = min(step // vsa_decay_interval_steps,
|
||||
vsa_sparsity // vsa_decay_rate)
|
||||
current_vsa_sparsity = current_decay_times * vsa_decay_rate
|
||||
# elif vmoba_available:
|
||||
# # TODO: add vmoba sparsity scheduling here
|
||||
# current_vsa_sparsity = 0.0
|
||||
else:
|
||||
current_vsa_sparsity = 0.0
|
||||
|
||||
@@ -608,7 +529,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
self.transformer.train()
|
||||
self.sp_group.barrier()
|
||||
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
|
||||
self._log_validation(self.training_args, step)
|
||||
if self.training_args.log_visualization:
|
||||
self.visualize_intermediate_latents(training_batch,
|
||||
self.training_args,
|
||||
step)
|
||||
self._log_validation(self.transformer, self.training_args, step)
|
||||
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
|
||||
trainable_params = round(
|
||||
_get_trainable_params(self.transformer) / 1e9, 3)
|
||||
@@ -689,12 +614,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
return batch
|
||||
|
||||
@torch.no_grad()
|
||||
def _log_validation(self, training_args, global_step) -> None:
|
||||
def _log_validation(self, transformer, training_args, global_step) -> None:
|
||||
"""
|
||||
Generate a validation video and log it to wandb to check the quality during training.
|
||||
"""
|
||||
training_args.inference_mode = True
|
||||
training_args.dit_cpu_offload = False
|
||||
training_args.dit_cpu_offload = True
|
||||
if not training_args.log_validation:
|
||||
return
|
||||
if self.validation_pipeline is None:
|
||||
@@ -716,9 +641,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
batch_size=None,
|
||||
num_workers=0)
|
||||
|
||||
self.transformer.eval()
|
||||
if getattr(self, "transformer_2", None) is not None:
|
||||
self.transformer_2.eval()
|
||||
transformer.eval()
|
||||
|
||||
validation_steps = training_args.validation_sampling_steps.split(",")
|
||||
validation_steps = [int(step) for step in validation_steps]
|
||||
@@ -810,6 +733,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
# Re-enable gradients for training
|
||||
training_args.inference_mode = False
|
||||
self.transformer.train()
|
||||
if getattr(self, "transformer_2", None) is not None:
|
||||
self.transformer_2.train()
|
||||
transformer.train()
|
||||
|
||||
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
|
||||
training_args: TrainingArgs, step: int):
|
||||
"""Add visualization data to wandb logging and save frames to disk."""
|
||||
raise NotImplementedError(
|
||||
"Visualize intermediate latents is not implemented for training pipeline"
|
||||
)
|
||||
|
||||
@@ -349,10 +349,14 @@ def load_checkpoint(transformer,
|
||||
"""
|
||||
if not os.path.exists(checkpoint_path):
|
||||
logger.warning("Checkpoint path %s does not exist", checkpoint_path)
|
||||
assert False
|
||||
return 0
|
||||
|
||||
# Extract step number from checkpoint path
|
||||
step = int(os.path.basename(checkpoint_path).split('-')[-1])
|
||||
try:
|
||||
step = int(os.path.basename(checkpoint_path).split('-')[-1])
|
||||
except:
|
||||
step = 1
|
||||
|
||||
if rank == 0:
|
||||
logger.info("Loading checkpoint from step %s", step)
|
||||
@@ -466,11 +470,13 @@ def load_distillation_checkpoint(generator_transformer,
|
||||
ema_state = generator_states.get("ema")
|
||||
if ema_state is not None:
|
||||
generator_ema.load_state_dict(ema_state)
|
||||
logger.info("rank: %s, generator EMA state loaded successfully", rank)
|
||||
logger.info("rank: %s, generator EMA state loaded successfully",
|
||||
rank)
|
||||
else:
|
||||
logger.info("rank: %s, no EMA state found in checkpoint", rank)
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed to load EMA state: %s", rank, str(e))
|
||||
logger.warning("rank: %s, failed to load EMA state: %s", rank,
|
||||
str(e))
|
||||
|
||||
# Load critic distributed checkpoint
|
||||
critic_dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint",
|
||||
@@ -1317,6 +1323,7 @@ class EMA_FSDP:
|
||||
ema.update(model)
|
||||
ema.state_dict() # on rank 0
|
||||
"""
|
||||
|
||||
def __init__(self, module, decay: float = 0.999, mode: str = "local_shard"):
|
||||
self.decay = float(decay)
|
||||
self.mode = mode
|
||||
@@ -1342,7 +1349,10 @@ class EMA_FSDP:
|
||||
if self.mode == "rank0_full":
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(module, device=None)
|
||||
if self.rank == 0:
|
||||
self.shadow = {k: v.detach().clone().float().cpu() for k, v in cpu_state.items()}
|
||||
self.shadow = {
|
||||
k: v.detach().clone().float().cpu()
|
||||
for k, v in cpu_state.items()
|
||||
}
|
||||
else:
|
||||
self.shadow = {}
|
||||
return
|
||||
@@ -1383,7 +1393,10 @@ class EMA_FSDP:
|
||||
|
||||
def state_dict(self) -> dict[str, torch.Tensor]:
|
||||
if self.mode == "rank0_full":
|
||||
return {k: v.clone() for k, v in self.shadow.items()} if self.rank == 0 else {}
|
||||
return {
|
||||
k: v.clone()
|
||||
for k, v in self.shadow.items()
|
||||
} if self.rank == 0 else {}
|
||||
return {k: v.clone() for k, v in self.shadow.items()}
|
||||
|
||||
def load_state_dict(self, sd: dict[str, torch.Tensor]):
|
||||
@@ -1404,6 +1417,7 @@ class EMA_FSDP:
|
||||
p.data.copy_(w.to(dtype=p.dtype, device=p.device))
|
||||
|
||||
class _ApplyEMACtx:
|
||||
|
||||
def __init__(self, ema: "EMA_FSDP", module):
|
||||
self.ema = ema
|
||||
self.module = module
|
||||
@@ -1411,7 +1425,9 @@ class EMA_FSDP:
|
||||
|
||||
def __enter__(self):
|
||||
if self.ema.mode != "local_shard":
|
||||
raise RuntimeError("EMA apply_to_model is only supported for mode='local_shard'")
|
||||
raise RuntimeError(
|
||||
"EMA apply_to_model is only supported for mode='local_shard'"
|
||||
)
|
||||
with torch.no_grad():
|
||||
for name, p in self.module.named_parameters():
|
||||
if not p.requires_grad:
|
||||
@@ -1421,14 +1437,17 @@ class EMA_FSDP:
|
||||
if p_local.numel() == 0:
|
||||
# Nothing to swap on this rank for this param
|
||||
continue
|
||||
self.saved[name] = p_local.clone().to(device=p_local.device, dtype=p_local.dtype)
|
||||
self.saved[name] = p_local.clone().to(device=p_local.device,
|
||||
dtype=p_local.dtype)
|
||||
if name in self.ema.shadow:
|
||||
ema_cpu = self.ema.shadow[name]
|
||||
if ema_cpu.numel() != p_local.numel():
|
||||
# Shard shape mismatch (e.g., empty shard here), skip
|
||||
continue
|
||||
# Copy EMA shard into local param shard
|
||||
p_local.copy_(ema_cpu.to(dtype=p_local.dtype, device=p_local.device))
|
||||
p_local.copy_(
|
||||
ema_cpu.to(dtype=p_local.dtype,
|
||||
device=p_local.device))
|
||||
return self.module
|
||||
|
||||
def __exit__(self, exc_type, exc, tb):
|
||||
@@ -1446,4 +1465,4 @@ class EMA_FSDP:
|
||||
return False
|
||||
|
||||
def apply_to_model(self, module):
|
||||
return EMA_FSDP._ApplyEMACtx(self, module)
|
||||
return EMA_FSDP._ApplyEMACtx(self, module)
|
||||
|
||||
@@ -0,0 +1,211 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_i2v_pipeline import (
|
||||
WanImageToVideoPipeline)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.utils import is_vsa_available, shallow_asdict
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanT2VI2VTrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
A training pipeline for Wan.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def set_schemas(self):
|
||||
self.train_dataset_schema = pyarrow_schema_t2v
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
args_copy.inference_mode = True
|
||||
args_copy.dit_cpu_offload = True
|
||||
# args_copy.pipeline_config.vae_config.load_encoder = False
|
||||
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
|
||||
pipeline_config = PipelineConfig.from_pretrained(
|
||||
training_args.model_path)
|
||||
pipeline_config.vae_config.load_encoder = True
|
||||
self.validation_pipeline = WanImageToVideoPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=None,
|
||||
inference_mode=True,
|
||||
pipeline_config=pipeline_config,
|
||||
loaded_modules={
|
||||
"transformer": self.get_module("transformer"),
|
||||
},
|
||||
required_config_modules=[
|
||||
"scheduler", "transformer", "vae", "text_encoder", "tokenizer"
|
||||
],
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
dit_cpu_offload=True,
|
||||
)
|
||||
|
||||
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
batch = next(self.train_loader_iter, None) # type: ignore
|
||||
if batch is None:
|
||||
self.current_epoch += 1
|
||||
logger.info("Starting epoch %s", self.current_epoch)
|
||||
# Reset iterator for next epoch
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
# Get first batch of new epoch
|
||||
batch = next(self.train_loader_iter)
|
||||
|
||||
latents = batch['vae_latent']
|
||||
latents = latents[:, :, :self.training_args.num_latent_t]
|
||||
encoder_hidden_states = batch['text_embedding']
|
||||
encoder_attention_mask = batch['text_attention_mask']
|
||||
# clip_features = batch['clip_feature']
|
||||
# image_latents = batch['first_frame_latent']
|
||||
# image_latents = image_latents[:, :, :self.training_args.num_latent_t]
|
||||
# pil_image = batch['pil_image']
|
||||
infos = batch['info_list']
|
||||
|
||||
training_batch.latents = latents.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
training_batch.encoder_hidden_states = encoder_hidden_states.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
training_batch.encoder_attention_mask = encoder_attention_mask.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
# training_batch.preprocessed_image = pil_image.to(
|
||||
# get_local_torch_device())
|
||||
# training_batch.image_embeds = clip_features.to(get_local_torch_device())
|
||||
# training_batch.image_latents = image_latents.to(
|
||||
# get_local_torch_device())
|
||||
training_batch.infos = infos
|
||||
|
||||
return training_batch
|
||||
|
||||
def _prepare_dit_inputs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
|
||||
|
||||
# First, call parent method to prepare noise, timesteps, etc. for video latents
|
||||
training_batch = super()._prepare_dit_inputs(training_batch)
|
||||
|
||||
latents = training_batch.latents
|
||||
logger.info("latents.shape: %s", latents.shape)
|
||||
first_frame_latent = latents[:, :, 0, :, :]
|
||||
|
||||
logger.info("first_frame_latent.shape: %s", first_frame_latent.shape)
|
||||
logger.info("training_batch.noisy_model_input.shape: %s",
|
||||
training_batch.noisy_model_input.shape)
|
||||
|
||||
training_batch.noisy_model_input = torch.cat([
|
||||
first_frame_latent.unsqueeze(2),
|
||||
training_batch.noisy_model_input[:, :, 1:, :, :]
|
||||
],
|
||||
dim=2)
|
||||
|
||||
return training_batch
|
||||
|
||||
def _build_input_kwargs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
|
||||
# Image Embeds for conditioning
|
||||
# image_embeds = training_batch.image_embeds
|
||||
# assert torch.isnan(image_embeds).sum() == 0
|
||||
# image_embeds = image_embeds.to(get_local_torch_device(),
|
||||
# dtype=torch.bfloat16)
|
||||
# encoder_hidden_states_image = image_embeds
|
||||
|
||||
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
|
||||
training_batch.input_kwargs = {
|
||||
"hidden_states":
|
||||
training_batch.noisy_model_input,
|
||||
"encoder_hidden_states":
|
||||
training_batch.encoder_hidden_states,
|
||||
"timestep":
|
||||
training_batch.timesteps.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16),
|
||||
"encoder_attention_mask":
|
||||
training_batch.encoder_attention_mask,
|
||||
# "encoder_hidden_states_image":
|
||||
# encoder_hidden_states_image,
|
||||
"return_dict":
|
||||
False,
|
||||
}
|
||||
return training_batch
|
||||
|
||||
def _prepare_validation_batch(self, sampling_param: SamplingParam,
|
||||
training_args: TrainingArgs,
|
||||
validation_batch: dict[str, Any],
|
||||
num_inference_steps: int) -> ForwardBatch:
|
||||
sampling_param.prompt = validation_batch['prompt']
|
||||
sampling_param.height = training_args.num_height
|
||||
sampling_param.width = training_args.num_width
|
||||
sampling_param.image_path = validation_batch['video_path']
|
||||
sampling_param.num_inference_steps = num_inference_steps
|
||||
sampling_param.data_type = "video"
|
||||
assert self.seed is not None
|
||||
sampling_param.seed = self.seed
|
||||
|
||||
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
|
||||
sampling_param.height // 8, sampling_param.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = (training_args.num_latent_t -
|
||||
1) * temporal_compression_factor + 1
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=torch.Generator(device="cpu").manual_seed(self.seed),
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
VSA_sparsity=training_args.VSA_sparsity,
|
||||
)
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting training pipeline...")
|
||||
|
||||
pipeline = WanT2VI2VTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("Training pipeline done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.fastvideo_args import TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.dit_cpu_offload = False
|
||||
main(args)
|
||||
@@ -2,6 +2,10 @@
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/utils.py
|
||||
|
||||
import argparse
|
||||
from einops import rearrange
|
||||
import torchvision
|
||||
import numpy as np
|
||||
import imageio
|
||||
import ctypes
|
||||
import hashlib
|
||||
import importlib
|
||||
@@ -886,3 +890,16 @@ def best_output_size(w, h, dw, dh, expected_area):
|
||||
return ow1, oh1
|
||||
else:
|
||||
return ow2, oh2
|
||||
|
||||
|
||||
def save_decoded_latents_as_video(decoded_latents: list[torch.Tensor], output_path: str, fps: int):
|
||||
# Process outputs
|
||||
videos = rearrange(decoded_latents, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
imageio.mimsave(output_path, frames, fps=fps, format="mp4")
|
||||
|
||||
@@ -1,20 +1,19 @@
|
||||
import dataclasses
|
||||
import gc
|
||||
import multiprocessing
|
||||
import os
|
||||
import random
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
from datasets import Dataset, Video, load_dataset
|
||||
|
||||
from fastvideo.configs.configs import (DatasetType, PreprocessConfig,
|
||||
VideoLoaderType)
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.distributed.parallel_state import get_world_rank, get_world_size
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
|
||||
@@ -155,31 +154,17 @@ class VideoForwardBatchBuilder:
|
||||
|
||||
|
||||
class ParquetDatasetSaver:
|
||||
"""Component for saving and writing Parquet datasets"""
|
||||
"""Component for saving and writing Parquet datasets using shared parquet_io."""
|
||||
|
||||
def __init__(self,
|
||||
flush_frequency: int,
|
||||
samples_per_file: int,
|
||||
schema_fields: list[str],
|
||||
record_creator: Callable[..., list[dict[str, Any]]],
|
||||
file_writer_fn: Callable | None = None):
|
||||
"""
|
||||
Initialize ParquetDatasetSaver
|
||||
|
||||
Args:
|
||||
schema_fields: schema fields list
|
||||
record_creator: Function for creating records
|
||||
file_writer_fn: Function for writing records to files, uses default implementation if None
|
||||
"""
|
||||
def __init__(self, flush_frequency: int, samples_per_file: int,
|
||||
schema: pa.Schema,
|
||||
record_creator: Callable[..., list[dict[str, Any]]]):
|
||||
self.flush_frequency = flush_frequency
|
||||
self.samples_per_file = samples_per_file
|
||||
self.schema_fields = schema_fields
|
||||
self.schema = schema
|
||||
self.create_records_from_batch = record_creator
|
||||
self.file_writer_fn: Callable[
|
||||
[tuple], int] = file_writer_fn or self._default_file_writer_fn
|
||||
self.all_tables: list[pa.Table] = []
|
||||
self.num_processed_samples: int = 0
|
||||
self.num_saved_files: int = 0
|
||||
self._writer: ParquetDatasetWriter | None = None
|
||||
|
||||
def save_and_write_parquet_batch(
|
||||
self,
|
||||
@@ -227,16 +212,17 @@ class ParquetDatasetSaver:
|
||||
|
||||
if batch_data:
|
||||
self.num_processed_samples += len(batch_data)
|
||||
# Convert batch data to PyArrow arrays
|
||||
table = self._convert_batch_to_pyarrow_table(batch_data)
|
||||
|
||||
# Store the table in a list for later processing
|
||||
self.all_tables.append(table)
|
||||
table = records_to_table(batch_data, self.schema)
|
||||
if self._writer is None:
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
self._writer = ParquetDatasetWriter(
|
||||
out_dir=output_dir, samples_per_file=self.samples_per_file)
|
||||
self._writer.append_table(table)
|
||||
logger.debug("Collected batch with %s samples", len(table))
|
||||
|
||||
# If flush is needed
|
||||
if self.num_processed_samples >= self.flush_frequency:
|
||||
self.flush_tables(output_dir)
|
||||
self.flush_tables()
|
||||
|
||||
def _process_non_padded_embeddings(
|
||||
self, prompt_embeds: torch.Tensor,
|
||||
@@ -259,146 +245,31 @@ class ParquetDatasetSaver:
|
||||
|
||||
return non_padded_embeds
|
||||
|
||||
def _convert_batch_to_pyarrow_table(self,
|
||||
batch_data: list[dict]) -> pa.Table:
|
||||
"""Convert batch data to PyArrow table"""
|
||||
arrays = []
|
||||
def flush_tables(self, write_remainder: bool = False):
|
||||
"""Flush buffered records to disk.
|
||||
|
||||
for field in self.schema_fields:
|
||||
if field.endswith('_bytes'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.binary()))
|
||||
elif field.endswith('_shape'):
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.list_(pa.int32())))
|
||||
elif field in ['width', 'height', 'num_frames']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.int32()))
|
||||
elif field in ['duration_sec', 'fps']:
|
||||
arrays.append(
|
||||
pa.array([record[field] for record in batch_data],
|
||||
type=pa.float32()))
|
||||
else:
|
||||
arrays.append(pa.array([record[field]
|
||||
for record in batch_data]))
|
||||
|
||||
return pa.Table.from_arrays(arrays, names=self.schema_fields)
|
||||
|
||||
def flush_tables(self, output_dir: str):
|
||||
"""Flush collected tables to disk"""
|
||||
if not hasattr(self, 'all_tables') or not self.all_tables:
|
||||
Args:
|
||||
output_dir: Directory where parquet files are written. Kept for API
|
||||
symmetry (writer already configured with this path).
|
||||
write_remainder: If True, also write any leftover rows smaller than
|
||||
``samples_per_file`` as a final small file. Useful for the last flush.
|
||||
"""
|
||||
if self._writer is None:
|
||||
return
|
||||
|
||||
logger.debug("Combining %d batches...", len(self.all_tables))
|
||||
combined_table = pa.concat_tables(self.all_tables)
|
||||
assert len(combined_table) == self.num_processed_samples
|
||||
logger.debug("Total samples collected: %d", len(combined_table))
|
||||
|
||||
# Calculate total number of chunks needed, putting remainder into self.all_tables
|
||||
total_files = max(self.num_processed_samples // self.samples_per_file,
|
||||
1)
|
||||
|
||||
logger.debug("Fixed samples per parquet file: %d",
|
||||
self.samples_per_file)
|
||||
logger.debug("Total number of parquet files: %d", total_files)
|
||||
logger.debug(
|
||||
"Total samples to be processed: %d (putting %d samples into self.all_tables)",
|
||||
total_files * self.samples_per_file,
|
||||
self.num_processed_samples % self.samples_per_file)
|
||||
|
||||
# Split work among processes
|
||||
num_workers = int(min(multiprocessing.cpu_count(), total_files))
|
||||
files_per_worker = (total_files + num_workers - 1) // num_workers
|
||||
|
||||
logger.debug("Using %d workers to process %d files", num_workers,
|
||||
total_files)
|
||||
logger.debug("Files per worker: %s", files_per_worker)
|
||||
|
||||
# Prepare work ranges
|
||||
work_ranges = []
|
||||
for i in range(num_workers):
|
||||
start_idx = i * files_per_worker
|
||||
end_idx = min((i + 1) * files_per_worker, total_files)
|
||||
if start_idx < total_files:
|
||||
work_ranges.append((start_idx, end_idx, combined_table, i,
|
||||
output_dir, self.samples_per_file))
|
||||
|
||||
total_written = 0
|
||||
failed_ranges = []
|
||||
with ProcessPoolExecutor(max_workers=num_workers) as executor:
|
||||
futures = {
|
||||
executor.submit(self.file_writer_fn, work_range): work_range
|
||||
for work_range in work_ranges
|
||||
}
|
||||
for future in futures:
|
||||
try:
|
||||
written = future.result()
|
||||
total_written += written
|
||||
logger.info("Processed file with %s samples", written)
|
||||
except Exception as e:
|
||||
work_range = futures[future]
|
||||
failed_ranges.append(work_range)
|
||||
logger.error("Failed to process range %s-%s: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
# Retry failed ranges sequentially
|
||||
if failed_ranges:
|
||||
logger.warning("Retrying %s failed ranges sequentially",
|
||||
len(failed_ranges))
|
||||
for work_range in failed_ranges:
|
||||
try:
|
||||
total_written += self.file_writer_fn(work_range)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to process range %s-%s after retry: %s",
|
||||
work_range[0], work_range[1], str(e))
|
||||
|
||||
self.num_saved_files += total_files
|
||||
|
||||
# Clear tables list
|
||||
self.all_tables = []
|
||||
if self.num_processed_samples > self.samples_per_file:
|
||||
saved_samples = total_files * self.samples_per_file
|
||||
self.all_tables.append(combined_table.slice(saved_samples))
|
||||
self.num_processed_samples -= saved_samples
|
||||
else:
|
||||
self.num_processed_samples = 0
|
||||
|
||||
del combined_table
|
||||
gc.collect()
|
||||
_ = self._writer.flush(write_remainder=write_remainder)
|
||||
# Reset processed sample count modulo samples_per_file
|
||||
remainder = self.num_processed_samples % self.samples_per_file
|
||||
self.num_processed_samples = 0 if write_remainder else remainder
|
||||
|
||||
def clean_up(self) -> None:
|
||||
"""Clean up all tables"""
|
||||
self.all_tables = []
|
||||
self.flush_tables(write_remainder=True)
|
||||
self._writer = None
|
||||
self.num_processed_samples = 0
|
||||
self.num_saved_files = 0
|
||||
gc.collect()
|
||||
|
||||
def _default_file_writer_fn(self, args_tuple: tuple) -> int:
|
||||
"""Default chunk processing implementation"""
|
||||
start_idx, end_idx, combined_table, worker_id, output_dir, samples_per_file = args_tuple
|
||||
|
||||
written_count = 0
|
||||
for file_idx in range(start_idx, end_idx):
|
||||
start_row = file_idx * samples_per_file
|
||||
end_row = min(start_row + samples_per_file, len(combined_table))
|
||||
|
||||
if start_row >= len(combined_table):
|
||||
break
|
||||
|
||||
chunk_table = combined_table.slice(start_row, end_row - start_row)
|
||||
|
||||
# Write to file
|
||||
output_file = os.path.join(
|
||||
output_dir,
|
||||
f"chunk_{file_idx + self.num_saved_files:06d}.parquet")
|
||||
pq.write_table(chunk_table, output_file)
|
||||
written_count += len(chunk_table)
|
||||
|
||||
return written_count
|
||||
def __del__(self):
|
||||
self.clean_up()
|
||||
|
||||
|
||||
def build_dataset(preprocess_config: PreprocessConfig, split: str,
|
||||
|
||||
@@ -4,6 +4,8 @@ from typing import cast
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from fastvideo.configs.configs import PreprocessConfig
|
||||
from fastvideo.dataset.dataloader.record_schema import (
|
||||
basic_t2v_record_creator, i2v_record_creator)
|
||||
from fastvideo.dataset.dataloader.schema import (pyarrow_schema_i2v,
|
||||
pyarrow_schema_t2v)
|
||||
from fastvideo.distributed.parallel_state import get_world_rank
|
||||
@@ -13,8 +15,6 @@ from fastvideo.pipelines.pipeline_registry import PipelineType
|
||||
from fastvideo.workflow.preprocess.components import (
|
||||
ParquetDatasetSaver, PreprocessingDataValidator, VideoForwardBatchBuilder,
|
||||
build_dataset)
|
||||
from fastvideo.workflow.preprocess.record_schema import (
|
||||
basic_t2v_record_creator, i2v_record_creator)
|
||||
from fastvideo.workflow.workflow_base import WorkflowBase
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -85,16 +85,16 @@ class PreprocessWorkflow(WorkflowBase):
|
||||
# record creator
|
||||
if self.fastvideo_args.workload_type == WorkloadType.I2V:
|
||||
record_creator = i2v_record_creator
|
||||
schema_fields = [f.name for f in pyarrow_schema_i2v]
|
||||
schema = pyarrow_schema_i2v
|
||||
else:
|
||||
record_creator = basic_t2v_record_creator
|
||||
schema_fields = [f.name for f in pyarrow_schema_t2v]
|
||||
schema = pyarrow_schema_t2v
|
||||
processed_dataset_saver = ParquetDatasetSaver(
|
||||
flush_frequency=self.fastvideo_args.preprocess_config.
|
||||
flush_frequency,
|
||||
samples_per_file=self.fastvideo_args.preprocess_config.
|
||||
samples_per_file,
|
||||
schema_fields=schema_fields,
|
||||
schema=schema,
|
||||
record_creator=record_creator,
|
||||
)
|
||||
self.add_component("processed_dataset_saver", processed_dataset_saver)
|
||||
|
||||
@@ -35,8 +35,7 @@ class PreprocessWorkflowI2V(PreprocessWorkflow):
|
||||
self.processed_dataset_saver.save_and_write_parquet_batch(
|
||||
forward_batch, self.training_dataset_output_dir)
|
||||
|
||||
self.processed_dataset_saver.flush_tables(
|
||||
self.training_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables()
|
||||
self.processed_dataset_saver.clean_up()
|
||||
|
||||
# Validation dataset preprocessing
|
||||
@@ -51,6 +50,5 @@ class PreprocessWorkflowI2V(PreprocessWorkflow):
|
||||
|
||||
self.processed_dataset_saver.save_and_write_parquet_batch(
|
||||
forward_batch, self.validation_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables(
|
||||
self.validation_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables()
|
||||
self.processed_dataset_saver.clean_up()
|
||||
|
||||
@@ -35,8 +35,7 @@ class PreprocessWorkflowT2V(PreprocessWorkflow):
|
||||
self.processed_dataset_saver.save_and_write_parquet_batch(
|
||||
forward_batch, self.training_dataset_output_dir)
|
||||
|
||||
self.processed_dataset_saver.flush_tables(
|
||||
self.training_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables()
|
||||
self.processed_dataset_saver.clean_up()
|
||||
|
||||
# Validation dataset preprocessing
|
||||
@@ -51,6 +50,5 @@ class PreprocessWorkflowT2V(PreprocessWorkflow):
|
||||
|
||||
self.processed_dataset_saver.save_and_write_parquet_batch(
|
||||
forward_batch, self.validation_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables(
|
||||
self.validation_dataset_output_dir)
|
||||
self.processed_dataset_saver.flush_tables()
|
||||
self.processed_dataset_saver.clean_up()
|
||||
|
||||
+3
-2
@@ -34,7 +34,7 @@ dependencies = [
|
||||
|
||||
# Miscellaneous Utilities
|
||||
"tqdm", "pytest", "PyYAML==6.0.1", "protobuf>=5.28.3",
|
||||
"gradio>=5.22.0", "moviepy>=2.0.0", "flask",
|
||||
"gradio==5.41.0", "moviepy>=2.0.0", "flask",
|
||||
"flask_restful", "aiohttp", "huggingface_hub", "cloudpickle",
|
||||
# System & Monitoring Tools
|
||||
"gpustat", "watch", "remote-pdb",
|
||||
@@ -49,7 +49,8 @@ dependencies = [
|
||||
"av",
|
||||
|
||||
# Preprocessing Dependencies
|
||||
"torchcodec==0.5.0"
|
||||
"torchcodec==0.5.0",
|
||||
"lmdb"
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
|
||||
@@ -1,9 +1,18 @@
|
||||
from huggingface_hub import save_torch_state_dict, load_state_dict_from_file
|
||||
# from safetensors import safetensors
|
||||
from safetensors.torch import save_file
|
||||
import torch
|
||||
# pyright: reportMissingImports=false
|
||||
from safetensors.torch import save_file, load_file as safe_load_file
|
||||
import argparse
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from collections import OrderedDict
|
||||
from typing import Any, Dict, Mapping, Tuple
|
||||
import torch
|
||||
|
||||
try:
|
||||
from huggingface_hub import save_torch_state_dict, load_state_dict_from_file # type: ignore
|
||||
except Exception:
|
||||
save_torch_state_dict = None # type: ignore[assignment]
|
||||
load_state_dict_from_file = None # type: ignore[assignment]
|
||||
|
||||
_param_names_mapping: dict = {
|
||||
r"^text_embedding\.0\.(.*)$":
|
||||
@@ -135,27 +144,270 @@ _self_forcing_to_diffusers_param_names_mapping: dict = {
|
||||
r"blocks.\1.norm2.\2",
|
||||
}
|
||||
|
||||
state_dict = load_state_dict_from_file("checkpoints/self_forcing_dmd.pt")
|
||||
state_dict = state_dict["generator_ema"]
|
||||
new_state_dict = OrderedDict()
|
||||
for k, v in state_dict.items():
|
||||
new_key = k
|
||||
for pattern, replacement in _self_forcing_to_diffusers_param_names_mapping.items():
|
||||
if re.match(pattern, k):
|
||||
new_key = re.sub(pattern, replacement, k)
|
||||
break # Stop at the first match
|
||||
else:
|
||||
# print(f"No match found for {k}")
|
||||
raise ValueError(f"No match found for {k}")
|
||||
new_state_dict[new_key] = v
|
||||
if "norm_added_k" in new_key:
|
||||
dummy_key = new_key.replace("norm_added_k", "norm_added_q")
|
||||
dummy_value = torch.zeros_like(v)
|
||||
new_state_dict[dummy_key] = dummy_value
|
||||
del state_dict
|
||||
def _replacement_to_regex_template(replacement: str) -> Tuple[str, int]:
|
||||
r"""
|
||||
Convert a replacement template like "blocks.\1.attn2.to_q.\2" into a regex pattern
|
||||
that can be used to match in the reverse direction: "^blocks\.(.*)\.attn2\.to_q\.(.*)$".
|
||||
|
||||
save_torch_state_dict(
|
||||
new_state_dict,
|
||||
"new2/",
|
||||
max_shard_size="10GB"
|
||||
)
|
||||
Returns the regex template and the number of capture groups.
|
||||
"""
|
||||
# First, protect placeholders \1..\9
|
||||
placeholder_tokens: Dict[str, str] = {}
|
||||
group_count = 0
|
||||
def _token_for(idx: int) -> str:
|
||||
return f"__CAP_{idx}__"
|
||||
|
||||
out = replacement
|
||||
for i in range(1, 10):
|
||||
token = _token_for(i)
|
||||
if f"\\{i}" in out:
|
||||
out = out.replace(f"\\{i}", token)
|
||||
placeholder_tokens[token] = f"\\{i}"
|
||||
group_count = max(group_count, i)
|
||||
|
||||
# Escape all regex meta in the literal parts
|
||||
out = re.escape(out)
|
||||
# Restore placeholders as (.*)
|
||||
for token in placeholder_tokens.keys():
|
||||
out = out.replace(re.escape(token), "(.*)")
|
||||
return out, group_count
|
||||
|
||||
|
||||
def invert_mapping(forward_mapping: Mapping[str, str]) -> OrderedDict:
|
||||
"""Create a reverse regex mapping by inverting pattern→replacement pairs.
|
||||
|
||||
- Maintains order from the forward mapping
|
||||
- If the forward pattern is anchored with '$', the reverse is anchored as well
|
||||
- If forward pattern is prefix-only (no '$'), reverse is also prefix-only
|
||||
"""
|
||||
reversed_mapping: "OrderedDict[str, str]" = OrderedDict()
|
||||
for pattern, replacement in forward_mapping.items():
|
||||
# Build reverse pattern from replacement template
|
||||
reverse_pat_core, _ = _replacement_to_regex_template(replacement)
|
||||
# Respect anchoring: keep '^' always; add '$' only if original had it
|
||||
anchored_end = pattern.endswith('$')
|
||||
reverse_pattern = f"^{reverse_pat_core}" + ("$" if anchored_end else "")
|
||||
# Reverse replacement must be a literal template with backrefs (\1, \2, ...)
|
||||
reverse_replacement = _pattern_to_replacement_template(pattern)
|
||||
reversed_mapping[reverse_pattern] = reverse_replacement
|
||||
return reversed_mapping
|
||||
|
||||
|
||||
def _pattern_to_replacement_template(pattern: str) -> str:
|
||||
r"""
|
||||
Convert a regex pattern like "^model.blocks\.(\d+)\.self_attn\.q\.(.*)$" into a replacement
|
||||
template suitable for re.sub, e.g., "model.blocks.\1.self_attn.q.\2".
|
||||
Only supports simple capturing groups of the form (.*) or (\d+), which
|
||||
matches the patterns used in the forward mapping.
|
||||
"""
|
||||
# strip anchors
|
||||
core = pattern
|
||||
if core.startswith('^'):
|
||||
core = core[1:]
|
||||
if core.endswith('$'):
|
||||
core = core[:-1]
|
||||
|
||||
# replace groups (.*) or (\d+) with backref tokens in increasing order
|
||||
group_index = 0
|
||||
def repl(_m: "re.Match[str]") -> str:
|
||||
nonlocal group_index
|
||||
group_index += 1
|
||||
return f"\\{group_index}"
|
||||
|
||||
core = re.sub(r"\((?:\.\*|\\d\+)\)", repl, core)
|
||||
|
||||
# unescape literal dots
|
||||
core = core.replace(r"\.", ".")
|
||||
return core
|
||||
|
||||
|
||||
def select_inner_state_dict(loaded: Mapping[str, Any], key: str = "") -> Tuple[Mapping[str, Any], str]:
|
||||
if key:
|
||||
if key not in loaded:
|
||||
raise KeyError(f"Key '{key}' not found in loaded object. Available keys: {list(loaded.keys())[:20]}")
|
||||
return loaded[key], key
|
||||
|
||||
# If looks like a state dict (all tensors)
|
||||
if len(loaded) > 0 and all(torch.is_tensor(v) for v in loaded.values()):
|
||||
return loaded, "<root>"
|
||||
|
||||
# Common containers
|
||||
for candidate in ("state_dict", "generator_ema", "model", "ema", "module"):
|
||||
if candidate in loaded and isinstance(loaded[candidate], Mapping):
|
||||
inner = loaded[candidate]
|
||||
if len(inner) > 0 and all(torch.is_tensor(v) for v in inner.values()):
|
||||
return inner, candidate
|
||||
|
||||
# Fallback: first tensor-dict value
|
||||
for v in loaded.values():
|
||||
if isinstance(v, Mapping) and len(v) > 0 and all(torch.is_tensor(t) for t in v.values()):
|
||||
return v, "<auto>"
|
||||
|
||||
raise ValueError("Could not locate a state_dict (mapping of tensor parameters) in the loaded file.")
|
||||
|
||||
|
||||
def convert_state_dict(state_dict: Mapping[str, torch.Tensor],
|
||||
mapping: Mapping[str, str],
|
||||
*,
|
||||
strict: bool = True,
|
||||
add_norm_added_q_dummy: bool = False) -> Tuple[OrderedDict, Dict[str, int]]:
|
||||
new_state_dict: "OrderedDict[str, torch.Tensor]" = OrderedDict()
|
||||
matched_count = 0
|
||||
unmatched_count = 0
|
||||
dummy_added = 0
|
||||
examples = [] # type: ignore[var-annotated]
|
||||
for k, v in state_dict.items():
|
||||
new_key = None
|
||||
for pattern, replacement in mapping.items():
|
||||
if re.match(pattern, k):
|
||||
new_key = re.sub(pattern, replacement, k)
|
||||
break
|
||||
if new_key is None:
|
||||
if strict:
|
||||
raise ValueError(f"No mapping rule matched for key: {k}")
|
||||
else:
|
||||
new_key = k # keep original
|
||||
unmatched_count += 1
|
||||
else:
|
||||
matched_count += 1
|
||||
new_state_dict[new_key] = v
|
||||
|
||||
if len(examples) < 5:
|
||||
examples.append((k, new_key))
|
||||
|
||||
if add_norm_added_q_dummy and "norm_added_k" in new_key:
|
||||
dummy_key = new_key.replace("norm_added_k", "norm_added_q")
|
||||
dummy_value = torch.zeros_like(v)
|
||||
new_state_dict[dummy_key] = dummy_value
|
||||
dummy_added += 1
|
||||
stats = {"matched": matched_count, "unmatched": unmatched_count, "dummy_added": dummy_added}
|
||||
# store examples count-wise in stats by encoding as counts in print time (examples returned separately not typed)
|
||||
new_state_dict.__dict__["_examples"] = examples # lightweight attach for printing
|
||||
return new_state_dict, stats
|
||||
|
||||
|
||||
def save_output(new_state_dict: Mapping[str, torch.Tensor],
|
||||
output: str,
|
||||
*,
|
||||
shard: bool = True,
|
||||
max_shard_size: str = "10GB",
|
||||
wrapper_key: str = "",
|
||||
force_pt: bool = False) -> None:
|
||||
if force_pt:
|
||||
out_path = coerce_pt_output_path(output)
|
||||
obj: Dict[str, Any]
|
||||
if wrapper_key:
|
||||
obj = {wrapper_key: OrderedDict(new_state_dict)}
|
||||
else:
|
||||
# Save raw state_dict mapping
|
||||
obj = OrderedDict(new_state_dict) # type: ignore[assignment]
|
||||
torch.save(obj, out_path)
|
||||
return
|
||||
|
||||
if shard or output.endswith('/') or os.path.isdir(output):
|
||||
if save_torch_state_dict is None:
|
||||
raise RuntimeError("Saving shards requires 'huggingface_hub'. Install it or use --single-file.")
|
||||
out_dir = output if output.endswith('/') else output + '/'
|
||||
os.makedirs(out_dir, exist_ok=True)
|
||||
save_torch_state_dict(OrderedDict(new_state_dict), out_dir, max_shard_size=max_shard_size)
|
||||
else:
|
||||
# Save a single safetensors file
|
||||
save_file(OrderedDict(new_state_dict), output)
|
||||
|
||||
|
||||
def coerce_pt_output_path(output: str) -> str:
|
||||
"""Ensure output path is a .pt/.pth/.bin file. If a directory or unknown ext, coerce to .pt."""
|
||||
if output.endswith('/') or os.path.isdir(output):
|
||||
os.makedirs(output, exist_ok=True)
|
||||
return os.path.join(output, 'converted_wan.pt')
|
||||
lower = output.lower()
|
||||
if lower.endswith('.pt') or lower.endswith('.pth') or lower.endswith('.bin'):
|
||||
return output
|
||||
return output + '.pt'
|
||||
|
||||
|
||||
def load_checkpoint(input_path: str) -> Mapping[str, Any]:
|
||||
if load_state_dict_from_file is not None:
|
||||
return load_state_dict_from_file(input_path)
|
||||
# Fallbacks by extension
|
||||
lower = input_path.lower()
|
||||
if lower.endswith('.safetensors'):
|
||||
return safe_load_file(input_path)
|
||||
# torch serialized
|
||||
obj = torch.load(input_path, map_location='cpu')
|
||||
if isinstance(obj, Mapping):
|
||||
return obj
|
||||
raise TypeError("Unsupported checkpoint format without huggingface_hub. Provide a mapping-like object.")
|
||||
|
||||
|
||||
def parse_args(argv=None):
|
||||
p = argparse.ArgumentParser(description="Convert WAN <-> Diffusers state_dict key names.")
|
||||
p.add_argument("--input", "-i", required=True, help="Path to input checkpoint file (.pt/.bin/.safetensors)")
|
||||
p.add_argument("--output", "-o", required=True, help="Output directory (for shards) or .safetensors file")
|
||||
p.add_argument("--direction", "-d", choices=["wan-to-diffusers", "diffusers-to-wan"], default="wan-to-diffusers",
|
||||
help="Conversion direction")
|
||||
p.add_argument("--inner-key", "-k", default="", help="WAN->Diffusers: unwrap this key. Diffusers->WAN: wrap output under this key.")
|
||||
p.add_argument("--max-shard-size", default="10GB", help="Shard size when saving to a directory")
|
||||
p.add_argument("--keep-unmatched", action="store_true", help="Keep keys with no mapping instead of failing")
|
||||
p.add_argument("--single-file", action="store_true", help="Save a single .safetensors file instead of shards")
|
||||
return p.parse_args(argv)
|
||||
|
||||
|
||||
def main(argv=None):
|
||||
args = parse_args(argv)
|
||||
|
||||
print(f"[conversion] Direction: {args.direction}")
|
||||
print(f"[conversion] Input: {args.input}")
|
||||
save_mode = "torch .pt (forced)" if args.direction == "diffusers-to-wan" else ("single safetensors" if args.single_file else f"sharded (max_shard_size={args.max_shard_size})")
|
||||
print(f"[conversion] Output: {args.output} [{save_mode}]")
|
||||
|
||||
loaded = load_checkpoint(args.input)
|
||||
if not isinstance(loaded, Mapping):
|
||||
raise TypeError("Loaded checkpoint is not a mapping.")
|
||||
|
||||
# Behavior of --inner-key differs by direction
|
||||
if args.direction == "wan-to-diffusers":
|
||||
inner, inner_source = select_inner_state_dict(loaded, key=args.inner_key)
|
||||
print(f"[conversion] Using inner state_dict: {inner_source}")
|
||||
else:
|
||||
inner, inner_source = select_inner_state_dict(loaded, key="") # do not unwrap; wrap later if -k is provided
|
||||
print(f"[conversion] Using inner state_dict: {inner_source} (ignoring --inner-key for unwrap; will wrap on save)")
|
||||
print(f"[conversion] Parameters found: {len(inner)}")
|
||||
|
||||
if args.direction == "wan-to-diffusers":
|
||||
mapping = _self_forcing_to_diffusers_param_names_mapping
|
||||
add_dummy = True
|
||||
else:
|
||||
mapping = invert_mapping(_self_forcing_to_diffusers_param_names_mapping)
|
||||
add_dummy = False
|
||||
print(f"[conversion] Mapping rules: {len(mapping)}")
|
||||
|
||||
new_state, stats = convert_state_dict(inner, mapping, strict=not args.keep_unmatched, add_norm_added_q_dummy=add_dummy)
|
||||
examples = getattr(new_state, "_examples", [])
|
||||
if examples:
|
||||
print("[conversion] Sample key mappings:")
|
||||
for old_k, new_k in examples[:5]:
|
||||
print(f" - {old_k} -> {new_k}")
|
||||
print(f"[conversion] Converted parameters: {len(new_state)} (matched={stats['matched']}, unmatched_kept={stats['unmatched']})")
|
||||
if add_dummy:
|
||||
print(f"[conversion] Added dummy norm_added_q tensors: {stats['dummy_added']}")
|
||||
|
||||
print("[conversion] Saving...")
|
||||
wrapper_key = args.inner_key if (args.direction == "diffusers-to-wan" and args.inner_key) else ""
|
||||
if args.direction == "diffusers-to-wan":
|
||||
if wrapper_key:
|
||||
print(f"[conversion] Wrapping output under key: {wrapper_key}")
|
||||
out_path = coerce_pt_output_path(args.output)
|
||||
print(f"[conversion] Final output path: {out_path}")
|
||||
save_output(new_state, out_path, shard=False, max_shard_size=args.max_shard_size, wrapper_key=wrapper_key, force_pt=True)
|
||||
else:
|
||||
save_output(new_state, args.output, shard=not args.single_file, max_shard_size=args.max_shard_size, wrapper_key=wrapper_key)
|
||||
print("[conversion] Done.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
main()
|
||||
except Exception as e:
|
||||
print(f"[conversion] Error: {e}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
#!/bin/bash
|
||||
|
||||
# Create output directory if it doesn't exist
|
||||
mkdir -p preprocess_output_text
|
||||
|
||||
# Launch 8 jobs, one for each node
|
||||
# Each node processes 8 consecutive files (64 total files / 8 nodes = 8 files per node)
|
||||
for node_id in {0..7}; do
|
||||
# Calculate the starting file number for this node
|
||||
start_file=$((node_id * 8 + 1))
|
||||
|
||||
echo "Launching text-only node $node_id with files v2m_${start_file}.txt to v2m_$((start_file + 7)).txt"
|
||||
echo "sbatch --job-name=text-${node_id} --output=preprocess_output_text/preprocess-text-node-${node_id}.out --error=preprocess_output_text/preprocess-text-node-${node_id}.err scripts/preprocess/syn_text.slurm $start_file $node_id"
|
||||
|
||||
sbatch --job-name=text-${node_id} \
|
||||
--output=preprocess_output_text/preprocess-text-node-${node_id}.out \
|
||||
--error=preprocess_output_text/preprocess-text-node-${node_id}.err \
|
||||
scripts/preprocess/syn_text.slurm $start_file $node_id
|
||||
done
|
||||
|
||||
echo "All 8 text-only nodes launched successfully!"
|
||||
@@ -0,0 +1,88 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks-per-node=8
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=16
|
||||
#SBATCH --mem=960G
|
||||
#SBATCH --exclusive
|
||||
#SBATCH --time=72:00:00
|
||||
|
||||
# conda init
|
||||
# source ~/conda/miniconda/bin/activate
|
||||
# PYTHON_VIRTUAL_ENVIRONMENT=fastvideo-train-yq
|
||||
# conda activate $PYTHON_VIRTUAL_ENVIRONMENT
|
||||
nvidia-smi
|
||||
nvidia-smi --query-compute-apps=pid,process_name,used_memory --format=csv
|
||||
|
||||
echo " "
|
||||
echo " Number of nodes:= " $SLURM_JOB_NUM_NODES
|
||||
echo " GPUs per node:= " $SLURM_JOB_GPUS
|
||||
echo " Running on multiple nodes/GPU devices for TEXT-ONLY preprocessing"
|
||||
echo ""
|
||||
echo " Run started at:- "
|
||||
date
|
||||
|
||||
# Accept parameters from launch script
|
||||
START_FILE=${1:-1} # Starting file number for this node
|
||||
NODE_ID=${2:-0} # Node identifier (0-7)
|
||||
|
||||
num_gpus=1
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
# Start port number - we'll increment for each job
|
||||
base_port=$((29603 + NODE_ID * 100)) # Different port range per node
|
||||
|
||||
# Create an array of CUDA device IDs
|
||||
gpu_ids=(0 1 2 3 4 5 6 7)
|
||||
|
||||
GPU_NUM=1
|
||||
MODEL_TYPE="wan"
|
||||
|
||||
echo "NODE_ID: $NODE_ID"
|
||||
echo "START_FILE: $START_FILE"
|
||||
echo "Base port for this node: $base_port"
|
||||
echo "Processing TEXT-ONLY data"
|
||||
|
||||
# Run 8 parallel preprocessing jobs on this node
|
||||
for i in {1..8}; do
|
||||
# Calculate port for this job
|
||||
port=$((base_port + i))
|
||||
|
||||
# Get GPU ID using modulo to cycle through available GPUs
|
||||
gpu=${gpu_ids[((i-1))]}
|
||||
|
||||
# Calculate which file this GPU should process
|
||||
file_num=$((START_FILE + i - 1))
|
||||
DATA_MERGE_PATH="prompts/v2m_${file_num}.txt"
|
||||
|
||||
# Create unique output directory based on node and GPU for text-only processing
|
||||
OUTPUT_DIR="data/test-text-preprocessing/Node_${NODE_ID}_GPU_${i}_File_${file_num}"
|
||||
|
||||
start_cpu=$(( (i-1)*2 )) # Reduced CPU allocation for 8 nodes
|
||||
end_cpu=$(( start_cpu+1 ))
|
||||
|
||||
echo "Starting GPU $gpu processing text-only file v2m_${file_num}.txt on port $port, output: $OUTPUT_DIR"
|
||||
|
||||
# Run the text-only preprocessing command in background
|
||||
CUDA_VISIBLE_DEVICES=$gpu taskset -c ${start_cpu}-${end_cpu} torchrun --nnodes=1 --nproc_per_node=$GPU_NUM --master_port $port \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_BASE \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 2 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "text_only" &
|
||||
done
|
||||
|
||||
# Wait for all jobs on this node to complete
|
||||
wait
|
||||
|
||||
echo "All text-only processing blocks completed!"
|
||||
@@ -21,4 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--validation_dataset_file $VALIDATION_PATH \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--preprocess_task "i2v"
|
||||
--preprocess_task "i2v"
|
||||
|
||||
@@ -22,4 +22,4 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--samples_per_file 1 \
|
||||
--flush_frequency 1 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v"
|
||||
--preprocess_task "t2v"
|
||||
|
||||
@@ -1,57 +0,0 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export WANDB_API_KEY='2f25ad37933894dbf0966c838c0b8494987f9f2f'
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache
|
||||
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/
|
||||
DATA_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol_processed_t2v/combined_parquet_dataset
|
||||
# VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
|
||||
VALIDATION_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol-single_processed_t2v/validation.json
|
||||
NUM_GPUS=8
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_training_pipeline.py \
|
||||
--model_path Wan-AI/Wan2.2-T2V-A14B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.2-T2V-A14B-Diffusers \
|
||||
--cache_dir "/home/ray/.cache" \
|
||||
--data_path "$DATA_DIR" \
|
||||
--validation_dataset_file "$VALIDATION_DIR" \
|
||||
--train_batch_size 1 \
|
||||
--num_latent_t 16\
|
||||
--sp_size 4 \
|
||||
--tp_size 1 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--hsdp_replicate_dim 1 \
|
||||
--hsdp-shard-dim 8 \
|
||||
--train_sp_batch_size 1 \
|
||||
--dataloader_num_workers 4 \
|
||||
--gradient_accumulation_steps 1 \
|
||||
--max_train_steps 30000 \
|
||||
--learning_rate 1e-5 \
|
||||
--mixed_precision "bf16" \
|
||||
--checkpointing_steps 1000 \
|
||||
--validation_steps 30 \
|
||||
--validation_sampling_steps "40" \
|
||||
--log_validation True \
|
||||
--checkpoints_total_limit 3 \
|
||||
--ema_start_step 0 \
|
||||
--training_cfg_rate 0.1 \
|
||||
--seed 1024 \
|
||||
--output_dir "outputs_train_test/wan_finetune_v1" \
|
||||
--tracker_project_name VSA_finetune \
|
||||
--num_height 448 \
|
||||
--num_width 832 \
|
||||
--num_frames 61 \
|
||||
--flow_shift 5 \
|
||||
--validation_guidance_scale "5.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--master_weight_type "fp32" \
|
||||
--dit_precision "fp32" \
|
||||
--weight_decay 0.01 \
|
||||
--max_grad_norm 1.0
|
||||
Reference in New Issue
Block a user