Compare commits

...
Author SHA1 Message Date
SolitaryThinker c4a0f789da wip 2025-10-31 20:54:49 +00:00
SolitaryThinker 5ac14938d2 add example 2025-10-31 20:54:00 +00:00
SolitaryThinker 9d239e9f8b prepare for wan2.2 self-forcing 2025-10-31 20:52:31 +00:00
Kaiqin Kong bdec816b31 move STA_configuration.py to fastvideo/attention/backends (#856) 2025-10-29 13:54:13 -07:00
William Lin 2cd2e57d2e [ci] fix causal ssim test (#848) 2025-10-26 19:33:07 -07:00
William Linandainsley 9370234294 [feat] Add gradio local inference demo (#847)
Co-authored-by: ainsley <jzhang2765@wisc.edu>
2025-10-26 07:01:33 -07:00
Jinzhe Pan 50da62e722 [bugfix] always force spawn instead of fork (#852) 2025-10-23 16:36:50 -07:00
William Lin 4f3e8751db [bugfix] [misc] Use training_state_checkpointing_steps in scripts/ (#846) 2025-10-19 20:20:53 -07:00
Jinzhe PanandXingyu Long f4c58894d9 [Feat] add ray support (#838)
Co-authored-by: Xingyu Long <xingyulong97@gmail.com>
2025-10-16 23:17:54 -07:00
Ohm-Rishabh 01c94ef385 [feat] unified trainer logging (#841) 2025-10-16 23:16:16 -07:00
Zhang Peiyuan 2415226d25 Update WeChat Link 2025-10-13 21:02:46 -07:00
Jiali Chen 404314d00f [Feature]Add video-to-video (V2V) pipeline (#829) 2025-10-12 21:53:05 -07:00
zyang6andkiritorl 87489f0872 Add wan2.1 functionality support for Ascend NPU platform (#810)
Co-authored-by: kiritorl <1021709528@qq.com>
2025-10-09 16:25:08 -07:00
Zhang Peiyuan 9ce7c8039e Update Wechat link 2025-10-06 15:01:19 -07:00
William Lin e1e25e95f9 [feature] Add torch profiler (#827) 2025-10-06 07:59:46 -07:00
92 changed files with 5368 additions and 604 deletions
+1 -1
View File
@@ -352,7 +352,7 @@ jobs:
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs && pytest ./fastvideo/entrypoints/ -vs"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
+1 -1
View File
@@ -7,7 +7,7 @@
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
<p align="center">
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/q46BbX6" target="_blank"> <b> WeChat </b> </a> |
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/tMwknPLY" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
+18
View File
@@ -0,0 +1,18 @@
<svg width="252" height="105" viewBox="0 0 252 105" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM100.768 13.1217L87.7028 29.4852H103.143L100.768 13.1217Z" fill="#356CFF"/>
<path d="M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM109.081 90.697L116.802 65.8487C116.802 65.8487 120.959 65.8487 132.242 65.8487C143.525 65.8487 137.586 78.5759 135.211 84.0304C133.307 88.4021 127.491 90.697 122.74 90.697C117.989 90.697 109.081 90.697 109.081 90.697Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944H159.747C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273L125.188 48.273L124 37.97L147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852H131.836C120.142 29.4852 125.897 1.00043 141.337 1.00043L173.188 1.00056Z" fill="#356CFF"/>
<path d="M179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056Z" fill="#356CFF"/>
<path d="M161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM237.948 77.9692C239.984 70.6965 240.917 65.242 228.446 65.242C215.975 65.242 211.818 71.9087 210.037 77.9692C208.255 84.0298 208.255 91.3025 219.538 91.3025C230.821 91.3025 235.911 85.2419 237.948 77.9692Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944M173.188 1.00056C173.188 1.00056 156.777 1.00043 141.337 1.00043M173.188 1.00056L141.337 1.00043M141.337 20.3944C146.088 20.3944 150.839 20.3944 159.747 20.3944M141.337 20.3944H159.747M159.747 20.3944C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273M148.463 48.273C139.556 48.273 125.188 48.273 125.188 48.273M148.463 48.273L125.188 48.273M125.188 48.273L124 37.97M124 37.97C124 37.97 141.931 37.97 147.87 37.97M124 37.97L147.87 37.97M147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852M151.433 29.4852C146.682 29.4852 138.962 29.4852 131.836 29.4852M151.433 29.4852H131.836M131.836 29.4852C120.142 29.4852 125.897 1.00043 141.337 1.00043M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057ZM96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM87.7028 29.4852L100.768 13.1217L103.143 29.4852H87.7028ZM89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457ZM108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM116.802 65.8487L109.081 90.697C109.081 90.697 117.989 90.697 122.74 90.697C127.491 90.697 133.307 88.4021 135.211 84.0304C137.586 78.5759 143.525 65.8487 132.242 65.8487C120.959 65.8487 116.802 65.8487 116.802 65.8487ZM179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056ZM161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457ZM230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM228.446 65.242C240.917 65.242 239.984 70.6965 237.948 77.9692C235.911 85.2419 230.821 91.3025 219.538 91.3025C208.255 91.3025 208.255 84.0298 210.037 77.9692C211.818 71.9087 215.975 65.242 228.446 65.242Z" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M15.2524 55.5451L21.191 100.999L24.7541 100.999L18.8156 55.5451L15.2524 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 55.5451L14.065 100.999L15.2527 100.999L9.31417 55.5451L8.12646 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 55.5451L6.93853 100.999L7.53239 100.999L1.59385 55.5451L1 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M15.2524 48.2724L30.0988 1H33.6619L18.8156 48.2724H15.2524Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 48.2724L22.9728 1H24.1605L9.31417 48.2724H8.12646Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 48.2724L15.8463 1H16.4402L1.59385 48.2724H1Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M85.3271 55.5457H67.5116L87 12.7363L44.3513 68.2729H58.6038L43.1636 101L85.3271 55.5457Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.18771" stroke-miterlimit="16"/>
</svg>

After

Width:  |  Height:  |  Size: 5.7 KiB

+6
View File
@@ -0,0 +1,6 @@
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
</svg>

After

Width:  |  Height:  |  Size: 691 B

+9
View File
@@ -72,6 +72,15 @@ We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
## STA Configuration Logic
Here is a diagram of how the window is configured and passed through the FastVideo pipeline:
<div align="center">
<img src="../../../docs/source/_static/images/STA_configuration.png" width="80%"/>
</div>
## Why is STA Fast?
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

+53
View File
@@ -0,0 +1,53 @@
# Profiling FastVideo
!!! warning
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down the inference.
## Profiling with PyTorch
FastVideo exposes a process-wide torch profiler that you can enable via environment variables. Set `FASTVIDEO_TORCH_PROFILER_DIR` to an absolute directory path to start collecting traces, and specify the regions you want recorded with `FASTVIDEO_TORCH_PROFILE_REGIONS`:
```bash
FASTVIDEO_TORCH_PROFILER_DIR=/mnt/traces/fastvideo \
FASTVIDEO_TORCH_PROFILE_REGIONS="profiler_region_model_loading,profiler_region_training_step"
```
All profiled regions must be registered in `fastvideo.profiler`; the current list includes:
- `profiler_region_model_loading` — pipeline/module loading
- `profiler_region_inference_pre_denoising`
- `profiler_region_inference_denoising`
- `profiler_region_inference_post_denoising`
- `profiler_region_training_checkpoint_saving`
- `profiler_region_training_dit`
- `profiler_region_training_validation`
- `profiler_region_training_epoch`
- `profiler_region_training_step`
- `profiler_region_training_backward`
- `profiler_region_training_optimizer`
- `profiler_region_distillation_teacher_forward`
- `profiler_region_distillation_student_forward`
- `profiler_region_distillation_loss`
- `profiler_region_distillation_update`
While profiling is enabled, FastVideo records additional annotations:
- `fastvideo.region::<name>` spans are emitted when entering a region.
- `fastvideo.profiler.enable_collection` / `fastvideo.profiler.disable_collection` events mark when torch profiler collection is toggled on or off.
Only one profiler instance is created per process; subsequent pipelines reuse the same controller. If you set `FASTVIDEO_TORCH_PROFILE_REGIONS` incorrectly (e.g. misspelled name), FastVideo logs a warning and ignores that entry.
Additional knobs:
- `FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES`
- `FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY`
- `FASTVIDEO_TORCH_PROFILER_WITH_STACK`
- `FASTVIDEO_TORCH_PROFILER_WITH_FLOPS`
Traces can be visualized using <https://ui.perfetto.dev/>.
### Best Practices
- Keep the profiled step count small; traces can be large and slow down job shutdown while the profiler flushes data.
- After profiling, clean up trace directories to avoid filling disks.
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `with self.profiler_controller.region("your_region"):` or the `@profile_region` decorator.
+1
View File
@@ -115,6 +115,7 @@ design/overview
contributing/overview
contributing/developer_env/index
contributing/profiling
:::
:::{toctree}
+44
View File
@@ -0,0 +1,44 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=2,
use_fsdp_inference=True,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
distributed_executor_backend="ray",
# image_encoder_cpu_offload=False,
)
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
if __name__ == "__main__":
main()
@@ -0,0 +1,43 @@
# NOTE: This is still a work in progress, and the checkpoints are not released yet.
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_t2v"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=True,
init_weights_from_safetensors="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
init_weights_from_safetensors_2="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
num_frame_per_block=7,
# image_encoder_cpu_offload=False,
)
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
if __name__ == "__main__":
main()
@@ -0,0 +1,36 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_wan2_1_Fun"
OUTPUT_NAME = "wan2.1_test"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers",
# "alibaba-pai/Wan2.2-Fun-A14B-Control",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
prompt = "一位年轻女性穿着一件粉色的连衣裙,裙子上有白色的装饰和粉色的纽扣。她的头发是紫色的,头上戴着一个红色的大蝴蝶结,显得非常可爱和精致。她还戴着一个红色的领结,整体造型充满了少女感和活力。她的表情温柔,双手轻轻交叉放在身前,姿态优雅。背景是简单的灰色,没有任何多余的装饰,使得人物更加突出。她的妆容清淡自然,突显了她的清新气质。整体画面给人一种甜美、梦幻的感觉,仿佛置身于童话世界中。"
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
# prompt = "A young woman with beautiful, clear eyes and blonde hair stands in the forest, wearing a white dress and a crown. Her expression is serene, reminiscent of a movie star, with fair and youthful skin. Her brown long hair flows in the wind. The video quality is very high, with a clear view. High quality, masterpiece, best quality, high resolution, ultra-fine, fantastical."
# negative_prompt = "Twisted body, limb deformities, text captions, comic, static, ugly, error, messy code."
image_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/8.png"
control_video_path = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/wan_fun/asset_Wan2_2/v1.0/pose.mp4"
video = generator.generate_video(prompt, negative_prompt=negative_prompt, image_path=image_path, video_path=control_video_path, output_path=OUTPUT_PATH, output_video_name=OUTPUT_NAME, save_video=True)
if __name__ == "__main__":
main()
+56
View File
@@ -0,0 +1,56 @@
# FastVideo Gradio Local Demo
This is a Gradio-based web interface for generating videos using the FastVideo framework. The demo allows users to create videos from text prompts with various customization options.
## Overview
The demo uses the FastVideo framework to generate videos based on text prompts. It provides a simple web interface built with Gradio that allows users to:
- Enter text prompts to generate videos
- Customize video parameters (dimensions, number of frames, etc.)
- Use negative prompts to guide the generation process
- Set or randomize seeds for reproducibility
---
## Usage
Run the demo with:
```bash
python examples/inference/gradio/local/gradio_local_demo.py
```
This will start a web server at `http://0.0.0.0:7860` where you can access the interface.
---
## Model Initialization
This demo initializes a `VideoGenerator` with the minimum required arguments for inference. Users can seamlessly adjust inference options between generations, including prompts, resolution, video length, *without ever needing to reload the model*.
## Video Generation
The core functionality is in the `generate_video` function, which:
1. Processes user inputs
2. Uses the FastVideo VideoGenerator from earlier to run inference (`generator.generate_video()`)
## Gradio Interface
The interface is built with several components:
- A text input for the prompt
- A video display for the result
- Inference options in a collapsible accordion:
- Height and width sliders
- Number of frames slider
- Guidance scale slider
- Negative prompt options
- Seed controls
### Inference Options
- **Height/Width**: Control the resolution of the generated video
- **Number of Frames**: Set how many frames to generate
- **Guidance Scale**: Control how closely the generation follows the prompt
- **Negative Prompt**: Specify what you don't want to see in the video
- **Seed**: Control randomness for reproducible results
@@ -0,0 +1,656 @@
import argparse
import os
import base64
import time
import gradio as gr
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
from copy import deepcopy
MODEL_PATH_MAPPING = {
"FastWan2.1-T2V-1.3B": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
# "FastWan2.2-TI2V-5B-FullAttn": "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers",
}
def create_timing_display(inference_time, total_time, stage_execution_times, num_frames):
dit_denoising_time = f"{stage_execution_times[5]:.2f}s" if len(stage_execution_times) > 5 else "N/A"
timing_html = f"""
<div style="margin: 10px 0;">
<h3 style="text-align: center; margin-bottom: 10px;">⏱️ Timing Breakdown</h3>
<div style="display: grid; grid-template-columns: repeat(5, 1fr); gap: 10px; margin-bottom: 10px;">
<div class="timing-card timing-card-highlight">
<div style="font-size: 20px;">🚀</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">DiT Denoising</div>
<div style="font-size: 16px; color: #ffa200; font-weight: bold;">{dit_denoising_time}</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🧠</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">E2E (w. vae/text encoder)</div>
<div style="font-size: 16px; color: #2563eb;">{inference_time:.2f}s</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🎬</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Video Encoding</div>
<div style="font-size: 16px; color: #dc2626;">N/A</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">🌐</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Network Transfer</div>
<div style="font-size: 16px; color: #059669;">N/A</div>
</div>
<div class="timing-card">
<div style="font-size: 20px;">📊</div>
<div style="font-weight: bold; margin: 3px 0; font-size: 14px;">Total Processing</div>
<div style="font-size: 18px; color: #0277bd;">{total_time:.2f}s</div>
</div>
</div>"""
if inference_time > 0:
fps = num_frames / inference_time
timing_html += f"""
<div class="performance-card" style="margin-top: 15px;">
<span style="font-weight: bold;">Generation Speed: </span>
<span style="font-size: 18px; color: #6366f1; font-weight: bold;">{fps:.1f} frames/second</span>
</div>"""
return timing_html + "</div>"
def setup_model_environment(model_path: str) -> None:
if "fullattn" in model_path.lower():
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
else:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
os.environ["FASTVIDEO_STAGE_LOGGING"] = "1"
def load_example_prompts():
def contains_chinese(text):
return any('\u4e00' <= char <= '\u9fff' for char in text)
def load_from_file(filepath):
prompts, labels = [], []
try:
with open(filepath, "r", encoding='utf-8') as f:
for line in f:
line = line.strip()
if line and not contains_chinese(line):
label = line[:100] + "..." if len(line) > 100 else line
labels.append(label)
prompts.append(line)
except Exception as e:
print(f"Warning: Could not read {filepath}: {e}")
return prompts, labels
examples, example_labels = load_from_file("examples/inference/gradio/local/prompts_final.txt")
if not examples:
examples = ["A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background."]
example_labels = ["Crowded rooftop bar at night"]
return examples, example_labels
def create_gradio_interface(default_params: dict[str, SamplingParam], generators: dict[str, VideoGenerator]):
def generate_video(
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
):
model_path = MODEL_PATH_MAPPING.get(model_selection, "FastVideo/FastWan2.1-T2V-1.3B-Diffusers")
setup_model_environment(model_path)
try:
if progress:
progress(0.1, desc="Loading model for local inference...")
generator = generators[model_path]
params = deepcopy(default_params[model_path])
total_start_time = time.time()
if progress:
progress(0.2, desc="Configuring parameters...")
params.prompt = prompt
params.seed = int(seed)
params.guidance_scale = guidance_scale
params.num_frames = int(num_frames)
params.height = int(height)
params.width = int(width)
if randomize_seed:
params.seed = torch.randint(0, 1000000, (1, )).item()
if use_negative_prompt and negative_prompt:
params.negative_prompt = negative_prompt
else:
params.negative_prompt = default_params[model_path].negative_prompt
if progress:
progress(0.4, desc="Generating video locally...")
output_dir = "outputs/"
os.makedirs(output_dir, exist_ok=True)
start_time = time.time()
result = generator.generate_video(prompt=prompt, sampling_param=params, save_video=True, return_frames=False)
inference_time = time.time() - start_time
logging_info = result.get("logging_info", None)
if logging_info:
stage_names = logging_info.get_execution_order()
stage_execution_times = [
logging_info.get_stage_info(stage_name).get("execution_time", 0.0)
for stage_name in stage_names
]
else:
stage_names = []
stage_execution_times = []
total_time = time.time() - total_start_time
timing_details=create_timing_display(inference_time=inference_time, total_time=total_time, stage_execution_times=stage_execution_times, num_frames=params.num_frames)
safe_prompt = params.prompt[:100].replace(' ', '_').replace('/', '_').replace('\\', '_')
video_filename = f"{params.prompt[:100]}.mp4"
output_path = os.path.join(output_dir, video_filename)
if progress:
progress(1.0, desc="Generation complete!")
return output_path, params.seed, timing_details
except Exception as e:
print(f"An error occurred during local generation: {e}")
return None, f"Generation failed: {str(e)}", ""
examples, example_labels = load_example_prompts()
theme = gr.themes.Base().set(
button_primary_background_fill="#2563eb",
button_primary_background_fill_hover="#1d4ed8",
button_primary_text_color="white",
slider_color="#2563eb",
checkbox_background_color_selected="#2563eb",
)
def get_default_values(model_name):
model_path = MODEL_PATH_MAPPING.get(model_name)
if model_path and model_path in default_params:
params = default_params[model_path]
return {
'height': params.height,
'width': params.width,
'num_frames': params.num_frames,
'guidance_scale': params.guidance_scale,
'seed': params.seed,
}
return {
'height': 448,
'width': 832,
'num_frames': 61,
'guidance_scale': 3.0,
'seed': 1024,
}
initial_values = get_default_values("FastWan2.1-T2V-1.3B")
with gr.Blocks(title="FastWan", theme=theme) as demo:
gr.Image("assets/full.svg", show_label=False, container=False, height=80)
gr.HTML("""
<div style="text-align: center; margin-bottom: 10px;">
<p style="font-size: 18px;"> Make Video Generation Go Blurrrrrrr </p>
<p style="font-size: 18px;"> <a href="https://github.com/hao-ai-lab/FastVideo/tree/main" target="_blank">Code</a> | <a href="https://hao-ai-lab.github.io/blogs/fastvideo_post_training/" target="_blank">Blog</a> | <a href="https://hao-ai-lab.github.io/FastVideo/" target="_blank">Docs</a> </p>
</div>
""")
with gr.Accordion("🎥 What Is FastVideo?", open=False):
gr.HTML("""
<div style="padding: 20px; line-height: 1.6;">
<p style="font-size: 16px; margin-bottom: 15px;">
FastVideo is an inference and post-training framework for diffusion models. It features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
</p>
</div>
""")
with gr.Row():
model_selection = gr.Dropdown(
choices=list(MODEL_PATH_MAPPING.keys()),
value="FastWan2.1-T2V-1.3B",
label="Select Model",
interactive=True
)
with gr.Row():
example_dropdown = gr.Dropdown(
choices=example_labels,
label="Example Prompts",
value=None,
interactive=True,
allow_custom_value=False
)
with gr.Row():
with gr.Column(scale=6):
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=3,
placeholder="Describe your scene...",
container=False,
lines=3,
autofocus=True,
)
with gr.Column(scale=1, min_width=120, elem_classes="center-button"):
run_button = gr.Button("Run", variant="primary", size="lg")
with gr.Row():
with gr.Column():
error_output = gr.Text(label="Error", visible=False)
timing_display = gr.Markdown(label="Timing Breakdown", visible=False)
with gr.Row(equal_height=True, elem_classes="main-content-row"):
with gr.Column(scale=1, elem_classes="advanced-options-column"):
with gr.Group():
gr.HTML("<div style='margin: 0 0 15px 0; text-align: center; font-size: 16px;'>Advanced Options</div>")
with gr.Row():
height = gr.Number(
label="Height",
value=initial_values['height'],
interactive=False,
container=True
)
width = gr.Number(
label="Width",
value=initial_values['width'],
interactive=False,
container=True
)
with gr.Row():
num_frames = gr.Number(
label="Number of Frames",
value=initial_values['num_frames'],
interactive=False,
container=True
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=initial_values['guidance_scale'],
)
with gr.Row():
use_negative_prompt = gr.Checkbox(
label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=3,
lines=3,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(
label="Seed",
minimum=0,
maximum=1000000,
step=1,
value=initial_values['seed'],
)
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
seed_output = gr.Number(label="Used Seed")
with gr.Column(scale=1, elem_classes="video-column"):
result = gr.Video(
label="Generated Video",
show_label=True,
height=466,
width=600,
container=True,
elem_classes="video-component"
)
gr.HTML("""
<style>
.center-button {
display: flex !important;
justify-content: center !important;
height: 100% !important;
padding-top: 1.4em !important;
}
.gradio-container {
max-width: 1200px !important;
margin: 0 auto !important;
}
.main {
max-width: 1200px !important;
margin: 0 auto !important;
}
.gr-form, .gr-box, .gr-group {
max-width: 1200px !important;
}
.gr-video {
max-width: 500px !important;
margin: 0 auto !important;
}
.main-content-row {
display: flex !important;
align-items: flex-start !important;
min-height: 500px !important;
gap: 20px !important;
}
.advanced-options-column,
.video-column {
display: flex !important;
flex-direction: column !important;
flex: 1 !important;
min-height: 400px !important;
align-items: stretch !important;
}
.video-column > * {
margin-top: 0 !important;
}
.video-column .gr-video,
.video-component {
margin-top: 0 !important;
padding-top: 0 !important;
}
.video-column .gr-video .gr-form {
margin-top: 0 !important;
}
.advanced-options-column .gr-group,
.video-column .gr-video {
margin-top: 0 !important;
vertical-align: top !important;
}
.advanced-options-column > *:last-child,
.video-column > *:last-child {
flex-grow: 0 !important;
}
@media (max-width: 1400px) {
.main-content-row {
min-height: 600px !important;
}
.advanced-options-column,
.video-column {
min-height: 600px !important;
}
}
@media (max-width: 1200px) {
.main-content-row {
flex-direction: column !important;
align-items: stretch !important;
}
.advanced-options-column,
.video-column {
min-height: auto !important;
width: 100% !important;
}
}
.timing-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 8px;
text-align: center;
min-height: 80px;
display: flex;
flex-direction: column;
justify-content: center;
}
.timing-card-highlight {
background: var(--background-fill-primary) !important;
border: 2px solid var(--color-accent) !important;
}
.performance-card {
background: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color) !important;
padding: 10px;
border-radius: 6px;
text-align: center;
}
.gr-number input[readonly] {
background-color: var(--background-fill-secondary) !important;
border: 1px solid var(--border-color-primary) !important;
color: var(--body-text-color-subdued) !important;
cursor: default !important;
text-align: center !important;
font-weight: 500 !important;
}
</style>
""")
def on_example_select(example_label):
if example_label and example_label in example_labels:
index = example_labels.index(example_label)
return examples[index]
return ""
example_dropdown.change(
fn=on_example_select,
inputs=example_dropdown,
outputs=prompt,
)
gr.HTML("""
<div style="text-align: center; margin-top: 10px; margin-bottom: 15px;">
<p style="font-size: 16px; margin: 0;">Note that this demo is meant to showcase FastWan's quality and that under a large number of requests, generation speed may be affected.</p>
</div>
""")
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=negative_prompt,
)
def on_model_selection_change(selected_model):
if not selected_model:
selected_model = "FastWan2.1-T2V-1.3B"
model_path = MODEL_PATH_MAPPING.get(selected_model)
if model_path and model_path in default_params:
params = default_params[model_path]
return (
gr.update(value=params.height),
gr.update(value=params.width),
gr.update(value=params.num_frames),
gr.update(value=params.guidance_scale),
gr.update(value=params.seed),
)
return (
gr.update(value=448),
gr.update(value=832),
gr.update(value=61),
gr.update(value=3.0),
gr.update(value=1024),
)
model_selection.change(
fn=on_model_selection_change,
inputs=model_selection,
outputs=[height, width, num_frames, guidance_scale, seed],
)
def handle_generation(*args, progress=None, request: gr.Request = None):
model_selection, prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, randomize_seed = args
result_path, seed_or_error, timing_details = generate_video(
prompt, negative_prompt, use_negative_prompt, seed, guidance_scale,
num_frames, height, width, randomize_seed, model_selection, progress
)
if result_path and os.path.exists(result_path):
return (
result_path,
seed_or_error,
gr.update(visible=False),
gr.update(visible=True, value=timing_details),
)
else:
return (
None,
seed_or_error,
gr.update(visible=True, value=seed_or_error),
gr.update(visible=False),
)
run_button.click(
fn=handle_generation,
inputs=[
model_selection,
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
randomize_seed,
],
outputs=[result, seed_output, error_output, timing_display],
concurrency_limit=20,
)
return demo
def main():
parser = argparse.ArgumentParser(description="FastVideo Gradio Local Demo")
parser.add_argument("--t2v_model_paths", type=str,
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
help="Comma separated list of paths to the T2V model(s)")
parser.add_argument("--host", type=str, default="0.0.0.0",
help="Host to bind to")
parser.add_argument("--port", type=int, default=7860,
help="Port to bind to")
args = parser.parse_args()
generators = {}
default_params = {}
model_paths = args.t2v_model_paths.split(",")
for model_path in model_paths:
print(f"Loading model: {model_path}")
setup_model_environment(model_path)
generators[model_path] = VideoGenerator.from_pretrained(model_path)
default_params[model_path] = SamplingParam.from_pretrained(model_path)
demo = create_gradio_interface(default_params, generators)
print(f"Starting Gradio frontend at http://{args.host}:{args.port}")
print(f"T2V Models: {args.t2v_model_paths}")
from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import HTMLResponse, FileResponse
import uvicorn
app = FastAPI()
@app.get("/logo.png")
def get_logo():
return FileResponse(
"assets/full.svg",
media_type="image/svg+xml",
headers={
"Cache-Control": "public, max-age=3600",
"Access-Control-Allow-Origin": "*"
}
)
@app.get("/favicon.ico")
def get_favicon():
favicon_path = "assets/icon-simple.svg"
if os.path.exists(favicon_path):
return FileResponse(
favicon_path,
media_type="image/svg+xml",
headers={
"Cache-Control": "public, max-age=3600",
"Access-Control-Allow-Origin": "*"
}
)
else:
raise HTTPException(status_code=404, detail="Favicon not found")
@app.get("/", response_class=HTMLResponse)
def index(request: Request):
base_url = str(request.base_url).rstrip('/')
return f"""
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>FastWan</title>
<meta name="title" content="FastWan">
<meta name="description" content="Make video generation go blurrrrrrr">
<meta name="keywords" content="FastVideo, video generation, AI, machine learning, FastWan">
<meta property="og:type" content="website">
<meta property="og:url" content="{base_url}/">
<meta property="og:title" content="FastWan">
<meta property="og:description" content="Make video generation go blurrrrrrr">
<meta property="og:image" content="{base_url}/logo.png">
<meta property="og:image:width" content="1200">
<meta property="og:image:height" content="630">
<meta property="og:site_name" content="FastWan">
<meta property="twitter:card" content="summary_large_image">
<meta property="twitter:url" content="{base_url}/">
<meta property="twitter:title" content="FastWan">
<meta property="twitter:description" content="Make video generation go blurrrrrrr">
<meta property="twitter:image" content="{base_url}/logo.png">
<link rel="icon" type="image/png" sizes="32x32" href="/favicon.ico">
<link rel="icon" type="image/png" sizes="16x16" href="/favicon.ico">
<link rel="apple-touch-icon" href="/favicon.ico">
<style>
body, html {{
margin: 0;
padding: 0;
height: 100%;
overflow: hidden;
}}
iframe {{
width: 100%;
height: 100vh;
border: none;
}}
</style>
</head>
<body>
<iframe src="/gradio" width="100%" height="100%" style="border: none;"></iframe>
</body>
</html>
"""
app = gr.mount_gradio_app(
app,
demo,
path="/gradio",
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("fastvideo-logos")]
)
uvicorn.run(app, host=args.host, port=args.port)
if __name__ == "__main__":
main()
@@ -0,0 +1,11 @@
A dynamic shot of a sleek black motorcycle accelerating down an empty highway at sunset. The bike's engine roars as it gains speed, smoke trailing from the tires. The rider, wearing a black leather jacket and helmet, leans forward with determination, gripping the handlebars tightly. The camera follows the motorcycle from a distance, capturing the dust kicked up behind it, then zooms in to show the intense focus on the rider's face. The background showcases the endless road stretching into the horizon with vibrant orange and pink hues of the setting sun. Medium shot transitioning to close-up.
A Jedi Master Yoda, recognizable by his green skin, large ears, and wise wrinkles, is performing on a small stage, strumming a guitar with great concentration. Yoda wears a casual robe and sits on a stool, his eyes closed as he plays, fully immersed in the music. The stage is dimly lit with spotlights highlighting Yoda, creating a mystical atmosphere. The background shows a live audience watching intently. Medium close-up shot focusing on Yoda's expressive face and hands moving gracefully over the guitar strings.
A cute, fluffy panda bear is preparing a meal in a cozy, modern kitchen. The panda is standing at a wooden countertop, wearing a white chef’s hat and apron. It skillfully stirs a pot on the stove with one hand while holding a spatula in the other. The kitchen is well-lit, with appliances and cabinets in pastel colors, creating a warm and inviting atmosphere. The panda moves gracefully, with a focused and determined expression, as steam rises from the pot. Medium shot focusing on the panda’s actions at the stove.
In a futuristic Tokyo rooftop during a heavy rainstorm, a robotic DJ stands behind a turntable, spinning vinyl records in a cyberpunk night setting. The robot has metallic, sleek body parts with glowing blue LED lights, and it moves gracefully with the beat. Raindrops create a shimmering effect as they hit the ground and the DJ. The surrounding environment features neon signs, towering skyscrapers, and a dark, misty atmosphere. The camera starts with a wide shot of the city skyline before zooming in on the DJ performing. Sci-fi, fantasy.
A realistic animated scene featuring a polar bear playing a guitar. The polar bear is standing upright, wearing a cozy fur vest and fingerless gloves. It holds the guitar with both hands, strumming the strings with one hand while plucking them with the other, showcasing natural, fluid motions. The polar bear's expressive face shows concentration and joy as it plays. The background is a snowy Arctic landscape with icebergs and a clear blue sky. The scene captures the bear from a mid-shot angle, focusing on its interaction with the guitar.
The scene opens to a breathtaking view of a tranquil ocean horizon at dusk, displaying a vibrant tapestry of oranges, pinks, and purples as the sun sets. In the foreground, tall, swaying palm trees frame the scene, their silhouettes stark against the colorful sky. The ocean itself shimmers with reflections of the sunset, creating a peaceful, almost ethereal atmosphere. A small boat can be seen in the distance, centered on the horizon, adding a sense of scale and solitude to the scene. The waves gently lap the shore, creating faint patterns on the sandy beach, which stretches across the foreground. Above, the sky is dotted with scattered clouds that catch the last light of the day, enhancing the drama and beauty of the scene. The overall mood is serene and contemplative, capturing a perfect moment of nature’s grandeur.
A large, modern semi-truck accelerating down an empty highway, gaining speed with each second. The truck's powerful engine roars as it moves forward, smoke billowing from the tires. The camera starts from a wide shot, capturing the truck in the distance, then smoothly zooms in to follow the vehicle as it speeds up. The truck's headlights illuminate the road ahead, casting a bright glow. The truck driver can be seen through the windshield, focused and determined. The background shows the vast openness of the highway stretching into the horizon under a clear blue sky. Medium to close-up shots of the truck as it accelerates.
Soft blue light pulses from the blade’s rune-etched hilt, illuminating nearby moss-covered roots and ferns. The surrounding trees are tall and gnarled, their branches curling like claws overhead. Fog swirls gently at ground level, parting slightly as a figure in a cloak approaches from the distance. Medium shot slowly zooming toward the sword, emphasizing its mystical aura.
The video opens with a tranquil scene in the heart of a dense forest, emphasizing two large, textured tree trunks in the foreground framing the view. Sunlight filters through the canopy above, casting intricate patterns of light and shadow on the trees and the ground. Between the tree trunks, a clear view of a calm, muddy river unfolds, its surface shimmering under the gentle sunlight. The riverbank is decorated with a variety of small bushes and vibrant foliage, subtly transitioning into the deep greens of tall, leafy plants. In the background, the dense forest looms, filled with dark, towering trees, their branches intertwining to form an intricate canopy. The scene is bathed in the soft glow of the sun, creating a serene and picturesque setting. Occasional sunbeams pierce through the foliage, adding a magical aura to the landscape. The vibrant reds and oranges of the smaller plants add contrast, bringing warmth to the earthy tones of the scenery. Overall, this harmonious blend of natural elements creates a peaceful and idyllic forest setting.
A lone figure stands on a large, moss-covered rock, surrounded by the soft rush of a nearby stream. The figure is wearing white sneakers and shorts, with a plaid shirt that hangs loosely in the breeze. The lighting creates dramatic shadows, enhancing the textures of the rock and the subtle movement of the water below. In the background, a waterfall cascades into the stream, completing this tranquil and serene nature scene.
In an industrial setting, a person leans casually against a railing, exuding a sense of confidence and composure. They are wearing a striking outfit, consisting of a vibrant, patterned jacket over a simple white crop top, creating a bold contrast. The atmosphere is infused with warm, ambient lighting that casts soft shadows on the concrete walls and metallic surfaces. Intricate wiring and pipes form an intricate backdrop, enhancing the urban aesthetic. Their relaxed posture and direct, engaging gaze suggest a sense of ease in this industrial environment. This scene encapsulates a blend of modern fashion and gritty, urban architecture, creating a visually compelling narrative.
+3 -1
View File
@@ -12,7 +12,7 @@ import torch
import fastvideo.envs as envs
from fastvideo.attention.backends.abstract import AttentionBackend
from fastvideo.logger import init_logger
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
logger = init_logger(__name__)
@@ -117,6 +117,8 @@ def _cached_get_attn_backend(
selected_backend = backend_name_to_enum(backend_by_env_var)
# get device-specific attn_backend
from fastvideo.platforms import current_platform
if selected_backend not in supported_attention_backends:
selected_backend = None
attention_cls = current_platform.get_attn_backend_cls(
@@ -2,13 +2,13 @@ from fastvideo.configs.models.encoders.base import (BaseEncoderOutput,
EncoderConfig,
ImageEncoderConfig,
TextEncoderConfig)
from fastvideo.configs.models.encoders.clip import (CLIPTextConfig,
CLIPVisionConfig)
from fastvideo.configs.models.encoders.clip import (
CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig)
from fastvideo.configs.models.encoders.llama import LlamaConfig
from fastvideo.configs.models.encoders.t5 import T5Config
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig", "LlamaConfig",
"T5Config"
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config"
]
+12
View File
@@ -77,6 +77,8 @@ class CLIPTextConfig(TextEncoderConfig):
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
enable_scale: bool = True
is_causal: bool = True
prefix: str = "clip"
@@ -87,4 +89,14 @@ class CLIPVisionConfig(ImageEncoderConfig):
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
enable_scale: bool = True
is_causal: bool = True
prefix: str = "clip"
@dataclass
class WAN2_1ControlCLIPVisionConfig(CLIPVisionConfig):
num_hidden_layers_override: int | None = 31
require_post_norm: bool | None = False
enable_scale: bool = False
is_causal: bool = False
+3 -1
View File
@@ -13,7 +13,7 @@ from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
SelfForcingWanT2V480PConfig)
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -27,6 +27,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers": WANV2VConfig,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V720PConfig,
@@ -36,6 +37,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers": SelfForcingWan2_2_T2V480PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
+22 -1
View File
@@ -7,7 +7,8 @@ import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
CLIPVisionConfig, T5Config)
CLIPVisionConfig, T5Config,
WAN2_1ControlCLIPVisionConfig)
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
@@ -100,6 +101,16 @@ class WanI2V720PConfig(WanI2V480PConfig):
flow_shift: float | None = 5.0
@dataclass
class WANV2VConfig(WanI2V480PConfig):
"""Configuration for WAN2.1 1.3B Control pipeline."""
image_encoder_config: EncoderConfig = field(
default_factory=WAN2_1ControlCLIPVisionConfig)
# CLIP encoder precision
image_encoder_precision: str = 'bf16'
@dataclass
class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
"""Base configuration for FastWan T2V 1.3B 480P pipeline architecture with DMD"""
@@ -165,3 +176,13 @@ class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 750, 500, 250])
warp_denoising_step: bool = True
@dataclass
class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
is_causal: bool = True
flow_shift: float | None = 12.0
boundary_ratio: float | None = 0.875
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 850, 700, 550, 350, 275, 200, 125])
warp_denoising_step: bool = True
+9
View File
@@ -18,6 +18,9 @@ class SamplingParam:
# Image inputs
image_path: str | None = None
# Video inputs
video_path: str | None = None
# Text inputs
prompt: str | list[str] | None = None
negative_prompt: str = "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"
@@ -200,6 +203,12 @@ class SamplingParam:
default=SamplingParam.image_path,
help="Path to input image for image-to-video generation",
)
parser.add_argument(
"--video_path",
type=str,
default=SamplingParam.video_path,
help="Path to input video for video-to-video generation",
)
parser.add_argument(
"--moba-config-path",
type=str,
+39 -16
View File
@@ -9,7 +9,7 @@ from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
# isort: off
from fastvideo.configs.sample.wan import (
FastWanT2V480PConfig,
FastWanT2V480P_SamplingParam,
Wan2_1_Fun_1_3B_InP_SamplingParam,
Wan2_2_I2V_A14B_SamplingParam,
Wan2_2_T2V_A14B_SamplingParam,
@@ -18,7 +18,9 @@ from fastvideo.configs.sample.wan import (
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam,
SelfForcingWanT2V480PConfig,
Wan2_1_Fun_1_3B_Control_SamplingParam,
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
)
# isort: on
from fastvideo.logger import init_logger
@@ -28,33 +30,50 @@ from fastvideo.utils import (maybe_download_model_index,
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
"FastVideo/FastHunyuan-diffusers":
FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo":
HunyuanSamplingParam,
"FastVideo/stepvideo-t2v-diffusers":
StepVideoT2VSamplingParam,
# Wan2.1
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers":
WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers":
WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers":
WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers":
WanI2V_14B_720P_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
Wan2_1_Fun_1_3B_Control_SamplingParam,
# Wan2.2
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers":
Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers":
Wan2_2_I2V_A14B_SamplingParam,
# FastWan2.1
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480PConfig,
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
FastWanT2V480P_SamplingParam,
# FastWan2.2
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
# Causal Self-Forcing Wan2.1
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers":
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
# Add other specific weight variants
}
@@ -64,6 +83,8 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -74,7 +95,9 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
"wanpipeline":
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
"stepvideo": StepVideoT2VSamplingParam
"wandmdpipeline": FastWanT2V480P_SamplingParam,
"wancausaldmdpipeline": SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
"stepvideo": StepVideoT2VSamplingParam,
# Other fallbacks by architecture
}
+32 -2
View File
@@ -97,7 +97,7 @@ class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
@dataclass
class FastWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
class FastWanT2V480P_SamplingParam(WanT2V_1_3B_SamplingParam):
# DMD parameters
# dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 757, 522])
num_inference_steps: int = 3
@@ -122,6 +122,17 @@ class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
num_inference_steps: int = 50
@dataclass
class Wan2_1_Fun_1_3B_Control_SamplingParam(SamplingParam):
fps: int = 16
num_frames: int = 49
height: int = 832
width: int = 480
guidance_scale: float = 6.0
teacache_params: WanTeaCacheParams = field(
default_factory=lambda: WanTeaCacheParams(teacache_thresh=0.1, ))
# =============================================
# ============= Wan2.2 TI2V Models =============
# =============================================
@@ -162,9 +173,28 @@ class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
# can be overridden during sampling
@dataclass
class Wan2_2_Fun_A14B_Control_SamplingParam(
Wan2_1_Fun_1_3B_Control_SamplingParam):
num_frames: int = 81
# =============================================
# ============= Causal Self-Forcing =============
# =============================================
@dataclass
class SelfForcingWanT2V480PConfig(WanT2V_1_3B_SamplingParam):
class SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam(
Wan2_1_Fun_1_3B_InP_SamplingParam):
pass
@dataclass
class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(
Wan2_2_T2V_A14B_SamplingParam):
guidance_scale: float = 2.0
guidance_scale_2: float = 2.0
num_inference_steps: int = 8
num_frames: int = 81
height: int = 448
width: int = 832
fps: int = 16
@@ -12,6 +12,7 @@ import tqdm
# Dataset
from torch.utils.data import Dataset, Sampler
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.platforms import current_platform
from fastvideo.dataset.utils import collate_rows_from_parquet_schema
from fastvideo.distributed import (get_sp_world_size, get_world_group,
@@ -342,6 +343,7 @@ def build_parquet_map_style_dataloader(
collate_fn=passthrough,
num_workers=num_data_workers,
pin_memory=True,
pin_memory_device=current_platform.device_name,
persistent_workers=num_data_workers > 0,
)
return dataset, loader
@@ -6,7 +6,6 @@ import os
import torch
from torch.distributed import ProcessGroup
from fastvideo.platforms import current_platform
from fastvideo.platforms.interface import CpuArchEnum
from .base_device_communicator import DeviceCommunicatorBase
@@ -22,6 +21,8 @@ class CpuCommunicator(DeviceCommunicatorBase):
super().__init__(cpu_group, device, device_group, unique_name)
self.dist_module = torch.distributed
from fastvideo.platforms import current_platform
if (current_platform.get_cpu_architecture()
== CpuArchEnum.X86) and hasattr(
torch.ops._C,
@@ -0,0 +1,71 @@
import torch
from torch.distributed import ProcessGroup
from fastvideo.distributed.device_communicators.base_device_communicator import (
DeviceCommunicatorBase)
class NpuCommunicator(DeviceCommunicatorBase):
def __init__(self,
cpu_group: ProcessGroup,
device: torch.device | None = None,
device_group: ProcessGroup | None = None,
unique_name: str = ""):
super().__init__(cpu_group, device, device_group, unique_name)
from fastvideo.distributed.device_communicators.pyhccl import (
PyHcclCommunicator)
self.pyhccl_comm: PyHcclCommunicator | None = None
if self.world_size > 1:
self.pyhccl_comm = PyHcclCommunicator(
group=self.cpu_group,
device=self.device,
)
def all_reduce(self, input_, op: torch.distributed.ReduceOp | None = None):
pyhccl_comm = self.pyhccl_comm
assert pyhccl_comm is not None, "pyhccl_comm should not be None"
out = pyhccl_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 pyhccl.
out = input_.clone()
torch.distributed.all_reduce(out, group=self.device_group, op=op)
return out
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank."""
if dst is None:
dst = (self.rank_in_group + 1) % self.world_size
pyhccl_comm = self.pyhccl_comm
if pyhccl_comm is not None and not pyhccl_comm.disabled:
pyhccl_comm.send(tensor, dst)
else:
torch.distributed.send(tensor, self.ranks[dst], self.device_group)
def recv(self,
size: torch.Size,
dtype: torch.dtype,
src: int | None = None) -> torch.Tensor:
"""Receives a tensor from the source rank."""
"""NOTE: `src` is the local rank of the source rank."""
if src is None:
src = (self.rank_in_group - 1) % self.world_size
tensor = torch.empty(size, dtype=dtype, device=self.device)
pyhccl_comm = self.pyhccl_comm
if pyhccl_comm is not None and not pyhccl_comm.disabled:
pyhccl_comm.recv(tensor, src)
else:
torch.distributed.recv(tensor, self.ranks[src], self.device_group)
return tensor
def destroy(self) -> None:
if self.pyhccl_comm is not None:
self.pyhccl_comm = None
@@ -0,0 +1,146 @@
import torch
import torch.distributed as dist
from torch.distributed import ProcessGroup, ReduceOp
from fastvideo.distributed.device_communicators.pyhccl_wrapper import (
HCCLLibrary, aclrtStream_t, buffer_type, hcclComm_t, hcclDataTypeEnum,
hcclRedOpTypeEnum, hcclUniqueId)
from fastvideo.distributed.utils import StatelessProcessGroup
from fastvideo.logger import init_logger
from fastvideo.utils import current_stream
logger = init_logger(__name__)
class PyHcclCommunicator:
def __init__(
self,
group: ProcessGroup | StatelessProcessGroup,
device: int | str | torch.device,
library_path: str | None = None,
):
"""
Args:
group: the process group to work on. If None, it will use the
default process group.
device: the device to bind the PyHcclCommunicator to. If None,
it will be bind to f"npu:{local_rank}".
library_path: the path to the HCCL library. If None, it will
use the default library path.
It is the caller's responsibility to make sure each communicator
is bind to a unique device.
"""
if not isinstance(group, StatelessProcessGroup):
assert dist.is_initialized()
assert dist.get_backend(group) != dist.Backend.HCCL, (
"PyHcclCommunicator should be attached to a non-HCCL group.")
# note: this rank is the rank in the group
self.rank = dist.get_rank(group)
self.world_size = dist.get_world_size(group)
else:
self.rank = group.rank
self.world_size = group.world_size
self.group = group
# if world_size == 1, no need to create communicator
if self.world_size == 1:
self.available = False
self.disabled = True
return
try:
self.hccl = HCCLLibrary(library_path)
except Exception:
logger.warning("disable hccl because of missing HCCL library")
# disable because of missing HCCL library
# e.g. in a non-NPU environment
self.available = False
self.disabled = True
return
self.available = True
self.disabled = False
logger.info("FastVideo is using pyhccl")
if isinstance(device, int):
device = torch.device(f"npu:{device}")
elif isinstance(device, str):
device = torch.device(device)
# now `device` is a `torch.device` object
assert isinstance(device, torch.device)
self.device = device
if self.rank == 0:
# get the unique id from HCCL
with torch.npu.device(device):
self.unique_id = self.hccl.hcclGetUniqueId()
else:
# construct an empty unique id
self.unique_id = hcclUniqueId()
if not isinstance(group, StatelessProcessGroup):
tensor = torch.ByteTensor(list(self.unique_id.internal))
ranks = dist.get_process_group_ranks(group)
# arg `src` in `broadcast` is the global rank
dist.broadcast(tensor, src=ranks[0], group=group)
byte_list = tensor.tolist()
for i, byte in enumerate(byte_list):
self.unique_id.internal[i] = byte
else:
self.unique_id = group.broadcast_obj(self.unique_id, src=0)
# hccl communicator and stream will use this device
# `torch.npu.device` is a context manager that changes the
# current npu device to the specified one
with torch.npu.device(device):
self.comm: hcclComm_t = self.hccl.hcclCommInitRank(
self.world_size, self.unique_id, self.rank)
stream = current_stream()
# A small all_reduce for warmup.
data = torch.zeros(1, device=device)
self.all_reduce(data)
stream.synchronize()
del data
def all_reduce(self,
in_tensor: torch.Tensor,
op: ReduceOp = ReduceOp.SUM,
stream=None) -> torch.Tensor:
if self.disabled:
return None
# hccl communicator created on a specific device
# will only work on tensors on the same device
# otherwise it will cause "illegal memory access"
assert in_tensor.device == self.device, (
f"this hccl communicator is created to work on {self.device}, "
f"but the input tensor is on {in_tensor.device}")
out_tensor = torch.empty_like(in_tensor)
if stream is None:
stream = current_stream()
self.hccl.hcclAllReduce(buffer_type(in_tensor.data_ptr()),
buffer_type(out_tensor.data_ptr()),
in_tensor.numel(),
hcclDataTypeEnum.from_torch(in_tensor.dtype),
hcclRedOpTypeEnum.from_torch(op), self.comm,
aclrtStream_t(stream.npu_stream))
return out_tensor
def broadcast(self, tensor: torch.Tensor, src: int, stream=None):
if self.disabled:
return
assert tensor.device == self.device, (
f"this hccl communicator is created to work on {self.device}, "
f"but the input tensor is on {tensor.device}")
if stream is None:
stream = current_stream()
buffer = buffer_type(tensor.data_ptr())
self.hccl.hcclBroadcast(buffer, tensor.numel(),
hcclDataTypeEnum.from_torch(tensor.dtype), src,
self.comm, aclrtStream_t(stream.npu_stream))
@@ -0,0 +1,208 @@
import ctypes
import platform
from dataclasses import dataclass
from typing import Any
import torch
from torch.distributed import ReduceOp
from fastvideo.logger import init_logger
from fastvideo.utils import find_hccl_library
logger = init_logger(__name__)
hcclResult_t = ctypes.c_int
hcclComm_t = ctypes.c_void_p
class hcclUniqueId(ctypes.Structure):
_fields_ = [("internal", ctypes.c_byte * 4108)]
aclrtStream_t = ctypes.c_void_p
buffer_type = ctypes.c_void_p
hcclDataType_t = ctypes.c_int
class hcclDataTypeEnum:
hcclInt8 = 0
hcclInt16 = 1
hcclInt32 = 2
hcclFloat16 = 3
hcclFloat32 = 4
hcclInt64 = 5
hcclUint64 = 6
hcclUint8 = 7
hcclUint16 = 8
hcclUint32 = 9
hcclFloat64 = 10
hcclBfloat16 = 11
hcclInt128 = 12
@classmethod
def from_torch(cls, dtype: torch.dtype) -> int:
if dtype == torch.int8:
return cls.hcclInt8
if dtype == torch.uint8:
return cls.hcclUint8
if dtype == torch.int32:
return cls.hcclInt32
if dtype == torch.int64:
return cls.hcclInt64
if dtype == torch.float16:
return cls.hcclFloat16
if dtype == torch.float32:
return cls.hcclFloat32
if dtype == torch.float64:
return cls.hcclFloat64
if dtype == torch.bfloat16:
return cls.hcclBfloat16
raise ValueError(f"Unsupported dtype: {dtype}")
hcclRedOp_t = ctypes.c_int
class hcclRedOpTypeEnum:
hcclSum = 0
hcclProd = 1
hcclMax = 2
hcclMin = 3
@classmethod
def from_torch(cls, op: ReduceOp) -> int:
if op == ReduceOp.SUM:
return cls.hcclSum
if op == ReduceOp.PRODUCT:
return cls.hcclProd
if op == ReduceOp.MAX:
return cls.hcclMax
if op == ReduceOp.MIN:
return cls.hcclMin
raise ValueError(f"Unsupported op: {op}")
@dataclass
class Function:
name: str
restype: Any
argtypes: list[Any]
class HCCLLibrary:
exported_functions = [
Function("HcclGetErrorString", ctypes.c_char_p, [hcclResult_t]),
Function("HcclGetRootInfo", hcclResult_t,
[ctypes.POINTER(hcclUniqueId)]),
Function("HcclCommInitRootInfo", hcclResult_t, [
ctypes.c_int,
ctypes.POINTER(hcclUniqueId),
ctypes.c_int,
ctypes.POINTER(hcclComm_t),
]),
Function("HcclAllReduce", hcclResult_t, [
buffer_type,
buffer_type,
ctypes.c_size_t,
hcclDataType_t,
hcclRedOp_t,
hcclComm_t,
aclrtStream_t,
]),
Function("HcclBroadcast", hcclResult_t, [
buffer_type,
ctypes.c_size_t,
hcclDataType_t,
ctypes.c_int,
hcclComm_t,
aclrtStream_t,
]),
Function("HcclCommDestroy", hcclResult_t, [hcclComm_t]),
]
# class attribute to store the mapping from the path to the library
# to avoid loading the same library multiple times
path_to_library_cache: dict[str, Any] = {}
# class attribute to store the mapping from library path
# to the correspongding directory
path_to_dict_mapping: dict[str, dict[str, Any]] = {}
def __init__(self, so_file: str | None = None):
so_file = so_file or find_hccl_library()
try:
if so_file not in HCCLLibrary.path_to_dict_mapping:
lib = ctypes.CDLL(so_file)
HCCLLibrary.path_to_library_cache[so_file] = lib
self.lib = HCCLLibrary.path_to_library_cache[so_file]
except Exception as e:
logger.error(
"Failed to load HCCL library from %s. "
"It is expected if you are not running on Ascend NPUs."
"Otherwise, the hccl library might not exist, be corrupted "
"or it does not support the current platform %s. "
"If you already have the library, please set the "
"environment variable HCCL_SO_PATH"
" to point to the correct hccl library path.", so_file,
platform.platform())
raise e
if so_file not in HCCLLibrary.path_to_dict_mapping:
_funcs: dict[str, Any] = {}
for func in HCCLLibrary.exported_functions:
f = getattr(self.lib, func.name)
f.restype = func.restype
f.argtypes = func.argtypes
_funcs[func.name] = f
HCCLLibrary.path_to_dict_mapping[so_file] = _funcs
self._funcs = HCCLLibrary.path_to_dict_mapping[so_file]
def hcclGetErrorString(self, result: hcclResult_t) -> str:
return self._funcs["HcclGetErrorString"](result).decode("utf-8")
def HCCL_CHECK(self, result: hcclResult_t) -> None:
if result != 0:
error_str = self.hcclGetErrorString(result)
raise RuntimeError(f"HCCL error: {error_str}")
def hcclGetUniqueId(self) -> hcclUniqueId:
unique_id = hcclUniqueId()
self.HCCL_CHECK(self._funcs["HcclGetRootInfo"](ctypes.byref(unique_id)))
return unique_id
def hcclCommInitRank(self, world_size: int, unique_id: hcclUniqueId,
rank: int) -> hcclComm_t:
comm = hcclComm_t()
self.HCCL_CHECK(self._funcs["HcclCommInitRootInfo"](
world_size, ctypes.byref(unique_id), rank, ctypes.byref(comm)))
return comm
def hcclAllReduce(self, sendbuff: buffer_type, recvbuff: buffer_type,
count: int, datatype: int, op: int, comm: hcclComm_t,
stream: aclrtStream_t) -> None:
self.HCCL_CHECK(self._funcs["HcclAllReduce"](sendbuff, recvbuff, count,
datatype, op, comm,
stream))
def hcclBroadcast(self, buf: buffer_type, count: int, datatype: int,
root: int, comm: hcclComm_t,
stream: aclrtStream_t) -> None:
self.HCCL_CHECK(self._funcs["HcclBroadcast"](buf, count, datatype, root,
comm, stream))
def hcclCommDestroy(self, comm: hcclComm_t) -> None:
self.HCCL_CHECK(self._funcs["HcclCommDestroy"](comm))
__all__ = [
"HCCLLibrary",
"hcclDataTypeEnum",
"hcclRedOpTypeEnum",
"hcclUniqueId",
"hcclComm_t",
"aclrtStream_t",
"buffer_type",
]
+37 -27
View File
@@ -45,7 +45,6 @@ from fastvideo.distributed.device_communicators.cpu_communicator import (
CpuCommunicator)
from fastvideo.distributed.utils import StatelessProcessGroup
from fastvideo.logger import init_logger
from fastvideo.platforms import current_platform
logger = init_logger(__name__)
@@ -190,7 +189,6 @@ class GroupCoordinator:
self.device = get_local_torch_device()
self.use_device_communicator = use_device_communicator
self.device_communicator: DeviceCommunicatorBase = None # type: ignore
if use_device_communicator and self.world_size > 1:
# Platform-aware device communicator selection
@@ -203,6 +201,15 @@ class GroupCoordinator:
device_group=self.device_group,
unique_name=self.unique_name,
)
elif current_platform.is_npu():
from fastvideo.distributed.device_communicators.npu_communicator import (
NpuCommunicator)
self.device_communicator = NpuCommunicator(
cpu_group=self.cpu_group,
device=self.device,
device_group=self.device_group,
unique_name=self.unique_name,
)
else:
# For MPS and CPU, use the CPU communicator
self.device_communicator = CpuCommunicator(
@@ -776,8 +783,13 @@ def init_distributed_environment(
):
# Determine the appropriate backend based on the platform
from fastvideo.platforms import current_platform
if backend == "nccl" and not current_platform.is_cuda_alike():
# Use gloo backend for non-CUDA platforms (MPS, CPU)
backend = "nccl"
if current_platform.is_cuda_alike():
logger.info("Using nccl backend for CUDA platform")
elif current_platform.is_npu():
backend = "hccl"
logger.info("Using hccl backend for NPU platform")
else:
backend = "gloo"
logger.info("Using gloo backend for %s platform",
current_platform.device_name)
@@ -791,21 +803,11 @@ def init_distributed_environment(
"distributed_init_method must be provided when initializing "
"distributed environment")
# For MPS, don't pass device_id as it doesn't support device indices
if current_platform.is_mps():
torch.distributed.init_process_group(
backend=backend,
init_method=distributed_init_method,
world_size=world_size,
rank=rank)
else:
# this backend is used for WORLD
torch.distributed.init_process_group(
backend=backend,
init_method=distributed_init_method,
world_size=world_size,
rank=rank,
device_id=device_id)
torch.distributed.init_process_group(
backend=backend,
init_method=distributed_init_method,
world_size=world_size,
rank=rank)
# set the local rank
# local_rank is not available in torch ProcessGroup,
# see https://github.com/pytorch/pytorch/issues/122816
@@ -948,9 +950,14 @@ def get_dp_rank() -> int:
def get_local_torch_device() -> torch.device:
"""Return the torch device for the current rank."""
return torch.device(f"cuda:{envs.LOCAL_RANK}"
) if current_platform.is_cuda_alike() else torch.device(
"mps")
from fastvideo.platforms import current_platform
if current_platform.is_npu():
device = torch.device(f"npu:{envs.LOCAL_RANK}")
elif current_platform.is_cuda_alike() or current_platform.is_cuda():
device = torch.device(f"cuda:{envs.LOCAL_RANK}")
else:
device = torch.device("mps")
return device
def maybe_init_distributed_environment_and_model_parallel(
@@ -968,7 +975,9 @@ def maybe_init_distributed_environment_and_model_parallel(
device = get_local_torch_device()
logger.info(
"Initializing distributed environment with world_size=%d, device=%s",
world_size, device)
world_size,
device,
local_main_process_only=False)
init_distributed_environment(
world_size=world_size,
@@ -979,10 +988,11 @@ def maybe_init_distributed_environment_and_model_parallel(
initialize_model_parallel(tensor_model_parallel_size=tp_size,
sequence_model_parallel_size=sp_size)
# Only set CUDA device if we're on a CUDA platform
if current_platform.is_cuda_alike():
device = torch.device(f"cuda:{local_rank}")
torch.cuda.set_device(device)
# set device if we're on a CUDA/NPU platform
from fastvideo.platforms import current_platform
device_type = current_platform.device_type
device = torch.device(f"{device_type}:{local_rank}")
current_platform.get_torch_device().set_device(device)
def model_parallel_is_initialized() -> bool:
+88 -38
View File
@@ -8,6 +8,7 @@ diffusion models.
import math
import os
import re
import time
from copy import deepcopy
from typing import Any
@@ -110,7 +111,7 @@ class VideoGenerator:
prompt: The prompt to use for generation (optional if prompt_txt is provided)
negative_prompt: The negative prompt to use (overrides the one in fastvideo_args)
output_path: Path to save the video (overrides the one in fastvideo_args)
output_video_name: Name of the video file to save. Default is the first 100 characters of the prompt.
prompt_path: Path to prompt file
save_video: Whether to save the video to disk
return_frames: Whether to return the raw frames
num_inference_steps: Number of denoising steps (overrides fastvideo_args)
@@ -127,8 +128,13 @@ class VideoGenerator:
Either the output dictionary, list of frames, or list of results for batch processing
"""
# Handle batch processing from text file
if self.fastvideo_args.prompt_txt is not None:
prompt_txt_path = self.fastvideo_args.prompt_txt
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
self.fastvideo_args.model_path)
sampling_param.update(kwargs)
if self.fastvideo_args.prompt_txt is not None or sampling_param.prompt_path is not None:
prompt_txt_path = sampling_param.prompt_path or self.fastvideo_args.prompt_txt
if not os.path.exists(prompt_txt_path):
raise FileNotFoundError(
f"Prompt text file not found: {prompt_txt_path}")
@@ -142,22 +148,19 @@ class VideoGenerator:
logger.info("Found %d prompts in %s", len(prompts), prompt_txt_path)
if sampling_param is not None:
original_output_video_name = sampling_param.output_video_name
else:
original_output_video_name = None
results = []
for i, batch_prompt in enumerate(prompts):
logger.info("Processing prompt %d/%d: %s...", i + 1,
len(prompts), batch_prompt[:100])
try:
# Generate video for this prompt using the same logic below
if sampling_param is not None and original_output_video_name is not None:
sampling_param.output_video_name = original_output_video_name + f"_{i}"
output_path = self._prepare_output_path(
sampling_param.output_path, batch_prompt)
kwargs["output_path"] = output_path
result = self._generate_single_video(
batch_prompt, sampling_param, **kwargs)
prompt=batch_prompt,
sampling_param=sampling_param,
**kwargs)
# Add prompt info to result
if isinstance(result, dict):
@@ -181,8 +184,73 @@ class VideoGenerator:
# Single prompt generation (original behavior)
if prompt is None:
raise ValueError("Either prompt or prompt_txt must be provided")
output_path = self._prepare_output_path(sampling_param.output_path,
prompt)
kwargs["output_path"] = output_path
return self._generate_single_video(prompt=prompt,
sampling_param=sampling_param,
**kwargs)
return self._generate_single_video(prompt, sampling_param, **kwargs)
def _prepare_output_path(
self,
output_path: str,
prompt: str,
) -> str:
"""Build a unique, sanitized .mp4 output file path.
- If `output_path` ends with .mp4 (case-insensitive), treat it as a file path.
- Otherwise, treat `output_path` as a directory and derive the filename
from the prompt.
- Invalid filename characters are removed; if the name changes, a
warning is logged.
- If the target path already exists, a numeric suffix is appended.
"""
def _sanitize_filename_component(name: str) -> str:
# Remove characters invalid on common filesystems, strip spaces/dots
sanitized = re.sub(r'[\/:*?"<>|]', '', name)
sanitized = sanitized.strip().strip('.')
sanitized = re.sub(r'\s+', ' ', sanitized)
return sanitized or "video"
base_path, extension = os.path.splitext(output_path)
extension_lower = extension.lower()
if extension_lower == ".mp4":
output_dir = os.path.dirname(output_path)
base_name = os.path.basename(
base_path) # filename without extension
sanitized_base = _sanitize_filename_component(base_name)
if sanitized_base != base_name:
logger.warning(
"The video name '%s' contained invalid characters. It has been renamed to '%s.mp4'",
os.path.basename(output_path),
sanitized_base,
)
video_name = f"{sanitized_base}.mp4"
else:
# Treat as directory; inform if an unexpected extension was provided.
if extension:
logger.info(
"Output path '%s' has non-mp4 extension '%s'; treating it as a directory and using a .mp4 filename derived from the prompt",
output_path,
extension,
)
output_dir = output_path
prompt_component = _sanitize_filename_component(prompt[:100])
video_name = f"{prompt_component}.mp4"
if output_dir:
os.makedirs(output_dir, exist_ok=True)
new_output_path = os.path.join(output_dir, video_name)
counter = 1
while os.path.exists(new_output_path):
name_part, ext_part = os.path.splitext(video_name)
new_video_name = f"{name_part}_{counter}{ext_part}"
new_output_path = os.path.join(output_dir, new_video_name)
counter += 1
return new_output_path
def _generate_single_video(
self,
@@ -200,15 +268,9 @@ class VideoGenerator:
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
else:
sampling_param = deepcopy(sampling_param)
kwargs["prompt"] = prompt
sampling_param.update(kwargs)
sampling_param = deepcopy(sampling_param)
output_path = kwargs["output_path"]
sampling_param.prompt = prompt
# Process negative prompt
if sampling_param.negative_prompt is not None:
sampling_param.negative_prompt = sampling_param.negative_prompt.strip(
@@ -277,7 +339,7 @@ class VideoGenerator:
height: {target_height}
width: {target_width}
video_length: {sampling_param.num_frames}
prompt: {prompt}
prompt: {sampling_param.prompt}
image_path: {sampling_param.image_path}
neg_prompt: {sampling_param.negative_prompt}
seed: {sampling_param.seed}
@@ -288,7 +350,7 @@ class VideoGenerator:
flow_shift: {fastvideo_args.pipeline_config.flow_shift}
embedded_guidance_scale: {fastvideo_args.pipeline_config.embedded_cfg_scale}
save_video: {sampling_param.save_video}
output_path: {sampling_param.output_path}
output_path: {output_path}
""" # type: ignore[attr-defined]
logger.info(debug_str)
@@ -298,13 +360,8 @@ class VideoGenerator:
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
extra={},
)
# Use prompt[:100] for video name
if batch.output_video_name is None:
batch.output_video_name = prompt[:100]
# Run inference
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
@@ -324,15 +381,8 @@ class VideoGenerator:
# Save video if requested
if batch.save_video:
output_path = batch.output_path
if output_path:
os.makedirs(output_path, exist_ok=True)
video_path = os.path.join(output_path,
f"{batch.output_video_name}.mp4")
imageio.mimsave(video_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", video_path)
else:
logger.warning("No output path provided, video not saved")
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
if batch.return_frames:
return frames
+56 -3
View File
@@ -14,20 +14,29 @@ if TYPE_CHECKING:
FASTVIDEO_CACHE_ROOT: str = os.path.expanduser("~/.cache/fastvideo")
FASTVIDEO_CONFIG_ROOT: str = os.path.expanduser("~/.config/fastvideo")
FASTVIDEO_CONFIGURE_LOGGING: int = 1
FASTVIDEO_RAY_PER_WORKER_GPUS: float = 1.0
FASTVIDEO_LOGGING_LEVEL: str = "INFO"
FASTVIDEO_LOGGING_PREFIX: str = ""
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_ATTENTION_CONFIG: str | None = None
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: str | None = None
NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: str | None = None
VERBOSE: bool = False
FASTVIDEO_TORCH_PROFILER_DIR: str | None = None
FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_STACK: bool = True
FASTVIDEO_TORCH_PROFILER_WITH_FLOPS: bool = False
FASTVIDEO_TORCH_PROFILE_REGIONS: str = ""
FASTVIDEO_SERVER_DEV_MODE: bool = False
FASTVIDEO_STAGE_LOGGING: bool = False
FASTVIDEO_HOST_IP: str = ""
FASTVIDEO_LOOPBACK_IP: str = ""
def get_default_cache_root() -> str:
@@ -113,6 +122,23 @@ environment_variables: dict[str, Callable[[], Any]] = {
os.path.join(get_default_cache_root(), "fastvideo"),
)),
# used in distributed environment to determine the ip address
# of the current node, when the node has multiple network interfaces.
# If you are using multi-node inference, you should set this differently
# on each node.
"FASTVIDEO_HOST_IP":
lambda: os.getenv("FASTVIDEO_HOST_IP", ""),
# Used to force set up loopback IP
"FASTVIDEO_LOOPBACK_IP":
lambda: os.getenv("FASTVIDEO_LOOPBACK_IP", ""),
# Number of GPUs per worker in Ray, if it is set to be a fraction,
# it allows ray to schedule multiple actors on a single GPU,
# so that users can colocate other actors on the same GPUs as FastVideo.
"FASTVIDEO_RAY_PER_WORKER_GPUS":
lambda: float(os.getenv("FASTVIDEO_RAY_PER_WORKER_GPUS", "1.0")),
# Interval in seconds to log a warning message when the ring buffer is full
"FASTVIDEO_RINGBUFFER_WARNING_INTERVAL":
lambda: int(os.environ.get("FASTVIDEO_RINGBUFFER_WARNING_INTERVAL", "60")),
@@ -186,9 +212,8 @@ environment_variables: dict[str, Callable[[], Any]] = {
os.path.expanduser(os.getenv("FASTVIDEO_ATTENTION_CONFIG", "."))),
# Use dedicated multiprocess context for workers.
# Both spawn and fork work
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "fork"),
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
# Enables torch profiler if set. Path to the directory where torch profiler
# traces are saved. Note that it must be an absolute path.
@@ -197,6 +222,34 @@ environment_variables: dict[str, Callable[[], Any]] = {
if os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", None) is None else os.
path.expanduser(os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", "."))),
# Enable torch profiler to record shapes if set
# FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES=1. If not set, torch profiler will
# not record shapes.
"FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES":
lambda: bool(
os.getenv("FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES", "0") != "0"),
# Enable torch profiler to profile memory if set
# FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY=1. If not set, torch profiler
# will not profile memory.
"FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY":
lambda: bool(
os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY", "0") != "0"),
# Enable torch profiler to profile stack if set
# FASTVIDEO_TORCH_PROFILER_WITH_STACK=1. If not set, torch profiler WILL
# profile stack by default.
"FASTVIDEO_TORCH_PROFILER_WITH_STACK":
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_STACK", "1") != "0"),
# Enable torch profiler to profile flops if set
# FASTVIDEO_TORCH_PROFILER_WITH_FLOPS=1. If not set, torch profiler will
# not profile flops.
"FASTVIDEO_TORCH_PROFILER_WITH_FLOPS":
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_FLOPS", "0") != "0"),
"FASTVIDEO_TORCH_PROFILE_REGIONS":
lambda: os.getenv("FASTVIDEO_TORCH_PROFILE_REGIONS", ""),
# If set, fastvideo will run in development mode, which will enable
# some additional endpoints for developing and debugging,
# e.g. `/reset_prefix_cache`
+17 -2
View File
@@ -7,15 +7,21 @@ import json
from contextlib import contextmanager
from dataclasses import field
from enum import Enum
from typing import Any
from typing import Any, TYPE_CHECKING
from fastvideo.configs.configs import PreprocessConfig
from fastvideo.configs.pipelines.base import PipelineConfig, STA_Mode
from fastvideo.configs.utils import clean_cli_args
from fastvideo.logger import init_logger
from fastvideo.platforms import current_platform
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean
if TYPE_CHECKING:
from ray.runtime_env import RuntimeEnv
from ray.util.placement_group import PlacementGroup
else:
RuntimeEnv = Any
PlacementGroup = Any
logger = init_logger(__name__)
@@ -91,6 +97,11 @@ class FastVideoArgs:
# Distributed executor backend
distributed_executor_backend: str = "mp"
# a few attributes for ray related
ray_placement_group: PlacementGroup | None = None
ray_runtime_env: RuntimeEnv | None = None
inference_mode: bool = True # if False == training mode
# HuggingFace specific parameters
@@ -421,6 +432,7 @@ class FastVideoArgs:
"--init-weights-from-safetensors-2",
type=str,
help="Path to safetensors file for initial weight loading")
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
@@ -505,6 +517,8 @@ class FastVideoArgs:
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
from fastvideo.platforms import current_platform
if current_platform.is_mps():
self.use_fsdp_inference = False
@@ -661,6 +675,7 @@ class TrainingArgs(FastVideoArgs):
validation_guidance_scale: str = ""
validation_steps: float = 0.0
log_validation: bool = False
trackers: list[str] = dataclasses.field(default_factory=list)
tracker_project_name: str = ""
wandb_run_name: str = ""
seed: int | None = None
+2 -1
View File
@@ -7,7 +7,6 @@ import torch.nn as nn
import torch.nn.functional as F
from fastvideo.layers.custom_op import CustomOp
from fastvideo.platforms import current_platform
@CustomOp.register("rms_norm")
@@ -34,6 +33,8 @@ class RMSNorm(CustomOp):
else var_hidden_size)
self.has_weight = has_weight
from fastvideo.platforms import current_platform
self.weight = torch.ones(hidden_size) if current_platform.is_cuda_alike(
) else torch.ones(hidden_size, dtype=dtype)
if self.has_weight:
+69 -10
View File
@@ -20,7 +20,7 @@ import fastvideo.envs as envs
from fastvideo.attention import (DistributedAttention,
LocalAttention)
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.distributed.parallel_state import get_sp_world_size, get_local_torch_device
from fastvideo.forward_context import get_forward_context
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
@@ -36,6 +36,26 @@ from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImag
from fastvideo.platforms import AttentionBackendEnum, current_platform
logger = init_logger(__name__)
class CacheAppend(torch.autograd.Function):
"""
KV cache with shape [batch, seq_len, heads, head_dim].
"""
@staticmethod
def forward(ctx, storage, active_cache, x, start, end):
# Ensure storage has the same dtype as x
storage.data[:, start:end] = x
ctx.save_for_backward(storage.to(x.dtype))
ctx.start = start
ctx.end = end
return storage[:, :end].to(x.dtype) # Ensure returned value has same dtype as input
@staticmethod
def backward(ctx, grad_output):
start = ctx.start
end = ctx.end
return None, grad_output[:, :start], grad_output[:, start:end], None, None
class CausalWanSelfAttention(nn.Module):
def __init__(self,
@@ -67,6 +87,10 @@ class CausalWanSelfAttention(nn.Module):
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
self.k_cache = None
self.v_cache = None
self.counter = 0
def forward(self,
q: torch.Tensor,
@@ -84,6 +108,19 @@ class CausalWanSelfAttention(nn.Module):
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
"""
if kv_cache is not None:
# Then we are in _forward_inference mode
if self.k_cache is None:
assert self.counter == 0
self.counter += 1
del self.k_cache
self.register_buffer("k_cache", torch.empty(1, self.max_attention_size, self.num_heads, self.head_dim, device=q.device, dtype=q.dtype), persistent=False)
self.k_cache.requires_grad_(True)
if self.v_cache is None:
del self.v_cache
self.register_buffer("v_cache", torch.empty(1, self.max_attention_size, self.num_heads, self.head_dim, device=v.device, dtype=v.dtype), persistent=False)
self.v_cache.requires_grad_(True)
if cache_start is None:
cache_start = current_start
@@ -128,6 +165,8 @@ class CausalWanSelfAttention(nn.Module):
num_new_tokens = roped_query.shape[1]
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
raise Exception("Not implemented")
# @TODO(Wei): This part has not been thoroughly tested yet. Use with caution.
# Calculate the number of new tokens added in this step
# Shift existing cache content left to discard oldest tokens
# Clone the source slice to avoid overlapping memory error
@@ -137,26 +176,44 @@ class CausalWanSelfAttention(nn.Module):
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# self.k_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
# self.k_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# self.v_cache[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
# self.v_cache[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache["local_end_index"].item() + current_end - \
kv_cache["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
# local_k = CacheAppend.apply(self.k_cache, kv_cache["k"], roped_key, local_start_index, local_end_index)
# local_v = CacheAppend.apply(self.v_cache, kv_cache["v"], v, local_start_index, local_end_index)
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
else:
# Assign new keys/values directly up to current_end
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache["k"] = kv_cache["k"].detach()
kv_cache["v"] = kv_cache["v"].detach()
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
# kv_cache["k"] = kv_cache["k"].detach()
# kv_cache["v"] = kv_cache["v"].detach()
# kv_cache["k"][:, local_start_index:local_end_index] = roped_key
# kv_cache["v"][:, local_start_index:local_end_index] = v
local_k = CacheAppend.apply(self.k_cache, kv_cache["k"], roped_key, local_start_index, local_end_index)
# logger.info("Is local_k meta tensor: %s, %s", local_k.is_meta, local_k.shape)
# logger.info("local_start_index: %d, local_end_index: %d, number of zeros in local k: %d", local_start_index, local_end_index, (local_k == 0).sum().item())
local_v = CacheAppend.apply(self.v_cache, kv_cache["v"], v, local_start_index, local_end_index)
x = self.attn(
roped_query,
kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
local_k,
local_v
)
# x = self.attn(
# roped_query,
# kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
# kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
# )
kv_cache["k"] = local_k
kv_cache["v"] = local_v
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
@@ -233,6 +290,9 @@ class CausalWanTransformerBlock(nn.Module):
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
self.null_shift = torch.tensor([0], device=get_local_torch_device())
self.null_scale = torch.tensor([0], device=get_local_torch_device())
def forward(
self,
hidden_states: torch.Tensor,
@@ -283,9 +343,8 @@ class CausalWanTransformerBlock(nn.Module):
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
hidden_states, attn_output, gate_msa, self.null_shift, self.null_scale)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
+4 -3
View File
@@ -685,9 +685,10 @@ class WanTransformer3DModel(CachableDiT):
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
encoder_hidden_states = encoder_hidden_states.to(
orig_dtype) if current_platform.is_mps(
) else encoder_hidden_states # cast to orig_dtype for MPS
if current_platform.is_mps() or current_platform.is_npu():
encoder_hidden_states = encoder_hidden_states.to(orig_dtype)
else:
encoder_hidden_states = encoder_hidden_states # cast to orig_dtype for MPS & NPU
assert encoder_hidden_states.dtype == orig_dtype
+2 -2
View File
@@ -140,7 +140,7 @@ class CLIPAttention(nn.Module):
"embed_dim must be divisible by num_heads "
f"(got `embed_dim`: {self.embed_dim} and `num_heads`:"
f" {self.num_heads}).")
self.scale = self.head_dim**-0.5
self.scale = self.head_dim**-0.5 if config.enable_scale else None
self.dropout = config.attention_dropout
self.qkv_proj = QKVParallelLinear(
@@ -166,7 +166,7 @@ class CLIPAttention(nn.Module):
self.head_dim,
self.num_heads_per_partition,
softmax_scale=self.scale,
causal=True,
causal=config.is_causal,
supported_attention_backends=config._supported_attention_backends)
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
+2 -1
View File
@@ -37,7 +37,6 @@ from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.loader.weight_utils import default_weight_loader
from fastvideo.platforms import current_platform
class AttentionType:
@@ -322,6 +321,8 @@ class T5Attention(nn.Module):
# Encoder/Decoder Self-Attention Layer, attn bias already cached.
assert attn_bias is not None
from fastvideo.platforms import current_platform
if attention_mask is not None:
attention_mask = attention_mask.view(
bs, 1, 1,
+21 -1
View File
@@ -29,7 +29,6 @@ from fastvideo.models.loader.weight_utils import (
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
pt_weights_iterator, safetensors_weights_iterator)
from fastvideo.models.registry import ModelRegistry
from fastvideo.platforms import current_platform
from fastvideo.utils import PRECISION_TO_TYPE
logger = init_logger(__name__)
@@ -251,6 +250,8 @@ class TextEncoderLoader(ComponentLoader):
use_cpu_offload = fastvideo_args.text_encoder_cpu_offload and len(
getattr(model_config, "_fsdp_shard_conditions", [])) > 0
from fastvideo.platforms import current_platform
if fastvideo_args.text_encoder_cpu_offload:
target_device = torch.device(
"mps") if current_platform.is_mps() else torch.device("cpu")
@@ -274,12 +275,27 @@ class TextEncoderLoader(ComponentLoader):
# Explicitly move model to target device after loading weights
model = model.to(target_device)
from fastvideo.platforms import current_platform
if use_cpu_offload:
# Disable FSDP for MPS as it's not compatible
if current_platform.is_mps():
logger.info(
"Disabling FSDP sharding for MPS platform as it's not compatible"
)
elif current_platform.is_npu():
mesh = init_device_mesh(
"npu",
mesh_shape=(1, dist.get_world_size()),
mesh_dim_names=("offload", "replicate"),
)
shard_model(
model,
cpu_offload=True,
reshard_after_forward=True,
mesh=mesh["offload"],
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=fastvideo_args.pin_cpu_memory)
else:
mesh = init_device_mesh(
"cuda",
@@ -325,6 +341,8 @@ class ImageEncoderLoader(TextEncoderLoader):
encoder_config = fastvideo_args.pipeline_config.image_encoder_config
encoder_config.update_model_arch(model_config)
from fastvideo.platforms import current_platform
if fastvideo_args.image_encoder_cpu_offload:
target_device = torch.device("mps") if current_platform.is_mps() else torch.device("cpu")
else:
@@ -379,6 +397,8 @@ class VAELoader(ComponentLoader):
vae_config = fastvideo_args.pipeline_config.vae_config
vae_config.update_model_arch(config)
from fastvideo.platforms import current_platform
if fastvideo_args.vae_cpu_offload:
target_device = torch.device("mps") if current_platform.is_mps() else torch.device("cpu")
else:
+11 -2
View File
@@ -108,8 +108,17 @@ def maybe_load_fsdp_model(
if not training_mode and not fsdp_inference:
hsdp_replicate_dim = world_size
hsdp_shard_dim = 1
device_mesh = init_device_mesh(
if current_platform.is_npu():
with torch.device("cpu"):
device_mesh = init_device_mesh(
"npu",
# (Replicate(), Shard(dim=0))
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
mesh_dim_names=("replicate", "shard"),
)
else:
device_mesh = init_device_mesh(
"cuda",
# (Replicate(), Shard(dim=0))
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
+2 -1
View File
@@ -26,7 +26,6 @@ from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.layers.activation import get_act_fn
from fastvideo.models.vaes.common import (DiagonalGaussianDistribution,
ParallelTiledVAE)
from fastvideo.platforms import current_platform
CACHE_T = 2
@@ -189,6 +188,8 @@ class WanCausalConv3d(nn.Conv3d):
self.padding = (0, 0, 0)
def forward(self, x, cache_x=None):
from fastvideo.platforms import current_platform
padding = list(self._padding)
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
+146 -29
View File
@@ -3,6 +3,7 @@
import os
import tempfile
from collections.abc import Callable
from typing import Any
from urllib.parse import unquote, urlparse
import imageio
@@ -11,6 +12,8 @@ import PIL.Image
import PIL.ImageOps
import requests
import torch
import torch.nn.functional as F
import torchvision.transforms.functional as TF
from packaging import version
if version.parse(version.parse(
@@ -131,12 +134,88 @@ def load_image(
return image
def _load_gif(gif_path: str) -> tuple[list[PIL.Image.Image], float | None]:
"""
Load frames from a GIF file.
Args:
gif_path: Path to the GIF file
Returns:
Tuple of (list of PIL images, original FPS or None)
"""
pil_images = []
original_fps = None
with PIL.Image.open(gif_path) as gif:
# Extract FPS from GIF metadata
if hasattr(gif, 'info') and 'duration' in gif.info:
duration_ms = gif.info['duration']
if duration_ms > 0:
original_fps = 1000.0 / duration_ms
# Extract all frames
try:
while True:
pil_images.append(gif.copy())
gif.seek(gif.tell() + 1)
except EOFError:
# End of GIF reached
pass
return pil_images, original_fps
def _load_video_with_ffmpeg(
video_path: str) -> tuple[list[PIL.Image.Image], float | None]:
"""
Load frames from a video file using ffmpeg.
Args:
video_path: Path to the video file
Returns:
Tuple of (list of PIL images, original FPS or None)
Raises:
AttributeError: If ffmpeg is not installed
"""
# Verify ffmpeg is available
try:
imageio.plugins.ffmpeg.get_exe()
except AttributeError as e:
raise AttributeError(
"Unable to find an ffmpeg installation on your machine. "
"Please install via `pip install imageio-ffmpeg`") from e
pil_images = []
original_fps = None
with imageio.get_reader(video_path) as reader:
# Try to extract FPS metadata
metadata = reader.get_meta_data()
original_fps = metadata.get('fps')
# Fallback: try format-specific metadata
if original_fps is None:
source_size = metadata.get('source_size', {})
if isinstance(source_size, dict):
original_fps = source_size.get('fps')
# Extract all frames
for frame in reader:
pil_images.append(PIL.Image.fromarray(frame))
return pil_images, original_fps
# adapted from diffusers.utils import load_video
def load_video(
video: str,
convert_method: Callable[[list[PIL.Image.Image]], list[PIL.Image.Image]]
| None = None,
) -> list[PIL.Image.Image]:
return_fps: bool = False,
) -> tuple[list[PIL.Image.Image], float | Any] | list[PIL.Image.Image]:
"""
Loads `video` to a list of PIL Image.
Args:
@@ -145,9 +224,12 @@ def load_video(
convert_method (Callable[[List[PIL.Image.Image]], List[PIL.Image.Image]], *optional*):
A conversion method to apply to the video after loading it. When set to `None` the images will be converted
to "RGB".
return_fps (`bool`, *optional*, defaults to `False`):
Whether to return the FPS of the video. If `True`, returns a tuple of (images, fps).
If `False`, returns only the list of images.
Returns:
`List[PIL.Image.Image]`:
The video as a list of PIL images.
`List[PIL.Image.Image]` or `Tuple[List[PIL.Image.Image], float | None]`:
The video as a list of PIL images. If `return_fps` is True, also returns the original FPS.
"""
is_url = video.startswith("http://") or video.startswith("https://")
is_file = os.path.isfile(video)
@@ -175,39 +257,27 @@ def load_video(
video_data = response.iter_content(chunk_size=8192)
for chunk in video_data:
temp_file.write(chunk)
video = video_path
was_tempfile_created = True
else:
video_path = video
pil_images = []
if video.endswith(".gif"):
gif = PIL.Image.open(video)
try:
while True:
pil_images.append(gif.copy())
gif.seek(gif.tell() + 1)
except EOFError:
pass
original_fps = None
else:
try:
imageio.plugins.ffmpeg.get_exe()
except AttributeError:
raise AttributeError(
"`Unable to find an ffmpeg installation on your machine. Please install via `pip install imageio-ffmpeg"
) from None
with imageio.get_reader(video) as reader:
# Read all frames
for frame in reader:
pil_images.append(PIL.Image.fromarray(frame))
if was_tempfile_created:
os.remove(video_path)
try:
if video_path.endswith(".gif"):
pil_images, original_fps = _load_gif(video_path)
else:
pil_images, original_fps = _load_video_with_ffmpeg(video_path)
finally:
# Clean up temporary file if it was created
if was_tempfile_created and os.path.exists(video_path):
os.remove(video_path)
if convert_method is not None:
pil_images = convert_method(pil_images)
return pil_images
return pil_images, original_fps if return_fps else pil_images
def get_default_height_width(
@@ -297,3 +367,50 @@ def resize(
else:
raise ValueError(f"resize_mode {resize_mode} is not supported")
return image
def create_default_image(width: int = 512, height: int = 512, color: tuple[int, int, int] = (0, 0, 0)) -> PIL.Image.Image:
"""
Create a default black PIL image.
Args:
width: Image width in pixels
height: Image height in pixels
color: RGB color tuple
Returns:
PIL.Image.Image: A new PIL image with specified dimensions and color
"""
return PIL.Image.new("RGB", (width, height), color=color)
def preprocess_reference_image_for_clip(image: PIL.Image.Image, device: torch.device) -> PIL.Image.Image:
"""
Preprocess reference image to match CLIP encoder requirements.
Applies normalization, resizing to 224x224, and denormalization to ensure
the image is in the correct format for CLIP processing.
Args:
image: Input PIL image
device: Target device for tensor operations
Returns:
Preprocessed PIL image ready for CLIP encoder
"""
# Convert PIL to tensor and normalize to [-1, 1] range
image_tensor = TF.to_tensor(image).sub_(0.5).div_(0.5).to(device)
# Resize to CLIP's expected input size (224x224) using bicubic interpolation
resized_tensor = F.interpolate(
image_tensor.unsqueeze(0),
size=(224, 224),
mode='bicubic',
align_corners=False
).squeeze(0)
# Denormalize back to [0, 1] range
denormalized_tensor = resized_tensor.mul_(0.5).add_(0.5)
return TF.to_pil_image(denormalized_tensor)
@@ -0,0 +1,83 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan video-to-video diffusion pipeline implementation.
This module contains an implementation of the Wan video-to-video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (
RefImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
VideoVAEEncodingStage, InputValidationStage, LatentPreparationStage,
TextEncodingStage, TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
logger = init_logger(__name__)
class WanVideoToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler", \
"image_encoder", "image_processor"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
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")],
))
if (self.get_module("image_encoder") is not None
and self.get_module("image_processor") is not None):
self.add_stage(
stage_name="ref_image_encoding_stage",
stage=RefImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
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="video_latent_preparation_stage",
stage=VideoVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanVideoToVideoPipeline
+25 -2
View File
@@ -14,12 +14,14 @@ import torch
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.distributed import (
maybe_init_distributed_environment_and_model_parallel)
maybe_init_distributed_environment_and_model_parallel, get_world_group)
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.profiler import get_or_create_profiler
from fastvideo.models.loader.component_loader import PipelineComponentLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages import PipelineStage
import fastvideo.envs as envs
from fastvideo.utils import (maybe_download_model,
verify_model_config_and_directory)
@@ -69,9 +71,18 @@ class ComposedPipelineBase(ABC):
maybe_init_distributed_environment_and_model_parallel(
fastvideo_args.tp_size, fastvideo_args.sp_size)
# Torch profiler. Enabled and configured through env vars:
# FASTVIDEO_TORCH_PROFILER_DIR=/path/to/save/trace
trace_dir = envs.FASTVIDEO_TORCH_PROFILER_DIR
self.profiler_controller = get_or_create_profiler(trace_dir)
self.profiler = self.profiler_controller.profiler
self.local_rank = get_world_group().local_rank
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
self.modules = self.load_modules(fastvideo_args, loaded_modules)
with self.profiler_controller.region("profiler_region_model_loading"):
self.modules = self.load_modules(fastvideo_args, loaded_modules)
def set_trainable(self) -> None:
# Only train DiT
@@ -352,6 +363,18 @@ class ComposedPipelineBase(ABC):
self._stage_name_mapping[stage_name] = stage
setattr(self, stage_name, stage)
def profile(self, is_start: bool = True):
if self.profiler is None:
raise RuntimeError("Profiler is not enabled.")
if is_start:
self.profiler.start()
else:
self.profiler.stop()
# only print profiler results on rank 0
if self.local_rank == 0:
print(self.profiler.key_averages().table(
sort_by="self_cuda_time_total"))
# TODO(will): don't hardcode no_grad
@torch.no_grad()
def forward(
@@ -86,6 +86,11 @@ class ForwardBatch:
prompt_path: str | None = None
output_path: str = "outputs/"
output_video_name: str | None = None
# Video inputs
video_path: str | None = None
video_latent: torch.Tensor | None = None
# Primary encoder embeddings
prompt_embeds: list[torch.Tensor] = field(default_factory=list)
negative_prompt_embeds: list[torch.Tensor] | None = None
+1
View File
@@ -21,6 +21,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanPipeline": "wan",
"WanDMDPipeline": "wan",
"WanImageToVideoPipeline": "wan",
"WanVideoToVideoPipeline": "wan",
"WanCausalDMDPipeline": "wan",
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
+5 -1
View File
@@ -14,7 +14,9 @@ from fastvideo.pipelines.stages.denoising import (DenoisingStage,
DmdDenoisingStage)
from fastvideo.pipelines.stages.encoding import EncodingStage
from fastvideo.pipelines.stages.image_encoding import (ImageEncodingStage,
ImageVAEEncodingStage)
RefImageEncodingStage,
ImageVAEEncodingStage,
VideoVAEEncodingStage)
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.latent_preparation import LatentPreparationStage
from fastvideo.pipelines.stages.stepvideo_encoding import (
@@ -35,7 +37,9 @@ __all__ = [
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
"RefImageEncodingStage",
"ImageVAEEncodingStage",
"VideoVAEEncodingStage",
"TextEncodingStage",
"StepvideoPromptEncodingStage",
]
+24 -1
View File
@@ -212,9 +212,18 @@ class CausalDMDDenosingStage(DenoisingStage):
start_index = 0
# DMD loop in causal blocks
# Optional per-block callback for streaming
on_block = None
try:
on_block = getattr(batch, "extra",
{}).get("on_block",
None) # type: ignore[attr-defined]
except Exception:
on_block = None
with self.progress_bar(total=len(block_sizes) *
len(timesteps)) as progress_bar:
for current_num_frames in block_sizes:
for block_idx, current_num_frames in enumerate(block_sizes):
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
# use BTCHW for DMD conversion routines
@@ -355,6 +364,20 @@ class CausalDMDDenosingStage(DenoisingStage):
)
start_index += current_num_frames
# Invoke callback with block metadata (no large tensor transfer required)
try:
if callable(on_block):
on_block(
block_index=block_idx,
total_blocks=len(block_sizes),
start_index=start_index - current_num_frames,
num_frames=current_num_frames,
latents=current_latents,
)
except Exception as e:
# Swallow callback errors so they don't break generation
logger.warning("on_block callback failed: %s", str(e))
batch.latents = latents
return batch
+10 -4
View File
@@ -283,9 +283,15 @@ class DenoisingStage(PipelineStage):
current_guidance_scale = batch.guidance_scale_2
assert current_model is not None, "current_model is None"
# Expand latents for I2V
# Expand latents for V2V/I2V
latent_model_input = latents.to(target_dtype)
if batch.image_latent is not None:
if batch.video_latent is not None:
latent_model_input = torch.cat([
latent_model_input, batch.video_latent,
torch.zeros_like(latents)
],
dim=1).to(target_dtype)
elif batch.image_latent is not None:
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
latent_model_input = torch.cat(
[latent_model_input, batch.image_latent],
@@ -569,7 +575,7 @@ class DenoisingStage(PipelineStage):
fastvideo_args: The inference arguments.
"""
# TODO(kevin): STA mask search, currently only support Wan2.1 with 69x768x1280
from fastvideo.STA_configuration import configure_sta
from fastvideo.attention.backends.STA_configuration import configure_sta
STA_mode = fastvideo_args.STA_mode
skip_time_steps = fastvideo_args.skip_time_steps
if batch.timesteps is None:
@@ -676,7 +682,7 @@ class DenoisingStage(PipelineStage):
raise NotImplementedError(
"STA mask search is not supported for this resolution")
from fastvideo.STA_configuration import save_mask_search_results
from fastvideo.attention.backends.STA_configuration import save_mask_search_results
if batch.mask_search_final_result_pos is not None and batch.prompt is not None:
save_mask_search_results(
[
+238 -8
View File
@@ -1,8 +1,12 @@
# SPDX-License-Identifier: Apache-2.0
"""
Image encoding stages for I2V diffusion pipelines.
Image and video encoding stages for diffusion pipelines.
This module contains implementations of image encoding stages for diffusion pipelines.
This module contains implementations of encoding stages for diffusion pipelines:
- ImageEncodingStage: Encodes images using image encoders (e.g., CLIP)
- RefImageEncodingStage: Encodes reference image for Wan2.1 control pipeline
- ImageVAEEncodingStage: Encodes images to latent space using VAE for I2V generation
- VideoVAEEncodingStage: Encodes videos to latent space using VAE for V2V and control tasks
"""
import PIL
@@ -14,7 +18,9 @@ from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.vaes.common import ParallelTiledVAE
from fastvideo.models.vision_utils import (get_default_height_width, normalize,
numpy_to_pt, pil_to_numpy, resize)
numpy_to_pt, pil_to_numpy, resize,
create_default_image,
preprocess_reference_image_for_clip)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
@@ -94,12 +100,60 @@ class ImageEncodingStage(PipelineStage):
return result
class RefImageEncodingStage(ImageEncodingStage):
"""
Stage for encoding reference image prompts into embeddings for Wan2.1 Control models.
This stage extends ImageEncodingStage with specialized preprocessing for reference images.
"""
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode the prompt into image encoder hidden states.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded prompt embeddings.
"""
self.image_encoder = self.image_encoder.to(get_local_torch_device())
image = batch.pil_image
if image is None:
image = create_default_image()
# Preprocess reference image for CLIP encoder
image_tensor = preprocess_reference_image_for_clip(
image, get_local_torch_device())
image_inputs = self.image_processor(images=image_tensor,
return_tensors="pt").to(
get_local_torch_device())
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs = self.image_encoder(**image_inputs)
image_embeds = outputs.last_hidden_state
batch.image_embeds.append(image_embeds)
if batch.pil_image is None:
batch.image_embeds = [
torch.zeros_like(x) for x in batch.image_embeds
]
return batch
class ImageVAEEncodingStage(PipelineStage):
"""
Stage for encoding pixel representations into latent space.
This stage handles the encoding of pixel representations into the final
input format (e.g., latents).
Stage for encoding image pixel representations into latent space.
This stage handles the encoding of image pixel representations into the final
input format (e.g., latents) for image-to-video generation.
"""
def __init__(self, vae: ParallelTiledVAE) -> None:
@@ -144,9 +198,9 @@ class ImageVAEEncodingStage(PipelineStage):
self.vae = self.vae.to(get_local_torch_device())
# Process single image for I2V
latent_height = height // self.vae.spatial_compression_ratio
latent_width = width // self.vae.spatial_compression_ratio
image = batch.pil_image
image = self.preprocess(
image,
@@ -296,3 +350,179 @@ class ImageVAEEncodingStage(PipelineStage):
result.add_check("image_latent", batch.image_latent,
[V.is_tensor, V.with_dims(5)])
return result
class VideoVAEEncodingStage(ImageVAEEncodingStage):
"""
Stage for encoding video pixel representations into latent space.
This stage handles the encoding of video pixel representations for video-to-video generation and control.
Inherits from ImageVAEEncodingStage to reuse common functionality.
"""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode video pixel representations into latent space.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded outputs.
"""
assert batch.video_latent is not None, "Video latent input is required for VideoVAEEncodingStage"
if fastvideo_args.mode == ExecutionMode.INFERENCE:
assert batch.height is not None and isinstance(batch.height, int)
assert batch.width is not None and isinstance(batch.width, int)
assert batch.num_frames is not None and isinstance(
batch.num_frames, int)
height = batch.height
width = batch.width
num_frames = batch.num_frames
elif fastvideo_args.mode == ExecutionMode.PREPROCESS:
assert batch.height is not None and isinstance(batch.height, list)
assert batch.width is not None and isinstance(batch.width, list)
assert batch.num_frames is not None and isinstance(
batch.num_frames, list)
num_frames = batch.num_frames[0]
height = batch.height[0]
width = batch.width[0]
self.vae = self.vae.to(get_local_torch_device())
# Prepare video tensor from control video
video_condition = self._prepare_control_video_tensor(
batch.video_latent, num_frames, height,
width).to(get_local_torch_device(), dtype=torch.float32)
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
# Encode control video
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
if not vae_autocast_enabled:
video_condition = video_condition.to(vae_dtype)
encoder_output = self.vae.encode(video_condition)
generator = batch.generator
if generator is None:
raise ValueError("Generator must be provided")
latent_condition = self.retrieve_latents(encoder_output, generator)
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latent_condition -= self.vae.shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.vae.scaling_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition = latent_condition * self.vae.scaling_factor
batch.video_latent = latent_condition
# Offload models if needed
if hasattr(self, 'maybe_free_model_hooks'):
self.maybe_free_model_hooks()
self.vae.to("cpu")
return batch
def _prepare_control_video_tensor(self, control_video, num_frames: int,
height: int, width: int) -> torch.Tensor:
"""
Prepare video tensor from control video input.
"""
if isinstance(control_video, list):
processed_frames = []
for i, frame in enumerate(control_video):
if i >= num_frames:
break
processed_frame = self.preprocess(
frame,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=height,
width=width).to(get_local_torch_device(),
dtype=torch.float32)
processed_frames.append(processed_frame)
if processed_frames:
video_tensor = torch.cat(
[f.unsqueeze(2) for f in processed_frames], dim=2)
else:
video_tensor = torch.zeros(1,
3,
0,
height,
width,
device=get_local_torch_device(),
dtype=torch.float32)
elif isinstance(control_video, torch.Tensor):
# Handle tensor input [batch, channels, frames, height, width]
video_tensor = control_video.to(get_local_torch_device(),
dtype=torch.float32)
if video_tensor.shape[2] > num_frames:
video_tensor = video_tensor[:, :, :num_frames]
else:
raise ValueError(
f"Unsupported control_video type: {type(control_video)}. "
"Expected list of PIL Images or torch.Tensor.")
# Pad with zeros if we have fewer frames than required
current_frames = video_tensor.shape[2]
if current_frames < num_frames:
padding_frames = num_frames - current_frames
zero_padding = torch.zeros(video_tensor.shape[0],
video_tensor.shape[1],
padding_frames,
height,
width,
device=video_tensor.device,
dtype=video_tensor.dtype)
video_tensor = torch.cat([video_tensor, zero_padding], dim=2)
return video_tensor
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify video encoding stage inputs."""
result = VerificationResult()
result.add_check("video_latent", batch.video_latent, V.not_none)
result.add_check("generator", batch.generator,
V.generator_or_list_generators)
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
result.add_check("height", batch.height, V.list_not_empty)
result.add_check("width", batch.width, V.list_not_empty)
result.add_check("num_frames", batch.num_frames, V.list_not_empty)
else:
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
result.add_check("num_frames", batch.num_frames, V.positive_int)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify video encoding stage outputs."""
result = VerificationResult()
result.add_check("video_latent", batch.video_latent,
[V.is_tensor, V.with_dims(5)])
return result
+50 -1
View File
@@ -9,7 +9,7 @@ from PIL import Image
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.vision_utils import load_image, load_video
from fastvideo.models.vision_utils import load_image, load_video, pil_to_numpy, numpy_to_pt, normalize, resize
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import (StageValidators,
@@ -135,6 +135,55 @@ class InputValidationStage(PipelineStage):
batch.width = ow
batch.pil_image = img
# for v2v, get control video from video path
if batch.video_path is not None:
pil_images, original_fps = load_video(batch.video_path,
return_fps=True)
logger.info("Loaded video with %s frames, original FPS: %s",
len(pil_images), original_fps)
# Get target parameters from batch
target_fps = batch.fps
target_num_frames = batch.num_frames
target_height = batch.height
target_width = batch.width
if target_fps is not None and original_fps is not None:
frame_skip = max(1, int(original_fps // target_fps))
if frame_skip > 1:
pil_images = pil_images[::frame_skip]
effective_fps = original_fps / frame_skip
logger.info(
"Resampled video from %.1f fps to %.1f fps (skip=%s)",
original_fps, effective_fps, frame_skip)
# Limit to target number of frames
if target_num_frames is not None and len(
pil_images) > target_num_frames:
pil_images = pil_images[:target_num_frames]
logger.info("Limited video to %s frames (from %s total)",
target_num_frames, len(pil_images))
# Resize each PIL image to target dimensions
resized_images = []
for pil_img in pil_images:
resized_img = resize(pil_img,
target_height,
target_width,
resize_mode="default",
resample="lanczos")
resized_images.append(resized_img)
# Convert PIL images to numpy array
video_numpy = pil_to_numpy(resized_images)
video_numpy = normalize(video_numpy)
video_tensor = numpy_to_pt(video_numpy)
# Rearrange to [C, T, H, W] and add batch dimension -> [B, C, T, H, W]
input_video = video_tensor.permute(1, 0, 2, 3).unsqueeze(0)
batch.video_latent = input_video
return batch
def verify_input(self, batch: ForwardBatch,
+29
View File
@@ -64,6 +64,29 @@ def mps_platform_plugin() -> str | None:
return "fastvideo.platforms.mps.MpsPlatform" if is_mps else None
def npu_platform_plugin() -> str | None:
is_npu = False
try:
import torch
# 导入 torch_npu 以初始化 NPU 后端
import torch_npu # noqa: F401
if torch.npu.is_available():
is_npu = True
logger.info("NPU is available")
except ImportError:
logger.error(
"NPU detection failed: PyTorch or PyTorch_NPU is not installed")
except AttributeError:
logger.error(
"NPU detection failed: PyTorch has no 'npu' attribute (use Ascend-adapted PyTorch)"
)
except Exception as e:
logger.error("NPU detection failed: unknown error - %s", str(e))
return "fastvideo.platforms.npu.NPUPlatform" if is_npu else None
def cpu_platform_plugin() -> str | None:
"""Detect if CPU platform should be used."""
# CPU is always available as a fallback
@@ -93,6 +116,7 @@ builtin_platform_plugins = {
'rocm': rocm_platform_plugin,
'mps': mps_platform_plugin,
'cpu': cpu_platform_plugin,
'npu': npu_platform_plugin,
}
@@ -115,6 +139,11 @@ def resolve_current_platform_cls_qualname() -> str:
if platform_cls_qualname is not None:
return platform_cls_qualname
# Fall back to NPU
platform_cls_qualname = npu_platform_plugin()
if platform_cls_qualname is not None:
return platform_cls_qualname
# Fall back to CPU as last resort
platform_cls_qualname = cpu_platform_plugin()
if platform_cls_qualname is not None:
+13 -1
View File
@@ -67,6 +67,7 @@ class CudaPlatformBase(Platform):
device_name: str = "cuda"
device_type: str = "cuda"
dispatch_key: str = "CUDA"
ray_device_key: str = "GPU"
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
@classmethod
@@ -107,6 +108,13 @@ class CudaPlatformBase(Platform):
torch.cuda.reset_peak_memory_stats(device)
return float(torch.cuda.max_memory_allocated(device))
@classmethod
def get_torch_device(cls):
"""
Return torch.cuda
"""
return torch.cuda
@classmethod
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None,
head_size: int, dtype: torch.dtype) -> str:
@@ -173,7 +181,11 @@ class CudaPlatformBase(Platform):
"Failed to import Video Sparse Attention backend: %s",
str(e))
raise ImportError(
"Video Sparse Attention backend is not installed. ") from e
"The Video Sparse Attention backend is not installed. "
"To install it, please follow the instructions at: "
"https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html "
) from e
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
try:
from csrc.attn.vmoba_attn.vmoba import ( # noqa: F401
+27 -6
View File
@@ -1,6 +1,3 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/platforms/interface.py
import enum
import random
from typing import NamedTuple
@@ -28,10 +25,12 @@ class PlatformEnum(enum.Enum):
CUDA = enum.auto()
ROCM = enum.auto()
TPU = enum.auto()
XPU = enum.auto()
CPU = enum.auto()
MPS = enum.auto()
OOT = enum.auto()
UNSPECIFIED = enum.auto()
NPU = enum.auto()
class CpuArchEnum(enum.Enum):
@@ -62,11 +61,18 @@ class Platform:
device_name: str
device_type: str
# available dispatch keys:
# check https://github.com/pytorch/pytorch/blob/313dac6c1ca0fa0cde32477509cce32089f8532a/torchgen/model.py#L134 # noqa
# use "CPU" as a fallback for platforms not registered in PyTorch
dispatch_key: str = "CPU"
# platform-agnostic way to specify the device control environment variable,
# .e.g. CUDA_VISIBLE_DEVICES for CUDA.
# hint: search for "get_visible_accelerator_ids_env_var" in
# https://github.com/ray-project/ray/tree/master/python/ray/_private/accelerators # noqa
device_control_env_var: str = "FASTVIDEO_DEVICE_CONTROL_ENV_VAR_PLACEHOLDER"
# available ray device keys:
# https://github.com/ray-project/ray/blob/10ba5adadcc49c60af2c358a33bb943fb491a171/python/ray/_private/ray_constants.py#L438 # noqa
# empty string means the device does not support ray
ray_device_key: str = ""
# The torch.compile backend for compiling simple and
# standalone functions. The default value is "inductor" to keep
# the same behavior as PyTorch.
@@ -76,6 +82,8 @@ class Platform:
supported_quantization: list[str] = []
additional_env_vars: list[str] = []
def is_cuda(self) -> bool:
return self._enum == PlatformEnum.CUDA
@@ -85,6 +93,9 @@ class Platform:
def is_tpu(self) -> bool:
return self._enum == PlatformEnum.TPU
def is_xpu(self) -> bool:
return self._enum == PlatformEnum.XPU
def is_cpu(self) -> bool:
return self._enum == PlatformEnum.CPU
@@ -98,6 +109,9 @@ class Platform:
def is_mps(self) -> bool:
return self._enum == PlatformEnum.MPS
def is_npu(self) -> bool:
return self._enum == PlatformEnum.NPU
@classmethod
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None,
head_size: int, dtype: torch.dtype) -> str:
@@ -157,6 +171,13 @@ class Platform:
"""
raise NotImplementedError
@classmethod
def get_torch_device(cls):
"""
Check if the current platform supports torch device.
"""
raise NotImplementedError
@classmethod
def inference_mode(cls):
"""A device-specific wrapper of `torch.inference_mode`.
+131
View File
@@ -0,0 +1,131 @@
import gc
from datetime import timedelta
import torch
from torch.distributed import ProcessGroup
from torch.distributed.distributed_c10d import PrefixStore
import fastvideo.envs as envs
from fastvideo.logger import init_logger
from fastvideo.platforms.interface import (AttentionBackendEnum, Platform,
PlatformEnum)
logger = init_logger(__name__)
class NPUPlatform(Platform):
_enum = PlatformEnum.NPU
device_name: str = "npu"
device_type: str = "npu"
simple_compile_backend: str = "eager" # Disable torch.compile()
ray_device_key: str = "NPU"
device_control_env_var: str = "ASCEND_RT_VISIBLE_DEVICES"
dispatch_key: str = "PrivateUse1"
def is_sleep_mode_available(self) -> bool:
return True
@classmethod
def get_device_capability(cls, device_id: int = 0):
return None
@classmethod
def get_device_name(cls, device_id: int = 0) -> str:
return torch.npu.get_device_name(device_id)
@classmethod
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
return True
@classmethod
def inference_mode(cls):
return torch.inference_mode()
@classmethod
def set_device(cls, device: torch.device):
torch.npu.set_device(device)
@classmethod
def empty_cache(cls):
torch.npu.empty_cache()
@classmethod
def synchronize(cls):
torch.npu.synchronize()
@classmethod
def mem_get_info(cls) -> tuple[int, int]:
return torch.npu.mem_get_info()
@classmethod
def clear_npu_memory(cls):
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@classmethod
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None,
head_size: int, dtype: torch.dtype) -> str:
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
envs.FASTVIDEO_ATTENTION_BACKEND)
if envs.FASTVIDEO_ATTENTION_BACKEND != "TORCH_SDPA":
logger.info("Ascend NPU only supports the Torch SDPA backend.")
else:
logger.info("Using Torch SDPA backend.")
return "fastvideo.attention.backends.sdpa.SDPABackend"
@classmethod
def get_current_memory_usage(cls,
device: torch.types.Device | None = None
) -> float:
torch.npu.reset_peak_memory_stats(device)
return torch.npu.max_memory_allocated(device)
@classmethod
def get_device_communicator_cls(cls) -> str:
return "fastvideo.distributed.device_communicators.npu_communicator.NpuCommunicator"
@classmethod
def is_pin_memory_available(cls):
return True
@classmethod
def get_torch_device(cls):
"""
Return torch.npu
"""
return torch.npu
@classmethod
def stateless_init_device_torch_dist_pg(
cls,
backend: str,
prefix_store: PrefixStore,
group_rank: int,
group_size: int,
timeout: timedelta,
) -> ProcessGroup:
from torch.distributed import is_hccl_available
from torch_npu._C._distributed_c10d import ProcessGroupHCCL
assert is_hccl_available()
options = ProcessGroup.Options(backend=backend)
pg: ProcessGroup = ProcessGroup(
prefix_store,
group_rank,
group_size,
options,
)
backend_options = ProcessGroupHCCL.Options()
backend_options._timeout = timeout
backend_class = ProcessGroupHCCL(prefix_store, group_rank, group_size,
backend_options)
device = torch.device("npu")
backend_class._set_sequence_number_for_group()
backend_type = ProcessGroup.BackendType.CUSTOM
pg._register_backend(device, backend_type, backend_class)
return pg
+1
View File
@@ -22,6 +22,7 @@ class RocmPlatform(Platform):
device_name: str = "rocm"
device_type: str = "cuda" # torch uses 'cuda' backend string
dispatch_key: str = "CUDA"
ray_device_key: str = "GPU"
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
@classmethod
+451
View File
@@ -0,0 +1,451 @@
# SPDX-License-Identifier: Apache-2.0
"""Utilities for managing the PyTorch profiler within FastVideo.
The profiler is shared across the process; this module adds a light-weight
controller that gates collection based on named *regions*. Regions may be
enabled through dedicated environment variables (e.g.
``FASTVIDEO_TORCH_PROFILE_MODEL_LOADING=1``) or via the consolidated
``FASTVIDEO_TORCH_PROFILE_REGIONS`` comma-separated list (e.g.
``FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_dit``).
Typical usage from client code::
controller = TorchProfilerController(profiler, activities)
with controller.region("training_dit"):
run_training_step()
To introduce a new region, register it via :func:`register_profiler_region`
and wrap the corresponding code in :meth:`TorchProfilerController.region`.
"""
from __future__ import annotations
import contextlib
from dataclasses import dataclass
from typing import Any
from collections.abc import Callable
import functools
from collections.abc import Iterable
import torch
import fastvideo.envs as envs
from fastvideo.logger import init_logger
logger = init_logger(__name__)
_GLOBAL_PROFILER: torch.profiler.profile | None = None
_GLOBAL_CONTROLLER: TorchProfilerController | None = None
@dataclass(frozen=True)
class ProfilerRegion:
"""Metadata describing a profiler region."""
name: str
description: str
default_enabled: bool = False
def __post_init__(self) -> None:
if not self.name or self.name.strip() != self.name:
raise ValueError(
f"Profiler region name must be non-empty without surrounding whitespace: {self.name!r}"
)
if not self.name.islower():
raise ValueError(
f"Profiler region name must be lower-case: {self.name!r}")
_REGISTERED_REGIONS: dict[str, ProfilerRegion] = {}
def _normalize_token(token: str) -> str:
return token.strip().lower()
def register_profiler_region(
name: str,
description: str,
*,
default_enabled: bool = False,
) -> None:
"""Register a profiler region so configuration can validate inputs."""
canonical = _normalize_token(name)
if canonical in _REGISTERED_REGIONS:
raise ValueError(f"Profiler region {name!r} is already registered")
region = ProfilerRegion(
name=canonical,
description=description,
default_enabled=bool(default_enabled),
)
_REGISTERED_REGIONS[canonical] = region
def resolve_profiler_region(name: str) -> ProfilerRegion | None:
"""Return the registered region matching ``name`` or ``None`` if absent."""
canonical = _normalize_token(name)
return _REGISTERED_REGIONS.get(canonical)
def list_profiler_regions() -> list[ProfilerRegion]:
"""Return all registered profiler regions sorted by canonical name."""
return [_REGISTERED_REGIONS[name] for name in sorted(_REGISTERED_REGIONS)]
_DEFAULT_ACTIVITIES: tuple[torch.profiler.ProfilerActivity, ...] = (
torch.profiler.ProfilerActivity.CPU,
torch.profiler.ProfilerActivity.CUDA,
)
def get_global_profiler() -> torch.profiler.profile | None:
"""Return the global profiler instance if one was created."""
return _GLOBAL_PROFILER
def set_global_profiler(profiler: torch.profiler.profile | None) -> None:
global _GLOBAL_PROFILER
_GLOBAL_PROFILER = profiler
def get_global_controller() -> TorchProfilerController | None:
return _GLOBAL_CONTROLLER
def set_global_controller(controller: TorchProfilerController | None) -> None:
global _GLOBAL_CONTROLLER
_GLOBAL_CONTROLLER = controller
register_profiler_region(
name="profiler_region_model_loading",
description="Module/model loading during pipeline initialization.",
default_enabled=False,
)
# register_profiler_region(
# name="profiler_region_inference_pre_denoising",
# description="Pre-denoising inference steps (conditioning, preprocessing).",
# )
# register_profiler_region(
# name="profiler_region_inference_denoising",
# description="The main inference denoising loop.",
# )
# register_profiler_region(
# name="profiler_region_inference_post_denoising",
# description=
# "Post-processing after denoising (decoder, conditioning restores).",
# )
register_profiler_region(
name="profiler_region_training_save_checkpoint",
description="Training save checkpoint operations.",
)
# general training related regions
register_profiler_region(
name="profiler_region_training_validation",
description="Validation loop during training.",
)
register_profiler_region(
name="profiler_region_training_train_one_step",
description="High-level step orchestration in the training loop.",
)
register_profiler_region(
name="profiler_region_training_train",
description="Single optimizer step including forward/backward passes.",
)
# distillation specific regions
register_profiler_region(
name="profiler_region_distillation_teacher_forward",
description="Teacher model forward pass in distillation pipelines.",
)
register_profiler_region(
name="profiler_region_distillation_student_forward",
description="Student model forward pass in distillation pipelines.",
)
register_profiler_region(
name="profiler_region_distillation_loss",
description="Distillation loss computation and aggregation.",
)
register_profiler_region(
name="profiler_region_distillation_update",
description="Parameter updates specific to distillation workflows.",
)
def get_or_create_profiler(trace_dir: str | None) -> TorchProfilerController:
"""Create or reuse the process-wide torch profiler controller."""
existing = get_global_controller()
if existing is not None:
if trace_dir:
logger.info("Reusing existing global torch profiler controller")
return existing
if not trace_dir:
logger.info("Torch profiler disabled; returning no-op controller")
return TorchProfilerController(None, _DEFAULT_ACTIVITIES, disabled=True)
logger.info("Profiling enabled. Traces will be saved to: %s", trace_dir)
logger.info(
"Profiler config: record_shapes=%s, profile_memory=%s, with_stack=%s, with_flops=%s",
envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES,
envs.FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY,
envs.FASTVIDEO_TORCH_PROFILER_WITH_STACK,
envs.FASTVIDEO_TORCH_PROFILER_WITH_FLOPS,
)
logger.info("FASTVIDEO_TORCH_PROFILE_REGIONS=%s",
envs.FASTVIDEO_TORCH_PROFILE_REGIONS)
profiler = torch.profiler.profile(
activities=_DEFAULT_ACTIVITIES,
record_shapes=envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES,
profile_memory=envs.FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY,
with_stack=envs.FASTVIDEO_TORCH_PROFILER_WITH_STACK,
with_flops=envs.FASTVIDEO_TORCH_PROFILER_WITH_FLOPS,
on_trace_ready=torch.profiler.tensorboard_trace_handler(trace_dir,
use_gzip=True),
)
controller = TorchProfilerController(profiler, _DEFAULT_ACTIVITIES)
controller.start()
logger.info("Torch profiler started")
return controller
@dataclass
class TorchProfilerConfig:
"""Configuration for torch profiler region control.
Use :meth:`from_env` to construct an instance with defaults inherited from
registered regions and optional overrides from the
``FASTVIDEO_TORCH_PROFILE_REGIONS`` environment variable. The resulting
``regions`` map is consumed by :class:`TorchProfilerController` to decide
when collection should be enabled.
"""
regions: dict[str, bool]
@classmethod
def from_env(cls) -> TorchProfilerConfig:
"""Build a configuration from process environment variables."""
requested_regions = {
token.strip()
for token in (getattr(envs, "FASTVIDEO_TORCH_PROFILE_REGIONS", "")
or "").split(",") if token.strip()
}
if not requested_regions:
available = ", ".join(region.name
for region in list_profiler_regions())
raise ValueError(
"FASTVIDEO_TORCH_PROFILE_REGIONS must list at least one region; "
f"available regions: {available}")
regions: dict[str, bool] = {}
available_regions = list_profiler_regions()
available_names = ", ".join(region.name for region in available_regions)
for token in requested_regions:
resolved = resolve_profiler_region(token)
if resolved is None:
logger.warning(
"Unknown profiler region '%s'; available regions: %s",
token, available_names)
continue
regions[resolved.name] = True
if not regions:
raise ValueError(
"FASTVIDEO_TORCH_PROFILE_REGIONS did not match any known regions; "
f"requested={sorted(requested_regions)}, available={available_names}"
)
return cls(regions=regions)
def __str__(self) -> str:
return f"TorchProfilerConfig(regions={self.regions})"
class TorchProfilerController:
"""Helper that toggles torch profiler collection for named regions.
Parameters
----------
profiler:
The shared :class:`torch.profiler.profile` instance, or ``None`` if
profiling is disabled.
activities:
Iterable of :class:`torch.profiler.ProfilerActivity` recorded by the
profiler.
config:
Optional :class:`TorchProfilerConfig`. If omitted, :meth:`from_env`
constructs one during initialization.
Examples
--------
Enabling an existing region from the command line::
FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_dit \
python fastvideo/training/wan_training_pipeline.py ...
Wrapping a code block in a custom region::
controller = TorchProfilerController(profiler, activities)
with controller.region("training_validation"):
run_validation_epoch()
Adding a new region requires three steps:
1. Define an env var in ``envs.py``.
2. Add a default entry to ``register_profiler_region`` in this module.
3. Wrap the target code in :meth:`region` using the new name.
"""
def __init__(
self,
profiler: Any,
activities: Iterable[torch.profiler.ProfilerActivity],
config: TorchProfilerConfig | None = None,
disabled: bool = False,
) -> None:
activities_tuple = tuple(activities)
existing = get_global_controller()
if existing is not None and not disabled:
raise RuntimeError(
"TorchProfilerController already initialized globally. Use get_global_controller()."
)
if disabled:
self._profiler = None
return
self._profiler = profiler
self._activities = activities_tuple
self._config = config or TorchProfilerConfig.from_env()
self._collection_enabled = False
self._active_region_depth = 0
logger.info(
"PROFILER: TorchProfilerController initialized with config: %s",
self._config)
set_global_profiler(self._profiler)
set_global_controller(self)
@property
def is_enabled(self) -> bool:
"""Return ``True`` when the underlying profiler is collecting."""
if self._profiler is None:
return False
return self._collection_enabled
def is_region_enabled(self, region: str) -> bool:
"""Return ``True`` if ``region`` should be collected."""
if self._profiler is None:
return False
return self._config.regions.get(region, False)
def _set_collection(self, enabled: bool) -> None:
if self._profiler is None:
return
if self._collection_enabled == enabled:
return
event = ("fastvideo.profiler.enable_collection"
if enabled else "fastvideo.profiler.disable_collection")
with torch.profiler.record_function(event):
self._profiler.toggle_collection_dynamic(enabled, self._activities)
self._collection_enabled = enabled
@contextlib.contextmanager
def region(self, region: str):
"""Context manager that enables profiling for ``region`` if configured."""
if self._profiler is None:
yield
return
if not self.is_region_enabled(region):
yield
return
with torch.profiler.record_function(f"fastvideo.region::{region}"):
self._active_region_depth += 1
if self._active_region_depth == 1:
logger.info(
"PROFILER: Setting collection to True (depth=%s) for region %s",
self._active_region_depth, region)
self._set_collection(True)
try:
yield
finally:
self._active_region_depth -= 1
logger.info("PROFILER: Decreasing active region depth to %s",
self._active_region_depth)
if self._active_region_depth == 0:
logger.info(
"PROFILER: Setting collection to False upon exiting region %s",
region)
self._set_collection(False)
def start(self) -> None:
"""Start the profiler and pause collection until a region is entered."""
logger.info("PROFILER: Starting profiler...")
if self._profiler is None:
return
self._profiler.start()
logger.info("PROFILER: Profiler started")
# Profiler starts with collection disabled by default.
logger.info("PROFILER: Setting collection to False")
self._set_collection(False)
logger.info("PROFILER: Profiler started with collection disabled")
def stop(self) -> None:
"""Stop the profiler after disabling collection and clearing state."""
if self._profiler is None:
return
logger.info("PROFILER: Stopping profiler...")
self._profiler.stop()
logger.info("PROFILER: Profiler stopped")
self._active_region_depth = 0
set_global_profiler(None)
set_global_controller(None)
@property
def has_profiler(self) -> bool:
"""Return ``True`` when a profiler instance is available."""
return self._profiler is not None
@property
def activities(self) -> tuple[torch.profiler.ProfilerActivity, ...]:
return tuple(self._activities)
@property
def profiler(self) -> torch.profiler.profile | None:
return self._profiler
def profile_region(
region: str) -> Callable[[Callable[..., Any]], Callable[..., Any]]:
"""Wrap a bound method so it runs inside a profiler region if available."""
def decorator(fn: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(fn)
def wrapped(self, *args, **kwargs):
controller = getattr(self, "profiler_controller", None)
if controller is None or not controller.has_profiler:
return fn(self, *args, **kwargs)
with controller.region(region):
return fn(self, *args, **kwargs)
return wrapped
return decorator
@@ -0,0 +1,78 @@
import os
from fastvideo.entrypoints.video_generator import VideoGenerator
def _new_video_generator() -> VideoGenerator:
# Bypass __init__ since we only test a pure helper method.
return VideoGenerator.__new__(VideoGenerator)
def test_prepare_output_path_file_sanitization(tmp_path):
vg = _new_video_generator()
target_dir = tmp_path / "dir"
raw_path = target_dir / "inv:al*id?.mp4"
result = vg._prepare_output_path(str(raw_path), prompt="ignored")
assert os.path.dirname(result) == str(target_dir)
assert os.path.basename(result) == "invalid.mp4"
assert os.path.isdir(target_dir)
def test_prepare_output_path_directory_prompt_derived(tmp_path):
vg = _new_video_generator()
out_dir = tmp_path / "outputs"
prompt = "Hello:/\\*?<>| world"
result = vg._prepare_output_path(str(out_dir), prompt=prompt)
assert os.path.dirname(result) == str(out_dir)
# spaces are preserved (collapsed) by sanitizer; here it becomes "Hello world.mp4"
assert os.path.basename(result) == "Hello world.mp4"
assert os.path.isdir(out_dir)
def test_prepare_output_path_non_mp4_treated_as_dir(tmp_path):
vg = _new_video_generator()
weird_dir = tmp_path / "foo.gif"
prompt = "My Video"
result = vg._prepare_output_path(str(weird_dir), prompt=prompt)
assert os.path.dirname(result) == str(weird_dir)
assert os.path.basename(result) == "My Video.mp4"
assert os.path.isdir(weird_dir)
def test_prepare_output_path_uniqueness_suffix(tmp_path):
vg = _new_video_generator()
out_dir = tmp_path / "outputs"
prompt = "Sample Name"
first = vg._prepare_output_path(str(out_dir), prompt=prompt)
# simulate existing file
os.makedirs(os.path.dirname(first), exist_ok=True)
with open(first, "wb") as f:
f.write(b"")
second = vg._prepare_output_path(str(out_dir), prompt=prompt)
assert os.path.basename(second) == "Sample Name_1.mp4"
# simulate second existing file as well
with open(second, "wb") as f:
f.write(b"")
third = vg._prepare_output_path(str(out_dir), prompt=prompt)
assert os.path.basename(third) == "Sample Name_2.mp4"
def test_prepare_output_path_empty_prompt_fallback(tmp_path):
vg = _new_video_generator()
out_dir = tmp_path / "outputs"
bad_prompt = ":/\\*?<>| .." # sanitizes to empty, should fallback to "video"
result = vg._prepare_output_path(str(out_dir), prompt=bad_prompt)
assert os.path.dirname(result) == str(out_dir)
assert os.path.basename(result) == "video.mp4"
+1 -1
View File
@@ -124,4 +124,4 @@ def run_self_forcing_tests():
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_unit_test():
run_test("pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ -vs")
run_test("pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ -vs")
@@ -19,6 +19,9 @@ if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
raise ValueError(f"Unsupported device for ssim tests: {device_name}")
# Base parameters from the shell script
@@ -69,7 +72,7 @@ def test_causal_similarity(prompt, ATTENTION_BACKEND, model_id):
base_output_dir = os.path.join(script_dir, 'generated_videos', model_id)
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
output_video_name = f"{prompt[:100]}.mp4"
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
@@ -119,7 +122,7 @@ def test_causal_similarity(prompt, ATTENTION_BACKEND, model_id):
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith('.mp4') and prompt[:100] in filename:
if filename.endswith('.mp4') and prompt[:100].strip() in filename:
reference_video_name = filename
break
@@ -19,6 +19,9 @@ if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
raise ValueError(f"Unsupported device for ssim tests: {device_name}")
# Base parameters from the shell script
HUNYUAN_PARAMS = {
@@ -115,7 +118,7 @@ def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
base_output_dir = os.path.join(script_dir, 'generated_videos', model_id)
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
output_video_name = f"{prompt[:100]}.mp4"
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
@@ -170,7 +173,7 @@ def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith('.mp4') and prompt[:100] in filename:
if filename.endswith('.mp4') and prompt[:100].strip() in filename:
reference_video_name = filename
break
@@ -216,7 +219,7 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
base_output_dir = os.path.join(script_dir, 'generated_videos', model_id)
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
output_video_name = f"{prompt[:100]}.mp4"
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
@@ -270,7 +273,7 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith('.mp4') and prompt[:100] in filename:
if filename.endswith('.mp4') and prompt[:100].strip() in filename:
reference_video_name = filename
break
@@ -102,7 +102,7 @@ def test_distributed_training():
process = subprocess.run(cmd, check=True)
summary_file = 'wandb/latest-run/files/wandb-summary.json'
summary_file = 'data/wan_finetune_test_VSA/tracker/wandb/latest-run/files/wandb-summary.json'
reference_wandb_summary = json.load(open(reference_wandb_summary_file))
wandb_summary = json.load(open(summary_file))
@@ -119,7 +119,7 @@ def test_distributed_training():
print(f"Process failed with return code: {process.returncode}")
raise subprocess.CalledProcessError(process.returncode, cmd, process.stdout, process.stderr)
summary_file = 'wandb/latest-run/files/wandb-summary.json'
summary_file = 'data/wan_finetune_test/tracker/wandb/latest-run/files/wandb-summary.json'
device_name = torch.cuda.get_device_name()
if "A40" in device_name:
@@ -83,7 +83,7 @@ def test_lora_training():
process = subprocess.run(cmd, check=True)
summary_file = 'wandb/latest-run/files/wandb-summary.json'
summary_file = '/workspace/tracker/wandb/latest-run/files/wandb-summary.json'
device_name = torch.cuda.get_device_name()
assert "L40S" in device_name, "Test must be run on L40S"
+32 -19
View File
@@ -43,8 +43,6 @@ from fastvideo.training.training_utils import (
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
import wandb # isort: skip
vsa_available = is_vsa_available()
logger = init_logger(__name__)
@@ -1342,14 +1340,20 @@ class DistillationPipeline(TrainingPipeline):
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
logs = {
f"validation_videos_{num_inference_steps}_steps": [
wandb.Video(filename, caption=caption)
for filename, caption in zip(
video_filenames, all_captions, strict=True)
]
}
wandb.log(logs, step=global_step)
artifacts = []
for filename, caption in zip(video_filenames,
all_captions,
strict=True):
video_artifact = self.tracker.video(filename,
caption=caption)
if video_artifact is not None:
artifacts.append(video_artifact)
if artifacts:
logs = {
f"validation_videos_{num_inference_steps}_steps":
artifacts
}
self.tracker.log_artifacts(logs, global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(step_videos, dst=0)
@@ -1363,8 +1367,8 @@ class DistillationPipeline(TrainingPipeline):
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
"""Add visualization data to wandb logging and save frames to disk."""
wandb_loss_dict = {}
"""Add visualization data to tracker logging and save frames to disk."""
tracker_loss_dict: dict[str, Any] = {}
dmd_latents_vis_dict = training_batch.dmd_latent_vis_dict
fake_score_latents_vis_dict = training_batch.fake_score_latent_vis_dict
fake_score_log_keys = ['generator_pred_video']
@@ -1394,8 +1398,10 @@ class DistillationPipeline(TrainingPipeline):
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[latent_key] = wandb.Video(
video_artifact = self.tracker.video(
video, fps=24, format="mp4") # change to 16 for Wan2.1
if video_artifact is not None:
tracker_loss_dict[latent_key] = video_artifact
# Clean up references
del video, latents
@@ -1425,14 +1431,16 @@ class DistillationPipeline(TrainingPipeline):
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[latent_key] = wandb.Video(
video_artifact = self.tracker.video(
video, fps=24, format="mp4") # change to 16 for Wan2.1
if video_artifact is not None:
tracker_loss_dict[latent_key] = video_artifact
# Clean up references
del video, latents
# Log to wandb
if self.global_rank == 0:
wandb.log(wandb_loss_dict, step=step)
# Log to tracker
if self.global_rank == 0 and tracker_loss_dict:
self.tracker.log_artifacts(tracker_loss_dict, step)
def train(self) -> None:
"""Main training loop with distillation-specific logging."""
@@ -1605,7 +1613,7 @@ class DistillationPipeline(TrainingPipeline):
}
log_data.update(faker_score_additional_logs)
wandb.log(log_data, step=step)
self.tracker.log(log_data, step)
# Save training state checkpoint (for resuming training)
if (self.training_args.training_state_checkpointing_steps > 0
@@ -1688,7 +1696,7 @@ class DistillationPipeline(TrainingPipeline):
step)
self._log_validation(self.transformer, self.training_args, step)
wandb.finish()
self.tracker.finish()
# Save final training state checkpoint
print("rank", self.global_rank,
@@ -1725,5 +1733,10 @@ class DistillationPipeline(TrainingPipeline):
self.save_ema_weights(self.training_args.output_dir,
self.training_args.max_train_steps)
if envs.FASTVIDEO_TORCH_PROFILER_DIR:
logger.info("Stopping profiler...")
self.profiler_controller.stop()
logger.info("Profiler stopped.")
if get_sp_group():
cleanup_dist_env_and_memory()
+7 -8
View File
@@ -1,13 +1,12 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import cast
from typing import Any, cast
import numpy as np
import torch
import torch.nn.functional as F
import wandb
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
from fastvideo.distributed import get_local_torch_device
@@ -361,8 +360,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
"""Add visualization data to wandb logging and save frames to disk."""
wandb_loss_dict = {}
tracker_loss_dict: dict[str, Any] = {}
latents_vis_dict = training_batch.latent_vis_dict
latent_log_keys = ['noisy_input', 'x0', 'pred_video']
for latent_key in latent_log_keys:
@@ -375,14 +373,15 @@ class ODEInitTrainingPipeline(TrainingPipeline):
video = pixel_latent.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[latent_key] = wandb.Video(
video_artifact = self.tracker.video(
video, fps=16, format="mp4") # change to 16 for Wan2.1
if video_artifact is not None:
tracker_loss_dict[latent_key] = video_artifact
# Clean up references
del video, pixel_latent, latent
# Log to wandb
if self.global_rank == 0:
wandb.log(wandb_loss_dict, step=step)
if self.global_rank == 0 and tracker_loss_dict:
self.tracker.log_artifacts(tracker_loss_dict, step)
def main(args) -> None:
@@ -11,7 +11,6 @@ from einops import rearrange
from tqdm.auto import tqdm
import fastvideo.envs as envs
import wandb
from fastvideo.distributed import (cleanup_dist_env_and_memory,
get_local_torch_device, get_sp_group,
get_world_group)
@@ -26,6 +25,7 @@ from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.training.training_utils import (EMA_FSDP,
save_distillation_checkpoint)
from fastvideo.utils import is_vsa_available, set_random_seed
from fastvideo.profiler import profile_region
logger = init_logger(__name__)
@@ -814,8 +814,8 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
"""Add visualization data to wandb logging and save frames to disk."""
wandb_loss_dict = {}
"""Add visualization data to tracker logging and save frames to disk."""
tracker_loss_dict: dict[str, Any] = {}
# Debug logging
if hasattr(training_batch, 'dmd_latent_vis_dict'):
@@ -867,8 +867,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[f"dmd_{latent_key}"] = wandb.Video(
video, fps=24, format="mp4")
video_artifact = self.tracker.video(video,
fps=24,
format="mp4")
if video_artifact is not None:
tracker_loss_dict[f"dmd_{latent_key}"] = video_artifact
del video, latents
# Process critic predictions
@@ -909,8 +912,12 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[f"critic_{latent_key}"] = wandb.Video(
video, fps=24, format="mp4")
video_artifact = self.tracker.video(video,
fps=24,
format="mp4")
if video_artifact is not None:
tracker_loss_dict[
f"critic_{latent_key}"] = video_artifact
del video, latents
# Log metadata
@@ -918,28 +925,29 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
training_batch,
'dmd_latent_vis_dict') and training_batch.dmd_latent_vis_dict:
if "generator_timestep" in training_batch.dmd_latent_vis_dict:
wandb_loss_dict[
tracker_loss_dict[
"generator_timestep"] = training_batch.dmd_latent_vis_dict[
"generator_timestep"].item()
if "dmd_timestep" in training_batch.dmd_latent_vis_dict:
wandb_loss_dict[
tracker_loss_dict[
"dmd_timestep"] = training_batch.dmd_latent_vis_dict[
"dmd_timestep"].item()
if hasattr(
training_batch, 'fake_score_latent_vis_dict'
) and training_batch.fake_score_latent_vis_dict and "fake_score_timestep" in training_batch.fake_score_latent_vis_dict:
wandb_loss_dict[
tracker_loss_dict[
"fake_score_timestep"] = training_batch.fake_score_latent_vis_dict[
"fake_score_timestep"].item()
# Log final dict contents
logger.info("Final wandb_loss_dict keys: %s",
list(wandb_loss_dict.keys()))
logger.info("Final tracker_loss_dict keys: %s",
list(tracker_loss_dict.keys()))
if self.global_rank == 0:
wandb.log(wandb_loss_dict, step=step)
if self.global_rank == 0 and tracker_loss_dict:
self.tracker.log_artifacts(tracker_loss_dict, step)
@profile_region("profiler_region_training_train")
def train(self) -> None:
"""Main training loop with self-forcing specific logging."""
assert self.training_args.seed is not None, "seed must be set"
@@ -1103,7 +1111,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
}
log_data.update(faker_score_additional_logs)
wandb.log(log_data, step=step)
self.tracker.log(log_data, step)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0 and self.training_args.log_visualization:
self.visualize_intermediate_latents(training_batch,
@@ -1185,7 +1193,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self._log_validation(self.transformer, self.training_args, step)
wandb.finish()
self.tracker.finish()
print("rank", self.global_rank,
"save final training state checkpoint at step",
@@ -1221,5 +1229,10 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
self.save_ema_weights(self.training_args.output_dir,
self.training_args.max_train_steps)
if envs.FASTVIDEO_TORCH_PROFILER_DIR:
logger.info("Stopping profiler...")
self.profiler_controller.stop()
logger.info("Profiler stopped.")
if get_sp_group():
cleanup_dist_env_and_memory()
+275
View File
@@ -0,0 +1,275 @@
"""Utilities for logging metrics and artifacts to external trackers.
This module is inspired by the trackers implementation in
https://github.com/huggingface/finetrainers and provides a minimal, shared
interface that can be used across all FastVideo training pipelines.
"""
from __future__ import annotations
import contextlib
import copy
import os
import pathlib
import time
from dataclasses import dataclass
from enum import Enum
from typing import Any
from collections.abc import Iterable, Iterator
from fastvideo.logger import init_logger
logger = init_logger(__name__)
@dataclass
class Timer:
"""Simple timer utility used by the trackers."""
name: str
_start_time: float | None = None
_end_time: float | None = None
def start(self) -> None:
self._start_time = time.perf_counter()
def end(self) -> None:
self._end_time = time.perf_counter()
@property
def elapsed_time(self) -> float:
if self._start_time is None:
raise RuntimeError(
"Timer.start() must be called before elapsed_time")
end_time = self._end_time if self._end_time is not None else time.perf_counter(
)
return end_time - self._start_time
class BaseTracker:
"""Base tracker implementation.
The default tracker stores timing information but does not emit any logs.
"""
def __init__(self) -> None:
self._timed_metrics: dict[str, float] = {}
@contextlib.contextmanager
def timed(
self,
name: str,
) -> Iterator[Timer]:
timer = Timer(name)
timer.start()
try:
yield timer
finally:
timer.end()
elapsed_time = timer.elapsed_time
if name in self._timed_metrics:
self._timed_metrics[name] += elapsed_time
else:
self._timed_metrics[name] = elapsed_time
def log(self, metrics: dict[str, Any],
step: int) -> None: # pragma: no cover - interface
"""Log metrics for the given step."""
# Merge timing metrics with provided metrics
metrics = {**self._timed_metrics, **metrics}
self._timed_metrics = {}
def log_artifacts(self, artifacts: dict[str, Any], step: int) -> None:
"""Log artifacts such as videos or images.
By default this is treated the same as :meth:`log`.
"""
if artifacts:
self.log(artifacts, step)
def finish(self) -> None: # pragma: no cover - interface
"""Finalize the tracker session."""
def video(
self,
data: Any,
*,
caption: str | None = None,
fps: int | None = None,
format: str | None = None,
) -> Any | None:
"""Create a tracker specific video artifact.
Trackers that do not support video artifacts should return ``None``.
"""
return None
class DummyTracker(BaseTracker):
"""Tracker implementation used when logging is disabled."""
def log(self, metrics: dict[str, Any],
step: int) -> None: # pragma: no cover - no-op
super().log(metrics, step)
def finish(self) -> None: # pragma: no cover - no-op
pass
class WandbTracker(BaseTracker):
"""Tracker implementation for Weights & Biases."""
def __init__(
self,
experiment_name: str,
log_dir: str,
*,
config: dict[str, Any] | None = None,
run_name: str | None = None,
) -> None:
super().__init__()
import wandb
pathlib.Path(log_dir).mkdir(parents=True, exist_ok=True)
self._wandb = wandb
self._run = wandb.init(
project=experiment_name,
dir=log_dir,
config=config,
name=run_name,
)
logger.info("Initialized Weights & Biases tracker")
def log(self, metrics: dict[str, Any], step: int) -> None:
metrics = {**self._timed_metrics, **metrics}
if metrics:
self._run.log(metrics, step=step)
self._timed_metrics = {}
def finish(self) -> None:
self._run.finish()
def video(
self,
data: Any,
*,
caption: str | None = None,
fps: int | None = None,
format: str | None = None,
) -> Any:
kwargs: dict[str, Any] = {}
if caption is not None:
kwargs["caption"] = caption
if fps is not None:
kwargs["fps"] = fps
if format is not None:
kwargs["format"] = format
else:
kwargs["format"] = "mp4"
return self._wandb.Video(data, **kwargs)
class SequentialTracker(BaseTracker):
"""A tracker that forwards logging calls to a sequence of trackers."""
def __init__(self, trackers: Iterable[BaseTracker]) -> None:
super().__init__()
self._trackers: list[BaseTracker] = list(trackers)
@contextlib.contextmanager
def timed(
self,
name: str,
) -> Iterator[Timer]:
with super().timed(name) as timer:
yield timer
for tracker in self._trackers:
tracker._timed_metrics = copy.deepcopy(self._timed_metrics)
def log(self, metrics: dict[str, Any], step: int) -> None:
for tracker in self._trackers:
tracker.log({**self._timed_metrics, **metrics}, step)
self._timed_metrics = {}
def log_artifacts(self, artifacts: dict[str, Any], step: int) -> None:
for tracker in self._trackers:
tracker.log_artifacts(artifacts, step)
self._timed_metrics = {}
def finish(self) -> None:
for tracker in self._trackers:
tracker.finish()
def video(
self,
data: Any,
*,
caption: str | None = None,
fps: int | None = None,
format: str | None = None,
) -> Any | None:
for tracker in self._trackers:
video = tracker.video(data, caption=caption, fps=fps, format=format)
if video is not None:
return video
return None
class Trackers(str, Enum):
NONE = "none"
WANDB = "wandb"
SUPPORTED_TRACKERS = {tracker.value for tracker in Trackers}
def initialize_trackers(
trackers: Iterable[str],
*,
experiment_name: str,
config: dict[str, Any] | None,
log_dir: str,
run_name: str | None = None,
) -> BaseTracker:
"""Create tracker instances based on ``trackers`` configuration."""
tracker_names = [tracker.lower() for tracker in trackers]
if not tracker_names:
return DummyTracker()
unsupported = [
name for name in tracker_names if name not in SUPPORTED_TRACKERS
]
if unsupported:
raise ValueError(
f"Unsupported tracker(s) provided: {unsupported}. Supported trackers: {sorted(SUPPORTED_TRACKERS)}"
)
tracker_instances: list[BaseTracker] = []
for tracker_name in tracker_names:
if tracker_name == Trackers.NONE.value:
tracker_instances.append(DummyTracker())
elif tracker_name == Trackers.WANDB.value:
tracker_instances.append(
WandbTracker(
experiment_name,
os.path.abspath(log_dir),
config=config,
run_name=run_name,
))
if not tracker_instances:
return DummyTracker()
if len(tracker_instances) == 1:
return tracker_instances[0]
return SequentialTracker(tracker_instances)
TrackerType = BaseTracker
+174 -129
View File
@@ -1,5 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
import dataclasses
from dataclasses import asdict
import math
import os
import time
@@ -7,7 +7,7 @@ from abc import ABC, abstractmethod
from collections import deque
from collections.abc import Iterator
from typing import Any
from fastvideo.profiler import profile_region
import imageio
import numpy as np
import torch
@@ -35,8 +35,11 @@ from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.pipelines import (ComposedPipelineBase, ForwardBatch,
LoRAPipeline, TrainingBatch)
from fastvideo.platforms import current_platform
from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.training.trackers import (DummyTracker, TrackerType,
initialize_trackers, Trackers)
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, count_trainable, get_scheduler,
@@ -45,8 +48,6 @@ from fastvideo.training.training_utils import (
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
set_random_seed, shallow_asdict)
import wandb # isort: skip
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
@@ -64,6 +65,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
train_loader_iter: Iterator[dict[str, Any]]
current_epoch: int = 0
train_transformer_2: bool = False
tracker: TrackerType
def __init__(
self,
@@ -79,6 +81,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
set_random_seed(fastvideo_args.seed) # for lora param init
super().__init__(model_path, fastvideo_args, required_config_modules,
loaded_modules) # type: ignore
self.tracker = DummyTracker()
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
raise RuntimeError(
@@ -208,12 +211,26 @@ class TrainingPipeline(LoRAPipeline, ABC):
# TODO(will): is there a cleaner way to track epochs?
self.current_epoch = 0
if self.global_rank == 0:
project = training_args.tracker_project_name or "fastvideo"
wandb_config = dataclasses.asdict(training_args)
wandb.init(project=project,
config=wandb_config,
name=training_args.wandb_run_name)
trackers = list(training_args.trackers)
if not trackers and training_args.tracker_project_name:
trackers.append(Trackers.WANDB.value)
if self.global_rank != 0:
trackers = []
tracker_log_dir = training_args.output_dir or os.getcwd()
if trackers:
tracker_log_dir = os.path.join(tracker_log_dir, "tracker")
tracker_config = asdict(training_args) if trackers else None
tracker_run_name = training_args.wandb_run_name or None
project = training_args.tracker_project_name or "fastvideo"
self.tracker = initialize_trackers(
trackers,
experiment_name=project,
config=tracker_config,
log_dir=tracker_log_dir,
run_name=tracker_run_name,
)
@abstractmethod
def initialize_validation_pipeline(self, training_args: TrainingArgs):
@@ -245,78 +262,83 @@ class TrainingPipeline(LoRAPipeline, ABC):
optimizer.zero_grad(set_to_none=True)
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
with self.tracker.timed("timing/get_next_batch"):
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
# latents, encoder_hidden_states, encoder_attention_mask, infos = batch
latents = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
infos = batch['info_list']
latents = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
infos = batch['info_list']
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.infos = infos
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.infos = infos
return training_batch
def _normalize_dit_input(self,
training_batch: TrainingBatch) -> TrainingBatch:
# TODO(will): support other models
training_batch.latents = normalize_dit_input('wan',
training_batch.latents,
self.get_module("vae"))
with self.tracker.timed("timing/normalize_input"):
training_batch.latents = normalize_dit_input(
'wan',
training_batch.latents,
self.get_module("vae"),
)
return training_batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
latents = training_batch.latents
batch_size = latents.shape[0]
noise = torch.randn(latents.shape,
generator=self.noise_gen_cuda,
device=latents.device,
dtype=latents.dtype)
timesteps = self._sample_timesteps(batch_size, latents.device)
assert self.training_args is not None, "training_args must be set"
with self.tracker.timed("timing/prepare_dit_inputs"):
latents = training_batch.latents
batch_size = latents.shape[0]
noise = torch.randn(latents.shape,
generator=self.noise_gen_cuda,
device=latents.device,
dtype=latents.dtype)
timesteps = self._sample_timesteps(batch_size, latents.device)
# Enable training for the model that will be trained next and disable the other
if self.train_transformer_2:
self._enable_training(self.transformer_2, self.optimizer_2)
self._disable_training(self.transformer, self.optimizer)
else:
self._enable_training(self.transformer, self.optimizer)
if self.transformer_2 is not None:
self._disable_training(self.transformer_2, self.optimizer_2)
# Enable training for the model that will be trained next and disable the other
if self.train_transformer_2:
self._enable_training(self.transformer_2, self.optimizer_2)
self._disable_training(self.transformer, self.optimizer)
else:
self._enable_training(self.transformer, self.optimizer)
if self.transformer_2 is not None:
self._disable_training(self.transformer_2, self.optimizer_2)
if self.training_args.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(
self.noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = (1.0 -
sigmas) * training_batch.latents + sigmas * noise
if self.training_args.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(
self.noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = (
1.0 - sigmas) * training_batch.latents + sigmas * noise
training_batch.noisy_model_input = noisy_model_input
training_batch.timesteps = timesteps
training_batch.sigmas = sigmas
training_batch.noise = noise
training_batch.raw_latent_shape = training_batch.latents.shape
training_batch.noisy_model_input = noisy_model_input
training_batch.timesteps = timesteps
training_batch.sigmas = sigmas
training_batch.noise = noise
training_batch.raw_latent_shape = training_batch.latents.shape
return training_batch
@@ -425,7 +447,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
# dtype=torch.bfloat16)
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
with set_forward_context(
with self.tracker.timed("timing/forward_backward"), set_forward_context(
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
model_pred = current_model(**input_kwargs)
@@ -446,8 +468,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
# local_main_process_only=False)
world_group = get_world_group()
world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
with self.tracker.timed("timing/reduce_loss"):
world_group = get_world_group()
world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
training_batch.total_loss += avg_loss.item()
return training_batch
@@ -459,25 +482,27 @@ class TrainingPipeline(LoRAPipeline, ABC):
# the following:
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
if max_grad_norm is not None:
# Only clip gradients for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
model_parts = [self.transformer_2]
else:
model_parts = [self.transformer]
with self.tracker.timed("timing/clip_grad_norm"):
# Only clip gradients for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
model_parts = [self.transformer_2]
else:
model_parts = [self.transformer]
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,
)
assert grad_norm is not float('nan') or grad_norm is not float(
'inf')
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
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,
)
assert grad_norm is not float('nan') or grad_norm is not float(
'inf')
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
else:
grad_norm = 0.0
training_batch.grad_norm = grad_norm
return training_batch
@profile_region("profiler_region_training_train_one_step")
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
training_batch = self._prepare_training(training_batch)
@@ -510,12 +535,13 @@ class TrainingPipeline(LoRAPipeline, ABC):
training_batch = self._clip_grad_norm(training_batch)
# Only step the optimizer and scheduler for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
self.optimizer_2.step()
self.lr_scheduler_2.step()
else:
self.optimizer.step()
self.lr_scheduler.step()
with self.tracker.timed("timing/optimizer_step"):
if self.train_transformer_2 and self.transformer_2 is not None:
self.optimizer_2.step()
self.lr_scheduler_2.step()
else:
self.optimizer.step()
self.lr_scheduler.step()
training_batch.total_loss = training_batch.total_loss
training_batch.grad_norm = training_batch.grad_norm
@@ -536,8 +562,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
logger.warning("Failed to load checkpoint, starting from step 0")
self.init_steps = 0
@profile_region("profiler_region_training_train")
def train(self) -> None:
assert self.seed is not None, "seed must be set"
assert self.training_args is not None, "training_args must be set"
set_random_seed(self.seed + self.global_rank)
logger.info('rank: %s: start training',
self.global_rank,
@@ -557,8 +585,8 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
self.noise_gen_cuda = torch.Generator(device="cuda").manual_seed(
self.seed)
self.noise_gen_cuda = torch.Generator(
device=current_platform.device_name).manual_seed(self.seed)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
@@ -620,48 +648,58 @@ class TrainingPipeline(LoRAPipeline, ABC):
})
progress_bar.update(1)
if self.global_rank == 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": self.lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
},
step=step,
)
metrics = {
"train_loss": loss,
"learning_rate": self.lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
}
self.tracker.log(metrics, step)
if step % self.training_args.training_state_checkpointing_steps == 0:
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir, step,
self.optimizer, self.train_dataloader,
self.lr_scheduler, self.noise_random_generator)
with self.profiler_controller.region(
"profiler_region_training_save_checkpoint"):
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir, step,
self.optimizer, self.train_dataloader,
self.lr_scheduler,
self.noise_random_generator)
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
if self.training_args.log_visualization:
self.visualize_intermediate_latents(training_batch,
self.training_args,
step)
self._log_validation(self.transformer, self.training_args, step)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
trainable_params = round(
count_trainable(self.transformer) / 1e9, 3)
logger.info(
"GPU memory usage after validation: %s MB, trainable params: %sB",
gpu_memory_usage, trainable_params)
with self.profiler_controller.region(
"profiler_region_training_validation"):
if self.training_args.log_visualization:
self.visualize_intermediate_latents(
training_batch, self.training_args, step)
self._log_validation(self.transformer, self.training_args,
step)
gpu_memory_usage = current_platform.get_torch_device(
).memory_allocated() / 1024**2
trainable_params = round(
count_trainable(self.transformer) / 1e9, 3)
logger.info(
"GPU memory usage after validation: %s MB, trainable params: %sB",
gpu_memory_usage, trainable_params)
wandb.finish()
self.tracker.finish()
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir,
self.training_args.max_train_steps, self.optimizer,
self.train_dataloader, self.lr_scheduler,
self.noise_random_generator)
if envs.FASTVIDEO_TORCH_PROFILER_DIR:
logger.info("Stopping profiler...")
self.profiler_controller.stop()
logger.info("Profiler stopped.")
if get_sp_group():
cleanup_dist_env_and_memory()
def _log_training_info(self) -> None:
assert self.training_args is not None, "training_args must be set"
total_batch_size = (self.world_size *
self.training_args.gradient_accumulation_steps /
self.training_args.sp_size *
@@ -687,7 +725,8 @@ class TrainingPipeline(LoRAPipeline, ABC):
logger.info(" Master weight dtype: %s",
self.transformer.parameters().__next__().dtype)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
gpu_memory_usage = current_platform.get_torch_device().memory_allocated(
) / 1024**2
logger.info("GPU memory usage before train_one_step: %s MB",
gpu_memory_usage)
logger.info("VSA validation sparsity: %s",
@@ -726,7 +765,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
"""
Generate a validation video and log it to wandb to check the quality during training.
Generate a validation video and log it to the configured tracker to check the quality during training.
"""
training_args.inference_mode = True
training_args.dit_cpu_offload = False
@@ -830,14 +869,20 @@ class TrainingPipeline(LoRAPipeline, ABC):
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
logs = {
f"validation_videos_{num_inference_steps}_steps": [
wandb.Video(filename, caption=caption)
for filename, caption in zip(
video_filenames, all_captions, strict=True)
]
}
wandb.log(logs, step=global_step)
artifacts = []
for filename, caption in zip(video_filenames,
all_captions,
strict=True):
video_artifact = self.tracker.video(filename,
caption=caption)
if video_artifact is not None:
artifacts.append(video_artifact)
if artifacts:
logs = {
f"validation_videos_{num_inference_steps}_steps":
artifacts
}
self.tracker.log_artifacts(logs, global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(step_videos, dst=0)
@@ -851,7 +896,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
"""Add visualization data to wandb logging and save frames to disk."""
"""Add visualization data to tracker logging and save frames to disk."""
raise NotImplementedError(
"Visualize intermediate latents is not implemented for training pipeline"
)
+284 -4
View File
@@ -7,26 +7,31 @@ import hashlib
import importlib
import importlib.util
import inspect
import ipaddress
import json
import math
import multiprocessing
from multiprocessing.context import BaseContext
import os
import signal
import socket
import sys
import tempfile
import warnings
import threading
import traceback
from collections.abc import Callable
from dataclasses import dataclass, fields, is_dataclass
from functools import lru_cache, partial, wraps
from typing import Any, TypeVar, cast
from pathlib import Path
from typing import Any, TextIO, TypeVar, cast
import cloudpickle
import filelock
import imageio
import numpy as np
import torch
import torchvision
import torchvision.utils as make_grid
import yaml
from diffusers.loaders.lora_base import (
_best_guess_weight_name) # watch out for potetential removal from diffusers
@@ -56,7 +61,7 @@ STR_ATTN_CONFIG_ENV_VAR: str = "FASTVIDEO_ATTENTION_CONFIG"
def find_nccl_library() -> str:
"""
We either use the library file specified by the `VLLM_NCCL_SO_PATH`
We either use the library file specified by the `FASTVIDEO_NCCL_SO_PATH`
environment variable, or we find the library file brought by PyTorch.
After importing `torch`, `libnccl.so.2` or `librccl.so.1` can be
found by `ctypes` automatically.
@@ -79,6 +84,28 @@ def find_nccl_library() -> str:
return str(so_file)
def find_hccl_library() -> str:
"""
We either use the library file specified by the `HCCL_SO_PATH`
environment variable, or we find the library file brought by PyTorch.
After importing `torch`, `libhccl.so` can be
found by `ctypes` automatically.
"""
so_file = envs.HCCL_SO_PATH
# manually load the nccl library
if so_file:
logger.info("Found hccl from environment variable HCCL_SO_PATH=%s",
so_file)
else:
if torch.version.cann is not None: # codespell:ignore cann
so_file = "libhccl.so"
else:
raise ValueError("HCCL only supports Ascend NPU backends.")
logger.info("Found hccl from library %s", so_file)
return so_file
prev_set_stream = torch.cuda.set_stream
_current_stream = None
@@ -898,9 +925,262 @@ def save_decoded_latents_as_video(decoded_latents: list[torch.Tensor],
videos = rearrange(decoded_latents, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
os.makedirs(os.path.dirname(output_path), exist_ok=True)
imageio.mimsave(output_path, frames, fps=fps, format="mp4")
def _format_bytes(num_bytes: int | float | None) -> str:
if num_bytes is None:
return "N/A"
return f"{num_bytes / (1024 ** 3):.2f} GB"
def log_torch_cuda_memory(
tag: str | None = None,
*,
log_fn: Callable[[str], None] | None = None,
log_file_path: str | os.PathLike[str] | None = "memory_trace.txt"
) -> None:
"""Log CUDA memory statistics via logger and append to a trace file."""
log_fn = log_fn or logger.info
prefix = f"[{tag}] " if tag else ""
if not torch.cuda.is_available():
message = f"{prefix}CUDA not available on this host."
log_fn(message)
_append_to_memory_trace(message, log_file_path)
return
try:
device_index = torch.cuda.current_device()
device_name = torch.cuda.get_device_name(device_index)
allocated = torch.cuda.memory_allocated(device_index)
reserved = torch.cuda.memory_reserved(device_index)
max_allocated = torch.cuda.max_memory_allocated(device_index)
max_reserved = torch.cuda.max_memory_reserved(device_index)
free_mem, total_mem = torch.cuda.mem_get_info(device_index)
except Exception as exc: # noqa: BLE001
message = f"{prefix}Unable to query CUDA memory stats: {exc}"
log_fn(message)
_append_to_memory_trace(message, log_file_path)
return
used_mem = total_mem - free_mem
stats = [
f"device={device_name} (index={device_index})",
f"allocated={_format_bytes(allocated)}",
f"reserved={_format_bytes(reserved)}",
f"max_allocated={_format_bytes(max_allocated)}",
f"max_reserved={_format_bytes(max_reserved)}",
f"used={_format_bytes(used_mem)}",
f"free={_format_bytes(free_mem)}",
f"total={_format_bytes(total_mem)}",
]
message = f"{prefix}CUDA memory stats: {' | '.join(stats)}"
log_fn(message)
_append_to_memory_trace(message, log_file_path)
def _append_to_memory_trace(
message: str, log_file_path: str | os.PathLike[str] | None) -> None:
if not log_file_path:
return
try:
path = Path(log_file_path)
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("a", encoding="utf-8") as outfile:
outfile.write(message + "\n")
except Exception: # noqa: BLE001
# Avoid raising/logging recursively if writing to the file fails.
pass
# TODO(xingyu): add adopted message for this
def get_ip() -> str:
host_ip = envs.FASTVIDEO_HOST_IP
if host_ip:
return host_ip
# IP is not set, try to get it from the network interface
# try ipv4
s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
try:
s.connect(("8.8.8.8", 80)) # Doesn't need to be reachable
return s.getsockname()[0]
except Exception:
pass
# try ipv6
try:
s = socket.socket(socket.AF_INET6, socket.SOCK_DGRAM)
# Google's public DNS server, see
# https://developers.google.com/speed/public-dns/docs/using#addresses
s.connect(("2001:4860:4860::8888", 80)) # Doesn't need to be reachable
return s.getsockname()[0]
except Exception:
pass
warnings.warn(
"Failed to get the IP address, using 0.0.0.0 by default."
"The value can be set by the environment variable FASTVIDEO_HOST_IP.",
stacklevel=2)
return "0.0.0.0"
def test_loopback_bind(address: str, family: socket.AddressFamily) -> bool:
try:
s = socket.socket(family, socket.SOCK_DGRAM)
s.bind((address, 0)) # Port 0 = auto assign
s.close()
return True
except OSError:
return False
def get_loopback_ip() -> str:
loopback_ip = envs.FASTVIDEO_LOOPBACK_IP
if loopback_ip:
return loopback_ip
# FASTVIDEO_LOOPBACK_IP is not set, try to get it based on network interface
if test_loopback_bind("127.0.0.1", socket.AF_INET):
return "127.0.0.1"
elif test_loopback_bind("::1", socket.AF_INET6):
return "::1"
else:
raise RuntimeError(
"Neither 127.0.0.1 nor ::1 are bound to a local interface. "
"Set the FASTVIDEO_LOOPBACK_IP environment variable explicitly.")
def is_valid_ipv6_address(address: str) -> bool:
try:
ipaddress.IPv6Address(address)
return True
except ValueError:
return False
def get_distributed_init_method(ip: str, port: int) -> str:
return get_tcp_uri(ip, port)
def get_tcp_uri(ip: str, port: int) -> str:
if is_valid_ipv6_address(ip):
return f"tcp://[{ip}]:{port}"
else:
return f"tcp://{ip}:{port}"
def get_open_port(port: int | None = None) -> int:
if port is not None:
while True:
try:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(("", port))
return port
except OSError:
port += 1 # Increment port number if already in use
logger.info("Port %d is already in use, trying port %d",
port - 1, port)
# try ipv4
try:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(("", 0))
return s.getsockname()[1]
except OSError:
# try ipv6
with socket.socket(socket.AF_INET6, socket.SOCK_STREAM) as s:
s.bind(("", 0))
return s.getsockname()[1]
def cuda_is_initialized() -> bool:
"""Check if CUDA is initialized."""
if not torch.cuda._is_compiled():
return False
return torch.cuda.is_initialized()
def xpu_is_initialized() -> bool:
"""Check if XPU is initialized."""
if not torch.xpu._is_compiled():
return False
return torch.xpu.is_initialized()
def force_spawn() -> None:
if os.environ.get("FASTVIDEO_WORKER_MULTIPROC_METHOD") == "fork":
logger.warning("We must use the `spawn` multiprocessing start method.")
os.environ["FASTVIDEO_WORKER_MULTIPROC_METHOD"] = "spawn"
def get_mp_context() -> BaseContext:
"""Get a multiprocessing context with a particular method (spawn or fork).
By default we follow the value of the FASTVIDEO_WORKER_MULTIPROC_METHOD to
determine the multiprocessing method (default is fork). However, under
certain conditions, we may enforce spawn and override the value of
FASTVIDEO_WORKER_MULTIPROC_METHOD.
"""
force_spawn()
mp_method = envs.FASTVIDEO_WORKER_MULTIPROC_METHOD
return multiprocessing.get_context(mp_method)
# ANSI color codes
CYAN = '\033[1;36m'
RESET = '\033[0;0m'
def _add_prefix(file: TextIO, worker_name: str, pid: int) -> None:
"""Prepend each output line with process-specific prefix"""
prefix = f"{CYAN}({worker_name} pid={pid}){RESET} "
file_write = file.write
def write_with_prefix(s: str):
if not s:
return
if file.start_new_line: # type: ignore[attr-defined]
file_write(prefix)
idx = 0
while (next_idx := s.find('\n', idx)) != -1:
next_idx += 1
file_write(s[idx:next_idx])
if next_idx == len(s):
file.start_new_line = True # type: ignore[attr-defined]
return
file_write(prefix)
idx = next_idx
file_write(s[idx:])
file.start_new_line = False # type: ignore[attr-defined]
file.start_new_line = True # type: ignore[attr-defined]
file.write = write_with_prefix # type: ignore[method-assign]
def decorate_logs(process_name: str | None = None) -> None:
"""
Adds a process-specific prefix to each line of output written to stdout and
stderr.
Args:
process_name: Optional; the name of the process to use in the prefix.
If not provided, the current process name from the multiprocessing
context is used.
"""
if process_name is None:
process_name = get_mp_context().current_process().name
pid = os.getpid()
_add_prefix(sys.stdout, process_name, pid)
_add_prefix(sys.stderr, process_name, pid)
+2 -2
View File
@@ -1,5 +1,5 @@
from .executor import Executor
from .gpu_worker import run_worker_process
from .multiproc_executor import MultiprocExecutor
from .ray_utils import initialize_ray_cluster
__all__ = ["Executor", "run_worker_process", "MultiprocExecutor"]
__all__ = ["Executor", "MultiprocExecutor", "initialize_ray_cluster"]
+5 -2
View File
@@ -23,11 +23,14 @@ class Executor(ABC):
def _init_executor(self) -> None:
raise NotImplementedError
@classmethod
def get_class(cls, fastvideo_args: FastVideoArgs) -> type["Executor"]:
@staticmethod
def get_class(fastvideo_args: FastVideoArgs) -> type["Executor"]:
if fastvideo_args.distributed_executor_backend == "mp":
from fastvideo.worker.multiproc_executor import MultiprocExecutor
return cast(type["Executor"], MultiprocExecutor)
elif fastvideo_args.distributed_executor_backend == "ray":
from fastvideo.worker.ray_distributed_executor import RayDistributedExecutor
return cast(type["Executor"], RayDistributedExecutor)
else:
raise ValueError(
f"Unsupported distributed executor backend: {fastvideo_args.distributed_executor_backend}"
+25 -161
View File
@@ -1,17 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
import contextlib
import faulthandler
import multiprocessing as mp
import os
import signal
import sys
from multiprocessing.connection import Connection
from typing import Any, TextIO, cast
from typing import Any, cast
import psutil
import torch
import fastvideo.envs as envs
from fastvideo.distributed import (
cleanup_dist_env_and_memory,
maybe_init_distributed_environment_and_model_parallel)
@@ -19,29 +11,18 @@ from fastvideo.distributed.parallel_state import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines import ForwardBatch, LoRAPipeline, build_pipeline
from fastvideo.platforms import current_platform
from fastvideo.utils import (get_exception_traceback,
kill_itself_when_parent_died)
logger = init_logger(__name__)
# ANSI color codes
CYAN = '\033[1;36m'
RESET = '\033[0;0m'
class Worker:
def __init__(self, fastvideo_args: FastVideoArgs, local_rank: int,
rank: int, pipe: Connection, master_port: int):
rank: int, distributed_init_method: str):
self.fastvideo_args = fastvideo_args
self.local_rank = local_rank
self.rank = rank
# TODO(will): don't hardcode this
self.distributed_init_method = "env://"
self.pipe = pipe
self.master_port = master_port
self.init_device()
self.distributed_init_method = distributed_init_method
# Init request dispatcher
# TODO(will): add request dispatcher: use TypeBasedDispatcher from
@@ -70,6 +51,8 @@ class Worker:
# Platform-agnostic device initialization
self.device = get_local_torch_device()
from fastvideo.platforms import current_platform
# _check_if_gpu_supports_dtype(self.model_config.dtype)
if current_platform.is_cuda_alike():
self.init_gpu_memory = torch.cuda.mem_get_info()[0]
@@ -77,15 +60,15 @@ class Worker:
# For MPS, we can't get memory info the same way
self.init_gpu_memory = 0
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = str(self.master_port)
os.environ["LOCAL_RANK"] = str(self.local_rank)
if self.fastvideo_args.distributed_executor_backend == "mp":
os.environ["LOCAL_RANK"] = str(self.local_rank)
os.environ["RANK"] = str(self.rank)
os.environ["WORLD_SIZE"] = str(self.fastvideo_args.num_gpus)
# Initialize the distributed environment.
maybe_init_distributed_environment_and_model_parallel(
self.fastvideo_args.tp_size, self.fastvideo_args.sp_size)
self.fastvideo_args.tp_size, self.fastvideo_args.sp_size,
self.distributed_init_method)
self.pipeline = build_pipeline(self.fastvideo_args)
@@ -94,11 +77,6 @@ class Worker:
output_batch = self.pipeline.forward(forward_batch, self.fastvideo_args)
return cast(ForwardBatch, output_batch)
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
self.pipeline.set_lora_adapter(lora_nickname, lora_path)
def shutdown(self) -> dict[str, Any]:
"""Gracefully shut down the worker process"""
logger.info("Worker %d shutting down...",
@@ -117,138 +95,24 @@ class Worker:
local_main_process_only=False)
return {"status": "shutdown_complete"}
def unmerge_lora_weights(self) -> None:
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> dict[str, Any]:
if isinstance(self.pipeline, LoRAPipeline):
self.pipeline.set_lora_adapter(lora_nickname, lora_path)
logger.info("Worker %d set LoRA adapter %s with path %s", self.rank,
lora_nickname, lora_path)
return {"status": "lora_adapter_set"}
return {"status": "failed: pipeline is not a LoRAPipeline"}
def unmerge_lora_weights(self) -> dict[str, Any]:
if isinstance(self.pipeline, LoRAPipeline):
self.pipeline.unmerge_lora_weights()
return {"status": "lora_adapter_unmerged"}
return {"status": "failed: pipeline is not a LoRAPipeline"}
def merge_lora_weights(self) -> None:
def merge_lora_weights(self) -> dict[str, Any]:
if isinstance(self.pipeline, LoRAPipeline):
self.pipeline.merge_lora_weights()
def event_loop(self) -> None:
"""Event loop for the worker."""
logger.info("Worker %d starting event loop...",
self.rank,
local_main_process_only=False)
while True:
try:
recv_rpc = self.pipe.recv()
method_name = recv_rpc.get('method')
# Handle shutdown request
if method_name == 'shutdown':
response = self.shutdown()
with contextlib.suppress(Exception):
self.pipe.send(response)
break # Exit the loop
# Handle regular RPC calls
if method_name == 'execute_forward':
forward_batch = recv_rpc['kwargs']['forward_batch']
fastvideo_args = recv_rpc['kwargs']['fastvideo_args']
output_batch = self.execute_forward(forward_batch,
fastvideo_args)
logging_info = None
if envs.FASTVIDEO_STAGE_LOGGING:
logging_info = output_batch.logging_info
self.pipe.send({
"output_batch": output_batch.output.cpu(),
"logging_info": logging_info
})
elif method_name == 'set_lora_adapter':
lora_nickname = recv_rpc['kwargs']['lora_nickname']
lora_path = recv_rpc['kwargs']['lora_path']
self.set_lora_adapter(lora_nickname, lora_path)
logger.info("Worker %d set LoRA adapter %s with path %s",
self.rank, lora_nickname, lora_path)
self.pipe.send({"status": "lora_adapter_set"})
elif method_name == 'unmerge_lora_weights':
self.unmerge_lora_weights()
logger.info("Worker %d unmerged LoRA weights", self.rank)
self.pipe.send({"status": "lora_adapter_unmerged"})
elif method_name == 'merge_lora_weights':
self.merge_lora_weights()
logger.info("Worker %d merged LoRA weights", self.rank)
self.pipe.send({"status": "lora_adapter_merged"})
else:
# Handle other methods dynamically if needed
args = recv_rpc.get('args', ())
kwargs = recv_rpc.get('kwargs', {})
if hasattr(self, method_name):
method = getattr(self, method_name)
result = method(*args, **kwargs)
self.pipe.send(result)
else:
self.pipe.send(
{"error": f"Unknown method: {method_name}"})
except KeyboardInterrupt:
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
def run_worker_process(fastvideo_args: FastVideoArgs, local_rank: int,
rank: int, pipe: Connection, master_port: int):
# Add process-specific prefix to stdout and stderr
process_name = mp.current_process().name
pid = os.getpid()
_add_prefix(sys.stdout, process_name, pid)
_add_prefix(sys.stderr, process_name, pid)
# Config the process
kill_itself_when_parent_died()
faulthandler.enable()
parent_process = psutil.Process().parent()
logger.info("Worker %d initializing...",
rank,
local_main_process_only=False)
try:
worker = Worker(fastvideo_args, local_rank, rank, pipe, master_port)
logger.info("Worker %d sending ready", rank)
pipe.send({
"status": "ready",
"local_rank": local_rank,
})
worker.event_loop()
except Exception:
traceback = get_exception_traceback()
logger.error("Worker %d hit an exception: %s", rank, traceback)
parent_process.send_signal(signal.SIGQUIT)
def _add_prefix(file: TextIO, worker_name: str, pid: int) -> None:
"""Prepend each output line with process-specific prefix"""
prefix = f"{CYAN}({worker_name} pid={pid}){RESET} "
file_write = file.write
def write_with_prefix(s: str):
if not s:
return
if file.start_new_line: # type: ignore[attr-defined]
file_write(prefix)
idx = 0
while (next_idx := s.find('\n', idx)) != -1:
next_idx += 1
file_write(s[idx:next_idx])
if next_idx == len(s):
file.start_new_line = True # type: ignore[attr-defined]
return
file_write(prefix)
idx = next_idx
file_write(s[idx:])
file.start_new_line = False # type: ignore[attr-defined]
file.start_new_line = True # type: ignore[attr-defined]
file.write = write_with_prefix # type: ignore[method-assign]
return {"status": "lora_adapter_merged"}
return {"status": "failed: pipeline is not a LoRAPipeline"}
+557 -64
View File
@@ -1,21 +1,30 @@
# SPDX-License-Identifier: Apache-2.0
import atexit
import contextlib
from dataclasses import dataclass
import faulthandler
import multiprocessing as mp
from multiprocessing.connection import Connection
import os
import signal
import socket
import time
from collections.abc import Callable
from multiprocessing.process import BaseProcess
from typing import Any
from typing import Any, cast
from collections.abc import Iterator
import psutil
from fastvideo.distributed.parallel_state import get_dp_group, get_tp_group
import fastvideo.envs as envs
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.utils import decorate_logs, get_distributed_init_method, get_exception_traceback, get_loopback_ip, get_mp_context, get_open_port, kill_itself_when_parent_died, force_spawn
from fastvideo.worker.executor import Executor
from fastvideo.worker.gpu_worker import run_worker_process
from fastvideo.worker.worker_base import WorkerWrapperBase
import torch
import torchvision
logger = init_logger(__name__)
@@ -26,51 +35,37 @@ class MultiprocExecutor(Executor):
self.world_size = self.fastvideo_args.num_gpus
self.shutting_down = False
# this will force the use of the `spawn` multiprocessing start if cuda
# is initialized
mp.set_start_method("spawn", force=True)
self.workers: list[BaseProcess] = []
self.worker_pipes = []
set_multiproc_executor_envs()
# Check if master_port is provided in fastvideo_args
if hasattr(
self.fastvideo_args,
'master_port') and self.fastvideo_args.master_port is not None:
self.master_port = self.fastvideo_args.master_port
logger.info("Using provided master port: %s", self.master_port)
else:
# Auto-find available port
import random
for port in range(29503 + random.randint(0, 10000), 65535):
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
if s.connect_ex(('localhost', port)) != 0:
self.master_port = port
break
else:
raise ValueError("No unused port found to use as master port")
logger.info("Auto-selected master port: %s", self.master_port)
master_port = get_open_port(self.fastvideo_args.master_port)
distributed_init_method = get_distributed_init_method(
get_loopback_ip(), master_port)
logger.info("Use master port: %s", master_port)
# Create pipes and start workers
for rank in range(self.world_size):
executor_pipe, worker_pipe = mp.Pipe(duplex=True)
self.worker_pipes.append(executor_pipe)
worker = mp.Process(target=run_worker_process,
name=f"FVWorkerProc-{rank}",
kwargs=dict(fastvideo_args=self.fastvideo_args,
local_rank=rank,
rank=rank,
pipe=worker_pipe,
master_port=self.master_port))
worker.start()
self.workers.append(worker)
unready_workers: list[UnreadyWorkerProcHandle] = []
success = False
try:
for rank in range(self.world_size):
unready_workers.append(
WorkerMultiprocProc.make_worker_process(
fastvideo_args=self.fastvideo_args,
local_rank=rank,
rank=rank,
distributed_init_method=distributed_init_method,
))
# Wait for all workers to be ready
for idx, pipe in enumerate(self.worker_pipes):
data = pipe.recv()
if data["status"] != "ready" or data["local_rank"] != idx:
raise RuntimeError(f"Worker {idx} failed to start")
logger.info("%d workers ready", self.world_size)
# Workers must be created before wait_for_ready to avoid
# deadlock, since worker.init_device() does a device sync.
self.workers = WorkerMultiprocProc.wait_for_ready(unready_workers)
success = True
finally:
if not success:
# Clean up the worker procs if there was a failure.
# Close death_writers first to signal workers to exit
self._ensure_worker_termination(
[uw.proc for uw in unready_workers])
# Register shutdown on exit
atexit.register(self.shutdown)
@@ -96,6 +91,61 @@ class MultiprocExecutor(Executor):
return result_batch
def execute_forward_streaming(
self, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> Iterator[dict[str, Any]]:
"""
Multiprocess streaming interface.
Broadcasts a streaming forward request to all workers. Rank 0 streams
events (progress/block/complete) back through its pipe; other ranks
run the forward pass and return a small completion status.
"""
# Send request to all workers
for worker in self.workers:
worker.pipe.send({
"method": "execute_forward_streaming",
"args": (),
"kwargs": {
"forward_batch": forward_batch,
"fastvideo_args": fastvideo_args,
},
})
done_workers: set[int] = set()
rank0_complete = False
pipes = [w.pipe for w in self.workers]
# Drain messages until all workers are done
while len(done_workers) < len(self.workers):
ready = mp.connection.wait(pipes)
for pipe in ready:
# Identify which worker sent this
idx = next(
(i for i, w in enumerate(self.workers) if w.pipe is pipe),
-1)
if idx < 0:
continue
msg = pipe.recv()
# Rank 0 streams events
if idx == 0 and isinstance(msg, dict):
msg_type = msg.get("type")
if msg_type in ("progress", "block", "complete"):
yield msg
if msg_type == "complete":
rank0_complete = True
done_workers.add(idx)
elif msg.get("status") == "done":
done_workers.add(idx)
else:
# Non-zero ranks: expect simple completion status
if isinstance(msg, dict) and msg.get("status") == "done":
done_workers.add(idx)
else:
# Any unexpected message from non-zero ranks counts as done
done_workers.add(idx)
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
@@ -129,12 +179,16 @@ class MultiprocExecutor(Executor):
kwargs = kwargs or {}
try:
for pipe in self.worker_pipes:
pipe.send({"method": method, "args": args, "kwargs": kwargs})
for worker in self.workers:
worker.pipe.send({
"method": method,
"args": args,
"kwargs": kwargs
})
responses = []
for pipe in self.worker_pipes:
response = pipe.recv()
for worker in self.workers:
response = worker.pipe.recv()
responses.append(response)
return responses
except TimeoutError as e:
@@ -145,8 +199,8 @@ class MultiprocExecutor(Executor):
logger.info(
"Received KeyboardInterrupt, sending SIGINT to all workers")
for worker in self.workers:
if worker.pid is not None:
os.kill(worker.pid, signal.SIGINT)
if worker.proc.pid is not None:
os.kill(worker.proc.pid, signal.SIGINT)
raise e
except Exception as e:
raise e
@@ -162,52 +216,84 @@ class MultiprocExecutor(Executor):
# First try gentle termination
try:
# Send termination message to all workers
for pipe in self.worker_pipes:
for worker in self.workers:
with contextlib.suppress(Exception):
pipe.send({"method": "shutdown", "args": (), "kwargs": {}})
worker.pipe.send({
"method": "shutdown",
"args": (),
"kwargs": {}
})
# Give workers some time to exit gracefully
start_time = time.perf_counter()
while time.perf_counter() - start_time < 5.0: # 5 seconds timeout
if all(not worker.is_alive() for worker in self.workers):
if all(not worker.proc.is_alive() for worker in self.workers):
break
time.sleep(0.1)
# Force terminate any remaining workers
for worker in self.workers:
if worker.is_alive():
worker.terminate()
if worker.proc.is_alive():
worker.proc.terminate()
# Final timeout for terminate
start_time = time.perf_counter()
while time.perf_counter() - start_time < 2.0: # 2 seconds timeout
if all(not worker.is_alive() for worker in self.workers):
if all(not worker.proc.is_alive() for worker in self.workers):
break
time.sleep(0.1)
# Kill if still alive
for worker in self.workers:
if worker.is_alive():
worker.kill()
worker.join(timeout=1.0)
if worker.proc.is_alive():
worker.proc.kill()
worker.proc.join(timeout=1.0)
except Exception as e:
logger.error("Error during shutdown: %s", e)
# Last resort, try to kill all workers
for worker in self.workers:
with contextlib.suppress(Exception):
if worker.is_alive():
worker.kill()
if worker.proc.is_alive():
worker.proc.kill()
# Clean up pipes
for pipe in self.worker_pipes:
for worker in self.workers:
with contextlib.suppress(Exception):
pipe.close()
worker.pipe.close()
self.workers = []
self.worker_pipes = []
logger.info("MultiprocExecutor shutdown complete")
@staticmethod
def _ensure_worker_termination(worker_procs: list[BaseProcess]):
"""Ensure that all worker processes are terminated. Assumes workers have
received termination requests. Waits for processing, then sends
termination and kill signals if needed."""
def wait_for_termination(procs: list[BaseProcess],
timeout: float) -> bool:
if not time:
# If we are in late stage shutdown, the interpreter may replace
# `time` with `None`.
return all(not proc.is_alive() for proc in procs)
start_time = time.time()
while time.time() - start_time < timeout:
if all(not proc.is_alive() for proc in procs):
return True
time.sleep(0.1)
return False
# Send SIGTERM if still running
active_procs = [proc for proc in worker_procs if proc.is_alive()]
for p in active_procs:
p.terminate()
if not wait_for_termination(active_procs, 4):
# Send SIGKILL if still running
active_procs = [p for p in active_procs if p.is_alive()]
for p in active_procs:
p.kill()
def __del__(self):
"""Ensure cleanup on garbage collection"""
self.shutdown()
@@ -219,3 +305,410 @@ class MultiprocExecutor(Executor):
def __exit__(self, exc_type, exc_val, exc_tb):
"""Ensure cleanup when exiting context"""
self.shutdown()
@dataclass
class UnreadyWorkerProcHandle:
"""WorkerProcess handle before READY."""
proc: BaseProcess
rank: int
pipe: Connection
ready_pipe: Connection
@dataclass
class WorkerProcHandle:
proc: BaseProcess
rank: int
pipe: Connection
@classmethod
def from_unready_handle(
cls, unready_handle: UnreadyWorkerProcHandle) -> "WorkerProcHandle":
return cls(
proc=unready_handle.proc,
rank=unready_handle.rank,
pipe=unready_handle.pipe,
)
class WorkerMultiprocProc:
"""Adapter that runs one Worker in busy loop."""
READY_STR = "READY"
def __init__(
self,
fastvideo_args: FastVideoArgs,
local_rank: int,
rank: int,
distributed_init_method: str,
pipe: Connection,
):
self.rank = rank
self.pipe = pipe
wrapper = WorkerWrapperBase(fastvideo_args=fastvideo_args,
rpc_rank=rank)
all_kwargs: list[dict] = [{} for _ in range(fastvideo_args.num_gpus)]
all_kwargs[rank] = {
"fastvideo_args": fastvideo_args,
"local_rank": local_rank,
"rank": rank,
"distributed_init_method": distributed_init_method,
}
wrapper.init_worker(all_kwargs)
self.worker = wrapper
# Initialize device
self.worker.init_device()
# Set process title and log prefix
self.setup_proc_title_and_log_prefix()
@staticmethod
def make_worker_process(
fastvideo_args: FastVideoArgs,
local_rank: int,
rank: int,
distributed_init_method: str,
) -> UnreadyWorkerProcHandle:
context = get_mp_context()
executor_pipe, worker_pipe = context.Pipe(duplex=True)
reader, writer = context.Pipe(duplex=False)
process_kwargs = {
"fastvideo_args": fastvideo_args,
"local_rank": local_rank,
"rank": rank,
"distributed_init_method": distributed_init_method,
"pipe": worker_pipe,
"ready_pipe": writer,
}
# Run EngineCore busy loop in background process.
proc = context.Process(target=WorkerMultiprocProc.worker_main,
kwargs=process_kwargs,
name=f"FVWorkerProc-{rank}",
daemon=True)
proc.start()
worker_pipe.close()
return UnreadyWorkerProcHandle(proc, rank, executor_pipe, reader)
@staticmethod
def worker_main(*args, **kwargs):
""" Worker initialization and execution loops.
This runs a background process """
# Signal handler used for graceful termination.
# SystemExit exception is only raised once to allow this and worker
# processes to terminate without error
shutdown_requested = False
def signal_handler(signum, frame):
nonlocal shutdown_requested
if not shutdown_requested:
shutdown_requested = True
raise SystemExit()
# Either SIGTERM or SIGINT will terminate the worker
signal.signal(signal.SIGTERM, signal_handler)
signal.signal(signal.SIGINT, signal_handler)
kill_itself_when_parent_died()
faulthandler.enable()
parent_process = psutil.Process().parent()
worker = None
ready_pipe = kwargs.pop("ready_pipe")
rank = kwargs.get("rank")
try:
worker = WorkerMultiprocProc(*args, **kwargs)
# Send READY once we know everything is loaded
ready_pipe.send({
"status": WorkerMultiprocProc.READY_STR,
})
ready_pipe.close()
ready_pipe = None
worker.worker_busy_loop()
except Exception:
if ready_pipe is not None:
logger.exception("WorkerMultiprocProc failed to start.")
else:
logger.exception("WorkerMultiprocProc failed.")
# The parent sends a SIGTERM to all worker processes if
# any worker dies. Set this value so we don't re-throw
# SystemExit() to avoid zmq exceptions in __del__.
shutdown_requested = True
traceback = get_exception_traceback()
logger.error("Worker %d hit an exception: %s", rank, traceback)
parent_process.send_signal(signal.SIGQUIT)
finally:
if ready_pipe is not None:
ready_pipe.close()
# Clean up once worker exits busy loop
if worker is not None:
worker.shutdown()
@staticmethod
def wait_for_ready(
unready_proc_handles: list[UnreadyWorkerProcHandle]
) -> list[WorkerProcHandle]:
e = Exception("WorkerMultiprocProc initialization failed due to "
"an exception in a background process. "
"See stack trace for root cause.")
pipes = {handle.ready_pipe: handle for handle in unready_proc_handles}
ready_proc_handles: list[WorkerProcHandle
| None] = ([None] * len(unready_proc_handles))
while pipes:
ready = mp.connection.wait(pipes.keys())
for pipe in ready:
assert isinstance(pipe, Connection)
try:
# Wait until the WorkerProc is ready.
unready_proc_handle = pipes.pop(pipe)
response: dict[str, Any] = pipe.recv()
if response["status"] != "READY":
raise e
ready_proc_handles[unready_proc_handle.rank] = (
WorkerProcHandle.from_unready_handle(
unready_proc_handle))
except EOFError:
e.__suppress_context__ = True
raise e from None
finally:
# Close connection.
pipe.close()
logger.info("%d workers ready", len(ready_proc_handles))
return cast(list[WorkerProcHandle], ready_proc_handles)
def shutdown(self) -> dict[str, Any]:
return self.worker.shutdown()
def worker_busy_loop(self) -> None:
"""Main busy loop for Multiprocessing Workers"""
while True:
logger.info("Worker %d starting event loop...", self.rank)
try:
rpc_call = self.pipe.recv()
method = rpc_call.get("method")
args = rpc_call.get("args", ())
kwargs = rpc_call.get("kwargs", {})
if isinstance(method, str):
if method == "shutdown":
response = self.shutdown()
with contextlib.suppress(Exception):
self.pipe.send(response)
break
if method == 'execute_forward':
forward_batch = kwargs['forward_batch']
fastvideo_args = kwargs['fastvideo_args']
output_batch = self.worker.execute_forward(
forward_batch, fastvideo_args)
logging_info = None
if envs.FASTVIDEO_STAGE_LOGGING:
logging_info = output_batch.logging_info
self.pipe.send({
"output_batch": output_batch.output.cpu(),
"logging_info": logging_info
})
if method == 'execute_forward_streaming':
forward_batch = kwargs['forward_batch']
fastvideo_args = kwargs['fastvideo_args']
# Install a lightweight per-block callback for streaming metadata from rank 0
try:
extra = getattr(forward_batch, 'extra', None)
if extra is None:
forward_batch.extra = {}
extra = forward_batch.extra
except Exception:
forward_batch.extra = {}
extra = forward_batch.extra
def _on_block(**evt):
# Only rank 0 streams events to the executor pipe
if self.rank != 0:
return
try:
block_index = int(evt.get('block_index', 0))
total_blocks = int(evt.get('total_blocks', 0))
num_frames = int(evt.get('num_frames', 0))
latents = evt.get('latents')
if latents is None:
# Fallback to meta event if no latents provided
self.pipe.send({
'type': 'block_meta',
'block_index': block_index,
'total_blocks': total_blocks,
'num_frames': num_frames,
})
return
# Decode latents to pixels using pipeline VAE
vae = None
try:
vae = self.worker.pipeline.get_module('vae')
except Exception:
vae = None
if vae is None:
# If no VAE, send meta only
self.pipe.send({
'type': 'block_meta',
'block_index': block_index,
'total_blocks': total_blocks,
'num_frames': num_frames,
})
return
# Prepare latents for VAE (apply scaling/shift like DecodingStage)
z = latents.permute(
0, 2, 1, 3, 4
) # [B,T,C,H,W] -> [B,C,T,H,W] expected by many VAEs
if hasattr(vae, 'scaling_factor'
) and vae.scaling_factor is not None:
sf = vae.scaling_factor
if isinstance(sf, torch.Tensor):
z = z / sf.to(z.device, z.dtype)
else:
z = z / sf
if hasattr(vae, 'shift_factor'
) and vae.shift_factor is not None:
shf = vae.shift_factor
if isinstance(shf, torch.Tensor):
z = z + shf.to(z.device, z.dtype)
else:
z = z + shf
with torch.autocast(device_type='cuda',
dtype=torch.bfloat16,
enabled=True):
pixels = vae.decode(
z
) # [B,C,T,H,W] in [-1,1] or already normalized depending on VAE
# Normalize to [0,1]
pixels = (pixels / 2 + 0.5).clamp(0, 1)
# Convert to uint8 HWC per-frame for first batch element
pixels = (pixels[0].permute(1, 2, 3, 0) *
255).to(torch.uint8).cpu().numpy()
# Now pixels shape: [T,H,W,C]
frames = [
pixels[i] for i in range(pixels.shape[0])
]
self.pipe.send({
'type': 'block',
'block_index': block_index,
'total_blocks': total_blocks,
'frames': frames,
})
except Exception:
# Fail-safe: don't break generation if decode fails
try:
self.pipe.send({
'type':
'block_meta',
'block_index':
int(evt.get('block_index', 0)),
'total_blocks':
int(evt.get('total_blocks', 0)),
'num_frames':
int(evt.get('num_frames', 0)),
})
except Exception:
pass
try:
extra['on_block'] = _on_block
except Exception:
pass
# Run forward on all ranks
output_batch = self.worker.execute_forward(
forward_batch, fastvideo_args)
if self.rank != 0:
# Non-zero ranks send a small completion message
self.pipe.send({"status": "done"})
else:
# Rank 0 builds frames and streams a single block then complete
samples = output_batch.output # Tensor [b,c,t,h,w]
try:
videos = samples.permute(2, 0, 1, 3, 4)
frames = []
for x in videos:
grid = torchvision.utils.make_grid(x,
nrow=6)
grid = grid.transpose(0, 1).transpose(
1, 2).squeeze(-1)
frame = (grid * 255).to(
torch.uint8).cpu().numpy()
frames.append(frame)
# Emit a single block with all frames
self.pipe.send({
"type": "block",
"block_index": 0,
"total_blocks": 1,
"frames": frames,
})
except Exception:
# If frame construction fails, still complete
pass
# Emit a complete event
self.pipe.send({
"type": "complete",
"result": {
"num_frames":
int(samples.shape[2])
if hasattr(samples, 'shape')
and len(samples.shape) >= 3 else None
},
})
else:
result = self.worker.execute_method(method, *args, **kwargs)
self.pipe.send(result)
except KeyboardInterrupt:
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
@staticmethod
def setup_proc_title_and_log_prefix() -> None:
dp_size = get_dp_group().world_size
dp_rank = get_dp_group().rank_in_group
tp_size = get_tp_group().world_size
tp_rank = get_tp_group().rank_in_group
process_name = "Worker"
if dp_size > 1:
process_name += f"_DP{dp_rank}"
if tp_size > 1:
process_name += f"_TP{tp_rank}"
decorate_logs(process_name)
def set_multiproc_executor_envs() -> None:
""" Set up environment variables that should be used when there are workers
in a multiprocessing environment. This should be called by the parent
process before worker processes are created"""
force_spawn()
@@ -0,0 +1,362 @@
# SPDX-License-Identifier: Apache-2.0
# Adapt from https://github.com/vllm-project/vllm/blob/releases/v0.11.0/vllm/executor/ray_distributed_executor.py
from collections import defaultdict
import os
import cloudpickle
import fastvideo.envs as envs
from dataclasses import dataclass
from typing import Any, TYPE_CHECKING
from collections.abc import Callable
from fastvideo.utils import get_ip, get_distributed_init_method, get_open_port
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.worker.executor import Executor
from fastvideo.worker.ray_utils import (
initialize_ray_cluster,
RayWorkerWrapper,
ray,
)
from fastvideo.worker.ray_env import get_env_vars_to_copy
from fastvideo.logger import init_logger
if ray is not None:
from ray.actor import ActorHandle
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
else:
ActorHandle = None
if TYPE_CHECKING:
from ray.util.placement_group import PlacementGroup
logger = init_logger(__name__)
@dataclass
class RayWorkerMetaData:
"""
Metadata for a Ray worker.
The order of ray worker creation can be random,
and we need to reset the rank after creating all workers.
"""
worker: ActorHandle
created_rank: int
adjusted_rank: int = -1
ip: str = ""
class RayDistributedExecutor(Executor):
"""Ray-based distributed executor"""
# These env vars are worker-specific, therefore are NOT copied
# from the driver to the workers
WORKER_SPECIFIC_ENV_VARS = {
"FASTVIDEO_HOST_IP",
"LOCAL_RANK",
"CUDA_VISIBLE_DEVICES",
}
# These non-vLLM env vars are copied from the driver to workers
ADDITIONAL_ENV_VARS = {"HF_TOKEN", "HUGGING_FACE_HUB_TOKEN"}
def _init_executor(self) -> None:
initialize_ray_cluster(self.fastvideo_args)
placement_group = self.fastvideo_args.ray_placement_group
# Disable Ray usage stats collection.
ray_usage = os.environ.get("RAY_USAGE_STATS_ENABLED", "0")
if ray_usage != "1":
os.environ["RAY_USAGE_STATS_ENABLED"] = "0"
self._init_workers_ray(placement_group)
# child class could overwrite this to return actual env vars.
def _get_env_vars_to_be_updated(self) -> list[dict[str, str]]:
return self._env_vars_for_all_workers
def _init_workers_ray(self, placement_group: "PlacementGroup",
**ray_remote_kwargs):
from fastvideo.platforms import current_platform
num_gpus = envs.FASTVIDEO_RAY_PER_WORKER_GPUS
# The remaining workers are the actual ray actors.
self.workers: list[RayWorkerWrapper] = []
# Create the workers.
# use the first N bundles that have GPU resources.
bundle_indices: list[int] = []
for bundle_id, bundle in enumerate(placement_group.bundle_specs):
if bundle.get(current_platform.ray_device_key, 0):
bundle_indices.append(bundle_id)
bundle_indices = bundle_indices[:self.fastvideo_args.num_gpus]
worker_metadata: list[RayWorkerMetaData] = []
driver_ip = get_ip()
for rank, bundle_id in enumerate(bundle_indices):
scheduling_strategy = PlacementGroupSchedulingStrategy(
placement_group=placement_group,
placement_group_capture_child_tasks=True,
placement_group_bundle_index=bundle_id,
)
if current_platform.ray_device_key == "GPU":
# NV+AMD GPUs, and Intel XPUs
worker = ray.remote(
num_cpus=0,
num_gpus=num_gpus,
scheduling_strategy=scheduling_strategy,
**ray_remote_kwargs,
)(RayWorkerWrapper).remote(fastvideo_args=self.fastvideo_args,
rpc_rank=rank)
else:
worker = ray.remote(
num_cpus=0,
num_gpus=0,
resources={current_platform.ray_device_key: num_gpus},
scheduling_strategy=scheduling_strategy,
**ray_remote_kwargs,
)(RayWorkerWrapper).remote(fastvideo_args=self.fastvideo_args,
rpc_rank=rank)
worker_metadata.append(
RayWorkerMetaData(worker=worker, created_rank=rank))
worker_ips = ray.get([
each.worker.get_node_ip.remote() # type: ignore[attr-defined]
for each in worker_metadata
])
for each, ip in zip(worker_metadata, worker_ips, strict=False):
each.ip = ip
logger.debug("workers: %s", worker_metadata)
ip_counts: dict[str, int] = {}
for ip in worker_ips:
ip_counts[ip] = ip_counts.get(ip, 0) + 1
def sort_by_driver_then_worker_ip(item: RayWorkerMetaData):
"""
Sort the workers based on 3 properties:
1. If the worker is on the same node as the driver (vllm engine),
it should be placed first.
2. Then, if the worker is on a node with fewer workers, it should
be placed first.
3. Finally, if the work is on a node with smaller IP address, it
should be placed first.
"""
ip = item.ip
return (0 if ip == driver_ip else 1, ip_counts[ip], ip)
# After sorting, the workers on the same node will be
# close to each other, and the workers on the driver
# node will be placed first.
sorted_worker_metadata = sorted(worker_metadata,
key=sort_by_driver_then_worker_ip)
start_rank = 0
for i, item in enumerate(sorted_worker_metadata):
item.adjusted_rank = i + start_rank
self.workers = [item.worker for item in sorted_worker_metadata]
rerank_mapping = {
item.created_rank: item.adjusted_rank
for item in sorted_worker_metadata
}
self._run_ray_workers("adjust_rank", rerank_mapping)
# Get the set of GPU IDs used on each node.
worker_node_and_gpu_ids = self._run_ray_workers("get_node_and_gpu_ids")
node_workers = defaultdict(list) # node id -> list of worker ranks
node_gpus = defaultdict(list) # node id -> list of gpu ids
for i, (node_id, gpu_ids) in enumerate(worker_node_and_gpu_ids):
node_workers[node_id].append(i)
# `gpu_ids` can be a list of strings or integers.
# convert them to integers for consistency.
# NOTE: gpu_ids can be larger than 9 (e.g. 16 GPUs),
# string sorting is not sufficient.
# see https://github.com/vllm-project/vllm/issues/5590
gpu_ids = [int(x) for x in gpu_ids]
node_gpus[node_id].extend(gpu_ids)
for node_id, gpu_ids in node_gpus.items():
node_gpus[node_id] = sorted(gpu_ids)
all_ips = set(worker_ips + [driver_ip])
n_ips = len(all_ips)
n_nodes = len(node_workers)
if n_nodes != n_ips:
raise RuntimeError(
f"Every node should have a unique IP address. Got {n_nodes}"
f" nodes with node ids {list(node_workers.keys())} and "
f"{n_ips} unique IP addresses {all_ips}. Please check your"
" network configuration. If you set `FASTVIDEO_HOST_IP`"
" environment variable, make sure it is unique for"
" each node.")
# Set environment variables for the driver and workers.
all_args_to_update_environment_variables: list[dict[str, str]] = [{
current_platform.device_control_env_var:
",".join(map(str, node_gpus[node_id])),
} for (node_id, _) in worker_node_and_gpu_ids]
# Environment variables to copy from driver to workers
env_vars_to_copy = get_env_vars_to_copy(
exclude_vars=self.WORKER_SPECIFIC_ENV_VARS,
additional_vars=set(current_platform.additional_env_vars).union(
self.ADDITIONAL_ENV_VARS),
destination="workers",
)
# Copy existing env vars to each worker's args
for args in all_args_to_update_environment_variables:
# TODO: refactor platform-specific env vars
for name in env_vars_to_copy:
if name in os.environ:
args[name] = os.environ[name]
self._env_vars_for_all_workers: list[dict[str, str]] = (
all_args_to_update_environment_variables)
self._run_ray_workers("update_environment_variables",
self._get_env_vars_to_be_updated())
if len(node_gpus) == 1:
# in single node case, we don't need to get the IP address.
# the loopback address is sufficient
# NOTE: a node may have several IP addresses, one for each
# network interface. `get_ip()` might return any of them,
# while they might not work for communication inside the node
# if the network setup is complicated. Using the loopback address
# solves this issue, as it always works for communication inside
# the node.
driver_ip = "127.0.0.1"
distributed_init_method = get_distributed_init_method(
driver_ip, get_open_port())
# Initialize the actual workers inside worker wrapper.
all_kwargs = []
for rank, (node_id, _) in enumerate(worker_node_and_gpu_ids):
local_rank = node_workers[node_id].index(rank)
kwargs = dict(
fastvideo_args=self.fastvideo_args,
local_rank=local_rank,
rank=rank,
distributed_init_method=distributed_init_method,
)
all_kwargs.append(kwargs)
self._run_ray_workers("init_worker", all_kwargs)
self._run_ray_workers("init_device")
# This is the list of workers that are rank 0 of each TP group EXCEPT
# global rank 0. These are the workers that will broadcast to the
# rest of the workers.
self.tp_driver_workers: list[RayWorkerWrapper] = []
# This is the list of workers that are not drivers and not the first
# worker in a TP group. These are the workers that will be
# broadcasted to.
self.non_driver_workers: list[RayWorkerWrapper] = []
# Enforce rank order for correct rank to return final output.
for index, worker in enumerate(self.workers):
# The driver worker is rank 0 and not in self.workers.
rank = index + 1
if rank % self.fastvideo_args.tp_size == 0:
self.tp_driver_workers.append(worker)
else:
self.non_driver_workers.append(worker)
def execute_forward(self, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
responses: list[ForwardBatch] = self.collective_rpc(
"execute_forward",
kwargs={
"forward_batch": forward_batch,
"fastvideo_args": fastvideo_args,
},
)
output = responses[0].output.cpu()
logging_info = None
if envs.FASTVIDEO_STAGE_LOGGING:
logging_info = responses[0].logging_info
result_batch = ForwardBatch(
data_type=forward_batch.data_type,
output=output,
logging_info=logging_info,
)
return result_batch
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
responses = self.collective_rpc("set_lora_adapter",
kwargs={
"lora_nickname": lora_nickname,
"lora_path": lora_path
})
for i, response in enumerate(responses):
if response["status"] != "lora_adapter_set":
raise RuntimeError(
f"Worker {i} failed to set LoRA adapter to {lora_path}")
def unmerge_lora_weights(self) -> None:
responses = self.collective_rpc("unmerge_lora_weights", kwargs={})
for i, response in enumerate(responses):
if response["status"] != "lora_adapter_unmerged":
raise RuntimeError(f"Worker {i} failed to unmerge LoRA weights")
def merge_lora_weights(self) -> None:
responses = self.collective_rpc("merge_lora_weights", kwargs={})
for i, response in enumerate(responses):
if response["status"] != "lora_adapter_merged":
raise RuntimeError(f"Worker {i} failed to merge LoRA weights")
def collective_rpc(self,
method: str | Callable,
timeout: float | None = None,
args: tuple = (),
kwargs: dict | None = None) -> list[Any]:
return self._run_ray_workers(method, *args, **(kwargs or {}))
def _run_ray_workers(
self,
method: str | Callable,
*args,
**kwargs,
) -> Any:
if isinstance(method, str):
sent_method = method
else:
sent_method = cloudpickle.dumps(method)
del method
# Start the ray workers first.
ray_workers = self.workers
ray_worker_outputs = [
worker.execute_method.remote(sent_method, *args, **kwargs)
for worker in ray_workers
]
# Get the results of the ray workers.
ray_worker_outputs = ray.get(ray_worker_outputs)
return ray_worker_outputs
def shutdown(self) -> None:
logger.info(
"Shutting down Ray distributed executor. If you see error log "
"from logging.cc regarding SIGTERM received, please ignore because "
"this is the expected termination process in Ray.")
import ray
for worker in self.workers:
ray.kill(worker)
self.workers = []
def __del__(self):
self.shutdown()
+79
View File
@@ -0,0 +1,79 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import json
import os
import fastvideo.envs as envs
from fastvideo.logger import init_logger
logger = init_logger(__name__)
CONFIG_HOME = envs.FASTVIDEO_CONFIG_ROOT
# This file contains a list of env vars that should not be copied
# from the driver to the Ray workers.
RAY_NON_CARRY_OVER_ENV_VARS_FILE = os.path.join(
CONFIG_HOME, "ray_non_carry_over_env_vars.json")
try:
if os.path.exists(RAY_NON_CARRY_OVER_ENV_VARS_FILE):
with open(RAY_NON_CARRY_OVER_ENV_VARS_FILE) as f:
RAY_NON_CARRY_OVER_ENV_VARS = set(json.load(f))
else:
RAY_NON_CARRY_OVER_ENV_VARS = set()
except json.JSONDecodeError:
logger.warning(
"Failed to parse %s. Using an empty set for non-carry-over env vars.",
RAY_NON_CARRY_OVER_ENV_VARS_FILE,
)
RAY_NON_CARRY_OVER_ENV_VARS = set()
def get_env_vars_to_copy(
exclude_vars: set[str] | None = None,
additional_vars: set[str] | None = None,
destination: str | None = None,
) -> set[str]:
"""
Get the environment variables to copy to downstream Ray actors.
Example use cases:
- Copy environment variables from RayDistributedExecutor to Ray workers.
- Copy environment variables from RayDPClient to Ray DPEngineCoreActor.
Args:
exclude_vars: A set of FastVideo defined environment variables to exclude
from copying.
additional_vars: A set of additional environment variables to copy.
If a variable is in both exclude_vars and additional_vars, it will
be excluded.
destination: The destination of the environment variables.
Returns:
A set of environment variables to copy.
"""
exclude_vars = exclude_vars or set()
additional_vars = additional_vars or set()
env_vars_to_copy = {
v
for v in set(envs.environment_variables).union(additional_vars)
if v not in exclude_vars and v not in RAY_NON_CARRY_OVER_ENV_VARS
}
to_destination = " to " + destination if destination is not None else ""
logger.info(
"RAY_NON_CARRY_OVER_ENV_VARS from config: %s",
RAY_NON_CARRY_OVER_ENV_VARS,
)
logger.info(
"Copying the following environment variables%s: %s",
to_destination,
[v for v in env_vars_to_copy if v in os.environ],
)
logger.info(
"If certain env vars should NOT be copied, add them to %s file",
RAY_NON_CARRY_OVER_ENV_VARS_FILE,
)
return env_vars_to_copy
+272
View File
@@ -0,0 +1,272 @@
# SPDX-License-Identifier: Apache-2.0
# Adapt from https://github.com/vllm-project/vllm/blob/releases/v0.11.0/vllm/executor/ray_utils.py
from collections import defaultdict
import os
import time
from fastvideo.utils import get_ip
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.worker.worker_base import WorkerWrapperBase
from fastvideo.logger import init_logger
logger = init_logger(__name__)
PG_WAIT_TIMEOUT = 1800
try:
import ray
from ray.util import placement_group_table
from ray.util.placement_group import PlacementGroup
try:
from ray._private.state import available_resources_per_node
except ImportError:
# Ray 2.9.x doesn't expose `available_resources_per_node`
from ray._private.state import state as _state
available_resources_per_node = _state._available_resources_per_node
class RayWorkerWrapper(WorkerWrapperBase):
"""Ray wrapper for fastvideo.worker.Worker, allowing Worker to be
lazily initialized after Ray sets CUDA_VISIBLE_DEVICES."""
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
def get_node_ip(self) -> str:
return get_ip()
def get_node_and_gpu_ids(self) -> tuple[str, list[int]]:
from fastvideo.platforms import current_platform
node_id = ray.get_runtime_context().get_node_id()
device_key = current_platform.ray_device_key
if not device_key:
raise RuntimeError("current platform %s does not support ray.",
current_platform.device_name)
gpu_ids = ray.get_runtime_context().get_accelerator_ids(
)[device_key]
return node_id, gpu_ids
def override_env_vars(self, vars: dict[str, str]):
os.environ.update(vars)
ray_import_err = None
except ImportError as e:
ray = None
ray_import_err = str(e)
def assert_ray_available() -> None:
"""Raise an exception if Ray is not available."""
if ray is None:
raise ValueError(f"Failed to import Ray: {ray_import_err}."
"Please install Ray with `pip install ray`.")
def _verify_bundles(placement_group: "PlacementGroup",
fastvideo_args: FastVideoArgs, device_str: str):
"""Verify a given placement group has bundles located in the right place.
There are 2 rules.
- Warn if all tensor parallel workers cannot fit in a single node.
- Fail if driver node is not included in a placement group.
"""
assert ray.is_initialized(), (
"Ray is not initialized although distributed-executor-backend is ray.")
pg_data = placement_group_table(placement_group)
# bundle_idx -> node_id
bundle_to_node_ids = pg_data["bundles_to_node_id"]
# bundle_idx -> bundle (e.g., {"GPU": 1})
bundles = pg_data["bundles"]
# node_id -> List of bundle (e.g., {"GPU": 1})
node_id_to_bundle: dict[str, list[dict[str, float]]] = defaultdict(list)
for bundle_idx, node_id in bundle_to_node_ids.items():
node_id_to_bundle[node_id].append(bundles[bundle_idx])
driver_node_id = ray.get_runtime_context().get_node_id()
if driver_node_id not in node_id_to_bundle:
raise RuntimeError(
f"driver node id {driver_node_id} is not included in a placement "
f"group {placement_group.id}. Node id -> bundles "
f"{node_id_to_bundle}. "
"You don't have enough GPUs available in a current node. Check "
"`ray status` and `ray list nodes` to see if you have available "
"GPUs in a node `{driver_node_id}` before starting an FastVideo engine."
)
for node_id, bundles in node_id_to_bundle.items():
if len(bundles) < fastvideo_args.tp_size:
logger.warning(
"tensor_parallel_size=%d "
"is bigger than a reserved number of %ss (%d "
"%ss) in a node %s. Tensor parallel workers can be "
"spread out to 2+ nodes which can degrade the performance "
"unless you have fast interconnect across nodes, like "
"Infiniband. To resolve this issue, make sure you have more "
"than %d GPUs available at each node.", fastvideo_args.tp_size,
device_str, len(bundles), device_str, node_id,
fastvideo_args.tp_size)
def _wait_until_pg_ready(current_placement_group: "PlacementGroup"):
"""Wait until a placement group is ready.
It prints the informative log messages if the placement group is
not created within time.
"""
# Wait until PG is ready - this will block until all
# requested resources are available, and will timeout
# if they cannot be provisioned.
placement_group_specs = current_placement_group.bundle_specs
s = time.time()
pg_ready_ref = current_placement_group.ready()
wait_interval = 10
while time.time() - s < PG_WAIT_TIMEOUT:
ready, _ = ray.wait([pg_ready_ref], timeout=wait_interval)
if len(ready) > 0:
break
# Exponential backoff for warning print.
wait_interval *= 2
logger.info(
"Waiting for creating a placement group of specs for "
"%d seconds. specs=%s. Check `ray status` and "
"`ray list nodes` to see if you have enough resources,"
" and make sure the IP addresses used by ray cluster"
" are the same as FASTVIDEO_HOST_IP environment variable"
" specified in each node if you are running on a multi-node.",
int(time.time() - s), placement_group_specs)
try:
ray.get(pg_ready_ref, timeout=0)
except ray.exceptions.GetTimeoutError:
raise ValueError(
"Cannot provide a placement group of "
f"{placement_group_specs=} within {PG_WAIT_TIMEOUT} seconds. See "
"`ray status` and `ray list nodes` to make sure the cluster has "
"enough resources.") from None
def initialize_ray_cluster(
fastvideo_args: FastVideoArgs,
ray_address: str | None = None,
):
"""Initialize the distributed cluster with Ray.
it will connect to the Ray cluster and create a placement group
for the workers, which includes the specification of the resources
for each distributed worker.
Args:
parallel_config: The configurations for parallel execution.
ray_address: The address of the Ray cluster. If None, uses
the default Ray cluster address.
"""
assert_ray_available()
from fastvideo.platforms import current_platform
if ray.is_initialized():
logger.info("Ray is already initialized. Skipping Ray initialization.")
elif current_platform.is_rocm() or current_platform.is_xpu():
# Try to connect existing ray instance and create a new one if not found
try:
ray.init("auto")
except ConnectionError:
logger.warning(
"No existing RAY instance detected. "
"A new instance will be launched with current node resources.")
ray.init(address=ray_address,
num_gpus=fastvideo_args.num_gpus,
runtime_env=fastvideo_args.ray_runtime_env)
else:
ray.init(address=ray_address,
runtime_env=fastvideo_args.ray_runtime_env)
device_str = current_platform.ray_device_key
if not device_str:
raise ValueError(
f"current platform {current_platform.device_name} does not "
"support ray.")
# Create or get the placement group for worker processes
if fastvideo_args.ray_placement_group:
current_placement_group = fastvideo_args.ray_placement_group
else:
current_placement_group = ray.util.get_current_placement_group()
if current_placement_group:
logger.info("Using the existing placement group")
# We are in a placement group
bundles = current_placement_group.bundle_specs
# Verify that we can use the placement group.
device_bundles = 0
for bundle in bundles:
bundle_devices = bundle.get(device_str, 0)
if bundle_devices > 1:
raise ValueError(
"Placement group bundle cannot have more than 1 "
f"{device_str}.")
if bundle_devices:
device_bundles += 1
if fastvideo_args.num_gpus > device_bundles:
raise ValueError(
f"The number of required {device_str}s exceeds the total "
f"number of available {device_str}s in the placement group. "
f"Required number of devices: {fastvideo_args.num_gpus}. "
f"Total number of devices: {device_bundles}.")
else:
logger.info("No current placement group found. "
"Creating a new placement group.")
num_devices_in_cluster = ray.cluster_resources().get(device_str, 0)
# Log a warning message and delay resource allocation failure response.
# Avoid immediate rejection to allow user-initiated placement group
# created and wait cluster to be ready
if fastvideo_args.num_gpus > num_devices_in_cluster:
logger.warning(
"The number of required %ss exceeds the total "
"number of available %ss in the placement group.", device_str,
device_str)
# Create a new placement group
placement_group_specs: list[dict[str, float]] = ([{
device_str: 1.0
} for _ in range(fastvideo_args.num_gpus)])
# FastVideo engine is also a worker to execute model with an accelerator,
# so it requires to have the device in a current node. Check if
# the current node has at least one device.
current_ip = get_ip()
current_node_id = ray.get_runtime_context().get_node_id()
current_node_resource = available_resources_per_node()[current_node_id]
if current_node_resource.get(device_str, 0) < 1:
raise ValueError(
f"Current node has no {device_str} available. "
f"{current_node_resource=}. FastVideo engine cannot start without "
f"{device_str}. Make sure you have at least 1 {device_str} "
f"available in a node {current_node_id=} {current_ip=}.")
# This way, at least bundle is required to be created in a current
# node.
placement_group_specs[0][f"node:{current_ip}"] = 0.001
# By default, Ray packs resources as much as possible.
current_placement_group = ray.util.placement_group(
placement_group_specs, strategy="PACK")
_wait_until_pg_ready(current_placement_group)
assert current_placement_group is not None
_verify_bundles(current_placement_group, fastvideo_args, device_str)
# Set the placement group in the fastvideo args
fastvideo_args.ray_placement_group = current_placement_group
def is_in_ray_actor():
"""Check if we are in a Ray actor."""
try:
import ray
return (ray.is_initialized()
and ray.get_runtime_context().get_actor_id() is not None)
except ImportError:
return False
+92
View File
@@ -0,0 +1,92 @@
import os
from typing import Any
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.worker.gpu_worker import Worker
from fastvideo.utils import (run_method, update_environment_variables)
logger = init_logger(__name__)
class WorkerWrapperBase:
"""
This class represents one process in an executor/engine. It is responsible
for lazily initializing the worker and handling the worker's lifecycle.
We first instantiate the WorkerWrapper, which remembers the worker module
and class name. Then, when we call `update_environment_variables`, and the
real initialization happens in `init_worker`.
"""
def __init__(
self,
fastvideo_args: FastVideoArgs,
rpc_rank: int = 0,
) -> None:
"""
Initialize the worker wrapper with the given fastvideo_args and rpc_rank.
Note: rpc_rank is the rank of the worker in the executor. In most cases,
it is also the rank of the worker in the distributed group. However,
when multiple executors work together, they can be different.
e.g. in the case of SPMD-style offline inference with TP=2,
users can launch 2 engines/executors, each with only 1 worker.
All workers have rpc_rank=0, but they have different ranks in the TP
group.
"""
self.rpc_rank = rpc_rank
self.worker: Worker | None = None
self.fastvideo_args: FastVideoArgs | None = None
# do not store this `fastvideo_args`, `init_worker` will set the final
# one.
def adjust_rank(self, rank_mapping: dict[int, int]) -> None:
"""
Adjust the rpc_rank based on the given mapping.
It is only used during the initialization of the executor,
to adjust the rpc_rank of workers after we create all workers.
"""
if self.rpc_rank in rank_mapping:
self.rpc_rank = rank_mapping[self.rpc_rank]
def update_environment_variables(self, envs_list: list[dict[str,
str]]) -> None:
envs = envs_list[self.rpc_rank]
key = 'CUDA_VISIBLE_DEVICES'
if key in envs and key in os.environ:
# overwriting CUDA_VISIBLE_DEVICES is desired behavior
# suppress the warning in `update_environment_variables`
del os.environ[key]
update_environment_variables(envs)
def init_worker(self, all_kwargs: list[dict[str, Any]]) -> None:
"""
Here we inject some common logic before initializing the worker.
Arguments are passed to the worker class constructor.
"""
kwargs = all_kwargs[self.rpc_rank]
self.fastvideo_args = kwargs.get("fastvideo_args")
assert self.fastvideo_args is not None, (
"fastvideo_args is required to initialize the worker")
self.worker = Worker(**kwargs)
assert self.worker is not None
def execute_method(self, method: str | bytes, *args, **kwargs):
try:
# method resolution order:
# if a method is defined in this class, it will be called directly.
# otherwise, since we define `__getattr__` and redirect attribute
# query to `self.worker`, the method will be called on the worker.
return run_method(self, method, args, kwargs)
except Exception as e:
# if the driver worker also execute methods,
# exceptions in the rest worker may cause deadlock in rpc like ray
# see https://github.com/vllm-project/vllm/issues/3455
# print the error and inform the user to solve the error
msg = (f"Error executing method {method!r}. "
"This might cause deadlock in distributed execution.")
logger.exception(msg)
raise e
def __getattr__(self, attr):
return getattr(self.worker, attr)
+37 -13
View File
@@ -15,29 +15,52 @@ classifiers = [
dependencies = [
# Core Libraries
"scipy==1.14.1", "six==1.16.0", "h5py==3.12.1", "requests>=2.32.2",
"scipy==1.14.1",
"six==1.16.0",
"h5py==3.12.1",
"requests>=2.32.2",
# Machine Learning & Transformers
"transformers>=4.46.1", "tokenizers>=0.20.1", "sentencepiece==0.2.0",
"timm==1.0.11", "peft>=0.15.0", "diffusers>=0.33.1",
"torch==2.7.1", "torchvision",
"transformers>=4.46.1",
"tokenizers>=0.20.1",
"sentencepiece==0.2.0",
"timm==1.0.11",
"peft>=0.15.0",
"diffusers>=0.33.1",
"torch==2.7.1",
"torchvision",
# Acceleration & Optimization
"accelerate==1.0.1",
# Computer Vision & Image Processing
"opencv-python==4.10.0.84", "pillow>=10.3.0", "imageio==2.36.0",
"imageio-ffmpeg==0.5.1", "einops",
"opencv-python==4.10.0.84",
"pillow>=10.3.0",
"imageio==2.36.0",
"imageio-ffmpeg==0.5.1",
"einops",
# Experiment Tracking & Logging
"wandb>=0.21.0", "loguru", "test-tube==0.7.5",
"wandb>=0.21.0",
"loguru",
"test-tube==0.7.5",
# Miscellaneous Utilities
"tqdm", "pytest", "PyYAML==6.0.1", "protobuf>=5.28.3",
"gradio==5.41.0", "moviepy>=2.0.0", "flask",
"flask_restful", "aiohttp", "huggingface_hub", "cloudpickle",
"tqdm",
"pytest",
"PyYAML==6.0.1",
"protobuf>=5.28.3",
"gradio==5.32.0",
"moviepy>=2.0.0",
"flask",
"flask_restful",
"aiohttp",
"huggingface_hub",
"cloudpickle",
# System & Monitoring Tools
"gpustat", "watch", "remote-pdb",
"gpustat",
"watch",
"remote-pdb",
# Kernel & Packaging
"wheel",
@@ -49,7 +72,8 @@ dependencies = [
"av",
# Preprocessing Dependencies
"torchcodec==0.5.0"
"torchcodec==0.5.0",
"ray>=2.49.1",
]
[tool.uv]
+1 -1
View File
@@ -32,7 +32,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--max_train_steps 30000 \
--learning_rate 2e-6 \
--mixed_precision "bf16" \
--checkpointing_steps 400 \
--training_state_checkpointing_steps 400 \
--validation_steps 100 \
--validation_sampling_steps "3" \
--log_validation \
+1 -1
View File
@@ -33,7 +33,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--max_train_steps 30000 \
--learning_rate 2e-6 \
--mixed_precision "bf16" \
--checkpointing_steps 400 \
--training_state_checkpointing_steps 400 \
--validation_steps 100 \
--validation_sampling_steps "3" \
--log_validation \
+1 -1
View File
@@ -30,7 +30,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--max_train_steps=5000 \
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=6000 \
--training_state_checkpointing_steps 6000 \
--validation_steps 200\
--validation_sampling_steps "2,4,8" \
--log_validation \
+1 -1
View File
@@ -36,7 +36,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--max_train_steps 30000 \
--learning_rate 1e-6 \
--mixed_precision "bf16" \
--checkpointing_steps 6000 \
--training_state_checkpointing_steps 6000 \
--validation_steps 100 \
--validation_sampling_steps "50" \
--log_validation \