Compare commits
15
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c4a0f789da | ||
|
|
5ac14938d2 | ||
|
|
9d239e9f8b | ||
|
|
bdec816b31 | ||
|
|
2cd2e57d2e | ||
|
|
9370234294 | ||
|
|
50da62e722 | ||
|
|
4f3e8751db | ||
|
|
f4c58894d9 | ||
|
|
01c94ef385 | ||
|
|
2415226d25 | ||
|
|
404314d00f | ||
|
|
87489f0872 | ||
|
|
9ce7c8039e | ||
|
|
e1e25e95f9 |
@@ -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 }}
|
||||
|
||||
@@ -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">
|
||||
|
||||
@@ -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 |
@@ -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 |
@@ -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 |
@@ -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.
|
||||
@@ -115,6 +115,7 @@ design/overview
|
||||
|
||||
contributing/overview
|
||||
contributing/developer_env/index
|
||||
contributing/profiling
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
|
||||
@@ -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()
|
||||
@@ -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.
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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`
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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(
|
||||
[
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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")
|
||||
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -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"
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
@@ -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"}
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
@@ -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]
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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 \
|
||||
|
||||
Reference in New Issue
Block a user