Compare commits

..
16 Commits
Author SHA1 Message Date
JerryZhou54 38ee9dc3b4 Complete Model Config design for VAEs 2025-04-23 19:08:11 +00:00
JerryZhou54 77b013fb8a Add model config for WanVAE 2025-04-23 19:01:13 +00:00
JerryZhou54 c31efe1234 Add model config for VAE 2025-04-23 18:59:25 +00:00
JerryZhou54 c056b89aea Add preliminary design for model config 2025-04-23 18:57:13 +00:00
Kevin Lin eac79b753f [V1] Worker improvements/cleanup (#361) 2025-04-22 00:32:43 -07:00
William Lin 4d58cf20d0 chore: Release FastVideo 0.0.2 and update python requirements (#360) 2025-04-21 14:10:48 -07:00
Kevin LinandWill Lin 52c93ecc9d [V1] Gradio demo with new API (#357)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-04-19 18:14:24 -07:00
William Lin 42d63166ac [V1] Process aware logging; improve logging msg (#356) 2025-04-19 15:03:29 -07:00
William Lin 6db20345a2 [V1] Worker cleanup; Logging clean up; enables isort again (#355) 2025-04-18 19:26:47 -07:00
William Lin ad27ea596c [sta] release 0.0.4 (#354) 2025-04-18 14:54:40 -07:00
William Lin 9aadb4bf8c [1/n] [v1] Add Worker abstractions for User API (#336) 2025-04-18 14:38:46 -07:00
Kevin Lin bd941df271 [Docs] Fix developer guide images (#353) 2025-04-17 22:32:19 -07:00
Yongqi Chen 8a73876d3b add STA to Wan v1 (#349) 2025-04-17 16:35:19 -07:00
Kevin Lin 1483a1138a [CLI] Fix duplicate --num-gpus (#352) 2025-04-17 13:01:48 -07:00
Wei Zhou 5e243d8292 Default to using original WanVAE's encoding/decoding algorithm (#351) 2025-04-17 13:00:25 -07:00
Kevin Lin b0c66d3200 [CI] Docker image improvements (#350) 2025-04-17 12:27:10 -07:00
77 changed files with 402427 additions and 621 deletions
+2
View File
@@ -8,12 +8,14 @@ on:
- main
paths:
- "docs/**/*.md"
- "fastvideo/v1/examples/**/*.py"
pull_request:
branches:
- main
types: [opened, ready_for_review, synchronize, reopened]
paths:
- "docs/**/*.md"
- "fastvideo/v1/examples/**/*.py"
# Allows you to run this workflow manually from the Actions tab
workflow_dispatch:
+4 -4
View File
@@ -116,7 +116,7 @@ jobs:
--gpu-count 1
--volume-size 100
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/encoders -s"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
@@ -164,7 +164,7 @@ jobs:
--gpu-count 1
--volume-size 100
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/vaes -s"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
@@ -212,7 +212,7 @@ jobs:
--gpu-count 1
--volume-size 100
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/transformers -s"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
@@ -261,7 +261,7 @@ jobs:
--disk-size 200
--volume-size 200
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/ssim -vs"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
- name: Terminate RunPod Instances
if: ${{ always() }}
+4 -1
View File
@@ -55,4 +55,7 @@ docs/source/getting_started/examples/
*.pkl
# Reference videos
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
# Static images
!docs/source/_static/images/**/*.png
+5 -4
View File
@@ -19,6 +19,7 @@ exclude: |
fastvideo/sample/.*|
fastvideo/train\.py|
fastvideo/utils/.*|
fastvideo/v1/examples/.*|
.github/workflows/fastvideo-publish.yml|
.github/workflows/sta-publish.yml
)
@@ -40,10 +41,10 @@ repos:
- id: codespell
additional_dependencies: ['tomli']
args: ['--toml', 'pyproject.toml']
# - repo: https://github.com/PyCQA/isort
# rev: 0a0b7a830386ba6a31c2ec8316849ae4d1b8240d # 6.0.0
# hooks:
# - id: isort
- repo: https://github.com/PyCQA/isort
rev: 6.0.1
hooks:
- id: isort
- repo: https://github.com/jackdewinter/pymarkdown
rev: v0.9.29
hooks:
+12 -1
View File
@@ -2,7 +2,7 @@ FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
ENV DEBIAN_FRONTEND=noninteractive
WORKDIR /app
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
@@ -34,4 +34,15 @@ RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
COPY . .
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
# Remove authentication headers
RUN git config --unset-all http.https://github.com/.extraheader || true
# Set up automatic conda environment activation for all shells
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
# Ensure .bashrc is sourced for SSH login shells
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
EXPOSE 22
+2 -2
View File
@@ -29,7 +29,7 @@ Dev in progress and highly experimental.
- ```2024/12/17```: `FastVideo` v1.0 is released.
## 🔧 Installation from source
The code is tested on Python 3.10.0, CUDA 12.4 and H100.
The code is tested on Python 3.10-3.12, CUDA 12.4 and H100.
```
# Clone FastVideo
@@ -47,7 +47,7 @@ To try Sliding Tile Attention (optional), please follow the instruction in [csrc
You can also install the Sliding Tile Attention package using
```
pip install st_attn==0.0.3
pip install st_attn==0.0.4
```
## 🚀 Inference
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -9,7 +9,7 @@ target = target.lower()
# Package metadata
PACKAGE_NAME = "st_attn"
VERSION = "0.0.3"
VERSION = "0.0.4"
AUTHOR = "Hao AI Lab"
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
+1 -1
View File
@@ -10,7 +10,7 @@
#ifdef TK_COMPILE_ATTN
extern torch::Tensor sta_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
);
#endif
@@ -4,8 +4,13 @@ import torch
from st_attn_cuda import sta_fwd
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True):
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, img_latent_shape='30*48*80'):
seq_length = q_all.shape[2]
img_latent_shape_mapping = {
'30x48x80':1,
'36x48x48':2,
'18x48x80':3,
}
if has_text:
assert q_all.shape[
2] >= 115200, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
@@ -17,8 +22,14 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
else:
assert q_all.shape[2] == 82944
if img_latent_shape == '36x48x48': # Stepvideo 204x768x68
assert q_all.shape[2] == 82944
elif img_latent_shape == '18x48x80': # Wan 69x768x1280
assert q_all.shape[2] == 69120
else:
raise ValueError(f"Unsupported {img_latent_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
kernel_aspect_ratio_flag = img_latent_shape_mapping[img_latent_shape]
hidden_states = torch.empty_like(q_all)
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
@@ -29,7 +40,7 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
head_index:head_index + 1],
hidden_states[batch:batch + 1, head_index:head_index + 1])
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text)
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
if has_text:
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True)
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
return hidden_states[:, :, :seq_length]
@@ -359,7 +359,7 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
#include <iostream>
torch::Tensor
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text)
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
@@ -558,123 +558,267 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
} else {
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
if (kernel_aspect_ratio_flag == 2){
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
}else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
}else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==6 && kernel_h_size == 3 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==6 && kernel_h_size == 3 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
}
else if (kernel_aspect_ratio_flag == 3) {
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 0, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 0, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 1, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,0, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,0, 3, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
}
else {
std::cout << "Unsupported kernel_aspect_ratio_flag: " << kernel_aspect_ratio_flag << std::endl;
}
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 18 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 27 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 40 KiB

+86
View File
@@ -50,3 +50,89 @@ pre-commit run --all-files
# Unit tests
pytest tests/
```
---
## 🐳 Using the FastVideo Docker Image
If you prefer a containerized development environment or want to avoid managing dependencies manually, you can use our prebuilt Docker image:
**Image:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
### Starting the container
```bash
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
```
This will:
- Start the container with GPU access
- Drop you into a shell with the `fastvideo-dev` Conda environment preconfigured
### Using the container
```bash
# Conda environment should already be active
# FastVideo package installed in editable mode
# Pull the latest changes from remote
cd /FastVideo
git pull
# Run linters and tests
pre-commit run --all-files
pytest tests/
```
---
## 📦 Developing FastVideo on RunPod
You can easily use the FastVideo Docker image as a custom container on [RunPod](https://www.runpod.io) for development or experimentation.
### Creating a new pod
Choose a GPU that supports CUDA 12.4
![RunPod CUDA selection](../_static/images/runpod_cuda.png)
When creating your pod template, use this image:
```
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
```
Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.io/pods/configuration/use-ssh)):
```bash
bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
```
![RunPod template configuration](../_static/images/runpod_template.png)
After deploying, the pod will take a few minutes to pull the image and start the SSH service.
![RunPod ssh](../_static/images/runpod_ssh.png)
### Working with the pod
After SSH'ing into your pod, you'll find the `fastvideo-dev` Conda environment already activated.
To pull in the latest changes from the GitHub repo:
```bash
cd /FastVideo
git pull
```
`If you have a persistent volume and want to keep your code changes, you can move /FastVideo to /workspace/FastVideo, or simply clone the repository there.`
Run your development workflows as usual:
```bash
# Run linters
pre-commit run --all-files
# Run tests
pytest tests/
```
+3 -1
View File
@@ -9,7 +9,7 @@ from typing import Optional
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
ROOT_DIR_RELATIVE = '../../../..'
EXAMPLE_DIR = ROOT_DIR / "examples"
EXAMPLE_DIR = ROOT_DIR / "fastvideo/v1/examples"
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/source/getting_started/examples"
@@ -208,6 +208,7 @@ def generate_examples():
glob_patterns = ["*.py", "*.md", "*.sh"]
# Find categorised examples
for category in category_indices:
print(category)
category_dir = EXAMPLE_DIR / category
globs = [category_dir.glob(pattern) for pattern in glob_patterns]
for path in itertools.chain(*globs):
@@ -228,6 +229,7 @@ def generate_examples():
# Generate the example documentation
for example in sorted(examples, key=lambda e: e.path.stem):
print(example)
doc_path = EXAMPLE_DOC_DIR / f"{example.path.stem}.md"
with open(doc_path, "w+") as f:
f.write(example.generate())
+17 -7
View File
@@ -2,12 +2,20 @@
# 🔧 Installation
FastVideo currently only supports Linux and CUDA GPUs. The code is tested on Python 3.10.0 and CUDA 12.4, primarily with NVIDIA H100 GPUs.
FastVideo currently only supports Linux and NVIDIA CUDA GPUs.
## Prerequisites
FastVideo has been tested on the following GPUs, but it should work on any GPUs that supports CUDA 12.4+, please create an issue if you discover any issues:
- RTX 4090
- A40
- L40S
- A100
- H100
- CUDA 12.4 installed and supported
- Linux operating system
## Requirements
- OS: Linux
- Python: 3.10-3.12
- CUDA 12.4+ (Untested on CUDA < 12.4)
## Installation Options
@@ -19,7 +27,9 @@ pip install fastvideo
### Option 2: Installation from Source
#### 1. Install Miniconda (if not already installed)
We recommend using a Python environment such as Conda.
#### 1. [Optional] Install Miniconda (if not already installed)
```bash
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
@@ -27,7 +37,7 @@ bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
#### 2. Create and activate a Conda environment for FastVideo
#### 2. [Optional] Create and activate a Conda environment for FastVideo
```bash
conda create -n fastvideo python=3.10 -y
@@ -56,7 +66,7 @@ pip install -e .
pip install flash-attn==2.7.0.post2 --no-build-isolation
```
### Sliding Tile Attention (STA)
### Sliding Tile Attention (STA) (Requires CUDA 12.4+ and H100)
To try Sliding Tile Attention (optional), please follow the instructions in [csrc/sliding_tile_attention/README.md](#sta-installation) to install STA.
+1
View File
@@ -69,6 +69,7 @@ sliding_tile_attention/demo
:caption: Inference
:maxdepth: 1
inference/wanvideo
inference/stepvideo
inference/hunyuanvideo
inference/fasthunyuan
+2
View File
@@ -1 +1,3 @@
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
__all__ = ["VideoGenerator"]
@@ -20,20 +20,39 @@ logger = init_logger(__name__)
# TODO(will-refactor): move this to a utils file
def dict_to_3d_list(mask_strategy,
t_max=50,
l_max=60,
h_max=24) -> List[List[List[Optional[torch.Tensor]]]]:
result = [[[None for _ in range(h_max)] for _ in range(l_max)]
for _ in range(t_max)]
if mask_strategy is None:
return result
def dict_to_3d_list(mask_strategy) -> List[List[List[Optional[torch.Tensor]]]]:
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
max_timesteps_idx = max(
timesteps_idx for timesteps_idx, layer_idx, head_idx in indices) + 1
max_layer_idx = max(layer_idx
for timesteps_idx, layer_idx, head_idx in indices) + 1
max_head_idx = max(head_idx
for timesteps_idx, layer_idx, head_idx in indices) + 1
result = [[[None for _ in range(max_head_idx)]
for _ in range(max_layer_idx)] for _ in range(max_timesteps_idx)]
for key, value in mask_strategy.items():
t, layer, h = map(int, key.split('_'))
result[t][layer][h] = value
timesteps_idx, layer_idx, head_idx = map(int, key.split('_'))
result[timesteps_idx][layer_idx][head_idx] = value
return result
class RangeDict(dict):
def __getitem__(self, item):
for key in self.keys():
if isinstance(key, tuple):
low, high = key
if low <= item <= high:
return super().__getitem__(key)
elif key == item:
return super().__getitem__(key)
raise KeyError(f"seq_len {item} not supported for STA")
class SlidingTileAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@@ -103,52 +122,73 @@ class SlidingTileAttentionImpl(AttentionImpl):
with open(config_file) as f:
mask_strategy = json.load(f)
mask_strategy = dict_to_3d_list(mask_strategy)
self.prefix = prefix
self.mask_strategy = mask_strategy
sp_group = get_sp_group()
self.sp_size = sp_group.world_size
# STA config
self.STA_base_tile_size = [6, 8, 8]
self.img_latent_shape_mapping = RangeDict({
(115200, 115456): '30x48x80',
82944: '36x48x48',
69120: '18x48x80',
})
self.full_window_mapping = {
'30x48x80': [5, 6, 10],
'36x48x48': [6, 6, 6],
'18x48x80': [3, 6, 10]
}
def tile(self, x: torch.Tensor) -> torch.Tensor:
x = rearrange(x,
"b (sp t h w) head d -> b (t sp h w) head d",
sp=self.sp_size,
t=30 // self.sp_size,
h=48,
w=80)
t=self.img_latent_shape_int[0] // self.sp_size,
h=self.img_latent_shape_int[1],
w=self.img_latent_shape_int[2])
return rearrange(
x,
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
n_t=5,
n_h=6,
n_w=10,
ts_t=6,
ts_h=8,
ts_w=8)
n_t=self.full_window_size[0],
n_h=self.full_window_size[1],
n_w=self.full_window_size[2],
ts_t=self.STA_base_tile_size[0],
ts_h=self.STA_base_tile_size[1],
ts_w=self.STA_base_tile_size[2])
def untile(self, x: torch.Tensor) -> torch.Tensor:
x = rearrange(
x,
"b (n_t n_h n_w ts_t ts_h ts_w) h d -> b (n_t ts_t n_h ts_h n_w ts_w) h d",
n_t=5,
n_h=6,
n_w=10,
ts_t=6,
ts_h=8,
ts_w=8)
n_t=self.full_window_size[0],
n_h=self.full_window_size[1],
n_w=self.full_window_size[2],
ts_t=self.STA_base_tile_size[0],
ts_h=self.STA_base_tile_size[1],
ts_w=self.STA_base_tile_size[2])
return rearrange(x,
"b (t sp h w) head d -> b (sp t h w) head d",
sp=self.sp_size,
t=30 // self.sp_size,
h=48,
w=80)
t=self.img_latent_shape_int[0] // self.sp_size,
h=self.img_latent_shape_int[1],
w=self.img_latent_shape_int[2])
def preprocess_qkv(
self,
qkv: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
img_sequence_length = qkv.shape[1]
self.img_latent_shape_str = self.img_latent_shape_mapping[
img_sequence_length]
self.full_window_size = self.full_window_mapping[
self.img_latent_shape_str]
self.img_latent_shape_int = list(
map(int, self.img_latent_shape_str.split('x')))
self.img_seq_length = self.img_latent_shape_int[
0] * self.img_latent_shape_int[1] * self.img_latent_shape_int[2]
return self.tile(qkv)
def postprocess_output(
@@ -173,11 +213,15 @@ class SlidingTileAttentionImpl(AttentionImpl):
timestep = attn_metadata.current_timestep
# pattern:'.double_blocks.0.attn.impl' or '.single_blocks.0.attn.impl'
layer_idx = int(self.prefix.split('.')[-3])
# TODO: remove hardcode
text_length = q.shape[1] - (30 * 48 * 80)
query = q.transpose(1, 2)
key = k.transpose(1, 2)
value = v.transpose(1, 2)
text_length = q.shape[1] - self.img_seq_length
has_text = text_length > 0
query = q.transpose(1, 2).contiguous()
key = k.transpose(1, 2).contiguous()
value = v.transpose(1, 2).contiguous()
head_num = query.size(1)
sp_group = get_sp_group()
@@ -187,7 +231,11 @@ class SlidingTileAttentionImpl(AttentionImpl):
self.mask_strategy[timestep][layer_idx][head_idx + start_head]
for head_idx in range(head_num)
]
hidden_states = sliding_tile_attention(query, key, value, windows,
text_length).transpose(1, 2)
# if has_text is False:
# from IPython import embed
# embed()
hidden_states = sliding_tile_attention(
query, key, value, windows, text_length, has_text,
self.img_latent_shape_str).transpose(1, 2)
return hidden_states
+5 -3
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from typing import List, Optional
from typing import Optional, Tuple
import torch
import torch.nn as nn
@@ -25,7 +25,8 @@ class DistributedAttention(nn.Module):
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[List[_Backend]] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = "",
**extra_impl_args) -> None:
super().__init__()
@@ -146,7 +147,8 @@ class LocalAttention(nn.Module):
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[List[_Backend]] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
**extra_impl_args) -> None:
super().__init__()
if softmax_scale is None:
+13 -2
View File
@@ -3,7 +3,8 @@
import os
from contextlib import contextmanager
from typing import Generator, List, Optional, Type, cast
from functools import cache
from typing import Generator, Optional, Tuple, Type, cast
import torch
@@ -81,7 +82,17 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[List[_Backend]] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
) -> Type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype,
supported_attention_backends)
@cache
def _cached_get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
) -> Type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
-8
View File
@@ -1,8 +0,0 @@
from fastvideo.v1.configs.hunyuan import HunyuanConfig, FastHunyuanConfig
from fastvideo.v1.configs.wan import WanT2V480PConfig, WanI2V480PConfig
from fastvideo.v1.configs.base import BaseConfig, SlidingTileAttnConfig
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig"
]
-70
View File
@@ -1,70 +0,0 @@
from dataclasses import dataclass, field
from typing import Optional, Dict, Any
@dataclass
class BaseConfig:
"""Base configuration for all pipeline architectures."""
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 125
fps: int = 24
# Video generation parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
seed: int = 1024
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
vae_scale_factor: Optional[int] = None
# DiT configuration
num_channels_latents: Optional[int] = None
# Image encoder configuration
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_len: int = -1
hidden_state_skip_layer: int = 0
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
neg_prompt: Optional[str] = None
# Additional parameters can be added as a dict
extra_params: Dict[str, Any] = field(default_factory=dict)
@dataclass
class SlidingTileAttnConfig(BaseConfig):
"""Configuration for sliding tile attention."""
# Override any BaseConfig defaults as needed
# Add sliding tile specific parameters
window_size: int = 16
stride: int = 8
# You can provide custom defaults for inherited fields
height: int = 576
width: int = 1024
# Additional configuration specific to sliding tile attention
pad_to_square: bool = False
use_overlap_optimization: bool = True
+7
View File
@@ -0,0 +1,7 @@
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
__all__ = [
"ArchConfig", "ModelConfig",
"VAEArchConfig", "VAEConfig"
]
+47
View File
@@ -0,0 +1,47 @@
from dataclasses import dataclass, fields
from typing import Dict, Any
# 1. ArchConfig contains all fields from diffuser's/transformer's config.json (i.e. all fields related to the architecture of the model)
# 2. ArchConfig should be inherited & overriden by each model arch_config
# 3. Any field in ArchConfig is fixed upon initialization, and should be hidden away from users
@dataclass
class ArchConfig:
pass
@dataclass
class ModelConfig:
# Every model config parameter can be categorized into either ArchConfig or everything else
# Diffuser/Transformer parameters
arch_config: ArchConfig = ArchConfig()
# FastVideo-specific parameters here
# i.e. STA, quantization, teacache
# This should be used only when loading from transformers/diffusers
def update_model_arch(
self,
source_model_dict: Dict[str, Any]
) -> None:
arch_config = self.arch_config
valid_fields = {f.name for f in fields(arch_config)}
for key, value in source_model_dict.items():
if key in valid_fields:
setattr(arch_config, key, value)
else:
raise AttributeError(f"{type(arch_config).__name__} has no field '{key}'")
def update_model_config(
self,
source_model_dict: Dict[str, Any]
) -> None:
assert "arch_config" not in source_model_dict.keys(), "Source model config shouldn't contain arch_config."
valid_fields = {f.name for f in fields(self)}
for key, value in source_model_dict.items():
if key in valid_fields:
setattr(self, key, value)
else:
print(f"{type(self).__name__} does not contain field '{key}'!")
raise AttributeError(f"Invalid field: {key}")
@@ -0,0 +1,7 @@
from fastvideo.v1.configs.models.vaes.hunyuanvae import HunyuanVAEConfig, HunyuanVAEArchConfig
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig, WanVAEArchConfig
__all__ = [
"HunyuanVAEConfig", "HunyuanVAEArchConfig",
"WanVAEConfig", "WanVAEArchConfig"
]
+36
View File
@@ -0,0 +1,36 @@
from dataclasses import dataclass
from typing import Union
import torch
from fastvideo.v1.configs.models import ArchConfig, ModelConfig
@dataclass
class VAEArchConfig(ArchConfig):
scaling_factor: Union[float, torch.tensor] = 0
temporal_compression_ratio: int = 4
spatial_compression_ratio: int = 8
@dataclass
class VAEConfig(ModelConfig):
arch_config: VAEArchConfig = VAEArchConfig()
# FastVideoVAE-specific parameters
load_encoder: bool = True
load_decoder: bool = True
tile_sample_min_height: int = 256
tile_sample_min_width: int = 256
tile_sample_min_num_frames: int = 16
tile_sample_stride_height: int = 192
tile_sample_stride_width: int = 192
tile_sample_stride_num_frames: int = 12
blend_num_frames: int = 0
use_tiling: bool = True
use_temporal_tiling: bool = True
use_parallel_tiling: bool = True
def __post_init__(self):
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
@@ -0,0 +1,37 @@
from dataclasses import dataclass
from typing import Tuple
from fastvideo.v1.configs.models.vaes.base import VAEConfig, VAEArchConfig
@dataclass
class HunyuanVAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 16
down_block_types: Tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
)
up_block_types: Tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
)
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512)
layers_per_block: int = 2
act_fn: str = "silu"
norm_num_groups: int = 32
scaling_factor: float = 0.476986
spatial_compression_ratio: int = 8
temporal_compression_ratio: int = 4
mid_block_add_attention: bool = True
def __post_init__(self):
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels)-1)
@dataclass
class HunyuanVAEConfig(VAEConfig):
arch_config: VAEArchConfig = HunyuanVAEArchConfig()
@@ -0,0 +1,72 @@
from dataclasses import dataclass
from typing import Tuple
import torch
from fastvideo.v1.configs.models.vaes.base import VAEConfig, VAEArchConfig
@dataclass
class WanVAEArchConfig(VAEArchConfig):
base_dim: int = 96
z_dim: int = 16
dim_mult: Tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
attn_scales: Tuple[float, ...] = ()
temperal_downsample: Tuple[bool, ...] = (False, True, True)
dropout: float = 0.0
latents_mean: Tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
)
latents_std: Tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
)
temporal_compression_ratio = 4
spatial_compression_ratio = 8
def __post_init__(self):
self.scaling_factor: torch.tensor = 1.0 / torch.tensor(self.latents_std).view(
1, self.z_dim, 1, 1, 1)
self.shift_factor: torch.tensor = torch.tensor(self.latents_mean).view(
1, self.z_dim, 1, 1, 1)
@dataclass
class WanVAEConfig(VAEConfig):
arch_config: VAEArchConfig = WanVAEArchConfig()
use_feature_cache: bool = True
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
def __post_init__(self):
self.blend_num_frames = (self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames) * 2
@@ -0,0 +1,9 @@
from fastvideo.v1.configs.pipelines.hunyuan import HunyuanConfig, FastHunyuanConfig
from fastvideo.v1.configs.pipelines.wan import WanT2V480PConfig, WanI2V480PConfig
from fastvideo.v1.configs.pipelines.base import BaseConfig, SlidingTileAttnConfig
from fastvideo.v1.configs.pipelines.registry import get_pipeline_config_cls_for_name
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "get_pipeline_config_cls_for_name"
]
+103
View File
@@ -0,0 +1,103 @@
from dataclasses import dataclass, asdict, fields
from typing import Optional, Dict, Any
import json
from fastvideo.v1.configs.models import ModelConfig, VAEConfig
from fastvideo.v1.utils import shallow_asdict
@dataclass
class BaseConfig:
"""Base configuration for all pipeline architectures."""
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 125
fps: int = 24
# Video generation parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
seed: int = 1024
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
# vae_scale_factor: Optional[int] = None # Deprecated
vae_config: VAEConfig = VAEConfig()
# DiT configuration
num_channels_latents: Optional[int] = None # Deprecated
# Image encoder configuration
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_len: int = -1 # Deprecated
hidden_state_skip_layer: int = 0 # Deprecated
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
neg_prompt: Optional[str] = None
def dump_to_json(self, file_path: str):
output_dict = shallow_asdict(self)
for key, value in output_dict.items():
if isinstance(value, ModelConfig):
model_dict = asdict(value)
# Model Arch Config should be hidden away from the users
model_dict.pop("arch_config")
output_dict[key] = model_dict
with open(file_path, "w") as f:
json.dump(output_dict, f, indent=2)
def load_from_json(self, file_path: str):
with open(file_path, "r") as f:
input_pipeline_dict = json.load(f)
self.update_pipeline_config(input_pipeline_dict)
def update_pipeline_config(
self,
source_pipeline_dict: Dict[str, Any]
) -> None:
for f in fields(self):
key = f.name
if key in source_pipeline_dict:
current_value = getattr(self, key)
new_value = source_pipeline_dict[key]
# If it's a nested ModelConfig, update it recursively
if isinstance(current_value, ModelConfig):
current_value.update_model_config(new_value)
else:
setattr(self, key, new_value)
@dataclass
class SlidingTileAttnConfig(BaseConfig):
"""Configuration for sliding tile attention."""
# Override any BaseConfig defaults as needed
# Add sliding tile specific parameters
window_size: int = 16
stride: int = 8
# You can provide custom defaults for inherited fields
height: int = 576
width: int = 1024
# Additional configuration specific to sliding tile attention
pad_to_square: bool = False
use_overlap_optimization: bool = True
@@ -1,12 +1,16 @@
from dataclasses import dataclass
from fastvideo.v1.configs.base import BaseConfig
from fastvideo.v1.configs.pipelines.base import BaseConfig
from fastvideo.v1.configs.models import VAEConfig
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
@dataclass
class HunyuanConfig(BaseConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
# VAE
vae_config: VAEConfig = HunyuanVAEConfig()
# Denoising stage
embedded_cfg_scale: int = 6
flow_shift: int = 7
@@ -26,6 +30,10 @@ class HunyuanConfig(BaseConfig):
text_encoder_precision_2: str = "fp16"
text_len_2: int = 77
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@dataclass
class FastHunyuanConfig(HunyuanConfig):
@@ -1,14 +1,15 @@
"""Registry for pipeline weight-specific configurations."""
import os
from typing import Dict, Type, Optional, Callable
from typing import Callable, Dict, Optional, Type
from fastvideo.v1.configs.base import BaseConfig
from fastvideo.v1.configs.hunyuan import HunyuanConfig, FastHunyuanConfig
from fastvideo.v1.configs.wan import WanT2V480PConfig, WanI2V480PConfig
from fastvideo.v1.configs.pipelines.base import BaseConfig
from fastvideo.v1.configs.pipelines.hunyuan import HunyuanConfig, FastHunyuanConfig
from fastvideo.v1.configs.pipelines.wan import WanT2V480PConfig, WanI2V480PConfig
from fastvideo.v1.utils import maybe_download_model_index, verify_model_config_and_directory
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import (maybe_download_model_index,
verify_model_config_and_directory)
logger = init_logger(__name__)
@@ -40,8 +41,8 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[BaseConfig]] = {
}
def get_pipeline_config_for_name(
pipeline_name_or_path: str) -> Optional[Type[BaseConfig]]:
def get_pipeline_config_cls_for_name(
pipeline_name_or_path: str) -> Optional[type[BaseConfig]]:
"""Get the appropriate config class for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
@@ -1,12 +1,19 @@
from dataclasses import dataclass
from fastvideo.v1.configs.base import BaseConfig
from fastvideo.v1.configs.pipelines.base import BaseConfig
from fastvideo.v1.configs.models import VAEConfig
from fastvideo.v1.configs.models.vaes import WanVAEConfig
@dataclass
class WanT2V480PConfig(BaseConfig):
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
# WanConfig-specific parameters with defaults
# VAE
vae_config: VAEConfig = WanVAEConfig()
vae_tiling: bool = False
vae_sp: bool = False
# Video parameters
height: int = 480
width: int = 832
@@ -30,6 +37,9 @@ class WanT2V480PConfig(BaseConfig):
# WanConfig-specific added parameters
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@dataclass
class WanI2V480PConfig(WanT2V480PConfig):
@@ -42,3 +52,7 @@ class WanI2V480PConfig(WanT2V480PConfig):
# Precision for each component
image_encoder_precision: str = "fp32"
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@@ -0,0 +1,41 @@
{
"height": 480,
"width": 832,
"num_frames": 81,
"fps": 16,
"num_inference_steps": 50,
"guidance_scale": 3.0,
"seed": 1024,
"guidance_rescale": 0.0,
"embedded_cfg_scale": 6.0,
"flow_shift": 3,
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
"load_encoder": false,
"load_decoder": true,
"tile_sample_min_height": 256,
"tile_sample_min_width": 256,
"tile_sample_min_num_frames": 16,
"tile_sample_stride_height": 192,
"tile_sample_stride_width": 192,
"tile_sample_stride_num_frames": 12,
"blend_num_frames": 8,
"use_tiling": false,
"use_temporal_tiling": false,
"use_parallel_tiling": false,
"use_feature_cache": true
},
"num_channels_latents": null,
"image_encoder_precision": "fp32",
"text_encoder_precision": "fp32",
"text_len": 512,
"hidden_state_skip_layer": 0,
"mask_strategy_file_path": null,
"enable_torch_compile": false,
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
}
@@ -0,0 +1,41 @@
{
"height": 480,
"width": 832,
"num_frames": 81,
"fps": 16,
"num_inference_steps": 40,
"guidance_scale": 5.0,
"seed": 1024,
"guidance_rescale": 0.0,
"embedded_cfg_scale": 6.0,
"flow_shift": 3,
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
"load_encoder": true,
"load_decoder": true,
"tile_sample_min_height": 256,
"tile_sample_min_width": 256,
"tile_sample_min_num_frames": 16,
"tile_sample_stride_height": 192,
"tile_sample_stride_width": 192,
"tile_sample_stride_num_frames": 12,
"blend_num_frames": 8,
"use_tiling": false,
"use_temporal_tiling": false,
"use_parallel_tiling": false,
"use_feature_cache": true
},
"num_channels_latents": null,
"image_encoder_precision": "fp32",
"text_encoder_precision": "fp32",
"text_len": 512,
"hidden_state_skip_layer": 0,
"mask_strategy_file_path": null,
"enable_torch_compile": false,
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
}
+16 -1
View File
@@ -1,5 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.v1.distributed.communication_op import *
from fastvideo.v1.distributed.parallel_state import *
from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size, get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size, get_world_group,
init_distributed_environment, initialize_model_parallel)
from fastvideo.v1.distributed.utils import *
__all__ = [
"init_distributed_environment",
"initialize_model_parallel",
"get_sequence_model_parallel_rank",
"get_sequence_model_parallel_world_size",
"get_tensor_model_parallel_rank",
"get_tensor_model_parallel_world_size",
"cleanup_dist_env_and_memory",
"get_world_group",
]
+1 -2
View File
@@ -9,8 +9,7 @@ from fastvideo.v1.utils import FlexibleArgumentParser
class CLISubcommand:
"""Base class for CLI subcommands"""
def __init__(self) -> None:
self.name = ""
name: str
def cmd(self, args: argparse.Namespace) -> None:
"""Execute the command with the given arguments"""
-4
View File
@@ -73,10 +73,6 @@ class GenerateSubcommand(CLISubcommand):
required=False,
help="Read CLI options from a config YAML file.")
generate_parser.add_argument("--num-gpus",
type=int,
default=1,
help="Number of GPUs to use")
generate_parser.add_argument("--master-port",
type=int,
default=None,
+290
View File
@@ -0,0 +1,290 @@
# SPDX-License-Identifier: Apache-2.0
"""
VideoGenerator module for FastVideo.
This module provides a consolidated interface for generating videos using
diffusion models.
"""
import os
import time
from dataclasses import asdict
from typing import Any, Callable, Dict, List, Optional, Union
import imageio
import numpy as np
import torch
import torchvision
from einops import rearrange
from fastvideo.v1.configs.pipelines import get_pipeline_config_cls_for_name, BaseConfig
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ForwardBatch
from fastvideo.v1.utils import align_to, shallow_asdict
from fastvideo.v1.worker.executor import Executor
logger = init_logger(__name__)
class VideoGenerator:
"""
A unified class for generating videos using diffusion models.
This class provides a simple interface for video generation with rich
customization options, similar to popular frameworks like HF Diffusers.
"""
def __init__(self, fastvideo_args: FastVideoArgs,
executor_class: type[Executor], log_stats: bool):
"""
Initialize the video generator.
Args:
pipeline: The pipeline to use for inference
fastvideo_args: The inference arguments
"""
self.fastvideo_args = fastvideo_args
self.executor = executor_class(fastvideo_args)
@classmethod
def from_pretrained(cls,
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[Union[str | BaseConfig]] = None,
**kwargs) -> "VideoGenerator":
"""
Create a video generator from a pretrained model.
Args:
model_path: Path or identifier for the pretrained model
device: Device to load the model on (e.g., "cuda", "cuda:0", "cpu")
torch_dtype: Data type for model weights (e.g., torch.float16)
**kwargs: Additional arguments to customize model loading
Returns:
The created video generator
Priority level: Default pipeline config < User's pipeline config < User's kwargs
"""
config = None
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, BaseConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = kwargs
else:
config_args = shallow_asdict(config)
config_args.update(kwargs)
fastvideo_args = FastVideoArgs(
model_path=model_path,
device_str=device or "cuda" if torch.cuda.is_available() else "cpu",
**config_args)
fastvideo_args.check_fastvideo_args()
return cls.from_fastvideo_args(fastvideo_args)
@classmethod
def from_fastvideo_args(cls,
fastvideo_args: FastVideoArgs) -> "VideoGenerator":
"""
Create a video generator with the specified arguments.
Args:
fastvideo_args: The inference arguments
Returns:
The created video generator
"""
# Initialize distributed environment if needed
# initialize_distributed_and_parallelism(fastvideo_args)
executor_class = Executor.get_class(fastvideo_args)
return cls(
fastvideo_args=fastvideo_args,
executor_class=executor_class,
log_stats=False, # TODO: implement
)
def generate_video(
self,
prompt: str,
image_path: Optional[str] = None,
negative_prompt: Optional[str] = None,
output_path: Optional[str] = None,
save_video: bool = True,
return_frames: bool = False,
num_inference_steps: Optional[int] = None,
guidance_scale: Optional[float] = None,
num_frames: Optional[int] = None,
height: Optional[int] = None,
width: Optional[int] = None,
fps: Optional[int] = None,
seed: Optional[int] = None,
callback: Optional[Callable[[int, int, torch.Tensor], None]] = None,
callback_steps: int = 1,
) -> Union[Dict[str, Any], List[np.ndarray]]:
"""
Generate a video based on the given prompt.
Args:
prompt: The prompt to use for generation
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)
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)
guidance_scale: Classifier-free guidance scale (overrides fastvideo_args)
num_frames: Number of frames to generate (overrides fastvideo_args)
height: Height of generated video (overrides fastvideo_args)
width: Width of generated video (overrides fastvideo_args)
fps: Frames per second for saved video (overrides fastvideo_args)
seed: Random seed for generation (overrides fastvideo_args)
callback: Callback function called after each step
callback_steps: Number of steps between each callback
Returns:
Either the output dictionary or the list of frames depending on return_frames
"""
# Create a copy of inference args to avoid modifying the original
fastvideo_args = self.fastvideo_args
# Override parameters if provided
if image_path is not None:
fastvideo_args.image_path = image_path
if negative_prompt is not None:
fastvideo_args.neg_prompt = negative_prompt
if num_inference_steps is not None:
fastvideo_args.num_inference_steps = num_inference_steps
if guidance_scale is not None:
fastvideo_args.guidance_scale = guidance_scale
if num_frames is not None:
fastvideo_args.num_frames = num_frames
if height is not None:
fastvideo_args.height = height
if width is not None:
fastvideo_args.width = width
if fps is not None:
fastvideo_args.fps = fps
if seed is not None:
fastvideo_args.seed = seed
# Validate inputs
if not isinstance(prompt, str):
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
# Process negative prompt
if fastvideo_args.neg_prompt is not None:
fastvideo_args.neg_prompt = fastvideo_args.neg_prompt.strip()
# Validate dimensions
if (fastvideo_args.height <= 0 or fastvideo_args.width <= 0
or fastvideo_args.num_frames <= 0):
raise ValueError(
f"Height, width, and num_frames must be positive integers, got "
f"height={fastvideo_args.height}, width={fastvideo_args.width}, "
f"num_frames={fastvideo_args.num_frames}")
if (fastvideo_args.num_frames - 1) % 4 != 0:
raise ValueError(
f"num_frames-1 must be a multiple of 4, got {fastvideo_args.num_frames}"
)
# Calculate sizes
target_height = align_to(fastvideo_args.height, 16)
target_width = align_to(fastvideo_args.width, 16)
# Calculate latent sizes
latents_size = [(fastvideo_args.num_frames - 1) // 4 + 1,
fastvideo_args.height // 8, fastvideo_args.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# Log parameters
debug_str = f"""
height: {target_height}
width: {target_width}
video_length: {fastvideo_args.num_frames}
prompt: {prompt}
neg_prompt: {fastvideo_args.neg_prompt}
seed: {fastvideo_args.seed}
infer_steps: {fastvideo_args.num_inference_steps}
num_videos_per_prompt: {fastvideo_args.num_videos}
guidance_scale: {fastvideo_args.guidance_scale}
n_tokens: {n_tokens}
flow_shift: {fastvideo_args.flow_shift}
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}"""
logger.info(debug_str)
# Prepare batch
device = torch.device(fastvideo_args.device_str)
batch = ForwardBatch(
prompt=prompt,
image_path=fastvideo_args.image_path,
negative_prompt=fastvideo_args.neg_prompt,
num_videos_per_prompt=fastvideo_args.num_videos,
height=fastvideo_args.height,
width=fastvideo_args.width,
num_frames=fastvideo_args.num_frames,
num_inference_steps=fastvideo_args.num_inference_steps,
guidance_scale=fastvideo_args.guidance_scale,
eta=0.0,
n_tokens=n_tokens,
data_type="video" if fastvideo_args.num_frames > 1 else "image",
device=device,
extra={},
)
# Run inference
start_time = time.time()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch
gen_time = time.time() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
# Process outputs
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
# Save video if requested
if save_video:
save_path = output_path or fastvideo_args.output_path
if save_path:
os.makedirs(os.path.dirname(save_path), exist_ok=True)
video_path = os.path.join(save_path, f"{prompt[:100]}.mp4")
imageio.mimsave(video_path, frames, fps=fastvideo_args.fps)
logger.info("Saved video to %s", video_path)
else:
logger.warning("No output path provided, video not saved")
if return_frames:
return frames
else:
return {
"samples": samples,
"prompts": prompt,
"size":
(target_height, target_width, fastvideo_args.num_frames),
"generation_time": gen_time
}
@@ -0,0 +1,28 @@
# Basic Video Generation Tutorial
The `VideoGenerator` class provides the primary Python interface for doing offline video generation, which is interacting with a diffusion pipeline without using a separate inference api server.
## Usage
The first script in this example shows the most basic usage of FastVideo. If you are new to Python and FastVideo, you should start here.
```bash
python fastvideo/v1/examples/inference/basic/basic.py
```
# Basic Walkthrough
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
```python
from fastvideo import VideoGenerator
generator = VideoGenerator.from_pretrained(
"FastVideo/FastHunyuan-Diffusers",
num_gpus=2,
)
prompt = "A beautiful woman in a red dress walking down a street"
video = generator.generate_video(prompt)
```
More to come! These examples and APIs are still under construction!
@@ -0,0 +1,25 @@
from fastvideo import VideoGenerator
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(
"FastVideo/FastHunyuan-Diffusers",
# if num_gpus > 1, FastVideo will automatically handle distributed setup
num_gpus=4,
)
# Generate videos with the same simple API, regardless of GPU count
prompt = "A beautiful woman in a red dress walking down a street"
video = generator.generate_video(prompt)
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = "A beautiful woman in a blue dress walking down a street"
video2 = generator.generate_video(prompt2)
if __name__ == "__main__":
main()
@@ -0,0 +1,81 @@
# FastVideo VideoGenerator Gradio 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
---
## Requirements
- Linux-based OS
- Python 3.10
- Cuda 12.4
- FastVideo
## Installation
```bash
pip install fastvideo
```
## Usage
Run the demo with:
```bash
python fastvideo/v1/examples/inference/gradio/gradio_demo.py
```
This will start a web server at `http://0.0.0.0:7860` where you can access the interface.
---
## Model Initialization
```python
args = FastVideoArgs(model_path="FastVideo/FastHunyuan-Diffusers", num_gpus=2)
generator = VideoGenerator.from_pretrained(
model_path=args.model_path,
num_gpus=args.num_gpus
)
```
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, or even the number of inference steps, *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()`)
3. Returns an output path that Gradio uses to display the generated 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
- Inference steps 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
- **Inference Steps**: More steps can improve quality but take longer
- **Negative Prompt**: Specify what you don't want to see in the video
- **Seed**: Control randomness for reproducible results
@@ -0,0 +1,141 @@
import os
import gradio as gr
import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo import VideoGenerator
if __name__ == "__main__":
args = FastVideoArgs(model_path="FastVideo/FastHunyuan-Diffusers", num_gpus=2)
generator = VideoGenerator.from_pretrained(
model_path=args.model_path,
num_gpus=args.num_gpus
)
def generate_video(
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed=False,
):
if randomize_seed:
seed = torch.randint(0, 1000000, (1, )).item()
if not use_negative_prompt:
negative_prompt = None
generator.generate_video(
prompt=prompt,
negative_prompt=negative_prompt,
num_inference_steps=num_inference_steps,
num_frames=num_frames,
height=height,
width=width,
guidance_scale=guidance_scale,
seed=seed
)
output_path = os.path.join(args.output_path, f"{prompt[:100]}.mp4")
return output_path, seed
examples = [
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
]
with gr.Blocks() as demo:
gr.Markdown("# FastVideo Inference Demo")
with gr.Group():
with gr.Row():
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=1,
placeholder="Enter your prompt",
container=False,
)
run_button = gr.Button("Run", scale=0)
result = gr.Video(label="Result", show_label=False)
with gr.Accordion("Advanced options", open=False):
with gr.Group():
with gr.Row():
height = gr.Slider(
label="Height",
minimum=256,
maximum=1024,
step=32,
value=args.height,
)
width = gr.Slider(label="Width", minimum=256, maximum=1024, step=32, value=args.width)
with gr.Row():
num_frames = gr.Slider(
label="Number of Frames",
minimum=21,
maximum=163,
value=45,
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=args.guidance_scale,
)
num_inference_steps = gr.Slider(
label="Inference Steps",
minimum=4,
maximum=100,
value=6,
)
with gr.Row():
use_negative_prompt = gr.Checkbox(label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=1,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed)
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
seed_output = gr.Number(label="Used Seed")
gr.Examples(examples=examples, inputs=prompt)
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=negative_prompt,
)
run_button.click(
fn=generate_video,
inputs=[
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed,
],
outputs=[result, seed_output],
)
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
+14 -8
View File
@@ -7,8 +7,10 @@ import dataclasses
from contextlib import contextmanager
from typing import List, Optional
from fastvideo.v1.utils import FlexibleArgumentParser
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import FlexibleArgumentParser
from fastvideo.v1.configs.models import VAEConfig
logger = init_logger(__name__)
@@ -19,7 +21,7 @@ class FastVideoArgs:
model_path: str
# Distributed executor backend
distributed_executor_backend: str = "torch"
distributed_executor_backend: str = "mp"
inference_mode: bool = True # if False == training mode
@@ -50,9 +52,10 @@ class FastVideoArgs:
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = False
vae_scale_factor: Optional[int] = None
vae_tiling: bool = True # Might change in between forward passes
vae_sp: bool = False # Might change in between forward passes
# vae_scale_factor: Optional[int] = None # Deprecated
vae_config: VAEConfig = VAEConfig()
# DiT configuration
num_channels_latents: Optional[int] = None
@@ -126,7 +129,7 @@ class FastVideoArgs:
parser.add_argument(
"--distributed-executor-backend",
type=str,
choices=["mp", "ray", "torch"],
choices=["mp"],
default=FastVideoArgs.distributed_executor_backend,
help="The distributed executor backend to use",
)
@@ -431,13 +434,16 @@ class FastVideoArgs:
return cls(**kwargs)
def check_inference_args(self) -> None:
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
if self.tp_size is None:
self.tp_size = self.num_gpus
if self.sp_size is None:
self.sp_size = self.num_gpus
if self.num_gpus < max(self.tp_size, self.sp_size):
self.num_gpus = max(self.tp_size, self.sp_size)
if self.tp_size != self.sp_size:
raise ValueError(
f"tp_size ({self.tp_size}) must be equal to sp_size ({self.sp_size})"
@@ -470,7 +476,7 @@ def prepare_fastvideo_args(argv: List[str]) -> FastVideoArgs:
FastVideoArgs.add_cli_args(parser)
raw_args = parser.parse_args(argv)
fastvideo_args = FastVideoArgs.from_cli_args(raw_args)
fastvideo_args.check_inference_args()
fastvideo_args.check_fastvideo_args()
global _current_fastvideo_args
_current_fastvideo_args = fastvideo_args
return fastvideo_args
+86 -2
View File
@@ -20,6 +20,13 @@ FASTVIDEO_LOGGING_CONFIG_PATH = envs.FASTVIDEO_LOGGING_CONFIG_PATH
FASTVIDEO_LOGGING_LEVEL = envs.FASTVIDEO_LOGGING_LEVEL
FASTVIDEO_LOGGING_PREFIX = envs.FASTVIDEO_LOGGING_PREFIX
RED = '\033[91m'
GREEN = '\033[92m'
RESET = '\033[0;0m'
_warned_local_main_process = False
_warned_main_process = False
_FORMAT = (f"{FASTVIDEO_LOGGING_PREFIX}%(levelname)s %(asctime)s "
"[%(filename)s:%(lineno)d] %(message)s")
_DATE_FORMAT = "%m-%d %H:%M:%S"
@@ -68,6 +75,68 @@ def _print_warning_once(logger: Logger, msg: str) -> None:
logger.warning(msg, stacklevel=2)
# TODO(will): add env variable to control this process-aware logging behavior
def _info(logger: Logger,
msg: object,
*args: Any,
main_process_only: bool = False,
local_main_process_only: bool = True,
**kwargs: Any) -> None:
"""Process-aware INFO level logging function.
This function controls logging behavior based on the process rank, allowing for
selective logging from specific processes in a distributed environment.
Args:
logger: The logger instance to use for logging
msg: The message format string to log
*args: Format string arguments
main_process_only: If True, only log if this is the global main process (RANK=0)
local_main_process_only: If True, only log if this is the local main process (LOCAL_RANK=0)
**kwargs: Additional keyword arguments to pass to the logger.log method
- stacklevel: Defaults to 2 to show the original caller's location
Note:
- When both main_process_only and local_main_process_only are True,
the message will be logged only if both conditions are met
- When both are False, the message will be logged from all processes
- By default, only logs from processes with LOCAL_RANK=0
"""
try:
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
except Exception:
local_rank = 0
rank = 0
is_main_process = rank == 0
is_local_main_process = local_rank == 0
if (main_process_only and is_main_process) or (local_main_process_only
and is_local_main_process):
logger.log(logging.INFO, msg, *args, **kwargs)
global _warned_local_main_process, _warned_main_process
if not _warned_local_main_process and local_main_process_only:
logger.warning(
'%s is_local_main_process is set to True, logging only from the local main process.%s',
GREEN,
RESET,
)
_warned_local_main_process = True
if not _warned_main_process and main_process_only:
logger.warning(
'%s is_main_process_only is set to True, logging only from the main process.%s',
GREEN,
RESET,
)
_warned_main_process = True
if not main_process_only and not local_main_process_only:
logger.log(logging.INFO, msg, *args, **kwargs)
class _FastvideoLogger(Logger):
"""
Note:
@@ -91,6 +160,20 @@ class _FastvideoLogger(Logger):
"""
_print_warning_once(self, msg)
def info( # type: ignore[override]
self,
msg: object,
*args: Any,
main_process_only: bool = False,
local_main_process_only: bool = True,
**kwargs: Any) -> None:
_info(self,
msg,
*args,
main_process_only=main_process_only,
local_main_process_only=local_main_process_only,
**kwargs)
def _configure_fastvideo_root_logger() -> None:
logging_config = dict[str, Any]()
@@ -128,7 +211,6 @@ def _configure_fastvideo_root_logger() -> None:
dictConfig(logging_config)
# TODO: add rank_zero_only log
def init_logger(name: str) -> _FastvideoLogger:
"""The main purpose of this function is to ensure that loggers are
retrieved in such a way that we can be sure the root fastvideo logger has
@@ -139,10 +221,12 @@ def init_logger(name: str) -> _FastvideoLogger:
methods_to_patch = {
"info_once": _print_info_once,
"warning_once": _print_warning_once,
"info": _info,
}
for method_name, method in methods_to_patch.items():
setattr(logger, method_name, MethodType(method, logger))
setattr(logger, method_name,
MethodType(method, logger)) # type: ignore[arg-type]
return cast(_FastvideoLogger, logger)
+5 -3
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import List, Optional, Union
from typing import List, Optional, Tuple, Union
import torch
from torch import nn
@@ -15,7 +15,9 @@ class BaseDiT(nn.Module, ABC):
_param_names_mapping: dict
hidden_size: int
num_attention_heads: int
_supported_attention_backends: List[_Backend] = []
# always supports torch_sdpa
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.TORCH_SDPA, )
def __init_subclass__(cls) -> None:
required_class_attrs = [
@@ -55,5 +57,5 @@ class BaseDiT(nn.Module, ABC):
)
@property
def supported_attention_backends(self) -> List[_Backend]:
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
return self._supported_attention_backends
+7 -9
View File
@@ -92,7 +92,7 @@ class MMDoubleStreamBlock(nn.Module):
num_attention_heads: int,
mlp_ratio: float,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[List[_Backend]] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
prefix: str = "",
):
super().__init__()
@@ -299,7 +299,7 @@ class MMSingleStreamBlock(nn.Module):
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[List[_Backend]] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
prefix: str = "",
):
super().__init__()
@@ -436,9 +436,8 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
lambda n, m: "single" in n and str.isdigit(n.split(".")[-1]),
lambda n, m: "refiner" in n and str.isdigit(n.split(".")[-1]),
]
_supported_attention_backends = [
_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN, _Backend.TORCH_SDPA
]
_supported_attention_backends = (_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
_param_names_mapping = {
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
@@ -572,7 +571,7 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
pooled_projection_dim: int = 768,
rope_theta: int = 256,
qk_norm: str = "rms_norm", #TODO(PY)
prefix="",
prefix="Hunyuan",
):
super().__init__()
hidden_size = attention_head_dim * num_attention_heads
@@ -896,9 +895,8 @@ class IndividualTokenRefinerBlock(nn.Module):
num_heads=num_attention_heads,
head_size=hidden_size // num_attention_heads,
# TODO: remove hardcode; remove STA
supported_attention_backends=[
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA
],
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA),
)
def forward(self, x, c):
+58 -51
View File
@@ -114,14 +114,14 @@ class WanSelfAttention(nn.Module):
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
# Scaled dot product attention
self.attn = LocalAttention(num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=[
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA
])
self.attn = LocalAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA))
def forward(self, x: torch.Tensor, context: torch.Tensor,
context_lens: int):
@@ -163,13 +163,14 @@ class WanT2VCrossAttention(WanSelfAttention):
class WanI2VCrossAttention(WanSelfAttention):
def __init__(
self,
dim: int,
num_heads: int,
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
supported_attention_backends: Optional[List[str]] = None) -> None:
self,
dim: int,
num_heads: int,
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps,
supported_attention_backends)
@@ -210,17 +211,17 @@ class WanI2VCrossAttention(WanSelfAttention):
class WanTransformerBlock(nn.Module):
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[List[_Backend]] = None,
):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
@@ -233,7 +234,8 @@ class WanTransformerBlock(nn.Module):
num_heads=num_heads,
head_size=dim // num_heads,
causal=False,
supported_attention_backends=supported_attention_backends)
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1")
self.hidden_dim = dim
self.num_attention_heads = num_heads
dim_head = dim // num_heads
@@ -351,7 +353,8 @@ class WanTransformer3DModel(BaseDiT):
_fsdp_shard_conditions = [
lambda n, m: "blocks" in n and str.isdigit(n.split(".")[-1]),
]
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
_supported_attention_backends = (_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
_param_names_mapping = {
r"^patch_embedding\.(.*)$":
r"patch_embedding.proj.\1",
@@ -391,25 +394,24 @@ class WanTransformer3DModel(BaseDiT):
r"blocks.\1.self_attn_residual_norm.norm.\2",
}
def __init__(
self,
patch_size: Tuple[int, int, int] = (1, 2, 2),
text_len=512,
num_attention_heads: int = 40,
attention_head_dim: int = 128,
in_channels: int = 16,
out_channels: int = 16,
text_dim: int = 4096,
freq_dim: int = 256,
ffn_dim: int = 13824,
num_layers: int = 40,
cross_attn_norm: bool = True,
qk_norm: str = "rms_norm_across_heads",
eps: float = 1e-6,
image_dim: Optional[int] = None,
added_kv_proj_dim: Optional[int] = None,
rope_max_seq_len: int = 1024,
) -> None:
def __init__(self,
patch_size: Tuple[int, int, int] = (1, 2, 2),
text_len=512,
num_attention_heads: int = 40,
attention_head_dim: int = 128,
in_channels: int = 16,
out_channels: int = 16,
text_dim: int = 4096,
freq_dim: int = 256,
ffn_dim: int = 13824,
num_layers: int = 40,
cross_attn_norm: bool = True,
qk_norm: str = "rms_norm_across_heads",
eps: float = 1e-6,
image_dim: Optional[int] = None,
added_kv_proj_dim: Optional[int] = None,
rope_max_seq_len: int = 1024,
prefix="Wan") -> None:
super().__init__()
inner_dim = num_attention_heads * attention_head_dim
@@ -436,11 +438,16 @@ class WanTransformer3DModel(BaseDiT):
# 3. Transformer blocks
self.blocks = nn.ModuleList([
WanTransformerBlock(inner_dim, ffn_dim, num_attention_heads,
qk_norm, cross_attn_norm, eps,
WanTransformerBlock(inner_dim,
ffn_dim,
num_attention_heads,
qk_norm,
cross_attn_norm,
eps,
added_kv_proj_dim,
self._supported_attention_backends)
for _ in range(num_layers)
self._supported_attention_backends,
prefix=f"{prefix}.blocks.{i}")
for i in range(num_layers)
])
# 4. Output norm & projection
+4 -3
View File
@@ -1,4 +1,4 @@
from typing import List
from typing import Tuple
from torch import nn
@@ -6,7 +6,8 @@ from fastvideo.v1.platforms import _Backend
class BaseEncoder(nn.Module):
_supported_attention_backends: List[_Backend] = []
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.TORCH_SDPA, )
def __init__(self, *args, **kwargs) -> None:
super().__init__()
@@ -19,5 +20,5 @@ class BaseEncoder(nn.Module):
pass
@property
def supported_attention_backends(self) -> List[_Backend]:
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
return self._supported_attention_backends
+2 -2
View File
@@ -469,7 +469,7 @@ class CLIPTextTransformer(nn.Module):
class CLIPTextModel(BaseEncoder):
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
def __init__(
self,
@@ -619,7 +619,7 @@ class CLIPVisionModel(BaseEncoder, SupportsQuant):
config_class = CLIPVisionConfig
main_input_name = "pixel_values"
packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]}
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
def __init__(
self,
+1 -1
View File
@@ -280,7 +280,7 @@ class LlamaDecoderLayer(nn.Module):
class LlamaModel(BaseEncoder):
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
def __init__(self,
config: LlamaConfig,
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import dataclasses
from dataclasses import asdict
import glob
import os
import time
@@ -310,9 +311,11 @@ class VAELoader(ComponentLoader):
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
config.pop("_diffusers_version")
vae_config = fastvideo_args.vae_config
vae_config.update_model_arch(config)
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(**config).to(fastvideo_args.device)
vae = vae_cls(vae_config).to(fastvideo_args.device)
# Find all safetensors files
safetensors_list = glob.glob(
@@ -322,7 +325,7 @@ class VAELoader(ComponentLoader):
safetensors_list
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
loaded = safetensors_load_file(safetensors_list[0])
vae.load_state_dict(loaded)
vae.load_state_dict(loaded, strict=False) # We might only load encoder or decoder
dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae = vae.eval().to(dtype)
@@ -343,6 +346,10 @@ class TransformerLoader(ComponentLoader):
"Only diffusers format is supported.")
model_config.pop("_diffusers_version")
# Config from Diffusers supercedes fastvideo's model config
# dit_config = fastvideo_args.dit_config
# model_config.update(dit_config)
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
# Find all safetensors files
+37 -19
View File
@@ -11,6 +11,7 @@ from diffusers.utils.torch_utils import randn_tensor
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size)
from fastvideo.v1.configs.models import VAEConfig
class ParallelTiledVAE(ABC):
@@ -20,29 +21,36 @@ class ParallelTiledVAE(ABC):
tile_sample_stride_height: int
tile_sample_stride_width: int
tile_sample_stride_num_frames: int
blend_num_frames: int
use_tiling: bool
use_temporal_tiling: bool
use_parallel_tiling: bool
temporal_compression_ratio: int
spatial_compression_ratio: int
scaling_factor: Union[float, torch.tensor]
def __init__(self, *args, **kwargs) -> None:
# Check if subclass has defined all required properties
required_attributes = [
'tile_sample_min_height', 'tile_sample_min_width',
'tile_sample_min_num_frames', 'tile_sample_stride_height',
'tile_sample_stride_width', 'tile_sample_stride_num_frames',
'spatial_compression_ratio', 'temporal_compression_ratio',
'use_tiling', 'use_temporal_tiling', 'use_parallel_tiling',
'scaling_factor'
]
def __init__(self, config: VAEConfig, **kwargs) -> None:
self.config = config
self.arch_config = config.arch_config
self.tile_sample_min_height = config.tile_sample_min_height
self.tile_sample_min_width = config.tile_sample_min_width
self.tile_sample_min_num_frames = config.tile_sample_min_num_frames
self.tile_sample_stride_height = config.tile_sample_stride_height
self.tile_sample_stride_width = config.tile_sample_stride_width
self.tile_sample_stride_num_frames = config.tile_sample_stride_num_frames
self.blend_num_frames = config.blend_num_frames
self.use_tiling = config.use_tiling
self.use_temporal_tiling = config.use_temporal_tiling
self.use_parallel_tiling = config.use_parallel_tiling
for attr in required_attributes:
if not hasattr(self, attr):
raise AttributeError(
f"Subclasses of ParallelVAE must define '{attr}' property")
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
@property
def temporal_compression_ratio(self) -> int:
return self.arch_config.temporal_compression_ratio
@property
def spatial_compression_ratio(self) -> int:
return self.arch_config.spatial_compression_ratio
@property
def scaling_factor(self) -> Union[float, torch.tensor]:
return self.arch_config.scaling_factor
@abstractmethod
def _encode(self, *args, **kwargs) -> torch.Tensor:
@@ -408,6 +416,10 @@ class ParallelTiledVAE(ABC):
tile_sample_stride_height: Optional[int] = None,
tile_sample_stride_width: Optional[int] = None,
tile_sample_stride_num_frames: Optional[int] = None,
blend_num_frames: Optional[int] = None,
use_tiling: Optional[bool] = None,
use_temporal_tiling: Optional[bool] = None,
use_parallel_tiling: Optional[bool] = None,
) -> None:
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
@@ -439,7 +451,13 @@ class ParallelTiledVAE(ABC):
self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height
self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width
self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
if blend_num_frames is not None:
self.blend_num_frames = blend_num_frames
else:
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
self.use_tiling = use_tiling or self.use_tiling
self.use_temporal_tiling = use_temporal_tiling or self.use_temporal_tiling
self.use_parallel_tiling = use_parallel_tiling or self.use_parallel_tiling
def disable_tiling(self) -> None:
r"""
+33 -75
View File
@@ -15,7 +15,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Optional, Tuple, Union
from typing import Optional, Tuple, Union, cast
import numpy as np
import torch
@@ -24,8 +24,8 @@ import torch.nn.functional as F
import torch.utils.checkpoint
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.models.utils import auto_attributes
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig, HunyuanVAEArchConfig
def prepare_causal_attention_mask(
@@ -773,95 +773,53 @@ class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
_supports_gradient_checkpointing = True
@auto_attributes
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
latent_channels: int = 16,
down_block_types: Tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
),
up_block_types: Tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
),
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
layers_per_block: int = 2,
act_fn: str = "silu",
norm_num_groups: int = 32,
scaling_factor: float = 0.476986,
spatial_compression_ratio: int = 8,
temporal_compression_ratio: int = 4,
mid_block_add_attention: bool = True,
load_encoder: bool = True,
load_decoder: bool = True,
config: HunyuanVAEConfig,
) -> None:
super().__init__()
ParallelTiledVAE.__init__(self, config)
arch_config: HunyuanVAEArchConfig = cast(HunyuanVAEArchConfig, config.arch_config)
# TODO(will): only pass in config. We do this by manually defining a
# config for hunyuan vae
self.block_out_channels = block_out_channels
if load_encoder:
self.block_out_channels = arch_config.block_out_channels
if config.load_encoder:
self.encoder = HunyuanVideoEncoder3D(
in_channels=in_channels,
out_channels=latent_channels,
down_block_types=down_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
in_channels=arch_config.in_channels,
out_channels=arch_config.latent_channels,
down_block_types=arch_config.down_block_types,
block_out_channels=arch_config.block_out_channels,
layers_per_block=arch_config.layers_per_block,
norm_num_groups=arch_config.norm_num_groups,
act_fn=arch_config.act_fn,
double_z=True,
mid_block_add_attention=mid_block_add_attention,
temporal_compression_ratio=temporal_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=arch_config.mid_block_add_attention,
temporal_compression_ratio=arch_config.temporal_compression_ratio,
spatial_compression_ratio=arch_config.spatial_compression_ratio,
)
self.quant_conv = nn.Conv3d(2 * latent_channels,
2 * latent_channels,
self.quant_conv = nn.Conv3d(2 * arch_config.latent_channels,
2 * arch_config.latent_channels,
kernel_size=1)
if load_decoder:
if config.load_decoder:
self.decoder = HunyuanVideoDecoder3D(
in_channels=latent_channels,
out_channels=out_channels,
up_block_types=up_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
time_compression_ratio=temporal_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=mid_block_add_attention,
in_channels=arch_config.latent_channels,
out_channels=arch_config.out_channels,
up_block_types=arch_config.up_block_types,
block_out_channels=arch_config.block_out_channels,
layers_per_block=arch_config.layers_per_block,
norm_num_groups=arch_config.norm_num_groups,
act_fn=arch_config.act_fn,
time_compression_ratio=arch_config.temporal_compression_ratio,
spatial_compression_ratio=arch_config.spatial_compression_ratio,
mid_block_add_attention=arch_config.mid_block_add_attention,
)
self.post_quant_conv = nn.Conv3d(latent_channels,
latent_channels,
self.post_quant_conv = nn.Conv3d(arch_config.latent_channels,
arch_config.latent_channels,
kernel_size=1)
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
# intermediate tiles together, the memory requirement can be lowered.
self.use_tiling = True
self.use_temporal_tiling = True
self.use_parallel_tiling = True
self.scaling_factor = scaling_factor
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 16
# The minimal distance between two spatial tiles
self.tile_sample_stride_height = 192
self.tile_sample_stride_width = 192
self.tile_sample_stride_num_frames = 12
ParallelTiledVAE.__init__(self)
def _encode(self, x: torch.Tensor) -> torch.Tensor:
x = self.encoder(x)
enc = self.quant_conv(x)
+34 -94
View File
@@ -14,17 +14,17 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Optional, Tuple, Union
import contextvars
from contextlib import contextmanager
from typing import Optional, Tuple, Union, cast
import torch
import torch.nn as nn
import torch.nn.functional as F
from contextlib import contextmanager
import contextvars
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.models.utils import auto_attributes
from fastvideo.v1.models.vaes.common import ParallelTiledVAE, DiagonalGaussianDistribution
from fastvideo.v1.configs.models.vaes import WanVAEConfig, WanVAEArchConfig
CACHE_T = 2
@@ -780,96 +780,34 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
_supports_gradient_checkpointing = False
@auto_attributes
def __init__(self,
base_dim: int = 96,
z_dim: int = 16,
dim_mult: Tuple[int, ...] = (1, 2, 4, 4),
num_res_blocks: int = 2,
attn_scales: Tuple[float, ...] = (),
temperal_downsample: Tuple[bool, ...] = (False, True, True),
dropout: float = 0.0,
latents_mean: Tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
),
latents_std: Tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
),
load_encoder: bool = True,
load_decoder: bool = True) -> None:
config: WanVAEConfig,
) -> None:
super().__init__()
ParallelTiledVAE.__init__(self, config)
self.z_dim = z_dim
self.temperal_downsample = list(temperal_downsample)
self.temperal_upsample = list(temperal_downsample)[::-1]
self.latents_mean = list(latents_mean)
self.latents_std = list(latents_std)
self.scaling_factor = 1.0 / torch.tensor(self.config.latents_std).view(
1, self.config.z_dim, 1, 1, 1)
self.shift_factor = torch.tensor(self.config.latents_mean).view(
1, self.config.z_dim, 1, 1, 1)
self.arch_config = cast(WanVAEArchConfig, self.arch_config)
self.z_dim = self.arch_config.z_dim
self.temperal_downsample = list(self.arch_config.temperal_downsample)
self.temperal_upsample = list(self.arch_config.temperal_downsample)[::-1]
self.latents_mean = list(self.arch_config.latents_mean)
self.latents_std = list(self.arch_config.latents_std)
self.shift_factor = self.arch_config.shift_factor
if load_encoder:
self.encoder = WanEncoder3d(base_dim, z_dim * 2, dim_mult,
num_res_blocks, attn_scales,
self.temperal_downsample, dropout)
self.quant_conv = WanCausalConv3d(z_dim * 2, z_dim * 2, 1)
self.post_quant_conv = WanCausalConv3d(z_dim, z_dim, 1)
if config.load_encoder:
self.encoder = WanEncoder3d(self.arch_config.base_dim, self.z_dim * 2, self.arch_config.dim_mult,
self.arch_config.num_res_blocks, self.arch_config.attn_scales,
self.temperal_downsample, self.arch_config.dropout)
self.quant_conv = WanCausalConv3d(self.z_dim * 2, self.z_dim * 2, 1)
self.post_quant_conv = WanCausalConv3d(self.z_dim, self.z_dim, 1)
if load_decoder:
self.decoder = WanDecoder3d(base_dim, z_dim, dim_mult,
num_res_blocks, attn_scales,
self.temperal_upsample, dropout)
if config.load_decoder:
self.decoder = WanDecoder3d(self.arch_config.base_dim, self.z_dim, self.arch_config.dim_mult,
self.arch_config.num_res_blocks, self.arch_config.attn_scales,
self.temperal_upsample, self.arch_config.dropout)
self.use_tiling = True
self.use_temporal_tiling = False
self.use_parallel_tiling = False
self.spatial_compression_ratio = 8
self.temporal_compression_ratio = 4
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 16
# The minimal distance between two spatial tiles
self.tile_sample_stride_height = 192
self.tile_sample_stride_width = 192
self.tile_sample_stride_num_frames = 12
# Whether to use the feature cache algorithm used by diffusers and Wan2.1
self.use_feature_cache = True # default to True for best performance
ParallelTiledVAE.__init__(self)
self.use_feature_cache = config.use_feature_cache
def clear_cache(self) -> None:
@@ -880,13 +818,15 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
count += 1
return count
self._conv_num = _count_conv3d(self.decoder)
self._conv_idx = 0
self._feat_map = [None] * self._conv_num
if self.config.load_decoder:
self._conv_num = _count_conv3d(self.decoder)
self._conv_idx = 0
self._feat_map = [None] * self._conv_num
# cache encode
self._enc_conv_num = _count_conv3d(self.encoder)
self._enc_conv_idx = 0
self._enc_feat_map = [None] * self._enc_conv_num
if self.config.load_encoder:
self._enc_conv_num = _count_conv3d(self.encoder)
self._enc_conv_idx = 0
self._enc_feat_map = [None] * self._enc_conv_num
def encode(self, x: torch.Tensor) -> torch.Tensor:
if self.use_feature_cache:
@@ -6,8 +6,6 @@ This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
from diffusers.image_processor import VaeImageProcessor
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
@@ -71,14 +69,6 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
"""
Initialize the pipeline.
"""
vae_scale_factor = 2**(len(self.get_module("vae").block_out_channels) -
1)
fastvideo_args.vae_scale_factor = vae_scale_factor
self.image_processor = VaeImageProcessor(
vae_scale_factor=vae_scale_factor)
self.add_module("image_processor", self.image_processor)
num_channels_latents = self.get_module("transformer").in_channels
fastvideo_args.num_channels_latents = num_channels_latents
@@ -7,8 +7,8 @@ This module contains implementations of image encoding stages for diffusion pipe
import torch
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vision_utils import load_image
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -7,8 +7,8 @@ This module contains implementations of prompt encoding stages for diffusion pip
import torch
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
+3 -2
View File
@@ -10,6 +10,7 @@ from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.utils import PRECISION_TO_TYPE
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
logger = init_logger(__name__)
@@ -22,8 +23,8 @@ class DecodingStage(PipelineStage):
output format (e.g., pixel values).
"""
def __init__(self, vae) -> None:
self.vae = vae
def __init__(self, vae: ParallelTiledVAE) -> None:
self.vae: ParallelTiledVAE = vae
def forward(
self,
+10 -6
View File
@@ -13,11 +13,12 @@ from tqdm.auto import tqdm
from fastvideo.v1.attention import get_attn_backend
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size)
get_sequence_model_parallel_world_size,
get_world_group)
from fastvideo.v1.distributed.communication_op import (
sequence_model_parallel_all_gather)
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
@@ -178,10 +179,9 @@ class DenoisingStage(PipelineStage):
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=[
supported_attention_backends=(
_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN,
_Backend.TORCH_SDPA
] # hack
_Backend.TORCH_SDPA) # hack
)
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
@@ -304,7 +304,11 @@ class DenoisingStage(PipelineStage):
Returns:
A tqdm progress bar.
"""
return tqdm(iterable=iterable, total=total)
local_rank = get_world_group().local_rank
if local_rank == 0:
return tqdm(iterable=iterable, total=total)
else:
return tqdm(iterable=iterable, total=total, disable=True)
def rescale_noise_cfg(self,
noise_cfg,
+3 -2
View File
@@ -15,6 +15,7 @@ from fastvideo.v1.models.vision_utils import (get_default_height_width,
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.utils import PRECISION_TO_TYPE
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
logger = init_logger(__name__)
@@ -27,8 +28,8 @@ class EncodingStage(PipelineStage):
input format (e.g., latents).
"""
def __init__(self, vae) -> None:
self.vae = vae
def __init__(self, vae: ParallelTiledVAE) -> None:
self.vae: ParallelTiledVAE = vae
def forward(
self,
@@ -6,7 +6,6 @@ from diffusers.utils.torch_utils import randn_tensor
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
@@ -21,10 +20,9 @@ class LatentPreparationStage(PipelineStage):
denoised during the diffusion process.
"""
def __init__(self, scheduler, vae=None) -> None:
def __init__(self, scheduler) -> None:
super().__init__()
self.scheduler = scheduler
self.vae = vae
def forward(
self,
@@ -44,7 +42,7 @@ class LatentPreparationStage(PipelineStage):
# Adjust video length based on VAE version if needed
if hasattr(self, 'adjust_video_length'):
batch = self.adjust_video_length(self.vae, batch, fastvideo_args)
batch = self.adjust_video_length(batch, fastvideo_args)
# Determine batch size
if isinstance(batch.prompt, list):
batch_size = len(batch.prompt)
@@ -70,15 +68,14 @@ class LatentPreparationStage(PipelineStage):
raise ValueError("Height and width must be provided")
assert fastvideo_args.num_channels_latents is not None
assert fastvideo_args.vae_scale_factor is not None
# Calculate latent shape
shape = (
batch_size,
fastvideo_args.num_channels_latents,
num_frames,
height // fastvideo_args.vae_scale_factor,
width // fastvideo_args.vae_scale_factor,
height // fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
width // fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
)
# Validate generator if it's a list
@@ -106,7 +103,7 @@ class LatentPreparationStage(PipelineStage):
return batch
def adjust_video_length(self, vae: ParallelTiledVAE, batch: ForwardBatch,
def adjust_video_length(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""
Adjust video length based on VAE version.
@@ -119,7 +116,7 @@ class LatentPreparationStage(PipelineStage):
The batch with adjusted video length.
"""
video_length = batch.num_frames
temporal_scale_factor = vae.temporal_compression_ratio if vae is not None else 4
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
# TODO
batch.num_frames = (video_length - 1) // temporal_scale_factor + 1
return batch
@@ -9,8 +9,8 @@ from typing import TypedDict
import torch
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
+1 -1
View File
@@ -6,8 +6,8 @@ This module contains implementations of prompt encoding stages for diffusion pip
"""
import torch
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
@@ -9,10 +9,13 @@ using the modular pipeline architecture.
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
# isort: off
from fastvideo.v1.pipelines.stages import (
CLIPImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
EncodingStage, InputValidationStage, LatentPreparationStage,
T5EncodingStage, TimestepPreparationStage)
# isort: on
logger = init_logger(__name__)
@@ -51,8 +54,7 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
@@ -69,9 +71,6 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
"""
Initialize the pipeline.
"""
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
fastvideo_args.vae_scale_factor = vae_scale_factor
num_channels_latents = self.get_module("transformer").out_channels
fastvideo_args.num_channels_latents = num_channels_latents
+1 -5
View File
@@ -47,8 +47,7 @@ class WanPipeline(ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
@@ -62,9 +61,6 @@ class WanPipeline(ComposedPipelineBase):
"""
Initialize the pipeline.
"""
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
fastvideo_args.vae_scale_factor = vae_scale_factor
num_channels_latents = self.get_module("transformer").in_channels
fastvideo_args.num_channels_latents = num_channels_latents
+12 -16
View File
@@ -8,8 +8,11 @@ import torch
from safetensors.torch import load_file
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.hunyuanvae import (
AutoencoderKLHunyuanVideo as MyHunyuanVAE)
# from fastvideo.v1.models.vaes.hunyuanvae import (
# AutoencoderKLHunyuanVideo as MyHunyuanVAE)
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.models.loader.component_loader import VAELoader
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
from fastvideo.v1.utils import maybe_download_model
logger = init_logger(__name__)
@@ -31,21 +34,14 @@ REFERENCE_LATENT = -105.51324462890625
@pytest.mark.usefixtures("distributed_setup")
def test_hunyuan_vae():
device = torch.device("cuda:0")
# Initialize the two model implementations
config = json.load(open(CONFIG_PATH))
config.pop("_class_name")
config.pop("_diffusers_version")
model = MyHunyuanVAE(**config).to(torch.bfloat16)
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
args.device = device
args.vae_config = HunyuanVAEConfig()
loaded = load_file(os.path.join(VAE_PATH,
"diffusion_pytorch_model.safetensors"))
model.load_state_dict(loaded)
# Set model to eval mode
model.eval()
# Move to GPU
model = model.to(device)
loader = VAELoader()
model = loader.load(VAE_PATH, "", args)
model.enable_tiling(tile_sample_min_height=32,
tile_sample_min_width=32,
+9 -11
View File
@@ -9,6 +9,7 @@ from diffusers import AutoencoderKLWan
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.utils import maybe_download_model
logger = init_logger(__name__)
@@ -30,6 +31,7 @@ def test_wan_vae():
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
args.device = device
args.vae_config = WanVAEConfig()
loader = VAELoader()
model2 = loader.load(VAE_PATH, "", args)
@@ -77,23 +79,19 @@ def test_wan_vae():
# Test decoding
logger.info("Testing decoding...")
latent1_tensor = latent1.mode()
latents_mean = (torch.tensor(model1.config.latents_mean).view(
mean1 = (torch.tensor(model1.config.latents_mean).view(
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype))
latents_std = 1.0 / torch.tensor(model1.config.latents_std).view(
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
std1 = (1.0 / torch.tensor(model1.config.latents_std).view(
1, model1.config.z_dim, 1, 1, 1)).to(input_tensor.device,
input_tensor.dtype)
latent1_tensor = latent1_tensor / latents_std + latents_mean
latent1_tensor = latent1_tensor / std1 + mean1
output1 = model1.decode(latent1_tensor).sample
mean2 = model2.config.arch_config.shift_factor.to(input_tensor.device, input_tensor.dtype)
std2 = model2.config.arch_config.scaling_factor.to(input_tensor.device, input_tensor.dtype)
latent2_tensor = latent2.mode()
latents_mean = (torch.tensor(model2.config.latents_mean).view(
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype))
latents_std = 1.0 / torch.tensor(model2.config.latents_std).view(
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype)
latent2_tensor = latent2_tensor / latents_std + latents_mean
latent2_tensor = latent2_tensor / std2 + mean2
output2 = model2.decode(latent2_tensor)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
+37 -10
View File
@@ -2,19 +2,23 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/utils.py
import argparse
import ctypes
import hashlib
import importlib
import inspect
import json
import math
import os
import signal
import sys
import tempfile
from functools import wraps, partial
from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast, Callable
from dataclasses import asdict, fields
import cloudpickle
import traceback
from dataclasses import asdict, fields, is_dataclass
from functools import partial, wraps
from typing import (Any, Callable, Dict, List, Optional, Tuple, Type, TypeVar,
Union, cast)
import cloudpickle
import filelock
import torch
import yaml
@@ -473,6 +477,7 @@ def maybe_download_model_index(model_name_or_path: str) -> Dict[str, Any]:
The parsed model_index.json as a dictionary
"""
import tempfile
from huggingface_hub import hf_hub_download
# If it's a local path, verify it directly
@@ -545,12 +550,34 @@ def run_method(obj: Any, method: Union[str, bytes, Callable], args: tuple[Any],
func = partial(method, obj) # type: ignore
return func(*args, **kwargs)
def shallow_asdict(obj):
if not is_dataclass(obj):
raise TypeError("Expected dataclass instance")
return {f.name: getattr(obj, f.name) for f in fields(obj)}
def diff_keys(a, b):
return [k for k in asdict(a) if asdict(a)[k] != asdict(b)[k]]
def kill_itself_when_parent_died() -> None:
# if sys.platform == "linux":
# sigkill this process when parent worker manager dies
PR_SET_PDEATHSIG = 1
libc = ctypes.CDLL("libc.so.6")
libc.prctl(PR_SET_PDEATHSIG, signal.SIGKILL)
# else:
# logger.warning("kill_itself_when_parent_died is only supported in linux.")
def update_in_place(target, source, ignore_fields=()):
for f in fields(target):
if hasattr(source, f.name) and f.name not in list(ignore_fields):
setattr(target, f.name, getattr(source, f.name))
def get_exception_traceback() -> str:
etype, value, tb = sys.exc_info()
err_str = "".join(traceback.format_exception(etype, value, tb))
return err_str
class TypeBasedDispatcher:
def __init__(self, mapping: List[Tuple[Type, Callable]]):
self._mapping = mapping
def __call__(self, obj: Any):
for ty, fn in self._mapping:
if isinstance(obj, ty):
return fn(obj)
raise ValueError(f"Invalid object: {obj}")
+78
View File
@@ -0,0 +1,78 @@
from abc import ABC, abstractmethod
from typing import (Any, Callable, Dict, List, Optional, Tuple, TypeVar, Union,
cast)
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.pipelines import ForwardBatch
from fastvideo.v1.utils import init_logger
logger = init_logger(__name__)
_R = TypeVar("_R")
class Executor(ABC):
def __init__(self, fastvideo_args: FastVideoArgs):
self.fastvideo_args = fastvideo_args
self._init_executor()
@abstractmethod
def _init_executor(self) -> None:
raise NotImplementedError
@classmethod
def get_class(cls, fastvideo_args: FastVideoArgs) -> type["Executor"]:
if fastvideo_args.distributed_executor_backend == "mp":
from fastvideo.v1.worker.multiproc_executor import MultiprocExecutor
return cast(type["Executor"], MultiprocExecutor)
else:
raise ValueError(
f"Unsupported distributed executor backend: {fastvideo_args.distributed_executor_backend}"
)
def execute_forward(
self,
forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
outputs: List[Dict[str,
Any]] = self.collective_rpc("execute_forward",
kwargs={
"forward_batch":
forward_batch,
"fastvideo_args":
fastvideo_args
})
return cast(ForwardBatch, outputs[0]["output_batch"])
@abstractmethod
def collective_rpc(self,
method: Union[str, Callable[..., _R]],
timeout: Optional[float] = None,
args: Tuple = (),
kwargs: Optional[Dict[str, Any]] = None) -> List[_R]:
"""
Execute an RPC call on all workers.
Args:
method: Name of the worker method to execute, or a callable that
is serialized and sent to all workers to execute.
If the method is a callable, it should accept an additional
`self` argument, in addition to the arguments passed in `args`
and `kwargs`. The `self` argument will be the worker object.
timeout: Maximum time in seconds to wait for execution. Raises a
:exc:`TimeoutError` on timeout. `None` means wait indefinitely.
args: Positional arguments to pass to the worker method.
kwargs: Keyword arguments to pass to the worker method.
Returns:
A list containing the results from each worker.
Note:
It is recommended to use this API to only pass control messages,
and set up data-plane communication to pass data.
"""
raise NotImplementedError
+238
View File
@@ -0,0 +1,238 @@
import contextlib
import faulthandler
import gc
import multiprocessing as mp
import os
import signal
import sys
from typing import Any, Dict, Optional, TextIO, cast
import psutil
import torch
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory,
init_distributed_environment,
initialize_model_parallel)
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ForwardBatch, build_pipeline
from fastvideo.v1.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):
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.init_device()
# Init request dispatcher
# TODO(will): add request dispatcher: use TypeBasedDispatcher from
# utils.py
# self._request_dispatcher = TypeBasedDispatcher(
# [
# (RpcReqInput, self.handle_rpc_request),
# (GenerateRequest, self.handle_generate_request),
# (ExpertDistributionReq, self.expert_distribution_handle),
# ]
# )
def init_device(self) -> None:
"""Initialize the device for the worker."""
assert self.fastvideo_args.device_str is not None
if self.fastvideo_args.device_str.startswith("cuda"):
# torch.distributed.all_reduce does not free the input tensor until
# the synchronization point. This causes the memory usage to grow
# as the number of all_reduce calls increases. This env var disables
# this behavior.
# Related issue:
# https://discuss.pytorch.org/t/cuda-allocation-lifetime-for-inputs-to-distributed-all-reduce/191573
os.environ["TORCH_NCCL_AVOID_RECORD_STREAMS"] = "1"
# This env var set by Ray causes exceptions with graph building.
os.environ.pop("NCCL_ASYNC_ERROR_HANDLING", None)
self.device = torch.device(f"cuda:{self.local_rank}")
torch.cuda.set_device(self.device)
# _check_if_gpu_supports_dtype(self.model_config.dtype)
gc.collect()
torch.cuda.empty_cache()
self.init_gpu_memory = torch.cuda.mem_get_info()[0]
else:
raise ValueError(
f"Unsupported device: {self.fastvideo_args.device_str}")
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
os.environ["LOCAL_RANK"] = str(self.local_rank)
os.environ["RANK"] = str(self.rank)
# Initialize the distributed environment.
init_worker_distributed_environment(self.fastvideo_args, self.rank,
self.distributed_init_method,
self.local_rank)
self.pipeline = build_pipeline(self.fastvideo_args)
def execute_forward(self, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
self.fastvideo_args.num_inference_steps = fastvideo_args.num_inference_steps
output_batch = self.pipeline.forward(forward_batch, self.fastvideo_args)
return cast(ForwardBatch, output_batch)
def shutdown(self) -> Dict[str, Any]:
"""Gracefully shut down the worker process"""
logger.info("Worker %d shutting down...",
self.rank,
local_main_process_only=False)
# Clean up resources
if hasattr(self, 'pipeline') and self.pipeline is not None:
# Clean up pipeline resources if needed
pass
# Release CUDA resources
if torch.cuda.is_available():
torch.cuda.empty_cache()
# Destroy the distributed environment
cleanup_dist_env_and_memory(shutdown_ray=False)
logger.info("Worker %d shutdown complete",
self.rank,
local_main_process_only=False)
return {"status": "shutdown_complete"}
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:
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)
self.pipe.send({"output_batch": output_batch.output.cpu()})
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}"})
def init_worker_distributed_environment(
fastvideo_args: FastVideoArgs,
rank: int,
distributed_init_method: Optional[str] = None,
local_rank: int = -1,
) -> None:
"""Initialize distributed environment and model parallelism."""
world_size = fastvideo_args.num_gpus
torch.cuda.set_device(local_rank)
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
device_str = f"cuda:{local_rank}"
fastvideo_args.device_str = device_str
fastvideo_args.device = torch.device(device_str)
assert fastvideo_args.sp_size is not None
assert fastvideo_args.tp_size is not None
initialize_model_parallel(
sequence_model_parallel_size=fastvideo_args.sp_size,
tensor_model_parallel_size=fastvideo_args.tp_size,
)
def run_worker_process(fastvideo_args: FastVideoArgs, local_rank: int,
rank: int, pipe):
# 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)
logger.info("Worker %d sending ready", rank)
pipe.send({
"status": "ready",
"local_rank": local_rank,
})
try:
worker.event_loop()
except KeyboardInterrupt:
logger.info(
"Worker %d received KeyboardInterrupt, shutting down...",
rank,
local_main_process_only=False)
worker.shutdown()
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]
+153
View File
@@ -0,0 +1,153 @@
import atexit
import contextlib
import multiprocessing as mp
import time
from multiprocessing.process import BaseProcess
from typing import Any, Callable, List, Optional, Union, cast
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.worker.executor import Executor
from fastvideo.v1.worker.gpu_worker import run_worker_process
logger = init_logger(__name__)
class MultiprocExecutor(Executor):
def _init_executor(self) -> None:
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 = []
# 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))
worker.start()
self.workers.append(worker)
logger.info("Workers: %s", self.workers)
# Wait for all workers to be ready
for idx, pipe in enumerate(self.worker_pipes):
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)
# Register shutdown on exit
atexit.register(self.shutdown)
def execute_forward(self, forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
responses = self.collective_rpc("execute_forward",
kwargs={
"forward_batch": forward_batch,
"fastvideo_args": fastvideo_args
})
return cast(ForwardBatch, responses[0]["output_batch"])
def collective_rpc(self,
method: Union[str, Callable],
timeout: Optional[float] = None,
args: tuple = (),
kwargs: Optional[dict] = None) -> list[Any]:
kwargs = kwargs or {}
try:
for pipe in self.worker_pipes:
pipe.send({"method": method, "args": args, "kwargs": kwargs})
responses = []
for pipe in self.worker_pipes:
response = pipe.recv()
responses.append(response)
return responses
except TimeoutError as e:
raise TimeoutError(f"RPC call to {method} timed out.") from e
except Exception as e:
# Re-raise any other exceptions
raise e
def shutdown(self) -> None:
"""Properly shut down the executor and its workers"""
if hasattr(self, 'shutting_down') and self.shutting_down:
return # Prevent multiple shutdown calls
logger.info("Shutting down MultiprocExecutor...")
self.shutting_down = True
# First try gentle termination
try:
# Send termination message to all workers
for pipe in self.worker_pipes:
with contextlib.suppress(Exception):
pipe.send({"method": "shutdown", "args": (), "kwargs": {}})
# Give workers some time to exit gracefully
start_time = time.time()
while time.time() - start_time < 5.0: # 5 seconds timeout
if all(not worker.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()
# Final timeout for terminate
start_time = time.time()
while time.time() - start_time < 2.0: # 2 seconds timeout
if all(not worker.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)
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()
# Clean up pipes
for pipe in self.worker_pipes:
with contextlib.suppress(Exception):
pipe.close()
self.workers = []
self.worker_pipes = []
logger.info("MultiprocExecutor shutdown complete")
def __del__(self):
"""Ensure cleanup on garbage collection"""
self.shutdown()
def __enter__(self):
"""Support for context manager protocol"""
return self
def __exit__(self, exc_type, exc_val, exc_tb):
"""Ensure cleanup when exiting context"""
self.shutdown()
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[project]
name = "fastvideo"
version = "0.0.1.dev1"
version = "0.0.2"
description = "FastVideo"
readme = "README.md"
requires-python = ">=3.8"
@@ -0,0 +1,24 @@
#!/bin/bash
num_gpus=4
export MODEL_BASE=FastVideo/FastHunyuan-Diffusers
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--sp_size $num_gpus \
--tp_size $num_gpus \
--height 720 \
--width 1280 \
--num_frames 125 \
--num_inference_steps 6 \
--guidance_scale 1 \
--embedded_cfg_scale 6 \
--flow_shift 17 \
--prompt_path ./assets/prompt.txt \
--seed 1024 \
--output_path outputs_video/ \
--model_path $MODEL_BASE \
--vae-sp
+27
View File
@@ -0,0 +1,27 @@
#!/bin/bash
num_gpus=2
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/v1/sample/v1_fastvideo_inference.py \
--sp_size $num_gpus \
--tp_size $num_gpus \
--height 768 \
--width 1280 \
--num_frames 69 \
--num_inference_steps 50 \
--fps 16 \
--guidance_scale 5.0 \
--prompt_path ./assets/prompt.txt \
--neg_prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--seed 12345 \
--output_path outputs_video/ \
--model_path $MODEL_BASE \
--vae-sp \
--text-encoder-precision "fp32" \
--use-cpu-offload