Compare commits

...
25 Commits
Author SHA1 Message Date
Will Lin 6b5b2dc6e7 mp training 2025-05-27 11:38:47 -07:00
Will Lin 91b7cc1be8 move utils into training_utils 2025-05-25 18:23:47 -07:00
Zihang-He d6365373b4 added gradient clipping 2025-05-25 23:54:10 +00:00
Will Lin 50fb94b902 gradient checking 2025-05-25 15:20:41 -07:00
Will Lin 2caa0d4d0b cleanup 2025-05-24 17:07:17 -07:00
Will Lin 24db823998 add validation 2025-05-24 16:27:05 -07:00
Will Lin 8c4704edf5 update train dependecies 2025-05-23 12:48:06 -07:00
JerryZhou54 dfba7ec833 Small fix 2025-05-23 19:03:58 +00:00
JerryZhou54 7a2e171f1b Small fix 2025-05-23 19:01:59 +00:00
JerryZhou54 a8aac6090a Integrate the new parquet dataloader into training pipeline 2025-05-23 18:27:52 +00:00
191d1be3b4 Will/training (#425)
Co-authored-by: Kevin Lin <42618777+kevin314@users.noreply.github.com>
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-05-23 02:06:54 -04:00
JerryZhou54 338ea1e5f2 Add script to upload preprocessed dataset to HF 2025-05-23 06:05:40 +00:00
JerryZhou54 42a2f272d5 Preprocessing Stage: 1. Save to parquets periodically 2. Allow resuming from the middle 3. Doesn't support multi-gpu for now. Data Loader Stage: 1. Allow multi-gpu dataloader 2. Doesn't drop files even if number of parquet files is not divisible by num_gpus 3. Able to resume training 2025-05-23 00:57:39 +00:00
JerryZhou54 982bfcfdc8 Fix small issues when loading mask 2025-05-21 22:46:32 +00:00
JerryZhou54 2ac06379a7 Finish data preprocessing and loading 2025-05-21 22:34:03 +00:00
JerryZhou54 baaa1673f7 Add data preprocessing script for WAN 2025-05-19 15:09:26 +00:00
William Lin ace6e971e5 [V1] Remove vLLM dependency (#413) 2025-05-18 00:57:28 -07:00
William Lin b4f6758253 [Teacache] allow None for forward_context batch when using teacache (#412) 2025-05-17 18:43:14 -07:00
William Lin 535d29b392 [Docs] Fix image (#407) 2025-05-12 14:30:05 -07:00
William Lin b4255517e0 [Docs] Add CLI docs (#406) 2025-05-12 14:16:48 -07:00
William Lin 6eeb60613f Release 0.1.0 (#405) 2025-05-12 11:54:23 -07:00
William Lin 53d2c7791f [V1] Update where num_frame rounding is done (#403) 2025-05-12 11:52:49 -07:00
William Lin 53cb693dca [V1] Docs Update (#402) 2025-05-12 11:52:09 -07:00
Kevin Lin 6f72d24876 [CLI] Default to pipeline config (#401) 2025-05-12 00:22:53 -07:00
William Lin d1459e9976 [V1] Update README (#400) 2025-05-11 22:32:03 -07:00
83 changed files with 6796 additions and 535 deletions
+78 -31
View File
@@ -2,62 +2,109 @@
<img src=assets/logo.jpg width="30%"/>
</div>
FastVideo is a lightweight framework for accelerating large video diffusion models.
**FastVideo is a unified framework for accelerated video generation.**
It features a clean, consistent API that works across popular video models, making it easier for developers to author new models and incorporate system- or kernel-level optimizations.
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
<p align="center">
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</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/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</b> </a> |
</p>
https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1
<div align="center">
<img src=assets/perf.png width="90%"/>
</div>
FastVideo currently offers: (with more to come)
## Key Features
- [NEW!] V1 inference API available. Full announcement coming soon!
- [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
FastVideo has the following features:
- State-of-the-art performance optimizations for inference
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
- [TeaCache](https://arxiv.org/pdf/2411.19108)
- [Sage Attention](https://arxiv.org/abs/2410.02367)
- Cutting edge models
- Wan2.1 T2V, I2V
- HunyuanVideo
- FastHunyuan: consistency distilled video diffusion models for 8x inference speedup.
- StepVideo T2V
- Distillation support
- Recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
Dev in progress and highly experimental.
## Change Log
- ```2025/02/20```: FastVideo now supports STA on [StepVideo](https://github.com/stepfun-ai/Step-Video-T2V) with 3.4X speedup!
- ```2025/02/18```: Release the inference code and kernel for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
- ```2025/01/13```: Support Lora finetuning for HunyuanVideo.
- ```2024/12/25```: Enable single 4090 inference for `FastHunyuan`, please rerun the installation steps to update the environment.
- ```2024/12/17```: `FastVideo` v0.0.1 is released.
## Getting Started
We recommend using an environment manager such as `Conda` to create a clean environment:
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
# Install FastVideo
pip install fastvideo
```
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) for more detailed installation instructions.
## Inference
### Generating Your First Video
Here's a minimal example to generate a video using the default settings. Create a file called `example.py` with the following code:
```python
from fastvideo import VideoGenerator
def main():
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1, # Adjust based on your hardware
)
# Define a prompt for your video
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
# Generate the video
video = generator.generate_video(
prompt,
return_frames=True, # Also return frames from this call (defaults to False)
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
if __name__ == '__main__':
main()
```
Run the script with:
```bash
python example.py
```
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html).
### Other docs:
- [Install FastVideo](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview.html)
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
### Inference
- [Quick Start](https://hao-ai-lab.github.io/FastVideo/inference/examples/basic.html)
- V1 Inference API Guide (Coming soon!)
### Distillation and Finetuning
## Distillation and Finetuning
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetuning.html)
### Deprecated APIs
- [V0 Inference (Deprecated)](https://hao-ai-lab.github.io/FastVideo/inference/v0_inference.html)
## 📑 Development Plan
<!-- - More distillation methods -->
<!-- - [ ] Add Distribution Matching Distillation -->
- More models support
<!-- - [ ] Add CogvideoX model -->
- [ ] Add StepVideo to V1
- [x] Add StepVideo to V1
- Optimization features
- [ ] Teacache in V1
- [ ] SageAttention in V1
- [x] Teacache in V1
- [x] SageAttention in V1
- Code updates
- [ ] V1 Configuration API
- [x] V1 Configuration API
- [ ] Support Training in V1
<!-- - [ ] fp8 support -->
<!-- - [ ] faster load model and save model support -->
BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 303 KiB

+46
View File
@@ -0,0 +1,46 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
DATA_DIR=./data
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
# --gradient_checkpointing\
# --pretrained_model_name_or_path hunyuanvideo-community/HunyuanVideo \
# --pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
torchrun --nnodes 1 --nproc_per_node 4\
fastvideo/v1/pipelines/training_pipeline.py\
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache"\
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
--train_batch_size=1\
--num_latent_t 1 \
--sp_size 4 \
--tp_size 4 \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=320\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "2,4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_16_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver \
--master_weight_type "bf16"
Binary file not shown.

After

Width:  |  Height:  |  Size: 303 KiB

+19 -14
View File
@@ -9,7 +9,7 @@
:::{raw} html
<p style="text-align:center">
<strong>FastVideo is a lightweight framework for accelerating large video diffusion models.
<strong>FastVideo is a unified framework for accelerated video generation.
</strong>
</p>
@@ -21,27 +21,31 @@
</p>
:::
FastVideo is a lightweight framework for accelerating large video diffusion models developed by the [Hao AI Lab](https://hao-ai-lab.github.io/).
It features a clean, consistent API that works across popular video models, making it easier for developers to author new models and incorporate system- or kernel-level optimizations.
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
<div style="text-align: center;">
<video controls width="800">
<source src="https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1" type="video/mp4">
Your browser does not support the video tag.
</video>
<img src=_static/images/perf.png width="100%"/>
</div>
FastVideo currently offers: (with more to come)
## Key Features
- [NEW!] V1 inference API available. Full announcement coming soon!
- [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
FastVideo has the following features:
- State-of-the-art performance optimizations for inference
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
- [TeaCache](https://arxiv.org/pdf/2411.19108)
- [Sage Attention](https://arxiv.org/abs/2410.02367)
- Cutting edge models
- Wan2.1 T2V, I2V
- HunyuanVideo
- FastHunyuan: consistency distilled video diffusion models for 8x inference speedup.
- StepVideo T2V
- Distillation support
- Recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
Dev in progress and highly experimental.
## Documentation
% How to start using FastVideo?
@@ -63,6 +67,7 @@ inference/configuration
inference/optimizations
inference/support_matrix
inference/examples/examples_inference_index
inference/cli
inference/add_pipeline
inference/v0_inference
:::
+77
View File
@@ -0,0 +1,77 @@
(inference-configuration)=
# Configuration
## Multi-GPU Setup
FastVideo automatically distributes the generation process when multiple GPUs are specified:
```python
# Will use 4 GPUs in parallel for faster generation
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=4,
)
```
## Customizing Generation
- `PipelineConfig`: Initialization time parameters
- `SamplingParam`: Generation time parameters
You can customize various parameters when generating videos using the `PipelineConfig` and `SamplingParam` class:
```python
from fastvideo import VideoGenerator, SamplingParam, PipelineConfig
def main():
model_name = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
config = PipelineConfig.from_pretrained(model_name)
config.vae_precision = "fp16"
config.use_cpu_offload = True
# Create the generator
generator = VideoGenerator.from_pretrained(
model_name,
num_gpus=1,
pipeline_config=config
)
# Create and customize sampling parameters
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
# How many frames to generate
sampling_param.num_frames = 45
# Video resolution (width, height)
sampling_param.width = 1024
sampling_param.height = 576
# How many steps we denoise the video (higher = better quality, slower generation)
sampling_param.num_inference_steps = 30
# How strongly the video conforms to the prompt (higher = more faithful to prompt)
sampling_param.guidance_scale = 7.5
# Random seed for reproducibility
sampling_param.seed = 42 # Optional, leave unset for random results
# Generate video with custom parameters
prompt = "A beautiful sunset over a calm ocean, with gentle waves."
video = generator.generate_video(
prompt,
sampling_param=sampling_param,
output_path="my_videos/", # Controls where videos are saved
return_frames=True, # Also return frames from this call (defaults to False)
save_video=True
)
# If return_frames=True, video contains the generated frames as a NumPy array
print(f"Generated {len(video)} frames")
if __name__ == '__main__':
main()
```
## Performance Optimization
For configuring optimizations, please see our [optimizations guide](#inference-optimizations)
+25 -30
View File
@@ -2,19 +2,11 @@
This page contains step-by-step instructions to get you quickly started with video generation using FastVideo.
## Table of Contents
- [Generating Your First Video](#generating-your-first-video)
- [Customizing Generation](#customizing-generation)
- [Available Models](#available-models)
- [Image-to-Video Generation](#image-to-video-generation)
- [Troubleshooting](#troubleshooting)
- [Advanced Configuration](#advanced-configuration)
- [Next Steps](#next-steps)
## Software Requirements
## Requirements
- **OS**: Linux (Tested on Ubuntu 22.04+)
- **Python**: 3.10-3.12
- **CUDA**: 12.4
- **GPU**: At least one NVIDIA GPU
## Installation
@@ -78,18 +70,25 @@ You can generate a video starting from an initial image:
```python
from fastvideo import VideoGenerator, SamplingParam
# Create the generator
generator = VideoGenerator.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
def main():
# Create the generator
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
# Set up parameters with an initial image
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
sampling_param.image_path = "path/to/your/image.jpg"
sampling_param.num_frames = 24
sampling_param.image_strength = 0.8 # How much to preserve the original image (0-1)
# Set up parameters with an initial image
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
sampling_param.num_frames = 107
sampling_param.image_strength = 0.8 # How much to preserve the original image (0-1)
# Generate video based on the image
prompt = "A photograph coming to life with gentle movement"
video = generator.generate_video(prompt, sampling_param=sampling_param)
# Generate video based on the image
prompt = "A photograph coming to life with gentle movement"
generator.generate_video(prompt, sampling_param=sampling_param,
output_path="my_videos/",
save_video=True)
if __name__ == '__main__':
main()
```
## Troubleshooting
@@ -100,7 +99,7 @@ Common issues and their solutions:
If you encounter CUDA out of memory errors:
- Reduce `num_frames` or video resolution
- Enable memory optimization with `enable_model_cpu_offload`
- Try a smaller model or use quantized versions
- Try a smaller model or use distilled versions
- Use `num_gpus` > 1 if multiple GPUs are available
### Slow Generation
@@ -116,14 +115,10 @@ If the generated video doesn't match your prompt:
- Experiment with different random seeds
- Try a different model
## Advanced Configuration
## Optimizations
## Next Steps
- Explore the [API Reference](../api/index.md) for detailed documentation
- Learn about [Advanced Inference Options](../inference/overview_back.md)
- See [Examples](../examples/index.md) for more usage scenarios
- Check out the [Model Training](../training/overview.md) guide to fine-tune models
- Join our [Community Discord](https://discord.gg/fastvideo) for support and sharing
- Learn about [Advanced Inference Configurations](#inference-configuration)
- Learn about using [Optimizations](#inference-optimizations)
- See [Examples](../examples/examples_inference_index.md) for more usage scenarios
- Join our [Community Discord](https://discord.gg/JA7cksDz86).
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
+1
View File
@@ -1,3 +1,4 @@
(inference-optimizations)=
# Optimizations
This page describes the various options for speeding up generation times in FastVideo.
+8 -8
View File
@@ -49,14 +49,14 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
* ❌
* ✅
* ✅
- * Wan T2V 1.4B
- * Wan T2V 1.3B
* `Wan-AI/Wan2.1-T2V-1.3B-Diffusers`
* 480P
* ✅
* ✅*
* ✅
- * Wan T2V 14B
* `Wan-AI/Wan2.1-T2V-1.3B-Diffusers`
* `Wan-AI/Wan2.1-T2V-14B-Diffusers`
* 480P, 720P
* ✅
* ✅*
@@ -67,18 +67,18 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
* ✅
* ✅*
* ✅
- * Wan T2V 720P
* `Wan-AI/Wan2.1-T2V-14B-Diffusers`
- * Wan I2V 720P
* `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers`
* 720P
* ✅
* ✅*
* ✅
- * StepVideo T2V
* Coming soon!
* `FastVideo/stepvideo-t2v-diffusers`
* 768px768px204f<br>544px992px204f<br>544px992px136f
*
*
*
* ❌
* ❌
* ✅
:::
**Note**: there are some known quality issues with Wan2.1 + Sliding Tile Attn. We are working on fixing this issue.
+122
View File
@@ -0,0 +1,122 @@
import argparse
import json
import os
import torch
import torch.distributed as dist
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import maybe_download_model, shallow_asdict
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo import PipelineConfig
from fastvideo.v1.pipelines.preprocess_pipeline import PreprocessPipeline
logger = init_logger(__name__)
BASE_MODEL_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
def main(args):
# Assume using torchrun
local_rank = int(os.getenv("RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
init_distributed_environment(world_size=world_size, rank=rank, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
pipeline_config = PipelineConfig.from_pretrained(MODEL_PATH)
kwargs = {
"use_cpu_offload": False,
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
}
pipeline_config_args = shallow_asdict(pipeline_config)
pipeline_config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=MODEL_PATH,
num_gpus=world_size,
device_str="cuda",
**pipeline_config_args,
)
fastvideo_args.check_fastvideo_args()
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
pipeline = PreprocessPipeline(MODEL_PATH, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--preprocess_video_batch_size",
type=int,
default=2,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--preprocess_text_batch_size",
type=int,
default=8,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--samples_per_file",
type=int,
default=64
)
parser.add_argument(
"--flush_frequency",
type=int,
default=256,
help="how often to save to parquet files"
)
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--max_height", type=int, default=480)
parser.add_argument("--max_width", type=int, default=848)
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("--dataset", default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
parser.add_argument("--speed_factor", type=float, default=1.0)
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
# text encoder & vae & diffusion model
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
main(args)
@@ -0,0 +1,199 @@
import argparse
import json
import os
import torch
import torch.distributed as dist
from accelerate.logging import get_logger
from diffusers.utils import export_to_video
from diffusers.video_processor import VideoProcessor
from torch.utils.data import DataLoader, Dataset
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
# from fastvideo.utils.load import load_text_encoder, load_vae
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader, TextEncoderLoader, TokenizerLoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.configs.models.encoders.t5 import T5Config
logger = get_logger(__name__)
class T5dataset(Dataset):
def __init__(
self,
json_path,
vae_debug,
):
self.json_path = json_path
self.vae_debug = vae_debug
with open(self.json_path, "r") as f:
train_dataset = json.load(f)
self.train_dataset = sorted(train_dataset, key=lambda x: x["latent_path"])
def __getitem__(self, idx):
caption = self.train_dataset[idx]["caption"]
filename = self.train_dataset[idx]["latent_path"].split(".")[0]
length = self.train_dataset[idx]["length"]
if self.vae_debug:
latents = torch.load(
os.path.join(args.output_dir, "latent", self.train_dataset[idx]["latent_path"]),
map_location="cpu",
)
else:
latents = []
return dict(caption=caption, latents=latents, filename=filename, length=length)
def __len__(self):
return len(self.train_dataset)
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
rank = int(os.getenv("RANK", 0))
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
print("world_size", world_size, "local rank", local_rank)
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
world_group = get_world_group()
# device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# torch.cuda.set_device(local_rank)
# if not dist.is_initialized():
# dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
videoprocessor = VideoProcessor(vae_scale_factor=8)
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "video"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "prompt_embed"), exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "prompt_attention_mask"), exist_ok=True)
vae_precision = "fp16"
text_encoder_precision = "fp32"
fastvideo_args = FastVideoArgs(model_path=args.model_path,
use_cpu_offload=False,
vae_precision=vae_precision,
text_encoder_precisions=(text_encoder_precision,))
fastvideo_args.device = device
fastvideo_args.device_str = f"cuda:{local_rank}"
# fastvideo_args.dit_config = HunyuanVideoConfig()
fastvideo_args.vae_config = WanVAEConfig()
fastvideo_args.text_encoder_configs = (T5Config(),)
# vae_loader = VAELoader()
# vae = vae_loader.load_vae()
text_encoder_loader = TextEncoderLoader()
tokenizer_loader = TokenizerLoader()
model_path = args.model_path
path = maybe_download_model(model_path)
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
ENCODER_PATH = os.path.join(path, "text_encoder")
TOKENIZER_PATH = os.path.join(path, "tokenizer")
print(ENCODER_PATH)
text_encoder = text_encoder_loader.load(ENCODER_PATH, "text_encoder", fastvideo_args)
tokenizer = tokenizer_loader.load(TOKENIZER_PATH, "tokenizer", fastvideo_args)
latents_json_path = os.path.join(args.output_dir, "videos2caption_temp.json")
train_dataset = T5dataset(latents_json_path, args.vae_debug)
# text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
# vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
# vae.enable_tiling()
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
json_data = []
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
# with torch.autocast("cuda", dtype=torch.float32):
print(data["caption"])
text_inputs = tokenizer(data["caption"], **fastvideo_args.text_encoder_configs[0].tokenizer_kwargs).to(
fastvideo_args.device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
outputs = text_encoder(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
)
from fastvideo.v1.configs.pipelines.wan import t5_postprocess_text
post_process_func = t5_postprocess_text
prompt_embeds = post_process_func(outputs)
prompt_attention_mask = attention_mask
if args.vae_debug:
latents = data["latents"]
video = vae.decode(latents.to(device), return_dict=False)[0]
video = videoprocessor.postprocess_video(video)
for idx, video_name in enumerate(data["filename"]):
prompt_embed_path = os.path.join(args.output_dir, "prompt_embed", video_name + ".pt")
video_path = os.path.join(args.output_dir, "video", video_name + ".mp4")
prompt_attention_mask_path = os.path.join(args.output_dir, "prompt_attention_mask",
video_name + ".pt")
# save latent
torch.save(prompt_embeds[idx], prompt_embed_path)
torch.save(prompt_attention_mask[idx], prompt_attention_mask_path)
print(f"sample {video_name} saved")
if args.vae_debug:
export_to_video(video[idx], video_path, fps=16)
item = {}
item["length"] = int(data["length"][idx])
item["latent_path"] = video_name + ".pt"
item["prompt_embed_path"] = video_name + ".pt"
item["prompt_attention_mask"] = video_name + ".pt"
item["caption"] = data["caption"][idx]
json_data.append(item)
dist.barrier()
local_data = json_data
gathered_data = [None] * world_size
dist.all_gather_object(gathered_data, local_data)
if local_rank == 0:
# os.remove(latents_json_path)
all_json_data = [item for sublist in gathered_data for item in sublist]
with open(os.path.join(args.output_dir, "videos2caption.json"), "w") as f:
json.dump(all_json_data, f, indent=4)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
# parser.add_argument("--model_type", type=str, default="mochi")
# text encoder & vae & diffusion model
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=1,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument("--vae_debug", action="store_true")
args = parser.parse_args()
main(args)
@@ -0,0 +1,151 @@
import argparse
import json
import os
import torch
# import torch.distributed as dist
# from accelerate.logging import get_logger
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.dataset import getdataset
# from fastvideo.utils.load import load_vae
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
logger = init_logger(__name__)
model_path = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
path = maybe_download_model(model_path)
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
VAE_PATH = os.path.join(path, "vae")
print(VAE_PATH)
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
rank = int(os.getenv("RANK", 0))
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
print("world_size", world_size, "local rank", local_rank)
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
world_group = get_world_group()
vae_precision = "fp16"
fastvideo_args = FastVideoArgs(model_path=VAE_PATH,
use_cpu_offload=False,
vae_precision=vae_precision)
fastvideo_args.device = device
# fastvideo_args.dit_config = HunyuanVideoConfig()
fastvideo_args.vae_config = WanVAEConfig()
train_dataset = getdataset(args)
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
# encoder_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# torch.cuda.set_device(local_rank)
# if not dist.is_initialized():
# dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
vae_loader = VAELoader()
vae = vae_loader.load(VAE_PATH, "vae", fastvideo_args)
# vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
# vae.enable_tiling()
os.makedirs(args.output_dir, exist_ok=True)
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
json_data = []
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
with torch.autocast("cuda", dtype=torch.float16):
latents = vae.encode(data["pixel_values"].to(device)).sample()
for idx, video_path in enumerate(data["path"]):
video_name = os.path.basename(video_path).split(".")[0]
latent_path = os.path.join(args.output_dir, "latent", video_name + ".pt")
torch.save(latents[idx].to(torch.bfloat16), latent_path)
item = {}
item["length"] = latents[idx].shape[1]
item["latent_path"] = video_name + ".pt"
item["caption"] = data["text"][idx]
json_data.append(item)
print(f"{video_name} processed")
world_group.barrier()
local_data = json_data
gathered_data = [None] * world_size
for i in range(world_size):
if local_rank == i:
world_group.broadcast_object(local_data, src=i)
else:
gathered_data[i] = world_group.broadcast_object(None, src=i)
gathered_data[local_rank] = json_data
print(gathered_data)
if local_rank == 0:
all_json_data = [item for sublist in gathered_data for item in sublist]
with open(os.path.join(args.output_dir, "videos2caption_temp.json"), "w") as f:
json.dump(all_json_data, f, indent=4)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
# parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--max_height", type=int, default=480)
parser.add_argument("--max_width", type=int, default=848)
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("--dataset", default="t2v")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
parser.add_argument("--speed_factor", type=float, default=1.0)
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
# text encoder & vae & diffusion model
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
main(args)
@@ -0,0 +1,115 @@
import argparse
import os
import torch
# import torch.distributed as dist
from accelerate.logging import get_logger
# from fastvideo.utils.load import load_text_encoder
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader, TextEncoderLoader, TokenizerLoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel, get_world_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.configs.models.encoders.t5 import T5Config
logger = get_logger(__name__)
def main(args):
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
rank = int(os.getenv("RANK", 0))
init_distributed_environment(rank=rank, world_size=world_size, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
print("world_size", world_size, "local rank", local_rank)
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
world_group = get_world_group()
vae_precision = "fp16"
text_encoder_precision = "fp32"
fastvideo_args = FastVideoArgs(model_path=args.model_path,
use_cpu_offload=False,
vae_precision=vae_precision,
text_encoder_precisions=(text_encoder_precision,))
fastvideo_args.device = device
fastvideo_args.device_str = f"cuda:{local_rank}"
# fastvideo_args.dit_config = HunyuanVideoConfig()
fastvideo_args.vae_config = WanVAEConfig()
fastvideo_args.text_encoder_configs = (T5Config(),)
# vae_loader = VAELoader()
# vae = vae_loader.load_vae()
text_encoder_loader = TextEncoderLoader()
tokenizer_loader = TokenizerLoader()
model_path = args.model_path
path = maybe_download_model(model_path)
# PIPELINE_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
ENCODER_PATH = os.path.join(path, "text_encoder")
TOKENIZER_PATH = os.path.join(path, "tokenizer")
print(ENCODER_PATH)
text_encoder = text_encoder_loader.load(ENCODER_PATH, "text_encoder", fastvideo_args)
tokenizer = tokenizer_loader.load(TOKENIZER_PATH, "tokenizer", fastvideo_args)
# text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
# autocast_type = torch.float16 if args.model_type == "hunyuan" else torch.bfloat16
# output_dir/validation/prompt_attention_mask
# output_dir/validation/prompt_embed
os.makedirs(os.path.join(args.output_dir, "validation"), exist_ok=True)
os.makedirs(
os.path.join(args.output_dir, "validation", "prompt_attention_mask"),
exist_ok=True,
)
os.makedirs(os.path.join(args.output_dir, "validation", "prompt_embed"), exist_ok=True)
with open(args.validation_prompt_txt, "r", encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
for prompt in prompts:
with torch.inference_mode():
# with torch.autocast("cuda", dtype=autocast_type):
text_inputs = tokenizer(prompt, **fastvideo_args.text_encoder_configs[0].tokenizer_kwargs).to(
fastvideo_args.device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
outputs = text_encoder(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
)
from fastvideo.v1.configs.pipelines.wan import t5_postprocess_text
post_process_func = t5_postprocess_text
prompt_embeds = post_process_func(outputs)
prompt_attention_mask = attention_mask
file_name = prompt.split(".")[0]
prompt_embed_path = os.path.join(args.output_dir, "validation", "prompt_embed", f"{file_name}.pt")
prompt_attention_mask_path = os.path.join(
args.output_dir,
"validation",
"prompt_attention_mask",
f"{file_name}.pt",
)
torch.save(prompt_embeds[0], prompt_embed_path)
torch.save(prompt_attention_mask[0], prompt_attention_mask_path)
print(f"sample {file_name} saved")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
args = parser.parse_args()
main(args)
+2 -1
View File
@@ -7,7 +7,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
# from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -38,6 +38,7 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
linear_range=0.5,
):
if linear_quadratic:
raise NotImplementedError("Linear quadratic schedule is not implemented")
linear_steps = int(num_train_timesteps * linear_range)
sigmas = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
+870
View File
@@ -0,0 +1,870 @@
# !/bin/python3
# isort: skip_file
import argparse
import math
import os
import time
from collections import deque
import torch
import torch.distributed as dist
import wandb
from accelerate.utils import set_seed
from diffusers.optimization import get_scheduler
from diffusers.utils import check_min_version
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import ShardingStrategy
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset,
latent_collate_function)
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.utils.checkpoint import (save_checkpoint, save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast,
sp_parallel_dataloader_wrapper)
from fastvideo.utils.dataset_utils import LengthGroupedSampler
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group,
get_sequence_parallel_state)
from fastvideo.utils.validation import log_validation
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.models.loader.component_loader import TransformerLoader, SchedulerLoader
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.31.0")
logger = init_logger(__name__)
BASE_MODEL_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
SCHEDULER_PATH = os.path.join(MODEL_PATH, "scheduler")
def reshard_fsdp(model):
for m in FSDP.fsdp_modules(model):
if m._has_params and m.sharding_strategy is not ShardingStrategy.NO_SHARD:
torch.distributed.fsdp._runtime_utils._reshard(m, m._handle, True)
def get_norm(model_pred, norms, gradient_accumulation_steps):
fro_norm = (
torch.linalg.matrix_norm(model_pred, ord="fro") / # codespell:ignore
gradient_accumulation_steps)
largest_singular_value = (torch.linalg.matrix_norm(model_pred, ord=2) /
gradient_accumulation_steps)
absolute_mean = torch.mean(
torch.abs(model_pred)) / gradient_accumulation_steps
absolute_max = torch.max(
torch.abs(model_pred)) / gradient_accumulation_steps
dist.all_reduce(fro_norm, op=dist.ReduceOp.AVG)
dist.all_reduce(largest_singular_value, op=dist.ReduceOp.AVG)
dist.all_reduce(absolute_mean, op=dist.ReduceOp.AVG)
norms["fro"] += torch.mean(fro_norm).item() # codespell:ignore
norms["largest singular value"] += torch.mean(largest_singular_value).item()
norms["absolute mean"] += absolute_mean.item()
norms["absolute max"] += absolute_max.item()
def distill_one_step(
transformer,
model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
num_euler_timesteps,
multiphase,
not_apply_cfg_solver,
distill_cfg,
ema_decay,
pred_decay_weight,
pred_decay_type,
hunyuan_teacher_disable_cfg,
):
total_loss = 0.0
optimizer.zero_grad()
model_pred_norm = {
"fro": 0.0, # codespell:ignore
"largest singular value": 0.0,
"absolute mean": 0.0,
"absolute max": 0.0,
}
for _ in range(gradient_accumulation_steps):
(
latents,
encoder_hidden_states,
latents_attention_mask,
encoder_attention_mask,
) = next(loader)
# model_input = normalize_dit_input(model_type, latents)
model_input = latents
noise = torch.randn_like(model_input)
bsz = model_input.shape[0]
index = torch.randint(0,
num_euler_timesteps, (bsz, ),
device=model_input.device).long()
if sp_size > 1:
broadcast(index)
# Add noise according to flow matching.
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index,
model_input.shape)
timesteps = (sigmas *
noise_scheduler.config.num_train_timesteps).view(-1)
# if squeeze to [], unsqueeze to [1]
timesteps_prev = (sigmas_prev *
noise_scheduler.config.num_train_timesteps).view(-1)
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
noisy_model_input = noisy_model_input.to(torch.bfloat16)
forward_batch = ForwardBatch(data_type="video", enable_teacache=False)
# Predict the noise residual
with torch.autocast("cuda", dtype=torch.bfloat16):
teacher_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if hunyuan_teacher_disable_cfg:
teacher_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
with torch.autograd.graph.save_on_cpu(pin_memory=True):
model_pred = transformer(**teacher_kwargs)
# if accelerator.is_main_process:
model_pred, end_index = solver.euler_style_multiphase_pred(
noisy_model_input, model_pred, index, multiphase)
with torch.no_grad():
w = distill_cfg
with torch.autocast("cuda", dtype=torch.bfloat16):
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
cond_teacher_output = teacher_transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
).float()
if not_apply_cfg_solver:
uncond_teacher_output = cond_teacher_output
else:
# Get teacher model prediction on noisy_latents and unconditional embedding
with torch.autocast("cuda", dtype=torch.bfloat16):
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
uncond_teacher_output = teacher_transformer(
noisy_model_input,
uncond_prompt_embed.unsqueeze(0).expand(
bsz, -1, -1),
timesteps,
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
return_dict=False,
).float()
teacher_output = uncond_teacher_output + w * (cond_teacher_output -
uncond_teacher_output)
x_prev = solver.euler_step(noisy_model_input, teacher_output,
index).to(torch.bfloat16)
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16):
if ema_transformer is not None:
target_pred = ema_transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
else:
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch):
with torch.autograd.graph.save_on_cpu(pin_memory=True):
target_pred = transformer(
x_prev,
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)
target, end_index = solver.euler_style_multiphase_pred(
x_prev, target_pred, index, multiphase, True)
huber_c = 0.001
# loss = loss.mean()
loss = (torch.mean(
torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) -
huber_c) / gradient_accumulation_steps)
if pred_decay_weight > 0:
if pred_decay_type == "l1":
pred_decay_loss = (
torch.mean(torch.sqrt(model_pred.float()**2)) *
pred_decay_weight / gradient_accumulation_steps)
loss += pred_decay_loss
elif pred_decay_type == "l2":
# essnetially k2?
pred_decay_loss = (torch.mean(model_pred.float()**2) *
pred_decay_weight /
gradient_accumulation_steps)
loss += pred_decay_loss
else:
assert NotImplementedError("pred_decay_type is not implemented")
# calculate model_pred norm and mean
get_norm(model_pred.detach().float(), model_pred_norm,
gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
dist.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
total_loss += avg_loss.item()
# update ema
if ema_transformer is not None:
reshard_fsdp(ema_transformer)
for p_averaged, p_model in zip(ema_transformer.parameters(),
transformer.parameters()):
with torch.no_grad():
p_averaged.copy_(
torch.lerp(p_averaged.detach(), p_model.detach(),
1 - ema_decay))
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
grad_norm = torch.nn.utils.clip_grad_norm_(transformer.parameters(),
max_norm=max_grad_norm)
optimizer.step()
lr_scheduler.step()
return total_loss, grad_norm.item(), model_pred_norm
def main(args):
torch.backends.cuda.matmul.allow_tf32 = True
local_rank = int(os.environ.get("LOCAL_RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
torch.cuda.set_device(rank)
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=args.sp_size,
sequence_model_parallel_size=args.sp_size)
fastvideo_args = FastVideoArgs(
model_path=MODEL_PATH,
num_gpus=world_size,
use_cpu_offload=False,
precision=args.master_weight_type,
dit_config=WanVideoConfig(),
device_str="cuda",
)
fastvideo_args.check_fastvideo_args()
device_str = f"cuda:{rank}"
device = torch.device(device_str)
fastvideo_args.device = device
# If passed along, set the training seed now. On GPU...
if args.seed is not None:
# TODO: t within the same seq parallel group should be the same. Noise should be different.
set_seed(args.seed + rank)
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weights to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
# Create model:
logger.info("--> loading model from %s", TRANSFORMER_PATH)
fastvideo_args.device = device
transformer_loader = TransformerLoader()
transformer = transformer_loader.load(TRANSFORMER_PATH, "", fastvideo_args)
transformer = transformer.train()
transformer.requires_grad_(True)
teacher_loader = TransformerLoader()
teacher_transformer = teacher_loader.load(TRANSFORMER_PATH, "",
fastvideo_args)
if args.use_ema:
ema_transformer = teacher_loader.load(TRANSFORMER_PATH, "",
fastvideo_args)
else:
ema_transformer = None
logger.info(
" Total training parameters = %s M",
sum(p.numel()
for p in transformer.parameters() if p.requires_grad) / 1e6)
logger.info("--> model loaded")
teacher_transformer.requires_grad_(False)
if args.use_ema:
ema_transformer.requires_grad_(False)
# scheduler
noise_scheduler_loader = SchedulerLoader()
noise_scheduler = noise_scheduler_loader.load(SCHEDULER_PATH, "",
fastvideo_args)
solver = EulerSolver(
noise_scheduler.sigmas.numpy()[::-1],
noise_scheduler.config.num_train_timesteps,
euler_timesteps=args.num_euler_timesteps,
)
solver.to(device)
params_to_optimize = transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
init_steps = 0
logger.info("optimizer: %s", optimizer)
# todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * world_size,
num_training_steps=args.max_train_steps * world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
args.cfg)
uncond_prompt_embed = train_dataset.uncond_prompt_embed
uncond_prompt_mask = train_dataset.uncond_prompt_mask
sampler = (LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
) if (args.group_frame or args.group_resolution) else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False))
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
collate_fn=latent_collate_function,
pin_memory=True,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
drop_last=True,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps *
args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps /
num_update_steps_per_epoch)
if rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
# Train!
total_batch_size = (world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
logger.info("***** Running training *****")
logger.info(" Num examples = %s", len(train_dataset))
logger.info(" Dataloader size = %s", len(train_dataloader))
logger.info(" Num Epochs = %s", args.num_train_epochs)
logger.info(" Resume training from step %s", init_steps)
logger.info(" Instantaneous batch size per device = %s",
args.train_batch_size)
logger.info(
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
total_batch_size)
logger.info(" Gradient Accumulation steps = %s",
args.gradient_accumulation_steps)
logger.info(" Total optimization steps = %s", args.max_train_steps)
logger.info(
" Total training parameters per FSDP shard = %s B",
sum(p.numel()
for p in transformer.parameters() if p.requires_grad) / 1e9)
# print dtype
logger.info(" Master weight dtype: %s",
transformer.parameters().__next__().dtype)
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader)
# log_validation(args, transformer, device,
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
def get_num_phases(multi_phased_distill_schedule, step):
# step-phase,step-phase
multi_phases = multi_phased_distill_schedule.split(",")
phase = multi_phases[-1].split("-")[-1]
for step_phases in multi_phases:
phase_step, phase = step_phases.split("-")
if step <= int(phase_step):
return int(phase)
return phase
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
assert args.multi_phased_distill_schedule is not None
num_phases = get_num_phases(args.multi_phased_distill_schedule, step)
loss, grad_norm, pred_norm = distill_one_step(
transformer,
args.model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
args.num_euler_timesteps,
num_phases,
args.not_apply_cfg_solver,
args.distill_cfg,
args.ema_decay,
args.pred_decay_weight,
args.pred_decay_type,
args.hunyuan_teacher_disable_cfg,
)
step_time = time.time() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"phases": num_phases,
})
progress_bar.update(1)
if rank <= 0:
wandb.log(
{
"train_loss":
loss,
"learning_rate":
lr_scheduler.get_last_lr()[0],
"step_time":
step_time,
"avg_step_time":
avg_step_time,
"grad_norm":
grad_norm,
"pred_fro_norm":
pred_norm["fro"], # codespell:ignore
"pred_largest_singular_value":
pred_norm["largest singular value"],
"pred_absolute_mean":
pred_norm["absolute mean"],
"pred_absolute_max":
pred_norm["absolute max"],
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(transformer, optimizer, rank,
args.output_dir, step)
else:
# Your existing checkpoint saving code
if args.use_ema:
save_checkpoint(ema_transformer, rank, args.output_dir,
step)
else:
save_checkpoint(transformer, rank, args.output_dir, step)
dist.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(
args,
transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=False,
)
if args.use_ema:
log_validation(
args,
ema_transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=True,
)
if args.use_lora:
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
args.max_train_steps)
else:
save_checkpoint(transformer, rank, args.output_dir,
args.max_train_steps)
if get_sequence_parallel_state():
destroy_sequence_parallel_group()
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model_type",
type=str,
default="mochi",
help="The type of model to train.")
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_height", type=int, default=480)
parser.add_argument("--num_width", type=int, default=848)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=10,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t",
type=int,
default=28,
help="Number of latent timesteps.")
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
# text encoder & vae & diffusion model
parser.add_argument("--pretrained_model_name_or_path", type=str)
parser.add_argument("--dit_model_name_or_path", type=str)
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
# diffusion setting
parser.add_argument("--ema_decay", type=float, default=0.95)
parser.add_argument("--ema_start_step", type=int, default=0)
parser.add_argument("--cfg", type=float, default=0.1)
# validation & logs
parser.add_argument("--validation_prompt_dir", type=str)
parser.add_argument("--validation_sampling_steps", type=str, default="64")
parser.add_argument("--validation_guidance_scale", type=str, default="4.5")
parser.add_argument("--validation_steps", type=float, default=64)
parser.add_argument("--log_validation", action="store_true")
parser.add_argument("--tracker_project_name", type=str, default=None)
parser.add_argument("--seed",
type=int,
default=None,
help="A seed for reproducible training.")
parser.add_argument(
"--output_dir",
type=str,
default=None,
help=
"The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--checkpoints_total_limit",
type=int,
default=None,
help=("Max number of checkpoints to store."),
)
parser.add_argument(
"--checkpointing_steps",
type=int,
default=500,
help=
("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."),
)
parser.add_argument("--shift", type=float, default=1.0)
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--resume_from_lora_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=
("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
parser.add_argument(
"--max_train_steps",
type=int,
default=None,
help=
"Total number of training steps to perform. If provided, overrides num_train_epochs.",
)
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
default=1,
help=
"Number of updates steps to accumulate before performing a backward/update pass.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-4,
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help=
"Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_warmup_steps",
type=int,
default=10,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument("--max_grad_norm",
default=1.0,
type=float,
help="Max gradient norm.")
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help=
"Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=
("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=
("Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
)
parser.add_argument(
"--use_cpu_offload",
action="store_true",
help=
"Whether to use CPU offload for param & gradient & optimizer states.",
)
parser.add_argument("--sp_size",
type=int,
default=1,
help="For sequence parallel")
parser.add_argument(
"--train_sp_batch_size",
type=int,
default=1,
help="Batch size for sequence parallel training",
)
parser.add_argument(
"--use_lora",
action="store_true",
default=False,
help="Whether to use LoRA for finetuning.",
)
parser.add_argument("--lora_alpha",
type=int,
default=256,
help="Alpha parameter for LoRA.")
parser.add_argument("--lora_rank",
type=int,
default=128,
help="LoRA rank parameter. ")
parser.add_argument("--fsdp_sharding_startegy", default="full")
# lr_scheduler
parser.add_argument(
"--lr_scheduler",
type=str,
default="constant",
help=
('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'),
)
parser.add_argument("--num_euler_timesteps", type=int, default=100)
parser.add_argument(
"--lr_num_cycles",
type=int,
default=1,
help="Number of cycles in the learning rate scheduler.",
)
parser.add_argument(
"--lr_power",
type=float,
default=1.0,
help="Power factor of the polynomial scheduler.",
)
parser.add_argument(
"--not_apply_cfg_solver",
action="store_true",
help="Whether to apply the cfg_solver.",
)
parser.add_argument("--distill_cfg",
type=float,
default=3.0,
help="Distillation coefficient.")
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
parser.add_argument("--scheduler_type",
type=str,
default="pcm",
help="The scheduler type to use.")
parser.add_argument(
"--linear_quadratic_threshold",
type=float,
default=0.025,
help="Threshold for linear quadratic scheduler.",
)
parser.add_argument(
"--linear_range",
type=float,
default=0.5,
help="Range for linear quadratic scheduler.",
)
parser.add_argument("--weight_decay",
type=float,
default=0.001,
help="Weight decay to apply.")
parser.add_argument("--use_ema",
action="store_true",
help="Whether to use EMA.")
parser.add_argument("--multi_phased_distill_schedule",
type=str,
default=None)
parser.add_argument("--pred_decay_weight", type=float, default=0.0)
parser.add_argument("--pred_decay_type", default="l1")
parser.add_argument("--hunyuan_teacher_disable_cfg", action="store_true")
parser.add_argument(
"--master_weight_type",
type=str,
default="fp32",
help="Weight type to use - fp32 or bf16.",
)
args = parser.parse_args()
main(args)
@@ -31,7 +31,7 @@ mochi_latents_std = torch.tensor([
mochi_scaling_factor = 1.0
def normalize_dit_input(model_type, latents):
def normalize_dit_input(model_type, latents, args=None):
if model_type == "mochi":
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
@@ -41,5 +41,16 @@ def normalize_dit_input(model_type, latents):
return latents * 0.476986
elif model_type == "hunyuan":
return latents * 0.476986
elif model_type == "wan":
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
vae_config = WanVAEConfig()
latents_mean = torch.tensor(vae_config.arch_config.latents_mean)
latents_std = 1.0 / torch.tensor(vae_config.arch_config.latents_std)
latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(device=latents.device)
latents_std = latents_std.view(1, -1, 1, 1, 1).to(device=latents.device)
latents = ((latents.float() - latents_mean) * latents_std).to(latents)
return latents
else:
raise NotImplementedError(f"model_type {model_type} not supported")
+39 -1
View File
@@ -11,6 +11,7 @@ from torch.distributed.checkpoint.optimizer import load_sharded_optimizer_state_
from torch.distributed.fsdp import FullOptimStateDictConfig, FullStateDictConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import StateDictType
import dataclasses
from fastvideo.utils.logging_ import main_print
@@ -44,13 +45,50 @@ def save_checkpoint_optimizer(model, optimizer, rank, output_dir, step, discrimi
optimizer_path = os.path.join(save_dir, "optimizer.pt")
torch.save(optim_state, optimizer_path)
else:
weight_path = os.path.join(save_dir, "discriminator_pytorch_model.safetensors")
weight_path = os.path.join(save_dstate_dictir, "discriminator_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
optimizer_path = os.path.join(save_dir, "discriminator_optimizer.pt")
torch.save(optim_state, optimizer_path)
main_print(f"--> checkpoint saved at step {step}")
def save_checkpoint_v1(transformer, rank, output_dir, step):
# from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
# from torch.distributed.fsdp import StateDictType, FullStateDictConfig
# Configure FSDP to save full state dict
FSDP.set_state_dict_type(
transformer,
state_dict_type=StateDictType.FULL_STATE_DICT,
state_dict_config=FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
)
# Now get the state dict
cpu_state = transformer.state_dict()
# Save it (only on rank 0 since we used rank0_only=True)
# if torch.distributed.get_rank() == 0:
# torch.save(state_dict, "model_checkpoint.pt")
if rank <= 0:
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
# weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.pt")
print(weight_path)
# save_file(cpu_state, weight_path)
torch.save(cpu_state, weight_path)
config_dict = transformer.hf_config
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
config_path = os.path.join(save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
main_print(f"--> checkpoint saved at step {step}")
def save_checkpoint(transformer, rank, output_dir, step):
main_print(f"--> saving checkpoint at step {step}")
with FSDP.state_dict_type(
-3
View File
@@ -84,9 +84,6 @@ class DistributedAttention(nn.Module):
# Check input shapes
assert q.dim() == 4 and k.dim() == 4 and v.dim(
) == 4, "Expected 4D tensors"
# assert bs = 1
assert q.shape[
0] == 1, "Batch size must be 1, and there should be no padding tokens"
batch_size, seq_len, num_heads, head_dim = q.shape
local_rank = get_sequence_model_parallel_rank()
world_size = get_sequence_model_parallel_world_size()
+1 -1
View File
@@ -2,7 +2,7 @@ from dataclasses import dataclass, field
from typing import Any, Optional, Tuple
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
+1 -1
View File
@@ -4,7 +4,7 @@ from typing import Any, Dict, List, Optional, Tuple
import torch
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
@@ -68,6 +68,9 @@ class HunyuanConfig(PipelineConfig):
embedded_cfg_scale: int = 6
flow_shift: int = 7
# Video parameters
use_cpu_offload: bool = True
# Text encoding stage
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LlamaConfig(), CLIPTextConfig()))
@@ -18,6 +18,9 @@ class StepVideoT2VConfig(PipelineConfig):
vae_tiling: bool = False
vae_sp: bool = False
# Video parameters
use_cpu_offload: bool = True
# Denoising stage
flow_shift: int = 13
timesteps_scale: bool = False
@@ -1,3 +0,0 @@
from fastvideo.v1.configs.quantization.base import QuantizationConfig
__all__ = ["QuantizationConfig"]
@@ -1,6 +0,0 @@
from dataclasses import dataclass
@dataclass
class QuantizationConfig:
pass
@@ -2,11 +2,12 @@ from torchvision import transforms
from torchvision.transforms import Lambda
from transformers import AutoTokenizer
from fastvideo.dataset.t2v_datasets import T2V_dataset
from fastvideo.dataset.transform import CenterCropResizeVideo, Normalize255, TemporalRandomCrop
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
def getdataset(args):
def getdataset(args, start_idx=0):
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
resize_topcrop = [
@@ -25,15 +26,15 @@ def getdataset(args):
norm_fun,
])
# tokenizer = AutoTokenizer.from_pretrained("/storage/ongoing/new/Open-Sora-Plan/cache_dir/mt5-xxl", cache_dir=args.cache_dir)
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name, cache_dir=args.cache_dir)
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name,
cache_dir=args.cache_dir)
if args.dataset == "t2v":
return T2V_dataset(
args,
transform=transform,
temporal_sample=temporal_sample,
tokenizer=tokenizer,
transform_topcrop=transform_topcrop,
)
return T2V_dataset(args,
transform=transform,
temporal_sample=temporal_sample,
tokenizer=tokenizer,
transform_topcrop=transform_topcrop,
start_idx=start_idx)
raise NotImplementedError(args.dataset)
@@ -44,7 +45,7 @@ if __name__ == "__main__":
from accelerate import Accelerator
from tqdm import tqdm
from fastvideo.dataset.t2v_datasets import dataset_prog
from fastvideo.v1.dataset.t2v_datasets import dataset_prog
args = type(
"args",
@@ -63,7 +64,8 @@ if __name__ == "__main__":
"interpolation_scale_h": 1,
"interpolation_scale_w": 1,
"cache_dir": "../cache_dir",
"image_data": "/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
"image_data":
"/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
"video_data": "1",
"train_fps": 24,
"drop_short_ratio": 1.0,
@@ -80,7 +82,10 @@ if __name__ == "__main__":
zero = 0
for idx in tqdm(range(num)):
image_data = dataset_prog.img_cap_list[idx]
caps = [i["cap"] if isinstance(i["cap"], list) else [i["cap"]] for i in image_data]
caps = [
i["cap"] if isinstance(i["cap"], list) else [i["cap"]]
for i in image_data
]
try:
caps = [[random.choice(i)] for i in caps]
except Exception as e:
+44
View File
@@ -0,0 +1,44 @@
# schema.py
"""
Unified data schema and format for saving and loading image/video data after
preprocessing.
It uses apache arrow in-memory format that can be consumed by modern data
frameworks that can handle parquet or lance file.
"""
import pyarrow as pa
pyarrow_schema = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
@@ -20,9 +20,11 @@ class LatentDataset(Dataset):
self.datase_dir_path = os.path.dirname(json_path)
self.video_dir = os.path.join(self.datase_dir_path, "video")
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
self.prompt_embed_dir = os.path.join(self.datase_dir_path, "prompt_embed")
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path, "prompt_attention_mask")
with open(self.json_path, "r") as f:
self.prompt_embed_dir = os.path.join(self.datase_dir_path,
"prompt_embed")
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path,
"prompt_attention_mask")
with open(self.json_path) as f:
self.data_anno = json.load(f)
# json.load(f) already keeps the order
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
@@ -31,12 +33,16 @@ class LatentDataset(Dataset):
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
# 256 zeros
self.uncond_prompt_mask = torch.zeros(256).bool()
self.lengths = [data_item["length"] if "length" in data_item else 1 for data_item in self.data_anno]
self.lengths = [
data_item["length"] if "length" in data_item else 1
for data_item in self.data_anno
]
def __getitem__(self, idx):
latent_file = self.data_anno[idx]["latent_path"]
prompt_embed_file = self.data_anno[idx]["prompt_embed_path"]
prompt_attention_mask_file = self.data_anno[idx]["prompt_attention_mask"]
prompt_attention_mask_file = self.data_anno[idx][
"prompt_attention_mask"]
# load
latent = torch.load(
os.path.join(self.latent_dir, latent_file),
@@ -54,7 +60,8 @@ class LatentDataset(Dataset):
weights_only=True,
)
prompt_attention_mask = torch.load(
os.path.join(self.prompt_attention_mask_dir, prompt_attention_mask_file),
os.path.join(self.prompt_attention_mask_dir,
prompt_attention_mask_file),
map_location="cpu",
weights_only=True,
)
@@ -104,8 +111,12 @@ def latent_collate_function(batch):
if __name__ == "__main__":
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt", num_latent_t=28)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=2, shuffle=False, collate_fn=latent_collate_function)
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt",
num_latent_t=28)
dataloader = torch.utils.data.DataLoader(dataset,
batch_size=2,
shuffle=False,
collate_fn=latent_collate_function)
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
print(
latent.shape,
+418
View File
@@ -0,0 +1,418 @@
import argparse
import os
import random
import time
import numpy as np
import pyarrow.parquet as pq
import torch
import tqdm
from torch import distributed as dist
from torch.utils.data import IterableDataset, get_worker_info
from torchdata.stateful_dataloader import StatefulDataLoader
# Path to your dataset
dataset_path = "/mnt/sharefs/users/hao.zhang/Vchitect-2M/Vchitect-2M-laten-93x512x512/train/"
class ParquetVideoTextDataset(IterableDataset):
"""Efficient loader for video-text data from a directory of Parquet files."""
def __init__(self,
path: str,
batch_size: int = 1024,
rank: int = 0,
world_size: int = 1,
cfg_rate: float = 0.0,
num_latent_t: int = 2):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.rank = rank
self.world_size = world_size
self.cfg_rate = cfg_rate
self.num_latent_t = num_latent_t
# Find all parquet files recursively
print(f"Scanning for parquet files in {self.path}")
self.parquet_files = []
for root, _, files in os.walk(self.path):
for file in files:
if file.endswith('.parquet'):
self.parquet_files.append(os.path.join(root, file))
# Sort files for consistent ordering
self.parquet_files.sort()
# Distribute files among workers
# drop last unenven files
print(f"Total files: {len(self.parquet_files)}")
total_files = len(self.parquet_files)
base_count = total_files // world_size
extra_files = total_files % world_size
if rank < extra_files:
start_idx = rank * (base_count + 1)
end_idx = start_idx + base_count + 1
else:
start_idx = rank * base_count + extra_files
end_idx = start_idx + base_count
self.parquet_files = self.parquet_files[start_idx:end_idx]
print(f"Files assigned to rank {rank}: {len(self.parquet_files)}")
if len(self.parquet_files) > 0:
print(f"First file: {self.parquet_files[0]}")
print(f"Last file: {self.parquet_files[-1]}")
# Initialize current file index
self.current_file_idx = 0
self.current_reader = None
self.current_batches = None
self.total_samples = 0
def _open_next_file(self):
"""Open the next parquet file for reading."""
num_workers = get_worker_info().num_workers
worker_id = get_worker_info().id
total_files = len(self.parquet_files)
base_count = total_files // num_workers
extra_files = total_files % num_workers
if worker_id < extra_files:
start_idx = worker_id * (base_count + 1)
end_idx = start_idx + base_count + 1
else:
start_idx = worker_id * base_count + extra_files
end_idx = start_idx + base_count
worker_parquet_files = self.parquet_files[start_idx:end_idx]
if self.current_file_idx >= len(worker_parquet_files):
print(
f"Rank {self.rank}, Worker {worker_id}: No more files to open (current_idx={self.current_file_idx}, total_files={len(worker_parquet_files)})"
)
return False
if self.current_reader is not None:
self.current_reader.close()
file_path = worker_parquet_files[self.current_file_idx]
print(
f"Rank {self.rank}, Worker {worker_id}: Opening file {self.current_file_idx + 1}/{len(worker_parquet_files)}: {file_path}"
)
try:
self.current_reader = pq.ParquetFile(file_path)
self.current_batches = self.current_reader.iter_batches(
batch_size=self.batch_size)
self.current_file_idx += 1
return True
except Exception as e:
print(f"Error opening file {file_path}: {str(e)}")
return False
def __iter__(self):
"""Iterate over the dataset in a streaming fashion."""
print(f"Rank {self.rank}: Starting iteration")
# First try to open a file
if not self._open_next_file():
print(f"Rank {self.rank}: Failed to open first file")
return
while True:
try:
# Get next batch from current file
batch = next(self.current_batches)
batch_dict = batch.to_pydict()
processed = self._process_batch(batch_dict)
# Update sample count
batch_size = len(processed["latents"])
self.total_samples += batch_size
# Print progress
if self.total_samples % 1000 == 0:
print(
f"Rank {self.rank}: Processed {self.total_samples} samples"
)
# Yield each item in the batch
for lat, emb, mask, info in zip(processed["latents"],
processed["embeddings"],
processed["masks"],
processed["info"]):
if lat.numel() == 0: # Split is validation
yield lat, emb, mask, info
else:
yield lat[:, -self.num_latent_t:], emb, mask, info
except StopIteration:
# Current file is exhausted, try next file
print(
f"Rank {self.rank}: Current file exhausted, trying next file"
)
self.current_batches = None
if not self._open_next_file():
print(
f"Rank {self.rank}: No more files to process. Total samples: {self.total_samples}"
)
break
except Exception as e:
print(f"Error processing batch: {str(e)}")
self.current_batches = None
if not self._open_next_file():
print(
f"Rank {self.rank}: Failed to open next file after error"
)
break
# Clean up
if self.current_reader is not None:
self.current_reader.close()
def _process_batch(self, batch):
"""Process a PyArrow batch into tensors."""
out = {"lat": [], "emb": [], "msk": [], "info": []}
for i in range(len(batch["vae_latent_bytes"])):
vae_latent_bytes = batch["vae_latent_bytes"][i]
vae_latent_shape = batch["vae_latent_shape"][i]
text_embedding_bytes = batch["text_embedding_bytes"][i]
text_embedding_shape = batch["text_embedding_shape"][i]
text_attention_mask_bytes = batch["text_attention_mask_bytes"][i]
text_attention_mask_shape = batch["text_attention_mask_shape"][i]
# Process latent
if not vae_latent_shape: # No VAE latent is stored. Split is validation
lat = np.array([])
else:
lat = np.frombuffer(vae_latent_bytes,
dtype=np.float32).reshape(vae_latent_shape)
# Make array writable
lat = np.copy(lat)
if random.random() < self.cfg_rate:
emb = np.zeros((512, 4096), dtype=np.float32)
else:
emb = np.frombuffer(
text_embedding_bytes,
dtype=np.float32).reshape(text_embedding_shape)
# Make array writable
emb = np.copy(emb)
if emb.shape[0] < 512:
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
padded_emb[:emb.shape[0], :] = emb
emb = padded_emb
elif emb.shape[0] > 512:
emb = emb[:512, :]
# Process mask
if len(text_attention_mask_bytes) > 0 and len(
text_attention_mask_shape) > 0:
msk = np.frombuffer(text_attention_mask_bytes,
dtype=np.uint8).astype(np.bool_)
msk = msk.reshape(1, -1)
# Make array writable
msk = np.copy(msk)
if msk.shape[1] < 512:
padded_msk = np.zeros((1, 512), dtype=np.bool_)
padded_msk[:, :msk.shape[1]] = msk
msk = padded_msk
elif msk.shape[1] > 512:
msk = msk[:, :512]
else:
msk = np.ones((1, 512), dtype=np.bool_)
# to string
file_name = str(batch["file_name"][i])
# Collect metadata
info = {
"width": batch["width"][i],
"height": batch["height"][i],
"num_frames": batch["num_frames"][i],
"duration_sec": batch["duration_sec"][i],
"fps": batch["fps"][i],
"file_name": batch["file_name"][i],
"caption": batch["caption"][i],
}
out["lat"].append(torch.from_numpy(lat))
out["emb"].append(torch.from_numpy(emb))
out["msk"].append(torch.from_numpy(msk))
out["info"].append(info)
return {
"latents": torch.stack(out["lat"]) if out["lat"] else None,
"embeddings": torch.stack(out["emb"]) if out["emb"] else None,
"masks": torch.stack(out["msk"]) if out["msk"] else None,
"info": out["info"]
}
def bind_cpu_cores(local_rank, cpu_per_process=16):
"""根据local_rank绑定固定cpu核。"""
start = local_rank * cpu_per_process
end = start + cpu_per_process
cores = list(range(start, end))
print(f"[Rank {local_rank}] Binding to CPU cores: {cores}")
os.sched_setaffinity(0, cores)
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description='Benchmark Parquet dataset loading speed')
parser.add_argument('--path',
type=str,
default=dataset_path,
help='Path to Parquet dataset')
parser.add_argument('--batch_size',
type=int,
default=4,
help='Batch size for DataLoader')
parser.add_argument('--num_batches',
type=int,
default=100,
help='Number of batches to benchmark')
parser.add_argument('--vae_debug', action="store_true")
args = parser.parse_args()
# Initialize distributed training
local_rank = int(os.environ.get("LOCAL_RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
rank = int(os.environ.get("RANK", 0))
# Initialize CUDA device first
if torch.cuda.is_available():
torch.cuda.set_device(local_rank)
device = torch.device(f"cuda:{local_rank}")
else:
device = torch.device("cpu")
# Initialize distributed training
if world_size > 1:
dist.init_process_group(backend="nccl",
init_method="env://",
world_size=world_size,
rank=rank)
print(
f"Initialized process: rank={rank}, local_rank={local_rank}, world_size={world_size}, device={device}"
)
# Bind CPU cores after distributed initialization
# bind_cpu_cores(local_rank, cpu_per_process=16)
# Create dataset
dataset = ParquetVideoTextDataset(
args.path,
batch_size=args.batch_size,
rank=rank,
world_size=world_size,
)
# Create DataLoader with proper settings
dataloader = StatefulDataLoader(
dataset,
batch_size=args.batch_size,
num_workers=1, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=True)
# Example of how to load dataloader state
# if os.path.exists("/workspace/FastVideo/dataloader_state.pt"):
# dataloader_state = torch.load("/workspace/FastVideo/dataloader_state.pt")
# dataloader.load_state_dict(dataloader_state[rank])
# Warm-up with synchronization
if rank == 0:
print("Warming up...")
for i, (latents, embeddings, masks, infos) in enumerate(dataloader):
# Example of how to save dataloader state
# if i == 30:
# dist.barrier()
# local_data = {rank: dataloader.state_dict()}
# gathered_data = [None] * world_size
# dist.all_gather_object(gathered_data, local_data)
# if rank == 0:
# global_state_dict = {}
# for d in gathered_data:
# global_state_dict.update(d)
# torch.save(global_state_dict, "dataloader_state.pt")
assert torch.sum(masks[0]).item() == torch.count_nonzero(
embeddings[0]).item() // 4096
if args.vae_debug:
from diffusers.utils import export_to_video
from diffusers.video_processor import VideoProcessor
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.models.loader.component_loader import VAELoader
VAE_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/vae"
fastvideo_args = FastVideoArgs(
model_path=VAE_PATH,
vae_config=WanVAEConfig(load_encoder=False),
vae_precision="fp32")
fastvideo_args.device = device
vae_loader = VAELoader()
vae = vae_loader.load(model_path=VAE_PATH,
architecture="",
fastvideo_args=fastvideo_args)
videoprocessor = VideoProcessor(vae_scale_factor=8)
with torch.inference_mode():
video = vae.decode(latents[0].unsqueeze(0).to(device))
video = videoprocessor.postprocess_video(video)
video_path = os.path.join("/workspace/FastVideo/debug_videos",
infos["caption"][0][:50] + ".mp4")
export_to_video(video[0], video_path, fps=16)
# Move data to device
# latents = latents.to(device)
# embeddings = embeddings.to(device)
if world_size > 1:
dist.barrier()
# Benchmark
if rank == 0:
print(f"Benchmarking with batch_size={args.batch_size}")
start_time = time.time()
total_samples = 0
for i, (latents, embeddings, masks,
infos) in enumerate(tqdm.tqdm(dataloader, total=args.num_batches)):
if i >= args.num_batches:
break
# Move data to device
latents = latents.to(device)
embeddings = embeddings.to(device)
# Calculate actual batch size
batch_size = latents.size(0)
total_samples += batch_size
# Print progress only from rank 0
if rank == 0 and (i + 1) % 10 == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
print(
f"Batch {i+1}/{args.num_batches}, Speed: {samples_per_sec:.2f} samples/sec"
)
# Final statistics
if world_size > 1:
dist.barrier()
if rank == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
print("\nBenchmark Results:")
print(f"Total time: {elapsed:.2f} seconds")
print(f"Total samples: {total_samples}")
print(f"Average speed: {samples_per_sec:.2f} samples/sec")
print(f"Time per batch: {elapsed/args.num_batches*1000:.2f} ms")
if world_size > 1:
dist.destroy_process_group()
@@ -46,7 +46,8 @@ class DataSetProg(metaclass=SingletonMeta):
for i in range(self.num_workers):
self.n_used_elements[i] = 0
per_worker = int(math.ceil(len(self.elements) / float(self.num_workers)))
per_worker = int(
math.ceil(len(self.elements) / float(self.num_workers)))
start = i * per_worker
end = min(start + per_worker, len(self.elements))
self.worker_elements[i] = self.elements[start:end]
@@ -57,7 +58,9 @@ class DataSetProg(metaclass=SingletonMeta):
else:
worker_id = work_info.id
idx = self.worker_elements[worker_id][self.n_used_elements[worker_id] % len(self.worker_elements[worker_id])]
idx = self.worker_elements[worker_id][
self.n_used_elements[worker_id] %
len(self.worker_elements[worker_id])]
self.n_used_elements[worker_id] += 1
return idx
@@ -65,7 +68,10 @@ class DataSetProg(metaclass=SingletonMeta):
dataset_prog = DataSetProg()
def filter_resolution(h, w, max_h_div_w_ratio=17 / 16, min_h_div_w_ratio=8 / 16):
def filter_resolution(h,
w,
max_h_div_w_ratio=17 / 16,
min_h_div_w_ratio=8 / 16):
if h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio:
return True
return False
@@ -73,7 +79,14 @@ def filter_resolution(h, w, max_h_div_w_ratio=17 / 16, min_h_div_w_ratio=8 / 16)
class T2V_dataset(Dataset):
def __init__(self, args, transform, temporal_sample, tokenizer, transform_topcrop):
def __init__(self,
args,
transform,
temporal_sample,
tokenizer,
transform_topcrop,
start_idx=0):
self.start_idx = start_idx
self.data = args.data_merge_path
self.num_frames = args.num_frames
self.train_fps = args.train_fps
@@ -102,7 +115,8 @@ class T2V_dataset(Dataset):
self.lengths = self.sample_num_frames
n_elements = len(cap_list)
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list, n_elements)
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
n_elements)
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
@@ -129,7 +143,8 @@ class T2V_dataset(Dataset):
video_path = dataset_prog.cap_list[idx]["path"]
assert os.path.exists(video_path), f"file {video_path} do not exist!"
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
torchvision_video, _, metadata = torchvision.io.read_video(video_path, output_format="TCHW")
torchvision_video, _, metadata = torchvision.io.read_video(
video_path, output_format="TCHW")
video = torchvision_video[frame_indices]
video = self.transform(video)
video = rearrange(video, "t c h w -> c t h w")
@@ -160,16 +175,17 @@ class T2V_dataset(Dataset):
)
input_ids = text_tokens_and_mask["input_ids"]
cond_mask = text_tokens_and_mask["attention_mask"]
return dict(
pixel_values=video,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=video_path,
)
return dict(pixel_values=video,
text=text,
input_ids=input_ids,
cond_mask=cond_mask,
path=video_path,
fps=dataset_prog.cap_list[idx]["fps"],
duration=dataset_prog.cap_list[idx]["duration"])
def get_image(self, idx):
image_data = dataset_prog.cap_list[idx] # [{'path': path, 'cap': cap}, ...]
image_data = dataset_prog.cap_list[
idx] # [{'path': path, 'cap': cap}, ...]
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
image = torch.from_numpy(np.array(image)) # [h, w, c]
@@ -178,13 +194,15 @@ class T2V_dataset(Dataset):
# h, w = i.shape[-2:]
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
image = (self.transform_topcrop(image) if "human_images" in image_data["path"] else self.transform(image)
image = (self.transform_topcrop(image) if "human_images"
in image_data["path"] else self.transform(image)
) # [1 C H W] -> num_img [1 C H W]
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
caps = (image_data["cap"] if isinstance(image_data["cap"], list) else [image_data["cap"]])
caps = (image_data["cap"]
if isinstance(image_data["cap"], list) else [image_data["cap"]])
caps = [random.choice(caps)]
text = caps
input_ids, cond_mask = [], []
@@ -238,10 +256,12 @@ class T2V_dataset(Dataset):
cnt_no_resolution += 1
continue
else:
if (resolution.get("height", None) is None or resolution.get("width", None) is None):
if (resolution.get("height", None) is None
or resolution.get("width", None) is None):
cnt_no_resolution += 1
continue
height, width = i["resolution"]["height"], i["resolution"]["width"]
height, width = i["resolution"]["height"], i["resolution"][
"width"]
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
is_pick = filter_resolution(
@@ -259,29 +279,34 @@ class T2V_dataset(Dataset):
i["num_frames"] = math.ceil(fps * duration)
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
if i["num_frames"] / fps > self.video_length_tolerance_range * (
self.num_frames / self.train_fps *
self.speed_factor): # too long video is not suitable for this training stage (self.num_frames)
self.num_frames / self.train_fps * self.speed_factor
): # too long video is not suitable for this training stage (self.num_frames)
cnt_too_long += 1
continue
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
frame_interval = fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, i["num_frames"], frame_interval).astype(int)
frame_indices = np.arange(start_frame_idx, i["num_frames"],
frame_interval).astype(int)
# comment out it to enable dynamic frames training
if (len(frame_indices) < self.num_frames and random.random() < self.drop_short_ratio):
if (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio):
cnt_too_short += 1
continue
# too long video will be temporal-crop randomly
if len(frame_indices) > self.num_frames:
begin_index, end_index = self.temporal_sample(len(frame_indices))
begin_index, end_index = self.temporal_sample(
len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
# frame_indices = frame_indices[:self.num_frames] # head crop
i["sample_frame_index"] = frame_indices.tolist()
new_cap_list.append(i)
i["sample_num_frames"] = len(i["sample_frame_index"]) # will use in dataloader(group sampler)
i["sample_num_frames"] = len(
i["sample_frame_index"]
) # will use in dataloader(group sampler)
sample_num_frames.append(i["sample_num_frames"])
elif path.endswith(".jpg"): # image
cnt_img += 1
@@ -290,13 +315,15 @@ class T2V_dataset(Dataset):
sample_num_frames.append(i["sample_num_frames"])
else:
raise NameError(
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image")
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
)
# import ipdb;ipdb.set_trace()
main_print(
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}")
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
)
return new_cap_list, sample_num_frames
def decord_read(self, path, frame_indices):
@@ -308,11 +335,14 @@ class T2V_dataset(Dataset):
def read_jsons(self, data):
cap_lists = []
with open(data, "r") as f:
folder_anno = [i.strip().split(",") for i in f.readlines() if len(i.strip()) > 0]
with open(data) as f:
folder_anno = [
i.strip().split(",") for i in f.readlines()
if len(i.strip()) > 0
]
print(folder_anno)
for folder, anno in folder_anno:
with open(anno, "r") as f:
with open(anno) as f:
sub_list = json.load(f)
for i in range(len(sub_list)):
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
@@ -320,5 +350,5 @@ class T2V_dataset(Dataset):
return cap_lists
def get_cap_list(self):
cap_lists = self.read_jsons(self.data)
cap_lists = self.read_jsons(self.data)[self.start_idx:]
return cap_lists
@@ -21,15 +21,19 @@ def center_crop_arr(pil_image, image_size):
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
"""
while min(*pil_image.size) >= 2 * image_size:
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size), resample=Image.BOX)
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size),
resample=Image.BOX)
scale = image_size / min(*pil_image.size)
pil_image = pil_image.resize(tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC)
pil_image = pil_image.resize(tuple(
round(x * scale) for x in pil_image.size),
resample=Image.BICUBIC)
arr = np.array(pil_image)
crop_y = (arr.shape[0] - image_size) // 2
crop_x = (arr.shape[1] - image_size) // 2
return Image.fromarray(arr[crop_y:crop_y + image_size, crop_x:crop_x + image_size])
return Image.fromarray(arr[crop_y:crop_y + image_size,
crop_x:crop_x + image_size])
def crop(clip, i, j, h, w):
@@ -44,7 +48,9 @@ def crop(clip, i, j, h, w):
def resize(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
)
return torch.nn.functional.interpolate(
clip,
size=target_size,
@@ -56,7 +62,9 @@ def resize(clip, target_size, interpolation_mode):
def resize_scale(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
)
H, W = clip.size(-2), clip.size(-1)
scale_ = target_size[0] / min(H, W)
return torch.nn.functional.interpolate(
@@ -166,7 +174,8 @@ def normalize_video(clip):
"""
_is_tensor_video_clip(clip)
if not clip.dtype == torch.uint8:
raise TypeError("clip tensor should have data type uint8. Got %s" % str(clip.dtype))
raise TypeError("clip tensor should have data type uint8. Got %s" %
str(clip.dtype))
# return clip.float().permute(3, 0, 1, 2) / 255.0
return clip.float() / 255.0
@@ -227,7 +236,9 @@ class RandomCropVideo:
th, tw = self.size
if h < th or w < tw:
raise ValueError(f"Required crop size {(th, tw)} is larger than input image size {(h, w)}")
raise ValueError(
f"Required crop size {(th, tw)} is larger than input image size {(h, w)}"
)
if w == tw and h == th:
return 0, 0, h, w
@@ -301,7 +312,9 @@ class LongSideResizeVideo:
else:
h = int(h * self.size / w)
w = self.size
resize_clip = resize(clip, target_size=(h, w), interpolation_mode=self.interpolation_mode)
resize_clip = resize(clip,
target_size=(h, w),
interpolation_mode=self.interpolation_mode)
return resize_clip
def __repr__(self) -> str:
@@ -321,7 +334,8 @@ class CenterCropResizeVideo:
interpolation_mode="bilinear",
):
if len(size) != 2:
raise ValueError(f"size should be tuple (height, width), instead got {size}")
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
self.top_crop = top_crop
self.interpolation_mode = interpolation_mode
@@ -335,7 +349,10 @@ class CenterCropResizeVideo:
size is (T, C, crop_size, crop_size)
"""
# clip_center_crop = center_crop_using_short_edge(clip)
clip_center_crop = center_crop_th_tw(clip, self.size[0], self.size[1], top_crop=self.top_crop)
clip_center_crop = center_crop_th_tw(clip,
self.size[0],
self.size[1],
top_crop=self.top_crop)
# import ipdb;ipdb.set_trace()
clip_center_crop_resize = resize(
clip_center_crop,
@@ -361,7 +378,8 @@ class UCFCenterCropVideo:
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(f"size should be tuple (height, width), instead got {size}")
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -376,7 +394,9 @@ class UCFCenterCropVideo:
torch.tensor: scale resized / center cropped video clip.
size is (T, C, crop_size, crop_size)
"""
clip_resize = resize_scale(clip=clip, target_size=self.size, interpolation_mode=self.interpolation_mode)
clip_resize = resize_scale(clip=clip,
target_size=self.size,
interpolation_mode=self.interpolation_mode)
clip_center_crop = center_crop(clip_resize, self.size)
return clip_center_crop
@@ -396,7 +416,8 @@ class KineticsRandomCropResizeVideo:
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(f"size should be tuple (height, width), instead got {size}")
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -405,7 +426,8 @@ class KineticsRandomCropResizeVideo:
def __call__(self, clip):
clip_random_crop = random_shift_crop(clip)
clip_resize = resize(clip_random_crop, self.size, self.interpolation_mode)
clip_resize = resize(clip_random_crop, self.size,
self.interpolation_mode)
return clip_resize
@@ -418,7 +440,8 @@ class CenterCropVideo:
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(f"size should be tuple (height, width), instead got {size}")
raise ValueError(
f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -514,7 +537,7 @@ class RandomHorizontalFlipVideo:
# ------------------------------------------------------------
# --------------------- Sampling ---------------------------
# ------------------------------------------------------------
class TemporalRandomCrop(object):
class TemporalRandomCrop:
"""Temporally crop the given frame indices at a random location.
Args:
@@ -531,7 +554,7 @@ class TemporalRandomCrop(object):
return begin_index, end_index
class DynamicSampleDuration(object):
class DynamicSampleDuration:
"""Temporally crop the given frame indices at a random location.
Args:
@@ -545,7 +568,8 @@ class DynamicSampleDuration(object):
def __call__(self, t, h, w):
if self.extra_1:
t = t - 1
truncate_t_list = list(range(t + 1))[t // 2:][::self.t_stride] # need half at least
truncate_t_list = list(
range(t + 1))[t // 2:][::self.t_stride] # need half at least
truncate_t = random.choice(truncate_t_list)
if self.extra_1:
truncate_t = truncate_t + 1
@@ -560,14 +584,18 @@ if __name__ == "__main__":
from torchvision import transforms
from torchvision.utils import save_image
vframes, aframes, info = io.read_video(filename="./v_Archery_g01_c03.avi", pts_unit="sec", output_format="TCHW")
vframes, aframes, info = io.read_video(filename="./v_Archery_g01_c03.avi",
pts_unit="sec",
output_format="TCHW")
trans = transforms.Compose([
Normalize255(),
RandomHorizontalFlipVideo(),
UCFCenterCropVideo(512),
# NormalizeVideo(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
transforms.Normalize(mean=[0.5, 0.5, 0.5],
std=[0.5, 0.5, 0.5],
inplace=True),
])
target_video_len = 32
@@ -582,7 +610,10 @@ if __name__ == "__main__":
# print(start_frame_ind)
# print(end_frame_ind)
assert end_frame_ind - start_frame_ind >= target_video_len
frame_indice = np.linspace(start_frame_ind, end_frame_ind - 1, target_video_len, dtype=int)
frame_indice = np.linspace(start_frame_ind,
end_frame_ind - 1,
target_video_len,
dtype=int)
print(frame_indice)
select_vframes = vframes[frame_indice]
@@ -593,11 +624,14 @@ if __name__ == "__main__":
print(select_vframes_trans.shape)
print(select_vframes_trans.dtype)
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) * 255).to(dtype=torch.uint8)
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) *
255).to(dtype=torch.uint8)
print(select_vframes_trans_int.dtype)
print(select_vframes_trans_int.permute(0, 2, 3, 1).shape)
io.write_video("./test.avi", select_vframes_trans_int.permute(0, 2, 3, 1), fps=8)
io.write_video("./test.avi",
select_vframes_trans_int.permute(0, 2, 3, 1),
fps=8)
for i in range(target_video_len):
save_image(
+10
View File
@@ -0,0 +1,10 @@
from huggingface_hub import HfApi, upload_folder
api = HfApi()
repo_id = "weizhou03/HD-Mixkit-Finetune-Wan" # customize this
api.create_repo(repo_id=repo_id, repo_type="dataset")
upload_folder(repo_id=repo_id,
folder_path="/workspace/data/HD-Mixkit-Finetune-Wan",
repo_type="dataset",
path_in_repo="")
+3 -1
View File
@@ -5,7 +5,8 @@ from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size, get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size, get_world_group,
init_distributed_environment, initialize_model_parallel)
init_distributed_environment, initialize_model_parallel,
model_parallel_is_initialized)
from fastvideo.v1.distributed.utils import *
__all__ = [
@@ -17,4 +18,5 @@ __all__ = [
"get_tensor_model_parallel_world_size",
"cleanup_dist_env_and_memory",
"get_world_group",
"model_parallel_is_initialized",
]
@@ -1,16 +1,182 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/device_communicators/base_device_communicator.py
from typing import Optional
from typing import Any, Optional, Tuple
import torch
import torch.distributed as dist
from torch.distributed import ProcessGroup
from torch import Tensor
from torch.distributed import ProcessGroup, ReduceOp
class DistributedAutograd:
"""Collection of autograd functions for distributed operations.
This class provides custom autograd functions for distributed operations like all_reduce,
all_gather, and all_to_all. Each operation is implemented as a static inner class with
proper forward and backward implementations.
"""
class AllReduce(torch.autograd.Function):
"""Differentiable all_reduce operation.
The gradient of all_reduce is another all_reduce operation since the operation
combines values from all ranks equally.
"""
@staticmethod
def forward(ctx: Any,
group: ProcessGroup,
input_: Tensor,
op: Optional[dist.ReduceOp] = None) -> Tensor:
ctx.group = group
ctx.op = op
output = input_.clone()
dist.all_reduce(output, group=group, op=op)
return output
@staticmethod
def backward(ctx: Any,
grad_output: Tensor) -> Tuple[None, Tensor, None]:
grad_output = grad_output.clone()
dist.all_reduce(grad_output, group=ctx.group, op=ctx.op)
return None, grad_output, None
class AllGather(torch.autograd.Function):
"""Differentiable all_gather operation.
The operation gathers tensors from all ranks and concatenates them along a specified dimension.
The backward pass uses reduce_scatter to efficiently distribute gradients back to source ranks.
"""
@staticmethod
def forward(ctx: Any, group: ProcessGroup, input_: Tensor,
world_size: int, dim: int) -> Tensor:
ctx.group = group
ctx.world_size = world_size
ctx.dim = dim
ctx.input_shape = input_.shape
input_size = input_.size()
output_size = (input_size[0] * world_size, ) + input_size[1:]
output_tensor = torch.empty(output_size,
dtype=input_.dtype,
device=input_.device)
dist.all_gather_into_tensor(output_tensor, input_, group=group)
output_tensor = output_tensor.reshape((world_size, ) + input_size)
output_tensor = output_tensor.movedim(0, dim)
output_tensor = output_tensor.reshape(input_size[:dim] +
(world_size *
input_size[dim], ) +
input_size[dim + 1:])
return output_tensor
@staticmethod
def backward(ctx: Any,
grad_output: Tensor) -> Tuple[None, Tensor, None, None]:
# Split the gradient tensor along the gathered dimension
dim_size = grad_output.size(ctx.dim) // ctx.world_size
grad_chunks = grad_output.reshape(grad_output.shape[:ctx.dim] +
(ctx.world_size, dim_size) +
grad_output.shape[ctx.dim + 1:])
grad_chunks = grad_chunks.movedim(ctx.dim, 0)
# Each rank only needs its corresponding gradient
grad_input = torch.empty(ctx.input_shape,
dtype=grad_output.dtype,
device=grad_output.device)
dist.reduce_scatter_tensor(grad_input,
grad_chunks.contiguous(),
group=ctx.group)
return None, grad_input, None, None
class AllToAll4D(torch.autograd.Function):
"""Differentiable all_to_all operation specialized for 4D tensors.
This operation is particularly useful for attention operations where we need to
redistribute data across ranks for efficient parallel processing.
The operation supports two modes:
1. scatter_dim=2, gather_dim=1: Used for redistributing attention heads
2. scatter_dim=1, gather_dim=2: Used for redistributing sequence dimensions
"""
@staticmethod
def forward(ctx: Any, group: ProcessGroup, input_: Tensor,
world_size: int, scatter_dim: int,
gather_dim: int) -> Tensor:
ctx.group = group
ctx.world_size = world_size
ctx.scatter_dim = scatter_dim
ctx.gather_dim = gather_dim
if world_size == 1:
return input_
assert input_.dim(
) == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
if scatter_dim == 2 and gather_dim == 1:
bs, shard_seqlen, hc, hs = input_.shape
seqlen = shard_seqlen * world_size
shard_hc = hc // world_size
input_t = input_.reshape(bs, shard_seqlen, world_size, shard_hc,
hs).transpose(0, 2).contiguous()
output = torch.empty_like(input_t)
dist.all_to_all_single(output, input_t, group=group)
output = output.reshape(seqlen, bs, shard_hc,
hs).transpose(0, 1).contiguous()
output = output.reshape(bs, seqlen, shard_hc, hs)
return output
elif scatter_dim == 1 and gather_dim == 2:
bs, seqlen, shard_hc, hs = input_.shape
hc = shard_hc * world_size
shard_seqlen = seqlen // world_size
input_t = input_.reshape(bs, world_size, shard_seqlen, shard_hc,
hs)
input_t = input_t.transpose(0, 3).transpose(0, 1).contiguous()
input_t = input_t.reshape(world_size, shard_hc, shard_seqlen,
bs, hs)
output = torch.empty_like(input_t)
dist.all_to_all_single(output, input_t, group=group)
output = output.reshape(hc, shard_seqlen, bs, hs)
output = output.transpose(0, 2).contiguous()
output = output.reshape(bs, shard_seqlen, hc, hs)
return output
else:
raise RuntimeError(
f"Invalid scatter_dim={scatter_dim}, gather_dim={gather_dim}. "
f"Only (scatter_dim=2, gather_dim=1) and (scatter_dim=1, gather_dim=2) are supported."
)
@staticmethod
def backward(
ctx: Any,
grad_output: Tensor) -> Tuple[None, Tensor, None, None, None]:
if ctx.world_size == 1:
return None, grad_output, None, None, None
# For backward pass, we swap scatter_dim and gather_dim
output = DistributedAutograd.AllToAll4D.apply(
ctx.group, grad_output, ctx.world_size, ctx.gather_dim,
ctx.scatter_dim)
return None, output, None, None, None
class DeviceCommunicatorBase:
"""
Base class for device-specific communicator.
Base class for device-specific communicator with autograd support.
It can use the `cpu_group` to initialize the communicator.
If the device has PyTorch integration (PyTorch can recognize its
communication backend), the `device_group` will also be given.
@@ -33,35 +199,28 @@ class DeviceCommunicatorBase:
self.rank_in_group = dist.get_group_rank(self.cpu_group,
self.global_rank)
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
dist.all_reduce(input_, group=self.device_group)
return input_
def all_reduce(self,
input_: torch.Tensor,
op: Optional[dist.ReduceOp] = ReduceOp.SUM) -> torch.Tensor:
"""Performs an all_reduce operation with gradient support."""
return DistributedAutograd.AllReduce.apply(self.device_group, input_,
op)
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
"""Performs an all_gather operation with gradient support."""
if dim < 0:
# Convert negative dim to positive.
dim += input_.dim()
input_size = input_.size()
# NOTE: we have to use concat-style all-gather here,
# stack-style all-gather has compatibility issues with
# torch.compile . see https://github.com/pytorch/pytorch/issues/138795
output_size = (input_size[0] * self.world_size, ) + input_size[1:]
# Allocate output tensor.
output_tensor = torch.empty(output_size,
dtype=input_.dtype,
device=input_.device)
# All-gather.
dist.all_gather_into_tensor(output_tensor,
input_,
group=self.device_group)
# Reshape
output_tensor = output_tensor.reshape((self.world_size, ) + input_size)
output_tensor = output_tensor.movedim(0, dim)
output_tensor = output_tensor.reshape(input_size[:dim] +
(self.world_size *
input_size[dim], ) +
input_size[dim + 1:])
return output_tensor
return DistributedAutograd.AllGather.apply(self.device_group, input_,
self.world_size, dim)
def all_to_all_4D(self,
input_: torch.Tensor,
scatter_dim: int = 2,
gather_dim: int = 1) -> torch.Tensor:
"""Performs a 4D all-to-all operation with gradient support."""
return DistributedAutograd.AllToAll4D.apply(self.device_group, input_,
self.world_size,
scatter_dim, gather_dim)
def gather(self,
input_: torch.Tensor,
@@ -95,81 +254,6 @@ class DeviceCommunicatorBase:
output_tensor = None
return output_tensor
def all_to_all_4D(self,
input_: torch.Tensor,
scatter_dim: int = 2,
gather_dim: int = 1) -> torch.Tensor:
"""Specialized all-to-all operation for 4D tensors (e.g., for QKV matrices).
Args:
input_ (torch.Tensor): 4D input tensor to be scattered and gathered.
scatter_dim (int, optional): Dimension along which to scatter. Defaults to 2.
gather_dim (int, optional): Dimension along which to gather. Defaults to 1.
Returns:
torch.Tensor: Output tensor after all-to-all operation.
"""
# Bypass the function if we are using only 1 GPU.
if self.world_size == 1:
return input_
assert input_.dim(
) == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
if scatter_dim == 2 and gather_dim == 1:
# input: (bs, seqlen/P, hc, hs) output: (bs, seqlen, hc/P, hs)
bs, shard_seqlen, hc, hs = input_.shape
seqlen = shard_seqlen * self.world_size
shard_hc = hc // self.world_size
# Reshape and transpose for scattering
input_t = (input_.reshape(bs, shard_seqlen, self.world_size,
shard_hc, hs).transpose(0,
2).contiguous())
output = torch.empty_like(input_t)
torch.distributed.all_to_all_single(output,
input_t,
group=self.device_group)
torch.cuda.synchronize()
# Reshape and transpose back
output = output.reshape(seqlen, bs, shard_hc,
hs).transpose(0, 1).contiguous().reshape(
bs, seqlen, shard_hc, hs)
return output
elif scatter_dim == 1 and gather_dim == 2:
# input: (bs, seqlen, hc/P, hs) output: (bs, seqlen/P, hc, hs)
bs, seqlen, shard_hc, hs = input_.shape
hc = shard_hc * self.world_size
shard_seqlen = seqlen // self.world_size
# Reshape and transpose for scattering
input_t = (input_.reshape(bs, self.world_size, shard_seqlen,
shard_hc, hs).transpose(0, 3).transpose(
0, 1).contiguous().reshape(
self.world_size, shard_hc,
shard_seqlen, bs, hs))
output = torch.empty_like(input_t)
torch.distributed.all_to_all_single(output,
input_t,
group=self.device_group)
torch.cuda.synchronize()
# Reshape and transpose back
output = output.reshape(hc, shard_seqlen, bs,
hs).transpose(0, 2).contiguous().reshape(
bs, shard_seqlen, hc, hs)
return output
else:
raise RuntimeError(
"scatter_dim must be 1 or 2 and gather_dim must be 1 or 2")
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank."""
@@ -29,17 +29,19 @@ class CudaCommunicator(DeviceCommunicatorBase):
device=self.device,
)
def all_reduce(self, input_):
def all_reduce(self,
input_,
op: Optional[torch.distributed.ReduceOp] = None):
pynccl_comm = self.pynccl_comm
assert pynccl_comm is not None
out = pynccl_comm.all_reduce(input_)
out = pynccl_comm.all_reduce(input_, op=op)
if out is None:
# fall back to the default all-reduce using PyTorch.
# this usually happens during testing.
# when we run the model, allreduce only happens for the TP
# group, where we always have either custom allreduce or pynccl.
out = input_.clone()
torch.distributed.all_reduce(out, group=self.device_group)
torch.distributed.all_reduce(out, group=self.device_group, op=op)
return out
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
+13 -5
View File
@@ -35,7 +35,7 @@ from unittest.mock import patch
import torch
import torch.distributed
from torch.distributed import Backend, ProcessGroup
from torch.distributed import Backend, ProcessGroup, ReduceOp
import fastvideo.v1.envs as envs
from fastvideo.v1.distributed.device_communicators.base_device_communicator import (
@@ -260,7 +260,11 @@ class GroupCoordinator:
with torch.cuda.stream(stream):
yield graph_capture_context
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
def all_reduce(
self,
input_: torch.Tensor,
op: Optional[torch.distributed.ReduceOp] = ReduceOp.SUM
) -> torch.Tensor:
"""
User-facing all-reduce function before we actually call the
all-reduce operation.
@@ -283,10 +287,14 @@ class GroupCoordinator:
return torch.ops.vllm.all_reduce(input_,
group_name=self.unique_name)
else:
return self._all_reduce_out_place(input_)
return self._all_reduce_out_place(input_, op=op)
def _all_reduce_out_place(self, input_: torch.Tensor) -> torch.Tensor:
return self.device_communicator.all_reduce(input_)
def _all_reduce_out_place(
self,
input_: torch.Tensor,
op: Optional[torch.distributed.ReduceOp] = ReduceOp.SUM
) -> torch.Tensor:
return self.device_communicator.all_reduce(input_, op=op)
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
world_size = self.world_size
+18 -5
View File
@@ -11,8 +11,11 @@ from fastvideo.v1.configs.sample.base import SamplingParam
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.v1.entrypoints.cli.utils import RaiseNotImplementedAction
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import FlexibleArgumentParser
logger = init_logger(__name__)
class GenerateSubcommand(CLISubcommand):
"""The `generate` subcommand for the FastVideo CLI"""
@@ -35,12 +38,22 @@ class GenerateSubcommand(CLISubcommand):
excluded_args = ['subparser', 'config', 'dispatch_function']
FastVideoArgs.from_cli_args(args)
filtered_args = {}
for k, v in vars(args).items():
if k not in excluded_args and v is not None:
filtered_args[k] = v
merged_args = {**filtered_args}
provided_args = {}
for k, v in vars(args).items():
if (k not in excluded_args and v is not None
and hasattr(args, '_provided') and k in args._provided):
provided_args[k] = v
if 'model_path' in vars(args) and args.model_path is not None:
provided_args['model_path'] = args.model_path
if 'prompt' in vars(args) and args.prompt is not None:
provided_args['prompt'] = args.prompt
merged_args = {**provided_args}
logger.info('CLI Args: %s', merged_args)
if 'model_path' not in merged_args or not merged_args['model_path']:
raise ValueError(
@@ -0,0 +1,29 @@
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from fastvideo.v1.pipelines.wan.wan_latent_pipeline import WanLatentPipeline
def main():
print("Starting data preprocessor")
pipeline = WanLatentPipeline.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
train_dataset = getdataset(args)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
)
for batch in train_dataloader:
pipeline(batch)
if __name__ == "__main__":
main()
+41 -6
View File
@@ -7,6 +7,7 @@ diffusion models.
"""
import gc
import math
import os
import time
from typing import Any, Dict, List, Optional, Union
@@ -180,12 +181,46 @@ class VideoGenerator:
f"height={sampling_param.height}, width={sampling_param.width}, "
f"num_frames={sampling_param.num_frames}")
if (
sampling_param.num_frames - 1
) % fastvideo_args.vae_config.arch_config.temporal_compression_ratio != 0:
raise ValueError(
f"num_frames-1 must be a multiple of {fastvideo_args.vae_config.arch_config.temporal_compression_ratio}, got {sampling_param.num_frames}"
)
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
num_frames = sampling_param.num_frames
num_gpus = fastvideo_args.num_gpus
use_temporal_scaling_frames = fastvideo_args.vae_config.use_temporal_scaling_frames
# Adjust number of frames based on number of GPUs
if use_temporal_scaling_frames:
orig_latent_num_frames = (num_frames -
1) // temporal_scale_factor + 1
else: # stepvideo only
orig_latent_num_frames = sampling_param.num_frames // 17 * 3
if orig_latent_num_frames % fastvideo_args.num_gpus != 0:
# Adjust latent frames to be divisible by number of GPUs
if sampling_param.num_frames_round_down:
# Ensure we have at least 1 batch per GPU
new_latent_num_frames = max(
1, (orig_latent_num_frames // num_gpus)) * num_gpus
else:
new_latent_num_frames = math.ceil(
orig_latent_num_frames / num_gpus) * num_gpus
if use_temporal_scaling_frames:
# Convert back to number of frames, ensuring num_frames-1 is a multiple of temporal_scale_factor
new_num_frames = (new_latent_num_frames -
1) * temporal_scale_factor + 1
else: # stepvideo only
# Find the least common multiple of 3 and num_gpus
divisor = math.lcm(3, num_gpus)
# Round up to the nearest multiple of this LCM
new_latent_num_frames = (
(new_latent_num_frames + divisor - 1) // divisor) * divisor
# Convert back to actual frames using the StepVideo formula
new_num_frames = new_latent_num_frames // 3 * 17
logger.info(
"Adjusting number of frames from %s to %s based on number of GPUs (%s)",
sampling_param.num_frames, new_num_frames,
fastvideo_args.num_gpus)
sampling_param.num_frames = new_num_frames
# Calculate sizes
target_height = align_to(sampling_param.height, 16)
+358 -1
View File
@@ -70,7 +70,7 @@ class FastVideoArgs:
# Text encoder configuration
DEFAULT_TEXT_ENCODER_PRECISIONS = (
"fp16",
"fp16",
# "fp16",
)
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
@@ -100,6 +100,10 @@ class FastVideoArgs:
device_str: Optional[str] = None
device = None
@property
def training_mode(self) -> bool:
return not self.inference_mode
def __post_init__(self):
pass
@@ -132,6 +136,13 @@ class FastVideoArgs:
help="The distributed executor backend to use",
)
parser.add_argument(
"--inference-mode",
action=StoreBoolean,
default=FastVideoArgs.inference_mode,
help="Whether to use inference mode",
)
# HuggingFace specific parameters
parser.add_argument(
"--trust-remote-code",
@@ -423,3 +434,349 @@ def get_current_fastvideo_args() -> FastVideoArgs:
# TODO(will): may need to handle this for CI.
raise ValueError("Current fastvideo args is not set.")
return _current_fastvideo_args
@dataclasses.dataclass
class TrainingArgs(FastVideoArgs):
data_path: str = ""
dataloader_num_workers: int = 0
num_height: int = 0
num_width: int = 0
num_frames: int = 0
train_batch_size: int = 0
num_latent_t: int = 0
group_frame: bool = False
group_resolution: bool = False
# text encoder & vae & diffusion model
pretrained_model_name_or_path: str = ""
dit_model_name_or_path: str = ""
cache_dir: str = ""
# diffusion setting
ema_decay: float = 0.0
ema_start_step: int = 0
cfg: float = 0.0
precondition_outputs: bool = False
# validation & logs
validation_prompt_dir: str = ""
validation_sampling_steps: str = ""
validation_guidance_scale: str = ""
validation_steps: float = 0.0
log_validation: bool = False
tracker_project_name: str = ""
# seed: int
# output
output_dir: str = ""
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: str = ""
resume_from_lora_checkpoint: str = ""
logging_dir: str = ""
# optimizer & scheduler
num_train_epochs: int = 0
max_train_steps: int = 0
gradient_accumulation_steps: int = 0
learning_rate: float = 0.0
scale_lr: bool = False
lr_scheduler: str = ""
lr_warmup_steps: int = 0
max_grad_norm: float = 0.0
gradient_checkpointing: bool = False
selective_checkpointing: float = 0.0
allow_tf32: bool = False
mixed_precision: str = ""
use_cpu_offload: bool = False
# fp16_full_eval: bool
# fp16_backend: str
train_sp_batch_size: int = 0
use_lora: bool = False
lora_alpha: int = 0
lora_rank: int = 0
fsdp_sharding_startegy: str = ""
weighting_scheme: str = ""
logit_mean: float = 0.0
logit_std: float = 1.0
mode_scale: float = 0.0
# lr_scheduler
lr_scheduler: str = ""
num_euler_timesteps: int = 0
lr_num_cycles: int = 0
lr_power: float = 0.0
not_apply_cfg_solver: bool = False
distill_cfg: float = 0.0
scheduler_type: str = ""
linear_quadratic_threshold: float = 0.0
linear_range: float = 0.0
weight_decay: float = 0.0
use_ema: bool = False
multi_phased_distill_schedule: str = ""
pred_decay_weight: float = 0.0
pred_decay_type: str = ""
hunyuan_teacher_disable_cfg: bool = False
# master_weight_type
master_weight_type: str = ""
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
# Get all fields from the dataclass
attrs = [attr.name for attr in dataclasses.fields(cls)]
# Create a dictionary of attribute values, with defaults for missing attributes
kwargs = {}
for attr in attrs:
# Handle renamed attributes or those with multiple CLI names
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
# Use getattr with default value from the dataclass for potentially missing attributes
else:
default_value = getattr(cls, attr, None)
kwargs[attr] = getattr(args, attr, default_value)
return cls(**kwargs)
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
parser.add_argument("--data-path",
type=str,
required=True,
help="Path to parquet files")
parser.add_argument("--dataloader-num-workers",
type=int,
required=True,
help="Number of workers for dataloader")
parser.add_argument("--num-height",
type=int,
required=True,
help="Number of heights")
parser.add_argument("--num-width",
type=int,
required=True,
help="Number of widths")
parser.add_argument("--num-frames",
type=int,
required=True,
help="Number of frames")
# Training batch and model configuration
parser.add_argument("--train-batch-size",
type=int,
required=True,
help="Training batch size")
parser.add_argument("--num-latent-t",
type=int,
required=True,
help="Number of latent time steps")
parser.add_argument("--group-frame",
action=StoreBoolean,
help="Whether to group frames during training")
parser.add_argument("--group-resolution",
action=StoreBoolean,
help="Whether to group resolutions during training")
# Model paths
parser.add_argument("--pretrained-model-name-or-path",
type=str,
required=True,
help="Path to pretrained model or model name")
parser.add_argument("--dit-model-name-or-path",
type=str,
required=False,
help="Path to DiT model or model name")
parser.add_argument("--cache-dir",
type=str,
help="Directory to cache models")
# Diffusion settings
parser.add_argument("--ema-decay",
type=float,
default=0.999,
help="EMA decay rate")
parser.add_argument("--ema-start-step",
type=int,
default=0,
help="Step to start EMA")
parser.add_argument("--cfg",
type=float,
help="Classifier-free guidance scale")
parser.add_argument(
"--precondition-outputs",
action=StoreBoolean,
help="Whether to precondition the outputs of the model")
# Validation and logging
parser.add_argument("--validation-prompt-dir",
type=str,
help="Directory containing validation prompts")
parser.add_argument("--validation-sampling-steps",
type=str,
help="Validation sampling steps")
parser.add_argument("--validation-guidance-scale",
type=str,
help="Validation guidance scale")
parser.add_argument("--validation-steps",
type=float,
help="Number of validation steps")
parser.add_argument("--log-validation",
action=StoreBoolean,
help="Whether to log validation results")
parser.add_argument("--tracker-project-name",
type=str,
help="Project name for tracking")
# Output configuration
parser.add_argument("--output-dir",
type=str,
required=True,
help="Output directory for checkpoints and logs")
parser.add_argument("--checkpoints-total-limit",
type=int,
help="Maximum number of checkpoints to keep")
parser.add_argument("--checkpointing-steps",
type=int,
help="Steps between checkpoints")
parser.add_argument("--resume-from-checkpoint",
type=str,
help="Path to checkpoint to resume from")
parser.add_argument("--resume-from-lora-checkpoint",
type=str,
help="Path to LoRA checkpoint to resume from")
parser.add_argument("--logging-dir",
type=str,
help="Directory for logging")
# Training configuration
parser.add_argument("--num-train-epochs",
type=int,
help="Number of training epochs")
parser.add_argument("--max-train-steps",
type=int,
help="Maximum number of training steps")
parser.add_argument("--gradient-accumulation-steps",
type=int,
help="Number of steps to accumulate gradients")
parser.add_argument("--learning-rate",
type=float,
required=True,
help="Learning rate")
parser.add_argument("--scale-lr",
action=StoreBoolean,
help="Whether to scale learning rate")
parser.add_argument("--lr-scheduler",
type=str,
default="constant",
help="Learning rate scheduler type")
parser.add_argument("--lr-warmup-steps",
type=int,
default=10,
help="Number of warmup steps for learning rate")
parser.add_argument("--max-grad-norm",
type=float,
help="Maximum gradient norm")
parser.add_argument("--gradient-checkpointing",
action=StoreBoolean,
help="Whether to use gradient checkpointing")
parser.add_argument("--selective-checkpointing",
type=float,
help="Selective checkpointing threshold")
parser.add_argument("--allow-tf32",
action=StoreBoolean,
help="Whether to allow TF32")
parser.add_argument("--mixed-precision",
type=str,
help="Mixed precision training type")
parser.add_argument("--train-sp-batch-size",
type=int,
help="Training spatial parallelism batch size")
# LoRA configuration
parser.add_argument("--use-lora",
action=StoreBoolean,
help="Whether to use LoRA")
parser.add_argument("--lora-alpha",
type=int,
help="LoRA alpha parameter")
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
parser.add_argument("--fsdp-sharding-strategy",
type=str,
help="FSDP sharding strategy")
parser.add_argument(
"--weighting_scheme",
type=str,
default="uniform",
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "uniform"],
)
parser.add_argument(
"--logit_mean",
type=float,
default=0.0,
help="mean to use when using the `'logit_normal'` weighting scheme.",
)
parser.add_argument(
"--logit_std",
type=float,
default=1.0,
help="std to use when using the `'logit_normal'` weighting scheme.",
)
parser.add_argument(
"--mode_scale",
type=float,
default=1.29,
help=
"Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
)
# Additional training parameters
parser.add_argument("--num-euler-timesteps",
type=int,
help="Number of Euler timesteps")
parser.add_argument("--lr-num-cycles",
type=int,
help="Number of learning rate cycles")
parser.add_argument("--lr-power",
type=float,
help="Learning rate power")
parser.add_argument("--not-apply-cfg-solver",
action=StoreBoolean,
help="Whether to not apply CFG solver")
parser.add_argument("--distill-cfg",
type=float,
help="Distillation CFG scale")
parser.add_argument("--scheduler-type", type=str, help="Scheduler type")
parser.add_argument("--linear-quadratic-threshold",
type=float,
help="Linear quadratic threshold")
parser.add_argument("--linear-range", type=float, help="Linear range")
parser.add_argument("--weight-decay", type=float, help="Weight decay")
parser.add_argument("--use-ema",
action=StoreBoolean,
help="Whether to use EMA")
parser.add_argument("--multi-phased-distill-schedule",
type=str,
help="Multi-phased distillation schedule")
parser.add_argument("--pred-decay-weight",
type=float,
help="Prediction decay weight")
parser.add_argument("--pred-decay-type",
type=str,
help="Prediction decay type")
parser.add_argument("--hunyuan-teacher-disable-cfg",
action=StoreBoolean,
help="Whether to disable CFG for Hunyuan teacher")
parser.add_argument("--master-weight-type",
type=str,
help="Master weight type")
return parser
-39
View File
@@ -9,7 +9,6 @@ import torch.nn.functional as F
# TODO (will): remove this dependency
from fastvideo.v1.layers.custom_op import CustomOp
from fastvideo.v1.platforms import current_platform
@CustomOp.register("silu_and_mul")
@@ -25,21 +24,12 @@ class SiluAndMul(CustomOp):
def __init__(self) -> None:
super().__init__()
if current_platform.is_cuda_alike() or current_platform.is_cpu():
self.op = torch.ops._C.silu_and_mul
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
d = x.shape[-1] // 2
return F.silu(x[..., :d]) * x[..., d:]
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
output_shape = (x.shape[:-1] + (d, ))
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
self.op(out, x)
return out
@CustomOp.register("gelu_and_mul")
class GeluAndMul(CustomOp):
@@ -57,24 +47,12 @@ class GeluAndMul(CustomOp):
self.approximate = approximate
if approximate not in ("none", "tanh"):
raise ValueError(f"Unknown approximate mode: {approximate}")
if current_platform.is_cuda_alike() or current_platform.is_cpu():
if approximate == "none":
self.op = torch.ops._C.gelu_and_mul
elif approximate == "tanh":
self.op = torch.ops._C.gelu_tanh_and_mul
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
d = x.shape[-1] // 2
return F.gelu(x[..., :d], approximate=self.approximate) * x[..., d:]
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
output_shape = (x.shape[:-1] + (d, ))
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
self.op(out, x)
return out
def extra_repr(self) -> str:
return f'approximate={repr(self.approximate)}'
@@ -84,8 +62,6 @@ class NewGELU(CustomOp):
def __init__(self):
super().__init__()
if current_platform.is_cuda_alike() or current_platform.is_cpu():
self.op = torch.ops._C.gelu_new
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
@@ -93,32 +69,17 @@ class NewGELU(CustomOp):
return 0.5 * x * (1.0 + torch.tanh(c *
(x + 0.044715 * torch.pow(x, 3.0))))
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
self.op(out, x)
return out
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
return self.op(x)
@CustomOp.register("quick_gelu")
class QuickGELU(CustomOp):
# https://github.com/huggingface/transformers/blob/main/src/transformers/activations.py#L90
def __init__(self):
super().__init__()
if current_platform.is_cuda_alike() or current_platform.is_cpu():
self.op = torch.ops._C.gelu_quick
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
"""PyTorch-native implementation equivalent to forward()."""
return x * torch.sigmoid(1.702 * x)
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
self.op(out, x)
return out
_ACTIVATION_REGISTRY = {
"gelu": nn.GELU,
+4
View File
@@ -50,6 +50,10 @@ class CustomOp(nn.Module):
return self.forward_native(*args, **kwargs)
def dispatch_forward(self) -> Callable:
# FIXME(will): for now, we always use the native implementation, since
# forward_cuda is using vllm's custom ops and it doesn't support
# backwards. We should add our own custom ops that support backwards.
return self.forward_native
# NOTE(woosuk): Here we assume that vLLM was built for only one
# specific backend. Currently, we do not support dynamic dispatching.
enabled = self.enabled()
+2 -29
View File
@@ -75,33 +75,6 @@ class RMSNorm(CustomOp):
else:
return x, residual
def forward_cuda(
self,
x: torch.Tensor,
residual: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
if self.variance_size_override is not None:
return self.forward_native(x, residual)
from vllm import _custom_ops as ops
if residual is not None:
ops.fused_add_rms_norm(
x,
residual,
self.weight.data,
self.variance_epsilon,
)
return x, residual
out = torch.empty_like(x)
ops.rms_norm(
out,
x,
self.weight.data,
self.variance_epsilon,
)
return out
def extra_repr(self) -> str:
s = f"hidden_size={self.weight.data.size(0)}"
s += f", eps={self.variance_epsilon}"
@@ -173,7 +146,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
# Apply normalization
normalized = self.norm(residual_output)
# Apply scale and shift
modulated = normalized * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
modulated = normalized * (1.0 + scale) + shift
return modulated, residual_output
@@ -209,4 +182,4 @@ class LayerNormScaleShift(nn.Module):
scale: torch.Tensor) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
normalized = self.norm(x)
return normalized * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
return normalized * (1.0 + scale) + shift
+17 -17
View File
@@ -7,16 +7,14 @@ from typing import Optional, Union
import torch
import torch.nn.functional as F
from torch.nn.parameter import Parameter
# TODO(will): remove this import by copying the definition from vLLM then
# manually import each quantization method we want to use. Refer to SGLang
from vllm.model_executor.layers.quantization.base_config import (
QuantizationConfig, QuantizeMethodBase)
from fastvideo.v1.distributed import (divide, get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
split_tensor_along_last_dim,
tensor_model_parallel_all_gather,
tensor_model_parallel_all_reduce)
from fastvideo.v1.layers.quantization.base_config import (QuantizationConfig,
QuantizeMethodBase)
from fastvideo.v1.logger import init_logger
# yapf: disable
from fastvideo.v1.models.parameter import (BasevLLMParameter,
@@ -545,19 +543,21 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
tp_size = get_tensor_model_parallel_world_size()
if isinstance(param, BlockQuantScaleParameter):
from vllm.model_executor.layers.quantization.fp8 import (
Fp8LinearMethod, Fp8MoEMethod)
assert self.quant_method is not None
assert isinstance(self.quant_method,
(Fp8LinearMethod, Fp8MoEMethod))
weight_block_size = self.quant_method.quant_config.weight_block_size
assert weight_block_size is not None
block_n, _ = weight_block_size[0], weight_block_size[1]
shard_offset = (
(sum(self.output_sizes[:loaded_shard_id]) + block_n - 1) //
block_n) // tp_size
shard_size = ((self.output_sizes[loaded_shard_id] + block_n - 1) //
block_n // tp_size)
raise NotImplementedError("FP8 is not implemented yet")
# FIXME(will): add fp8 support
# from vllm.model_executor.layers.quantization.fp8 import (
# Fp8LinearMethod, Fp8MoEMethod)
# assert self.quant_method is not None
# assert isinstance(self.quant_method,
# (Fp8LinearMethod, Fp8MoEMethod))
# weight_block_size = self.quant_method.quant_config.weight_block_size
# assert weight_block_size is not None
# block_n, _ = weight_block_size[0], weight_block_size[1]
# shard_offset = (
# (sum(self.output_sizes[:loaded_shard_id]) + block_n - 1) //
# block_n) // tp_size
# shard_size = ((self.output_sizes[loaded_shard_id] + block_n - 1) //
# block_n // tp_size)
else:
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // tp_size
shard_size = self.output_sizes[loaded_shard_id] // tp_size
@@ -0,0 +1,65 @@
from typing import Literal, get_args
from fastvideo.v1.layers.quantization.base_config import QuantizationConfig
QuantizationMethods = Literal[None]
QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods))
# The customized quantization methods which will be added to this dict.
_CUSTOMIZED_METHOD_TO_QUANT_CONFIG = {}
def register_quantization_config(quantization: str):
"""Register a customized vllm quantization config.
When a quantization method is not supported by vllm, you can register a customized
quantization config to support it.
Args:
quantization (str): The quantization method name.
Examples:
>>> from fastvideo.v1.layers.quantization import register_quantization_config
>>> from fastvideo.v1.layers.quantization import get_quantization_config
>>> from fastvideo.v1.layers.quantization.base_config import QuantizationConfig
>>>
>>> @register_quantization_config("my_quant")
... class MyQuantConfig(QuantizationConfig):
... pass
>>>
>>> get_quantization_config("my_quant")
<class 'MyQuantConfig'>
""" # noqa: E501
def _wrapper(quant_config_cls):
if quantization in QUANTIZATION_METHODS:
raise ValueError(
f"The quantization method `{quantization}` is already exists.")
if not issubclass(quant_config_cls, QuantizationConfig):
raise ValueError("The quantization config must be a subclass of "
"`QuantizationConfig`.")
_CUSTOMIZED_METHOD_TO_QUANT_CONFIG[quantization] = quant_config_cls
QUANTIZATION_METHODS.append(quantization)
return quant_config_cls
return _wrapper
def get_quantization_config(quantization: str) -> type[QuantizationConfig]:
if quantization not in QUANTIZATION_METHODS:
raise ValueError(f"Invalid quantization method: {quantization}")
method_to_config: dict[str, type[QuantizationConfig]] = {}
# Update the `method_to_config` with customized quantization methods.
method_to_config.update(_CUSTOMIZED_METHOD_TO_QUANT_CONFIG)
return method_to_config[quantization]
all = [
"QuantizationMethods",
"QuantizationConfig",
"get_quantization_config",
"QUANTIZATION_METHODS",
]
@@ -0,0 +1,151 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/quantization/base_config.py
import inspect
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Optional
import torch
from torch import nn
if TYPE_CHECKING:
from fastvideo.v1.layers.quantization import QuantizationMethods
else:
QuantizationMethods = str
class QuantizeMethodBase(ABC):
"""Base class for different quantized methods."""
@abstractmethod
def create_weights(self, layer: torch.nn.Module, *weight_args,
**extra_weight_attrs):
"""Create weights for a layer.
The weights will be set as attributes of the layer."""
raise NotImplementedError
@abstractmethod
def apply(self, layer: torch.nn.Module, *args, **kwargs) -> torch.Tensor:
"""Apply the weights in layer to the input tensor.
Expects create_weights to have been called before on the layer."""
raise NotImplementedError
# Not required functions
def embedding(self, layer: torch.nn.Module, *args,
**kwargs) -> torch.Tensor:
"""Gather embeddings in the layer based on indices in the input tensor.
Expects create_weights to have been called before on the layer."""
raise NotImplementedError
def process_weights_after_loading(self, layer: nn.Module) -> None:
"""Process the weight after loading.
This can be used for example, to transpose weights for computation.
"""
return
def method_has_implemented_embedding(
method_class: type[QuantizeMethodBase]) -> bool:
"""
Not all quant methods have embedding implemented, so we need to check that
it exists for our given method. We check this by making sure the function
has been changed from the base implementation.
"""
base_embedding = inspect.getattr_static(QuantizeMethodBase, "embedding",
None)
class_embedding = inspect.getattr_static(method_class, "embedding", None)
return (class_embedding is not None
and class_embedding is not base_embedding)
class QuantizationConfig(ABC):
"""Base class for quantization configs."""
def __init__(self):
super().__init__()
# mapping is updated by models as they initialize
self.packed_modules_mapping: dict[str, list[str]] = dict()
@abstractmethod
def get_name(self) -> QuantizationMethods:
"""Name of the quantization method."""
raise NotImplementedError
@abstractmethod
def get_supported_act_dtypes(self) -> list[torch.dtype]:
"""List of supported activation dtypes."""
raise NotImplementedError
@classmethod
@abstractmethod
def get_min_capability(cls) -> int:
"""Minimum GPU capability to support the quantization method.
E.g., 70 for Volta, 75 for Turing, 80 for Ampere.
This requirement is due to the custom CUDA kernels used by the
quantization method.
"""
raise NotImplementedError
@staticmethod
@abstractmethod
def get_config_filenames() -> list[str]:
"""List of filenames to search for in the model directory."""
raise NotImplementedError
@classmethod
@abstractmethod
def from_config(cls, config: dict[str, Any]) -> "QuantizationConfig":
"""Create a config class from the model's quantization config."""
raise NotImplementedError
@classmethod
def override_quantization_method(
cls, hf_quant_cfg, user_quant) -> Optional[QuantizationMethods]:
"""
Detects if this quantization method can support a given checkpoint
format by overriding the user specified quantization method --
this method should only be overwritten by subclasses in exceptional
circumstances
"""
return None
@staticmethod
def get_from_keys(config: dict[str, Any], keys: list[str]) -> Any:
"""Get a value from the model's quantization config."""
for key in keys:
if key in config:
return config[key]
raise ValueError(f"Cannot find any of {keys} in the model's "
"quantization config.")
@staticmethod
def get_from_keys_or(config: dict[str, Any], keys: list[str],
default: Any) -> Any:
"""Get a optional value from the model's quantization config."""
try:
return QuantizationConfig.get_from_keys(config, keys)
except ValueError:
return default
@abstractmethod
def get_quant_method(self, layer: torch.nn.Module,
prefix: str) -> Optional[QuantizeMethodBase]:
"""Get the quantize method to use for the quantized layer.
Args:
layer: The layer for the quant method.
prefix: The full name of the layer in the state dict
Returns:
The quantize method. None if the given layer doesn't support quant
method.
"""
raise NotImplementedError
def get_cache_scale(self, name: str) -> Optional[str]:
return None
-22
View File
@@ -152,28 +152,6 @@ class RotaryEmbedding(CustomOp):
key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape)
return query, key
def forward_cuda(
self,
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
offsets: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
from vllm import _custom_ops as ops
self.cos_sin_cache = self.cos_sin_cache.to(query.device,
dtype=query.dtype)
# ops.rotary_embedding()/batched_rotary_embedding()
# are in-place operations that update the query and key tensors.
if offsets is not None:
ops.batched_rotary_embedding(positions, query, key, self.head_size,
self.cos_sin_cache, self.is_neox_style,
self.rotary_dim, offsets)
else:
ops.rotary_embedding(positions, query, key, self.head_size,
self.cos_sin_cache, self.is_neox_style)
return query, key
def extra_repr(self) -> str:
s = f"head_size={self.head_size}, rotary_dim={self.rotary_dim}"
s += f", max_position_embeddings={self.max_position_embeddings}"
@@ -6,12 +6,12 @@ from typing import List, Optional, Sequence, Tuple
import torch
import torch.nn.functional as F
from torch.nn.parameter import Parameter, UninitializedParameter
from vllm.model_executor.layers.quantization.base_config import (
QuantizationConfig, QuantizeMethodBase, method_has_implemented_embedding)
from fastvideo.v1.distributed import (divide, get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
tensor_model_parallel_all_reduce)
from fastvideo.v1.layers.quantization.base_config import (
QuantizationConfig, QuantizeMethodBase, method_has_implemented_embedding)
from fastvideo.v1.models.parameter import BasevLLMParameter
from fastvideo.v1.models.utils import set_weight_attrs
from fastvideo.v1.platforms import current_platform
+6 -4
View File
@@ -33,9 +33,11 @@ class BaseDiT(nn.Module, ABC):
f"Subclasses of BaseDiT must define '{attr}' class variable"
)
def __init__(self, config: DiTConfig, **kwargs) -> None:
def __init__(self, config: DiTConfig, hf_config: dict[str, Any],
**kwargs) -> None:
super().__init__()
self.config = config
self.hf_config = hf_config
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
@@ -95,8 +97,6 @@ class CachableDiT(BaseDiT):
if self.config.prefix == "wan":
self.use_ret_steps = self.config.cache_config.use_ret_steps
self.is_even = False
self.previous_e0_even: torch.Tensor | None = None
self.previous_e0_odd: torch.Tensor | None = None
self.previous_residual_even: torch.Tensor | None = None
self.previous_residual_odd: torch.Tensor | None = None
self.accumulated_rel_l1_distance_even = 0
@@ -106,7 +106,9 @@ class CachableDiT(BaseDiT):
else:
self.accumulated_rel_l1_distance = 0
self.previous_modulated_input = None
self.previous_residual = None
self.previous_resiual = None
self.previous_e0_even: torch.Tensor | None = None
self.previous_e0_odd: torch.Tensor | None = None
def maybe_cache_states(self, hidden_states: torch.Tensor,
original_hidden_states: torch.Tensor) -> None:
+6 -6
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from typing import List, Optional, Tuple, Union
from typing import Any, List, Optional, Tuple, Union
import numpy as np
import torch
@@ -442,8 +442,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
)._supported_attention_backends
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
def __init__(self, config: HunyuanVideoConfig):
super().__init__(config=config)
def __init__(self, config: HunyuanVideoConfig, hf_config: dict[str, Any]):
super().__init__(config=config, hf_config=hf_config)
self.patch_size = [
config.patch_size_t, config.patch_size, config.patch_size
@@ -562,8 +562,7 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
"""
forward_context = get_forward_context()
forward_batch = forward_context.forward_batch
assert forward_batch is not None
enable_teacache = forward_batch.enable_teacache
enable_teacache = forward_batch is not None and forward_batch.enable_teacache
if guidance is None:
guidance = torch.tensor([6016.0],
@@ -661,7 +660,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
forward_context = get_forward_context()
forward_batch = forward_context.forward_batch
assert forward_batch is not None
if forward_batch is None:
return False
current_timestep = forward_context.current_timestep
enable_teacache = forward_batch.enable_teacache
+4 -3
View File
@@ -10,7 +10,7 @@
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
from typing import Dict, Optional, Tuple
from typing import Any, Dict, Optional, Tuple
import torch
from einops import rearrange, repeat
@@ -462,8 +462,9 @@ class StepVideoModel(BaseDiT):
_supported_attention_backends = StepVideoConfig(
)._supported_attention_backends
def __init__(self, config: StepVideoConfig) -> None:
super().__init__(config=config)
def __init__(self, config: StepVideoConfig, hf_config: dict[str,
Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
self.num_attention_heads = config.num_attention_heads
self.attention_head_dim = config.attention_head_dim
self.in_channels = config.in_channels
+7 -8
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import List, Optional, Tuple, Union
from typing import Any, List, Optional, Tuple, Union
import numpy as np
import torch
@@ -298,7 +298,7 @@ class WanTransformerBlock(nn.Module):
hidden_states = hidden_states.squeeze(1)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
assert orig_dtype != torch.float32
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
@@ -360,8 +360,9 @@ class WanTransformer3DModel(CachableDiT):
)._supported_attention_backends
_param_names_mapping = WanVideoConfig()._param_names_mapping
def __init__(self, config: WanVideoConfig) -> None:
super().__init__(config=config)
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
@@ -424,8 +425,7 @@ class WanTransformer3DModel(CachableDiT):
guidance=None,
**kwargs) -> torch.Tensor:
forward_batch = get_forward_context().forward_batch
assert forward_batch is not None
enable_teacache = forward_batch.enable_teacache
enable_teacache = forward_batch is not None and forward_batch.enable_teacache
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
@@ -525,8 +525,7 @@ class WanTransformer3DModel(CachableDiT):
forward_context = get_forward_context()
forward_batch = forward_context.forward_batch
assert forward_batch is not None
if not forward_batch.enable_teacache:
if forward_batch is None or not forward_batch.enable_teacache:
return False
teacache_params = forward_batch.teacache_params
assert teacache_params is not None, "teacache_params is not initialized"
+1 -1
View File
@@ -13,12 +13,12 @@ from fastvideo.v1.attention import LocalAttention
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
CLIPTextConfig,
CLIPVisionConfig)
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.distributed import (divide,
get_tensor_model_parallel_world_size)
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
RowParallelLinear)
from fastvideo.v1.layers.quantization import QuantizationConfig
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.encoders.base import ImageEncoder, TextEncoder
from fastvideo.v1.models.encoders.vision import resolve_visual_encoder_outputs
+1 -1
View File
@@ -32,12 +32,12 @@ from torch import nn
from fastvideo.v1.attention import LocalAttention
# from ..utils import (extract_layer_index)
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput, LlamaConfig
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.distributed import get_tensor_model_parallel_world_size
from fastvideo.v1.layers.activation import SiluAndMul
from fastvideo.v1.layers.layernorm import RMSNorm
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
QKVParallelLinear, RowParallelLinear)
from fastvideo.v1.layers.quantization import QuantizationConfig
from fastvideo.v1.layers.rotary_embedding import get_rope
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.v1.models.encoders.base import TextEncoder
+1 -1
View File
@@ -28,13 +28,13 @@ import torch.nn.functional as F
from torch import nn
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput, T5Config
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.distributed import (get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size)
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.layers.layernorm import RMSNorm
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
QKVParallelLinear, RowParallelLinear)
from fastvideo.v1.layers.quantization import QuantizationConfig
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.v1.models.encoders.base import TextEncoder
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
+22 -3
View File
@@ -6,6 +6,7 @@ import json
import os
import time
from abc import ABC, abstractmethod
from copy import deepcopy
from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
import torch
@@ -366,6 +367,7 @@ class TransformerLoader(ComponentLoader):
fastvideo_args: FastVideoArgs):
"""Load the transformer based on the model path, architecture, and inference args."""
config = get_diffusers_config(model=model_path)
hf_config = deepcopy(config)
cls_name = config.pop("_class_name")
if cls_name is None:
raise ValueError(
@@ -392,13 +394,30 @@ class TransformerLoader(ComponentLoader):
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
# Load the model using FSDP loader
logger.info("Loading model from %s", cls_name)
logger.info("Loading model from %s, default_dtype: %s", cls_name, default_dtype)
# model = load_fsdp_model(model_cls=model_cls,
# init_params={
# "config": dit_config,
# "hf_config": hf_config
# },
# weight_dir_list=safetensors_list,
# device=fastvideo_args.device,
# cpu_offload=fastvideo_args.use_cpu_offload,
# default_dtype=default_dtype)
model = load_fsdp_model(model_cls=model_cls,
init_params={"config": dit_config},
init_params={
"config": dit_config,
"hf_config": hf_config
},
weight_dir_list=safetensors_list,
device=fastvideo_args.device,
cpu_offload=fastvideo_args.use_cpu_offload,
default_dtype=default_dtype)
default_dtype=default_dtype,
# TODO(will): make these configurable
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
)
if fastvideo_args.enable_torch_compile:
logger.info("Torch Compile enabled for DiT")
for n, m in reversed(list(model.named_modules())):
+30 -4
View File
@@ -14,13 +14,16 @@ from typing import (Any, Callable, DefaultDict, Dict, Generator, Hashable, List,
import torch
from torch import nn
from torch.distributed import DeviceMesh, init_device_mesh
from torch.distributed._composable.fsdp import CPUOffloadPolicy, fully_shard
from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard, MixedPrecisionPolicy
from torch.distributed._tensor import distribute_tensor
from torch.nn.modules.module import _IncompatibleKeys
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
# TODO(PY): move this to utils elsewhere
@@ -86,16 +89,29 @@ def get_param_names_mapping(
# TODO(PY): add compile option
# param_dtype: torch.dtype,
# reduce_dtype: torch.dtype,
# output_dtype: torch.dtype,
# pp_enabled: bool = False,
# cpu_offload: bool = False,
def load_fsdp_model(
model_cls: Type[nn.Module],
init_params: Dict[str, Any],
weight_dir_list: List[str],
device: torch.device,
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
cpu_offload: bool = False,
default_dtype: Optional[torch.dtype] = torch.bfloat16,
output_dtype: Optional[torch.dtype] = None,
) -> torch.nn.Module:
mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=True)
# with set_default_dtype(default_dtype), torch.device("meta"):
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
device_mesh = init_device_mesh(
"cuda",
mesh_shape=(get_sequence_model_parallel_world_size(), ),
@@ -104,6 +120,7 @@ def load_fsdp_model(
shard_model(model,
cpu_offload=cpu_offload,
reshard_after_forward=True,
mp_policy=mp_policy,
dp_mesh=device_mesh["dp"])
weight_iterator = safetensors_weights_iterator(weight_dir_list)
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
@@ -121,6 +138,7 @@ def load_fsdp_model(
f"Unexpected param or buffer {n} on meta device.")
for p in model.parameters():
p.requires_grad = False
# set_state_dict(model, StateDictType.LOCAL_STATE_DICT)
return model
@@ -129,6 +147,7 @@ def shard_model(
*,
cpu_offload: bool,
reshard_after_forward: bool = True,
mp_policy: Optional[MixedPrecisionPolicy] = None,
dp_mesh: Optional[DeviceMesh] = None,
) -> None:
"""
@@ -156,14 +175,17 @@ def shard_model(
"""
fsdp_kwargs = {
"reshard_after_forward": reshard_after_forward,
"mesh": dp_mesh
"mesh": dp_mesh,
"mp_policy": mp_policy,
}
if cpu_offload:
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
# Shard the model with FSDP, iterating in reverse to start with
# iterating in reverse to start with
# lowest-level modules first
num_layers_sharded = 0
# TODO(will): don't reshard after forward for the last layer to save on the
# all-gather that will immediately happen Shard the model with FSDP,
for n, m in reversed(list(model.named_modules())):
if any([
shard_condition(n, m)
@@ -210,6 +232,10 @@ def load_fsdp_model_from_full_model_state_dict(
NotImplementedError: If got FSDP with more than 1D.
"""
meta_sharded_sd = model.state_dict()
# s = fully_shard.state(model)
# logger.info(f"type(s): {type(s)}")
# logger.info(f"s: {s}")
# import pdb; pdb.set_trace()
sharded_sd = {}
to_merge_params: DefaultDict[Hashable, Dict[Any, Any]] = defaultdict(dict)
+158 -14
View File
@@ -5,19 +5,25 @@ Base class for composed pipelines.
This module defines the base class for pipelines that are composed of multiple stages.
"""
import argparse
import os
from abc import ABC, abstractmethod
from copy import deepcopy
from typing import Any, Dict, List, Optional, cast
from typing import Any, Dict, List, Optional, Union, cast
import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.pipelines import (PipelineConfig,
get_pipeline_config_cls_for_name)
from fastvideo.v1.distributed import (init_distributed_environment,
initialize_model_parallel,
model_parallel_is_initialized)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import PipelineStage
from fastvideo.v1.utils import (maybe_download_model,
from fastvideo.v1.utils import (maybe_download_model, shallow_asdict,
verify_model_config_and_directory)
logger = init_logger(__name__)
@@ -39,15 +45,20 @@ class ComposedPipelineBase(ABC):
def __init__(self,
model_path: str,
fastvideo_args: FastVideoArgs,
config: Optional[Dict[str, Any]] = None):
config: Optional[Dict[str, Any]] = None,
required_config_modules: Optional[List[str]] = None):
"""
Initialize the pipeline. After __init__, the pipeline should be ready to
use. The pipeline should be stateless and not hold any batch state.
"""
self.fastvideo_args = fastvideo_args
self.model_path = model_path
self._stages: List[PipelineStage] = []
self._stage_name_mapping: Dict[str, PipelineStage] = {}
if required_config_modules is not None:
self._required_config_modules = required_config_modules
if self._required_config_modules is None:
raise NotImplementedError(
"Subclass must set _required_config_modules")
@@ -59,16 +70,135 @@ class ComposedPipelineBase(ABC):
else:
self.config = config
self.maybe_init_distributed_environment(fastvideo_args)
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
self.modules = self.load_modules(fastvideo_args)
if fastvideo_args.training_mode:
if fastvideo_args.log_validation:
self.initialize_validation_pipeline(fastvideo_args)
self.initialize_training_pipeline(fastvideo_args)
self.initialize_pipeline(fastvideo_args)
logger.info("Creating pipeline stages...")
self.create_pipeline_stages(fastvideo_args)
# logger.info("Creating pipeline stages...")
# self.create_pipeline_stages(fastvideo_args)
def get_module(self, module_name: str) -> Any:
if fastvideo_args.training_mode:
logger.info("Creating training pipeline stages...")
self.create_training_stages(fastvideo_args)
else:
logger.info("Creating pipeline stages...")
self.create_pipeline_stages(fastvideo_args)
def initialize_training_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"if training_mode is True, the pipeline must implement this method")
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"if log_validation is True, the pipeline must implement this method"
)
@classmethod
def from_pretrained(cls,
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[
Union[str
| PipelineConfig]] = None,
args: Optional[argparse.Namespace] = None,
required_config_modules: Optional[List[str]] = None,
**kwargs) -> "ComposedPipelineBase":
config = None
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, PipelineConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = kwargs
else:
config_args = shallow_asdict(config)
config_args.update(kwargs)
if args.inference_mode:
fastvideo_args = FastVideoArgs(model_path=model_path,
device_str=device or "cuda" if
torch.cuda.is_available() else "cpu",
**config_args)
fastvideo_args.model_path = model_path
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
) else "cpu"
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
else:
assert args is not None, "args must be provided for training mode"
fastvideo_args = TrainingArgs.from_cli_args(args)
# TODO(will): fix this so that its not so ugly
fastvideo_args.model_path = model_path
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
) else "cpu"
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
# we use cpu offload for training
fastvideo_args.use_cpu_offload = False
# make sure we are in training mode
fastvideo_args.inference_mode = False
# we hijack the precision to be the master weight type so that the
# 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.precision = fastvideo_args.master_weight_type
assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
fastvideo_args.check_fastvideo_args()
logger.info(f"fastvideo_args in from_pretrained: {fastvideo_args}")
return cls(model_path,
fastvideo_args,
required_config_modules=required_config_modules)
def maybe_init_distributed_environment(self, fastvideo_args: FastVideoArgs):
if model_parallel_is_initialized():
return
local_rank = int(os.environ.get("LOCAL_RANK", -1))
world_size = int(os.environ.get("WORLD_SIZE", -1))
rank = int(os.environ.get("RANK", -1))
if local_rank == -1 or world_size == -1 or rank == -1:
raise ValueError(
"Local rank, world size, and rank must be set. Use torchrun to launch the script."
)
torch.cuda.set_device(local_rank)
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
initialize_model_parallel(
tensor_model_parallel_size=fastvideo_args.tp_size,
sequence_model_parallel_size=fastvideo_args.sp_size)
device = torch.device(f"cuda:{local_rank}")
fastvideo_args.device = device
def get_module(self, module_name: str, default_value: Any = None) -> Any:
if module_name not in self.modules:
return default_value
return self.modules[module_name]
def add_module(self, module_name: str, module: Any):
@@ -114,6 +244,19 @@ class ComposedPipelineBase(ABC):
"""
raise NotImplementedError
# @abstractmethod
# def create_validation_stages(self, fastvideo_args: FastVideoArgs):
# """
# Create the validation pipeline stages.
# """
# raise NotImplementedError
def create_training_stages(self, fastvideo_args: FastVideoArgs):
"""
Create the training pipeline stages.
"""
raise NotImplementedError
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
@@ -136,19 +279,21 @@ class ComposedPipelineBase(ABC):
modules_config
) > 1, "model_index.json must contain at least one pipeline module"
required_modules = [
"vae", "text_encoder", "transformer", "scheduler", "tokenizer"
]
for module_name in required_modules:
for module_name in self.required_config_modules:
if module_name not in modules_config:
raise ValueError(
f"model_index.json must contain a {module_name} module")
logger.info("Diffusers config passed sanity checks")
# all the component models used by the pipeline
required_modules = self.required_config_modules
logger.info("Loading required modules: %s", required_modules)
modules = {}
for module_name, (transformers_or_diffusers,
architecture) in modules_config.items():
if module_name not in required_modules:
logger.info("Skipping module %s", module_name)
continue
component_model_path = os.path.join(self.model_path, module_name)
module = PipelineComponentLoader.load_module(
module_name=module_name,
@@ -164,7 +309,6 @@ class ComposedPipelineBase(ABC):
logger.warning("Overwriting module %s", module_name)
modules[module_name] = module
required_modules = self.required_config_modules
# Check if all required modules were loaded
for module_name in required_modules:
if module_name not in modules or modules[module_name] is None:
@@ -198,7 +342,7 @@ class ComposedPipelineBase(ABC):
# Execute each stage
logger.info("Running pipeline stages: %s",
self._stage_name_mapping.keys())
logger.info("Batch: %s", batch)
# logger.info("Batch: %s", batch)
for stage in self.stages:
batch = stage(batch, fastvideo_args)
+303
View File
@@ -0,0 +1,303 @@
import os
import sys
import time
from collections import deque
import torch
from tqdm.auto import tqdm
# import torch.distributed as dist
import wandb
from fastvideo.utils.checkpoint import save_checkpoint
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper
from fastvideo.utils.validation import log_validation
from fastvideo.v1.distributed import cleanup_dist_env_and_memory, get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
logger = init_logger(__name__)
class WanTrainingPipeline(ComposedPipelineBase): # == distill_one_step
_required_config_modules = ["scheduler", "transformer"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
):
device = fastvideo_args.device
local_rank = int(os.environ.get("LOCAL_RANK", -1))
rank = int(os.environ.get("RANK", -1))
assert rank != -1
assert local_rank != -1
sp_group = get_sp_group()
world_size = sp_group.world_size
rank = sp_group.rank
args = fastvideo_args
transformer = self.get_module("transformer")
teacher_transformer = self.get_module("teacher_transformer")
ema_transformer = None
assert not fastvideo_args.use_ema, "ema is not supported now"
assert teacher_transformer is not None
assert transformer is not None
train_dataset = self.train_dataset
train_dataloader = self.train_dataloader
init_steps = self.init_steps
lr_scheduler = self.lr_scheduler
optimizer = self.optimizer
noise_scheduler = self.noise_scheduler
solver = self.solver
noise_random_generator = None
uncond_prompt_embed = self.uncond_prompt_embed
uncond_prompt_mask = self.uncond_prompt_mask
# Train!
total_batch_size = (world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
logger.info("***** Running training *****")
logger.info(" Num examples = %s", len(train_dataset))
logger.info(" Dataloader size = %s", len(train_dataloader))
logger.info(" Num Epochs = %s", args.num_train_epochs)
logger.info(" Resume training from step %s", init_steps)
logger.info(" Instantaneous batch size per device = %s",
args.train_batch_size)
logger.info(
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
total_batch_size)
logger.info(" Gradient Accumulation steps = %s",
args.gradient_accumulation_steps)
logger.info(" Total optimization steps = %s", args.max_train_steps)
logger.info(
" Total training parameters per FSDP shard = %s B",
sum(p.numel()
for p in transformer.parameters() if p.requires_grad) / 1e9)
# print dtype
logger.info(" Master weight dtype: %s",
transformer.parameters().__next__().dtype)
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
# loader = self.get_module("train_dataloader")
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader)
# log_validation(args, transformer, device,
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
def get_num_phases(multi_phased_distill_schedule, step):
# step-phase,step-phase
multi_phases = multi_phased_distill_schedule.split(",")
phase = multi_phases[-1].split("-")[-1]
for step_phases in multi_phases:
phase_step, phase = step_phases.split("-")
if step <= int(phase_step):
return int(phase)
return phase
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
assert args.multi_phased_distill_schedule is not None
num_phases = get_num_phases(args.multi_phased_distill_schedule,
step)
loss, grad_norm, pred_norm = self.distill_one_step(
transformer,
args.model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
args.num_euler_timesteps,
num_phases,
args.not_apply_cfg_solver,
args.distill_cfg,
args.ema_decay,
args.pred_decay_weight,
args.pred_decay_type,
args.hunyuan_teacher_disable_cfg,
)
step_time = time.time() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"phases": num_phases,
})
progress_bar.update(1)
if rank <= 0:
wandb.log(
{
"train_loss":
loss,
"learning_rate":
lr_scheduler.get_last_lr()[0],
"step_time":
step_time,
"avg_step_time":
avg_step_time,
"grad_norm":
grad_norm,
"pred_fro_norm":
pred_norm["fro"], # codespell:ignore
"pred_largest_singular_value":
pred_norm["largest singular value"],
"pred_absolute_mean":
pred_norm["absolute mean"],
"pred_absolute_max":
pred_norm["absolute max"],
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
raise NotImplementedError("lora is not supported now")
# save_lora_checkpoint(transformer, optimizer, rank,
# args.output_dir, step)
else:
# Your existing checkpoint saving code
if args.use_ema:
raise NotImplementedError("ema is not supported now")
save_checkpoint(ema_transformer, rank, args.output_dir,
step)
else:
save_checkpoint(transformer, rank, args.output_dir,
step)
sp_group.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(
args,
transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=False,
)
if args.use_ema:
log_validation(
args,
ema_transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.
linear_quadratic_threshold,
linear_range=args.linear_range,
ema=True,
)
if args.use_lora:
raise NotImplementedError("lora is not supported now")
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
args.max_train_steps)
else:
save_checkpoint(transformer, rank, args.output_dir,
args.max_train_steps)
if get_sp_group():
cleanup_dist_env_and_memory()
def main(args):
logger.info("Starting training pipeline...")
pipeline = WanTrainingPipeline.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers", args=args)
args = pipeline.fastvideo_args
pipeline.forward(None, args)
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
print(args)
main(args)
@@ -0,0 +1,567 @@
# SPDX-License-Identifier: Apache-2.0
"""
T2V Data Preprocessing pipeline implementation.
This module contains an implementation of the T2V Data Preprocessing pipeline
using the modular pipeline architecture.
"""
import gc
import multiprocessing
import os
from concurrent.futures import ProcessPoolExecutor
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.v1.dataset import getdataset
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import TextEncodingStage
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class PreprocessPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
args,
):
# Initialize class variables for data sharing
self.video_data = {} # Store video metadata and paths
self.latent_data = {} # Store latent tensors
self.preprocess_validation_text(fastvideo_args, args)
self.preprocess_video_and_text(fastvideo_args, args)
def preprocess_video_and_text(self, fastvideo_args: FastVideoArgs, args):
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(combined_parquet_dir, exist_ok=True)
local_rank = int(os.getenv("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
# Get how many samples have already been processed
start_idx = 0
for root, _, files in os.walk(combined_parquet_dir):
for file in files:
if file.endswith('.parquet'):
table = pq.read_table(os.path.join(root, file))
start_idx += table.num_rows
# Loading dataset
train_dataset = getdataset(args, start_idx=start_idx)
sampler = DistributedSampler(train_dataset,
rank=local_rank,
num_replicas=world_size,
shuffle=False)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
num_processed_samples = 0
# Add progress bar for video preprocessing
pbar = tqdm(train_dataloader,
desc="Processing videos",
unit="batch",
disable=local_rank != 0)
for batch_idx, data in enumerate(pbar):
if data is None:
continue
with torch.inference_mode():
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, pixel_values in enumerate(data["pixel_values"]):
if not torch.all(
pixel_values == 0): # Check if all values are zero
valid_indices.append(i)
num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples
valid_data = {
"pixel_values":
torch.stack(
[data["pixel_values"][i] for i in valid_indices]),
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
"fps": [data["fps"][i] for i in valid_indices],
"duration": [data["duration"][i] for i in valid_indices],
}
# VAE
with torch.autocast("cuda", dtype=torch.float32):
latents = self.get_module("vae").encode(
valid_data["pixel_values"].to(
fastvideo_args.device)).mean
# Get corresponding captions for this batch
batch_captions = valid_data["text"]
batch = ForwardBatch(
data_type="video",
prompt=batch_captions,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
prompt_embeds, prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
assert prompt_embeds.shape[0] == prompt_attention_mask.shape[0]
# Remove padding from prompt_embeds using attention mask for all batches
# Get sequence lengths from attention masks (number of 1s)
seq_lens = prompt_attention_mask.sum(dim=1)
# Create a list to store non-padded embeddings and masks
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_mask = non_padded_masks
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
# Get the corresponding latent and info using video name
latent = latents[idx].cpu()
video_name = os.path.basename(video_path).split(".")[0]
height, width = valid_data["pixel_values"][idx].shape[-2:]
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
).astype(np.uint8)
# Create record for Parquet dataset
record = {
"id": video_name,
"vae_latent_bytes": vae_latent.tobytes(),
"vae_latent_shape": list(vae_latent.shape),
"vae_latent_dtype": str(vae_latent.dtype),
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape":
list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx],
"media_type": "video",
"width": width,
"height": height,
"num_frames": latents[idx].shape[1],
"duration_sec": float(valid_data["duration"][idx]),
"fps": float(valid_data["fps"][idx]),
}
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 = [
pa.array([record["id"] for record in batch_data]),
pa.array(
[record["vae_latent_bytes"] for record in batch_data],
type=pa.binary()),
pa.array(
[record["vae_latent_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array(
[record["vae_latent_dtype"] for record in batch_data]),
pa.array([
record["text_embedding_bytes"] for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_embedding_shape"] for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_embedding_dtype"] for record in batch_data
]),
pa.array([
record["text_attention_mask_bytes"]
for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_attention_mask_shape"]
for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_attention_mask_dtype"]
for record in batch_data
]),
pa.array([record["file_name"] for record in batch_data]),
pa.array([record["caption"] for record in batch_data]),
pa.array([record["media_type"] for record in batch_data]),
pa.array([record["width"] for record in batch_data],
type=pa.int32()),
pa.array([record["height"] for record in batch_data],
type=pa.int32()),
pa.array([record["num_frames"] for record in batch_data],
type=pa.int32()),
pa.array([record["duration_sec"] for record in batch_data],
type=pa.float32()),
pa.array([record["fps"] for record in batch_data],
type=pa.float32()),
]
table = pa.Table.from_arrays(
arrays, names=[f.name for f in pyarrow_schema])
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info(f"Collected batch with {len(table)} samples")
if num_processed_samples >= args.flush_frequency:
assert hasattr(self, 'all_tables') and self.all_tables
print(f"Combining {len(self.all_tables)} batches...")
combined_table = pa.concat_tables(self.all_tables)
assert len(combined_table) == num_processed_samples
print(f"Total samples collected: {len(combined_table)}")
# Calculate total number of chunks needed, discarding remainder
total_chunks = max(
num_processed_samples // args.samples_per_file, 1)
print(
f"Fixed samples per parquet file: {args.samples_per_file}")
print(f"Total number of parquet files: {total_chunks}")
print(
f"Total samples to be processed: {total_chunks * args.samples_per_file} (discarding {num_processed_samples % args.samples_per_file} samples)"
)
# Split work among processes
num_workers = int(min(multiprocessing.cpu_count(),
total_chunks))
chunks_per_worker = (total_chunks + num_workers -
1) // num_workers
print(
f"Using {num_workers} workers to process {total_chunks} chunks"
)
logger.info(f"Chunks per worker: {chunks_per_worker}")
# Prepare work ranges
work_ranges = []
for i in range(num_workers):
start_idx = i * chunks_per_worker
end_idx = min((i + 1) * chunks_per_worker, total_chunks)
if start_idx < total_chunks:
work_ranges.append(
(start_idx, end_idx, combined_table, i,
combined_parquet_dir, args.samples_per_file))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
for work_range in work_ranges
}
for future in tqdm(futures, desc="Processing chunks"):
try:
written = future.result()
total_written += written
logger.info(
f"Processed chunk with {written} samples")
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]}: {str(e)}"
)
# Retry failed ranges sequentially
if failed_ranges:
logger.warning(
f"Retrying {len(failed_ranges)} failed ranges sequentially"
)
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(
work_range)
except Exception as e:
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]} after retry: {str(e)}"
)
logger.info(f"Total samples written: {total_written}")
num_processed_samples = 0
self.all_tables = []
def preprocess_validation_text(self, fastvideo_args: FastVideoArgs, args):
# Create Parquet dataset directory for validation
validation_parquet_dir = os.path.join(args.output_dir,
"validation_parquet_dataset")
os.makedirs(validation_parquet_dir, exist_ok=True)
# Initialize Parquet dataset
validation_parquet_path = os.path.join(validation_parquet_dir,
"data.parquet")
with open(args.validation_prompt_txt, encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for validation text preprocessing
pbar = tqdm(enumerate(prompts),
desc="Processing validation prompts",
unit="prompt")
for prompt_idx, prompt in pbar:
with torch.inference_mode():
# Text Encoder
batch = ForwardBatch(
data_type="video",
prompt=prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
prompt_embeds = result_batch.prompt_embeds[0]
prompt_attention_mask = result_batch.prompt_attention_mask[0]
file_name = prompt.split(".")[0]
# Get the sequence length from attention mask (number of 1s)
seq_len = prompt_attention_mask.sum().item()
# Slice the embeddings to keep only the non-padding parts
text_embedding = prompt_embeds[0, :seq_len].cpu().numpy()
text_attention_mask = prompt_attention_mask[
0, :seq_len].cpu().numpy().astype(np.uint8)
# Log the shapes after removing padding
logger.info(
f"Shape after removing padding - Embeddings: {text_embedding.shape}, Mask: {text_attention_mask.shape}"
)
# Create record for Parquet dataset
record = {
"id": file_name,
"vae_latent_bytes": b"", # Not available for validation
"vae_latent_shape": [],
"vae_latent_dtype": "",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"text_attention_mask_bytes": text_attention_mask.tobytes(),
"text_attention_mask_shape": list(text_attention_mask.shape),
"text_attention_mask_dtype": str(text_attention_mask.dtype),
"file_name": file_name,
"caption": prompt,
"media_type": "video",
"width": 0, # Not available for validation
"height": 0, # Not available for validation
"num_frames": 0, # Not available for validation
"duration_sec": 0.0, # Not available for validation
"fps": 0.0, # Not available for validation
}
batch_data.append(record)
logger.info(f"Saved validation sample: {file_name}")
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 = [
pa.array([record["id"] for record in batch_data]),
pa.array([record["vae_latent_bytes"] for record in batch_data],
type=pa.binary()),
pa.array([record["vae_latent_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array([record["vae_latent_dtype"] for record in batch_data]),
pa.array(
[record["text_embedding_bytes"] for record in batch_data],
type=pa.binary()),
pa.array(
[record["text_embedding_shape"] for record in batch_data],
type=pa.list_(pa.int32())),
pa.array(
[record["text_embedding_dtype"] for record in batch_data]),
pa.array([
record["text_attention_mask_bytes"] for record in batch_data
],
type=pa.binary()),
pa.array([
record["text_attention_mask_shape"] for record in batch_data
],
type=pa.list_(pa.int32())),
pa.array([
record["text_attention_mask_dtype"] for record in batch_data
]),
pa.array([record["file_name"] for record in batch_data]),
pa.array([record["caption"] for record in batch_data]),
pa.array([record["media_type"] for record in batch_data]),
pa.array([record["width"] for record in batch_data],
type=pa.int32()),
pa.array([record["height"] for record in batch_data],
type=pa.int32()),
pa.array([record["num_frames"] for record in batch_data],
type=pa.int32()),
pa.array([record["duration_sec"] for record in batch_data],
type=pa.float32()),
pa.array([record["fps"] for record in batch_data],
type=pa.float32()),
]
table = pa.Table.from_arrays(arrays,
names=[f.name for f in pyarrow_schema])
write_pbar.update(1)
write_pbar.close()
logger.info(f"Total validation samples: {len(table)}")
work_range = (0, 1, table, 0, validation_parquet_dir, len(table))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=1) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
}
for future in tqdm(futures, desc="Processing chunks"):
try:
total_written += future.result()
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]}: {str(e)}"
)
# Retry failed ranges sequentially
if failed_ranges:
logger.warning(
f"Retrying {len(failed_ranges)} failed ranges sequentially")
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(work_range)
except Exception as e:
logger.error(
f"Failed to process range {work_range[0]}-{work_range[1]} after retry: {str(e)}"
)
logger.info(f"Total validation samples written: {total_written}")
# Clear memory
del table
gc.collect() # Force garbage collection
@staticmethod
def process_chunk_range(args):
start_idx, end_idx, table, worker_id, output_dir, samples_per_file = args
try:
total_written = 0
num_samples = len(table)
# Create worker-specific subdirectory
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
os.makedirs(worker_dir, exist_ok=True)
# Check how many files there are already in the dir, and update i accordingly
num_parquets = 0
for root, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
for i in range(start_idx, end_idx):
start_sample = i * samples_per_file
end_sample = min((i + 1) * samples_per_file, num_samples)
chunk = table.slice(start_sample, end_sample - start_sample)
# Create chunk file in worker's directory
chunk_path = os.path.join(
worker_dir, f"data_chunk_{i + num_parquets}.parquet")
temp_path = chunk_path + '.tmp'
try:
# Write to temporary file
pq.write_table(chunk, temp_path, compression='zstd')
# Rename temporary file to final file
if os.path.exists(chunk_path):
os.remove(
chunk_path) # Remove existing file if it exists
os.rename(temp_path, chunk_path)
total_written += len(chunk)
except Exception as e:
# Clean up temporary file if it exists
if os.path.exists(temp_path):
os.remove(temp_path)
raise e
return total_written
except Exception as e:
logger.error(
f"Error processing chunks {start_idx}-{end_idx} for worker {worker_id}: {str(e)}"
)
raise
EntryClass = PreprocessPipeline
+4 -2
View File
@@ -74,7 +74,8 @@ class DenoisingStage(PipelineStage):
)
# Setup precision and autocast settings
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
# target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
@@ -83,6 +84,7 @@ class DenoisingStage(PipelineStage):
), get_sequence_model_parallel_rank()
sp_group = world_size > 1
if sp_group:
# b c t h w -> b t n s h w
latents = rearrange(batch.latents,
"b t (n s) h w -> b t n s h w",
n=world_size).contiguous()
@@ -188,7 +190,7 @@ class DenoisingStage(PipelineStage):
# Predict noise residual
with torch.autocast(device_type="cuda",
dtype=target_dtype,
dtype=torch.bfloat16,
enabled=autocast_enabled):
# TODO(will-refactor): all of this should be in the stage's init
@@ -3,8 +3,6 @@
Input validation stage for diffusion pipelines.
"""
import math
import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
@@ -88,24 +86,4 @@ class InputValidationStage(PipelineStage):
f"Guidance scale must be positive, but got {batch.guidance_scale}"
)
# Adjust number of frames based on number of GPUs
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
orig_latent_num_frames = (batch.num_frames -
1) // temporal_scale_factor + 1
if orig_latent_num_frames % fastvideo_args.num_gpus != 0:
if batch.num_frames_round_down:
new_latent_num_frames = math.floor(
orig_latent_num_frames /
fastvideo_args.num_gpus) * fastvideo_args.num_gpus
else:
new_latent_num_frames = math.ceil(
orig_latent_num_frames /
fastvideo_args.num_gpus) * fastvideo_args.num_gpus
new_num_frames = (new_latent_num_frames -
1) * temporal_scale_factor + 1
logger.info(
"Adjusting number of frames from %s to %s based on number of GPUs (%s)",
batch.num_frames, new_num_frames, fastvideo_args.num_gpus)
batch.num_frames = new_num_frames
return batch
+14 -4
View File
@@ -63,10 +63,15 @@ class TextEncodingStage(PipelineStage):
if fastvideo_args.use_cpu_offload:
text_encoder = text_encoder.to(fastvideo_args.device)
assert isinstance(batch.prompt, str)
text = preprocess_func(batch.prompt)
text_inputs = tokenizer(text, **encoder_config.tokenizer_kwargs).to(
fastvideo_args.device)
assert isinstance(batch.prompt, (str, list))
if isinstance(batch.prompt, str):
batch.prompt = [batch.prompt]
texts = []
for prompt_str in batch.prompt:
texts.append(preprocess_func(prompt_str))
text_inputs = tokenizer(texts,
**encoder_config.tokenizer_kwargs).to(
fastvideo_args.device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
with set_forward_context(current_timestep=0, attn_metadata=None):
@@ -78,6 +83,8 @@ class TextEncodingStage(PipelineStage):
prompt_embeds = postprocess_func(outputs)
batch.prompt_embeds.append(prompt_embeds)
if batch.prompt_attention_mask is not None:
batch.prompt_attention_mask.append(attention_mask)
if batch.do_classifier_free_guidance:
assert isinstance(batch.negative_prompt, str)
@@ -98,6 +105,9 @@ class TextEncodingStage(PipelineStage):
assert batch.negative_prompt_embeds is not None
batch.negative_prompt_embeds.append(negative_prompt_embeds)
if batch.negative_attention_mask is not None:
batch.negative_attention_mask.append(
negative_attention_mask)
if fastvideo_args.use_cpu_offload:
text_encoder.to('cpu')
+824
View File
@@ -0,0 +1,824 @@
import gc
import os
import sys
import time
import traceback
from abc import ABC, abstractmethod
from collections import deque
from copy import deepcopy
import imageio
import numpy as np
import torch
import torchvision
from diffusers.optimization import get_scheduler
from einops import rearrange
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
# import torch.distributed as dist
import wandb
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.utils.checkpoint import save_checkpoint_v1
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.parquet_datasets import ParquetVideoTextDataset
from fastvideo.v1.distributed import cleanup_dist_env_and_memory, get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.training_utils import (
_clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas)
from fastvideo.v1.pipelines.wan.wan_pipeline import WanValidationPipeline
logger = init_logger(__name__)
# Manual gradient checking flag - set to True to enable gradient verification
ENABLE_GRADIENT_CHECK = False
GRADIENT_CHECK_DTYPE = torch.bfloat16
class TrainingPipeline(ComposedPipelineBase, ABC):
"""
A pipeline for training a model. All training pipelines should inherit from this class.
All reusable components and code should be implemented in this class.
"""
_required_config_modules = ["scheduler", "transformer"]
def initialize_training_pipeline(self, fastvideo_args: TrainingArgs):
logger.info("Initializing training pipeline...")
self.device = fastvideo_args.device
self.sp_group = get_sp_group()
self.world_size = self.sp_group.world_size
self.rank = self.sp_group.rank
self.local_rank = self.sp_group.local_rank
self.transformer = self.get_module("transformer")
assert self.transformer is not None
self.transformer.requires_grad_(True)
self.transformer.train()
args = fastvideo_args
noise_scheduler = self.modules["scheduler"]
params_to_optimize = self.transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
init_steps = 0
logger.info("optimizer: %s", optimizer)
# todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * self.world_size,
num_training_steps=args.max_train_steps * self.world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = ParquetVideoTextDataset(
args.data_path,
batch_size=args.train_batch_size,
rank=self.rank,
world_size=self.world_size,
cfg_rate=args.cfg,
num_latent_t=args.num_latent_t)
train_dataloader = StatefulDataLoader(
train_dataset,
batch_size=args.train_batch_size,
num_workers=args.
dataloader_num_workers, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=True)
self.lr_scheduler = lr_scheduler
self.train_dataset = train_dataset
self.train_dataloader = train_dataloader
self.init_steps = init_steps
self.optimizer = optimizer
self.noise_scheduler = noise_scheduler
# self.noise_random_generator = noise_random_generator
# num_update_steps_per_epoch = math.ceil(
# len(train_dataloader) / args.gradient_accumulation_steps *
# args.sp_size / args.train_sp_batch_size)
# args.num_train_epochs = math.ceil(args.max_train_steps /
# num_update_steps_per_epoch)
if self.rank <= 0:
project = args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=args)
@abstractmethod
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
raise NotImplementedError(
"Training pipelines must implement this method")
@abstractmethod
def train_one_step(self, transformer, model_type, optimizer, lr_scheduler,
loader, noise_scheduler, noise_random_generator,
gradient_accumulation_steps, sp_size,
precondition_outputs, max_grad_norm, weighting_scheme,
logit_mean, logit_std, mode_scale):
"""
Train one step of the model.
"""
raise NotImplementedError(
"Training pipeline must implement this method")
def log_validation(self, transformer, fastvideo_args, global_step):
fastvideo_args.inference_mode = True
fastvideo_args.use_cpu_offload = False
if not fastvideo_args.log_validation:
return
if self.validation_pipeline is None:
raise ValueError("Validation pipeline is not set")
# Create sampling parameters if not provided
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
# Prepare validation prompts
print('fastvideo_args.validation_prompt_dir',
fastvideo_args.validation_prompt_dir)
validation_dataset = ParquetVideoTextDataset(
fastvideo_args.validation_prompt_dir,
batch_size=1,
rank=0,
world_size=1,
cfg_rate=0,
num_latent_t=args.num_latent_t)
validation_dataloader = StatefulDataLoader(
validation_dataset,
batch_size=1,
num_workers=1, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
drop_last=False)
transformer.requires_grad_(False)
for p in transformer.parameters():
p.requires_grad = False
transformer.eval()
# Add the transformer to the validation pipeline
self.validation_pipeline.add_module("transformer", transformer)
self.validation_pipeline.latent_preparation_stage.transformer = transformer
self.validation_pipeline.denoising_stage.transformer = transformer
# Process each validation prompt
videos = []
captions = []
for _, embeddings, masks, infos in validation_dataloader:
logger.info(f"infos: {infos}")
caption = infos['caption']
captions.append(caption)
prompt_embeds = embeddings.to(fastvideo_args.device).to(torch.bfloat16)
prompt_attention_mask = masks.to(fastvideo_args.device).to(torch.bfloat16)
# Calculate sizes
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]
logger.info('embed dtype', prompt_embeds.dtype)
# Prepare batch for validation
# print('shape of embeddings', prompt_embeds.shape)
batch = ForwardBatch(
# **shallow_asdict(sampling_param),
data_type="video",
latents=None,
# seed=sampling_param.seed,
# data_type="video",
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
# make sure we use the same height, width, and num_frames as the training pipeline
height=args.num_height,
width=args.num_width,
num_frames=args.num_frames,
# num_inference_steps=fastvideo_args.validation_sampling_steps,
num_inference_steps=50,
# guidance_scale=fastvideo_args.validation_guidance_scale,
guidance_scale=1,
n_tokens=n_tokens,
do_classifier_free_guidance=False,
eta=0.0,
extra={},
)
# Run validation inference
with torch.autocast("cuda", dtype=torch.bfloat16):
with torch.inference_mode():
output_batch = self.validation_pipeline.forward(
batch, fastvideo_args)
samples = output_batch.output
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
videos.append(frames)
# Log validation results
rank = int(os.environ.get("RANK", 0))
if rank == 0:
video_filenames = []
video_captions = []
for i, video in enumerate(videos):
caption = captions[i]
filename = os.path.join(
fastvideo_args.output_dir,
f"validation_step_{global_step}_video_{i}.mp4")
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
video_captions.append(
caption) # Store the caption for each video
logs = {
"validation_videos": [
wandb.Video(filename,
caption=caption) for filename, caption in zip(
video_filenames, video_captions)
]
}
wandb.log(logs, step=global_step)
# Re-enable gradients for training
transformer.requires_grad_(True)
transformer.train()
gc.collect()
torch.cuda.empty_cache()
def gradient_check_parameters(self,
transformer,
latents,
encoder_hidden_states,
encoder_attention_mask,
timesteps,
target,
eps=5e-2,
max_params_to_check=2000):
"""
Verify gradients using finite differences for FSDP models with GRADIENT_CHECK_DTYPE.
Uses standard tolerances for GRADIENT_CHECK_DTYPE precision.
"""
# Move all inputs to CPU and clear GPU memory
inputs_cpu = {
'latents': latents.cpu(),
'encoder_hidden_states': encoder_hidden_states.cpu(),
'encoder_attention_mask': encoder_attention_mask.cpu(),
'timesteps': timesteps.cpu(),
'target': target.cpu()
}
del latents, encoder_hidden_states, encoder_attention_mask, timesteps, target
torch.cuda.empty_cache()
def compute_loss():
# Move inputs to GPU, compute loss, cleanup
inputs_gpu = {
k:
v.to(self.fastvideo_args.device,
dtype=GRADIENT_CHECK_DTYPE
if k != 'encoder_attention_mask' else None)
for k, v in inputs_cpu.items()
}
# Use GRADIENT_CHECK_DTYPE for more accurate gradient checking
# with torch.autocast(enabled=False, device_type="cuda"):
with torch.autocast("cuda", dtype=GRADIENT_CHECK_DTYPE):
with set_forward_context(
current_timestep=inputs_gpu['timesteps'],
attn_metadata=None):
model_pred = transformer(
hidden_states=inputs_gpu['latents'],
encoder_hidden_states=inputs_gpu[
'encoder_hidden_states'],
timestep=inputs_gpu['timesteps'],
encoder_attention_mask=inputs_gpu[
'encoder_attention_mask'],
return_dict=False)[0]
if self.fastvideo_args.precondition_outputs:
sigmas = get_sigmas(self.noise_scheduler,
inputs_gpu['latents'].device,
inputs_gpu['timesteps'],
n_dim=inputs_gpu['latents'].ndim,
dtype=inputs_gpu['latents'].dtype)
model_pred = inputs_gpu['latents'] - model_pred * sigmas
target_adjusted = inputs_gpu['target']
else:
target_adjusted = inputs_gpu['target']
loss = torch.mean((model_pred - target_adjusted)**2)
# Cleanup and return
loss_cpu = loss.cpu()
del inputs_gpu, model_pred, target_adjusted
if 'sigmas' in locals(): del sigmas
torch.cuda.empty_cache()
return loss_cpu.to(self.fastvideo_args.device)
try:
# Get analytical gradients
transformer.zero_grad()
analytical_loss = compute_loss()
analytical_loss.backward()
# Check gradients for selected parameters
absolute_errors = []
param_count = 0
for name, param in transformer.named_parameters():
if not (param.requires_grad and param.grad is not None
and param_count < max_params_to_check
and param.grad.abs().max() > 5e-4):
continue
# Get local parameter and gradient tensors
local_param = param._local_tensor if hasattr(
param, '_local_tensor') else param
local_grad = param.grad._local_tensor if hasattr(
param.grad, '_local_tensor') else param.grad
# Find first significant gradient element
flat_param = local_param.data.view(-1)
flat_grad = local_grad.view(-1)
check_idx = next((i for i in range(min(10, flat_param.numel()))
if abs(flat_grad[i]) > 1e-4), 0)
# Store original values
orig_value = flat_param[check_idx].item()
analytical_grad = flat_grad[check_idx].item()
# Compute numerical gradient
for delta in [eps, -eps]:
with torch.no_grad():
flat_param[check_idx] = orig_value + delta
loss = compute_loss()
if delta > 0: loss_plus = loss.item()
else: loss_minus = loss.item()
# Restore parameter and compute error
with torch.no_grad():
flat_param[check_idx] = orig_value
numerical_grad = (loss_plus - loss_minus) / (2 * eps)
abs_error = abs(analytical_grad - numerical_grad)
rel_error = abs_error / max(abs(analytical_grad),
abs(numerical_grad), 1e-3)
absolute_errors.append(abs_error)
logger.info(
f"{name}[{check_idx}]: analytical={analytical_grad:.6f}, "
f"numerical={numerical_grad:.6f}, abs_error={abs_error:.2e}, rel_error={rel_error:.2%}"
)
# param_count += 1
# Compute and log statistics
if absolute_errors:
min_err, max_err, mean_err = min(absolute_errors), max(
absolute_errors
), sum(absolute_errors) / len(absolute_errors)
logger.info(
f"Gradient check stats: min={min_err:.2e}, max={max_err:.2e}, mean={mean_err:.2e}"
)
if self.rank <= 0:
wandb.log({
"grad_check/min_abs_error":
min_err,
"grad_check/max_abs_error":
max_err,
"grad_check/mean_abs_error":
mean_err,
"grad_check/analytical_loss":
analytical_loss.item(),
})
return max_err
return float('inf')
except Exception as e:
logger.error(f"Gradient check failed: {e}")
traceback.print_exc()
return float('inf')
def setup_gradient_check(self, args, loader_iter, noise_scheduler,
noise_random_generator):
"""
Setup and perform gradient check on a fresh batch.
Args:
args: Training arguments
loader_iter: Data loader iterator
noise_scheduler: Noise scheduler for diffusion
noise_random_generator: Random number generator for noise
Returns:
float or None: Maximum gradient error or None if check is disabled/fails
"""
if not ENABLE_GRADIENT_CHECK:
return None
try:
# Get a fresh batch and process it exactly like train_one_step
check_latents, check_encoder_hidden_states, check_encoder_attention_mask, check_infos = next(
loader_iter)
# Process exactly like in train_one_step but use GRADIENT_CHECK_DTYPE
check_latents = check_latents.to(self.fastvideo_args.device,
dtype=GRADIENT_CHECK_DTYPE)
check_encoder_hidden_states = check_encoder_hidden_states.to(
self.fastvideo_args.device, dtype=GRADIENT_CHECK_DTYPE)
check_latents = normalize_dit_input("wan", check_latents)
batch_size = check_latents.shape[0]
check_noise = torch.randn_like(check_latents)
check_u = compute_density_for_timestep_sampling(
weighting_scheme=args.weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=args.logit_mean,
logit_std=args.logit_std,
mode_scale=args.mode_scale,
)
check_indices = (check_u *
noise_scheduler.config.num_train_timesteps).long()
check_timesteps = noise_scheduler.timesteps[check_indices].to(
device=check_latents.device)
check_sigmas = get_sigmas(
noise_scheduler,
check_latents.device,
check_timesteps,
n_dim=check_latents.ndim,
dtype=check_latents.dtype,
)
check_noisy_model_input = (
1.0 - check_sigmas) * check_latents + check_sigmas * check_noise
# Compute target exactly like train_one_step
if args.precondition_outputs:
check_target = check_latents
else:
check_target = check_noise - check_latents
# Perform gradient check with the exact same inputs as training
max_grad_error = self.gradient_check_parameters(
transformer=self.transformer,
latents=
check_noisy_model_input, # Use noisy input like in training
encoder_hidden_states=check_encoder_hidden_states,
encoder_attention_mask=check_encoder_attention_mask,
timesteps=check_timesteps,
target=check_target,
max_params_to_check=100 # Check more parameters
)
if max_grad_error > 5e-2:
logger.error(
f"❌ Large gradient error detected: {max_grad_error:.2e}")
else:
logger.info(
f"✅ Gradient check passed: max error {max_grad_error:.2e}")
return max_grad_error
except Exception as e:
logger.error(f"Gradient check setup failed: {e}")
traceback.print_exc()
return None
class WanTrainingPipeline(TrainingPipeline):
"""
A training pipeline for Wan.
"""
_required_config_modules = ["scheduler", "transformer"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
pass
def create_training_stages(self, fastvideo_args: FastVideoArgs):
pass
def initialize_validation_pipeline(self, fastvideo_args: FastVideoArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(fastvideo_args)
args_copy.inference_mode = True
args_copy.vae_config.load_encoder = False
# TODO(will): clean this up
args_copy.precision = "bf16"
validation_pipeline = WanValidationPipeline.from_pretrained(
args.model_path, args=args_copy)
self.validation_pipeline = validation_pipeline
def train_one_step(
self,
transformer,
model_type,
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
precondition_outputs,
max_grad_norm,
weighting_scheme,
logit_mean,
logit_std,
mode_scale,
):
self.modules["transformer"].requires_grad_(True)
self.modules["transformer"].train()
total_loss = 0.0
optimizer.zero_grad()
for _ in range(gradient_accumulation_steps):
(
latents,
encoder_hidden_states,
encoder_attention_mask,
infos,
) = next(loader_iter)
latents = latents.to(self.fastvideo_args.device,
dtype=torch.bfloat16)
encoder_hidden_states = encoder_hidden_states.to(
self.fastvideo_args.device, dtype=torch.bfloat16)
latents = normalize_dit_input(model_type, latents)
batch_size = latents.shape[0]
noise = torch.randn_like(latents)
u = compute_density_for_timestep_sampling(
weighting_scheme=weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
logit_mean=logit_mean,
logit_std=logit_std,
mode_scale=mode_scale,
)
indices = (u * noise_scheduler.config.num_train_timesteps).long()
timesteps = noise_scheduler.timesteps[indices].to(
device=latents.device)
if sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
sp_group.broadcast(timesteps, src=0)
sigmas = get_sigmas(
noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
print('device before forward ',
next(transformer.named_parameters())[1].device)
with torch.autocast("cuda", dtype=torch.bfloat16):
input_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if 'hunyuan' in model_type:
input_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
with set_forward_context(current_timestep=timesteps,
attn_metadata=None):
model_pred = transformer(**input_kwargs)[0]
if precondition_outputs:
model_pred = noisy_model_input - model_pred * sigmas
if precondition_outputs:
target = latents
else:
target = noise - latents
loss = (torch.mean((model_pred.float() - target.float())**2) /
gradient_accumulation_steps)
print('device before backwardin context',
next(transformer.named_parameters())[1].device)
print('device before backward out context',
next(transformer.named_parameters())[1].device)
loss.backward()
print('device after backward out context',
next(transformer.named_parameters())[1].device)
avg_loss = loss.detach().clone()
sp_group = get_sp_group()
sp_group.all_reduce(avg_loss, op=torch.distributed.ReduceOp.AVG)
total_loss += avg_loss.item()
model_parts = [self.transformer]
grad_norm = _clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
foreach=None,
)
optimizer.step()
print('device after optimizer step',
next(transformer.named_parameters())[1].device)
lr_scheduler.step()
print('device after scheduler step',
next(transformer.named_parameters())[1].device)
return total_loss, grad_norm.item()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
):
args = fastvideo_args
self.fastvideo_args = args
train_dataloader = self.train_dataloader
init_steps = self.init_steps
lr_scheduler = self.lr_scheduler
optimizer = self.optimizer
noise_scheduler = self.noise_scheduler
noise_random_generator = None
from diffusers import FlowMatchEulerDiscreteScheduler
noise_scheduler = FlowMatchEulerDiscreteScheduler()
# Train!
total_batch_size = (self.world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
logger.info("***** Running training *****")
# logger.info(f" Num examples = {len(train_dataset)}")
# logger.info(f" Dataloader size = {len(train_dataloader)}")
# logger.info(f" Num Epochs = {args.num_train_epochs}")
logger.info(f" Resume training from step {init_steps}")
logger.info(
f" Instantaneous batch size per device = {args.train_batch_size}")
logger.info(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
logger.info(
f" Gradient Accumulation steps = {args.gradient_accumulation_steps}"
)
logger.info(f" Total optimization steps = {args.max_train_steps}")
logger.info(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in self.transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
logger.info(
f" Master weight dtype: {self.transformer.parameters().__next__().dtype}"
)
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=self.local_rank > 0,
)
loader_iter = iter(train_dataloader)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader_iter)
# get gpu memory usage
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info(
f"GPU memory usage before train_one_step: {gpu_memory_usage} MB")
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.perf_counter()
loss, grad_norm = self.train_one_step(
self.transformer,
# args.model_type,
"wan",
optimizer,
lr_scheduler,
loader_iter,
noise_scheduler,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.precondition_outputs,
args.max_grad_norm,
args.weighting_scheme,
args.logit_mean,
args.logit_std,
args.mode_scale,
)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info(
f"GPU memory usage after train_one_step: {gpu_memory_usage} MB")
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
# Manual gradient checking - only at first step
if step == 1 and ENABLE_GRADIENT_CHECK:
logger.info(f"Performing gradient check at step {step}")
self.setup_gradient_check(args, loader_iter, noise_scheduler,
noise_random_generator)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
if self.rank <= 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
raise NotImplementedError("LoRA is not supported now")
# Save LoRA weights
save_lora_checkpoint(transformer, optimizer, rank,
args.output_dir, step, pipe)
else:
# Your existing checkpoint saving code
save_checkpoint_v1(self.transformer, self.rank,
args.output_dir, step)
self.transformer.train()
self.sp_group.barrier()
if args.log_validation and step % args.validation_steps == 0:
self.log_validation(self.transformer, args, step)
if args.use_lora:
raise NotImplementedError("LoRA is not supported now")
# save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps, pipe)
else:
save_checkpoint_v1(self.transformer, self.rank, args.output_dir,
args.max_train_steps)
if get_sp_group():
cleanup_dist_env_and_memory()
def main(args):
logger.info("Starting training pipeline...")
pipeline = WanTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.fastvideo_args
pipeline.forward(None, args)
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.v1.fastvideo_args import TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.use_cpu_offload = False
print(args)
main(args)
+278
View File
@@ -0,0 +1,278 @@
import math
from typing import List, Optional, Tuple, Union
import torch
import torch.distributed as dist
import torch.distributed.tensor
from fastvideo.v1.logger import init_logger
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = False
def compute_density_for_timestep_sampling(
weighting_scheme: str,
batch_size: int,
generator,
logit_mean: float = None,
logit_std: float = None,
mode_scale: float = None,
):
"""
Compute the density for sampling the timesteps when doing SD3 training.
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
"""
if weighting_scheme == "logit_normal":
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
u = torch.normal(
mean=logit_mean,
std=logit_std,
size=(batch_size, ),
device="cpu",
generator=generator,
)
u = torch.nn.functional.sigmoid(u)
elif weighting_scheme == "mode":
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2)**2 - 1 + u)
else:
u = torch.rand(size=(batch_size, ), device="cpu", generator=generator)
return u
def get_sigmas(noise_scheduler,
device,
timesteps,
n_dim=4,
dtype=torch.float32):
sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype)
schedule_timesteps = noise_scheduler.timesteps.to(device)
timesteps = timesteps.to(device)
step_indices = [(schedule_timesteps == t).nonzero().item()
for t in timesteps]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < n_dim:
sigma = sigma.unsqueeze(-1)
return sigma
logger = init_logger(__name__)
def _clip_grad_norm_while_handling_failing_dtensor_cases(
parameters: Union[torch.Tensor, List[torch.Tensor]],
max_norm: float,
norm_type: float = 2.0,
error_if_nonfinite: bool = False,
foreach: Optional[bool] = None,
pp_mesh: Optional[torch.distributed.device_mesh.DeviceMesh] = None,
) -> Optional[torch.Tensor]:
global _HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES
if not _HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES:
try:
return clip_grad_norm_(parameters, max_norm, norm_type,
error_if_nonfinite, foreach, pp_mesh)
except NotImplementedError as e:
if "DTensor does not support cross-mesh operation" in str(e):
# https://github.com/pytorch/pytorch/issues/134212
logger.warning(
"DTensor does not support cross-mesh operation. If you haven't fully tensor-parallelized your "
"model, while combining other parallelisms such as FSDP, it could be the reason for this error. "
"Gradient clipping will be skipped and gradient norm will not be logged."
)
except Exception as e:
logger.warning(
f"An error occurred while clipping gradients: {e}. Gradient clipping will be skipped and gradient "
f"norm will not be logged.")
_HAS_ERRORED_CLIP_GRAD_NORM_WHILE_HANDLING_FAILING_DTENSOR_CASES = True
return None
# Copied from https://github.com/pytorch/torchtitan/blob/4a169701555ab9bd6ca3769f9650ae3386b84c6e/torchtitan/utils.py#L362
@torch.no_grad()
def clip_grad_norm_(
parameters: Union[torch.Tensor, List[torch.Tensor]],
max_norm: float,
norm_type: float = 2.0,
error_if_nonfinite: bool = False,
foreach: Optional[bool] = None,
pp_mesh: Optional[torch.distributed.device_mesh.DeviceMesh] = None,
) -> torch.Tensor:
r"""
Clip the gradient norm of parameters.
Gradient norm clipping requires computing the gradient norm over the entire model.
`torch.nn.utils.clip_grad_norm_` only computes gradient norm along DP/FSDP/TP dimensions.
We need to manually reduce the gradient norm across PP stages.
See https://github.com/pytorch/torchtitan/issues/596 for details.
Args:
parameters (`torch.Tensor` or `List[torch.Tensor]`):
Tensors that will have gradients normalized.
max_norm (`float`):
Maximum norm of the gradients after clipping.
norm_type (`float`, defaults to `2.0`):
Type of p-norm to use. Can be `inf` for infinity norm.
error_if_nonfinite (`bool`, defaults to `False`):
If `True`, an error is thrown if the total norm of the gradients from `parameters` is `nan`, `inf`, or `-inf`.
foreach (`bool`, defaults to `None`):
Use the faster foreach-based implementation. If `None`, use the foreach implementation for CUDA and CPU native tensors
and silently fall back to the slow implementation for other device types.
pp_mesh (`torch.distributed.device_mesh.DeviceMesh`, defaults to `None`):
Pipeline parallel device mesh. If not `None`, will reduce gradient norm across PP stages.
Returns:
`torch.Tensor`:
Total norm of the gradients
"""
grads = [p.grad for p in parameters if p.grad is not None]
# TODO(aryan): Wait for next Pytorch release to use `torch.nn.utils.get_total_norm`
# total_norm = torch.nn.utils.get_total_norm(grads, norm_type, error_if_nonfinite, foreach)
total_norm = _get_total_norm(grads, norm_type, error_if_nonfinite, foreach)
# If total_norm is a DTensor, the placements must be `torch.distributed._tensor.ops.math_ops._NormPartial`.
# We can simply reduce the DTensor to get the total norm in this tensor's process group
# and then convert it to a local tensor.
# It has two purposes:
# 1. to make sure the total norm is computed correctly when PP is used (see below)
# 2. to return a reduced total_norm tensor whose .item() would return the correct value
if isinstance(total_norm, torch.distributed.tensor.DTensor):
# Will reach here if any non-PP parallelism is used.
# If only using PP, total_norm will be a local tensor.
total_norm = total_norm.full_tensor()
if pp_mesh is not None:
if math.isinf(norm_type):
dist.all_reduce(total_norm,
op=dist.ReduceOp.MAX,
group=pp_mesh.get_group())
else:
total_norm **= norm_type
dist.all_reduce(total_norm,
op=dist.ReduceOp.SUM,
group=pp_mesh.get_group())
total_norm **= 1.0 / norm_type
_clip_grads_with_norm_(parameters, max_norm, total_norm, foreach)
return total_norm
@torch.no_grad()
def _clip_grads_with_norm_(
parameters: Union[torch.Tensor, List[torch.Tensor]],
max_norm: float,
total_norm: torch.Tensor,
foreach: Optional[bool] = None,
) -> None:
if isinstance(parameters, torch.Tensor):
parameters = [parameters]
grads = [p.grad for p in parameters if p.grad is not None]
max_norm = float(max_norm)
if len(grads) == 0:
return
grouped_grads: dict[Tuple[torch.device, torch.dtype],
Tuple[List[List[torch.Tensor]],
List[int]]] = (_group_tensors_by_device_and_dtype(
[grads])) # type: ignore[assignment]
clip_coef = max_norm / (total_norm + 1e-6)
# Note: multiplying by the clamped coef is redundant when the coef is clamped to 1, but doing so
# avoids a `if clip_coef < 1:` conditional which can require a CPU <=> device synchronization
# when the gradients do not reside in CPU memory.
clip_coef_clamped = torch.clamp(clip_coef, max=1.0)
for (device, _), ([device_grads], _) in grouped_grads.items():
if (foreach is None and _has_foreach_support(device_grads, device)) or (
foreach and _device_has_foreach_support(device)):
torch._foreach_mul_(device_grads, clip_coef_clamped.to(device))
elif foreach:
raise RuntimeError(
f"foreach=True was passed, but can't use the foreach API on {device.type} tensors"
)
else:
clip_coef_clamped_device = clip_coef_clamped.to(device)
for g in device_grads:
g.mul_(clip_coef_clamped_device)
def _get_total_norm(
tensors: Union[torch.Tensor, List[torch.Tensor]],
norm_type: float = 2.0,
error_if_nonfinite: bool = False,
foreach: Optional[bool] = None,
) -> torch.Tensor:
if isinstance(tensors, torch.Tensor):
tensors = [tensors]
else:
tensors = list(tensors)
norm_type = float(norm_type)
if len(tensors) == 0:
return torch.tensor(0.0)
first_device = tensors[0].device
grouped_tensors: dict[tuple[torch.device, torch.dtype],
tuple[list[list[torch.Tensor]], list[int]]] = (
_group_tensors_by_device_and_dtype(
[tensors] # type: ignore[list-item]
)) # type: ignore[assignment]
norms: List[torch.Tensor] = []
for (device, _), ([device_tensors], _) in grouped_tensors.items():
local_tensors = [
t.to_local()
if isinstance(t, torch.distributed.tensor.DTensor) else t
for t in device_tensors
]
if (foreach is None and _has_foreach_support(local_tensors, device)
) or (foreach and _device_has_foreach_support(device)):
norms.extend(torch._foreach_norm(local_tensors, norm_type))
elif foreach:
raise RuntimeError(
f"foreach=True was passed, but can't use the foreach API on {device.type} tensors"
)
else:
norms.extend(
[torch.linalg.vector_norm(g, norm_type) for g in local_tensors])
total_norm = torch.linalg.vector_norm(
torch.stack([norm.to(first_device) for norm in norms]), norm_type)
if error_if_nonfinite and torch.logical_or(total_norm.isnan(),
total_norm.isinf()):
raise RuntimeError(
f"The total norm of order {norm_type} for gradients from "
"`parameters` is non-finite, so it cannot be clipped. To disable "
"this error and scale the gradients by the non-finite norm anyway, "
"set `error_if_nonfinite=False`")
return total_norm
def _get_foreach_kernels_supported_devices() -> list[str]:
r"""Return the device type list that supports foreach kernels."""
return ["cuda", "xpu", torch._C._get_privateuse1_backend_name()]
@torch.no_grad()
def _group_tensors_by_device_and_dtype(
tensorlistlist: List[List[Optional[torch.Tensor]]],
with_indices: bool = False,
) -> dict[tuple[torch.device, torch.dtype], tuple[
List[List[Optional[torch.Tensor]]], List[int]]]:
return torch._C._group_tensors_by_device_and_dtype(tensorlistlist,
with_indices)
def _device_has_foreach_support(device: torch.device) -> bool:
return device.type in (_get_foreach_kernels_supported_devices() +
["cpu"]) and not torch.jit.is_scripting()
def _has_foreach_support(tensors: List[torch.Tensor],
device: torch.device) -> bool:
return _device_has_foreach_support(device) and all(
t is None or type(t) in [torch.Tensor] for t in tensors)
@@ -0,0 +1,19 @@
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
class WanLatentPipeline(ComposedPipelineBase):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
# def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
pass
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs):
logger.info("WAN Latent Pipeline forward")
pass
+27 -2
View File
@@ -15,7 +15,6 @@ from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
TextEncodingStage,
TimestepPreparationStage)
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
@@ -48,7 +47,33 @@ class WanPipeline(ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
class WanValidationPipeline(ComposedPipelineBase):
"""
Validation pipeline for Wan2.1, assumes that the input are preprocess latents.
"""
_required_config_modules = ["vae", "scheduler"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
+6 -6
View File
@@ -21,10 +21,10 @@ def cuda_platform_plugin() -> Optional[str]:
pynvml = import_pynvml() # type: ignore[no-untyped-call]
pynvml.nvmlInit()
try:
# NOTE: Edge case: vllm cpu build on a GPU machine.
# NOTE: Edge case: fastvideo cpu build on a GPU machine.
# Third-party pynvml can be imported in cpu build,
# we need to check if vllm is built with cpu too.
# Otherwise, vllm will always activate cuda plugin
# we need to check if fastvideo is built with cpu too.
# Otherwise, fastvideo will always activate cuda plugin
# on a GPU machine, even if in a cpu build.
is_cuda = (pynvml.nvmlDeviceGetCount() > 0)
finally:
@@ -72,12 +72,12 @@ if TYPE_CHECKING:
def __getattr__(name: str):
if name == 'current_platform':
# lazy init current_platform.
# 1. out-of-tree platform plugins need `from vllm.platforms import
# 1. out-of-tree platform plugins need `from fastvideo.platforms import
# Platform` so that they can inherit `Platform` class. Therefore,
# we cannot resolve `current_platform` during the import of
# `vllm.platforms`.
# `fastvideo.platforms`.
# 2. when users use out-of-tree platform plugins, they might run
# `import vllm`, some vllm internal code might access
# `import fastvideo`, some fastvideo internal code might access
# `current_platform` during the import, and we need to make sure
# `current_platform` is only resolved after the plugins are loaded
# (we have tests for this, if any developer violate this, they will
-4
View File
@@ -9,10 +9,6 @@ from functools import lru_cache, wraps
from typing import Callable, List, Optional, Tuple, TypeVar, Union
import torch
# NOTE(will): this import is necessary to trigger the registration of the custom
# ops from vllm, which we use
# import custom ops, trigger op registration
import vllm._C # noqa
from typing_extensions import ParamSpec
import fastvideo.v1.envs as envs
+27 -2
View File
@@ -180,8 +180,33 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
else:
processed_args.append(arg)
return super().parse_args( # type: ignore[no-any-return]
processed_args, namespace)
namespace = super().parse_args(processed_args, namespace)
# Track which arguments were explicitly provided
namespace._provided = set()
i = 0
while i < len(args):
arg = args[i]
if arg.startswith('--'):
# Handle --key=value format
if '=' in arg:
key = arg.split('=')[0][2:].replace('-', '_')
namespace._provided.add(key)
i += 1
# Handle --key value format
else:
key = arg[2:].replace('-', '_')
namespace._provided.add(key)
# Skip the value if there is one
if i + 1 < len(args) and not args[i + 1].startswith('-'):
i += 2
else:
i += 1
else:
i += 1
return namespace # type: ignore[no-any-return]
def _pull_args_from_config(self, args: List[str]) -> List[str]:
"""Method to pull arguments specified in the config file
+8
View File
@@ -154,6 +154,14 @@ class Worker:
logger.error(
"Worker %d in loop received KeyboardInterrupt, aborting forward pass",
self.rank)
try:
self.pipe.send(
{"error": "Operation aborted by KeyboardInterrupt"})
logger.info("Worker %d sent error response after interrupt",
self.rank)
except Exception as e:
logger.error("Worker %d failed to send error response: %s",
self.rank, str(e))
continue
@@ -42,7 +42,6 @@ class MultiprocExecutor(Executor):
pipe=worker_pipe))
worker.start()
self.workers.append(worker)
logger.info("Workers: %s", self.workers)
# Wait for all workers to be ready
for idx, pipe in enumerate(self.worker_pipes):
+953
View File
@@ -0,0 +1,953 @@
# !/bin/python3
# isort: skip_file
import argparse
import math
import os
import time
from collections import deque
from copy import deepcopy
import torch
import torch.distributed as dist
import wandb
from accelerate.utils import set_seed
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.optimization import get_scheduler
from diffusers.utils import check_min_version
from peft import LoraConfig
# from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
# from torch.distributed.fsdp import ShardingStrategy
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset,
latent_collate_function)
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint,
save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast,
sp_parallel_dataloader_wrapper)
from fastvideo.utils.dataset_utils import LengthGroupedSampler
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing)
# from fastvideo.utils.load import load_transformer
# from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group, get_sequence_parallel_state,
# initialize_sequence_parallel_state)
from fastvideo.utils.validation import log_validation
from fastvideo.v1.models.loader.component_loader import TransformerLoader
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.forward_context import set_forward_context
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.31.0")
def main_print(content):
if int(os.environ["LOCAL_RANK"]) <= 0:
print(content)
# def reshard_fsdp(model):
# for m in FSDP.fsdp_modules(model):
# if m._has_params and m.sharding_strategy is not ShardingStrategy.NO_SHARD:
# torch.distributed.fsdp._runtime_utils._reshard(m, m._handle, True)
def get_norm(model_pred, norms, gradient_accumulation_steps):
fro_norm = (
torch.linalg.matrix_norm(model_pred, ord="fro") / # codespell:ignore
gradient_accumulation_steps)
largest_singular_value = (torch.linalg.matrix_norm(model_pred, ord=2) /
gradient_accumulation_steps)
absolute_mean = torch.mean(
torch.abs(model_pred)) / gradient_accumulation_steps
absolute_max = torch.max(
torch.abs(model_pred)) / gradient_accumulation_steps
dist.all_reduce(fro_norm, op=dist.ReduceOp.AVG)
dist.all_reduce(largest_singular_value, op=dist.ReduceOp.AVG)
dist.all_reduce(absolute_mean, op=dist.ReduceOp.AVG)
norms["fro"] += torch.mean(fro_norm).item() # codespell:ignore
norms["largest singular value"] += torch.mean(largest_singular_value).item()
norms["absolute mean"] += absolute_mean.item()
norms["absolute max"] += absolute_max.item()
def distill_one_step(
transformer,
model_type,
teacher_transformer,
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
gradient_accumulation_steps,
sp_size,
max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
num_euler_timesteps,
multiphase,
not_apply_cfg_solver,
distill_cfg,
ema_decay,
pred_decay_weight,
pred_decay_type,
hunyuan_teacher_disable_cfg,
):
total_loss = 0.0
optimizer.zero_grad()
model_pred_norm = {
"fro": 0.0, # codespell:ignore
"largest singular value": 0.0,
"absolute mean": 0.0,
"absolute max": 0.0,
}
for _ in range(gradient_accumulation_steps):
(
latents,
encoder_hidden_states,
latents_attention_mask,
encoder_attention_mask,
) = next(loader)
model_input = normalize_dit_input(model_type, latents)
noise = torch.randn_like(model_input)
bsz = model_input.shape[0]
index = torch.randint(0,
num_euler_timesteps, (bsz, ),
device=model_input.device).long()
if sp_size > 1:
broadcast(index)
# Add noise according to flow matching.
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index,
model_input.shape)
timesteps = (sigmas *
noise_scheduler.config.num_train_timesteps).view(-1)
# if squeeze to [], unsqueeze to [1]
timesteps_prev = (sigmas_prev *
noise_scheduler.config.num_train_timesteps).view(-1)
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
# Predict the noise residual
with torch.autocast("cuda", dtype=torch.bfloat16):
print(f"--> noisy_model_input.dtype: {noisy_model_input.dtype}")
print(f"--> noisy_model_input.shape: {noisy_model_input.shape}")
print(
f"--> encoder_hidden_states.dtype: {encoder_hidden_states.dtype}"
)
print(
f"--> encoder_hidden_states.shape: {encoder_hidden_states.shape}"
)
print(f"--> timesteps.dtype: {timesteps.dtype}")
print(f"--> timesteps.shape: {timesteps.shape}")
print(
f"--> encoder_attention_mask.dtype: {encoder_attention_mask.dtype}"
)
print(
f"--> encoder_attention_mask.shape: {encoder_attention_mask.shape}"
)
noisy_model_input = noisy_model_input.to(dtype=torch.bfloat16)
teacher_kwargs = {
"hidden_states": noisy_model_input,
"encoder_hidden_states": encoder_hidden_states,
"timestep": timesteps,
"encoder_attention_mask": encoder_attention_mask, # B, L
"return_dict": False,
}
if hunyuan_teacher_disable_cfg:
teacher_kwargs["guidance"] = torch.tensor(
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
# batch = ForwardBatch(
# enable_teacache=False,
# )
with set_forward_context(current_timestep=0,
attn_metadata=None,
forward_batch=None,
fastvideo_args=None):
model_pred = transformer(**teacher_kwargs)[0]
print(f"--> model_pred shape: {model_pred.shape}")
huber_c = 0.001
target = torch.randn_like(model_pred)
loss = (torch.mean(
torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) -
huber_c) / gradient_accumulation_steps)
loss.backward()
print(f"--> loss: {loss.item()}")
assert False, "stop here"
# if accelerator.is_main_process:
model_pred, end_index = solver.euler_style_multiphase_pred(
noisy_model_input, model_pred, index, multiphase)
with torch.no_grad():
w = distill_cfg
with torch.autocast("cuda", dtype=torch.bfloat16):
cond_teacher_output = teacher_transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
)[0].float()
if not_apply_cfg_solver:
uncond_teacher_output = cond_teacher_output
else:
# Get teacher model prediction on noisy_latents and unconditional embedding
with torch.autocast("cuda", dtype=torch.bfloat16):
uncond_teacher_output = teacher_transformer(
noisy_model_input,
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
timesteps,
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
return_dict=False,
)[0].float()
teacher_output = uncond_teacher_output + w * (cond_teacher_output -
uncond_teacher_output)
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
with torch.no_grad():
with torch.autocast("cuda", dtype=torch.bfloat16):
if ema_transformer is not None:
target_pred = ema_transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
else:
target_pred = transformer(
x_prev.float(),
encoder_hidden_states,
timesteps_prev,
encoder_attention_mask, # B, L
return_dict=False,
)[0]
target, end_index = solver.euler_style_multiphase_pred(
x_prev, target_pred, index, multiphase, True)
huber_c = 0.001
# loss = loss.mean()
loss = (torch.mean(
torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) -
huber_c) / gradient_accumulation_steps)
if pred_decay_weight > 0:
if pred_decay_type == "l1":
pred_decay_loss = (
torch.mean(torch.sqrt(model_pred.float()**2)) *
pred_decay_weight / gradient_accumulation_steps)
loss += pred_decay_loss
elif pred_decay_type == "l2":
# essnetially k2?
pred_decay_loss = (torch.mean(model_pred.float()**2) *
pred_decay_weight /
gradient_accumulation_steps)
loss += pred_decay_loss
else:
assert NotImplementedError("pred_decay_type is not implemented")
# calculate model_pred norm and mean
get_norm(model_pred.detach().float(), model_pred_norm,
gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
dist.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
total_loss += avg_loss.item()
# update ema
if ema_transformer is not None:
reshard_fsdp(ema_transformer)
for p_averaged, p_model in zip(ema_transformer.parameters(),
transformer.parameters()):
with torch.no_grad():
p_averaged.copy_(
torch.lerp(p_averaged.detach(), p_model.detach(),
1 - ema_decay))
grad_norm = transformer.clip_grad_norm_(max_grad_norm)
optimizer.step()
lr_scheduler.step()
return total_loss, grad_norm.item(), model_pred_norm
def main(args):
torch.backends.cuda.matmul.allow_tf32 = True
local_rank = int(os.environ.get("LOCAL_RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.environ.get("WORLD_SIZE", 1))
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=args.sp_size,
sequence_model_parallel_size=args.sp_size)
# If passed along, set the training seed now. On GPU...
if args.seed is not None:
# TODO: t within the same seq parallel group should be the same. Noise should be different.
set_seed(args.seed + rank)
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weights to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
# Create model:
print(
f"--> local_rank: {local_rank}, rank: {rank}, world_size: {world_size}")
main_print(f"--> using model pipeline {args.pretrained_model_name_or_path}")
model_path = maybe_download_model(args.pretrained_model_name_or_path)
main_print(f"--> loading model from {model_path}")
transformer_path = os.path.join(model_path, "transformer")
main_print(f"--> loading transformer from {transformer_path}")
precision = torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16
precision_str = "fp32" if precision == torch.float32 else "bf16"
print(f"--> precision: {precision_str}")
print(f"--> precision: {precision}")
# transformer_path = os.path.join(args.pretrained_model_name_or_path, "transformer")
fastvideo_args = FastVideoArgs(model_path=transformer_path,
use_cpu_offload=False,
precision=precision_str)
# fastvideo_args.dit_config = HunyuanVideoConfig()
fastvideo_args.dit_config = WanVideoConfig()
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
fastvideo_args.device_str = f"cuda:{local_rank}"
device = fastvideo_args.device
torch.cuda.set_device(device)
loader = TransformerLoader()
print(f"--> loading transformer to device {device} on rank {rank}")
transformer = loader.load(transformer_path, "",
fastvideo_args).to(device, dtype=precision)
# teacher_transformer = deepcopy(transformer)
if args.use_ema:
raise NotImplementedError("EMA is not supported for v1 distillation.")
ema_transformer = deepcopy(transformer)
else:
ema_transformer = None
if args.use_lora:
raise NotImplementedError("LoRA is not supported for v1 distillation.")
assert args.model_type == "mochi", "LoRA is only supported for Mochi model."
transformer.requires_grad_(False)
transformer_lora_config = LoraConfig(
r=args.lora_rank,
lora_alpha=args.lora_alpha,
init_lora_weights=True,
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
)
transformer.add_adapter(transformer_lora_config)
main_print(
f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M"
)
main_print(
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
)
if args.use_lora:
raise NotImplementedError("LoRA is not supported for v1 distillation.")
transformer.config.lora_rank = args.lora_rank
transformer.config.lora_alpha = args.lora_alpha
transformer.config.lora_target_modules = [
"to_k", "to_q", "to_v", "to_out.0"
]
transformer._no_split_modules = no_split_modules
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](
transformer)
main_print("--> model loaded")
if args.gradient_checkpointing:
apply_fsdp_checkpointing(transformer, no_split_modules,
args.selective_checkpointing)
apply_fsdp_checkpointing(teacher_transformer, no_split_modules,
args.selective_checkpointing)
if args.use_ema:
apply_fsdp_checkpointing(ema_transformer, no_split_modules,
args.selective_checkpointing)
# Set model as trainable.
transformer.train()
transformer.requires_grad_(True)
# teacher_transformer.requires_grad_(False)
if args.use_ema:
ema_transformer.requires_grad_(False)
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
if args.scheduler_type == "pcm_linear_quadratic":
linear_steps = int(noise_scheduler.config.num_train_timesteps *
args.linear_range)
sigmas = linear_quadratic_schedule(
noise_scheduler.config.num_train_timesteps,
args.linear_quadratic_threshold,
linear_steps,
)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
else:
sigmas = noise_scheduler.sigmas
solver = EulerSolver(
sigmas.numpy()[::-1],
noise_scheduler.config.num_train_timesteps,
euler_timesteps=args.num_euler_timesteps,
)
solver.to(device)
params_to_optimize = transformer.parameters()
# l = list(params_to_optimize)
# for p in params_to_optimize:
# main_print(type(p))
# main_print(f"--> p: {p.shape}")
# main_print(f"--> p: {p.dtype}")
# main_print(f"--> p: {p.device}")
# main_print(f"--> p: {p.requires_grad}")
# main_print('=------------------------')
# break
# print(f"--> params_to_optimize: {list(params_to_optimize)}")
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
# print(f"--> params_to_optimize2: {params_to_optimize}")
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=args.weight_decay,
eps=1e-8,
)
# optimizer = None
init_steps = 0
if args.resume_from_lora_checkpoint:
transformer, optimizer, init_steps = resume_lora_optimizer(
transformer, args.resume_from_lora_checkpoint, optimizer)
main_print(f"optimizer: {optimizer}")
# todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * world_size,
num_training_steps=args.max_train_steps * world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
args.cfg)
uncond_prompt_embed = train_dataset.uncond_prompt_embed
uncond_prompt_mask = train_dataset.uncond_prompt_mask
sampler = (LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
) if (args.group_frame or args.group_resolution) else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False))
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
collate_fn=latent_collate_function,
pin_memory=True,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
drop_last=True,
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader) / args.gradient_accumulation_steps *
args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps /
num_update_steps_per_epoch)
# if rank <= 0:
# project = args.tracker_project_name or "fastvideo"
# wandb.init(project=project, config=args)
# Train!
total_batch_size = (world_size * args.gradient_accumulation_steps /
args.sp_size * args.train_sp_batch_size)
main_print("***** Running training *****")
main_print(f" Num examples = {len(train_dataset)}")
main_print(f" Dataloader size = {len(train_dataloader)}")
main_print(f" Num Epochs = {args.num_train_epochs}")
main_print(f" Resume training from step {init_steps}")
main_print(
f" Instantaneous batch size per device = {args.train_batch_size}")
main_print(
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
)
main_print(
f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
main_print(f" Total optimization steps = {args.max_train_steps}")
main_print(
f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B"
)
# print dtype
main_print(
f" Master weight dtype: {transformer.parameters().__next__().dtype}")
# Potentially load in the weights and states from a previous save
if args.resume_from_checkpoint:
assert NotImplementedError(
"resume_from_checkpoint is not supported now.")
# TODO
progress_bar = tqdm(
range(0, args.max_train_steps),
initial=init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
step_times = deque(maxlen=100)
# todo future
for i in range(init_steps):
next(loader)
# log_validation(args, transformer, device,
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
def get_num_phases(multi_phased_distill_schedule, step):
# step-phase,step-phase
multi_phases = multi_phased_distill_schedule.split(",")
phase = multi_phases[-1].split("-")[-1]
for step_phases in multi_phases:
phase_step, phase = step_phases.split("-")
if step <= int(phase_step):
return int(phase)
return phase
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
assert args.multi_phased_distill_schedule is not None
num_phases = get_num_phases(args.multi_phased_distill_schedule, step)
loss, grad_norm, pred_norm = distill_one_step(
transformer,
args.model_type,
None, # teacher_transformer
ema_transformer,
optimizer,
lr_scheduler,
loader,
noise_scheduler,
solver,
noise_random_generator,
args.gradient_accumulation_steps,
args.sp_size,
args.max_grad_norm,
uncond_prompt_embed,
uncond_prompt_mask,
args.num_euler_timesteps,
num_phases,
args.not_apply_cfg_solver,
args.distill_cfg,
args.ema_decay,
args.pred_decay_weight,
args.pred_decay_type,
args.hunyuan_teacher_disable_cfg,
)
step_time = time.time() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"phases": num_phases,
})
progress_bar.update(1)
if rank <= 0:
wandb.log(
{
"train_loss":
loss,
"learning_rate":
lr_scheduler.get_last_lr()[0],
"step_time":
step_time,
"avg_step_time":
avg_step_time,
"grad_norm":
grad_norm,
"pred_fro_norm":
pred_norm["fro"], # codespell:ignore
"pred_largest_singular_value":
pred_norm["largest singular value"],
"pred_absolute_mean":
pred_norm["absolute mean"],
"pred_absolute_max":
pred_norm["absolute max"],
},
step=step,
)
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(transformer, optimizer, rank,
args.output_dir, step)
else:
# Your existing checkpoint saving code
if args.use_ema:
save_checkpoint(ema_transformer, rank, args.output_dir,
step)
else:
save_checkpoint(transformer, rank, args.output_dir, step)
dist.barrier()
if args.log_validation and step % args.validation_steps == 0:
log_validation(
args,
transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=False,
)
if args.use_ema:
log_validation(
args,
ema_transformer,
device,
torch.bfloat16,
step,
scheduler_type=args.scheduler_type,
shift=args.shift,
num_euler_timesteps=args.num_euler_timesteps,
linear_quadratic_threshold=args.linear_quadratic_threshold,
linear_range=args.linear_range,
ema=True,
)
if args.use_lora:
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir,
args.max_train_steps)
else:
save_checkpoint(transformer, rank, args.output_dir,
args.max_train_steps)
if get_sequence_parallel_state():
destroy_sequence_parallel_group()
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model_type",
type=str,
default="mochi",
help="The type of model to train.")
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_height", type=int, default=480)
parser.add_argument("--num_width", type=int, default=848)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=10,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--train_batch_size",
type=int,
default=16,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--num_latent_t",
type=int,
default=28,
help="Number of latent timesteps.")
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
# text encoder & vae & diffusion model
parser.add_argument("--pretrained_model_name_or_path", type=str)
# parser.add_argument("--dit_model_name_or_path", type=str)
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
# diffusion setting
parser.add_argument("--ema_decay", type=float, default=0.95)
parser.add_argument("--ema_start_step", type=int, default=0)
parser.add_argument("--cfg", type=float, default=0.1)
# validation & logs
parser.add_argument("--validation_prompt_dir", type=str)
parser.add_argument("--validation_sampling_steps", type=str, default="64")
parser.add_argument("--validation_guidance_scale", type=str, default="4.5")
parser.add_argument("--validation_steps", type=float, default=64)
parser.add_argument("--log_validation", action="store_true")
parser.add_argument("--tracker_project_name", type=str, default=None)
parser.add_argument("--seed",
type=int,
default=None,
help="A seed for reproducible training.")
parser.add_argument(
"--output_dir",
type=str,
default=None,
help=
"The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--checkpoints_total_limit",
type=int,
default=None,
help=("Max number of checkpoints to store."),
)
parser.add_argument(
"--checkpointing_steps",
type=int,
default=500,
help=
("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
" training using `--resume_from_checkpoint`."),
)
parser.add_argument("--shift", type=float, default=1.0)
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--resume_from_lora_checkpoint",
type=str,
default=None,
help=
("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=
("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
parser.add_argument(
"--max_train_steps",
type=int,
default=None,
help=
"Total number of training steps to perform. If provided, overrides num_train_epochs.",
)
parser.add_argument(
"--gradient_accumulation_steps",
type=int,
default=1,
help=
"Number of updates steps to accumulate before performing a backward/update pass.",
)
parser.add_argument(
"--learning_rate",
type=float,
default=1e-4,
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help=
"Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_warmup_steps",
type=int,
default=10,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument("--max_grad_norm",
default=1.0,
type=float,
help="Max gradient norm.")
parser.add_argument(
"--gradient_checkpointing",
action="store_true",
help=
"Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=
("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument(
"--mixed_precision",
type=str,
default=None,
choices=["no", "fp16", "bf16"],
help=
("Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
)
parser.add_argument(
"--use_cpu_offload",
action="store_true",
help=
"Whether to use CPU offload for param & gradient & optimizer states.",
)
parser.add_argument("--sp_size",
type=int,
default=1,
help="For sequence parallel")
parser.add_argument(
"--train_sp_batch_size",
type=int,
default=1,
help="Batch size for sequence parallel training",
)
parser.add_argument(
"--use_lora",
action="store_true",
default=False,
help="Whether to use LoRA for finetuning.",
)
parser.add_argument("--lora_alpha",
type=int,
default=256,
help="Alpha parameter for LoRA.")
parser.add_argument("--lora_rank",
type=int,
default=128,
help="LoRA rank parameter. ")
parser.add_argument("--fsdp_sharding_startegy", default="full")
# lr_scheduler
parser.add_argument(
"--lr_scheduler",
type=str,
default="constant",
help=
('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'),
)
parser.add_argument("--num_euler_timesteps", type=int, default=100)
parser.add_argument(
"--lr_num_cycles",
type=int,
default=1,
help="Number of cycles in the learning rate scheduler.",
)
parser.add_argument(
"--lr_power",
type=float,
default=1.0,
help="Power factor of the polynomial scheduler.",
)
parser.add_argument(
"--not_apply_cfg_solver",
action="store_true",
help="Whether to apply the cfg_solver.",
)
parser.add_argument("--distill_cfg",
type=float,
default=3.0,
help="Distillation coefficient.")
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
parser.add_argument("--scheduler_type",
type=str,
default="pcm",
help="The scheduler type to use.")
parser.add_argument(
"--linear_quadratic_threshold",
type=float,
default=0.025,
help="Threshold for linear quadratic scheduler.",
)
parser.add_argument(
"--linear_range",
type=float,
default=0.5,
help="Range for linear quadratic scheduler.",
)
parser.add_argument("--weight_decay",
type=float,
default=0.001,
help="Weight decay to apply.")
parser.add_argument("--use_ema",
action="store_true",
help="Whether to use EMA.")
parser.add_argument("--multi_phased_distill_schedule",
type=str,
default=None)
parser.add_argument("--pred_decay_weight", type=float, default=0.0)
parser.add_argument("--pred_decay_type", default="l1")
parser.add_argument("--hunyuan_teacher_disable_cfg", action="store_true")
parser.add_argument(
"--master_weight_type",
type=str,
default="fp32",
help="Weight type to use - fp32 or bf16.",
)
args = parser.parse_args()
main(args)
+9 -6
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "fastvideo"
version = "0.0.5"
version = "0.1.1"
description = "FastVideo"
readme = "README.md"
requires-python = ">=3.8"
@@ -19,12 +19,9 @@ dependencies = [
# Machine Learning & Transformers
"transformers>=4.46.1", "tokenizers>=0.20.1", "sentencepiece==0.2.0",
"timm==1.0.11", "peft==0.13.2", "diffusers>=0.33.0", "bitsandbytes",
"timm==1.0.11", "peft==0.13.2", "diffusers>=0.33.1", "bitsandbytes",
"torch==2.6.0", "torchvision",
# vLLM
"vllm>=0.7.3",
# Acceleration & Optimization
"accelerate==1.0.1",
@@ -50,6 +47,12 @@ dependencies = [
# flash-attn: pip install flash-attn==2.7.4.post1 --no-cache-dir --no-build-isolation
train = [
"torchdata",
"pyarrow",
"datasets",
]
lint = [
"pre-commit==4.0.1",
]
@@ -60,7 +63,7 @@ test = [
"pytest",
]
dev = [ "fastvideo[lint]", "fastvideo[test]", ]
dev = [ "fastvideo[lint]", "fastvideo[test]", "fastvideo[train]", ]
[project.scripts]
fastvideo = "fastvideo.v1.entrypoints.cli.main:main"
@@ -1,11 +1,11 @@
import json
from pathlib import Path
import csv
import cv2
def get_video_info(video_path, prompt_text):
"""Extract video information using OpenCV and corresponding prompt text"""
def get_video_info(video_path, metadata):
"""Extract video information using OpenCV and corresponding metadata"""
cap = cv2.VideoCapture(str(video_path))
if not cap.isOpened():
@@ -23,60 +23,66 @@ def get_video_info(video_path, prompt_text):
return {
"path": video_path.name,
"title": metadata.get("Video Title", ""),
"description": metadata.get("Video Description", ""),
"video_url": metadata.get("Video URL", ""),
"download_url": metadata.get("Download URL", ""),
"resolution": {
"width": width,
"height": height
},
"fps": fps,
"duration": duration,
"cap": [prompt_text]
"cap": [metadata.get("Video Description", "")]
}
def read_prompt_file(prompt_path):
"""Read and return the content of a prompt file"""
def read_csv_file(csv_path):
"""Read and return the content of a CSV file"""
try:
with open(prompt_path, 'r', encoding='utf-8') as f:
return f.read().strip()
with open(csv_path, 'r', encoding='utf-8') as f:
reader = csv.DictReader(f)
return list(reader)
except Exception as e:
print(f"Error reading prompt file {prompt_path}: {e}")
print(f"Error reading CSV file {csv_path}: {e}")
return None
def process_videos_and_prompts(video_dir_path, prompt_dir_path, verbose=False):
"""Process videos and their corresponding prompt files
def process_videos_from_csv(video_dir_path, csv_path, verbose=False):
"""Process videos using metadata from CSV file
Args:
video_dir_path (str): Path to directory containing video files
prompt_dir_path (str): Path to directory containing prompt files
csv_path (str): Path to CSV file containing video metadata
verbose (bool): Whether to print verbose processing information
"""
video_dir = Path(video_dir_path)
prompt_dir = Path(prompt_dir_path)
csv_data = read_csv_file(csv_path)
processed_data = []
# Ensure directories exist
if not video_dir.exists() or not prompt_dir.exists():
print(f"Error: One or both directories do not exist:\nVideos: {video_dir}\nPrompts: {prompt_dir}")
if not video_dir.exists():
print(f"Error: Video directory does not exist: {video_dir}")
return []
if csv_data is None:
return []
# Process each video file
for video_file in video_dir.glob('*.mp4'):
video_name = video_file.stem
prompt_file = prompt_dir / f"{video_name}.txt"
# Check if corresponding prompt file exists
if not prompt_file.exists():
print(f"Warning: No prompt file found for video {video_name}")
for row in csv_data:
video_filename = row.get("Filename")
if not video_filename:
continue
# Read prompt content
prompt_text = read_prompt_file(prompt_file)
if prompt_text is None:
video_file = video_dir / video_filename
# Check if video file exists
if not video_file.exists():
print(f"Warning: Video file not found: {video_filename}")
continue
# Process video and add to results
video_info = get_video_info(video_file, prompt_text)
video_info = get_video_info(video_file, row)
if video_info:
processed_data.append(video_info)
@@ -105,9 +111,9 @@ def parse_args():
"""Parse command line arguments"""
import argparse
parser = argparse.ArgumentParser(description='Process videos and their corresponding prompt files')
parser = argparse.ArgumentParser(description='Process videos using metadata from CSV file')
parser.add_argument('--video_dir', '-v', required=True, help='Directory containing video files')
parser.add_argument('--prompt_dir', '-p', required=True, help='Directory containing prompt text files')
parser.add_argument('--csv_path', '-c', required=True, help='Path to CSV file containing video metadata')
parser.add_argument('--output_path',
'-o',
required=True,
@@ -121,8 +127,8 @@ if __name__ == "__main__":
# Parse command line arguments
args = parse_args()
# Process videos and prompts
processed_videos = process_videos_and_prompts(args.video_dir, args.prompt_dir, args.verbose)
# Process videos from CSV
processed_videos = process_videos_from_csv(args.video_dir, args.csv_path, args.verbose)
if processed_videos:
# Save results
+11 -3
View File
@@ -24,9 +24,9 @@ def is_16_9_ratio(width: int, height: int, tolerance: float = 0.1) -> bool:
def resize_video(args_tuple):
"""
Resize a single video file.
args_tuple: (input_file, output_dir, width, height, fps)
args_tuple: (input_file, output_dir, width, height, fps, num_frames)
"""
input_file, output_dir, width, height, fps = args_tuple
input_file, output_dir, width, height, fps, num_frames = args_tuple
video = None
resized = None
output_file = output_dir / f"{input_file.name}"
@@ -39,6 +39,13 @@ def resize_video(args_tuple):
if not is_16_9_ratio(video.w, video.h):
return (input_file.name, "skipped", "Not 16:9")
# Calculate target duration based on num_frames and fps
target_duration = num_frames / fps
# Trim video if it's longer than target duration
if video.duration > target_duration:
video = video.subclip(0, target_duration)
def process_frame(frame):
frame_float = frame.astype(float) / 255.0
resized = resize(frame_float, (height, width, 3), mode='reflect', anti_aliasing=True, preserve_range=True)
@@ -75,7 +82,7 @@ def process_folder(args):
print(f"Target: {args.width}x{args.height} at {args.fps}fps")
# Prepare arguments for parallel processing
process_args = [(video_file, output_path, args.width, args.height, args.fps) for video_file in video_files]
process_args = [(video_file, output_path, args.width, args.height, args.fps, args.num_frames) for video_file in video_files]
successful = 0
skipped = 0
@@ -115,6 +122,7 @@ def parse_args():
parser.add_argument('--width', type=int, default=1280, help='Target width in pixels (default: 848)')
parser.add_argument('--height', type=int, default=720, help='Target height in pixels (default: 480)')
parser.add_argument('--fps', type=int, default=30, help='Target frames per second (default: 30)')
parser.add_argument('--num_frames', type=int, default=163, help='Target number of frames (default: 163)')
parser.add_argument('--max_workers',
type=int,
default=4,
+45
View File
@@ -0,0 +1,45 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
DATA_DIR=/workspace/data
num_gpus=1
IP=127.0.0.1
torchrun --nnodes 1 --nproc_per_node $num_gpus \
--node_rank=0 \
--rdzv_id=456 \
--rdzv_backend=c10d \
--rdzv_endpoint=$IP:29500 \
fastvideo/distill_wan.py\
--seed 42\
--cache_dir "$DATA_DIR/.cache"\
--data_json_path "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/HD-Mixkit-Finetune-Hunyuan/validation"\
--train_batch_size=1 \
--num_latent_t 1 \
--sp_size $num_gpus \
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=320\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--master_weight_type="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "2,4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_16_HD"\
--tracker_project_name Hunyuan_Distill \
--num_height 720 \
--num_width 1280 \
--num_frames 125 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver
+27
View File
@@ -0,0 +1,27 @@
#!/bin/bash
num_gpus=2
export FASTVIDEO_ATTENTION_BACKEND=
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/entrypoints/data_preprocessor.py \
--sp_size $num_gpus \
--tp_size $num_gpus \
--height 480 \
--width 832 \
--num_frames 77 \
--num_inference_steps 50 \
--fps 16 \
--guidance_scale 3.0 \
--prompt_path ./assets/prompt.txt \
--neg_prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--seed 1024 \
--output_path outputs_video/ \
--model_path $MODEL_BASE \
--vae-sp \
--text-encoder-precision "fp32" \
--use-cpu-offload
+25
View File
@@ -0,0 +1,25 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
TEXT_ENCODER_PATH="/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/tokenizer"
MODEL_TYPE="wan"
DATA_MERGE_PATH="/workspace/data/Mixkit-Src/merge.txt"
OUTPUT_DIR="/workspace/data/HD-Mixkit-Finetune-Wan"
VALIDATION_PATH="assets/prompt.txt"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size=4 \
--max_height=480 \
--max_width=832 \
--num_frames=81 \
--dataloader_num_workers 1 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--text_encoder_name $TEXT_ENCODER_PATH \
--train_fps 16 \
--validation_prompt_txt $VALIDATION_PATH \
--samples_per_file 108 \
--flush_frequency 108
@@ -0,0 +1,30 @@
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
# MODEL_PATH="/home/ray/.cache/huggingface/hub/models--Wan-AI--Wan2.1-T2V-1.3B-Diffusers/snapshots/0fad780a534b6463e45facd96134c9f345acfa5b"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_MERGE_PATH="data/cats_480/merge.txt"
OUTPUT_DIR="data/cats_480_latents/"
VALIDATION_PATH="assets/prompt.txt"
# torchrun --nproc_per_node=$GPU_NUM \
# fastvideo/data_preprocess/preprocess_vae_latents_v1.py \
# --model_path $MODEL_PATH \
# --data_merge_path $DATA_MERGE_PATH \
# --train_batch_size=1 \
# --max_height=480 \
# --max_width=832 \
# --num_frames=81 \
# --dataloader_num_workers 1 \
# --output_dir=$OUTPUT_DIR \
# --train_fps 16
# torchrun --nproc_per_node=$GPU_NUM \
# fastvideo/data_preprocess/preprocess_text_embeddings_v1.py \
# --model_path $MODEL_PATH \
# --output_dir=$OUTPUT_DIR
torchrun --nproc_per_node=1 \
fastvideo/data_preprocess/preprocess_validation_text_embeddings_v1.py \
--model_path $MODEL_PATH \
--output_dir=$OUTPUT_DIR \
--validation_prompt_txt $VALIDATION_PATH
+53
View File
@@ -0,0 +1,53 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR=data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset
VALIDATION_DIR=data/HD-Mixkit-Finetune-Wan/validation_parquet_dataset
NUM_GPUS=1
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
# --gradient_checkpointing\
# --pretrained_model_name_or_path hunyuanvideo-community/HunyuanVideo \
# --pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
fastvideo/v1/pipelines/training_pipeline.py\
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_prompt_dir "$VALIDATION_DIR"\
--train_batch_size=1\
--num_latent_t 4 \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 1\
--gradient_accumulation_steps=1\
--max_train_steps=120 \
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=60 \
--validation_steps 20\
--validation_sampling_steps "2,4,8" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_finetune"\
--tracker_project_name wan_finetune \
--num_height 480 \
--num_width 832 \
--num_frames 81 \
--shift 3 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 0.01 \
--not_apply_cfg_solver \
--master_weight_type "fp32" \
--max_grad_norm 1.0