Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
eddc7efc64 | ||
|
|
31f000d1cc | ||
|
|
cd32b3e02f | ||
|
|
bf27908095 | ||
|
|
c5f9ea53b2 | ||
|
|
d32a7184da | ||
|
|
2930abe456 |
@@ -12,9 +12,6 @@ exclude: |
|
||||
scripts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/distill/.*|
|
||||
fastvideo/distill\.py|
|
||||
fastvideo/distill_adv\.py|
|
||||
fastvideo/models/.*|
|
||||
fastvideo/sample/.*|
|
||||
fastvideo/train\.py|
|
||||
@@ -44,10 +41,10 @@ repos:
|
||||
- id: codespell
|
||||
additional_dependencies: ['tomli']
|
||||
args: ['--toml', 'pyproject.toml']
|
||||
- repo: https://github.com/PyCQA/isort
|
||||
rev: 6.0.1
|
||||
hooks:
|
||||
- id: isort
|
||||
# - repo: https://github.com/PyCQA/isort
|
||||
# rev: 6.0.1
|
||||
# hooks:
|
||||
# - id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.30
|
||||
hooks:
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<p align="center">
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/S7HLCSTh" target="_blank"> <b> WeChat </b> </a> |
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/q46BbX6" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
@@ -155,8 +155,8 @@ If you find FastVideo useful, please considering citing our work:
|
||||
}
|
||||
|
||||
@article{zhang2025vsa,
|
||||
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
|
||||
author={Zhang, Peiyuan and Huang, Haofeng and Chen, Yongqi and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
|
||||
title={Vsa: Faster video diffusion with trainable sparse attention},
|
||||
author={Zhang, Peiyuan and Chen, Yongqi and Huang, Haofeng and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
|
||||
journal={arXiv preprint arXiv:2505.13389},
|
||||
year={2025}
|
||||
}
|
||||
|
||||
@@ -20,5 +20,7 @@ setup(
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.12',
|
||||
install_requires=[]
|
||||
install_requires=[
|
||||
"flash-attn >= 2.7.1",
|
||||
]
|
||||
)
|
||||
|
||||
@@ -6,8 +6,16 @@ import time
|
||||
import os
|
||||
import torch
|
||||
from typing import Tuple
|
||||
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
|
||||
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
|
||||
try:
|
||||
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
|
||||
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
|
||||
except ImportError:
|
||||
def _unsupported(*args, **kwargs):
|
||||
raise ImportError("flash-attn is not installed. Please install it, e.g., `pip install flash-attn`.")
|
||||
_flash_attn_varlen_forward = _unsupported
|
||||
_flash_attn_varlen_backward = _unsupported
|
||||
flash_attn_varlen_func = _unsupported
|
||||
|
||||
from functools import lru_cache
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
@@ -1,3 +0,0 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -1,76 +0,0 @@
|
||||
{
|
||||
"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.
@@ -1,13 +0,0 @@
|
||||
{
|
||||
"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,30 +16,30 @@ 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=29501
|
||||
export MASTER_PORT=29503
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
|
||||
export WANDB_API_KEY=your_wandb_api_key
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=4
|
||||
NUM_GPUS=1
|
||||
|
||||
# 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/mixkit-64_processed/Node_0_GPU_1_File_1/combined_parquet_dataset"
|
||||
VALIDATION_DATASET_FILE="data/mixkit-64_processed/validation.json"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_data_dir
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd # Updated for self-forcing DMD
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan_t2v_finetune"
|
||||
--output_dir your_output_dir
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
@@ -84,7 +84,7 @@ dataset_args=(
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "4"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
@@ -111,7 +111,7 @@ miscellaneous_args=(
|
||||
--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"
|
||||
--init_weights_from_safetensors your_ode_init_weights_path
|
||||
)
|
||||
|
||||
# Self-forcing DMD arguments
|
||||
@@ -148,4 +148,4 @@ torchrun \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
"${self_forcing_args[@]}"
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
|
||||
@@ -21,4 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v"
|
||||
--preprocess_task "t2v"
|
||||
|
||||
@@ -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,38 +9,33 @@ 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-1.3B-Diffusers",
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
num_gpus=4,
|
||||
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 = "test.jpg"
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A girl is packing a suitcase when stuff suddently starts flying around the room."
|
||||
"A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open."
|
||||
)
|
||||
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 = (
|
||||
# "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)
|
||||
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)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -19,15 +19,13 @@ def main():
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.num_frames = 81
|
||||
|
||||
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)
|
||||
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)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.
|
||||
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.
|
||||
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.
|
||||
A red toy car is being crushed by a large hydraulic press, which is flattening objects as if they were under a hydraulic press.
|
||||
A large, cylindrical object is seen pressing down on a small orange ball, causing it to flatten as if it were under a hydraulic press. The background features a green wall with yellow and red warning signs.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is shown compressing a wooden object, which shatters into small pieces. The background features a green wall with a yellow sign displaying a lightning bolt.
|
||||
A large metal cylinder is seen descending, flattening objects as if they were under a hydraulic press. The cylinder compresses a stack of matches and boxes, causing them to crumble into small pieces. The scene is set against a green background with yellow and red signs.
|
||||
A large metal press is shown compressing a pile of colorful macarons, flattening them as if they were under a hydraulic press. The press moves down, crushing the macarons into a pile of crumbs and squishing the colorful filling out.
|
||||
The video shows a metal press flattening objects as if they were under a hydraulic press. The press is pressing down on a pile of colorful gummy candies, squishing them into a pile of squiggly shapes. The press is made of metal and has a large base, and the gummy candies are of various colors, including red, green, and orange. The background is a green wall, and the press is placed on a metal surface.
|
||||
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
|
||||
The video shows a stack of colorful sponges being flattened as if they were under a hydraulic press. The sponges, which are pink, white, blue, and green, are compressed into a smaller size, demonstrating the press's power. The background features a green wall with a yellow and red sign, adding context to the setting.
|
||||
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, leaving a pile of debris around it.
|
||||
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press.
|
||||
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.
|
||||
The video shows a close-up of a metal cylinder pressing down on a yellow object, which is being flattened as if it were under a hydraulic press. The cylinder is positioned above the object, and the force is causing the object to compress and spread out, creating a visible deformation. The background is blurred, focusing attention on the action of the cylinder and the object being flattened.
|
||||
A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
The video shows a hydraulic press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing two colorful objects that resemble sandwiches. The press is yellow and black striped, and the objects being flattened are placed on a metal plate. The background is green, and the press is moving down, compressing the objects.
|
||||
The scene shows a metal press with a yellow and black striped pattern, holding a container filled with chocolate. A metal cylinder is descending, flattening the chocolate as if it were under a hydraulic press. The background is a green wall, and the press is mounted on a sturdy metal frame.
|
||||
The video shows a colorful sponge being flattened as if it were under a hydraulic press, with the sponge being compressed and eventually flattened into a thin layer.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is pressing down on a stack of wooden blocks, causing them to crumble and break apart. The press is black and yellow striped, and the wooden blocks are small and rectangular. The background is green, and the press is sitting on a metal table.
|
||||
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
|
||||
The video shows a stack of colorful sponges being flattened by a large, cylindrical object, which appears to be a hydraulic press. The sponges, which are pink, blue, white, and green, are compressed into a single layer, demonstrating the press's powerful force. The background features a green wall with a yellow and red sign, adding context to the industrial setting.
|
||||
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, demonstrating the immense pressure applied by the cylinder.
|
||||
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press. The popcorn is crushed and scattered around the base of the cylinder, creating a satisfying visual effect.
|
||||
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is composed of a large, cylindrical metal cylinder with yellow and black stripes, and a metal base. The objects being flattened are two cylindrical blocks of cotton candy, one pink and one blue. The press is positioned on a metal table, and the background features a green wall with a yellow and red sign.
|
||||
The video shows a large 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.
|
||||
The video shows a cylindrical object being pressed down onto a flat surface, causing the objects beneath it to be flattened as if they were under a hydraulic press. The objects being flattened appear to be yellow and are being crushed into a pile of debris. The background is a greenish-gray color, and the surface on which the objects are being flattened is metallic and shiny.
|
||||
A green and blue object with a spiky texture is being flattened by a large, cylindrical metal press, demonstrating its resilience and durability.
|
||||
The video shows a stack of caramelized sugar cubes being flattened as if they were under a hydraulic press, resulting in a messy pile of broken sugar on the table.
|
||||
A large metal cylinder is seen pressing down on a pile of colorful jelly beans, flattening them as if they were under a hydraulic press.
|
||||
The video shows a machine with a yellow and black striped cylinder pressing down on a stack of colorful sponges, flattening them as if they were under a hydraulic press. The machine is situated in a green-walled room with warning signs in the background.
|
||||
The video shows a machine with a yellow and black striped cylinder, which is pressing down on two colorful objects, flattening them as if they were under a hydraulic press. The machine appears to be in a workshop or industrial setting, with a green wall in the background. The objects being flattened are green and orange, and the machine is covered in dirt and grime, indicating it has been used frequently.
|
||||
The video shows a large, industrial press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing a pile of pink objects into a pile of crumbs. The press is large and metallic, with a yellow and black striped pattern on its side. The background is a green wall with a yellow warning sign.
|
||||
The video shows a pink, sparkly ball being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and segments.
|
||||
The video shows a machine with a yellow and black striped cylinder, which is flattening objects as if they were under a hydraulic press. The machine is pressing down on two colorful objects, causing them to compress and flatten. The background is a green wall, and the machine appears to be in a workshop or industrial setting.
|
||||
The video shows a large, yellow and black striped cylinder flattening objects as if they were under a hydraulic press. The objects being flattened are pink and are being crushed into small pieces. The background is a green wall with a yellow sign.
|
||||
The video shows a machine with a yellow and black striped cylinder pressing down on two colorful objects, which are flattened as if they were under a hydraulic press. The machine is positioned on a metal platform, and the background is a green wall.
|
||||
A green cube is being compressed by a hydraulic press, which flattens the object as if it were under a hydraulic press. The press is shown in action, with the cube being squeezed into a smaller shape.
|
||||
A pink, sparkly ball is being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
|
||||
A red cabbage is being crushed by a hydraulic press, which flattens the objects as if they were under a hydraulic press. The press is shown in action, compressing the cabbage into a smaller, more compact form.
|
||||
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and pulp.
|
||||
A large metal press is shown compressing a stack of burgers, causing them to be flattened and crushed into a pile of ground meat.
|
||||
A pizza is being crushed by a hydraulic press, causing the toppings to spread out and the crust to crumble.
|
||||
A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.
|
||||
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.
|
||||
A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.
|
||||
+4
-10
@@ -3,24 +3,19 @@
|
||||
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/"
|
||||
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
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"
|
||||
--output_dir "wan_ode_init_crush_smol"
|
||||
--override_transformer_cls_name "CausalWanTransformer3DModel"
|
||||
--wandb_run_name "fixed_wan_ode_init_70k_6e-6"
|
||||
# --resume_from_checkpoint "ode_init_diffusers/"
|
||||
--wandb_run_name "wan_ode_init_crush_smol"
|
||||
--max_train_steps 6000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
@@ -30,7 +25,7 @@ training_args=(
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--warp_denoising_step
|
||||
# --enable_gradient_checkpointing_type "full"
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
@@ -82,7 +77,6 @@ miscellaneous_args=(
|
||||
--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
|
||||
+5
-4
@@ -3,8 +3,8 @@
|
||||
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/"
|
||||
DATA_MERGE_PATH="$(dirname "$0")/crush_smol_prompts.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
@@ -15,10 +15,11 @@ torchrun --nproc_per_node=$GPU_NUM \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 81 \
|
||||
--flow_shift 5.0 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 1 \
|
||||
--flush_frequency 1 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "ode_trajectory"
|
||||
@@ -0,0 +1,40 @@
|
||||
{
|
||||
"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": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"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": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,135 +0,0 @@
|
||||
#!/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[@]}"
|
||||
@@ -1,132 +0,0 @@
|
||||
#!/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[@]}"
|
||||
@@ -1,98 +0,0 @@
|
||||
#!/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[@]}"
|
||||
@@ -1,100 +0,0 @@
|
||||
#!/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[@]}"
|
||||
@@ -1,76 +0,0 @@
|
||||
{
|
||||
"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
|
||||
}
|
||||
]
|
||||
}
|
||||
+44
-42
@@ -1,21 +1,22 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=ode_vidprom2k
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --nodes=8
|
||||
#SBATCH --ntasks=8
|
||||
#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 --output=VSA_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=VSA_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv2
|
||||
conda activate your_env
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
@@ -29,45 +30,41 @@ 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=VIDEO_SPARSE_ATTN
|
||||
# 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"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
# 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
|
||||
--tracker_project_name wan_t2v_VSA
|
||||
--output_dir "checkpoints/wan_t2v_finetune_VSA"
|
||||
--max_train_steps 4000
|
||||
--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"
|
||||
--num_frames 81
|
||||
# --enable_gradient_checkpointing_type "full" # if OOM enable this
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--num_gpus 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 8
|
||||
--hsdp_replicate_dim 64
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
@@ -79,25 +76,25 @@ model_args=(
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
--validation_guidance_scale "5.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-6
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 2000
|
||||
--weight_decay 1e-4
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -106,26 +103,31 @@ 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"
|
||||
--flow_shift 1
|
||||
--seed 1000
|
||||
)
|
||||
|
||||
# VSA arguments
|
||||
vsa_args=(
|
||||
--VSA_decay_rate 0.03 \
|
||||
--VSA_decay_interval_steps 50 \
|
||||
--VSA_sparsity 0.9 \
|
||||
)
|
||||
|
||||
# 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 \
|
||||
--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/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${vsa_args[@]}"
|
||||
@@ -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_old"
|
||||
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/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 50
|
||||
--validation_steps 200
|
||||
--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_old/"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
GPU_NUM=2 # 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 8 \
|
||||
--preprocess.preprocess_video_batch_size 2 \
|
||||
--preprocess.dataloader_num_workers 0 \
|
||||
--preprocess.max_height 480 \
|
||||
--preprocess.max_width 832 \
|
||||
|
||||
@@ -28,4 +28,4 @@
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,94 +0,0 @@
|
||||
#!/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[@]}"
|
||||
@@ -1,24 +0,0 @@
|
||||
#!/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"
|
||||
@@ -5,7 +5,6 @@ from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from flash_attn.bert_padding import pad_input
|
||||
|
||||
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
|
||||
process_moba_output)
|
||||
@@ -134,6 +133,8 @@ class VMOBAAttentionImpl(AttentionImpl):
|
||||
**extra_impl_args) -> None:
|
||||
self.prefix = prefix
|
||||
self.layer_idx = self._get_layer_idx(prefix)
|
||||
from flash_attn.bert_padding import pad_input
|
||||
self.pad_input = pad_input
|
||||
|
||||
def _get_layer_idx(self, prefix: str) -> int | None:
|
||||
match = re.search(r"blocks\.(\d+)", prefix)
|
||||
@@ -169,7 +170,6 @@ class VMOBAAttentionImpl(AttentionImpl):
|
||||
moba_chunk_size = attn_metadata.st_chunk_size
|
||||
moba_topk = attn_metadata.st_topk
|
||||
|
||||
# torch.distributed.breakpoint()
|
||||
query, chunk_size = process_moba_input(query,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
@@ -205,8 +205,8 @@ class VMOBAAttentionImpl(AttentionImpl):
|
||||
simsum_threshold=attn_metadata.moba_threshold,
|
||||
threshold_type=attn_metadata.moba_threshold_type,
|
||||
)
|
||||
hidden_states = pad_input(hidden_states, indices_q, batch_size,
|
||||
sequence_length)
|
||||
hidden_states = self.pad_input(hidden_states, indices_q, batch_size,
|
||||
sequence_length)
|
||||
hidden_states = process_moba_output(hidden_states,
|
||||
attn_metadata.patch_resolution,
|
||||
moba_chunk_size)
|
||||
|
||||
@@ -27,6 +27,7 @@ class DiTArchConfig(ArchConfig):
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 0
|
||||
exclude_lora_layers: list[str] = field(default_factory=list)
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self._compile_conditions:
|
||||
|
||||
@@ -4,12 +4,13 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.wan import (WanI2V480PConfig, WanI2V720PConfig,
|
||||
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
|
||||
WanI2V480PConfig, WanI2V720PConfig,
|
||||
WanT2V480PConfig, WanT2V720PConfig)
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
"SelfForcingWanT2V480PConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -45,8 +45,6 @@ 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)
|
||||
@@ -87,6 +85,10 @@ class PipelineConfig:
|
||||
# DMD parameters
|
||||
dmd_denoising_steps: list[int] | None = field(default=None)
|
||||
|
||||
# Wan2.2 TI2V parameters
|
||||
ti2v_task: bool = False
|
||||
boundary_ratio: float | None = None
|
||||
|
||||
# Compilation
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
@@ -213,24 +215,6 @@ 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")
|
||||
@@ -262,9 +246,7 @@ 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))
|
||||
|
||||
|
||||
@@ -11,9 +11,9 @@ from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
# isort: off
|
||||
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, SelfForcingWanT2V480PConfig)
|
||||
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
|
||||
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
|
||||
SelfForcingWanT2V480PConfig)
|
||||
# isort: on
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
@@ -48,8 +48,8 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"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(),
|
||||
"stepvideo": lambda id: "stepvideo" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
@@ -61,8 +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,
|
||||
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
|
||||
"stepvideo": StepVideoT2VConfig
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
@@ -37,7 +37,7 @@ def getdataset(args) -> VideoCaptionMergedDataset:
|
||||
temporal_sample=temporal_sample,
|
||||
transform_topcrop=transform_topcrop,
|
||||
seed=args.seed)
|
||||
|
||||
|
||||
|
||||
def gettextdataset(args) -> TextDataset:
|
||||
return TextDataset(data_merge_path=args.data_merge_path,
|
||||
|
||||
@@ -1,43 +0,0 @@
|
||||
# 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)
|
||||
}
|
||||
|
||||
@@ -1,75 +0,0 @@
|
||||
# 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
|
||||
@@ -3,9 +3,6 @@ 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) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
|
||||
@@ -158,7 +158,6 @@ class FastVideoArgs:
|
||||
"transformer": True,
|
||||
"vae": True,
|
||||
})
|
||||
override_transformer_cls_name: str | None = None
|
||||
|
||||
# # DMD parameters
|
||||
# dmd_denoising_steps: List[int] | None = field(default=None)
|
||||
@@ -397,12 +396,6 @@ 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)
|
||||
|
||||
@@ -705,7 +698,6 @@ 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
|
||||
|
||||
@@ -276,4 +276,4 @@ class LayerNormScaleShift(nn.Module):
|
||||
if self.compute_dtype == torch.float32:
|
||||
output = output.to(x.dtype)
|
||||
|
||||
return output
|
||||
return output
|
||||
@@ -147,8 +147,8 @@ class CausalWanSelfAttention(nn.Module):
|
||||
# Assign new keys/values directly up to current_end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
# kv_cache["k"] = kv_cache["k"].detach()
|
||||
# kv_cache["v"] = kv_cache["v"].detach()
|
||||
kv_cache["k"] = kv_cache["k"].detach()
|
||||
kv_cache["v"] = kv_cache["v"].detach()
|
||||
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
@@ -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,10 +415,6 @@ 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
|
||||
|
||||
@@ -464,11 +460,11 @@ 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,
|
||||
# TODO(will): make these configurable
|
||||
default_dtype=default_dtype,
|
||||
param_dtype=torch.bfloat16,
|
||||
reduce_dtype=torch.float32,
|
||||
output_dtype=None,
|
||||
@@ -477,11 +473,9 @@ 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)
|
||||
|
||||
logger.info("Converting model to dtype: %s", default_dtype)
|
||||
model = model.to(default_dtype)
|
||||
assert next(model.parameters()).dtype == default_dtype, "Model dtype does not match default dtype"
|
||||
|
||||
model = model.eval()
|
||||
return model
|
||||
|
||||
|
||||
@@ -88,6 +88,7 @@ def maybe_load_fsdp_model(
|
||||
mp_policy=mp_policy,
|
||||
)
|
||||
|
||||
logger.info("Loading model with default_dtype: %s", default_dtype)
|
||||
with set_default_dtype(default_dtype), torch.device("meta"):
|
||||
model = model_cls(**init_params)
|
||||
|
||||
@@ -126,7 +127,6 @@ def maybe_load_fsdp_model(
|
||||
model,
|
||||
weight_iterator,
|
||||
device,
|
||||
# param_dtype,
|
||||
default_dtype,
|
||||
strict=True,
|
||||
cpu_offload=cpu_offload,
|
||||
@@ -139,8 +139,6 @@ 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
|
||||
|
||||
|
||||
|
||||
@@ -64,8 +64,15 @@ class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
|
||||
def step(self, model_output: torch.FloatTensor, timestep: torch.FloatTensor, sample: torch.FloatTensor, to_final=False, return_dict=False, **kwargs):
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
elif timestep.ndim == 0:
|
||||
# handles the case where timestep is a scalar, this occurs when we
|
||||
# use this scheduler for ODE trajectory
|
||||
timestep = timestep.unsqueeze(0)
|
||||
|
||||
self.sigmas = self.sigmas.to(model_output.device)
|
||||
self.timesteps = self.timesteps.to(model_output.device)
|
||||
timestep = timestep.to(model_output.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)
|
||||
|
||||
@@ -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.double().to(device)
|
||||
noise_input_latent = noise_input_latent.double().to(device)
|
||||
sigmas = scheduler.sigmas.double().to(device)
|
||||
timesteps = scheduler.timesteps.double().to(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)
|
||||
timestep_id = torch.argmin(
|
||||
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
@@ -7,8 +7,6 @@ This module wires the causal DMD denoising stage into the modular pipeline.
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
|
||||
# isort: off
|
||||
|
||||
@@ -40,7 +40,7 @@ class ComposedPipelineBase(ABC):
|
||||
_extra_config_module_map: dict[str, str] = {}
|
||||
training_args: TrainingArgs | None = None
|
||||
fastvideo_args: FastVideoArgs | TrainingArgs | None = None
|
||||
modules: dict[str, torch.nn.Module] = {}
|
||||
modules: dict[str, Any] = {}
|
||||
post_init_called: bool = False
|
||||
|
||||
# TODO(will): args should support both inference args and training args
|
||||
@@ -121,25 +121,14 @@ class ComposedPipelineBase(ABC):
|
||||
model_path: str,
|
||||
device: str | None = None,
|
||||
torch_dtype: torch.dtype | None = None,
|
||||
pipeline_config: PipelineConfig | None = None,
|
||||
pipeline_config: str | 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.
|
||||
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
|
||||
|
||||
Load a pipeline from a pretrained model.
|
||||
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
|
||||
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
|
||||
"""
|
||||
@@ -147,18 +136,9 @@ 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():
|
||||
@@ -169,8 +149,7 @@ 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.
|
||||
fastvideo_args.pipeline_config.dit_precision = 'fp32'
|
||||
# assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
|
||||
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
|
||||
|
||||
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
|
||||
|
||||
|
||||
@@ -151,7 +151,7 @@ class ForwardBatch:
|
||||
output: torch.Tensor | None = None
|
||||
return_trajectory_latents: bool = False
|
||||
return_trajectory_decoded: bool = False
|
||||
trajectory_timesteps: list[int] | None = None
|
||||
trajectory_timesteps: list[torch.Tensor] | None = None
|
||||
trajectory_latents: torch.Tensor | None = None
|
||||
trajectory_decoded: list[torch.Tensor] | None = None
|
||||
|
||||
@@ -246,7 +246,6 @@ 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)
|
||||
|
||||
|
||||
|
||||
@@ -12,25 +12,28 @@ 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.dataset import gettextdataset
|
||||
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
|
||||
records_to_table)
|
||||
from fastvideo.dataset.dataloader.record_schema import (
|
||||
ode_text_only_record_creator)
|
||||
from fastvideo.dataset.dataloader.schema import (
|
||||
pyarrow_schema_ode_trajectory_text_only)
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
|
||||
SelfForcingFlowMatchScheduler)
|
||||
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,
|
||||
@@ -40,128 +43,6 @@ 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."""
|
||||
|
||||
@@ -170,28 +51,22 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
]
|
||||
preprocess_dataloader: StatefulDataLoader
|
||||
preprocess_loader_iter: Iterator[dict[str, Any]]
|
||||
pbar: Any
|
||||
num_processed_samples: int
|
||||
|
||||
def get_schema_fields(self):
|
||||
"""Get the schema fields for ODE Trajectory pipeline."""
|
||||
return [f.name for f in pyarrow_schema_ode_trajectory]
|
||||
def get_pyarrow_schema(self) -> pa.Schema:
|
||||
"""Return the PyArrow schema for ODE Trajectory pipeline."""
|
||||
return pyarrow_schema_ode_trajectory_text_only
|
||||
|
||||
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(
|
||||
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=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())
|
||||
@@ -200,9 +75,6 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
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")))
|
||||
@@ -213,56 +85,47 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
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):
|
||||
def preprocess_text_and_trajectory(self, fastvideo_args: FastVideoArgs,
|
||||
args):
|
||||
"""Preprocess text-only data and generate trajectory information."""
|
||||
|
||||
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)
|
||||
# For text-only processing, we only need text data
|
||||
# Filter out samples without text
|
||||
valid_indices = []
|
||||
for i, pixel_values in enumerate(data["pixel_values"]):
|
||||
if not torch.all(
|
||||
pixel_values == 0): # Check if all values are zero
|
||||
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
|
||||
# Create new batch with only valid samples (text-only)
|
||||
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)
|
||||
# Add fps and duration if available in data
|
||||
if "fps" in data:
|
||||
valid_data["fps"] = [data["fps"][i] for i in valid_indices]
|
||||
if "duration" in data:
|
||||
valid_data["duration"] = [
|
||||
data["duration"][i] for i in valid_indices
|
||||
]
|
||||
|
||||
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,
|
||||
@@ -274,27 +137,6 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
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
|
||||
@@ -315,52 +157,34 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
trajectory_latents = []
|
||||
trajectory_timesteps = []
|
||||
trajectory_decoded = []
|
||||
|
||||
for i, (prompt_embed, prompt_attention_mask) in enumerate(
|
||||
zip(prompt_embeds, prompt_attention_masks, strict=False)):
|
||||
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,
|
||||
)
|
||||
|
||||
# Collect the trajectory data (text-to-video generation)
|
||||
batch = ForwardBatch(**shallow_asdict(sampling_params), )
|
||||
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.num_inference_steps = 48
|
||||
batch.return_trajectory_latents = True
|
||||
# Enabling this will save the decoded trajectory videos.
|
||||
# Used for debugging.
|
||||
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(
|
||||
@@ -369,252 +193,97 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
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)}")
|
||||
# Prepare extra features for text-only processing
|
||||
extra_features = {
|
||||
"trajectory_latents": trajectory_latents,
|
||||
"trajectory_timesteps": 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 = []
|
||||
if batch.return_trajectory_decoded:
|
||||
for i, decoded_frames in enumerate(trajectory_decoded):
|
||||
for j, decoded_frame in enumerate(decoded_frames):
|
||||
save_decoded_latents_as_video(
|
||||
decoded_frame,
|
||||
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
|
||||
args.train_fps)
|
||||
|
||||
# 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]
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data: list[dict[str, Any]] = []
|
||||
|
||||
# Convert tensors to numpy arrays
|
||||
vae_latent = latent.cpu().numpy()
|
||||
text_embedding = prompt_embeds[idx].cpu().numpy()
|
||||
# Add progress bar for saving outputs
|
||||
save_pbar = tqdm(enumerate(valid_data["path"]),
|
||||
desc="Saving outputs",
|
||||
unit="item",
|
||||
leave=False)
|
||||
|
||||
# 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}")
|
||||
for idx, video_path in save_pbar:
|
||||
video_name = os.path.basename(video_path).split(".")[0]
|
||||
|
||||
# Convert tensors to numpy arrays
|
||||
text_embedding = prompt_embeds[idx].cpu().numpy()
|
||||
|
||||
# Get extra features for this sample
|
||||
sample_extra_features = {}
|
||||
if extra_features:
|
||||
for key, value in extra_features.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
sample_extra_features[key] = value[idx].cpu(
|
||||
).float().numpy()
|
||||
).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]
|
||||
assert isinstance(value, list)
|
||||
if isinstance(value[idx], torch.Tensor):
|
||||
sample_extra_features[key] = value[idx].cpu(
|
||||
).float().numpy()
|
||||
else:
|
||||
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)
|
||||
# Create record for Parquet dataset (text-only ODE schema)
|
||||
record: dict[str, Any] = ode_text_only_record_creator(
|
||||
video_name=video_name,
|
||||
text_embedding=text_embedding,
|
||||
caption=valid_data["text"][idx],
|
||||
trajectory_latents=sample_extra_features[
|
||||
"trajectory_latents"],
|
||||
trajectory_timesteps=sample_extra_features[
|
||||
"trajectory_timesteps"],
|
||||
)
|
||||
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]))
|
||||
if batch_data:
|
||||
write_pbar = tqdm(total=1,
|
||||
desc="Writing to Parquet dataset",
|
||||
unit="batch")
|
||||
table = records_to_table(batch_data,
|
||||
self.get_pyarrow_schema())
|
||||
write_pbar.update(1)
|
||||
write_pbar.close()
|
||||
|
||||
table = pa.Table.from_arrays(arrays,
|
||||
names=self.get_schema_fields())
|
||||
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)
|
||||
|
||||
# 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))
|
||||
|
||||
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
|
||||
|
||||
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
|
||||
# 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)
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
|
||||
if not self.post_init_called:
|
||||
@@ -628,7 +297,7 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
os.makedirs(self.combined_parquet_dir, exist_ok=True)
|
||||
|
||||
# Loading dataset
|
||||
train_dataset = getdataset(args)
|
||||
train_dataset = gettextdataset(args)
|
||||
|
||||
self.preprocess_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
@@ -648,7 +317,7 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
|
||||
# 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)
|
||||
self.preprocess_text_and_trajectory(fastvideo_args, args)
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_ODE_Trajectory
|
||||
|
||||
@@ -56,6 +56,10 @@ def main(args) -> None:
|
||||
PreprocessPipeline = PreprocessPipeline_I2V
|
||||
elif args.preprocess_task == "text_only":
|
||||
PreprocessPipeline = PreprocessPipeline_Text
|
||||
elif args.preprocess_task == "ode_trajectory":
|
||||
assert args.flow_shift is not None, "flow_shift is required for ode_trajectory"
|
||||
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
|
||||
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
|
||||
else:
|
||||
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
|
||||
f"Valid options: t2v, i2v, ode_trajectory, text_only")
|
||||
@@ -101,10 +105,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("--flow_shift", type=float, default=None)
|
||||
parser.add_argument("--preprocess_task",
|
||||
type=str,
|
||||
default="t2v",
|
||||
choices=["t2v", "i2v", "text_only"],
|
||||
choices=["t2v", "i2v", "text_only", "ode_trajectory"],
|
||||
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)
|
||||
|
||||
@@ -78,8 +78,6 @@ 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)
|
||||
|
||||
|
||||
@@ -53,7 +53,20 @@ class DecodingStage(PipelineStage):
|
||||
@torch.no_grad()
|
||||
def decode(self, latents: torch.Tensor,
|
||||
fastvideo_args: FastVideoArgs) -> torch.Tensor:
|
||||
"""Decode latents into pixel space."""
|
||||
"""
|
||||
Decode latent representations into pixel space using VAE.
|
||||
|
||||
Args:
|
||||
latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents)
|
||||
fastvideo_args: Configuration containing:
|
||||
- disable_autocast: Whether to disable automatic mixed precision (default: False)
|
||||
- pipeline_config.vae_precision: VAE computation precision ("fp32", "fp16", "bf16")
|
||||
- pipeline_config.vae_tiling: Whether to enable VAE tiling for memory efficiency
|
||||
|
||||
Returns:
|
||||
Decoded video tensor with shape (batch, channels, frames, height, width),
|
||||
normalized to [0, 1] range and moved to CPU as float32
|
||||
"""
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
latents = latents.to(get_local_torch_device())
|
||||
|
||||
@@ -103,12 +116,26 @@ class DecodingStage(PipelineStage):
|
||||
"""
|
||||
Decode latent representations into pixel space.
|
||||
|
||||
This method processes the batch through the VAE decoder, converting latent
|
||||
representations to pixel-space video/images. It also optionally decodes
|
||||
trajectory latents for visualization purposes.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
batch: The current batch containing:
|
||||
- latents: Tensor to decode (batch, channels, frames, height_latents, width_latents)
|
||||
- return_trajectory_decoded (optional): Flag to decode trajectory latents
|
||||
- trajectory_latents (optional): Latents at different timesteps
|
||||
- trajectory_timesteps (optional): Corresponding timesteps
|
||||
fastvideo_args: Configuration containing:
|
||||
- output_type: "latent" to skip decoding, otherwise decode to pixels
|
||||
- vae_cpu_offload: Whether to offload VAE to CPU after decoding
|
||||
- model_loaded: Track VAE loading state
|
||||
- model_paths: Path to VAE model if loading needed
|
||||
|
||||
Returns:
|
||||
The batch with decoded outputs.
|
||||
Modified batch with:
|
||||
- output: Decoded frames (batch, channels, frames, height, width) as CPU float32
|
||||
- trajectory_decoded (if requested): List of decoded frames per timestep
|
||||
"""
|
||||
# load vae if not already loaded (used for memory constrained devices)
|
||||
pipeline = self.pipeline() if self.pipeline else None
|
||||
@@ -128,15 +155,13 @@ class DecodingStage(PipelineStage):
|
||||
# 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]
|
||||
# batch.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}")
|
||||
logger.info("decoding trajectory latent for timestep: %s",
|
||||
cur_timestep)
|
||||
decoded_frames = self.decode(cur_latent, fastvideo_args)
|
||||
batch.trajectory_decoded.append(decoded_frames.cpu().float())
|
||||
|
||||
|
||||
@@ -140,12 +140,11 @@ class DenoisingStage(PipelineStage):
|
||||
latents = latents[:, :, rank_in_sp_group, :, :, :]
|
||||
batch.latents = latents
|
||||
if batch.image_latent is not None:
|
||||
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
|
||||
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
|
||||
@@ -205,14 +204,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.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_ratio = fastvideo_args.pipeline_config.dit_config.boundary_ratio
|
||||
if batch.boundary_ratio is not None:
|
||||
logger.info("Overriding boundary ratio from %s to %s",
|
||||
boundary_ratio, batch.boundary_ratio)
|
||||
boundary_ratio = batch.boundary_ratio
|
||||
|
||||
boundary_timestep *= self.scheduler.num_train_timesteps
|
||||
if boundary_ratio is not None:
|
||||
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
|
||||
else:
|
||||
boundary_timestep = None
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
@@ -254,7 +253,8 @@ class DenoisingStage(PipelineStage):
|
||||
patch_size[2])
|
||||
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
|
||||
|
||||
trajectory_timesteps: list[int] = []
|
||||
# Initialize lists for ODE trajectory
|
||||
trajectory_timesteps: list[torch.Tensor] = []
|
||||
trajectory_latents: list[torch.Tensor] = []
|
||||
|
||||
# Run denoising loop
|
||||
@@ -284,27 +284,14 @@ class DenoisingStage(PipelineStage):
|
||||
|
||||
# Expand latents for I2V
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if batch.image_latent is not None and not fastvideo_args.pipeline_config.t2v_as_i2v_task:
|
||||
if batch.image_latent is not None:
|
||||
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()
|
||||
@@ -319,13 +306,6 @@ 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 = (
|
||||
@@ -460,7 +440,6 @@ class DenoisingStage(PipelineStage):
|
||||
# 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
|
||||
@@ -474,30 +453,25 @@ class DenoisingStage(PipelineStage):
|
||||
trajectory_tensor: torch.Tensor | None = None
|
||||
if trajectory_latents:
|
||||
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
|
||||
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
|
||||
dim=0)
|
||||
else:
|
||||
trajectory_tensor = None
|
||||
trajectory_timesteps_tensor = None
|
||||
|
||||
# Gather results if using sequence parallelism
|
||||
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()
|
||||
if trajectory_tensor is not None and trajectory_timesteps_tensor is not None:
|
||||
batch.trajectory_timesteps = trajectory_timesteps_tensor.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,81 +105,6 @@ 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,
|
||||
@@ -232,28 +157,57 @@ class ImageVAEEncodingStage(PipelineStage):
|
||||
# (B, C, H, W) -> (B, C, 1, H, W)
|
||||
image = image.unsqueeze(2)
|
||||
|
||||
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)
|
||||
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)
|
||||
|
||||
latent_condition = self.encode_tensor(video_condition, fastvideo_args,
|
||||
batch.generator)
|
||||
# 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
|
||||
|
||||
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,
|
||||
|
||||
@@ -35,15 +35,9 @@ 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(total_num_videos)]
|
||||
seeds = [seed + i for i in range(num_videos_per_prompt)]
|
||||
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:
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.info("ROCm platform is unavailable: %s", e)
|
||||
|
||||
return "fastvideo.platforms.rocm.RocmPlatform" if is_rocm else None
|
||||
|
||||
|
||||
Binary file not shown.
@@ -0,0 +1,184 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
from huggingface_hub import snapshot_download
|
||||
import subprocess
|
||||
import sys
|
||||
from fastvideo.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
|
||||
import shutil
|
||||
|
||||
# Import the training pipeline
|
||||
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
|
||||
|
||||
NUM_NODES = "1"
|
||||
MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
|
||||
# preprocessing
|
||||
DATA_DIR = "data"
|
||||
LOCAL_RAW_DATA_DIR = Path(os.path.join(DATA_DIR, "crush-smol"))
|
||||
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
|
||||
PREPROCESSING_ENTRY_FILE_PATH = "fastvideo/pipelines/preprocess/v1_preprocess.py"
|
||||
|
||||
LOCAL_PREPROCESSED_DATA_DIR = Path(os.path.join(DATA_DIR, "crush-smol_processed_t2v"))
|
||||
|
||||
|
||||
# training
|
||||
NUM_GPUS_PER_NODE_TRAINING = "4"
|
||||
TRAINING_ENTRY_FILE_PATH = "fastvideo/training/wan_distillation_pipeline.py"
|
||||
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "combined_parquet_dataset")
|
||||
LOCAL_VALIDATION_DATASET_FILE = "examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
LOCAL_OUTPUT_DIR = Path(os.path.join(DATA_DIR, "outputs"))
|
||||
|
||||
def download_data():
|
||||
# create the data dir if it doesn't exist
|
||||
data_dir = Path(DATA_DIR)
|
||||
|
||||
print(f"Creating data directory at {data_dir}")
|
||||
os.makedirs(data_dir, exist_ok=True)
|
||||
|
||||
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
|
||||
try:
|
||||
result = snapshot_download(
|
||||
repo_id="wlsaidhi/crush-smol-merged",
|
||||
local_dir=str(LOCAL_RAW_DATA_DIR),
|
||||
repo_type="dataset",
|
||||
resume_download=True,
|
||||
token=os.environ.get("HF_TOKEN"), # In case authentication is needed
|
||||
)
|
||||
print(f"Download completed successfully. Files downloaded to: {result}")
|
||||
|
||||
# Verify the download
|
||||
if not LOCAL_RAW_DATA_DIR.exists():
|
||||
raise RuntimeError(f"Download appeared to succeed but {LOCAL_RAW_DATA_DIR} does not exist")
|
||||
|
||||
# List downloaded files
|
||||
print("Downloaded files:")
|
||||
for file in LOCAL_RAW_DATA_DIR.rglob("*"):
|
||||
if file.is_file():
|
||||
print(f" - {file.relative_to(LOCAL_RAW_DATA_DIR)}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during download: {str(e)}")
|
||||
raise
|
||||
|
||||
|
||||
def run_preprocessing():
|
||||
# remove the local_preprocessed_data_dir if it exists
|
||||
if LOCAL_PREPROCESSED_DATA_DIR.exists():
|
||||
print(f"Removing local_preprocessed_data_dir: {LOCAL_PREPROCESSED_DATA_DIR}")
|
||||
shutil.rmtree(LOCAL_PREPROCESSED_DATA_DIR)
|
||||
|
||||
# Run torchrun command
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
|
||||
PREPROCESSING_ENTRY_FILE_PATH,
|
||||
"--model_path", MODEL_PATH,
|
||||
"--seed", "42",
|
||||
"--data_merge_path", os.path.join(LOCAL_RAW_DATA_DIR, "merge.txt"),
|
||||
"--preprocess_video_batch_size", "1",
|
||||
"--max_height", "480",
|
||||
"--max_width", "832",
|
||||
"--num_frames", "81",
|
||||
"--dataloader_num_workers", "0",
|
||||
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
|
||||
"--train_fps", "16",
|
||||
"--samples_per_file", "1",
|
||||
"--flush_frequency", "1",
|
||||
"--video_length_tolerance_range", "5",
|
||||
"--preprocess_task", "t2v",
|
||||
]
|
||||
|
||||
process = subprocess.run(cmd, check=True)
|
||||
|
||||
|
||||
def run_training():
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE_TRAINING,
|
||||
TRAINING_ENTRY_FILE_PATH,
|
||||
"--model_path", MODEL_PATH,
|
||||
"--inference_mode", "False",
|
||||
"--pretrained_model_name_or_path", MODEL_PATH,
|
||||
"--data_path", LOCAL_TRAINING_DATA_DIR,
|
||||
"--validation_dataset_file", LOCAL_VALIDATION_DATASET_FILE,
|
||||
"--train_batch_size", "1",
|
||||
"--num_latent_t", "8",
|
||||
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--sp_size", "1",
|
||||
"--tp_size", "1",
|
||||
"--hsdp_replicate_dim", "1",
|
||||
"--hsdp_shard_dim", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--train_sp_batch_size", "1",
|
||||
"--dataloader_num_workers", "10",
|
||||
"--gradient_accumulation_steps", "1",
|
||||
"--max_train_steps", "501",
|
||||
"--learning_rate", "2e-6",
|
||||
"--fake_score_learning_rate", "2e-6",
|
||||
"--mixed_precision", "bf16",
|
||||
"--training_state_checkpointing_steps", "1000",
|
||||
"--weight_only_checkpointing_steps", "1000",
|
||||
"--validation_steps", "50",
|
||||
"--validation_sampling_steps", "3",
|
||||
"--log_validation",
|
||||
"--checkpoints_total_limit", "3",
|
||||
"--ema_start_step", "0",
|
||||
"--training_cfg_rate", "0.0",
|
||||
"--output_dir", LOCAL_OUTPUT_DIR,
|
||||
"--tracker_project_name", "ci_wan_t2v_dmd_overfit",
|
||||
"--num_height", "480",
|
||||
"--num_width", "832",
|
||||
"--num_frames", "81",
|
||||
"--flow_shift", "8",
|
||||
"--validation_guidance_scale", "6.0",
|
||||
"--weight_decay", "0.01",
|
||||
"--generator_update_interval", "5",
|
||||
"--dmd_denoising_steps", "1000,757,522",
|
||||
"--min_timestep_ratio", "0.02",
|
||||
"--max_timestep_ratio", "0.98",
|
||||
"--seed", "1000",
|
||||
"--real_score_guidance_scale", "3.5",
|
||||
"--dit_precision", "fp32",
|
||||
"--max_grad_norm", "1.0",
|
||||
"--enable_gradient_checkpointing_type", "full",
|
||||
]
|
||||
|
||||
print(f"Running training with command: {cmd}")
|
||||
process = subprocess.run(cmd, check=True)
|
||||
|
||||
|
||||
def test_e2e_overfit_single_sample():
|
||||
os.environ["WANDB_MODE"] = "online"
|
||||
|
||||
download_data()
|
||||
run_preprocessing()
|
||||
run_training()
|
||||
|
||||
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_v0.mp4")
|
||||
print(f"reference_video_file: {reference_video_file}")
|
||||
final_validation_video_file = os.path.join(LOCAL_OUTPUT_DIR, "validation_step_900_inference_steps_50_video_0.mp4")
|
||||
print(f"final_validation_video_file: {final_validation_video_file}")
|
||||
|
||||
|
||||
# Ensure both files exist
|
||||
assert os.path.exists(reference_video_file), f"Reference video not found at {reference_video_file}"
|
||||
assert os.path.exists(final_validation_video_file), f"Validation video not found at {final_validation_video_file}"
|
||||
|
||||
# Compute SSIM
|
||||
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
|
||||
reference_video_file,
|
||||
final_validation_video_file,
|
||||
use_ms_ssim=True # Using MS-SSIM for better quality assessment
|
||||
)
|
||||
|
||||
print("\n===== SSIM Results for Step 900 Validation =====")
|
||||
print(f"Mean MS-SSIM: {mean_ssim:.4f}")
|
||||
print(f"Min MS-SSIM: {min_ssim:.4f}")
|
||||
print(f"Max MS-SSIM: {max_ssim:.4f}")
|
||||
|
||||
assert max_ssim > 0.5, f"Max SSIM is below 0.5: {max_ssim}"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_e2e_overfit_single_sample()
|
||||
@@ -62,6 +62,11 @@ def download_data():
|
||||
|
||||
|
||||
def run_preprocessing():
|
||||
# remove the local_preprocessed_data_dir if it exists
|
||||
if LOCAL_PREPROCESSED_DATA_DIR.exists():
|
||||
print(f"Removing local_preprocessed_data_dir: {LOCAL_PREPROCESSED_DATA_DIR}")
|
||||
shutil.rmtree(LOCAL_PREPROCESSED_DATA_DIR)
|
||||
|
||||
# Run torchrun command
|
||||
cmd = [
|
||||
"torchrun",
|
||||
|
||||
@@ -1 +1 @@
|
||||
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":0.50390625,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.19960195198655128,"_runtime":107.325113071}
|
||||
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":1.260593056678772,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.2620866410434246,"_runtime":107.325113071}
|
||||
@@ -110,7 +110,7 @@ def test_distributed_training():
|
||||
'avg_step_time': 1.0,
|
||||
'grad_norm': 0.1,
|
||||
'step_time': 1.0,
|
||||
'train_loss': 0.001
|
||||
'train_loss': 0.005
|
||||
}
|
||||
|
||||
failures = []
|
||||
|
||||
@@ -1 +1 @@
|
||||
{"train_loss":0.10545400530099869,"_step":5,"_wandb":{"runtime":35},"step_time":2.2575189135968685,"grad_norm":0.53125,"_runtime":35.701958502,"avg_step_time":2.4754387199878694,"_timestamp":1.7525528420185745e+09,"learning_rate":1e-06,"vsa_sparsity":0,"validation_videos_8_steps":{"videos":[{"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.","_type":"video-file","sha256":"ee83e83df073a648f89dcd288cccaed9af765e11fc35c77a6e0cb2ebaa1be5b0","size":475248,"path":"media/videos/validation_videos_8_steps_0_ee83e83df073a648f89d.mp4"},{"size":341490,"path":"media/videos/validation_videos_8_steps_0_d0b2758549d5c82845ca.mp4","caption":"A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","_type":"video-file","sha256":"d0b2758549d5c82845ca3c1ac0db6812877631495e77ac041acdc8ea31f0a5ee"},{"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.","_type":"video-file","sha256":"08638381f1607d6ab10684772be38a58b18628fd402a208b25f6af96e765d454","size":436814,"path":"media/videos/validation_videos_8_steps_0_08638381f1607d6ab106.mp4"}],"captions":["A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.","A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","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."],"_type":"videos","count":3}}
|
||||
{"train_loss":0.1021774671971798,"_step":5,"_wandb":{"runtime":35},"step_time":2.2575189135968685,"grad_norm":0.11582941561937332,"_runtime":35.701958502,"avg_step_time":2.4754387199878694,"_timestamp":1.7525528420185745e+09,"learning_rate":1e-06,"vsa_sparsity":0,"validation_videos_8_steps":{"videos":[{"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.","_type":"video-file","sha256":"ee83e83df073a648f89dcd288cccaed9af765e11fc35c77a6e0cb2ebaa1be5b0","size":475248,"path":"media/videos/validation_videos_8_steps_0_ee83e83df073a648f89d.mp4"},{"size":341490,"path":"media/videos/validation_videos_8_steps_0_d0b2758549d5c82845ca.mp4","caption":"A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","_type":"video-file","sha256":"d0b2758549d5c82845ca3c1ac0db6812877631495e77ac041acdc8ea31f0a5ee"},{"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.","_type":"video-file","sha256":"08638381f1607d6ab10684772be38a58b18628fd402a208b25f6af96e765d454","size":436814,"path":"media/videos/validation_videos_8_steps_0_08638381f1607d6ab106.mp4"}],"captions":["A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.","A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.","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."],"_type":"videos","count":3}}
|
||||
@@ -64,5 +64,3 @@ def test_parquet_dataset_saver_flush_and_last(tmp_path: Path):
|
||||
assert len(files2) == 2
|
||||
total = sum(pq.read_table(str(f)).num_rows for f in files2)
|
||||
assert total == 5
|
||||
|
||||
|
||||
|
||||
@@ -92,9 +92,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
shift=self.timestep_shift)
|
||||
|
||||
if training_args.real_score_model_path:
|
||||
logger.info(
|
||||
f"Loading real score transformer from: {training_args.real_score_model_path}"
|
||||
)
|
||||
logger.info("Loading real score transformer from: %s",
|
||||
training_args.real_score_model_path)
|
||||
self.real_score_transformer = self.load_module_from_path(
|
||||
training_args.real_score_model_path, "transformer",
|
||||
training_args)
|
||||
@@ -103,9 +102,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"real_score_transformer")
|
||||
|
||||
if training_args.fake_score_model_path:
|
||||
logger.info(
|
||||
f"Loading fake score transformer from: {training_args.fake_score_model_path}"
|
||||
)
|
||||
logger.info("Loading fake score transformer from: %s",
|
||||
training_args.fake_score_model_path)
|
||||
self.fake_score_transformer = self.load_module_from_path(
|
||||
training_args.fake_score_model_path, "transformer",
|
||||
training_args)
|
||||
@@ -194,14 +192,13 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
|
||||
|
||||
self.generator_ema = None
|
||||
self.generator_ema: EMA_FSDP | None = None
|
||||
if (self.training_args.ema_decay
|
||||
is not None) and (self.training_args.ema_decay > 0.0):
|
||||
self.generator_ema = EMA_FSDP(self.transformer,
|
||||
decay=self.training_args.ema_decay)
|
||||
logger.info(
|
||||
f"Initialized generator EMA with decay={self.training_args.ema_decay}"
|
||||
)
|
||||
logger.info("Initialized generator EMA with decay=%s",
|
||||
self.training_args.ema_decay)
|
||||
else:
|
||||
logger.info("Generator EMA disabled (ema_decay <= 0.0)")
|
||||
|
||||
@@ -218,7 +215,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
Returns:
|
||||
The loaded module
|
||||
"""
|
||||
logger.info(f"Loading {module_type} from custom path: {model_path}")
|
||||
logger.info("Loading %s from custom path: %s", module_type, model_path)
|
||||
# Set flag to prevent custom weight loading for teacher/critic models
|
||||
training_args._loading_teacher_critic_model = True
|
||||
|
||||
@@ -228,7 +225,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
# Download the model if it's a Hugging Face model ID
|
||||
local_model_path = maybe_download_model(model_path)
|
||||
logger.info(f"Model downloaded/found at: {local_model_path}")
|
||||
logger.info("Model downloaded/found at: %s", local_model_path)
|
||||
config = verify_model_config_and_directory(local_model_path)
|
||||
|
||||
if module_type not in config:
|
||||
@@ -237,7 +234,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
extra_module = self._extra_config_module_map[module_type]
|
||||
if extra_module in config:
|
||||
module_type = extra_module
|
||||
logger.info(f"Using {extra_module} for {module_type}")
|
||||
logger.info("Using %s for %s", extra_module,
|
||||
module_type)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Module {module_type} not found in config at {local_model_path}"
|
||||
@@ -262,8 +260,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
fastvideo_args=training_args,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
f"Successfully loaded {module_type} from {component_path}")
|
||||
logger.info("Successfully loaded %s from %s", module_type,
|
||||
component_path)
|
||||
return module
|
||||
finally:
|
||||
# Always clean up the flag
|
||||
@@ -293,7 +291,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
return model
|
||||
return model
|
||||
|
||||
def get_ema_model_copy(self):
|
||||
def get_ema_model_copy(self) -> torch.nn.Module | None:
|
||||
"""Get a copy of the model with EMA weights applied."""
|
||||
if self.generator_ema is not None:
|
||||
ema_model = copy.deepcopy(self.transformer)
|
||||
@@ -301,7 +299,7 @@ class DistillationPipeline(TrainingPipeline):
|
||||
return ema_model
|
||||
return None
|
||||
|
||||
def is_ema_ready(self, current_step: int = None):
|
||||
def is_ema_ready(self, current_step: int | None = None):
|
||||
"""Check if EMA is ready for use (after ema_start_step)."""
|
||||
if current_step is None:
|
||||
current_step = getattr(self, 'current_trainstep', 0)
|
||||
@@ -350,14 +348,14 @@ class DistillationPipeline(TrainingPipeline):
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
|
||||
logger.info(f"EMA weights saved to {weight_path}")
|
||||
logger.info("EMA weights saved to %s", weight_path)
|
||||
|
||||
del ema_model
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to save EMA weights: {str(e)}")
|
||||
logger.error("Failed to save EMA weights: %s", str(e))
|
||||
|
||||
def get_ema_stats(self):
|
||||
def get_ema_stats(self) -> dict[str, Any]:
|
||||
"""Get EMA statistics for monitoring."""
|
||||
if self.generator_ema is None:
|
||||
return {
|
||||
@@ -571,7 +569,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
|
||||
noisy_latent = self.noise_scheduler.add_noise(
|
||||
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
|
||||
timestep).detach().unflatten(0, (1, generator_pred_video.shape[1]))
|
||||
timestep).detach().unflatten(0,
|
||||
(1, generator_pred_video.shape[1]))
|
||||
|
||||
# fake_score_transformer forward
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
@@ -730,11 +729,12 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"encoder_hidden_states": training_batch.encoder_hidden_states,
|
||||
"encoder_attention_mask": training_batch.encoder_attention_mask,
|
||||
}
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
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
|
||||
|
||||
training_batch.dmd_latent_vis_dict = {}
|
||||
training_batch.fake_score_latent_vis_dict = {}
|
||||
@@ -1251,9 +1251,8 @@ class DistillationPipeline(TrainingPipeline):
|
||||
(self.generator_ema is None) and (self.training_args.ema_decay > 0):
|
||||
self.generator_ema = EMA_FSDP(
|
||||
self.transformer, decay=self.training_args.ema_decay)
|
||||
logger.info(
|
||||
f"Created generator EMA at step {step} with decay={self.training_args.ema_decay}"
|
||||
)
|
||||
logger.info("Created generator EMA at step %s with decay=%s",
|
||||
step, self.training_args.ema_decay)
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
training_batch = self.train_one_step(training_batch)
|
||||
|
||||
@@ -1,476 +0,0 @@
|
||||
# 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
@@ -44,7 +44,8 @@ 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
|
||||
|
||||
@@ -124,7 +125,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,
|
||||
@@ -487,7 +488,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
current_vsa_sparsity = current_decay_times * vsa_decay_rate
|
||||
# elif vmoba_available:
|
||||
# # TODO: add vmoba sparsity scheduling here
|
||||
# current_vsa_sparsity = 0.0
|
||||
# pass
|
||||
else:
|
||||
current_vsa_sparsity = 0.0
|
||||
|
||||
@@ -529,10 +530,6 @@ 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:
|
||||
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(
|
||||
@@ -733,11 +730,4 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
|
||||
# Re-enable gradients for training
|
||||
training_args.inference_mode = False
|
||||
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"
|
||||
)
|
||||
transformer.train()
|
||||
@@ -349,14 +349,10 @@ 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
|
||||
try:
|
||||
step = int(os.path.basename(checkpoint_path).split('-')[-1])
|
||||
except:
|
||||
step = 1
|
||||
step = int(os.path.basename(checkpoint_path).split('-')[-1])
|
||||
|
||||
if rank == 0:
|
||||
logger.info("Loading checkpoint from step %s", step)
|
||||
@@ -1302,6 +1298,10 @@ def get_scheduler(
|
||||
last_epoch=last_epoch)
|
||||
|
||||
|
||||
def count_trainable(model: torch.nn.Module) -> int:
|
||||
return sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
|
||||
|
||||
class EMA_FSDP:
|
||||
"""
|
||||
FSDP2-friendly EMA with two modes:
|
||||
@@ -1323,7 +1323,6 @@ 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
|
||||
@@ -1417,7 +1416,6 @@ 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
|
||||
|
||||
@@ -4,10 +4,10 @@ from copy import deepcopy
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import WanCausalDMDPipeline
|
||||
from fastvideo.training.self_forcing_distillation_pipeline import SelfForcingDistillationPipeline
|
||||
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
|
||||
WanCausalDMDPipeline)
|
||||
from fastvideo.training.self_forcing_distillation_pipeline import (
|
||||
SelfForcingDistillationPipeline)
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
@@ -69,4 +69,4 @@ if __name__ == "__main__":
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
main(args)
|
||||
|
||||
@@ -1,211 +0,0 @@
|
||||
# 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)
|
||||
+6
-5
@@ -2,10 +2,6 @@
|
||||
# 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
|
||||
@@ -27,10 +23,14 @@ from typing import Any, TypeVar, cast
|
||||
|
||||
import cloudpickle
|
||||
import filelock
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
import yaml
|
||||
from diffusers.loaders.lora_base import (
|
||||
_best_guess_weight_name) # watch out for potetential removal from diffusers
|
||||
from einops import rearrange
|
||||
from huggingface_hub import snapshot_download
|
||||
from remote_pdb import RemotePdb
|
||||
from torch.distributed.fsdp import MixedPrecisionPolicy
|
||||
@@ -892,7 +892,8 @@ def best_output_size(w, h, dw, dh, expected_area):
|
||||
return ow2, oh2
|
||||
|
||||
|
||||
def save_decoded_latents_as_video(decoded_latents: list[torch.Tensor], output_path: str, fps: int):
|
||||
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 = []
|
||||
|
||||
+1
-2
@@ -49,8 +49,7 @@ dependencies = [
|
||||
"av",
|
||||
|
||||
# Preprocessing Dependencies
|
||||
"torchcodec==0.5.0",
|
||||
"lmdb"
|
||||
"torchcodec==0.5.0"
|
||||
]
|
||||
|
||||
[tool.uv]
|
||||
|
||||
@@ -1,18 +1,9 @@
|
||||
# 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
|
||||
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
|
||||
|
||||
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]
|
||||
import re
|
||||
from collections import OrderedDict
|
||||
|
||||
_param_names_mapping: dict = {
|
||||
r"^text_embedding\.0\.(.*)$":
|
||||
@@ -144,270 +135,27 @@ _self_forcing_to_diffusers_param_names_mapping: dict = {
|
||||
r"blocks.\1.norm2.\2",
|
||||
}
|
||||
|
||||
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\.(.*)$".
|
||||
|
||||
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)
|
||||
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:
|
||||
# Save a single safetensors file
|
||||
save_file(OrderedDict(new_state_dict), output)
|
||||
# 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 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)
|
||||
save_torch_state_dict(
|
||||
new_state_dict,
|
||||
"new2/",
|
||||
max_shard_size="10GB"
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user