Compare commits
19
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
38ee9dc3b4 | ||
|
|
77b013fb8a | ||
|
|
c31efe1234 | ||
|
|
c056b89aea | ||
|
|
eac79b753f | ||
|
|
4d58cf20d0 | ||
|
|
52c93ecc9d | ||
|
|
42d63166ac | ||
|
|
6db20345a2 | ||
|
|
ad27ea596c | ||
|
|
9aadb4bf8c | ||
|
|
bd941df271 | ||
|
|
8a73876d3b | ||
|
|
1483a1138a | ||
|
|
5e243d8292 | ||
|
|
b0c66d3200 | ||
|
|
c86da2c736 | ||
|
|
bae2a19dcf | ||
|
|
057686f59d |
@@ -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:
|
||||
|
||||
@@ -6,6 +6,7 @@ on:
|
||||
- main
|
||||
paths:
|
||||
- 'pyproject.toml' # Trigger when pyproject.toml changes
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
@@ -41,7 +42,7 @@ jobs:
|
||||
|
||||
build-publish-main:
|
||||
needs: check-version-change
|
||||
if: needs.check-version-change.outputs.version-changed == 'true'
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
@@ -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() }}
|
||||
|
||||
@@ -6,6 +6,7 @@ on:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/sliding_tile_attention/setup.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
@@ -43,7 +44,7 @@ jobs:
|
||||
build_wheels:
|
||||
name: Build Wheel
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
@@ -57,6 +58,29 @@ jobs:
|
||||
cuda-version: ['12.4.1', '12.5.1', '12.6.3']
|
||||
|
||||
steps:
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
@@ -145,7 +169,7 @@ jobs:
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels, check-version-change]
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ubuntu-22.04
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
+4
-1
@@ -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
|
||||
@@ -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
@@ -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
|
||||
@@ -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
@@ -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"
|
||||
|
||||
@@ -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 |
@@ -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
|
||||
|
||||

|
||||
|
||||
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"
|
||||
```
|
||||
|
||||

|
||||
|
||||
After deploying, the pod will take a few minutes to pull the image and start the SSH service.
|
||||
|
||||

|
||||
|
||||
### 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/
|
||||
```
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -69,6 +69,7 @@ sliding_tile_attention/demo
|
||||
:caption: Inference
|
||||
:maxdepth: 1
|
||||
|
||||
inference/wanvideo
|
||||
inference/stepvideo
|
||||
inference/hunyuanvideo
|
||||
inference/fasthunyuan
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
|
||||
|
||||
__all__ = ["VideoGenerator"]
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing import (TYPE_CHECKING, Any, Dict, Generic, Optional, Protocol, Set,
|
||||
Type, TypeVar)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
import torch
|
||||
@@ -154,7 +154,7 @@ class AttentionMetadataBuilder(ABC, Generic[T]):
|
||||
self,
|
||||
current_timestep: int,
|
||||
forward_batch: "ForwardBatch",
|
||||
inference_args: "InferenceArgs",
|
||||
fastvideo_args: "FastVideoArgs",
|
||||
) -> T:
|
||||
"""Build attention metadata with on-device tensors."""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -12,7 +12,7 @@ from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
@@ -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
|
||||
@@ -77,7 +96,7 @@ class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
self,
|
||||
current_timestep: int,
|
||||
forward_batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> SlidingTileAttentionMetadata:
|
||||
|
||||
return SlidingTileAttentionMetadata(current_timestep=current_timestep, )
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
@@ -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"
|
||||
]
|
||||
@@ -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"
|
||||
]
|
||||
@@ -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
|
||||
@@ -0,0 +1,47 @@
|
||||
from dataclasses import dataclass
|
||||
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
|
||||
num_inference_steps: int = 50
|
||||
|
||||
# Text encoding stage
|
||||
hidden_state_skip_layer: int = 2
|
||||
text_len: int = 256
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precision: str = "fp16"
|
||||
|
||||
# HunyuanConfig-specific added parameters
|
||||
# Secondary text encoder
|
||||
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):
|
||||
"""Configuration specifically optimized for FastHunyuan weights."""
|
||||
|
||||
# Override HunyuanConfig defaults
|
||||
num_inference_steps: int = 6
|
||||
flow_shift: int = 17
|
||||
|
||||
# No need to re-specify guidance_scale or embedded_cfg_scale as they
|
||||
# already have the desired values from HunyuanConfig
|
||||
@@ -0,0 +1,78 @@
|
||||
"""Registry for pipeline weight-specific configurations."""
|
||||
|
||||
import os
|
||||
from typing import Callable, Dict, Optional, Type
|
||||
|
||||
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.logger import init_logger
|
||||
from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
verify_model_config_and_directory)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Registry maps specific model weights to their config classes
|
||||
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[BaseConfig]] = {
|
||||
"FastVideo/FastHunyuan-Diffusers": FastHunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
# For determining pipeline type from model ID
|
||||
PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
|
||||
"hunyuan": lambda id: "hunyuan" in id.lower(),
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
# Fallback configs when exact match isn't found but architecture is detected
|
||||
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[BaseConfig]] = {
|
||||
"hunyuan":
|
||||
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"wanpipeline":
|
||||
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
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):
|
||||
config = verify_model_config_and_directory(pipeline_name_or_path)
|
||||
logger.warning(
|
||||
"FastVideo may not correctly identify the optimal config for this model, as the local directory may have been renamed."
|
||||
)
|
||||
else:
|
||||
config = maybe_download_model_index(pipeline_name_or_path)
|
||||
|
||||
pipeline_name = config["_class_name"]
|
||||
|
||||
# First try exact match for specific weights
|
||||
if pipeline_name_or_path in WEIGHT_CONFIG_REGISTRY:
|
||||
return WEIGHT_CONFIG_REGISTRY[pipeline_name_or_path]
|
||||
|
||||
# Try partial matches (for local paths that might include the weight ID)
|
||||
for registered_id, config_class in WEIGHT_CONFIG_REGISTRY.items():
|
||||
if registered_id in pipeline_name_or_path:
|
||||
return config_class
|
||||
|
||||
# If no match, try to use the fallback config
|
||||
fallback_config = None
|
||||
print(pipeline_name)
|
||||
# Try to determine pipeline architecture for fallback
|
||||
for pipeline_type, detector in PIPELINE_DETECTOR.items():
|
||||
if detector(pipeline_name.lower()):
|
||||
fallback_config = PIPELINE_FALLBACK_CONFIG.get(pipeline_type)
|
||||
break
|
||||
|
||||
logger.warning("No match found for pipeline %s, using fallback config %s.",
|
||||
pipeline_name_or_path, fallback_config)
|
||||
return fallback_config
|
||||
@@ -0,0 +1,58 @@
|
||||
from dataclasses import dataclass
|
||||
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
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 3.0
|
||||
neg_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
flow_shift: int = 3
|
||||
num_inference_steps: int = 50
|
||||
|
||||
# Text encoding stage
|
||||
text_len: int = 512
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precision: str = "fp32"
|
||||
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
@dataclass
|
||||
class WanI2V480PConfig(WanT2V480PConfig):
|
||||
"""Base configuration for Wan I2V 14B 480P pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 40
|
||||
|
||||
# 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"
|
||||
}
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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"""
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import List, cast
|
||||
|
||||
from fastvideo.v1.entrypoints.cli import utils
|
||||
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
@@ -73,16 +73,12 @@ 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,
|
||||
help="Port for the master process")
|
||||
|
||||
generate_parser = InferenceArgs.add_cli_args(generate_parser)
|
||||
generate_parser = FastVideoArgs.add_cli_args(generate_parser)
|
||||
|
||||
return cast(FlexibleArgumentParser, generate_parser)
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -4,23 +4,35 @@
|
||||
|
||||
import argparse
|
||||
import dataclasses
|
||||
from contextlib import contextmanager
|
||||
from typing import List, Optional
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
from fastvideo.v1.configs.models import VAEConfig
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class InferenceArgs:
|
||||
class FastVideoArgs:
|
||||
# Model and path configuration
|
||||
model_path: str
|
||||
|
||||
# Distributed executor backend
|
||||
distributed_executor_backend: str = "mp"
|
||||
|
||||
inference_mode: bool = True # if False == training mode
|
||||
|
||||
# HuggingFace specific parameters
|
||||
trust_remote_code: bool = False
|
||||
revision: Optional[str] = None
|
||||
|
||||
# Parallelism
|
||||
tp_size: int = 1
|
||||
sp_size: int = 1
|
||||
num_gpus: int = 1
|
||||
tp_size: Optional[int] = None
|
||||
sp_size: Optional[int] = None
|
||||
dist_timeout: Optional[int] = None # timeout for torch.distributed
|
||||
|
||||
# Video generation parameters
|
||||
@@ -40,9 +52,10 @@ class InferenceArgs:
|
||||
|
||||
# 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
|
||||
@@ -112,40 +125,55 @@ class InferenceArgs:
|
||||
help="Directory containing StepVideo model",
|
||||
)
|
||||
|
||||
# distributed_executor_backend
|
||||
parser.add_argument(
|
||||
"--distributed-executor-backend",
|
||||
type=str,
|
||||
choices=["mp"],
|
||||
default=FastVideoArgs.distributed_executor_backend,
|
||||
help="The distributed executor backend to use",
|
||||
)
|
||||
|
||||
# HuggingFace specific parameters
|
||||
parser.add_argument(
|
||||
"--trust-remote-code",
|
||||
action="store_true",
|
||||
default=InferenceArgs.trust_remote_code,
|
||||
default=FastVideoArgs.trust_remote_code,
|
||||
help="Trust remote code when loading HuggingFace models",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--revision",
|
||||
type=str,
|
||||
default=InferenceArgs.revision,
|
||||
default=FastVideoArgs.revision,
|
||||
help=
|
||||
"The specific model version to use (can be a branch name, tag name, or commit id)",
|
||||
)
|
||||
|
||||
# Parallelism
|
||||
parser.add_argument(
|
||||
"--num-gpus",
|
||||
type=int,
|
||||
default=FastVideoArgs.num_gpus,
|
||||
help="The number of GPUs to use.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tensor-parallel-size",
|
||||
"--tp-size",
|
||||
type=int,
|
||||
default=InferenceArgs.tp_size,
|
||||
default=FastVideoArgs.tp_size,
|
||||
help="The tensor parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sequence-parallel-size",
|
||||
"--sp-size",
|
||||
type=int,
|
||||
default=InferenceArgs.sp_size,
|
||||
default=FastVideoArgs.sp_size,
|
||||
help="The sequence parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dist-timeout",
|
||||
type=int,
|
||||
default=InferenceArgs.dist_timeout,
|
||||
default=FastVideoArgs.dist_timeout,
|
||||
help="Set timeout for torch.distributed initialization.",
|
||||
)
|
||||
|
||||
@@ -153,56 +181,56 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--height",
|
||||
type=int,
|
||||
default=InferenceArgs.height,
|
||||
default=FastVideoArgs.height,
|
||||
help="Height of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--width",
|
||||
type=int,
|
||||
default=InferenceArgs.width,
|
||||
default=FastVideoArgs.width,
|
||||
help="Width of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-frames",
|
||||
type=int,
|
||||
default=InferenceArgs.num_frames,
|
||||
default=FastVideoArgs.num_frames,
|
||||
help="Number of frames to generate",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-inference-steps",
|
||||
type=int,
|
||||
default=InferenceArgs.num_inference_steps,
|
||||
default=FastVideoArgs.num_inference_steps,
|
||||
help="Number of inference steps",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-scale",
|
||||
type=float,
|
||||
default=InferenceArgs.guidance_scale,
|
||||
default=FastVideoArgs.guidance_scale,
|
||||
help="Guidance scale for classifier-free guidance",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-rescale",
|
||||
type=float,
|
||||
default=InferenceArgs.guidance_rescale,
|
||||
default=FastVideoArgs.guidance_rescale,
|
||||
help="Guidance rescale for classifier-free guidance",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--embedded-cfg-scale",
|
||||
type=float,
|
||||
default=InferenceArgs.embedded_cfg_scale,
|
||||
default=FastVideoArgs.embedded_cfg_scale,
|
||||
help="Embedded CFG scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flow-shift",
|
||||
"--shift",
|
||||
type=float,
|
||||
default=InferenceArgs.flow_shift,
|
||||
default=FastVideoArgs.flow_shift,
|
||||
help="Flow shift parameter",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-type",
|
||||
type=str,
|
||||
default=InferenceArgs.output_type,
|
||||
default=FastVideoArgs.output_type,
|
||||
choices=["pil"],
|
||||
help="Output type for the generated video",
|
||||
)
|
||||
@@ -210,7 +238,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--precision",
|
||||
type=str,
|
||||
default=InferenceArgs.precision,
|
||||
default=FastVideoArgs.precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for the model",
|
||||
)
|
||||
@@ -219,14 +247,14 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--vae-precision",
|
||||
type=str,
|
||||
default=InferenceArgs.vae_precision,
|
||||
default=FastVideoArgs.vae_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for VAE",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-tiling",
|
||||
action="store_true",
|
||||
default=InferenceArgs.vae_tiling,
|
||||
default=FastVideoArgs.vae_tiling,
|
||||
help="Enable VAE tiling",
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -238,14 +266,14 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision",
|
||||
type=str,
|
||||
default=InferenceArgs.text_encoder_precision,
|
||||
default=FastVideoArgs.text_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-len",
|
||||
type=int,
|
||||
default=InferenceArgs.text_len,
|
||||
default=FastVideoArgs.text_len,
|
||||
help="Maximum text length",
|
||||
)
|
||||
|
||||
@@ -253,7 +281,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--image-encoder-precision",
|
||||
type=str,
|
||||
default=InferenceArgs.image_encoder_precision,
|
||||
default=FastVideoArgs.image_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for image encoder",
|
||||
)
|
||||
@@ -263,14 +291,14 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision-2",
|
||||
type=str,
|
||||
default=InferenceArgs.text_encoder_precision_2,
|
||||
default=FastVideoArgs.text_encoder_precision_2,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for secondary text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-len-2",
|
||||
type=int,
|
||||
default=InferenceArgs.text_len_2,
|
||||
default=FastVideoArgs.text_len_2,
|
||||
help="Maximum secondary text length",
|
||||
)
|
||||
|
||||
@@ -278,13 +306,13 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--flow-solver",
|
||||
type=str,
|
||||
default=InferenceArgs.flow_solver,
|
||||
default=FastVideoArgs.flow_solver,
|
||||
help="Solver for flow matching",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--denoise-type",
|
||||
type=str,
|
||||
default=InferenceArgs.denoise_type,
|
||||
default=FastVideoArgs.denoise_type,
|
||||
help="Denoise type for noised inputs",
|
||||
)
|
||||
|
||||
@@ -305,7 +333,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--scheduler-type",
|
||||
type=str,
|
||||
default=InferenceArgs.scheduler_type,
|
||||
default=FastVideoArgs.scheduler_type,
|
||||
help="Type of scheduler to use",
|
||||
)
|
||||
|
||||
@@ -313,19 +341,19 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--neg-prompt",
|
||||
type=str,
|
||||
default=InferenceArgs.neg_prompt,
|
||||
default=FastVideoArgs.neg_prompt,
|
||||
help="Negative prompt for sampling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-videos",
|
||||
type=int,
|
||||
default=InferenceArgs.num_videos,
|
||||
default=FastVideoArgs.num_videos,
|
||||
help="Number of videos to generate per prompt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fps",
|
||||
type=int,
|
||||
default=InferenceArgs.fps,
|
||||
default=FastVideoArgs.fps,
|
||||
help="Frames per second for output video",
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -344,7 +372,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--log-level",
|
||||
type=str,
|
||||
default=InferenceArgs.log_level,
|
||||
default=FastVideoArgs.log_level,
|
||||
help="The logging level of all loggers.",
|
||||
)
|
||||
|
||||
@@ -368,20 +396,20 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
type=str,
|
||||
default=InferenceArgs.output_path,
|
||||
default=FastVideoArgs.output_path,
|
||||
help="Directory to save generated videos",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed",
|
||||
type=int,
|
||||
default=InferenceArgs.seed,
|
||||
default=FastVideoArgs.seed,
|
||||
help="Random seed for reproducibility",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "InferenceArgs":
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs":
|
||||
args.tp_size = args.tensor_parallel_size
|
||||
args.sp_size = args.sequence_parallel_size
|
||||
args.flow_shift = getattr(args, "shift", args.flow_shift)
|
||||
@@ -406,8 +434,20 @@ class InferenceArgs:
|
||||
|
||||
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})"
|
||||
)
|
||||
|
||||
# Validate VAE spatial parallelism with VAE tiling
|
||||
if self.vae_sp and not self.vae_tiling:
|
||||
@@ -418,10 +458,10 @@ class InferenceArgs:
|
||||
raise ValueError("prompt_path must be a text file")
|
||||
|
||||
|
||||
_inference_args = None
|
||||
_current_fastvideo_args = None
|
||||
|
||||
|
||||
def prepare_inference_args(argv: List[str]) -> InferenceArgs:
|
||||
def prepare_fastvideo_args(argv: List[str]) -> FastVideoArgs:
|
||||
"""
|
||||
Prepare the inference arguments from the command line arguments.
|
||||
|
||||
@@ -433,26 +473,38 @@ def prepare_inference_args(argv: List[str]) -> InferenceArgs:
|
||||
The inference arguments.
|
||||
"""
|
||||
parser = FlexibleArgumentParser()
|
||||
InferenceArgs.add_cli_args(parser)
|
||||
FastVideoArgs.add_cli_args(parser)
|
||||
raw_args = parser.parse_args(argv)
|
||||
inference_args = InferenceArgs.from_cli_args(raw_args)
|
||||
inference_args.check_inference_args()
|
||||
global _inference_args
|
||||
_inference_args = inference_args
|
||||
return inference_args
|
||||
fastvideo_args = FastVideoArgs.from_cli_args(raw_args)
|
||||
fastvideo_args.check_fastvideo_args()
|
||||
global _current_fastvideo_args
|
||||
_current_fastvideo_args = fastvideo_args
|
||||
return fastvideo_args
|
||||
|
||||
|
||||
def get_inference_args() -> InferenceArgs:
|
||||
global _inference_args
|
||||
if _inference_args is None:
|
||||
raise ValueError("Inference arguments not set")
|
||||
return _inference_args
|
||||
@contextmanager
|
||||
def set_current_fastvideo_args(fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Temporarily set the current fastvideo config.
|
||||
Used during model initialization.
|
||||
We save the current fastvideo config in a global variable,
|
||||
so that all modules can access it, e.g. custom ops
|
||||
can access the fastvideo config to determine how to dispatch.
|
||||
"""
|
||||
global _current_fastvideo_args
|
||||
old_fastvideo_args = _current_fastvideo_args
|
||||
try:
|
||||
_current_fastvideo_args = fastvideo_args
|
||||
yield
|
||||
finally:
|
||||
_current_fastvideo_args = old_fastvideo_args
|
||||
|
||||
|
||||
class DeprecatedAction(argparse.Action):
|
||||
|
||||
def __init__(self, option_strings, dest, nargs=0, **kwargs):
|
||||
super().__init__(option_strings, dest, nargs=nargs, **kwargs)
|
||||
|
||||
def __call__(self, parser, namespace, values, option_string=None):
|
||||
raise ValueError(self.help)
|
||||
def get_current_fastvideo_args() -> FastVideoArgs:
|
||||
if _current_fastvideo_args is None:
|
||||
# in ci, usually when we test custom ops/modules directly,
|
||||
# we don't set the fastvideo config. In that case, we set a default
|
||||
# config.
|
||||
# TODO(will): may need to handle this for CI.
|
||||
raise ValueError("Current fastvideo args is not set.")
|
||||
return _current_fastvideo_args
|
||||
@@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -52,7 +52,7 @@ def get_forward_context() -> ForwardContext:
|
||||
@contextmanager
|
||||
def set_forward_context(current_timestep,
|
||||
attn_metadata,
|
||||
inference_args: Optional[InferenceArgs] = None):
|
||||
fastvideo_args: Optional[FastVideoArgs] = None):
|
||||
"""A context manager that stores the current forward context,
|
||||
can be attention metadata, etc.
|
||||
Here we can inject common logic for every model forward pass.
|
||||
|
||||
@@ -10,7 +10,7 @@ from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
|
||||
build_pipeline)
|
||||
@@ -28,29 +28,29 @@ class InferenceEngine:
|
||||
def __init__(
|
||||
self,
|
||||
pipeline: ComposedPipelineBase,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
):
|
||||
"""
|
||||
Initialize the inference engine.
|
||||
|
||||
Args:
|
||||
pipeline: The pipeline to use for inference.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
default_negative_prompt: The default negative prompt to use.
|
||||
"""
|
||||
self.pipeline = pipeline
|
||||
self.inference_args = inference_args
|
||||
self.fastvideo_args = fastvideo_args
|
||||
|
||||
@classmethod
|
||||
def create_engine(
|
||||
cls,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> "InferenceEngine":
|
||||
"""
|
||||
Create an inference engine with the specified arguments.
|
||||
|
||||
Args:
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
model_loader_cls: The model loader class to use. If None, it will be
|
||||
determined from the model type.
|
||||
pipeline_type: The type of pipeline to create. If None, it will be
|
||||
@@ -71,16 +71,16 @@ class InferenceEngine:
|
||||
# this way for training we can just do pipeline_cls.from_pretrained(
|
||||
# checkpoint_path) and have it handle everything.
|
||||
# TODO(Peiyuan): Then maybe we should only pass in model path and device, not the entire inference args?
|
||||
pipeline = build_pipeline(inference_args)
|
||||
pipeline = build_pipeline(fastvideo_args)
|
||||
logger.info("Pipeline Ready")
|
||||
|
||||
# Create the inference engine
|
||||
return cls(pipeline, inference_args)
|
||||
return cls(pipeline, fastvideo_args)
|
||||
|
||||
def run(
|
||||
self,
|
||||
prompt: str,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Run inference with the pipeline.
|
||||
@@ -96,17 +96,17 @@ class InferenceEngine:
|
||||
"""
|
||||
out_dict: Dict[str, Any] = dict()
|
||||
|
||||
num_videos_per_prompt = inference_args.num_videos
|
||||
seed = inference_args.seed
|
||||
height = inference_args.height
|
||||
width = inference_args.width
|
||||
video_length = inference_args.num_frames
|
||||
negative_prompt = inference_args.neg_prompt
|
||||
infer_steps = inference_args.num_inference_steps
|
||||
guidance_scale = inference_args.guidance_scale
|
||||
flow_shift = inference_args.flow_shift
|
||||
embedded_guidance_scale = inference_args.embedded_cfg_scale
|
||||
image_path = inference_args.image_path
|
||||
num_videos_per_prompt = fastvideo_args.num_videos
|
||||
seed = fastvideo_args.seed
|
||||
height = fastvideo_args.height
|
||||
width = fastvideo_args.width
|
||||
video_length = fastvideo_args.num_frames
|
||||
negative_prompt = fastvideo_args.neg_prompt
|
||||
infer_steps = fastvideo_args.num_inference_steps
|
||||
guidance_scale = fastvideo_args.guidance_scale
|
||||
flow_shift = fastvideo_args.flow_shift
|
||||
embedded_guidance_scale = fastvideo_args.embedded_cfg_scale
|
||||
image_path = fastvideo_args.image_path
|
||||
|
||||
# ========================================================================
|
||||
# Arguments: target_width, target_height, target_video_length
|
||||
@@ -161,21 +161,21 @@ class InferenceEngine:
|
||||
# return
|
||||
# sp_group = get_sp_group()
|
||||
# local_rank = sp_group.rank
|
||||
device = torch.device(inference_args.device_str)
|
||||
device = torch.device(fastvideo_args.device_str)
|
||||
batch = ForwardBatch(
|
||||
image_path=image_path,
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
height=inference_args.height,
|
||||
width=inference_args.width,
|
||||
num_frames=inference_args.num_frames,
|
||||
num_inference_steps=inference_args.num_inference_steps,
|
||||
guidance_scale=inference_args.guidance_scale,
|
||||
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,
|
||||
# generator=generator,
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
data_type="video" if inference_args.num_frames > 1 else "image",
|
||||
data_type="video" if fastvideo_args.num_frames > 1 else "image",
|
||||
device=device,
|
||||
extra={}, # Any additional parameters
|
||||
)
|
||||
@@ -184,7 +184,7 @@ class InferenceEngine:
|
||||
print(batch)
|
||||
print('===============================================')
|
||||
print('===============================================')
|
||||
print(inference_args)
|
||||
print(fastvideo_args)
|
||||
|
||||
# ========================================================================
|
||||
# Pipeline inference
|
||||
@@ -192,7 +192,7 @@ class InferenceEngine:
|
||||
start_time = time.time()
|
||||
samples = self.pipeline.forward(
|
||||
batch=batch,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
).output
|
||||
# TODO(will): fix and move to hunyuan stage
|
||||
# out_dict["seeds"] = batch.seeds
|
||||
|
||||
+86
-2
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -83,13 +83,13 @@ def get_hf_config(
|
||||
|
||||
def get_diffusers_config(
|
||||
model: str,
|
||||
inference_args: Optional[dict] = None,
|
||||
fastvideo_args: Optional[dict] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Gets a configuration for the given diffusers model.
|
||||
|
||||
Args:
|
||||
model: The model name or path.
|
||||
inference_args: Optional inference arguments to override in the config.
|
||||
fastvideo_args: Optional inference arguments to override in the config.
|
||||
|
||||
Returns:
|
||||
The loaded configuration.
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import dataclasses
|
||||
from dataclasses import asdict
|
||||
import glob
|
||||
import os
|
||||
import time
|
||||
@@ -13,7 +14,7 @@ from safetensors.torch import load_file as safetensors_load_file
|
||||
from transformers import AutoImageProcessor, AutoTokenizer, PretrainedConfig
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.hf_transformer_utils import (get_diffusers_config,
|
||||
get_hf_config)
|
||||
@@ -36,14 +37,14 @@ class ComponentLoader(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Load the component based on the model path, architecture, and inference args.
|
||||
|
||||
Args:
|
||||
model_path: Path to the component model
|
||||
architecture: Architecture of the component model
|
||||
inference_args: Inference arguments
|
||||
fastvideo_args: Inference arguments
|
||||
|
||||
Returns:
|
||||
The loaded component
|
||||
@@ -199,20 +200,20 @@ class TextEncoderLoader(ComponentLoader):
|
||||
yield from self._get_weights_iterator(source)
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the text encoders based on the model path, architecture, and inference args."""
|
||||
model_config: PretrainedConfig = get_hf_config(
|
||||
model=model_path,
|
||||
trust_remote_code=inference_args.trust_remote_code,
|
||||
revision=inference_args.revision,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
model_override_args=None,
|
||||
)
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
target_device = torch.device(inference_args.device_str)
|
||||
target_device = torch.device(fastvideo_args.device_str)
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, model_config, target_device,
|
||||
inference_args.text_encoder_precision)
|
||||
fastvideo_args.text_encoder_precision)
|
||||
|
||||
def load_model(self,
|
||||
model_path: str,
|
||||
@@ -249,27 +250,27 @@ class TextEncoderLoader(ComponentLoader):
|
||||
class ImageEncoderLoader(TextEncoderLoader):
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the text encoders based on the model path, architecture, and inference args."""
|
||||
model_config: PretrainedConfig = get_hf_config(
|
||||
model=model_path,
|
||||
trust_remote_code=inference_args.trust_remote_code,
|
||||
revision=inference_args.revision,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
model_override_args=None,
|
||||
)
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
target_device = torch.device(inference_args.device_str)
|
||||
target_device = torch.device(fastvideo_args.device_str)
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, model_config, target_device,
|
||||
inference_args.image_encoder_precision)
|
||||
fastvideo_args.image_encoder_precision)
|
||||
|
||||
|
||||
class ImageProcessorLoader(ComponentLoader):
|
||||
"""Loader for image processor."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the image processor based on the model path, architecture, and inference args."""
|
||||
logger.info("Loading image processor from %s", model_path)
|
||||
|
||||
@@ -283,7 +284,7 @@ class TokenizerLoader(ComponentLoader):
|
||||
"""Loader for tokenizers."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the tokenizer based on the model path, architecture, and inference args."""
|
||||
logger.info("Loading tokenizer from %s", model_path)
|
||||
|
||||
@@ -301,7 +302,7 @@ class VAELoader(ComponentLoader):
|
||||
"""Loader for VAE."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the VAE based on the model path, architecture, and inference args."""
|
||||
# TODO(will): move this to a constants file
|
||||
config = get_diffusers_config(model=model_path)
|
||||
@@ -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(inference_args.device)
|
||||
vae = vae_cls(vae_config).to(fastvideo_args.device)
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
@@ -322,8 +325,8 @@ 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)
|
||||
dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
|
||||
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)
|
||||
|
||||
return vae
|
||||
@@ -333,7 +336,7 @@ class TransformerLoader(ComponentLoader):
|
||||
"""Loader for transformer."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the transformer based on the model path, architecture, and inference args."""
|
||||
model_config = get_diffusers_config(model=model_path)
|
||||
cls_name = model_config.pop("_class_name")
|
||||
@@ -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
|
||||
@@ -354,16 +361,16 @@ class TransformerLoader(ComponentLoader):
|
||||
logger.info("Loading model from %s safetensors files in %s",
|
||||
len(safetensors_list), model_path)
|
||||
|
||||
# initialize_sequence_parallel_group(inference_args.sp_size)
|
||||
default_dtype = PRECISION_TO_TYPE[inference_args.precision]
|
||||
# initialize_sequence_parallel_group(fastvideo_args.sp_size)
|
||||
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s", cls_name)
|
||||
model = load_fsdp_model(model_cls=model_cls,
|
||||
init_params=model_config,
|
||||
weight_dir_list=safetensors_list,
|
||||
device=inference_args.device,
|
||||
cpu_offload=inference_args.use_cpu_offload,
|
||||
device=fastvideo_args.device,
|
||||
cpu_offload=fastvideo_args.use_cpu_offload,
|
||||
default_dtype=default_dtype)
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
@@ -380,7 +387,7 @@ class SchedulerLoader(ComponentLoader):
|
||||
"""Loader for scheduler."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the scheduler based on the model path, architecture, and inference args."""
|
||||
config = get_diffusers_config(model=model_path)
|
||||
|
||||
@@ -391,8 +398,8 @@ class SchedulerLoader(ComponentLoader):
|
||||
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
scheduler = scheduler_cls(**config)
|
||||
if inference_args.flow_shift is not None:
|
||||
scheduler.set_shift(inference_args.flow_shift)
|
||||
if fastvideo_args.flow_shift is not None:
|
||||
scheduler.set_shift(fastvideo_args.flow_shift)
|
||||
|
||||
return scheduler
|
||||
|
||||
@@ -405,7 +412,7 @@ class GenericComponentLoader(ComponentLoader):
|
||||
self.library = library
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load a generic component based on the model path, architecture, and inference args."""
|
||||
logger.warning("Using generic loader for %s with library %s",
|
||||
model_path, self.library)
|
||||
@@ -415,8 +422,8 @@ class GenericComponentLoader(ComponentLoader):
|
||||
|
||||
model = AutoModel.from_pretrained(
|
||||
model_path,
|
||||
trust_remote_code=inference_args.trust_remote_code,
|
||||
revision=inference_args.revision,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
)
|
||||
logger.info("Loaded generic transformers model: %s",
|
||||
model.__class__.__name__)
|
||||
@@ -443,7 +450,7 @@ class PipelineComponentLoader:
|
||||
@staticmethod
|
||||
def load_module(module_name: str, component_model_path: str,
|
||||
transformers_or_diffusers: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Load a pipeline module.
|
||||
|
||||
@@ -452,7 +459,7 @@ class PipelineComponentLoader:
|
||||
component_model_path: Path to the component model
|
||||
transformers_or_diffusers: Whether the module is from transformers or diffusers
|
||||
architecture: Architecture of the component model
|
||||
inference_args: Inference arguments
|
||||
fastvideo_args: Inference arguments
|
||||
|
||||
Returns:
|
||||
The loaded module
|
||||
@@ -469,4 +476,4 @@ class PipelineComponentLoader:
|
||||
transformers_or_diffusers)
|
||||
|
||||
# Load the module
|
||||
return loader.load(component_model_path, architecture, inference_args)
|
||||
return loader.load(component_model_path, architecture, fastvideo_args)
|
||||
|
||||
@@ -8,7 +8,7 @@ from diffusers.utils import BaseOutput
|
||||
|
||||
|
||||
class BaseScheduler(ABC):
|
||||
timesteps: torch.tensor
|
||||
timesteps: torch.Tensor
|
||||
order: int
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
@@ -38,9 +38,9 @@ class BaseScheduler(ABC):
|
||||
@abstractmethod
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
model_output: torch.Tensor,
|
||||
timestep: Union[int, torch.Tensor],
|
||||
sample: torch.Tensor,
|
||||
return_dict: bool = True,
|
||||
) -> Union[BaseOutput, Tuple]:
|
||||
pass
|
||||
|
||||
@@ -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"""
|
||||
|
||||
@@ -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)
|
||||
|
||||
+315
-114
@@ -14,19 +14,39 @@
|
||||
# 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
|
||||
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.models.vaes.common import ParallelTiledVAE, DiagonalGaussianDistribution
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig, WanVAEArchConfig
|
||||
|
||||
CACHE_T = 2
|
||||
|
||||
is_first_frame = contextvars.ContextVar("is_first_frame", default=False)
|
||||
feat_cache = contextvars.ContextVar("feat_cache", default=None)
|
||||
feat_idx = contextvars.ContextVar("feat_idx", default=0)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def forward_context(first_frame_arg=False,
|
||||
feat_cache_arg=None,
|
||||
feat_idx_arg=None):
|
||||
is_first_frame_token = is_first_frame.set(first_frame_arg)
|
||||
feat_cache_token = feat_cache.set(feat_cache_arg)
|
||||
feat_idx_token = feat_idx.set(feat_idx_arg)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
is_first_frame.reset(is_first_frame_token)
|
||||
feat_cache.reset(feat_cache_token)
|
||||
feat_idx.reset(feat_idx_token)
|
||||
|
||||
|
||||
class WanCausalConv3d(nn.Conv3d):
|
||||
r"""
|
||||
@@ -60,12 +80,17 @@ class WanCausalConv3d(nn.Conv3d):
|
||||
)
|
||||
self.padding: Tuple[int, int, int]
|
||||
# Set up causal padding
|
||||
self._padding = (self.padding[2], self.padding[2], self.padding[1],
|
||||
self.padding[1], 2 * self.padding[0], 0)
|
||||
self._padding: Tuple[int, ...] = (self.padding[2], self.padding[2],
|
||||
self.padding[1], self.padding[1],
|
||||
2 * self.padding[0], 0)
|
||||
self.padding = (0, 0, 0)
|
||||
|
||||
def forward(self, x):
|
||||
def forward(self, x, cache_x=None):
|
||||
padding = list(self._padding)
|
||||
if cache_x is not None and self._padding[4] > 0:
|
||||
cache_x = cache_x.to(x.device)
|
||||
x = torch.cat([cache_x, x], dim=2)
|
||||
padding[4] -= cache_x.shape[2]
|
||||
x = F.pad(x, padding)
|
||||
return super().forward(x)
|
||||
|
||||
@@ -157,28 +182,82 @@ class WanResample(nn.Module):
|
||||
self.time_conv = WanCausalConv3d(dim,
|
||||
dim, (3, 1, 1),
|
||||
stride=(2, 1, 1),
|
||||
padding=(1, 0, 0))
|
||||
padding=(0, 0, 0))
|
||||
|
||||
else:
|
||||
self.resample = nn.Identity()
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
def forward(self, x):
|
||||
b, c, t, h, w = x.size()
|
||||
first_frame = is_first_frame.get()
|
||||
if first_frame:
|
||||
assert t == 1
|
||||
if self.mode == "upsample3d" and not first_frame and hasattr(
|
||||
self, "time_conv"):
|
||||
x = self.time_conv(x)
|
||||
x = x.reshape(b, 2, c, t, h, w)
|
||||
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3)
|
||||
x = x.reshape(b, c, t * 2, h, w)
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if self.mode == "upsample3d":
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
if _feat_cache[idx] is None:
|
||||
_feat_cache[idx] = "Rep"
|
||||
_feat_idx += 1
|
||||
else:
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[
|
||||
idx] is not None and _feat_cache[idx] != "Rep":
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
if cache_x.shape[2] < 2 and _feat_cache[
|
||||
idx] is not None and _feat_cache[idx] == "Rep":
|
||||
cache_x = torch.cat([
|
||||
torch.zeros_like(cache_x).to(cache_x.device),
|
||||
cache_x
|
||||
],
|
||||
dim=2)
|
||||
if _feat_cache[idx] == "Rep":
|
||||
x = self.time_conv(x)
|
||||
else:
|
||||
x = self.time_conv(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
|
||||
x = x.reshape(b, 2, c, t, h, w)
|
||||
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]),
|
||||
3)
|
||||
x = x.reshape(b, c, t * 2, h, w)
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
elif not first_frame and hasattr(self, "time_conv"):
|
||||
x = self.time_conv(x)
|
||||
x = x.reshape(b, 2, c, t, h, w)
|
||||
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3)
|
||||
x = x.reshape(b, c, t * 2, h, w)
|
||||
t = x.shape[2]
|
||||
x = x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)
|
||||
x = self.resample(x)
|
||||
x = x.view(b, t, x.size(1), x.size(2), x.size(3)).permute(0, 2, 1, 3, 4)
|
||||
if self.mode == "downsample3d" and not first_frame and hasattr(
|
||||
self, "time_conv"):
|
||||
x = self.time_conv(x)
|
||||
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if self.mode == "downsample3d":
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
if _feat_cache[idx] is None:
|
||||
_feat_cache[idx] = x.clone()
|
||||
_feat_idx += 1
|
||||
else:
|
||||
cache_x = x[:, :, -1:, :, :].clone()
|
||||
x = self.time_conv(
|
||||
torch.cat([_feat_cache[idx][:, :, -1:, :, :], x], 2))
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
elif not first_frame and hasattr(self, "time_conv"):
|
||||
x = self.time_conv(x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -222,7 +301,25 @@ class WanResidualBlock(nn.Module):
|
||||
x = self.norm1(x)
|
||||
x = self.nonlinearity(x)
|
||||
|
||||
x = self.conv1(x)
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
|
||||
x = self.conv1(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv1(x)
|
||||
|
||||
# Second normalization and activation
|
||||
x = self.norm2(x)
|
||||
@@ -231,7 +328,25 @@ class WanResidualBlock(nn.Module):
|
||||
# Dropout
|
||||
x = self.dropout(x)
|
||||
|
||||
x = self.conv2(x)
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
|
||||
x = self.conv2(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv2(x)
|
||||
|
||||
# Add residual connection
|
||||
return x + h
|
||||
@@ -400,15 +515,30 @@ class WanEncoder3d(nn.Module):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
x = self.conv_in(x)
|
||||
def forward(self, x):
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_in(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv_in(x)
|
||||
|
||||
## downsamples
|
||||
for layer in self.down_blocks:
|
||||
if isinstance(layer, WanResample):
|
||||
x = layer(x, first_frame=first_frame)
|
||||
else:
|
||||
x = layer(x)
|
||||
x = layer(x)
|
||||
|
||||
## middle
|
||||
x = self.mid_block(x)
|
||||
@@ -416,7 +546,26 @@ class WanEncoder3d(nn.Module):
|
||||
## head
|
||||
x = self.norm_out(x)
|
||||
x = self.nonlinearity(x)
|
||||
x = self.conv_out(x)
|
||||
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_out(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -465,7 +614,7 @@ class WanUpBlock(nn.Module):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
def forward(self, x):
|
||||
"""
|
||||
Forward pass through the upsampling block.
|
||||
|
||||
@@ -481,7 +630,7 @@ class WanUpBlock(nn.Module):
|
||||
x = resnet(x)
|
||||
|
||||
if self.upsamplers is not None:
|
||||
x = self.upsamplers[0](x, first_frame=first_frame)
|
||||
x = self.upsamplers[0](x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -569,21 +718,57 @@ class WanDecoder3d(nn.Module):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
def forward(self, x):
|
||||
## conv1
|
||||
x = self.conv_in(x)
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_in(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv_in(x)
|
||||
|
||||
## middle
|
||||
x = self.mid_block(x)
|
||||
|
||||
## upsamples
|
||||
for up_block in self.up_blocks:
|
||||
x = up_block(x, first_frame=first_frame)
|
||||
x = up_block(x)
|
||||
|
||||
## head
|
||||
x = self.norm_out(x)
|
||||
x = self.nonlinearity(x)
|
||||
x = self.conv_out(x)
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_out(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -595,96 +780,89 @@ 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
|
||||
self.use_feature_cache = config.use_feature_cache
|
||||
|
||||
# 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
|
||||
def clear_cache(self) -> None:
|
||||
|
||||
# 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 _count_conv3d(model) -> int:
|
||||
count = 0
|
||||
for m in model.modules():
|
||||
if isinstance(m, WanCausalConv3d):
|
||||
count += 1
|
||||
return count
|
||||
|
||||
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
|
||||
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:
|
||||
self.clear_cache()
|
||||
with forward_context(feat_cache_arg=self._enc_feat_map,
|
||||
feat_idx_arg=self._enc_conv_idx):
|
||||
t = x.shape[2]
|
||||
iter_ = 1 + (t - 1) // 4
|
||||
for i in range(iter_):
|
||||
feat_idx.set(0)
|
||||
if i == 0:
|
||||
out = self.encoder(x[:, :, :1, :, :])
|
||||
else:
|
||||
out_ = self.encoder(x[:, :,
|
||||
1 + 4 * (i - 1):1 + 4 * i, :, :])
|
||||
out = torch.cat([out, out_], 2)
|
||||
enc = self.quant_conv(out)
|
||||
mu, logvar = enc[:, :self.z_dim, :, :, :], enc[:,
|
||||
self.z_dim:, :, :, :]
|
||||
enc = torch.cat([mu, logvar], dim=1)
|
||||
enc = DiagonalGaussianDistribution(enc)
|
||||
self.clear_cache()
|
||||
else:
|
||||
for block in self.encoder.down_blocks:
|
||||
if isinstance(block,
|
||||
WanResample) and block.mode == "downsample3d":
|
||||
_padding = list(block.time_conv._padding)
|
||||
_padding[4] = 2
|
||||
block.time_conv._padding = tuple(_padding)
|
||||
enc = ParallelTiledVAE.encode(self, x)
|
||||
|
||||
return enc
|
||||
|
||||
def _encode(self, x: torch.Tensor, first_frame=False) -> torch.Tensor:
|
||||
out = self.encoder(x, first_frame=first_frame)
|
||||
with forward_context(first_frame_arg=first_frame):
|
||||
out = self.encoder(x)
|
||||
enc = self.quant_conv(out)
|
||||
mu, logvar = enc[:, :self.z_dim, :, :, :], enc[:, self.z_dim:, :, :, :]
|
||||
enc = torch.cat([mu, logvar], dim=1)
|
||||
@@ -708,9 +886,32 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
enc = torch.cat([first_frame, enc], dim=2)
|
||||
return enc
|
||||
|
||||
def decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
if self.use_feature_cache:
|
||||
self.clear_cache()
|
||||
iter_ = z.shape[2]
|
||||
x = self.post_quant_conv(z)
|
||||
with forward_context(feat_cache_arg=self._feat_map,
|
||||
feat_idx_arg=self._conv_idx):
|
||||
for i in range(iter_):
|
||||
feat_idx.set(0)
|
||||
if i == 0:
|
||||
out = self.decoder(x[:, :, i:i + 1, :, :])
|
||||
else:
|
||||
out_ = self.decoder(x[:, :, i:i + 1, :, :])
|
||||
out = torch.cat([out, out_], 2)
|
||||
|
||||
out = torch.clamp(out, min=-1.0, max=1.0)
|
||||
self.clear_cache()
|
||||
else:
|
||||
out = ParallelTiledVAE.decode(self, z)
|
||||
|
||||
return out
|
||||
|
||||
def _decode(self, z: torch.Tensor, first_frame=False) -> torch.Tensor:
|
||||
x = self.post_quant_conv(z)
|
||||
out = self.decoder(x, first_frame=first_frame)
|
||||
with forward_context(first_frame_arg=first_frame):
|
||||
out = self.decoder(x)
|
||||
|
||||
out = torch.clamp(out, min=-1.0, max=1.0)
|
||||
|
||||
|
||||
@@ -40,7 +40,7 @@ from fastvideo.v1.pipelines.stages import (
|
||||
ConditioningStage,
|
||||
# Import other required stages
|
||||
)
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
class YourCustomPipeline(ComposedPipelineBase):
|
||||
@@ -53,7 +53,7 @@ class YourCustomPipeline(ComposedPipelineBase):
|
||||
# Add other required modules
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
# Add and configure pipeline stages
|
||||
self.add_stage(
|
||||
stage_name="input_validation_stage",
|
||||
@@ -61,14 +61,14 @@ class YourCustomPipeline(ComposedPipelineBase):
|
||||
)
|
||||
# Add more stages as needed
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# Initialize pipeline-specific components
|
||||
pass
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, inference_args: InferenceArgs) -> ForwardBatch:
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
# Implement your pipeline's forward pass
|
||||
batch = self.input_validation_stage(batch, inference_args)
|
||||
batch = self.input_validation_stage(batch, fastvideo_args)
|
||||
# Add more stage executions
|
||||
return batch
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ Diffusion pipelines for fastvideo.v1.
|
||||
This package contains diffusion pipelines for generating videos and images.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -16,7 +16,7 @@ from fastvideo.v1.utils import (maybe_download_model,
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
|
||||
def build_pipeline(fastvideo_args: FastVideoArgs) -> ComposedPipelineBase:
|
||||
"""
|
||||
Only works with valid hf diffusers configs. (model_index.json)
|
||||
We want to build a pipeline based on the inference args mode_path:
|
||||
@@ -25,9 +25,9 @@ def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
|
||||
3. based on the config, determine the pipeline class
|
||||
"""
|
||||
# Get pipeline type
|
||||
model_path = inference_args.model_path
|
||||
model_path = fastvideo_args.model_path
|
||||
model_path = maybe_download_model(model_path)
|
||||
# inference_args.downloaded_model_path = model_path
|
||||
# fastvideo_args.downloaded_model_path = model_path
|
||||
logger.info("Model path: %s", model_path)
|
||||
config = verify_model_config_and_directory(model_path)
|
||||
|
||||
@@ -41,7 +41,7 @@ def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
|
||||
pipeline_architecture)
|
||||
|
||||
# instantiate the pipeline
|
||||
pipeline = pipeline_cls(model_path, inference_args, config)
|
||||
pipeline = pipeline_cls(model_path, fastvideo_args, config)
|
||||
logger.info("Pipeline instantiated")
|
||||
|
||||
# pipeline is now initialized and ready to use
|
||||
|
||||
@@ -12,7 +12,7 @@ from typing import Any, Dict, List, Optional, cast
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -38,7 +38,7 @@ class ComposedPipelineBase(ABC):
|
||||
# TODO(will): args should support both inference args and training args
|
||||
def __init__(self,
|
||||
model_path: str,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
config: Optional[Dict[str, Any]] = None):
|
||||
"""
|
||||
Initialize the pipeline. After __init__, the pipeline should be ready to
|
||||
@@ -61,12 +61,12 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
# Load modules directly in initialization
|
||||
logger.info("Loading pipeline modules...")
|
||||
self.modules = self.load_modules(inference_args)
|
||||
self.modules = self.load_modules(fastvideo_args)
|
||||
|
||||
self.initialize_pipeline(inference_args)
|
||||
self.initialize_pipeline(fastvideo_args)
|
||||
|
||||
logger.info("Creating pipeline stages...")
|
||||
self.create_pipeline_stages(inference_args)
|
||||
self.create_pipeline_stages(fastvideo_args)
|
||||
|
||||
def get_module(self, module_name: str) -> Any:
|
||||
return self.modules[module_name]
|
||||
@@ -77,7 +77,7 @@ class ComposedPipelineBase(ABC):
|
||||
def _load_config(self, model_path: str) -> Dict[str, Any]:
|
||||
model_path = maybe_download_model(self.model_path)
|
||||
self.model_path = model_path
|
||||
# inference_args.downloaded_model_path = model_path
|
||||
# fastvideo_args.downloaded_model_path = model_path
|
||||
logger.info("Model path: %s", model_path)
|
||||
config = verify_model_config_and_directory(model_path)
|
||||
return cast(Dict[str, Any], config)
|
||||
@@ -108,20 +108,20 @@ class ComposedPipelineBase(ABC):
|
||||
return self._stages
|
||||
|
||||
@abstractmethod
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Create the pipeline stages.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def load_modules(self, inference_args: InferenceArgs) -> Dict[str, Any]:
|
||||
def load_modules(self, fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
"""
|
||||
Load the modules from the config.
|
||||
"""
|
||||
@@ -156,7 +156,7 @@ class ComposedPipelineBase(ABC):
|
||||
component_model_path=component_model_path,
|
||||
transformers_or_diffusers=transformers_or_diffusers,
|
||||
architecture=architecture,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
logger.info("Loaded module %s from %s", module_name,
|
||||
component_model_path)
|
||||
@@ -185,14 +185,14 @@ class ComposedPipelineBase(ABC):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Generate a video or image using the pipeline.
|
||||
|
||||
Args:
|
||||
batch: The batch to generate from.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
Returns:
|
||||
ForwardBatch: The batch with the generated video or image.
|
||||
"""
|
||||
@@ -201,7 +201,7 @@ class ComposedPipelineBase(ABC):
|
||||
self._stage_name_mapping.keys())
|
||||
logger.info("Batch: %s", batch)
|
||||
for stage in self.stages:
|
||||
batch = stage(batch, inference_args)
|
||||
batch = stage(batch, fastvideo_args)
|
||||
|
||||
# Return the output
|
||||
return batch
|
||||
|
||||
@@ -6,9 +6,7 @@ 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.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (CLIPTextEncodingStage,
|
||||
@@ -30,7 +28,7 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
|
||||
"transformer", "scheduler"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
@@ -67,20 +65,12 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = 2**(len(self.get_module("vae").block_out_channels) -
|
||||
1)
|
||||
inference_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
|
||||
inference_args.num_channels_latents = num_channels_latents
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = HunyuanVideoPipeline
|
||||
|
||||
@@ -22,7 +22,7 @@ class ForwardBatch:
|
||||
execution, allowing methods to update specific components without needing
|
||||
to manage numerous individual parameters.
|
||||
"""
|
||||
# TODO(will): double check that args are separate from inference_args
|
||||
# TODO(will): double check that args are separate from fastvideo_args
|
||||
# properly. Also maybe think about providing an abstraction for pipeline
|
||||
# specific arguments.
|
||||
data_type: str
|
||||
|
||||
@@ -12,7 +12,7 @@ from abc import ABC, abstractmethod
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
@@ -45,7 +45,7 @@ class PipelineStage(ABC):
|
||||
def __call__(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Execute the stage's processing on the batch with optional logging.
|
||||
@@ -53,7 +53,7 @@ class PipelineStage(ABC):
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The updated batch information after this stage's processing.
|
||||
@@ -65,7 +65,7 @@ class PipelineStage(ABC):
|
||||
|
||||
try:
|
||||
# Call the actual implementation
|
||||
result = self._call_implementation(batch, inference_args)
|
||||
result = self._call_implementation(batch, fastvideo_args)
|
||||
|
||||
execution_time = time.time() - start_time
|
||||
self._logger.info("[%s] Execution completed in %s ms",
|
||||
@@ -85,13 +85,13 @@ class PipelineStage(ABC):
|
||||
else:
|
||||
# Just call the implementation directly if logging is disabled
|
||||
# TODO(will): Also handle backward
|
||||
return self.forward(batch, inference_args)
|
||||
return self.forward(batch, fastvideo_args)
|
||||
|
||||
@abstractmethod
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Forward pass of the stage's processing.
|
||||
@@ -101,7 +101,7 @@ class PipelineStage(ABC):
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The updated batch information after this stage's processing.
|
||||
@@ -111,6 +111,6 @@ class PipelineStage(ABC):
|
||||
def backward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -7,8 +7,8 @@ This module contains implementations of image encoding stages for diffusion pipe
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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
|
||||
@@ -40,19 +40,19 @@ class CLIPImageEncodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into image encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder = self.image_encoder.to(batch.device)
|
||||
|
||||
image = load_image(batch.image_path)
|
||||
@@ -64,7 +64,7 @@ class CLIPImageEncodingStage(PipelineStage):
|
||||
|
||||
batch.image_embeds.append(image_embeds)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
@@ -7,8 +7,8 @@ This module contains implementations of prompt encoding stages for diffusion pip
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -39,19 +39,19 @@ class CLIPTextEncodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
@@ -85,7 +85,7 @@ class CLIPTextEncodingStage(PipelineStage):
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
batch.negative_prompt_embeds.append(negative_prompt_embeds)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ Conditioning stage for diffusion pipelines.
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.pipelines.stages.base import PipelineStage
|
||||
@@ -24,14 +24,14 @@ class ConditioningStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Apply conditioning to the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with applied conditioning.
|
||||
|
||||
@@ -5,11 +5,12 @@ Decoding stage for diffusion pipelines.
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.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,20 +23,20 @@ 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,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Decode latent representations into pixel space.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with decoded outputs.
|
||||
@@ -46,13 +47,13 @@ class DecodingStage(PipelineStage):
|
||||
raise ValueError("Latents must be provided")
|
||||
|
||||
# Skip decoding if output type is latent
|
||||
if inference_args.output_type == "latent":
|
||||
if fastvideo_args.output_type == "latent":
|
||||
image = latents
|
||||
else:
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32
|
||||
) and not inference_args.disable_autocast
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
@@ -73,9 +74,9 @@ class DecodingStage(PipelineStage):
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if inference_args.vae_tiling:
|
||||
if fastvideo_args.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if inference_args.vae_sp:
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
|
||||
@@ -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.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -51,20 +52,20 @@ class DenoisingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Run the denoising loop.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
# If use cpu offload, need to load the model back into gpu again
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.transformer = self.transformer.to(batch.device)
|
||||
# Prepare extra step kwargs for scheduler
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
@@ -76,9 +77,9 @@ class DenoisingStage(PipelineStage):
|
||||
)
|
||||
|
||||
# Setup precision and autocast settings
|
||||
target_dtype = PRECISION_TO_TYPE[inference_args.precision]
|
||||
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not inference_args.disable_autocast
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Handle sequence parallelism if enabled
|
||||
world_size, rank = get_sequence_model_parallel_world_size(
|
||||
@@ -161,11 +162,11 @@ class DenoisingStage(PipelineStage):
|
||||
# Prepare inputs for transformer
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (torch.tensor(
|
||||
[inference_args.embedded_cfg_scale] *
|
||||
[fastvideo_args.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=batch.device,
|
||||
).to(target_dtype) * 1000.0 if inference_args.embedded_cfg_scale
|
||||
).to(target_dtype) * 1000.0 if fastvideo_args.embedded_cfg_scale
|
||||
is not None else None)
|
||||
|
||||
# Predict noise residual
|
||||
@@ -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(
|
||||
@@ -193,7 +193,7 @@ class DenoisingStage(PipelineStage):
|
||||
attn_metadata = self.attn_metadata_builder.build(
|
||||
current_timestep=i,
|
||||
forward_batch=batch,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
@@ -203,11 +203,11 @@ class DenoisingStage(PipelineStage):
|
||||
# TODO(will): finalize the interface. vLLM uses this to
|
||||
# support torch dynamo compilation. They pass in
|
||||
# attn_metadata, vllm_config, and num_tokens. We can pass in
|
||||
# inference_args or training_args, and attn_metadata.
|
||||
# fastvideo_args or training_args, and attn_metadata.
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
# inference_args=inference_args
|
||||
# fastvideo_args=fastvideo_args
|
||||
):
|
||||
# Run transformer
|
||||
noise_pred = self.transformer(
|
||||
@@ -223,7 +223,7 @@ class DenoisingStage(PipelineStage):
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
# inference_args=inference_args
|
||||
# fastvideo_args=fastvideo_args
|
||||
):
|
||||
# Run transformer
|
||||
noise_pred_uncond = self.transformer(
|
||||
@@ -267,7 +267,7 @@ class DenoisingStage(PipelineStage):
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.transformer.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing import Optional
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
load_image, normalize,
|
||||
@@ -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,20 +28,20 @@ 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,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode pixel representations into latent space.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded outputs.
|
||||
@@ -62,7 +63,7 @@ class EncodingStage(PipelineStage):
|
||||
video_condition = torch.cat([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1],
|
||||
inference_args.num_frames - 1, batch.height,
|
||||
fastvideo_args.num_frames - 1, batch.height,
|
||||
batch.width)
|
||||
],
|
||||
dim=2)
|
||||
@@ -70,17 +71,17 @@ class EncodingStage(PipelineStage):
|
||||
dtype=torch.float32)
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not inference_args.disable_autocast
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Encode Image
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if inference_args.vae_tiling:
|
||||
if fastvideo_args.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if inference_args.vae_sp:
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
video_condition = video_condition.to(vae_dtype)
|
||||
@@ -106,9 +107,9 @@ class EncodingStage(PipelineStage):
|
||||
else:
|
||||
latent_condition = latent_condition * self.vae.scaling_factor
|
||||
|
||||
mask_lat_size = torch.ones(1, 1, inference_args.num_frames,
|
||||
mask_lat_size = torch.ones(1, 1, fastvideo_args.num_frames,
|
||||
latent_height, latent_width)
|
||||
mask_lat_size[:, :, list(range(1, inference_args.num_frames))] = 0
|
||||
mask_lat_size[:, :, list(range(1, fastvideo_args.num_frames))] = 0
|
||||
first_frame_mask = mask_lat_size[:, :, 0:1]
|
||||
first_frame_mask = torch.repeat_interleave(
|
||||
first_frame_mask,
|
||||
|
||||
@@ -5,7 +5,7 @@ Input validation stage for diffusion pipelines.
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.pipelines.stages.base import PipelineStage
|
||||
@@ -22,10 +22,10 @@ class InputValidationStage(PipelineStage):
|
||||
"""
|
||||
|
||||
def _generate_seeds(self, batch: ForwardBatch,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Generate seeds for the inference"""
|
||||
seed = inference_args.seed
|
||||
num_videos_per_prompt = inference_args.num_videos
|
||||
seed = fastvideo_args.seed
|
||||
num_videos_per_prompt = fastvideo_args.num_videos
|
||||
|
||||
seeds = [seed + i for i in range(num_videos_per_prompt)]
|
||||
batch.seeds = seeds
|
||||
@@ -37,19 +37,19 @@ class InputValidationStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Validate and prepare inputs.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The validated batch information.
|
||||
"""
|
||||
self._generate_seeds(batch, inference_args)
|
||||
self._generate_seeds(batch, fastvideo_args)
|
||||
|
||||
# Ensure prompt is properly formatted
|
||||
if batch.prompt is None and batch.prompt_embeds is None:
|
||||
@@ -91,6 +91,6 @@ class InputValidationStage(PipelineStage):
|
||||
|
||||
# Set data type if not already set
|
||||
if batch.data_type is None:
|
||||
batch.data_type = inference_args.precision
|
||||
batch.data_type = fastvideo_args.precision
|
||||
|
||||
return batch
|
||||
|
||||
@@ -4,9 +4,8 @@ Latent preparation stage for diffusion pipelines.
|
||||
"""
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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,22 +20,21 @@ 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,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Prepare initial latent variables for the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with prepared latent variables.
|
||||
@@ -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, inference_args)
|
||||
batch = self.adjust_video_length(batch, fastvideo_args)
|
||||
# Determine batch size
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
@@ -69,16 +67,15 @@ class LatentPreparationStage(PipelineStage):
|
||||
if height is None or width is None:
|
||||
raise ValueError("Height and width must be provided")
|
||||
|
||||
assert inference_args.num_channels_latents is not None
|
||||
assert inference_args.vae_scale_factor is not None
|
||||
assert fastvideo_args.num_channels_latents is not None
|
||||
|
||||
# Calculate latent shape
|
||||
shape = (
|
||||
batch_size,
|
||||
inference_args.num_channels_latents,
|
||||
fastvideo_args.num_channels_latents,
|
||||
num_frames,
|
||||
height // inference_args.vae_scale_factor,
|
||||
width // inference_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,20 +103,20 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
return batch
|
||||
|
||||
def adjust_video_length(self, vae: ParallelTiledVAE, batch: ForwardBatch,
|
||||
inference_args: InferenceArgs) -> ForwardBatch:
|
||||
def adjust_video_length(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""
|
||||
Adjust video length based on VAE version.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
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.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -61,19 +61,19 @@ class LlamaEncodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text = prompt_template_video["template"].format(batch.prompt)
|
||||
@@ -123,7 +123,7 @@ class LlamaEncodingStage(PipelineStage):
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
batch.negative_prompt_embeds.append(negative_last_hidden_state)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
@@ -6,8 +6,8 @@ This module contains implementations of prompt encoding stages for diffusion pip
|
||||
"""
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -38,19 +38,19 @@ class T5EncodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text = batch.prompt
|
||||
@@ -109,7 +109,7 @@ class T5EncodingStage(PipelineStage):
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
batch.negative_prompt_embeds.append(neg_prompt_embeds)
|
||||
|
||||
if inference_args.use_cpu_offload:
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ This module contains implementations of timestep preparation stages for diffusio
|
||||
|
||||
import inspect
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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.pipelines.stages.base import PipelineStage
|
||||
@@ -29,14 +29,14 @@ class TimestepPreparationStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Prepare timesteps for the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with prepared timesteps.
|
||||
|
||||
@@ -6,13 +6,16 @@ This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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__)
|
||||
|
||||
@@ -24,7 +27,7 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
|
||||
"image_encoder", "image_processor"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
@@ -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")))
|
||||
@@ -65,15 +67,12 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
|
||||
inference_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
num_channels_latents = self.get_module("transformer").out_channels
|
||||
inference_args.num_channels_latents = num_channels_latents
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = WanImageToVideoPipeline
|
||||
|
||||
@@ -6,7 +6,7 @@ This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
@@ -26,7 +26,7 @@ class WanPipeline(ComposedPipelineBase):
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
@@ -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(
|
||||
@@ -58,15 +57,12 @@ class WanPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
|
||||
inference_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
num_channels_latents = self.get_module("transformer").in_channels
|
||||
inference_args.num_channels_latents = num_channels_latents
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = WanPipeline
|
||||
|
||||
@@ -11,12 +11,12 @@ from einops import rearrange
|
||||
|
||||
from fastvideo.v1.distributed import (init_distributed_environment,
|
||||
initialize_model_parallel)
|
||||
from fastvideo.v1.inference_args import InferenceArgs, prepare_inference_args
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, prepare_fastvideo_args
|
||||
# Fix the import path
|
||||
from fastvideo.v1.inference_engine import InferenceEngine
|
||||
|
||||
|
||||
def initialize_distributed_and_parallelism(inference_args: InferenceArgs):
|
||||
def initialize_distributed_and_parallelism(fastvideo_args: FastVideoArgs):
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
@@ -25,33 +25,33 @@ def initialize_distributed_and_parallelism(inference_args: InferenceArgs):
|
||||
rank=rank,
|
||||
local_rank=local_rank)
|
||||
device_str = f"cuda:{local_rank}"
|
||||
inference_args.device_str = device_str
|
||||
inference_args.device = torch.device(device_str)
|
||||
assert inference_args.sp_size is not None
|
||||
assert inference_args.tp_size is not None
|
||||
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=inference_args.sp_size,
|
||||
tensor_model_parallel_size=inference_args.tp_size,
|
||||
sequence_model_parallel_size=fastvideo_args.sp_size,
|
||||
tensor_model_parallel_size=fastvideo_args.tp_size,
|
||||
)
|
||||
|
||||
|
||||
def main(inference_args: InferenceArgs):
|
||||
initialize_distributed_and_parallelism(inference_args)
|
||||
engine = InferenceEngine.create_engine(inference_args, )
|
||||
def main(fastvideo_args: FastVideoArgs):
|
||||
initialize_distributed_and_parallelism(fastvideo_args)
|
||||
engine = InferenceEngine.create_engine(fastvideo_args, )
|
||||
|
||||
if inference_args.prompt_path is not None:
|
||||
with open(inference_args.prompt_path) as f:
|
||||
if fastvideo_args.prompt_path is not None:
|
||||
with open(fastvideo_args.prompt_path) as f:
|
||||
prompts = [line.strip() for line in f.readlines()]
|
||||
else:
|
||||
if inference_args.prompt is None:
|
||||
if fastvideo_args.prompt is None:
|
||||
raise ValueError("prompt or prompt_path is required")
|
||||
prompts = [inference_args.prompt]
|
||||
prompts = [fastvideo_args.prompt]
|
||||
|
||||
# Process each prompt
|
||||
for prompt in prompts:
|
||||
outputs = engine.run(
|
||||
prompt=prompt,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
|
||||
# Process outputs
|
||||
@@ -63,13 +63,13 @@ def main(inference_args: InferenceArgs):
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
|
||||
# Save video
|
||||
os.makedirs(os.path.dirname(inference_args.output_path), exist_ok=True)
|
||||
imageio.mimsave(os.path.join(inference_args.output_path,
|
||||
os.makedirs(os.path.dirname(fastvideo_args.output_path), exist_ok=True)
|
||||
imageio.mimsave(os.path.join(fastvideo_args.output_path,
|
||||
f"{prompt[:100]}.mp4"),
|
||||
frames,
|
||||
fps=inference_args.fps)
|
||||
fps=fastvideo_args.fps)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
inference_args = prepare_inference_args(sys.argv[1:])
|
||||
main(inference_args)
|
||||
fastvideo_args = prepare_fastvideo_args(sys.argv[1:])
|
||||
main(fastvideo_args)
|
||||
|
||||
@@ -11,7 +11,7 @@ from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
# from fastvideo.v1.models.hunyuan.text_encoder import load_text_encoder, load_tokenizer
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
|
||||
@@ -38,7 +38,7 @@ def test_clip_encoder():
|
||||
- Load models with the same weights and parameters
|
||||
- Produce nearly identical outputs for the same input prompts
|
||||
"""
|
||||
args = InferenceArgs(model_path="openai/clip-vit-large-patch14",
|
||||
args = FastVideoArgs(model_path="openai/clip-vit-large-patch14",
|
||||
precision="float16")
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ from transformers import AutoConfig
|
||||
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
@@ -38,7 +38,7 @@ def test_llama_encoder():
|
||||
- Load models with the same weights and parameters
|
||||
- Produce nearly identical outputs for the same input prompts
|
||||
"""
|
||||
args = InferenceArgs(model_path="meta-llama/Llama-2-7b-hf",
|
||||
args = FastVideoArgs(model_path="meta-llama/Llama-2-7b-hf",
|
||||
precision="float16")
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -7,7 +7,7 @@ import torch
|
||||
from diffusers import WanTransformer3DModel
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TransformerLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
@@ -29,7 +29,7 @@ def test_wan_transformer():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = InferenceArgs(model_path=TRANSFORMER_PATH,
|
||||
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
|
||||
use_cpu_offload=False,
|
||||
precision=precision_str)
|
||||
args.device = device
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -6,9 +6,10 @@ import pytest
|
||||
import torch
|
||||
from diffusers import AutoencoderKLWan
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
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__)
|
||||
@@ -28,11 +29,13 @@ def test_wan_vae():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.bfloat16
|
||||
precision_str = "bf16"
|
||||
args = InferenceArgs(model_path=VAE_PATH, vae_precision=precision_str)
|
||||
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)
|
||||
assert model2.use_feature_cache # Default to use the original WanVAE algorithm
|
||||
|
||||
model1 = AutoencoderKLWan.from_pretrained(
|
||||
VAE_PATH, torch_dtype=precision).to(device).eval()
|
||||
@@ -48,43 +51,48 @@ def test_wan_vae():
|
||||
32,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
latent_tensor = torch.randn(batch_size,
|
||||
16,
|
||||
21,
|
||||
32,
|
||||
32,
|
||||
device=device,
|
||||
dtype=precision)
|
||||
# latent_tensor = torch.randn(batch_size,
|
||||
# 16,
|
||||
# 21,
|
||||
# 32,
|
||||
# 32,
|
||||
# device=device,
|
||||
# dtype=precision)
|
||||
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
# Test encoding
|
||||
logger.info("Testing encoding...")
|
||||
latent1 = model1.encode(input_tensor).latent_dist.mean
|
||||
latent1 = model1.encode(input_tensor).latent_dist
|
||||
print("--------------------------------")
|
||||
latent2 = model2.encode(input_tensor).mean
|
||||
latent2 = model2.encode(input_tensor)
|
||||
# Check if latents have the same shape
|
||||
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
|
||||
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
|
||||
assert latent1.mean.shape == latent2.mean.shape, f"Latent shapes don't match: {latent1.mean.shape} vs {latent2.mean.shape}"
|
||||
# Check if latents are similar
|
||||
max_diff_encode = torch.max(torch.abs(latent1 - latent2))
|
||||
mean_diff_encode = torch.mean(torch.abs(latent1 - latent2))
|
||||
max_diff_encode = torch.max(torch.abs(latent1.mean - latent2.mean))
|
||||
mean_diff_encode = torch.mean(torch.abs(latent1.mean - latent2.mean))
|
||||
logger.info("Maximum difference between encoded latents: %s",
|
||||
max_diff_encode.item())
|
||||
logger.info("Mean difference between encoded latents: %s",
|
||||
mean_diff_encode.item())
|
||||
assert mean_diff_encode < 5e-1, f"Encoded latents differ significantly: mean diff = {mean_diff_encode.item()}"
|
||||
assert max_diff_encode < 1e-5, f"Encoded latents differ significantly: max diff = {mean_diff_encode.item()}"
|
||||
# Test decoding
|
||||
logger.info("Testing decoding...")
|
||||
latents_mean = (torch.tensor(model1.config.latents_mean).view(
|
||||
1, model1.config.z_dim, 1, 1, 1).to(latent_tensor.device,
|
||||
latent_tensor.dtype))
|
||||
latents_std = 1.0 / torch.tensor(model1.config.latents_std).view(
|
||||
1, model1.config.z_dim, 1, 1, 1).to(latent_tensor.device,
|
||||
latent_tensor.dtype)
|
||||
latent_tensor = latent_tensor / latents_std + latents_mean
|
||||
output2 = model2.decode(latent_tensor)
|
||||
output1 = model1.decode(latent_tensor).sample
|
||||
latent1_tensor = latent1.mode()
|
||||
mean1 = (torch.tensor(model1.config.latents_mean).view(
|
||||
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
|
||||
input_tensor.dtype))
|
||||
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 / 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()
|
||||
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}"
|
||||
|
||||
@@ -95,4 +103,4 @@ def test_wan_vae():
|
||||
max_diff_decode.item())
|
||||
logger.info("Mean difference between decoded outputs: %s",
|
||||
mean_diff_decode.item())
|
||||
assert mean_diff_decode < 1e-1, f"Decoded outputs differ significantly: mean diff = {mean_diff_decode.item()}"
|
||||
assert max_diff_decode < 1e-5, f"Decoded outputs differ significantly: max diff = {mean_diff_decode.item()}"
|
||||
+131
-7
@@ -2,17 +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
|
||||
from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast
|
||||
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
|
||||
@@ -25,7 +31,7 @@ logger = init_logger(__name__)
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
# TODO(will): used to convert inference_args.precision to torch.dtype. Find a
|
||||
# TODO(will): used to convert fastvideo_args.precision to torch.dtype. Find a
|
||||
# cleaner way to do this.
|
||||
PRECISION_TO_TYPE = {
|
||||
"fp32": torch.float32,
|
||||
@@ -364,9 +370,7 @@ def import_pynvml():
|
||||
install FastVideo. It provides a Python module named `pynvml`.
|
||||
- `pynvml` (https://pypi.org/project/pynvml/): An unofficial wrapper.
|
||||
Prior to version 12.0, it also provides a Python module `pynvml`,
|
||||
and therefore conflicts with the official one. What's worse,
|
||||
the module is a Python package, and has higher priority than
|
||||
the official one which is a standalone Python file.
|
||||
and therefore conflicts with the official one which is a standalone Python file.
|
||||
This causes errors when both of them are installed.
|
||||
Starting from version 12.0, it migrates to a new module
|
||||
named `pynvml_utils` to avoid the conflict.
|
||||
@@ -383,12 +387,15 @@ def import_pynvml():
|
||||
|
||||
|
||||
def maybe_download_model(model_path: str,
|
||||
local_dir: Optional[str] = None) -> str:
|
||||
local_dir: Optional[str] = None,
|
||||
download: bool = True) -> str:
|
||||
"""
|
||||
Check if the model path is a Hugging Face Hub model ID and download it if needed.
|
||||
|
||||
Args:
|
||||
model_path: Local path or Hugging Face Hub model ID
|
||||
local_dir: Local directory to save the model
|
||||
download: Whether to download the model from Hugging Face Hub
|
||||
|
||||
Returns:
|
||||
Local path to the model
|
||||
@@ -457,3 +464,120 @@ def verify_model_config_and_directory(model_path: str) -> Dict[str, Any]:
|
||||
|
||||
logger.info("Diffusers version: %s", config["_diffusers_version"])
|
||||
return cast(Dict[str, Any], config)
|
||||
|
||||
|
||||
def maybe_download_model_index(model_name_or_path: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Download and extract just the model_index.json for a Hugging Face model.
|
||||
|
||||
Args:
|
||||
model_name_or_path: Path or HF Hub model ID
|
||||
|
||||
Returns:
|
||||
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
|
||||
if os.path.exists(model_name_or_path):
|
||||
return verify_model_config_and_directory(model_name_or_path)
|
||||
|
||||
# For remote models, download just the model_index.json
|
||||
try:
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
# Download just the model_index.json file
|
||||
model_index_path = hf_hub_download(repo_id=model_name_or_path,
|
||||
filename="model_index.json",
|
||||
local_dir=tmp_dir)
|
||||
|
||||
# Load the model_index.json
|
||||
with open(model_index_path) as f:
|
||||
config: Dict[str, Any] = json.load(f)
|
||||
|
||||
# Verify it has the required fields
|
||||
if "_class_name" not in config:
|
||||
raise ValueError(
|
||||
f"model_index.json for {model_name_or_path} does not contain _class_name field"
|
||||
)
|
||||
|
||||
if "_diffusers_version" not in config:
|
||||
raise ValueError(
|
||||
f"model_index.json for {model_name_or_path} does not contain _diffusers_version field"
|
||||
)
|
||||
|
||||
# Add the pipeline name for downstream use
|
||||
config["pipeline_name"] = config["_class_name"]
|
||||
|
||||
logger.info("Downloaded model_index.json for %s, pipeline: %s",
|
||||
model_name_or_path, config["_class_name"])
|
||||
return config
|
||||
|
||||
except Exception as e:
|
||||
raise ValueError(
|
||||
f"Failed to download or parse model_index.json for {model_name_or_path}: {e}"
|
||||
) from e
|
||||
|
||||
|
||||
def update_environment_variables(envs: Dict[str, str]):
|
||||
for k, v in envs.items():
|
||||
if k in os.environ and os.environ[k] != v:
|
||||
logger.warning(
|
||||
"Overwriting environment variable %s "
|
||||
"from '%s' to '%s'", k, os.environ[k], v)
|
||||
os.environ[k] = v
|
||||
|
||||
|
||||
def run_method(obj: Any, method: Union[str, bytes, Callable], args: tuple[Any],
|
||||
kwargs: dict[str, Any]) -> Any:
|
||||
"""
|
||||
Run a method of an object with the given arguments and keyword arguments.
|
||||
If the method is string, it will be converted to a method using getattr.
|
||||
If the method is serialized bytes and will be deserialized using
|
||||
cloudpickle.
|
||||
If the method is a callable, it will be called directly.
|
||||
"""
|
||||
if isinstance(method, bytes):
|
||||
func = partial(cloudpickle.loads(method), obj)
|
||||
elif isinstance(method, str):
|
||||
try:
|
||||
func = getattr(obj, method)
|
||||
except AttributeError:
|
||||
raise NotImplementedError(f"Method {method!r} is not"
|
||||
" implemented.") from None
|
||||
else:
|
||||
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 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 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}")
|
||||
|
||||
@@ -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
|
||||
@@ -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]
|
||||
@@ -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()
|
||||
+3
-3
@@ -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"
|
||||
@@ -36,8 +36,8 @@ dependencies = [
|
||||
"wandb==0.18.5", "loguru", "test-tube==0.7.5",
|
||||
|
||||
# Miscellaneous Utilities
|
||||
"tqdm==4.66.5", "PyYAML==6.0.1", "idna==3.6", "protobuf==5.28.3",
|
||||
"gradio==5.3.0", "moviepy==1.0.3", "flask",
|
||||
"tqdm==4.66.5", "PyYAML==6.0.1", "protobuf==5.28.3",
|
||||
"gradio>=5.22.0", "moviepy==1.0.3", "flask",
|
||||
"flask_restful", "aiohttp", "huggingface_hub", "cloudpickle",
|
||||
# System & Monitoring Tools
|
||||
"gpustat", "watch",
|
||||
|
||||
@@ -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
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user