Compare commits

...
26 Commits
Author SHA1 Message Date
William Lin 59ab481eb1 release 0.0.5 (#399) 2025-05-11 16:23:11 -07:00
William Lin 0cf001986a [Docs] More docs update (#394) 2025-05-11 16:19:47 -07:00
applesaucethebunandBrayden Zhong 51956369a5 [Misc] Replace instances of time.time() with time.perf_counter() (#396)
Co-authored-by: Brayden Zhong <b8zhong@uwaterloo.ca>
2025-05-11 15:56:39 -07:00
Kevin Lin 4b0970cbbf [CI] Set volume_size required to false (#398) 2025-05-11 15:55:34 -07:00
Kevin Lin 94bf47a572 [CI] Set default disk size (#397) 2025-05-11 15:34:17 -07:00
River (Zihang He)andWill Lin 1a3ac9074b Zihang stepvideo v1 (#389)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-05-11 14:41:42 -07:00
Kevin Lin fb0581d5b0 [CLI] Update cli to support new api/model config (#384) 2025-05-11 13:50:33 -07:00
Kevin Lin 6c74ab4132 [CI] Use python 3.10/3.11 for SSIM test (#392) 2025-05-10 13:48:29 -07:00
William Lin 9a91021c56 [Docs] Add collect_env.py and various docs update (#393) 2025-05-10 13:42:42 -07:00
Wei Zhou 51c94d6a73 [Misc] Small Fixes & Features (#390) 2025-05-07 15:40:49 -07:00
William Lin dba38dbc03 Release 0.0.4 (#388) 2025-05-07 02:40:42 -07:00
William Lin 2034cc3c4f [misc] Improve worker cleanup (#387) 2025-05-06 22:45:22 -07:00
William Lin c69afce2f6 [Docs] Update for V1 (#381) 2025-05-06 16:34:44 -07:00
Wei Zhou b08e758eb3 Small Fixes & Features (#378) 2025-05-06 15:52:34 -07:00
William Lin 3f3462d7ce Cleanup Teacache params (#386) 2025-05-06 15:52:05 -07:00
William Lin c9c47dd89c Add Teacache to V1 (#371) 2025-05-06 11:58:56 -07:00
William Lin 9c4ef7c2f1 [Lint] fix (#382) 2025-05-06 00:42:55 -07:00
Kevin Lin f25eb4b905 [CI] Add write permissions to build-image workflow (#379) 2025-05-05 19:21:03 -07:00
Kevin Lin 048d55ccbb [CI] Add new images for different Python versions (#377) 2025-05-05 10:28:33 -07:00
Wei Zhou a271c55fe4 Fix FSDP issues when using cpu_offload flag (#376) 2025-05-02 15:54:52 -07:00
William Lin f663ae0d8a change gradio example to use model configs (#375) 2025-05-02 10:42:03 -07:00
William Lin c0911aa3dd release 0.0.3 (#374) 2025-05-02 03:50:41 -07:00
William Lin 5f59687ae7 Fix model config for python 3.11+ (#373) 2025-05-02 03:18:56 -07:00
Wei Zhou 5adbc81cdc Refactor encoder (#370) 2025-05-02 02:41:12 -07:00
Kevin Lin f26d5c37c1 Update SSIM tests to use new API (#369) 2025-05-01 04:14:24 -07:00
Wei Zhou 6a4ef42378 [V1] Model config (#358) 2025-04-30 14:48:22 -07:00
151 changed files with 8669 additions and 2644 deletions
+106
View File
@@ -0,0 +1,106 @@
name: Build Image Template
on:
workflow_call:
inputs:
python_version:
required: true
type: string
dockerfile_path:
required: true
type: string
tag_suffix:
required: true
type: string
jobs:
build-and-push:
runs-on: ubuntu-latest
permissions:
contents: read
packages: write
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Free up disk space
run: |
# Display initial space
echo "Initial disk space:"
df -h
# Remove large directories directly
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/*
# Clean Docker
docker system prune -af --volumes
# Display available space after cleanup
echo "Disk space after cleanup:"
df -h
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Login to GitHub Container Registry
uses: docker/login-action@v3
with:
registry: ghcr.io
username: ${{ github.repository_owner }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Prepare tags
id: prepare-tags
run: |
SHORT_SHA=$(echo ${{ github.sha }} | cut -c1-7)
TAGS="type=raw,value=${{ inputs.tag_suffix }}-latest"
TAGS="${TAGS}\ntype=raw,value=${{ inputs.tag_suffix }}-sha-${SHORT_SHA}"
# Set Python 3.10 as the default image
if [[ "${{ inputs.python_version }}" == "3.10" ]]; then
TAGS="${TAGS}\ntype=raw,value=latest"
fi
{
echo "tags<<EOF"
echo -e "$TAGS"
echo "EOF"
} >> $GITHUB_OUTPUT
- name: Extract metadata for Docker
id: meta
uses: docker/metadata-action@v5
with:
images: ghcr.io/${{ github.repository }}/fastvideo-dev
tags: ${{ steps.prepare-tags.outputs.tags }}
- name: Build and push Docker image
id: build-push
uses: docker/build-push-action@v6
with:
context: .
file: ${{ inputs.dockerfile_path }}
push: true
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}
cache-from: type=gha
cache-to: type=gha,mode=max
- name: Success message
run: |
echo "✅ Python ${{ inputs.python_version }} image successfully built and pushed to ghcr.io/${{ github.repository }}/fastvideo-dev:${{ inputs.tag_suffix }}-latest"
echo "To run tests with this image, manually trigger the 'Run Tests' workflow."
+48 -74
View File
@@ -1,78 +1,52 @@
name: Build and Push Docker Image
name: Build and Push Docker Images
on:
workflow_dispatch: # Only manual triggers
workflow_dispatch:
inputs:
python_3_10:
description: 'Build Python 3.10 image'
required: false
default: false
type: boolean
python_3_11:
description: 'Build Python 3.11 image'
required: false
default: false
type: boolean
python_3_12:
description: 'Build Python 3.12 image'
required: false
default: false
type: boolean
permissions:
contents: read
packages: write
jobs:
build-and-push:
runs-on: ubuntu-latest
permissions:
contents: read
packages: write
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Free up disk space
run: |
# Display initial space
echo "Initial disk space:"
df -h
# Remove large directories directly
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/*
# Clean Docker
docker system prune -af --volumes
# Display available space after cleanup
echo "Disk space after cleanup:"
df -h
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Login to GitHub Container Registry
uses: docker/login-action@v3
with:
registry: ghcr.io
username: ${{ github.repository_owner }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Extract metadata for Docker
id: meta
uses: docker/metadata-action@v5
with:
images: ghcr.io/${{ github.repository }}/fastvideo-dev
tags: |
type=raw,value=latest
type=sha,format=short
- name: Build and push Docker image
id: build-push
uses: docker/build-push-action@v6
with:
context: .
push: true
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}
cache-from: type=gha
cache-to: type=gha,mode=max
- name: Success message
run: |
echo "✅ Image successfully built and pushed to ghcr.io/${{ github.repository }}/fastvideo-dev:latest"
echo "To run tests with this image, manually trigger the 'Run Tests' workflow."
build-python-3-10:
if: ${{ github.event.inputs.python_3_10 == 'true' }}
uses: ./.github/workflows/build-image-template.yml
with:
python_version: '3.10'
dockerfile_path: docker/Dockerfile.python3.10
tag_suffix: py3.10
secrets: inherit
build-python-3-11:
if: ${{ github.event.inputs.python_3_11 == 'true' }}
uses: ./.github/workflows/build-image-template.yml
with:
python_version: '3.11'
dockerfile_path: docker/Dockerfile.python3.11
tag_suffix: py3.11
secrets: inherit
build-python-3-12:
if: ${{ github.event.inputs.python_3_12 == 'true' }}
uses: ./.github/workflows/build-image-template.yml
with:
python_version: '3.12'
dockerfile_path: docker/Dockerfile.python3.12
tag_suffix: py3.12
secrets: inherit
+58 -170
View File
@@ -83,193 +83,81 @@ jobs:
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.encoder-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_encoder_test == 'true')
runs-on: ubuntu-latest
environment: runpod-runners
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Set up SSH key
run: |
mkdir -p ~/.ssh
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
- name: Install dependencies
run: pip install requests
- name: Run tests on RunPod
env:
JOB_ID: "encoder-test"
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: 30
run: >-
python .github/scripts/runpod_api.py
--gpu-type "NVIDIA A40"
--gpu-count 1
--volume-size 100
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
env:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
JOB_ID: "encoder-test"
run: python .github/scripts/runpod_cleanup.py
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "encoder-test"
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
vae-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.vae-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_vae_test == 'true')
runs-on: ubuntu-latest
environment: runpod-runners
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Set up SSH key
run: |
mkdir -p ~/.ssh
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
- name: Install dependencies
run: pip install requests
- name: Run tests on RunPod
env:
JOB_ID: "vae-test"
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: 30
run: >-
python .github/scripts/runpod_api.py
--gpu-type "NVIDIA A40"
--gpu-count 1
--volume-size 100
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
env:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
JOB_ID: "vae-test"
run: python .github/scripts/runpod_cleanup.py
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "vae-test"
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
transformer-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.transformer-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_transformer_test == 'true')
runs-on: ubuntu-latest
environment: runpod-runners
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Set up SSH key
run: |
mkdir -p ~/.ssh
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
- name: Install dependencies
run: pip install requests
- name: Run tests on RunPod
env:
JOB_ID: "transformer-test"
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: 30
run: >-
python .github/scripts/runpod_api.py
--gpu-type "NVIDIA L40S"
--gpu-count 1
--volume-size 100
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
env:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
JOB_ID: "transformer-test"
run: python .github/scripts/runpod_cleanup.py
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "transformer-test"
gpu_type: "NVIDIA L40S"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
ssim-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
runs-on: ubuntu-latest
environment: runpod-runners
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Set up SSH key
run: |
mkdir -p ~/.ssh
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
- name: Install dependencies
run: pip install requests
- name: Run tests on RunPod
env:
JOB_ID: "ssim-test"
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: 45
run: >-
python .github/scripts/runpod_api.py
--gpu-type "NVIDIA A40"
--gpu-count 2
--disk-size 200
--volume-size 200
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
- name: Terminate RunPod Instances
if: ${{ always() }}
env:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
JOB_ID: "ssim-test"
run: python .github/scripts/runpod_cleanup.py
strategy:
fail-fast: false
matrix:
python-version: [
{version: "3.10", tag: "latest"},
{version: "3.11", tag: "py3.11-latest"},
{version: "3.12", tag: "py3.12-latest"}
]
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "ssim-test-py${{ matrix.python-version.version }}"
gpu_type: "NVIDIA A40"
gpu_count: 2
volume_size: 200
disk_size: 200
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
timeout_minutes: 60
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
runpod-cleanup:
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
@@ -289,7 +177,7 @@ jobs:
- name: Cleanup all RunPod instances
env:
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test"]' # JSON array of job IDs
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12"]'
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
run: python .github/scripts/runpod_cleanup.py
+91
View File
@@ -0,0 +1,91 @@
name: RunPod Test
on:
workflow_call:
inputs:
job_id:
required: true
type: string
description: "Unique identifier for this test job"
gpu_type:
required: true
type: string
description: "GPU type to use (e.g. NVIDIA A40, NVIDIA L40S)"
gpu_count:
required: true
type: number
description: "Number of GPUs to use"
volume_size:
required: false
type: number
default: 20
description: "Volume size in GB"
disk_size:
required: false
type: number
default: 20
description: "Disk size in GB"
image:
required: true
type: string
description: "Docker image to use"
test_command:
required: true
type: string
description: "Command to run tests"
timeout_minutes:
required: false
type: number
default: 30
description: "Timeout in minutes"
secrets:
RUNPOD_API_KEY:
required: true
RUNPOD_PRIVATE_KEY:
required: true
jobs:
run-test:
runs-on: ubuntu-latest
environment: runpod-runners
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Set up SSH key
run: |
mkdir -p ~/.ssh
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
- name: Install dependencies
run: pip install requests
- name: Run tests on RunPod
env:
JOB_ID: ${{ inputs.job_id }}
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: ${{ inputs.timeout_minutes }}
run: >-
python .github/scripts/runpod_api.py
--gpu-type "${{ inputs.gpu_type }}"
--gpu-count ${{ inputs.gpu_count }}
--volume-size ${{ inputs.volume_size }}
--disk-size ${{ inputs.disk_size }}
--image "${{ inputs.image }}"
--test-command "${{ inputs.test_command }}"
- name: Terminate RunPod Instances
if: ${{ always() }}
env:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
JOB_ID: ${{ inputs.job_id }}
run: python .github/scripts/runpod_cleanup.py
+4 -2
View File
@@ -19,9 +19,11 @@ exclude: |
fastvideo/sample/.*|
fastvideo/train\.py|
fastvideo/utils/.*|
fastvideo/v1/examples/.*|
examples/.*|
.github/workflows/fastvideo-publish.yml|
.github/workflows/sta-publish.yml
.github/workflows/sta-publish.yml|
.github/workflows/build-image-template.yml|
docs/source/inference/support_matrix.md
)
repos:
- repo: https://github.com/google/yapf
+772
View File
@@ -0,0 +1,772 @@
# SPDX-License-Identifier: Apache-2.0
# type: ignore
# ruff: noqa
# code borrowed from https://github.com/pytorch/pytorch/blob/main/torch/utils/collect_env.py
# and vllm: https://github.com/vllm-project/vllm/blob/main/vllm/collect_env.py
import datetime
import locale
import os
import re
import subprocess
import sys
# Unlike the rest of the PyTorch this file must be python2 compliant.
# This script outputs relevant system environment info
# Run it with `python collect_env.py` or `python -m torch.utils.collect_env`
from collections import namedtuple
from fastvideo.v1.envs import environment_variables
try:
import torch
TORCH_AVAILABLE = True
except (ImportError, NameError, AttributeError, OSError):
TORCH_AVAILABLE = False
# System Environment Information
SystemEnv = namedtuple(
'SystemEnv',
[
'torch_version',
'is_debug_build',
'cuda_compiled_version',
'gcc_version',
'clang_version',
'cmake_version',
'os',
'libc_version',
'python_version',
'python_platform',
'is_cuda_available',
'cuda_runtime_version',
'cuda_module_loading',
'nvidia_driver_version',
'nvidia_gpu_models',
'cudnn_version',
'pip_version', # 'pip' or 'pip3'
'pip_packages',
'conda_packages',
'hip_compiled_version',
'hip_runtime_version',
'miopen_runtime_version',
'caching_allocator_config',
'is_xnnpack_available',
'cpu_info',
'fastvideo_version',
'fastvideo_build_flags',
'gpu_topo',
'env_vars',
])
DEFAULT_CONDA_PATTERNS = {
"torch",
"numpy",
"cudatoolkit",
"soumith",
"mkl",
"magma",
"triton",
"optree",
"nccl",
"transformers",
"zmq",
"nvidia",
"pynvml",
}
DEFAULT_PIP_PATTERNS = {
"torch",
"numpy",
"mypy",
"flake8",
"triton",
"optree",
"onnx",
"nccl",
"transformers",
"zmq",
"nvidia",
"pynvml",
}
def run(command):
"""Return (return-code, stdout, stderr)."""
shell = True if type(command) is str else False
p = subprocess.Popen(command,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
shell=shell)
raw_output, raw_err = p.communicate()
rc = p.returncode
if get_platform() == 'win32':
enc = 'oem'
else:
enc = locale.getpreferredencoding()
output = raw_output.decode(enc)
if command == 'nvidia-smi topo -m':
# don't remove the leading whitespace of `nvidia-smi topo -m`
# because they are meaningful
output = output.rstrip()
else:
output = output.strip()
err = raw_err.decode(enc)
return rc, output, err.strip()
def run_and_read_all(run_lambda, command):
"""Run command using run_lambda; reads and returns entire output if rc is 0."""
rc, out, _ = run_lambda(command)
if rc != 0:
return None
return out
def run_and_parse_first_match(run_lambda, command, regex):
"""Run command using run_lambda, returns the first regex match if it exists."""
rc, out, _ = run_lambda(command)
if rc != 0:
return None
match = re.search(regex, out)
if match is None:
return None
return match.group(1)
def run_and_return_first_line(run_lambda, command):
"""Run command using run_lambda and returns first line if output is not empty."""
rc, out, _ = run_lambda(command)
if rc != 0:
return None
return out.split('\n')[0]
def get_conda_packages(run_lambda, patterns=None):
if patterns is None:
patterns = DEFAULT_CONDA_PATTERNS
conda = os.environ.get('CONDA_EXE', 'conda')
out = run_and_read_all(run_lambda, "{} list".format(conda))
if out is None:
return out
return "\n".join(line for line in out.splitlines()
if not line.startswith("#") and any(name in line
for name in patterns))
def get_gcc_version(run_lambda):
return run_and_parse_first_match(run_lambda, 'gcc --version', r'gcc (.*)')
def get_clang_version(run_lambda):
return run_and_parse_first_match(run_lambda, 'clang --version',
r'clang version (.*)')
def get_cmake_version(run_lambda):
return run_and_parse_first_match(run_lambda, 'cmake --version',
r'cmake (.*)')
def get_nvidia_driver_version(run_lambda):
if get_platform() == 'darwin':
cmd = 'kextstat | grep -i cuda'
return run_and_parse_first_match(run_lambda, cmd,
r'com[.]nvidia[.]CUDA [(](.*?)[)]')
smi = get_nvidia_smi()
return run_and_parse_first_match(run_lambda, smi, r'Driver Version: (.*?) ')
def get_gpu_info(run_lambda):
if get_platform() == 'darwin' or (TORCH_AVAILABLE and hasattr(
torch.version, 'hip') and torch.version.hip is not None):
if TORCH_AVAILABLE and torch.cuda.is_available():
if torch.version.hip is not None:
prop = torch.cuda.get_device_properties(0)
if hasattr(prop, "gcnArchName"):
gcnArch = " ({})".format(prop.gcnArchName)
else:
gcnArch = "NoGCNArchNameOnOldPyTorch"
else:
gcnArch = ""
return torch.cuda.get_device_name(None) + gcnArch
return None
smi = get_nvidia_smi()
uuid_regex = re.compile(r' \(UUID: .+?\)')
rc, out, _ = run_lambda(smi + ' -L')
if rc != 0:
return None
# Anonymize GPUs by removing their UUID
return re.sub(uuid_regex, '', out)
def get_running_cuda_version(run_lambda):
return run_and_parse_first_match(run_lambda, 'nvcc --version',
r'release .+ V(.*)')
def get_cudnn_version(run_lambda):
"""Return a list of libcudnn.so; it's hard to tell which one is being used."""
if get_platform() == 'win32':
system_root = os.environ.get('SYSTEMROOT', 'C:\\Windows')
cuda_path = os.environ.get('CUDA_PATH', "%CUDA_PATH%")
where_cmd = os.path.join(system_root, 'System32', 'where')
cudnn_cmd = '{} /R "{}\\bin" cudnn*.dll'.format(where_cmd, cuda_path)
elif get_platform() == 'darwin':
# CUDA libraries and drivers can be found in /usr/local/cuda/. See
# https://docs.nvidia.com/cuda/cuda-installation-guide-mac-os-x/index.html#install
# https://docs.nvidia.com/deeplearning/sdk/cudnn-install/index.html#installmac
# Use CUDNN_LIBRARY when cudnn library is installed elsewhere.
cudnn_cmd = 'ls /usr/local/cuda/lib/libcudnn*'
else:
cudnn_cmd = 'ldconfig -p | grep libcudnn | rev | cut -d" " -f1 | rev'
rc, out, _ = run_lambda(cudnn_cmd)
# find will return 1 if there are permission errors or if not found
if len(out) == 0 or (rc != 1 and rc != 0):
l = os.environ.get('CUDNN_LIBRARY')
if l is not None and os.path.isfile(l):
return os.path.realpath(l)
return None
files_set = set()
for fn in out.split('\n'):
fn = os.path.realpath(fn) # eliminate symbolic links
if os.path.isfile(fn):
files_set.add(fn)
if not files_set:
return None
# Alphabetize the result because the order is non-deterministic otherwise
files = sorted(files_set)
if len(files) == 1:
return files[0]
result = '\n'.join(files)
return 'Probably one of the following:\n{}'.format(result)
def get_nvidia_smi():
# Note: nvidia-smi is currently available only on Windows and Linux
smi = 'nvidia-smi'
if get_platform() == 'win32':
system_root = os.environ.get('SYSTEMROOT', 'C:\\Windows')
program_files_root = os.environ.get('PROGRAMFILES', 'C:\\Program Files')
legacy_path = os.path.join(program_files_root, 'NVIDIA Corporation',
'NVSMI', smi)
new_path = os.path.join(system_root, 'System32', smi)
smis = [new_path, legacy_path]
for candidate_smi in smis:
if os.path.exists(candidate_smi):
smi = '"{}"'.format(candidate_smi)
break
return smi
def get_fastvideo_version():
return ""
from fastvideo import __version__, __version_tuple__
if __version__ == "dev":
return "N/A (dev)"
version_str = __version_tuple__[-1]
if isinstance(version_str, str) and version_str.startswith('g'):
# it's a dev build
if '.' in version_str:
# it's a dev build containing local changes
git_sha = version_str.split('.')[0][1:]
date = version_str.split('.')[-1][1:]
return f"{__version__} (git sha: {git_sha}, date: {date})"
else:
# it's a dev build without local changes
git_sha = version_str[1:] # type: ignore
return f"{__version__} (git sha: {git_sha})"
return __version__
def summarize_fastvideo_build_flags():
# This could be a static method if the flags are constant, or dynamic if you need to check environment variables, etc.
return 'CUDA Archs: {}; ROCm: {}; Neuron: {}'.format(
os.environ.get('TORCH_CUDA_ARCH_LIST', 'Not Set'),
'Enabled' if os.environ.get('ROCM_HOME') else 'Disabled',
'Enabled' if os.environ.get('NEURON_CORES') else 'Disabled',
)
def get_gpu_topo(run_lambda):
output = None
if get_platform() == 'linux':
output = run_and_read_all(run_lambda, 'nvidia-smi topo -m')
if output is None:
output = run_and_read_all(run_lambda, 'rocm-smi --showtopo')
return output
# example outputs of CPU infos
# * linux
# Architecture: x86_64
# CPU op-mode(s): 32-bit, 64-bit
# Address sizes: 46 bits physical, 48 bits virtual
# Byte Order: Little Endian
# CPU(s): 128
# On-line CPU(s) list: 0-127
# Vendor ID: GenuineIntel
# Model name: Intel(R) Xeon(R) Platinum 8375C CPU @ 2.90GHz
# CPU family: 6
# Model: 106
# Thread(s) per core: 2
# Core(s) per socket: 32
# Socket(s): 2
# Stepping: 6
# BogoMIPS: 5799.78
# Flags: fpu vme de pse tsc msr pae mce cx8 apic sep mtrr pge mca cmov pat pse36 clflush mmx fxsr
# sse sse2 ss ht syscall nx pdpe1gb rdtscp lm constant_tsc arch_perfmon rep_good nopl
# xtopology nonstop_tsc cpuid aperfmperf tsc_known_freq pni pclmulqdq monitor ssse3 fma cx16
# pcid sse4_1 sse4_2 x2apic movbe popcnt tsc_deadline_timer aes xsave avx f16c rdrand
# hypervisor lahf_lm abm 3dnowprefetch invpcid_single ssbd ibrs ibpb stibp ibrs_enhanced
# fsgsbase tsc_adjust bmi1 avx2 smep bmi2 erms invpcid avx512f avx512dq rdseed adx smap
# avx512ifma clflushopt clwb avx512cd sha_ni avx512bw avx512vl xsaveopt xsavec xgetbv1
# xsaves wbnoinvd ida arat avx512vbmi pku ospke avx512_vbmi2 gfni vaes vpclmulqdq
# avx512_vnni avx512_bitalg tme avx512_vpopcntdq rdpid md_clear flush_l1d arch_capabilities
# Virtualization features:
# Hypervisor vendor: KVM
# Virtualization type: full
# Caches (sum of all):
# L1d: 3 MiB (64 instances)
# L1i: 2 MiB (64 instances)
# L2: 80 MiB (64 instances)
# L3: 108 MiB (2 instances)
# NUMA:
# NUMA node(s): 2
# NUMA node0 CPU(s): 0-31,64-95
# NUMA node1 CPU(s): 32-63,96-127
# Vulnerabilities:
# Itlb multihit: Not affected
# L1tf: Not affected
# Mds: Not affected
# Meltdown: Not affected
# Mmio stale data: Vulnerable: Clear CPU buffers attempted, no microcode; SMT Host state unknown
# Retbleed: Not affected
# Spec store bypass: Mitigation; Speculative Store Bypass disabled via prctl and seccomp
# Spectre v1: Mitigation; usercopy/swapgs barriers and __user pointer sanitization
# Spectre v2: Mitigation; Enhanced IBRS, IBPB conditional, RSB filling, PBRSB-eIBRS SW sequence
# Srbds: Not affected
# Tsx async abort: Not affected
# * win32
# Architecture=9
# CurrentClockSpeed=2900
# DeviceID=CPU0
# Family=179
# L2CacheSize=40960
# L2CacheSpeed=
# Manufacturer=GenuineIntel
# MaxClockSpeed=2900
# Name=Intel(R) Xeon(R) Platinum 8375C CPU @ 2.90GHz
# ProcessorType=3
# Revision=27142
#
# Architecture=9
# CurrentClockSpeed=2900
# DeviceID=CPU1
# Family=179
# L2CacheSize=40960
# L2CacheSpeed=
# Manufacturer=GenuineIntel
# MaxClockSpeed=2900
# Name=Intel(R) Xeon(R) Platinum 8375C CPU @ 2.90GHz
# ProcessorType=3
# Revision=27142
def get_cpu_info(run_lambda):
rc, out, err = 0, '', ''
if get_platform() == 'linux':
rc, out, err = run_lambda('lscpu')
elif get_platform() == 'win32':
rc, out, err = run_lambda(
'wmic cpu get Name,Manufacturer,Family,Architecture,ProcessorType,DeviceID, \
CurrentClockSpeed,MaxClockSpeed,L2CacheSize,L2CacheSpeed,Revision /VALUE'
)
elif get_platform() == 'darwin':
rc, out, err = run_lambda("sysctl -n machdep.cpu.brand_string")
cpu_info = 'None'
if rc == 0:
cpu_info = out
else:
cpu_info = err
return cpu_info
def get_platform():
if sys.platform.startswith('linux'):
return 'linux'
elif sys.platform.startswith('win32'):
return 'win32'
elif sys.platform.startswith('cygwin'):
return 'cygwin'
elif sys.platform.startswith('darwin'):
return 'darwin'
else:
return sys.platform
def get_mac_version(run_lambda):
return run_and_parse_first_match(run_lambda, 'sw_vers -productVersion',
r'(.*)')
def get_windows_version(run_lambda):
system_root = os.environ.get('SYSTEMROOT', 'C:\\Windows')
wmic_cmd = os.path.join(system_root, 'System32', 'Wbem', 'wmic')
findstr_cmd = os.path.join(system_root, 'System32', 'findstr')
return run_and_read_all(
run_lambda,
'{} os get Caption | {} /v Caption'.format(wmic_cmd, findstr_cmd))
def get_lsb_version(run_lambda):
return run_and_parse_first_match(run_lambda, 'lsb_release -a',
r'Description:\t(.*)')
def check_release_file(run_lambda):
return run_and_parse_first_match(run_lambda, 'cat /etc/*-release',
r'PRETTY_NAME="(.*)"')
def get_os(run_lambda):
from platform import machine
platform = get_platform()
if platform == 'win32' or platform == 'cygwin':
return get_windows_version(run_lambda)
if platform == 'darwin':
version = get_mac_version(run_lambda)
if version is None:
return None
return 'macOS {} ({})'.format(version, machine())
if platform == 'linux':
# Ubuntu/Debian based
desc = get_lsb_version(run_lambda)
if desc is not None:
return '{} ({})'.format(desc, machine())
# Try reading /etc/*-release
desc = check_release_file(run_lambda)
if desc is not None:
return '{} ({})'.format(desc, machine())
return '{} ({})'.format(platform, machine())
# Unknown platform
return platform
def get_python_platform():
import platform
return platform.platform()
def get_libc_version():
import platform
if get_platform() != 'linux':
return 'N/A'
return '-'.join(platform.libc_ver())
def get_pip_packages(run_lambda, patterns=None):
"""Return `pip list` output. Note: will also find conda-installed pytorch and numpy packages."""
if patterns is None:
patterns = DEFAULT_PIP_PATTERNS
def run_with_pip():
try:
import importlib.util
pip_spec = importlib.util.find_spec('pip')
pip_available = pip_spec is not None
except ImportError:
pip_available = False
if pip_available:
cmd = [sys.executable, '-mpip', 'list', '--format=freeze']
elif os.environ.get("UV") is not None:
print("uv is set")
cmd = ["uv", "pip", "list", "--format=freeze"]
else:
raise RuntimeError(
"Could not collect pip list output (pip or uv module not available)"
)
out = run_and_read_all(run_lambda, cmd)
return "\n".join(line for line in out.splitlines()
if any(name in line for name in patterns))
pip_version = 'pip3' if sys.version[0] == '3' else 'pip'
out = run_with_pip()
return pip_version, out
def get_cachingallocator_config():
ca_config = os.environ.get('PYTORCH_CUDA_ALLOC_CONF', '')
return ca_config
def get_cuda_module_loading_config():
if TORCH_AVAILABLE and torch.cuda.is_available():
torch.cuda.init()
config = os.environ.get('CUDA_MODULE_LOADING', '')
return config
else:
return "N/A"
def is_xnnpack_available():
if TORCH_AVAILABLE:
import torch.backends.xnnpack
return str(torch.backends.xnnpack.enabled) # type: ignore[attr-defined]
else:
return "N/A"
def get_env_vars():
env_vars = ''
secret_terms = ('secret', 'token', 'api', 'access', 'password')
report_prefix = ("TORCH", "NCCL", "PYTORCH", "CUDA", "CUBLAS", "CUDNN",
"OMP_", "MKL_", "NVIDIA")
for k, v in os.environ.items():
if any(term in k.lower() for term in secret_terms):
continue
if k in environment_variables:
env_vars = env_vars + "{}={}".format(k, v) + "\n"
if k.startswith(report_prefix):
env_vars = env_vars + "{}={}".format(k, v) + "\n"
return env_vars
def get_env_info():
run_lambda = run
pip_version, pip_list_output = get_pip_packages(run_lambda)
if TORCH_AVAILABLE:
version_str = torch.__version__
debug_mode_str = str(torch.version.debug)
cuda_available_str = str(torch.cuda.is_available())
cuda_version_str = torch.version.cuda
if not hasattr(torch.version,
'hip') or torch.version.hip is None: # cuda version
hip_compiled_version = hip_runtime_version = miopen_runtime_version = 'N/A'
else: # HIP version
def get_version_or_na(cfg, prefix):
_lst = [s.rsplit(None, 1)[-1] for s in cfg if prefix in s]
return _lst[0] if _lst else 'N/A'
cfg = torch._C._show_config().split('\n')
hip_runtime_version = get_version_or_na(cfg, 'HIP Runtime')
miopen_runtime_version = get_version_or_na(cfg, 'MIOpen')
cuda_version_str = 'N/A'
hip_compiled_version = torch.version.hip
else:
version_str = debug_mode_str = cuda_available_str = cuda_version_str = 'N/A'
hip_compiled_version = hip_runtime_version = miopen_runtime_version = 'N/A'
sys_version = sys.version.replace("\n", " ")
conda_packages = get_conda_packages(run_lambda)
fastvideo_version = get_fastvideo_version()
fastvideo_build_flags = summarize_fastvideo_build_flags()
gpu_topo = get_gpu_topo(run_lambda)
return SystemEnv(
torch_version=version_str,
is_debug_build=debug_mode_str,
python_version='{} ({}-bit runtime)'.format(
sys_version,
sys.maxsize.bit_length() + 1),
python_platform=get_python_platform(),
is_cuda_available=cuda_available_str,
cuda_compiled_version=cuda_version_str,
cuda_runtime_version=get_running_cuda_version(run_lambda),
cuda_module_loading=get_cuda_module_loading_config(),
nvidia_gpu_models=get_gpu_info(run_lambda),
nvidia_driver_version=get_nvidia_driver_version(run_lambda),
cudnn_version=get_cudnn_version(run_lambda),
hip_compiled_version=hip_compiled_version,
hip_runtime_version=hip_runtime_version,
miopen_runtime_version=miopen_runtime_version,
pip_version=pip_version,
pip_packages=pip_list_output,
conda_packages=conda_packages,
os=get_os(run_lambda),
libc_version=get_libc_version(),
gcc_version=get_gcc_version(run_lambda),
clang_version=get_clang_version(run_lambda),
cmake_version=get_cmake_version(run_lambda),
caching_allocator_config=get_cachingallocator_config(),
is_xnnpack_available=is_xnnpack_available(),
cpu_info=get_cpu_info(run_lambda),
fastvideo_version=fastvideo_version,
fastvideo_build_flags=fastvideo_build_flags,
gpu_topo=gpu_topo,
env_vars=get_env_vars(),
)
env_info_fmt = """
PyTorch version: {torch_version}
Is debug build: {is_debug_build}
CUDA used to build PyTorch: {cuda_compiled_version}
ROCM used to build PyTorch: {hip_compiled_version}
OS: {os}
GCC version: {gcc_version}
Clang version: {clang_version}
CMake version: {cmake_version}
Libc version: {libc_version}
Python version: {python_version}
Python platform: {python_platform}
Is CUDA available: {is_cuda_available}
CUDA runtime version: {cuda_runtime_version}
CUDA_MODULE_LOADING set to: {cuda_module_loading}
GPU models and configuration: {nvidia_gpu_models}
Nvidia driver version: {nvidia_driver_version}
cuDNN version: {cudnn_version}
HIP runtime version: {hip_runtime_version}
MIOpen runtime version: {miopen_runtime_version}
Is XNNPACK available: {is_xnnpack_available}
CPU:
{cpu_info}
Versions of relevant libraries:
{pip_packages}
{conda_packages}
""".strip()
# both the above code and the following code use `strip()` to
# remove leading/trailing whitespaces, so we need to add a newline
# in between to separate the two sections
env_info_fmt += "\n"
env_info_fmt += """
FastVideo Version: {fastvideo_version}
FastVideo Build Flags:
{fastvideo_build_flags}
GPU Topology:
{gpu_topo}
{env_vars}
""".strip()
def pretty_str(envinfo):
def replace_nones(dct, replacement='Could not collect'):
for key in dct.keys():
if dct[key] is not None:
continue
dct[key] = replacement
return dct
def replace_bools(dct, true='Yes', false='No'):
for key in dct.keys():
if dct[key] is True:
dct[key] = true
elif dct[key] is False:
dct[key] = false
return dct
def prepend(text, tag='[prepend]'):
lines = text.split('\n')
updated_lines = [tag + line for line in lines]
return '\n'.join(updated_lines)
def replace_if_empty(text, replacement='No relevant packages'):
if text is not None and len(text) == 0:
return replacement
return text
def maybe_start_on_next_line(string):
# If `string` is multiline, prepend a \n to it.
if string is not None and len(string.split('\n')) > 1:
return '\n{}\n'.format(string)
return string
mutable_dict = envinfo._asdict()
# If nvidia_gpu_models is multiline, start on the next line
mutable_dict['nvidia_gpu_models'] = \
maybe_start_on_next_line(envinfo.nvidia_gpu_models)
# If the machine doesn't have CUDA, report some fields as 'No CUDA'
dynamic_cuda_fields = [
'cuda_runtime_version',
'nvidia_gpu_models',
'nvidia_driver_version',
]
all_cuda_fields = dynamic_cuda_fields + ['cudnn_version']
all_dynamic_cuda_fields_missing = all(mutable_dict[field] is None
for field in dynamic_cuda_fields)
if TORCH_AVAILABLE and not torch.cuda.is_available(
) and all_dynamic_cuda_fields_missing:
for field in all_cuda_fields:
mutable_dict[field] = 'No CUDA'
if envinfo.cuda_compiled_version is None:
mutable_dict['cuda_compiled_version'] = 'None'
# Replace True with Yes, False with No
mutable_dict = replace_bools(mutable_dict)
# Replace all None objects with 'Could not collect'
mutable_dict = replace_nones(mutable_dict)
# If either of these are '', replace with 'No relevant packages'
mutable_dict['pip_packages'] = replace_if_empty(
mutable_dict['pip_packages'])
mutable_dict['conda_packages'] = replace_if_empty(
mutable_dict['conda_packages'])
# Tag conda and pip packages with a prefix
# If they were previously None, they'll show up as ie '[conda] Could not collect'
if mutable_dict['pip_packages']:
mutable_dict['pip_packages'] = prepend(
mutable_dict['pip_packages'], '[{}] '.format(envinfo.pip_version))
if mutable_dict['conda_packages']:
mutable_dict['conda_packages'] = prepend(mutable_dict['conda_packages'],
'[conda] ')
mutable_dict['cpu_info'] = envinfo.cpu_info
return env_info_fmt.format(**mutable_dict)
def get_pretty_env_info():
return pretty_str(get_env_info())
def main():
print("Collecting environment information...")
output = get_pretty_env_info()
print(output)
if TORCH_AVAILABLE and hasattr(torch, 'utils') and hasattr(
torch.utils, '_crash_handler'):
minidump_dir = torch.utils._crash_handler.DEFAULT_MINIDUMP_DIR
if sys.platform == "linux" and os.path.exists(minidump_dir):
dumps = [
os.path.join(minidump_dir, dump)
for dump in os.listdir(minidump_dir)
]
latest = max(dumps, key=os.path.getctime)
ctime = os.path.getctime(latest)
creation_time = datetime.datetime.fromtimestamp(ctime).strftime(
'%Y-%m-%d %H:%M:%S')
msg = "\n*** Detected a minidump at {} created on {}, ".format(latest, creation_time) + \
"if this is related to your bug please include it when you file a report ***"
print(msg, file=sys.stderr)
if __name__ == '__main__':
main()
+1 -1
View File
@@ -29,7 +29,7 @@ RUN echo "# Placeholder" > README.md
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.0.post2 --no-build-isolation && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
conda clean -afy
COPY . .
+48
View File
@@ -0,0 +1,48 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
ENV DEBIAN_FRONTEND=noninteractive
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
git \
ca-certificates \
openssh-server \
&& rm -rf /var/lib/apt/lists/*
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
rm Miniconda3-latest-Linux-x86_64.sh
ENV PATH=/opt/conda/bin:$PATH
RUN conda create --name fastvideo-dev python=3.11.11 -y
SHELL ["/bin/bash", "-c"]
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
conda clean -afy
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
+48
View File
@@ -0,0 +1,48 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
ENV DEBIAN_FRONTEND=noninteractive
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
git \
ca-certificates \
openssh-server \
&& rm -rf /var/lib/apt/lists/*
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
rm Miniconda3-latest-Linux-x86_64.sh
ENV PATH=/opt/conda/bin:$PATH
RUN conda create --name fastvideo-dev python=3.12.9 -y
SHELL ["/bin/bash", "-c"]
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
conda clean -afy
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
+6 -16
View File
@@ -1,25 +1,15 @@
sphinx==6.2.1
sphinx-argparse==0.4.0
sphinx-book-theme==1.0.1
sphinx==7.4.7
sphinx-argparse==0.5.2
sphinx-autodoc2==0.5.0
sphinx-book-theme==1.1.4
sphinx-copybutton==0.5.2
sphinx-design==0.6.1
sphinx-togglebutton==0.3.2
myst-parser==3.0.1
msgspec
cloudpickle
commonmark # Required by sphinx-argparse when using :markdownhelp:
# packages to install to build the documentation
cachetools
pydantic >= 2.8
-f https://download.pytorch.org/whl/cpu
torch
py-cpuinfo
transformers
mistral_common >= 1.5.4
aiohttp
starlette
openai # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
fastapi # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
partial-json-parser # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
requests
zmq
torch
+19
View File
@@ -0,0 +1,19 @@
# Summary
## Video Generator
```{autodoc2-summary}
fastvideo.VideoGenerator
```
## Initialization Configuration
```{autodoc2-summary}
fastvideo.v1.configs.pipelines.PipelineConfig
```
## Sampling Configuration
```{autodoc2-summary}
fastvideo.v1.configs.sample.SamplingParam
```
+22
View File
@@ -0,0 +1,22 @@
# type: ignore
# SPDX-License-Identifier: Apache-2.0
from docutils import nodes
from myst_parser.parsers.sphinx_ import MystParser
from sphinx.ext.napoleon import docstring
class NapoleonParser(MystParser):
def parse(self, input_string: str, document: nodes.document) -> None:
# Get the Sphinx configuration
config = document.settings.env.config
parsed_content = str(
docstring.GoogleDocstring(
str(docstring.NumpyDocstring(input_string, config)),
config,
))
return super().parse(parsed_content, document)
Parser = NapoleonParser
+62 -44
View File
@@ -13,17 +13,19 @@
# documentation root, use os.path.abspath to make it absolute, like shown here.
import datetime
import inspect
import logging
import os
import re
import sys
from pathlib import Path
from typing import Optional
import requests
from sphinx.ext import autodoc
logger = logging.getLogger(__name__)
sys.path.append(os.path.abspath("../.."))
REPO_ROOT = Path(__file__).resolve().parent.parent.parent
print(os.path.abspath(REPO_ROOT))
sys.path.append(os.path.abspath(REPO_ROOT))
# -- Project information -----------------------------------------------------
@@ -41,8 +43,7 @@ extensions = [
"sphinx.ext.linkcode",
"sphinx.ext.intersphinx",
"sphinx_copybutton",
"sphinx.ext.autodoc",
"sphinx.ext.autosummary",
"autodoc2",
"myst_parser",
"sphinxarg.ext",
"sphinx_design",
@@ -50,6 +51,31 @@ extensions = [
]
myst_enable_extensions = [
"colon_fence",
"fieldlist",
]
autodoc2_packages = [
{
"path": "../../fastvideo",
"exclude_dirs": ["__pycache__", "third_party"],
},
]
autodoc2_output_dir = "api"
autodoc2_render_plugin = "myst"
autodoc2_hidden_objects = ["dunder", "private", "inherited"]
autodoc2_docstring_parser_regexes = [
(".*", "docs.source.autodoc2_docstring_parser"),
]
autodoc2_sort_names = True
autodoc2_index_template = None
autodoc2_skip_module_regexes = [
"fastvideo.dataset",
"fastvideo.distill",
"fastvideo.data_preprocess",
"fastvideo.models",
"fastvideo.sample",
"fastvideo.utils",
"fastvideo.distill_adv",
"fastvideo.train",
]
# Add any paths that contain templates here, relative to this directory.
@@ -78,6 +104,11 @@ html_theme_options = {
'repository_url': 'https://github.com/hao-ai-lab/FastVideo/',
'use_repository_button': True,
'use_edit_page_button': True,
# Prevents the full API being added to the left sidebar of every page.
# Reduces build time by 2.5x and reduces build size from ~225MB to ~95MB.
'collapse_navbar': True,
# Makes API visible in the right sidebar on API reference pages.
'show_toc_level': 3,
}
# Add any paths that contain custom static files (such as style sheets) here,
# relative to this directory. They are copied after the builtin static files,
@@ -160,38 +191,38 @@ def linkcode_resolve(domain, info):
return None
if not info['module']:
return None
module = info['module']
# try to determine the correct file and line number to link to
obj = sys.modules[module]
# Get path from module name
file = Path(f"{info['module'].replace('.', '/')}.py")
path = REPO_ROOT / file
if not path.exists():
path = REPO_ROOT / file.with_suffix("") / "__init__.py"
if not path.exists():
return None
# get as specific as we can
lineno: int = 0
filename: str = ""
try:
for part in info['fullname'].split('.'):
obj = getattr(obj, part)
# Get the line number of the object
with open(path) as f:
lines = f.readlines()
name = info['fullname'].split(".")[-1]
pattern = fr"^( {{4}})*((def|class) )?{name}\b.*"
for lineno, line in enumerate(lines, 1):
if not line or line.startswith("#"):
continue
if re.match(pattern, line):
break
if not (inspect.isclass(obj) or inspect.isfunction(obj)
or inspect.ismethod(obj)):
obj = obj.__class__ # type: ignore[assignment]
# If the line number is not found, return None
if lineno == len(lines):
return None
lineno = inspect.getsourcelines(obj)[1]
filename = (inspect.getsourcefile(obj)
or f"{filename}.py").split("FastVideo/", 1)[1]
except Exception:
# For some things, like a class member, won't work, so
# we'll use the line number of the parent (the class)
pass
if filename.startswith("checkouts/"):
# If the line number is found, create the URL
filename = path.relative_to(REPO_ROOT)
if "checkouts" in path.parts:
# a PR build on readthedocs
pr_number = filename.split("/")[1]
filename = filename.split("/", 2)[2]
pr_number = REPO_ROOT.name
base, branch = get_repo_base_and_branch(pr_number)
if base and branch:
return f"https://github.com/{base}/blob/{branch}/{filename}#L{lineno}"
# Otherwise, link to the source file on the main branch
return f"https://github.com/hao-ai-lab/FastVideo/blob/main/{filename}#L{lineno}"
@@ -203,6 +234,8 @@ autodoc_mock_imports = [
"cpuinfo",
"cv2",
"torch",
"huggingface_hub",
"torchvision",
"transformers",
"psutil",
"prometheus_client",
@@ -231,18 +264,6 @@ for mock_target in autodoc_mock_imports:
"been loaded into sys.modules when the sphinx build starts.",
mock_target)
class MockedClassDocumenter(autodoc.ClassDocumenter):
"""Remove note about base class when a class is derived from object."""
def add_line(self, line: str, source: str, *lineno: int) -> None:
if line == " Bases: :py:class:`object`":
return
super().add_line(line, source, *lineno)
autodoc.ClassDocumenter = MockedClassDocumenter
intersphinx_mapping = {
"python": ("https://docs.python.org/3", None),
"typing_extensions":
@@ -254,7 +275,4 @@ intersphinx_mapping = {
"psutil": ("https://psutil.readthedocs.io/en/stable", None),
}
autodoc_preserve_defaults = True
autodoc_warningiserror = True
navigation_with_keys = False
@@ -1,3 +1,4 @@
(docker)=
# 🐳 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:
+1 -1
View File
@@ -39,7 +39,7 @@ Now you can install FastVideo and setup git hooks for running linting. By using
pip install -e .[dev]
# Can also install flash-attn (optional)
pip install flash-attn==2.7.0.post2 --no-build-isolation
pip install flash-attn==2.7.4.post1 --no-build-isolation
# Linting, formatting and static type checking
pre-commit install --hook-type pre-commit --hook-type commit-msg
+1 -1
View File
@@ -9,7 +9,7 @@ from typing import Optional
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
ROOT_DIR_RELATIVE = '../../../..'
EXAMPLE_DIR = ROOT_DIR / "fastvideo/v1/examples"
EXAMPLE_DIR = ROOT_DIR / "examples"
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/source/getting_started/examples"
+59 -34
View File
@@ -4,32 +4,19 @@
FastVideo currently only supports Linux and NVIDIA CUDA GPUs.
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
## Requirements
- OS: Linux
- Python: 3.10-3.12
- CUDA 12.4+
- **OS: Linux**
- **Python: 3.10-3.12**
- **CUDA 12.4**
- **At least 1 NVIDIA GPU**
## Installation Options
## Set up using Python
### Create a new Python environment
### Option 1: Quick Install
```bash
pip install fastvideo
```
### Option 2: Installation from Source
We recommend using a Python environment such as Conda.
#### 1. [Optional] Install Miniconda (if not already installed)
#### Conda
You can create a new python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html)
##### 1. Install Miniconda (if not already installed)
```bash
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
@@ -37,38 +24,77 @@ bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
#### 2. [Optional] Create and activate a Conda environment for FastVideo
##### 2. Create and activate a Conda environment for FastVideo
```bash
conda create -n fastvideo python=3.10 -y
# (Recommended) Create a new conda environment.
conda create -n fastvideo python=3.12 -y
conda activate fastvideo
```
#### 3. Clone the FastVideo repository
:::{note}
[PyTorch has deprecated the conda release channel](https://github.com/pytorch/pytorch/issues/138506). If you use `conda`, please only use it to create Python environment rather than installing packages.
:::
#### uv
:::{tip}
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
:::
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
```console
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools` in the environment.
uv venv --python 3.12 --seed
source .venv/bin/activate
```
### Installation
```bash
pip install fastvideo
# or if you are using uv
uv pip install fastvideo
```
Also optionally install flash-attn:
```bash
pip install flash-attn==2.7.4.post1 --no-build-isolation
```
### Installation from Source
#### 1. Clone the FastVideo repository
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
```
#### 4. Install FastVideo
#### 2. Install FastVideo
Basic installation:
```bash
pip install -e .
# or if you are using uv
uv pip install -e .
```
## Optional Dependencies
### Optional Dependencies
### Flash Attention
#### Flash Attention
```bash
pip install flash-attn==2.7.0.post2 --no-build-isolation
pip install flash-attn==2.7.4.post1 --no-build-isolation
```
### 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.
## Set up using Docker
We also have prebuilt docker images with FastVideo dependencies pre-installed:
[Docker Images](#docker)
## Development Environment Setup
@@ -78,8 +104,7 @@ If you're planning to contribute to FastVideo please see the following page:
## Hardware Requirements
### For Basic Inference
- NVIDIA GPU with CUDA support
- Minimum 20GB VRAM for quantized models (e.g., single RTX 4090)
- NVIDIA GPU with CUDA 12.4 support
### For Lora Finetuning
- 40GB GPU memory each for 2 GPUs with lora
+83
View File
@@ -0,0 +1,83 @@
# V1 API
FastVideo's V1 API provides a streamlined interface for video generation tasks with powerful customization options. This page documents the primary components of the API.
## Video Generator
This class will be the primary Python API for generating videos and images.
```{autodoc2-summary}
fastvideo.VideoGenerator
```
`````{py:class} VideoGenerator(fastvideo_args: fastvideo.v1.fastvideo_args.FastVideoArgs, executor_class: type[fastvideo.v1.worker.executor.Executor], log_stats: bool)
:canonical: fastvideo.v1.entrypoints.video_generator.VideoGenerator
```{autodoc2-docstring} fastvideo.v1.entrypoints.video_generator.VideoGenerator
:parser: docs.source.autodoc2_docstring_parser
```
`VideoGenerator.from_pretrained()` should be the primary way of creating a new video generator.
````{py:method} from_pretrained(model_path: str, device: typing.Optional[str] = None, torch_dtype: typing.Optional[torch.dtype] = None, pipeline_config: typing.Optional[typing.Union[str | fastvideo.v1.configs.pipelines.PipelineConfig]] = None, **kwargs) -> fastvideo.v1.entrypoints.video_generator.VideoGenerator
:canonical: fastvideo.v1.entrypoints.video_generator.VideoGenerator.from_pretrained
:classmethod:
```{autodoc2-docstring} fastvideo.v1.entrypoints.video_generator.VideoGenerator.from_pretrained
:parser: docs.source.autodoc2_docstring_parser
```
## Configuring FastVideo
The follow two classes `PipelineConfig` and `SamplingParam` are used to configure initialization and sampling parameters, respectively.
### PipelineConfig
```{autodoc2-summary}
fastvideo.PipelineConfig
```
`````{py:class} PipelineConfig
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig
:parser: docs.source.autodoc2_docstring_parser
```
````{py:method} from_pretrained(model_path: str) -> fastvideo.v1.configs.pipelines.base.PipelineConfig
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig.from_pretrained
:classmethod:
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig.from_pretrained
:parser: docs.source.autodoc2_docstring_parser
```
````{py:method} dump_to_json(file_path: str)
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig.dump_to_json
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig.dump_to_json
:parser: docs.source.autodoc2_docstring_parser
```
### SamplingParam
```{autodoc2-summary}
fastvideo.SamplingParam
```
`````{py:class} SamplingParam
:canonical: fastvideo.v1.configs.sample.base.SamplingParam
```{autodoc2-docstring} fastvideo.v1.configs.sample.base.SamplingParam
:parser: docs.source.autodoc2_docstring_parser
```
````{py:method} from_pretrained(model_path: str) -> fastvideo.v1.configs.sample.base.SamplingParam
:canonical: fastvideo.v1.configs.sample.base.SamplingParam.from_pretrained
:classmethod:
```{autodoc2-docstring} fastvideo.v1.configs.sample.base.SamplingParam.from_pretrained
:parser: docs.source.autodoc2_docstring_parser
```
+14 -2
View File
@@ -51,14 +51,19 @@ Dev in progress and highly experimental.
:maxdepth: 1
getting_started/installation
<!-- getting_started/examples/examples_index -->
<!-- getting_started/v1_api -->
:::
:::{toctree}
:caption: Inference
:maxdepth: 1
inference/inference_quick_start
inference/configuration
inference/optimizations
inference/support_matrix
inference/examples/examples_inference_index
inference/add_pipeline
inference/v0_inference
:::
@@ -93,7 +98,14 @@ design/overview
contributing/overview
contributing/developer_env/index
contributing/add_pipeline
:::
:::{toctree}
:caption: API Reference
:maxdepth: 2
<!-- api/summary -->
api/fastvideo/fastvideo
:::
## Indices and tables
@@ -1,6 +1,6 @@
(add-pipeline)=
# 🏗️ Adding a New Diffusion Pipeline
# 🏗️ Adding a New Pipeline
This guide explains how to implement a custom diffusion pipeline in FastVideo, leveraging the framework's modular architecture for high-performance video generation.
+151
View File
@@ -0,0 +1,151 @@
# FastVideo CLI Inference
The FastVideo CLI provides a quick way to access the FastVideo inference pipeline for video generation. For more advanced usage,
see the Python interface [here](https://hao-ai-lab.github.io/FastVideo/inference/examples/basic.html).
## Basic Usage
The basic command to generate a video is:
```bash
fastvideo generate --model-path {MODEL_PATH} --prompt {PROMPT}
```
### Required Parameters
- `--model-path {MODEL_PATH}`: Path to the model or model ID
- `--prompt {PROMPT}`: Text description for the video you want to generate
## Common Arguments
To see all the options, you can use the `--help` flag:
```bash
fastvideo generate --help
```
### Hardware Configuration
- `--num-gpus {NUM_GPUS}`: Number of GPUs to use
- `--tp-size {TP_SIZE}`: Tensor parallelism size (Typically should match the number of GPUs)
- `--sp-size {SP_SIZE}`: Sequence parallelism size (Typically should match the number of GPUs)
#### Video Configuration
- `--height {HEIGHT}`: Height of the generated video
- `--width {WIDTH}`: Width of the generated video
- `--num-frames {NUM_FRAMES}`: Number of frames to generate
- `--fps {FPS}`: Frames per second for the saved video
#### Generation Parameters
- `--num-inference-steps {STEPS}`: Number of denoising steps
- `--negative-prompt {PROMPT}`: Negative prompt to guide generation away from certain concepts
- `--seed {SEED}`: Random seed for reproducible generation
#### Output Options
- `--output-path {PATH}`: Directory to save the generated video
- `--save-video`: Whether to save the video to disk
- `--return-frames`: Whether to return the raw frames
## Using Configuration Files
Instead of specifying all parameters on the command line, you can use a configuration file:
```bash
fastvideo generate --config {CONFIG_FILE_PATH}
```
The config file should be in JSON or YAML format with the same parameter names as the CLI options. Command-line arguments will take precedence over settings in the configuration file, allowing you to override specific values while keeping the rest from the config file.
Example configuration file (config.json):
```json
{
"model_path": "FastVideo/FastHunyuan-diffusers",
"prompt": "A beautiful woman in a red dress walking down a street",
"output_path": "outputs/",
"num_gpus": 2,
"sp_size": 2,
"tp_size": 2,
"num_frames": 45,
"height": 720,
"width": 1280,
"num_inference_steps": 6,
"seed": 1024,
"fps": 24,
"precision": "bf16",
"vae_precision": "fp16",
"vae_tiling": true,
"vae_sp": true,
"vae_config": {
"load_encoder": false,
"load_decoder": true,
"tile_sample_min_height": 256,
"tile_sample_min_width": 256
},
"text_encoder_precisions": [
"fp16",
"fp16"
],
"mask_strategy_file_path": null,
"enable_torch_compile": false
}
```
Or using YAML format (config.yaml):
```yaml
model_path: "FastVideo/FastHunyuan-diffusers"
prompt: "A beautiful woman in a red dress walking down a street"
output_path: "outputs/"
num_gpus: 2
sp_size: 2
tp_size: 2
num_frames: 45
height: 720
width: 1280
num_inference_steps: 6
seed: 1024
fps: 24
precision: "bf16"
vae_precision: "fp16"
vae_tiling: true
vae_sp: true
vae_config:
load_encoder: false
load_decoder: true
tile_sample_min_height: 256
tile_sample_min_width: 256
text_encoder_precisions:
- "fp16"
- "fp16"
mask_strategy_file_path: null
enable_torch_compile: false
```
## Examples
Generating a simple video:
```bash
fastvideo generate --model-path FastVideo/FastHunyuan-diffusers --prompt "A cat playing with a ball of yarn" --num-frames 45 --height 720 --width 1280 --num-inference-steps 6 --seed 1024 --output-path outputs/
```
Using a negative prompt to avoid certain elements:
```bash
fastvideo generate --model-path FastVideo/FastHunyuan-diffusers --prompt "A beautiful forest landscape" --negative-prompt "people, buildings, roads"
```
Combining command line arguments and a configuration file:
```bash
fastvideo generate --config config.json --prompt "A capybara lounging in a hammock"
```
## Troubleshooting
- If you encounter CUDA out-of-memory errors, try reducing the video dimensions or number of frames, or the number of inference steps.
- For reproducible results, set the same seed value between runs.
-33
View File
@@ -1,33 +0,0 @@
(fasthunyuan)=
# FastHunyuan
## Inference FastHunyuan on single RTX4090
We now support NF4 and LLM-INT8 quantized inference using BitsAndBytes for FastHunyuan. With NF4 quantization, inference can be performed on a single RTX 4090 GPU, requiring just 20GB of VRAM.
```bash
# Download the model weight
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan-diffusers --local_dir=data/FastHunyuan-diffusers --repo_type=model
# CLI inference
bash scripts/inference/inference_hunyuan_hf_quantization.sh
```
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
| Configuration | Memory to Init Transformer | Peak Memory After Init Pipeline (Denoise) | Diffusion Time | End-to-End Time |
|--------------------------------|----------------------------|--------------------------------------------|----------------|-----------------|
| BF16 + Pipeline CPU Offload | 23.883G | 33.744G | 81s | 121.5s |
| INT8 + Pipeline CPU Offload | 13.911G | 27.979G | 88s | 116.7s |
| NF4 + Pipeline CPU Offload | 9.453G | 19.26G | 78s | 114.5s |
For improved quality in generated videos, we recommend using a GPU with 80GB of memory to run the BF16 model with the original Hunyuan pipeline. To execute the inference, use the following section:
## FastHunyuan
```bash
# Download the model weight
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan --local_dir=data/FastHunyuan --repo_type=model
# CLI inference
bash scripts/inference/inference_hunyuan.sh
```
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
-9
View File
@@ -1,9 +0,0 @@
(fastmochi)=
# FastMochi
```bash
# Download the model weight
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastMochi-diffusers --local_dir=data/FastMochi-diffusers --repo_type=model
# CLI inference
bash scripts/inference/inference_mochi_sp.sh
-18
View File
@@ -1,18 +0,0 @@
(hunyuanvideo)=
# HunyuanVideo
## Inference HunyuanVideo with Sliding Tile Attention
First, download the model:
```bash
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
```
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
```bash
sh scripts/inference/inference_hunyuan_STA.sh
```
## Video Demos using STA + Teacache
Visit our [demo website](https://fast-video.github.io/) to explore our complete collection of examples. We shorten a single video generation process from 945s to 317s on H100.
@@ -0,0 +1,129 @@
# Inference Quick Start
This page contains step-by-step instructions to get you quickly started with video generation using FastVideo.
## Table of Contents
- [Generating Your First Video](#generating-your-first-video)
- [Customizing Generation](#customizing-generation)
- [Available Models](#available-models)
- [Image-to-Video Generation](#image-to-video-generation)
- [Troubleshooting](#troubleshooting)
- [Advanced Configuration](#advanced-configuration)
- [Next Steps](#next-steps)
## Software Requirements
- **OS**: Linux (Tested on Ubuntu 22.04+)
- **Python**: 3.10-3.12
- **CUDA**: 12.4
## Installation
We recommend using an environment manager such as `Conda` to create a clean environment:
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
# Install FastVideo
pip install fastvideo
```
For advanced installation options, see the [Installation Guide](installation.md).
## Generating Your First Video
Here's a minimal example to generate a video using the default settings. Create a file called `example.py` with the following code:
```python
from fastvideo import VideoGenerator
def main():
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1, # Adjust based on your hardware
)
# Define a prompt for your video
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
# Generate the video
video = generator.generate_video(
prompt,
return_frames=True, # Also return frames from this call (defaults to False)
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
if __name__ == '__main__':
main()
```
Run the script with:
```bash
python example.py
```
The generated video will be saved in the current directory under `my_videos/`.
## Available Models
Please see the [support matrix](#support-matrix) for the list of supported models and their available optimizations.
## Image-to-Video Generation
You can generate a video starting from an initial image:
```python
from fastvideo import VideoGenerator, SamplingParam
# Create the generator
generator = VideoGenerator.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
# Set up parameters with an initial image
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
sampling_param.image_path = "path/to/your/image.jpg"
sampling_param.num_frames = 24
sampling_param.image_strength = 0.8 # How much to preserve the original image (0-1)
# Generate video based on the image
prompt = "A photograph coming to life with gentle movement"
video = generator.generate_video(prompt, sampling_param=sampling_param)
```
## Troubleshooting
Common issues and their solutions:
### Out of Memory Errors
If you encounter CUDA out of memory errors:
- Reduce `num_frames` or video resolution
- Enable memory optimization with `enable_model_cpu_offload`
- Try a smaller model or use quantized versions
- Use `num_gpus` > 1 if multiple GPUs are available
### Slow Generation
To speed up generation:
- Reduce `num_inference_steps` (20-30 is usually sufficient)
- Use half precision (`fp16`) for the VAE
- Use multiple GPUs if available
### Unexpected Results
If the generated video doesn't match your prompt:
- Try increasing `guidance_scale` (7.0-9.0 works well)
- Make your prompt more detailed and specific
- Experiment with different random seeds
- Try a different model
## Advanced Configuration
## Optimizations
## Next Steps
- Explore the [API Reference](../api/index.md) for detailed documentation
- Learn about [Advanced Inference Options](../inference/overview_back.md)
- See [Examples](../examples/index.md) for more usage scenarios
- Check out the [Model Training](../training/overview.md) guide to fine-tune models
- Join our [Community Discord](https://discord.gg/fastvideo) for support and sharing
+147
View File
@@ -0,0 +1,147 @@
# Optimizations
This page describes the various options for speeding up generation times in FastVideo.
## Table of Contents
- Optimized Attention Backends
- [Flash Attention](#optimizations-flash)
- [Sliding Tile Attention](#optimizations-sta)
- [Sage Attention](#optimizations-sage)
- Caching Techniques
- [TeaCache](#optimizations-teacache)
(optimizations-backends)=
## Attention Backends
### Available Backends
- Torch SDPA: `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`
- Flash Attention 2 and 3: `FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN`
- Sliding Tile Attention: `FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN`
- Sage Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN`
### Configuring Backends
There are two ways to configure the attention backend in FastVideo.
#### 1. In Python
In python, set the `FASTVIDEO_ATTENTION_BACKEND` environment variable before instantiating `VideoGenerator` like this:
```python
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLIDING_TILE_ATTN"
```
#### 2. In CLI
You can also set the environment variable on the command line:
```bash
FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN python example.py
```
(optimizations-flash)=
### Flash Attention
**`FLASH_ATTN`**
We recommend always installing [Flash Attention 2](https://github.com/Dao-AILab/flash-attention):
```bash
pip install flash-attn==2.7.4.post1 --no-build-isolation
```
And if using a Hopper+ GPU (ie H100), installing [Flash Attention 3](https://github.com/Dao-AILab/flash-attention?tab=readme-ov-file#flashattention-3-beta-release) by compiling it from source (takes about 10 minutes for me):
```bash
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention
cd hopper
pip install ninja
python setup.py install
```
:::{note}
FastVideo will automatically detect and use `FA3` if it is installed when using `FLASH_ATTN` backend.
:::
(optimizations-sta)=
### Sliding Tile Attention
**`SLIDING_TILE_ATTN`**
```bash
pip install st_attn==0.0.4
```
Please see [this page](#sta-installation) for more installation instructions.
(optimizations-sage)=
### Sage Attention
**`SAGE_ATTN`**
To use [SageAttention](https://github.com/thu-ml/SageAttention) 2.1.1, please compile from source:
```bash
git clone https://github.com/thu-ml/SageAttention.git
cd sageattention
python setup.py install # or pip install -e .
```
(optimizations-teacache)=
## Teacache
TeaCache is an optimization technique supported in FastVideo that can significantly speed up video generation by skipping redundant calculations across diffusion steps. This guide explains how to enable and configure TeaCache for optimal performance in FastVideo.
### What is TeaCache?
See the official [TeaCache](https://github.com/ali-vilab/TeaCache) repo and their paper for more details.
### How to Enable TeaCache
Enabling TeaCache is straightforward - simply add the `enable_teacache=True` parameter to your `generate_video()` call:
```python
# ... previous code
generator.generate_video(
prompt="Your prompt here",
sampling_param=params,
enable_teacache=True
)
# more code ...
```
### Complete Example
At the bottom is a complete example of using TeaCache for faster video generation. You can run it using the following command:
```bash
python examples/inference/optimizations/teacache_example.py
```
### Advanced Configuration
While TeaCache works well with default settings, you can fine-tune its behavior by adjusting the threshold value:
1. Lower threshold values (e.g., 0.1) will result in more skipped calculations and faster generation with slightly more potential for quality degradation
2. Higher threshold values (e.g., 0.15-0.23) will skip fewer calculations but maintain quality closer to the original
Note that the optimal threshold depends on your specific model and content.
## Benchmarking different optimizations
To benchmark the performance improvement, try generating the same video with and without TeaCache enabled and compare the generation times:
```python
# Without TeaCache
start_time = time.perf_counter()
generator.generate_video(prompt="Your prompt", enable_teacache=False)
standard_time = time.perf_counter() - start_time
# With TeaCache
start_time = time.perf_counter()
generator.generate_video(prompt="Your prompt", enable_teacache=True)
teacache_time = time.perf_counter() - start_time
print(f"Standard generation: {standard_time:.2f} seconds")
print(f"TeaCache generation: {teacache_time:.2f} seconds")
print(f"Speedup: {standard_time/teacache_time:.2f}x")
```
Note: If you want to benchmark different attention backends, you'll need to reinstantiate `VideoGenerator`.
-16
View File
@@ -1,16 +0,0 @@
(stepvideo)=
# StepVideo
## Inference StepVideo with Sliding Tile Attention
First, download the model:
```
python scripts/huggingface/download_hf.py --repo_id=stepfun-ai/stepvideo-t2v --local_dir=data/stepvideo-t2v --repo_type=model
```
Use the following scripts to run inference for StepVideo. When using STA for inference, the generated videos will have dimensions of 204×768×768 (currently, this is the only supported shape).
```bash
sh scripts/inference/inference_stepvideo_STA.sh # Inference stepvideo with STA
sh scripts/inference/inference_stepvideo.sh # Inference original stepvideo
```
+92
View File
@@ -0,0 +1,92 @@
(support-matrix)=
# Compatibility Matrix
The table below shows every supported model and optimizations supported for them.
The symbols used have the following meanings:
- ✅ = Full compatibility
- ❌ = No compatibility
## Models x Optimization
The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods and FastVideo will use the optimal default parameters when initializing and generating videos.
:::{raw} html
<style>
/* Make smaller to try to improve readability */
td {
font-size: 0.9rem;
text-align: center;
}
th {
text-align: center;
font-size: 0.9rem;
}
</style>
:::
:::{list-table}
:header-rows: 1
:stub-columns: 3
:widths: auto
:class: vertical-table-header
- * Model Name
* HuggingFace Model ID
* Resolutions
* TeaCache
* Sliding Tile Attn
* Sage Attn
- * HunyuanVideo
* `hunyuanvideo-community/HunyuanVideo`
* 720px1280p<br>544px960p
* ❌
* ✅
* ✅
- * FastHunyuan
* `FastVideo/FastHunyuan-diffusers`
* 720px1280p<br>544px960p
* ❌
* ✅
* ✅
- * Wan T2V 1.4B
* `Wan-AI/Wan2.1-T2V-1.3B-Diffusers`
* 480P
* ✅
* ✅*
* ✅
- * Wan T2V 14B
* `Wan-AI/Wan2.1-T2V-1.3B-Diffusers`
* 480P, 720P
* ✅
* ✅*
* ✅
- * Wan I2V 480P
* `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers`
* 480P
* ✅
* ✅*
* ✅
- * Wan T2V 720P
* `Wan-AI/Wan2.1-T2V-14B-Diffusers`
* 720P
* ✅
* ✅*
* ✅
- * StepVideo T2V
* Coming soon!
* 768px768px204f<br>544px992px204f<br>544px992px136f
*
*
*
:::
**Note**: there are some known quality issues with Wan2.1 + Sliding Tile Attn. We are working on fixing this issue.
## Special requirements
### StepVideo T2V
- The self-attention in text-encoder (step_llm) only supports CUDA capabilities sm_80 sm_86 and sm_90
### Sliding Tile Attention
- Currently only Hopper GPUs (H100s) are supported.
-44
View File
@@ -1,44 +0,0 @@
(wanvideo)=
# WanVideo
## Inference T2V with WanVideo
First, download the model:
```bash
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-T2V-1.3B-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
```
or
```bash
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-T2V-14B-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
```
Then run the inference using:
```bash
sh scripts/inference/v1_inference_wan.sh
```
Remember to set `MODEL_BASE` and `num_gpus` accordingly.
## Inference I2V with WanVideo
First, download the model:
```bash
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-I2V-14B-480P-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
```
or
```bash
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-I2V-14B-720P-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
```
Then run the inference using:
```bash
sh scripts/inference/v1_inference_wan_i2v.sh
```
Remember to set `MODEL_BASE` and `num_gpus` accordingly.
+40 -2
View File
@@ -1,3 +1,41 @@
# Basic
# 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.
The class provides the main python interface for using FastVideo's inference pipeline.
## Requirements
- At least a single NVIDIA GPU with CUDA 12.4.
- Python 3.10-3.12
## Installation
If you have not installed FastVideo, please following these [instructions](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) first.
## 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
# if you have not cloned the directory:
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
python 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
def main():
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1,
)
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
video = generator.generate_video(prompt)
if __name__ == "__main__":
main()
```
+41 -1
View File
@@ -1 +1,41 @@
print('Hello, world!')
from fastvideo import VideoGenerator
# from fastvideo.v1.configs.sample import SamplingParam
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
# if num_gpus > 1, FastVideo will automatically handle distributed setup
num_gpus=1,
)
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
video2 = generator.generate_video(prompt2)
if __name__ == "__main__":
main()
+62
View File
@@ -0,0 +1,62 @@
from fastvideo import VideoGenerator
def main():
# This is the config class for the model initialization
config = PipelineConfig.from_pretrained("FastVideo/FastHunyuan-Diffusers")
# can be used to dump the config to a yaml file
config.dump_to_yaml("config.yaml")
print(config)
# {
# 'vae_config': {
# 'scale_factor': 8,
# 'sp': True,
# 'tiling': True,
# 'precision': 'fp16'
# },
# 'text_encoder_config': {
# 'precision': 'fp16'
# },
# 'dit_config': {
# 'precision': 'fp16'
# },
# 'inference_args': {
# 'guidance_scale': 7.5,
# 'num_inference_steps': 5,
# 'seed': 1024,
# 'guidance_rescale': 0.0,
# 'flow_shift': 17,
# 'num_inference_steps': 5,
# }
# }
config.vae_config.scale_factor = 16
# FastVideo will automatically used 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",
num_gpus=4,
config=config,
# or
config_path="config.yaml",
)
sampling_param = SamplingParam.from_pretrained(
"FastVideo/FastHunyuan-Diffusers")
sampling_param.num_inference_steps = 5
# 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,
sampling_param=sampling_param,
num_inference_steps=6)
video2 = generator.generate_video(prompt2)
prompt2 = "A beautiful woman in a blue dress walking down a street"
if __name__ == "__main__":
main()
@@ -1,4 +1,4 @@
# FastVideo VideoGenerator Gradio Demo
# FastVideo 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.
@@ -13,25 +13,12 @@ The demo uses the FastVideo framework to generate videos based on text prompts.
---
## 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
python 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.
@@ -40,15 +27,6 @@ This will start a web server at `http://0.0.0.0:7860` where you can access the i
## 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
@@ -78,4 +56,4 @@ The interface is built with several components:
- **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
- **Seed**: Control randomness for reproducible results
+169
View File
@@ -0,0 +1,169 @@
import argparse
import os
from copy import deepcopy
import gradio as gr
import torch
from fastvideo import VideoGenerator
from fastvideo.v1.configs.sample.base import SamplingParam
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="FastVideo Gradio Demo")
parser.add_argument("--model_path",
type=str,
default="FastVideo/FastHunyuan-diffusers",
help="Path to the model")
parser.add_argument("--num_gpus",
type=int,
default=1,
help="Number of GPUs to use")
parser.add_argument("--output_path",
type=str,
default="outputs",
help="Path to save generated videos")
parsed_args = parser.parse_args()
# args = FastVideoArgs(model_path="FastVideo/FastHunyuan-Diffusers", num_gpus=2)
generator = VideoGenerator.from_pretrained(
model_path=parsed_args.model_path, num_gpus=parsed_args.num_gpus)
default_params = SamplingParam.from_pretrained(parsed_args.model_path)
def generate_video(
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed=False,
):
params = deepcopy(default_params)
params.prompt = prompt
params.negative_prompt = negative_prompt
params.seed = seed
params.guidance_scale = guidance_scale
params.num_frames = num_frames
params.height = height
params.width = width
params.num_inference_steps = num_inference_steps
if randomize_seed:
params.seed = torch.randint(0, 1000000, (1, )).item()
if not use_negative_prompt:
params.negative_prompt = None
generator.generate_video(prompt=prompt, sampling_param=params)
output_path = os.path.join(parsed_args.output_path,
f"{params.prompt[:100]}.mp4")
return output_path, params.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=default_params.height,
)
width = gr.Slider(label="Width",
minimum=256,
maximum=1024,
step=32,
value=default_params.width)
with gr.Row():
num_frames = gr.Slider(
label="Number of Frames",
minimum=21,
maximum=163,
value=default_params.num_frames,
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=default_params.guidance_scale,
)
num_inference_steps = gr.Slider(
label="Inference Steps",
minimum=4,
maximum=100,
value=default_params.num_inference_steps,
)
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=default_params.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=default_params.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)
@@ -0,0 +1,9 @@
# Optimization Examples
```bash
python examples/inference/optimizations/attention_example.py
```
```bash
python examples/inference/optimizations/teacache_example.py
```
@@ -0,0 +1,33 @@
import os
import time
from fastvideo import VideoGenerator
def main():
# set the attention backend
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
start_time = time.perf_counter()
gen = VideoGenerator.from_pretrained(
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1,
)
load_time = time.perf_counter() - start_time
print(f"Model loading time: {load_time:.2f} seconds")
gen_start_time = time.perf_counter()
gen.generate_video(
prompt=
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.",
seed=1024,
output_path="example_outputs/")
generation_time = time.perf_counter() - gen_start_time
print(f"Video generation time: {generation_time:.2f} seconds")
total_time = time.perf_counter() - start_time
print(f"Total execution time: {total_time:.2f} seconds")
if __name__ == "__main__":
main()
@@ -0,0 +1,44 @@
import time
from fastvideo import VideoGenerator, SamplingParam
def main():
start_time = time.perf_counter()
gen = VideoGenerator.from_pretrained(
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1,
use_cpu_offload=False,
)
load_time = time.perf_counter() - start_time
print(f"Model loading time: {load_time:.2f} seconds")
gen_start_time = time.perf_counter()
params = SamplingParam.from_pretrained(
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
)
# this controls the threshold for the tea cache
params.teacache_params.teacache_thresh = 0.08
gen.generate_video(
prompt=
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.",
sampling_param=params,
height=480,
width=832,
num_frames=61, # 85 ,77
num_inference_steps=50,
enable_teacache=True,
seed=1024,
output_path="example_outputs/")
generation_time = time.perf_counter() - gen_start_time
print(f"Video generation time: {generation_time:.2f} seconds")
total_time = time.perf_counter() - start_time
print(f"Total execution time: {total_time:.2f} seconds")
if __name__ == "__main__":
main()
+3 -1
View File
@@ -1,3 +1,5 @@
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
__all__ = ["VideoGenerator"]
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam"]
+2 -2
View File
@@ -452,7 +452,7 @@ def main(args):
return phase
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
start_time = time.perf_counter()
assert args.multi_phased_distill_schedule is not None
num_phases = get_num_phases(args.multi_phased_distill_schedule, step)
@@ -482,7 +482,7 @@ def main(args):
args.hunyuan_teacher_disable_cfg,
)
step_time = time.time() - start_time
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
+2 -2
View File
@@ -517,7 +517,7 @@ def main(args):
for step in range(init_steps + 1, args.max_train_steps + 1):
assert args.multi_phased_distill_schedule is not None
num_phases = get_num_phases(args.multi_phased_distill_schedule, step)
start_time = time.time()
start_time = time.perf_counter()
(
generator_loss,
generator_grad_norm,
@@ -547,7 +547,7 @@ def main(args):
args.discriminator_head_stride,
)
step_time = time.time() - start_time
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
+2 -2
View File
@@ -452,7 +452,7 @@ class HunyuanVideoSampler(Inference):
# ========================================================================
# Pipeline inference
# ========================================================================
start_time = time.time()
start_time = time.perf_counter()
samples = self.pipeline(
prompt=prompt,
height=target_height,
@@ -476,7 +476,7 @@ class HunyuanVideoSampler(Inference):
out_dict["samples"] = samples
out_dict["prompts"] = prompt
gen_time = time.time() - start_time
gen_time = time.perf_counter() - start_time
logger.info(f"Success, time: {gen_time}")
return out_dict
+2 -2
View File
@@ -364,7 +364,7 @@ def main(args):
for i in range(init_steps):
next(loader)
for step in range(init_steps + 1, args.max_train_steps + 1):
start_time = time.time()
start_time = time.perf_counter()
loss, grad_norm = train_one_step(
transformer,
args.model_type,
@@ -383,7 +383,7 @@ def main(args):
args.mode_scale,
)
step_time = time.time() - start_time
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
@@ -56,20 +56,6 @@ class AttentionMetadata:
# Current step of diffusion process
current_timestep: int
# @property
# @abstractmethod
# def inference_metadata(self) -> Optional["AttentionMetadata"]:
# """Return the attention metadata that's required to run prefill
# attention."""
# pass
# @property
# @abstractmethod
# def training_metadata(self) -> Optional["AttentionMetadata"]:
# """Return the attention metadata that's required to run decode
# attention."""
# pass
def asdict_zerocopy(self,
skip_fields: Optional[Set[str]] = None
) -> Dict[str, Any]:
@@ -86,55 +72,6 @@ class AttentionMetadata:
T = TypeVar("T", bound=AttentionMetadata)
# class AttentionState(ABC, Generic[T]):
# """Holds attention backend-specific objects reused during the
# lifetime of the model runner."""
# @abstractmethod
# def __init__(self, runner: "ModelRunnerBase"):
# ...
# @abstractmethod
# @contextmanager
# def graph_capture(self, max_batch_size: int):
# """Context manager used when capturing CUDA graphs."""
# yield
# @abstractmethod
# def graph_clone(self, batch_size: int) -> "AttentionState[T]":
# """Clone attention state to save in CUDA graph metadata."""
# ...
# @abstractmethod
# def graph_capture_get_metadata_for_batch(
# self,
# batch_size: int,
# is_encoder_decoder_model: bool = False) -> T:
# """Get attention metadata for CUDA graph capture of batch_size."""
# ...
# @abstractmethod
# def get_graph_input_buffers(
# self,
# attn_metadata: T,
# is_encoder_decoder_model: bool = False) -> Dict[str, Any]:
# """Get attention-specific input buffers for CUDA graph capture."""
# ...
# @abstractmethod
# def prepare_graph_input_buffers(
# self,
# input_buffers: Dict[str, Any],
# attn_metadata: T,
# is_encoder_decoder_model: bool = False) -> None:
# """In-place modify input buffers dict for CUDA graph replay."""
# ...
# @abstractmethod
# def begin_forward(self, model_input: "ModelRunnerInputBase") -> None:
# """Prepare state for forward pass."""
# ...
class AttentionMetadataBuilder(ABC, Generic[T]):
"""Abstract class for attention metadata builders."""
-9
View File
@@ -1,9 +0,0 @@
from fastvideo.v1.configs.base import BaseConfig, SlidingTileAttnConfig
from fastvideo.v1.configs.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.v1.configs.registry import get_pipeline_config_cls_for_name
from fastvideo.v1.configs.wan import WanI2V480PConfig, WanT2V480PConfig
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "get_pipeline_config_cls_for_name"
]
-67
View File
@@ -1,67 +0,0 @@
from dataclasses import dataclass
from typing import Optional
@dataclass
class BaseConfig:
"""Base configuration for all pipeline architectures."""
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 125
fps: int = 24
# Video generation parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
seed: int = 1024
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
vae_scale_factor: Optional[int] = None
# DiT configuration
num_channels_latents: Optional[int] = None
# Image encoder configuration
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_len: int = -1
hidden_state_skip_layer: int = 0
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
neg_prompt: Optional[str] = None
@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
+48
View File
@@ -0,0 +1,48 @@
{
"embedded_cfg_scale": 6,
"flow_shift": 17,
"use_cpu_offload": false,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_tiling": true,
"vae_sp": true,
"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": 4,
"use_tiling": true,
"use_temporal_tiling": true,
"use_parallel_tiling": true
},
"dit_config": {
"prefix": "Hunyuan",
"quant_config": null
},
"text_encoder_precisions": [
"fp16",
"fp16"
],
"text_encoder_configs": [
{
"prefix": "llama",
"quant_config": null,
"lora_config": null
},
{
"prefix": "clip",
"quant_config": null,
"lora_config": null,
"num_hidden_layers_override": null,
"require_post_norm": null
}
],
"mask_strategy_file_path": null,
"enable_torch_compile": false
}
-40
View File
@@ -1,40 +0,0 @@
from dataclasses import dataclass
from fastvideo.v1.configs.base import BaseConfig
@dataclass
class HunyuanConfig(BaseConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
# 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
@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
+6
View File
@@ -0,0 +1,6 @@
from fastvideo.v1.configs.models.base import ModelConfig
from fastvideo.v1.configs.models.dits.base import DiTConfig
from fastvideo.v1.configs.models.encoders.base import EncoderConfig
from fastvideo.v1.configs.models.vaes.base import VAEConfig
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
+72
View File
@@ -0,0 +1,72 @@
from dataclasses import dataclass, field, fields
from typing import Any, Dict
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
# 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 & overridden 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 = field(default_factory=ArchConfig)
# FastVideo-specific parameters here
# i.e. STA, quantization, teacache
def __getattr__(self, name):
# Only called if 'name' is not found in ModelConfig directly
if hasattr(self.arch_config, name):
return getattr(self.arch_config, name)
raise AttributeError(
f"'{type(self).__name__}' object has no attribute '{name}'")
def __getstate__(self):
# Return a dictionary of attributes to pickle
# Convert to dict and exclude any problematic attributes
state = self.__dict__.copy()
return state
def __setstate__(self, state):
# Restore instance attributes from the unpickled state
self.__dict__.update(state)
# 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}'")
if hasattr(arch_config, "__post_init__"):
arch_config.__post_init__()
def update_model_config(self, source_model_dict: Dict[str, Any]) -> None:
assert "arch_config" not in source_model_dict, "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:
logger.warning("%s does not contain field '%s'!",
type(self).__name__, key)
raise AttributeError(f"Invalid field: {key}")
if hasattr(self, "__post_init__"):
self.__post_init__()
@@ -0,0 +1,5 @@
from fastvideo.v1.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.v1.configs.models.dits.stepvideo import StepVideoConfig
from fastvideo.v1.configs.models.dits.wanvideo import WanVideoConfig
__all__ = ["HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig"]
+56
View File
@@ -0,0 +1,56 @@
from dataclasses import dataclass, field
from typing import Any, Optional, Tuple
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
@dataclass
class DiTArchConfig(ArchConfig):
_fsdp_shard_conditions: list = field(default_factory=list)
_compile_conditions: list = field(default_factory=list)
_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
hidden_size: int = 0
num_attention_heads: int = 0
num_channels_latents: int = 0
def __post_init__(self) -> None:
if not self._compile_conditions:
self._compile_conditions = self._fsdp_shard_conditions.copy()
@dataclass
class DiTConfig(ModelConfig):
arch_config: DiTArchConfig = field(default_factory=DiTArchConfig)
# FastVideoDiT-specific parameters
prefix: str = ""
quant_config: Optional[QuantizationConfig] = None
@staticmethod
def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any:
"""Add CLI arguments for DiTConfig fields"""
parser.add_argument(
f"--{prefix}.prefix",
type=str,
dest=f"{prefix.replace('-', '_')}.prefix",
default=DiTConfig.prefix,
help="Prefix for the DiT model",
)
parser.add_argument(
f"--{prefix}.quant-config",
type=str,
dest=f"{prefix.replace('-', '_')}.quant_config",
default=None,
help="Quantization configuration for the DiT model",
)
return parser
@@ -0,0 +1,177 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
import torch
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double" in n and str.isdigit(n.split(".")[-1])
def is_single_block(n: str, m) -> bool:
return "single" in n and str.isdigit(n.split(".")[-1])
def is_refiner_block(n: str, m) -> bool:
return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
return n.split(".")[-1] == "txt_in"
@dataclass
class HunyuanVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(
default_factory=lambda:
[is_double_block, is_single_block, is_refiner_block])
_compile_conditions: list = field(
default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
_param_names_mapping: dict = field(
default_factory=lambda: {
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"txt_in.t_embedder.mlp.fc_out.\1",
r"^context_embedder\.proj_in\.(.*)$":
r"txt_in.input_embedder.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"txt_in.c_embedder.fc_in.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"txt_in.c_embedder.fc_out.\1",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$":
r"txt_in.refiner_blocks.\1.norm1.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$":
r"txt_in.refiner_blocks.\1.norm2.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 0, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 1, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 2, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
# 3. x_embedder mapping:
r"^x_embedder\.proj\.(.*)$":
r"img_in.proj.\1",
# 4. Top-level time_text_embed mappings:
r"^time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"time_in.mlp.fc_in.\1",
r"^time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"time_in.mlp.fc_out.\1",
r"^time_text_embed\.guidance_embedder\.linear_1\.(.*)$":
r"guidance_in.mlp.fc_in.\1",
r"^time_text_embed\.guidance_embedder\.linear_2\.(.*)$":
r"guidance_in.mlp.fc_out.\1",
r"^time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"vector_in.fc_in.\1",
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"vector_in.fc_out.\1",
# 5. transformer_blocks mapping:
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
r"double_blocks.\1.img_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
r"double_blocks.\1.txt_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"double_blocks.\1.img_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"double_blocks.\1.img_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"double_blocks.\1.img_attn_proj.\2",
# Corrected: merge attn.to_add_out into the main projection.
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
r"double_blocks.\1.txt_attn_proj.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
r"double_blocks.\1.txt_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
r"double_blocks.\1.txt_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_out.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_out.\2",
# 6. single_transformer_blocks mapping:
r"^single_transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"single_blocks.\1.q_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"single_blocks.\1.k_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"single_blocks.\1.linear1.\2", 0, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"single_blocks.\1.linear1.\2", 1, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"single_blocks.\1.linear1.\2", 2, 4),
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$":
(r"single_blocks.\1.linear1.\2", 3, 4),
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$":
r"single_blocks.\1.linear2.\2",
r"^single_transformer_blocks\.(\d+)\.norm\.linear\.(.*)$":
r"single_blocks.\1.modulation.linear.\2",
# 7. Final layers mapping:
r"^norm_out\.linear\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
r"^proj_out\.(.*)$":
r"final_layer.linear.\1",
})
patch_size: int = 2
patch_size_t: int = 1
in_channels: int = 16
out_channels: int = 16
num_attention_heads: int = 24
attention_head_dim: int = 128
mlp_ratio: float = 4.0
num_layers: int = 20
num_single_layers: int = 40
num_refiner_layers: int = 2
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56)
guidance_embeds: bool = False
dtype: Optional[torch.dtype] = None
text_embed_dim: int = 4096
pooled_projection_dim: int = 768
rope_theta: int = 256
qk_norm: str = "rms_norm"
def __post_init__(self):
super().__post_init__()
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
self.num_channels_latents: int = self.in_channels
@dataclass
class HunyuanVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=HunyuanVideoArchConfig)
prefix: str = "Hunyuan"
@@ -0,0 +1,65 @@
from dataclasses import dataclass, field
from typing import List, Optional, Tuple, Union
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class StepVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
_param_names_mapping: dict = field(
default_factory=lambda: {
# transformer block
r"^transformer_blocks\.(\d+)\.norm1\.(weight|bias)$":
r"transformer_blocks.\1.norm1.norm.\2",
r"^transformer_blocks\.(\d+)\.norm2\.(weight|bias)$":
r"transformer_blocks.\1.norm2.norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.weight$":
r"transformer_blocks.\1.ff.fc_in.weight",
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.weight$":
r"transformer_blocks.\1.ff.fc_out.weight",
# adanorm block
r"^adaln_single\.emb\.timestep_embedder\.linear_1\.(weight|bias)$":
r"adaln_single.emb.mlp.fc_in.\1",
r"^adaln_single\.emb\.timestep_embedder\.linear_2\.(weight|bias)$":
r"adaln_single.emb.mlp.fc_out.\1",
# caption projection
r"^caption_projection\.linear_1\.(weight|bias)$":
r"caption_projection.fc_in.\1",
r"^caption_projection\.linear_2\.(weight|bias)$":
r"caption_projection.fc_out.\1",
})
num_attention_heads: int = 48
attention_head_dim: int = 128
in_channels: int = 64
out_channels: Optional[int] = 64
num_layers: int = 48
dropout: float = 0.0
patch_size: int = 1
norm_type: str = "ada_norm_single"
norm_elementwise_affine: bool = False
norm_eps: float = 1e-6
caption_channels: Optional[Union[int, List[int], Tuple[int, ...]]] = field(
default_factory=lambda: [6144, 1024])
attention_type: Optional[str] = "torch"
use_additional_conditions: Optional[bool] = False
def __post_init__(self):
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.out_channels = self.in_channels if self.out_channels is None else self.out_channels
self.num_channels_latents = self.out_channels
@dataclass
class StepVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=StepVideoArchConfig)
prefix: str = "StepVideo"
@@ -0,0 +1,83 @@
from dataclasses import dataclass, field
from typing import Optional, Tuple
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class WanVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
_param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$":
r"patch_embedding.proj.\1",
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$":
r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_out.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
r"blocks.\1.ffn.fc_out.\2",
r"blocks\.(\d+)\.norm2\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
})
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
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.in_channels if self.added_kv_proj_dim is None else self.out_channels
@dataclass
class WanVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=WanVideoArchConfig)
prefix: str = "Wan"
@@ -0,0 +1,14 @@
from fastvideo.v1.configs.models.encoders.base import (BaseEncoderOutput,
EncoderConfig,
ImageEncoderConfig,
TextEncoderConfig)
from fastvideo.v1.configs.models.encoders.clip import (CLIPTextConfig,
CLIPVisionConfig)
from fastvideo.v1.configs.models.encoders.llama import LlamaConfig
from fastvideo.v1.configs.models.encoders.t5 import T5Config
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig", "LlamaConfig",
"T5Config"
]
@@ -0,0 +1,75 @@
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Tuple
import torch
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
@dataclass
class EncoderArchConfig(ArchConfig):
architectures: List[str] = field(default_factory=lambda: [])
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
output_hidden_states: bool = False
use_return_dict: bool = True
@dataclass
class TextEncoderArchConfig(EncoderArchConfig):
vocab_size: int = 0
hidden_size: int = 0
num_hidden_layers: int = 0
num_attention_heads: int = 0
pad_token_id: int = 0
eos_token_id: int = 0
text_len: int = 0
hidden_state_skip_layer: int = 0
decoder_start_token_id: int = 0
output_past: bool = True
scalable_attention: bool = True
tie_word_embeddings: bool = False
tokenizer_kwargs: Dict[str, Any] = field(default_factory=dict)
def __post_init__(self) -> None:
self.tokenizer_kwargs = {
"truncation": True,
"max_length": self.text_len,
"return_tensors": "pt",
}
@dataclass
class ImageEncoderArchConfig(EncoderArchConfig):
pass
@dataclass
class BaseEncoderOutput:
last_hidden_state: Optional[torch.FloatTensor] = None
pooler_output: Optional[torch.FloatTensor] = None
hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
attentions: Optional[Tuple[torch.FloatTensor, ...]] = None
attention_mask: Optional[torch.Tensor] = None
@dataclass
class EncoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=EncoderArchConfig)
prefix: str = ""
quant_config: Optional[QuantizationConfig] = None
lora_config: Optional[Any] = None
@dataclass
class TextEncoderConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
@dataclass
class ImageEncoderConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=ImageEncoderArchConfig)
@@ -0,0 +1,66 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
ImageEncoderConfig,
TextEncoderArchConfig,
TextEncoderConfig)
@dataclass
class CLIPTextArchConfig(TextEncoderArchConfig):
vocab_size: int = 49408
hidden_size: int = 512
intermediate_size: int = 2048
projection_dim: int = 512
num_hidden_layers: int = 12
num_attention_heads: int = 8
max_position_embeddings: int = 77
hidden_act: str = "quick_gelu"
layer_norm_eps: float = 1e-5
dropout: float = 0.0
attention_dropout: float = 0.0
initializer_range: float = 0.02
initializer_factor: float = 1.0
pad_token_id: int = 1
bos_token_id: int = 49406
eos_token_id: int = 49407
text_len: int = 77
@dataclass
class CLIPVisionArchConfig(ImageEncoderArchConfig):
hidden_size: int = 768
intermediate_size: int = 3072
projection_dim: int = 512
num_hidden_layers: int = 12
num_attention_heads: int = 12
num_channels: int = 3
image_size: int = 224
patch_size: int = 32
hidden_act: str = "quick_gelu"
layer_norm_eps: float = 1e-5
dropout: float = 0.0
attention_dropout: float = 0.0
initializer_range: float = 0.02
initializer_factor: float = 1.0
@dataclass
class CLIPTextConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(
default_factory=CLIPTextArchConfig)
num_hidden_layers_override: Optional[int] = None
require_post_norm: Optional[bool] = None
prefix: str = "clip"
@dataclass
class CLIPVisionConfig(ImageEncoderConfig):
arch_config: ImageEncoderArchConfig = field(
default_factory=CLIPVisionArchConfig)
num_hidden_layers_override: Optional[int] = None
require_post_norm: Optional[bool] = None
prefix: str = "clip"
@@ -0,0 +1,40 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
@dataclass
class LlamaArchConfig(TextEncoderArchConfig):
vocab_size: int = 32000
hidden_size: int = 4096
intermediate_size: int = 11008
num_hidden_layers: int = 32
num_attention_heads: int = 32
num_key_value_heads: Optional[int] = None
hidden_act: str = "silu"
max_position_embeddings: int = 2048
initializer_range: float = 0.02
rms_norm_eps: float = 1e-6
use_cache: bool = True
pad_token_id: int = 0
bos_token_id: int = 1
eos_token_id: int = 2
pretraining_tp: int = 1
tie_word_embeddings: bool = False
rope_theta: float = 10000.0
rope_scaling: Optional[float] = None
attention_bias: bool = False
attention_dropout: float = 0.0
mlp_bias: bool = False
head_dim: Optional[int] = None
hidden_state_skip_layer: int = 2
text_len: int = 256
@dataclass
class LlamaConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(default_factory=LlamaArchConfig)
prefix: str = "llama"
@@ -0,0 +1,55 @@
from dataclasses import dataclass, field
from typing import Optional
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
@dataclass
class T5ArchConfig(TextEncoderArchConfig):
vocab_size: int = 32128
d_model: int = 512
d_kv: int = 64
d_ff: int = 2048
num_layers: int = 6
num_decoder_layers: Optional[int] = None
num_heads: int = 8
relative_attention_num_buckets: int = 32
relative_attention_max_distance: int = 128
dropout_rate: float = 0.1
layer_norm_epsilon: float = 1e-6
initializer_factor: float = 1.0
feed_forward_proj: str = "relu"
dense_act_fn: str = ""
is_gated_act: bool = False
is_encoder_decoder: bool = True
use_cache: bool = True
pad_token_id: int = 0
eos_token_id: int = 1
classifier_dropout: float = 0.0
text_len: int = 512
# Referenced from https://github.com/huggingface/transformers/blob/main/src/transformers/models/t5/configuration_t5.py
def __post_init__(self):
super().__post_init__()
act_info = self.feed_forward_proj.split("-")
self.dense_act_fn: str = act_info[-1]
self.is_gated_act: bool = act_info[0] == "gated"
if self.feed_forward_proj == "gated-gelu":
self.dense_act_fn = "gelu_new"
self.tokenizer_kwargs = {
"padding": "max_length",
"truncation": True,
"max_length": self.text_len,
"add_special_tokens": True,
"return_attention_mask": True,
"return_tensors": "pt",
}
@dataclass
class T5Config(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(default_factory=T5ArchConfig)
prefix: str = "t5"
@@ -0,0 +1,9 @@
from fastvideo.v1.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.v1.configs.models.vaes.stepvideovae import StepVideoVAEConfig
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig
__all__ = [
"HunyuanVAEConfig",
"WanVAEConfig",
"StepVideoVAEConfig",
]
+130
View File
@@ -0,0 +1,130 @@
from dataclasses import dataclass, field
from typing import Any, Union
import torch
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.utils import StoreBoolean
@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 = field(default_factory=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
use_temporal_scaling_frames: bool = True
def __post_init__(self):
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
@staticmethod
def add_cli_args(parser: Any, prefix: str = "vae-config") -> Any:
"""Add CLI arguments for VAEConfig fields"""
parser.add_argument(
f"--{prefix}.load-encoder",
action=StoreBoolean,
dest=f"{prefix.replace('-', '_')}.load_encoder",
default=VAEConfig.load_encoder,
help="Whether to load the VAE encoder",
)
parser.add_argument(
f"--{prefix}.load-decoder",
action=StoreBoolean,
dest=f"{prefix.replace('-', '_')}.load_decoder",
default=VAEConfig.load_decoder,
help="Whether to load the VAE decoder",
)
parser.add_argument(
f"--{prefix}.tile-sample-min-height",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_min_height",
default=VAEConfig.tile_sample_min_height,
help="Minimum height for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.tile-sample-min-width",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_min_width",
default=VAEConfig.tile_sample_min_width,
help="Minimum width for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.tile-sample-min-num-frames",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_min_num_frames",
default=VAEConfig.tile_sample_min_num_frames,
help="Minimum number of frames for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.tile-sample-stride-height",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_stride_height",
default=VAEConfig.tile_sample_stride_height,
help="Stride height for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.tile-sample-stride-width",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_stride_width",
default=VAEConfig.tile_sample_stride_width,
help="Stride width for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.tile-sample-stride-num-frames",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_stride_num_frames",
default=VAEConfig.tile_sample_stride_num_frames,
help="Stride number of frames for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.blend-num-frames",
type=int,
dest=f"{prefix.replace('-', '_')}.blend_num_frames",
default=VAEConfig.blend_num_frames,
help="Number of frames to blend for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.use-tiling",
action=StoreBoolean,
dest=f"{prefix.replace('-', '_')}.use_tiling",
default=VAEConfig.use_tiling,
help="Whether to use tiling for VAE",
)
parser.add_argument(
f"--{prefix}.use-temporal-tiling",
action=StoreBoolean,
dest=f"{prefix.replace('-', '_')}.use_temporal_tiling",
default=VAEConfig.use_temporal_tiling,
help="Whether to use temporal tiling for VAE",
)
parser.add_argument(
f"--{prefix}.use-parallel-tiling",
action=StoreBoolean,
dest=f"{prefix.replace('-', '_')}.use_parallel_tiling",
default=VAEConfig.use_parallel_tiling,
help="Whether to use parallel tiling for VAE",
)
return parser
@@ -0,0 +1,40 @@
from dataclasses import dataclass, field
from typing import Tuple
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
@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 = field(default_factory=HunyuanVAEArchConfig)
@@ -0,0 +1,28 @@
from dataclasses import dataclass, field
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class StepVideoVAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
z_channels: int = 64
num_res_blocks: int = 2
version: int = 2
frame_len: int = 17
world_size: int = 1
spatial_compression_ratio: int = 16
temporal_compression_ratio: int = 8
scaling_factor: float = 1.0
@dataclass
class StepVideoVAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=StepVideoVAEArchConfig)
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
use_temporal_scaling_frames: bool = False
@@ -0,0 +1,75 @@
from dataclasses import dataclass, field
from typing import Tuple
import torch
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
@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 = field(default_factory=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,18 @@
from fastvideo.v1.configs.pipelines.base import (PipelineConfig,
SlidingTileAttnConfig)
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
HunyuanConfig)
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_for_name)
from fastvideo.v1.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
WanI2V720PConfig,
WanT2V480PConfig,
WanT2V720PConfig)
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"get_pipeline_config_cls_for_name"
]
+151
View File
@@ -0,0 +1,151 @@
import json
from dataclasses import asdict, dataclass, field, fields
from typing import Any, Callable, Dict, Optional, Tuple, cast
import torch
from fastvideo.v1.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
VAEConfig)
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import shallow_asdict
logger = init_logger(__name__)
def preprocess_text(prompt: str) -> str:
return prompt
def postprocess_text(output: BaseEncoderOutput) -> torch.tensor:
raise NotImplementedError
@dataclass
class PipelineConfig:
"""Base configuration for all pipeline architectures."""
# Video generation parameters
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_config: VAEConfig = field(default_factory=VAEConfig)
# DiT configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
# Text encoder configuration
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp16", ))
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), ))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
...] = field(default_factory=lambda:
(postprocess_text, ))
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
# Compilation
enable_torch_compile: bool = False
@classmethod
def from_pretrained(cls, model_path: str) -> "PipelineConfig":
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_for_name)
pipeline_config_cls = get_pipeline_config_cls_for_name(model_path)
if pipeline_config_cls is not None:
pipeline_config = pipeline_config_cls()
else:
logger.warning(
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
model_path)
pipeline_config = cls()
return cast(PipelineConfig, pipeline_config)
def dump_to_json(self, file_path: str):
output_dict = shallow_asdict(self)
del_keys = []
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
elif isinstance(value, tuple) and all(
isinstance(v, ModelConfig) for v in value):
model_dicts = []
for v in value:
model_dict = asdict(v)
# Model Arch Config should be hidden away from the users
model_dict.pop("arch_config")
model_dicts.append(model_dict)
output_dict[key] = model_dicts
elif isinstance(value, tuple) and all(callable(f) for f in value):
# Skip dumping functions
del_keys.append(key)
for key in del_keys:
output_dict.pop(key, None)
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) 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)
elif isinstance(current_value, tuple) and all(
isinstance(v, ModelConfig) for v in current_value):
assert len(current_value) == len(
new_value
), "Users shouldn't delete or add text encoder config objects in your json"
for target_config, source_config in zip(
current_value, new_value):
target_config.update_model_config(source_config)
else:
setattr(self, key, new_value)
if hasattr(self, "__post_init__"):
self.__post_init__()
@dataclass
class SlidingTileAttnConfig(PipelineConfig):
"""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
+100
View File
@@ -0,0 +1,100 @@
from dataclasses import dataclass, field
from typing import Callable, Tuple, TypedDict
import torch
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
CLIPTextConfig, LlamaConfig)
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
PROMPT_TEMPLATE_ENCODE_VIDEO = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
class PromptTemplate(TypedDict):
template: str
crop_start: int
prompt_template_video: PromptTemplate = {
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
"crop_start": 95,
}
def llama_preprocess_text(prompt: str) -> str:
return prompt_template_video["template"].format(prompt)
def llama_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
hidden_state_skip_layer = 2
assert outputs.hidden_states is not None
hidden_states: Tuple[torch.Tensor, ...] = outputs.hidden_states
last_hidden_state: torch.tensor = hidden_states[-(hidden_state_skip_layer +
1)]
crop_start = prompt_template_video.get("crop_start", -1)
last_hidden_state = last_hidden_state[:, crop_start:]
return last_hidden_state
def clip_preprocess_text(prompt: str) -> str:
return prompt
def clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
pooler_output: torch.tensor = outputs.pooler_output
return pooler_output
@dataclass
class HunyuanConfig(PipelineConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
# DiT
dit_config: DiTConfig = field(default_factory=HunyuanVideoConfig)
# VAE
vae_config: VAEConfig = field(default_factory=HunyuanVAEConfig)
# Denoising stage
embedded_cfg_scale: int = 6
flow_shift: int = 7
# Text encoding stage
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LlamaConfig(), CLIPTextConfig()))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (llama_preprocess_text, clip_preprocess_text))
postprocess_text_funcs: Tuple[
Callable[[BaseEncoderOutput], torch.tensor],
...] = field(default_factory=lambda:
(llama_postprocess_text, clip_postprocess_text))
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp16"))
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
flow_shift: int = 17
# No need to re-specify guidance_scale or embedded_cfg_scale as they
# already have the desired values from HunyuanConfig
@@ -3,9 +3,14 @@
import os
from typing import Callable, Dict, Optional, Type
from fastvideo.v1.configs.base import BaseConfig
from fastvideo.v1.configs.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.v1.configs.wan import WanI2V480PConfig, WanT2V480PConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
HunyuanConfig)
from fastvideo.v1.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
WanI2V720PConfig,
WanT2V480PConfig,
WanT2V720PConfig)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import (maybe_download_model_index,
verify_model_config_and_directory)
@@ -13,11 +18,14 @@ from fastvideo.v1.utils import (maybe_download_model_index,
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[BaseConfig]] = {
"FastVideo/FastHunyuan-Diffusers": FastHunyuanConfig,
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
"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
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V720PConfig,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
# Add other specific weight variants
}
@@ -26,22 +34,24 @@ 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(),
"stepvideo": lambda id: "stepvideo" 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]] = {
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
"hunyuan":
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
"wanpipeline":
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
"stepvideo": StepVideoT2VConfig
# Other fallbacks by architecture
}
def get_pipeline_config_cls_for_name(
pipeline_name_or_path: str) -> Optional[type[BaseConfig]]:
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
"""Get the appropriate config class for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
@@ -65,7 +75,6 @@ def get_pipeline_config_cls_for_name(
# 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()):
@@ -0,0 +1,29 @@
from dataclasses import dataclass, field
from fastvideo.v1.configs.models import DiTConfig, VAEConfig
from fastvideo.v1.configs.models.dits import StepVideoConfig
from fastvideo.v1.configs.models.vaes import StepVideoVAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
@dataclass
class StepVideoT2VConfig(PipelineConfig):
"""Base configuration for StepVideo pipeline architecture."""
# WanConfig-specific parameters with defaults
# DiT
dit_config: DiTConfig = field(default_factory=StepVideoConfig)
# VAE
vae_config: VAEConfig = field(default_factory=StepVideoVAEConfig)
vae_tiling: bool = False
vae_sp: bool = False
# Denoising stage
flow_shift: int = 13
timesteps_scale: bool = False
pos_magic: str = "超高清、HDR 视频、环境光、杜比全景声、画面稳定、流畅动作、逼真的细节、专业级构图、超现实主义、自然、生动、超细节、清晰。"
neg_magic: str = "画面暗、低分辨率、不良手、文本、缺少手指、多余的手指、裁剪、低质量、颗粒状、签名、水印、用户名、模糊。"
# Precision for each component
precision: str = "bf16"
vae_precision: str = "bf16"
+99
View File
@@ -0,0 +1,99 @@
from dataclasses import dataclass, field
from typing import Callable, Tuple
import torch
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
CLIPVisionConfig, T5Config)
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
mask: torch.tensor = outputs.attention_mask
hidden_state: torch.tensor = outputs.last_hidden_state
seq_lens = mask.gt(0).sum(dim=1).long()
assert torch.isnan(hidden_state).sum() == 0
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens)]
prompt_embeds_tensor: torch.tensor = torch.stack([
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
for u in prompt_embeds
],
dim=0)
return prompt_embeds_tensor
@dataclass
class WanT2V480PConfig(PipelineConfig):
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
# WanConfig-specific parameters with defaults
# DiT
dit_config: DiTConfig = field(default_factory=WanVideoConfig)
# VAE
vae_config: VAEConfig = field(default_factory=WanVAEConfig)
vae_tiling: bool = False
vae_sp: bool = False
# Video parameters
use_cpu_offload: bool = True
# Denoising stage
flow_shift: int = 3
# Text encoding stage
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (T5Config(), ))
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
...] = field(default_factory=lambda:
(t5_postprocess_text, ))
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp32", ))
# WanConfig-specific added parameters
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@dataclass
class WanT2V720PConfig(WanT2V480PConfig):
"""Base configuration for Wan T2V 14B 720P pipeline architecture."""
# WanConfig-specific parameters with defaults
# Denoising stage
flow_shift: int = 5
@dataclass
class WanI2V480PConfig(WanT2V480PConfig):
"""Base configuration for Wan I2V 14B 480P pipeline architecture."""
# WanConfig-specific parameters with defaults
# Precision for each component
image_encoder_config: EncoderConfig = field(
default_factory=CLIPVisionConfig)
image_encoder_precision: str = "fp32"
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@dataclass
class WanI2V720PConfig(WanI2V480PConfig):
"""Base configuration for Wan I2V 14B 720P pipeline architecture."""
# WanConfig-specific parameters with defaults
# Denoising stage
flow_shift: int = 5
@@ -0,0 +1,3 @@
from fastvideo.v1.configs.quantization.base import QuantizationConfig
__all__ = ["QuantizationConfig"]
@@ -0,0 +1,6 @@
from dataclasses import dataclass
@dataclass
class QuantizationConfig:
pass
+3
View File
@@ -0,0 +1,3 @@
from fastvideo.v1.configs.sample.base import SamplingParam
__all__ = ["SamplingParam"]
+191
View File
@@ -0,0 +1,191 @@
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Union
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@dataclass
class SamplingParam:
"""
Sampling parameters for video generation.
"""
# All fields below are copied from ForwardBatch
data_type: str = "video"
# Image inputs
image_path: Optional[str] = None
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
negative_prompt: Optional[str] = None
prompt_path: Optional[str] = None
output_path: str = "outputs/"
# Batch info
num_videos_per_prompt: int = 1
seed: int = 1024
# Original dimensions (before VAE scaling)
num_frames: int = 125
num_frames_round_down: bool = False # Whether to round down num_frames if it's not divisible by num_gpus
height: int = 720
width: int = 1280
fps: int = 24
# Denoising parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
# TeaCache parameters
enable_teacache: bool = False
# Misc
save_video: bool = True
return_frames: bool = False
def __post_init__(self) -> None:
self.data_type = "video" if self.num_frames > 1 else "image"
def check_sampling_param(self):
if self.prompt_path and not self.prompt_path.endswith(".txt"):
raise ValueError("prompt_path must be a txt file")
def update(self, source_dict: Dict[str, Any]) -> None:
for key, value in source_dict.items():
if hasattr(self, key):
setattr(self, key, value)
else:
logger.exception("%s has no attribute %s",
type(self).__name__, key)
self.__post_init__()
@classmethod
def from_pretrained(cls, model_path: str) -> "SamplingParam":
from fastvideo.v1.configs.sample.registry import (
get_sampling_param_cls_for_name)
sampling_cls = get_sampling_param_cls_for_name(model_path)
if sampling_cls is not None:
sampling_param: SamplingParam = sampling_cls()
else:
logger.warning(
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
model_path)
sampling_param = cls()
return sampling_param
@staticmethod
def add_cli_args(parser: Any) -> Any:
"""Add CLI arguments for SamplingParam fields"""
parser.add_argument(
"--prompt",
type=str,
default=SamplingParam.prompt,
help="Text prompt for video generation",
)
parser.add_argument(
"--negative-prompt",
type=str,
default=SamplingParam.negative_prompt,
help="Negative text prompt for video generation",
)
parser.add_argument(
"--prompt-path",
type=str,
default=SamplingParam.prompt_path,
help="Path to a text file containing the prompt",
)
parser.add_argument(
"--output-path",
type=str,
default=SamplingParam.output_path,
help="Path to save the generated video",
)
parser.add_argument(
"--num-videos-per-prompt",
type=int,
default=SamplingParam.num_videos_per_prompt,
help="Number of videos to generate per prompt",
)
parser.add_argument(
"--seed",
type=int,
default=SamplingParam.seed,
help="Random seed for generation",
)
parser.add_argument(
"--num-frames",
type=int,
default=SamplingParam.num_frames,
help="Number of frames to generate",
)
parser.add_argument(
"--height",
type=int,
default=SamplingParam.height,
help="Height of generated video",
)
parser.add_argument(
"--width",
type=int,
default=SamplingParam.width,
help="Width of generated video",
)
parser.add_argument(
"--fps",
type=int,
default=SamplingParam.fps,
help="Frames per second for saved video",
)
parser.add_argument(
"--num-inference-steps",
type=int,
default=SamplingParam.num_inference_steps,
help="Number of denoising steps",
)
parser.add_argument(
"--guidance-scale",
type=float,
default=SamplingParam.guidance_scale,
help="Classifier-free guidance scale",
)
parser.add_argument(
"--guidance-rescale",
type=float,
default=SamplingParam.guidance_rescale,
help="Guidance rescale factor",
)
parser.add_argument(
"--save-video",
action="store_true",
default=SamplingParam.save_video,
help="Whether to save the video to disk",
)
parser.add_argument(
"--no-save-video",
action="store_false",
dest="save_video",
help="Don't save the video to disk",
)
parser.add_argument(
"--return-frames",
action="store_true",
default=SamplingParam.return_frames,
help="Whether to return the raw frames",
)
parser.add_argument(
"--image-path",
type=str,
default=SamplingParam.image_path,
help="Path to input image for image-to-video generation",
)
return parser
@dataclass
class CacheParams:
cache_type: str = "none"
+29
View File
@@ -0,0 +1,29 @@
from dataclasses import dataclass, field
from fastvideo.v1.configs.sample.base import SamplingParam
from fastvideo.v1.configs.sample.teacache import TeaCacheParams
@dataclass
class HunyuanSamplingParam(SamplingParam):
num_inference_steps: int = 50
num_frames: int = 125
height: int = 720
width: int = 1280
fps: int = 24
guidance_scale: float = 1.0
teacache_params: TeaCacheParams = field(
default_factory=lambda: TeaCacheParams(
teacache_thresh=0.15,
coefficients=[
7.33226126e+02, -4.01131952e+02, 6.75869174e+01,
-3.14987800e+00, 9.61237896e-02
]))
@dataclass
class FastHunyuanSamplingParam(HunyuanSamplingParam):
num_inference_steps: int = 6
+83
View File
@@ -0,0 +1,83 @@
import os
from typing import Any, Callable, Dict, Optional
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.v1.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.v1.configs.sample.wan import (WanI2V_14B_480P_SamplingParam,
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam)
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
SAMPLING_PARAM_REGISTRY: Dict[str, Any] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
# Add other specific weight variants
}
# For determining pipeline type from model ID
SAMPLING_PARAM_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(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
# Add other pipeline architecture detectors
}
# Fallback configs when exact match isn't found but architecture is detected
SAMPLING_FALLBACK_PARAM: Dict[str, Any] = {
"hunyuan":
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
"wanpipeline":
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
"stepvideo": StepVideoT2VSamplingParam
# Other fallbacks by architecture
}
def get_sampling_param_cls_for_name(
pipeline_name_or_path: str) -> Optional[Any]:
"""Get the appropriate sampling param 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 sampling param 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 SAMPLING_PARAM_REGISTRY:
return SAMPLING_PARAM_REGISTRY[pipeline_name_or_path]
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in SAMPLING_PARAM_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
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in SAMPLING_PARAM_DETECTOR.items():
if detector(pipeline_name.lower()):
fallback_config = SAMPLING_FALLBACK_PARAM.get(pipeline_type)
break
logger.warning(
"No match found for pipeline %s, using fallback sampling param %s.",
pipeline_name_or_path, fallback_config)
return fallback_config
+19
View File
@@ -0,0 +1,19 @@
from dataclasses import dataclass
from fastvideo.v1.configs.sample.base import SamplingParam
@dataclass
class StepVideoT2VSamplingParam(SamplingParam):
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 81
# Denoising stage
guidance_scale: float = 9.0
num_inference_steps: int = 50
# neg magic and pos magic
# pos_magic: str = "超高清、HDR 视频、环境光、杜比全景声、画面稳定、流畅动作、逼真的细节、专业级构图、超现实主义、自然、生动、超细节、清晰。"
# neg_magic: str = "画面暗、低分辨率、不良手、文本、缺少手指、多余的手指、裁剪、低质量、颗粒状、签名、水印、用户名、模糊。"
+40
View File
@@ -0,0 +1,40 @@
from dataclasses import dataclass, field
from fastvideo.v1.configs.sample.base import CacheParams
@dataclass
class TeaCacheParams(CacheParams):
cache_type: str = "teacache"
teacache_thresh: float = 0.0
coefficients: list[float] = field(default_factory=list)
@dataclass
class WanTeaCacheParams(CacheParams):
# Unfortunately, TeaCache is very different for Wan than other models
cache_type: str = "teacache"
teacache_thresh: float = 0.0
use_ret_steps: bool = True
ret_steps_coeffs: list[float] = field(default_factory=list)
non_ret_steps_coeffs: list[float] = field(default_factory=list)
@property
def coefficients(self) -> list[float]:
if self.use_ret_steps:
return self.ret_steps_coeffs
else:
return self.non_ret_steps_coeffs
@property
def ret_steps(self) -> int:
if self.use_ret_steps:
return 5 * 2
else:
return 1 * 2
def get_cutoff_steps(self, num_inference_steps: int) -> int:
if self.use_ret_steps:
return num_inference_steps * 2
else:
return num_inference_steps * 2 - 2
+95
View File
@@ -0,0 +1,95 @@
from dataclasses import dataclass, field
from fastvideo.v1.configs.sample.base import SamplingParam
from fastvideo.v1.configs.sample.teacache import WanTeaCacheParams
@dataclass
class WanT2V_1_3B_SamplingParam(SamplingParam):
# Video parameters
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
# Denoising stage
guidance_scale: float = 3.0
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
num_inference_steps: int = 50
teacache_params: WanTeaCacheParams = field(
default_factory=lambda: WanTeaCacheParams(
teacache_thresh=0.08,
ret_steps_coeffs=[
-5.21862437e+04, 9.23041404e+03, -5.28275948e+02,
1.36987616e+01, -4.99875664e-02
],
non_ret_steps_coeffs=[
2.39676752e+03, -1.31110545e+03, 2.01331979e+02,
-8.29855975e+00, 1.37887774e-01
]))
@dataclass
class WanT2V_14B_SamplingParam(SamplingParam):
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 81
fps: int = 16
# Denoising stage
guidance_scale: float = 5.0
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
num_inference_steps: int = 50
teacache_params: WanTeaCacheParams = field(
default_factory=lambda: WanTeaCacheParams(
teacache_thresh=0.20,
use_ret_steps=False,
ret_steps_coeffs=[
-3.03318725e+05, 4.90537029e+04, -2.65530556e+03,
5.87365115e+01, -3.15583525e-01
],
non_ret_steps_coeffs=[
-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429,
-13.02252404
]))
@dataclass
class WanI2V_14B_480P_SamplingParam(WanT2V_1_3B_SamplingParam):
# Denoising stage
guidance_scale: float = 5.0
num_inference_steps: int = 40
teacache_params: WanTeaCacheParams = field(
default_factory=lambda: WanTeaCacheParams(
teacache_thresh=0.26,
ret_steps_coeffs=[
-3.03318725e+05, 4.90537029e+04, -2.65530556e+03,
5.87365115e+01, -3.15583525e-01
],
non_ret_steps_coeffs=[
-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429,
-13.02252404
]))
@dataclass
class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
# Denoising stage
guidance_scale: float = 5.0
num_inference_steps: int = 40
teacache_params: WanTeaCacheParams = field(
default_factory=lambda: WanTeaCacheParams(
teacache_thresh=0.3,
ret_steps_coeffs=[
-3.03318725e+05, 4.90537029e+04, -2.65530556e+03,
5.87365115e+01, -3.15583525e-01
],
non_ret_steps_coeffs=[
-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429,
-13.02252404
]))
-45
View File
@@ -1,45 +0,0 @@
from dataclasses import dataclass
from fastvideo.v1.configs.base import BaseConfig
@dataclass
class WanT2V480PConfig(BaseConfig):
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
# WanConfig-specific parameters with defaults
# 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
@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"
@@ -0,0 +1,41 @@
{
"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
},
"dit_config": {
"prefix": "Wan",
"quant_config": null
},
"text_encoder_precisions": [
"fp32"
],
"text_encoder_configs": [
{
"prefix": "t5",
"quant_config": null,
"lora_config": null
}
],
"mask_strategy_file_path": null,
"enable_torch_compile": false
}
@@ -0,0 +1,49 @@
{
"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
},
"dit_config": {
"prefix": "Wan",
"quant_config": null
},
"text_encoder_precisions": [
"fp32"
],
"text_encoder_configs": [
{
"prefix": "t5",
"quant_config": null,
"lora_config": null
}
],
"mask_strategy_file_path": null,
"enable_torch_compile": false,
"image_encoder_config": {
"prefix": "clip",
"quant_config": null,
"lora_config": null,
"num_hidden_layers_override": null,
"require_post_norm": null
},
"image_encoder_precision": "fp32"
}
@@ -1,16 +0,0 @@
num_gpus: 4
model_path: FastVideo/FastHunyuan-diffusers
master_port: 29503
sp_size: 4
tp_size: 4
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/
vae-sp: True
+3 -3
View File
@@ -94,14 +94,14 @@ class StatelessProcessGroup:
key = f"send_to/{dst}/{self.send_dst_counter[dst]}"
self.store.set(key, pickle.dumps(obj))
self.send_dst_counter[dst] += 1
self.entries.append((key, time.time()))
self.entries.append((key, time.perf_counter()))
def expire_data(self) -> None:
"""Expire data that is older than `data_expiration_seconds` seconds."""
while self.entries:
# check the oldest entry
key, timestamp = self.entries[0]
if time.time() - timestamp > self.data_expiration_seconds:
if time.perf_counter() - timestamp > self.data_expiration_seconds:
self.store.delete_key(key)
self.entries.popleft()
else:
@@ -125,7 +125,7 @@ class StatelessProcessGroup:
f"{self.broadcast_send_counter}")
self.store.set(key, pickle.dumps(obj))
self.broadcast_send_counter += 1
self.entries.append((key, time.time()))
self.entries.append((key, time.perf_counter()))
return obj
else:
key = (f"broadcast_from/{src}/"
+99 -34
View File
@@ -2,10 +2,14 @@
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
import argparse
from typing import List, cast
import dataclasses
import os
from typing import Any, Dict, List, Optional, cast
from fastvideo.v1.entrypoints.cli import utils
from fastvideo import PipelineConfig, VideoGenerator
from fastvideo.v1.configs.sample.base import SamplingParam
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.v1.entrypoints.cli.utils import RaiseNotImplementedAction
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.utils import FlexibleArgumentParser
@@ -16,45 +20,70 @@ class GenerateSubcommand(CLISubcommand):
def __init__(self) -> None:
self.name = "generate"
super().__init__()
self.init_arg_names = self._get_init_arg_names()
self.generation_arg_names = self._get_generation_arg_names()
def _get_init_arg_names(self) -> List[str]:
"""Get names of arguments for VideoGenerator initialization"""
return ["num_gpus", "tp_size", "sp_size", "model_path"]
def _get_generation_arg_names(self) -> List[str]:
"""Get names of arguments for generate_video method"""
return [field.name for field in dataclasses.fields(SamplingParam)]
def cmd(self, args: argparse.Namespace) -> None:
excluded_args = [
'subparser', 'config', 'num_gpus', 'master_port',
'dispatch_function'
]
excluded_args = ['subparser', 'config', 'dispatch_function']
# Create a filtered dictionary of arguments
filtered_args = {
FastVideoArgs.from_cli_args(args)
filtered_args = {}
for k, v in vars(args).items():
if k not in excluded_args and v is not None:
filtered_args[k] = v
merged_args = {**filtered_args}
if 'model_path' not in merged_args or not merged_args['model_path']:
raise ValueError(
"model_path must be provided either in config file or via --model-path"
)
if 'prompt' not in merged_args or not merged_args['prompt']:
raise ValueError(
"prompt must be provided either in config file or via --prompt")
init_args = {
k: v
for k, v in vars(args).items()
if k not in excluded_args and v is not None
for k, v in merged_args.items() if k in self.init_arg_names
}
generation_args = {
k: v
for k, v in merged_args.items() if k in self.generation_arg_names
}
main_args = []
pipeline_config = PipelineConfig.from_pretrained(
merged_args['model_path'])
for key, value in filtered_args.items():
# Convert underscores to dashes in argument names
arg_name = f"--{key.replace('_', '-')}"
update_config_from_args(pipeline_config.dit_config, merged_args,
"dit_config")
update_config_from_args(pipeline_config.vae_config, merged_args,
"vae_config")
update_config_from_args(pipeline_config, merged_args)
# Handle boolean flags
if isinstance(value, bool):
if value:
main_args.append(arg_name)
else:
main_args.append(arg_name)
main_args.append(str(value))
model_path = init_args.pop('model_path')
prompt = generation_args.pop('prompt')
utils.launch_distributed(args.num_gpus,
main_args,
master_port=args.master_port)
generator = VideoGenerator.from_pretrained(
model_path=model_path, **init_args, pipeline_config=pipeline_config)
generator.generate_video(prompt=prompt, **generation_args)
def validate(self, args: argparse.Namespace) -> None:
"""Validate the arguments for this command"""
if args.num_gpus is not None and args.num_gpus <= 0:
raise ValueError("Number of gpus must be positive")
if args.master_port is not None and (args.master_port < 1024
or args.master_port > 65535):
raise ValueError("Master port must be between 1024 and 65535")
if args.config and not os.path.exists(args.config):
raise ValueError(f"Config file not found: {args.config}")
def subparser_init(
self,
@@ -63,7 +92,7 @@ class GenerateSubcommand(CLISubcommand):
"generate",
help="Run inference on a model",
usage=
"fastvideo generate --model-path MODEL_PATH_OR_ID --prompt PROMPT [OPTIONS]"
"fastvideo generate (--model-path MODEL_PATH_OR_ID --prompt PROMPT) | --config CONFIG_FILE [OPTIONS]"
)
generate_parser.add_argument(
@@ -71,17 +100,53 @@ class GenerateSubcommand(CLISubcommand):
type=str,
default='',
required=False,
help="Read CLI options from a config YAML file.")
generate_parser.add_argument("--master-port",
type=int,
default=None,
help="Port for the master process")
help=
"Read CLI options from a config JSON or YAML file. If provided, --model-path and --prompt are optional."
)
generate_parser = FastVideoArgs.add_cli_args(generate_parser)
generate_parser = SamplingParam.add_cli_args(generate_parser)
generate_parser.add_argument(
"--text-encoder-configs",
action=RaiseNotImplementedAction,
help=
"JSON array of text encoder configurations (NOT YET IMPLEMENTED)",
)
return cast(FlexibleArgumentParser, generate_parser)
def cmd_init() -> List[CLISubcommand]:
return [GenerateSubcommand()]
def update_config_from_args(config: Any,
args_dict: Dict[str, Any],
prefix: Optional[str] = None) -> None:
"""
Update configuration object from arguments dictionary.
Args:
config: The configuration object to update
args_dict: Dictionary containing arguments
prefix: Prefix for the configuration parameters in the args_dict.
If None, assumes direct attribute mapping without prefix.
"""
# Handle top-level attributes (no prefix)
if prefix is None:
for key, value in args_dict.items():
if hasattr(config, key) and value is not None:
if key == "text_encoder_precisions" and isinstance(value, list):
setattr(config, key, tuple(value))
else:
setattr(config, key, value)
return
# Handle nested attributes with prefix
prefix_with_dot = f"{prefix}."
for key, value in args_dict.items():
if key.startswith(prefix_with_dot) and value is not None:
attr_name = key[len(prefix_with_dot):]
if hasattr(config, attr_name):
setattr(config, attr_name, value)
+8
View File
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import os
import subprocess
import sys
@@ -10,6 +11,13 @@ from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class RaiseNotImplementedAction(argparse.Action):
def __call__(self, parser, namespace, values, option_string=None):
raise NotImplementedError(
f"The {option_string} option is not yet implemented")
def launch_distributed(num_gpus: int,
args: List[str],
master_port: Optional[int] = None) -> int:
+75 -82
View File
@@ -6,10 +6,10 @@ This module provides a consolidated interface for generating videos using
diffusion models.
"""
import gc
import os
import time
from dataclasses import asdict
from typing import Any, Callable, Dict, List, Optional, Union
from typing import Any, Dict, List, Optional, Union
import imageio
import numpy as np
@@ -17,11 +17,13 @@ import torch
import torchvision
from einops import rearrange
from fastvideo.v1.configs import get_pipeline_config_cls_for_name
from fastvideo.v1.configs.pipelines import (PipelineConfig,
get_pipeline_config_cls_for_name)
from fastvideo.v1.configs.sample import SamplingParam
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
from fastvideo.v1.utils import align_to, shallow_asdict
from fastvideo.v1.worker.executor import Executor
logger = init_logger(__name__)
@@ -52,6 +54,9 @@ class VideoGenerator:
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[
Union[str
| PipelineConfig]] = None,
**kwargs) -> "VideoGenerator":
"""
Create a video generator from a pretrained model.
@@ -64,22 +69,29 @@ class VideoGenerator:
Returns:
The created video generator
Priority level: Default pipeline config < User's pipeline config < User's kwargs
"""
config = None
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, PipelineConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = {}
config_args = kwargs
else:
config_args = asdict(config)
# override config_args with kwargs
config_args.update(kwargs)
config_args = shallow_asdict(config)
config_args.update(kwargs)
fastvideo_args = FastVideoArgs(
model_path=model_path,
@@ -115,19 +127,8 @@ class VideoGenerator:
def generate_video(
self,
prompt: str,
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,
sampling_param: Optional[SamplingParam] = None,
**kwargs,
) -> Union[Dict[str, Any], List[np.ndarray]]:
"""
Generate a video based on the given prompt.
@@ -154,96 +155,80 @@ class VideoGenerator:
# Create a copy of inference args to avoid modifying the original
fastvideo_args = self.fastvideo_args
# Override parameters if provided
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()
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
kwargs["prompt"] = prompt
sampling_param.update(kwargs)
# Process negative prompt
if fastvideo_args.neg_prompt is not None:
fastvideo_args.neg_prompt = fastvideo_args.neg_prompt.strip()
if sampling_param.negative_prompt is not None:
sampling_param.negative_prompt = sampling_param.negative_prompt.strip(
)
# Validate dimensions
if (fastvideo_args.height <= 0 or fastvideo_args.width <= 0
or fastvideo_args.num_frames <= 0):
if (sampling_param.height <= 0 or sampling_param.width <= 0
or sampling_param.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}")
f"height={sampling_param.height}, width={sampling_param.width}, "
f"num_frames={sampling_param.num_frames}")
if (fastvideo_args.num_frames - 1) % 4 != 0:
if (
sampling_param.num_frames - 1
) % fastvideo_args.vae_config.arch_config.temporal_compression_ratio != 0:
raise ValueError(
f"num_frames-1 must be a multiple of 4, got {fastvideo_args.num_frames}"
f"num_frames-1 must be a multiple of {fastvideo_args.vae_config.arch_config.temporal_compression_ratio}, got {sampling_param.num_frames}"
)
# Calculate sizes
target_height = align_to(fastvideo_args.height, 16)
target_width = align_to(fastvideo_args.width, 16)
target_height = align_to(sampling_param.height, 16)
target_width = align_to(sampling_param.width, 16)
# Calculate latent sizes
latents_size = [(fastvideo_args.num_frames - 1) // 4 + 1,
fastvideo_args.height // 8, fastvideo_args.width // 8]
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# Log parameters
debug_str = f"""
height: {target_height}
width: {target_width}
video_length: {fastvideo_args.num_frames}
video_length: {sampling_param.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}
neg_prompt: {sampling_param.negative_prompt}
seed: {sampling_param.seed}
infer_steps: {sampling_param.num_inference_steps}
num_videos_per_prompt: {sampling_param.num_videos_per_prompt}
guidance_scale: {sampling_param.guidance_scale}
n_tokens: {n_tokens}
flow_shift: {fastvideo_args.flow_shift}
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}"""
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}
save_video: {sampling_param.save_video}
output_path: {sampling_param.output_path}
""" # type: ignore[attr-defined]
logger.info(debug_str)
# Prepare batch
device = torch.device(fastvideo_args.device_str)
batch = ForwardBatch(
prompt=prompt,
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,
**shallow_asdict(sampling_param),
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()
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch
gen_time = time.time() - start_time
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
# Process outputs
@@ -255,23 +240,31 @@ class VideoGenerator:
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 batch.save_video:
save_path = batch.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)
imageio.mimsave(video_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", video_path)
else:
logger.warning("No output path provided, video not saved")
if return_frames:
if batch.return_frames:
return frames
else:
return {
"samples": samples,
"prompts": prompt,
"size":
(target_height, target_width, fastvideo_args.num_frames),
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time
}
def shutdown(self):
"""
Shutdown the video generator.
"""
self.executor.shutdown()
del self.executor
gc.collect()
torch.cuda.empty_cache()
@@ -1,28 +0,0 @@
# 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!
@@ -1,25 +0,0 @@
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()
@@ -1,141 +0,0 @@
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)
+95 -177
View File
@@ -5,19 +5,32 @@
import argparse
import dataclasses
from contextlib import contextmanager
from typing import List, Optional
from dataclasses import field
from typing import Any, Callable, List, Optional, Tuple
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import FlexibleArgumentParser
from fastvideo.v1.utils import FlexibleArgumentParser, StoreBoolean
logger = init_logger(__name__)
def preprocess_text(prompt: str) -> str:
return prompt
def postprocess_text(output: Any) -> Any:
raise NotImplementedError
@dataclasses.dataclass
class FastVideoArgs:
# Model and path configuration
model_path: str
# Cache strategy
cache_strategy: str = "none"
# Distributed executor backend
distributed_executor_backend: str = "mp"
@@ -34,67 +47,56 @@ class FastVideoArgs:
dist_timeout: Optional[int] = None # timeout for torch.distributed
# Video generation parameters
height: int = 720
width: int = 1280
num_frames: int = 117
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
output_type: str = "pil"
# Model configuration
# DiT configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
precision: str = "bf16"
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = False
vae_scale_factor: Optional[int] = None
# DiT configuration
num_channels_latents: 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 = field(default_factory=VAEConfig)
# Image encoder configuration
image_encoder_precision: str = "fp32"
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_len: int = 256
hidden_state_skip_layer: int = 2
# Secondary text encoder
text_encoder_precision_2: str = "fp16"
text_len_2: int = 77
# Flow Matching parameters
flow_solver: str = "euler"
denoise_type: str = "flow" # Deprecated. Will use scheduler_config.json
DEFAULT_TEXT_ENCODER_PRECISIONS = (
"fp16",
"fp16",
)
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), ))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: Tuple[Callable[[Any], Any], ...] = field(
default_factory=lambda: (postprocess_text, ))
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
# Scheduler options
scheduler_type: str = "euler" # Deprecated. Will use the param in scheduler_config.json
neg_prompt: Optional[str] = None
num_videos: int = 1
fps: int = 24
use_cpu_offload: bool = False
disable_autocast: bool = False
# StepVideo specific parameters
pos_magic: Optional[str] = None
neg_magic: Optional[str] = None
timesteps_scale: Optional[bool] = None
# Logging
log_level: str = "info"
# Inference parameters
image_path: Optional[str] = None
prompt: Optional[str] = None
prompt_path: Optional[str] = None
output_path: str = "outputs/"
seed: int = 1024
device_str: Optional[str] = None
device = None
@@ -107,7 +109,6 @@ class FastVideoArgs:
parser.add_argument(
"--model-path",
type=str,
required=True,
help=
"The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
)
@@ -134,7 +135,7 @@ class FastVideoArgs:
# HuggingFace specific parameters
parser.add_argument(
"--trust-remote-code",
action="store_true",
action=StoreBoolean,
default=FastVideoArgs.trust_remote_code,
help="Trust remote code when loading HuggingFace models",
)
@@ -174,43 +175,6 @@ class FastVideoArgs:
help="Set timeout for torch.distributed initialization.",
)
# Video generation parameters
parser.add_argument(
"--height",
type=int,
default=FastVideoArgs.height,
help="Height of generated video",
)
parser.add_argument(
"--width",
type=int,
default=FastVideoArgs.width,
help="Width of generated video",
)
parser.add_argument(
"--num-frames",
type=int,
default=FastVideoArgs.num_frames,
help="Number of frames to generate",
)
parser.add_argument(
"--num-inference-steps",
type=int,
default=FastVideoArgs.num_inference_steps,
help="Number of inference steps",
)
parser.add_argument(
"--guidance-scale",
type=float,
default=FastVideoArgs.guidance_scale,
help="Guidance scale for classifier-free guidance",
)
parser.add_argument(
"--guidance-rescale",
type=float,
default=FastVideoArgs.guidance_rescale,
help="Guidance rescale for classifier-free guidance",
)
parser.add_argument(
"--embedded-cfg-scale",
type=float,
@@ -250,28 +214,23 @@ class FastVideoArgs:
)
parser.add_argument(
"--vae-tiling",
action="store_true",
action=StoreBoolean,
default=FastVideoArgs.vae_tiling,
help="Enable VAE tiling",
)
parser.add_argument(
"--vae-sp",
action="store_true",
action=StoreBoolean,
help="Enable VAE spatial parallelism",
)
parser.add_argument(
"--text-encoder-precision",
"--text-encoder-precisions",
nargs="+",
type=str,
default=FastVideoArgs.text_encoder_precision,
default=FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS,
choices=["fp32", "fp16", "bf16"],
help="Precision for text encoder",
)
parser.add_argument(
"--text-len",
type=int,
default=FastVideoArgs.text_len,
help="Maximum text length",
help="Precision for each text encoder",
)
# Image encoder config
@@ -283,36 +242,6 @@ class FastVideoArgs:
help="Precision for image encoder",
)
# Secondary text encoder
parser.add_argument(
"--text-encoder-precision-2",
type=str,
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=FastVideoArgs.text_len_2,
help="Maximum secondary text length",
)
# Flow Matching parameters
parser.add_argument(
"--flow-solver",
type=str,
default=FastVideoArgs.flow_solver,
help="Solver for flow matching",
)
parser.add_argument(
"--denoise-type",
type=str,
default=FastVideoArgs.denoise_type,
help="Denoise type for noised inputs",
)
# STA (Spatial-Temporal Attention) parameters
parser.add_argument(
"--mask-strategy-file-path",
@@ -321,50 +250,42 @@ class FastVideoArgs:
)
parser.add_argument(
"--enable-torch-compile",
action="store_true",
action=StoreBoolean,
help=
"Use torch.compile for speeding up STA inference without teacache",
)
# Scheduler options
parser.add_argument(
"--scheduler-type",
type=str,
default=FastVideoArgs.scheduler_type,
help="Type of scheduler to use",
)
# HunYuan specific parameters
parser.add_argument(
"--neg-prompt",
type=str,
default=FastVideoArgs.neg_prompt,
help="Negative prompt for sampling",
)
parser.add_argument(
"--num-videos",
type=int,
default=FastVideoArgs.num_videos,
help="Number of videos to generate per prompt",
)
parser.add_argument(
"--fps",
type=int,
default=FastVideoArgs.fps,
help="Frames per second for output video",
)
parser.add_argument(
"--use-cpu-offload",
action="store_true",
action=StoreBoolean,
help="Use CPU offload for the model load",
)
parser.add_argument(
"--disable-autocast",
action="store_true",
action=StoreBoolean,
help=
"Disable autocast for denoising loop and vae decoding in pipeline sampling",
)
parser.add_argument(
"--pos_magic",
type=str,
default=FastVideoArgs.pos_magic,
help="Positive magic prompt for sampling",
)
parser.add_argument(
"--neg_magic",
type=str,
default=FastVideoArgs.neg_magic,
help="Negative magic prompt for sampling",
)
parser.add_argument(
"--timesteps_scale",
type=bool,
default=FastVideoArgs.timesteps_scale,
help="Bool for applying scheduler scale in set_timesteps",
)
# Logging
parser.add_argument(
"--log-level",
@@ -373,35 +294,13 @@ class FastVideoArgs:
help="The logging level of all loggers.",
)
# Inference parameters
prompt_group = parser.add_mutually_exclusive_group(required=True)
prompt_group.add_argument(
"--prompt",
type=str,
help="Text prompt for video generation",
)
prompt_group.add_argument(
"--prompt-path",
type=str,
help="Path to a text file containing the prompt",
)
# Add VAE configuration arguments
from fastvideo.v1.configs.models.vaes.base import VAEConfig
VAEConfig.add_cli_args(parser)
parser.add_argument("--image-path",
type=str,
help="Path to the image for I2V generation")
parser.add_argument(
"--output-path",
type=str,
default=FastVideoArgs.output_path,
help="Directory to save generated videos",
)
parser.add_argument(
"--seed",
type=int,
default=FastVideoArgs.seed,
help="Random seed for reproducibility",
)
# Add DiT configuration arguments
from fastvideo.v1.configs.models.dits.base import DiTConfig
DiTConfig.add_cli_args(parser)
return parser
@@ -451,8 +350,27 @@ class FastVideoArgs:
raise ValueError(
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
)
if self.prompt_path and not self.prompt_path.endswith(".txt"):
raise ValueError("prompt_path must be a text file")
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
)
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
raise ValueError(
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
if self.enable_torch_compile and self.num_gpus > 1:
logger.warning(
"Currently torch compile does not work with multi-gpu. Setting enable_torch_compile to False"
)
self.enable_torch_compile = False
_current_fastvideo_args = None
+7 -5
View File
@@ -11,6 +11,7 @@ import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
if TYPE_CHECKING:
from fastvideo.v1.attention import AttentionMetadata
@@ -30,11 +31,13 @@ batchsize_forward_time: defaultdict = defaultdict(list)
#
@dataclass
class ForwardContext:
current_timestep: int
# TODO(will): check this arg
# copy from vllm_config.compilation_config.static_forward_context
# attn_layers: Dict[str, Any]
# TODO: extend to support per-layer dynamic forward context
attn_metadata: "AttentionMetadata" # set dynamically for each forward pass
forward_batch: Optional[ForwardBatch] = None
_forward_context: Optional[ForwardContext] = None
@@ -52,6 +55,7 @@ def get_forward_context() -> ForwardContext:
@contextmanager
def set_forward_context(current_timestep,
attn_metadata,
forward_batch: Optional[ForwardBatch] = None,
fastvideo_args: Optional[FastVideoArgs] = None):
"""A context manager that stores the current forward context,
can be attention metadata, etc.
@@ -63,7 +67,9 @@ def set_forward_context(current_timestep,
forward_start_time = time.perf_counter()
global _forward_context
prev_context = _forward_context
_forward_context = ForwardContext(attn_metadata=attn_metadata)
_forward_context = ForwardContext(current_timestep=current_timestep,
attn_metadata=attn_metadata,
forward_batch=forward_batch)
try:
yield
finally:
@@ -76,10 +82,6 @@ def set_forward_context(current_timestep,
else:
# for v1 attention backends
batchsize = attn_metadata.num_input_tokens
# we use synchronous scheduling right now,
# adding a sync point here should not affect
# scheduling of the next batch
torch.cuda.synchronize()
now = time.perf_counter()
# time measurement is in milliseconds
batchsize_forward_time[batchsize].append(
+4 -2
View File
@@ -1,3 +1,5 @@
# type: ignore
# SPDX-License-Identifier: Apache-2.0
"""
Inference module for diffusion models.
@@ -189,7 +191,7 @@ class InferenceEngine:
# ========================================================================
# Pipeline inference
# ========================================================================
start_time = time.time()
start_time = time.perf_counter()
samples = self.pipeline.forward(
batch=batch,
fastvideo_args=fastvideo_args,
@@ -199,7 +201,7 @@ class InferenceEngine:
out_dict["samples"] = samples
out_dict["prompts"] = prompt
gen_time = time.time() - start_time
gen_time = time.perf_counter() - start_time
logger.info("Success, time: %s", gen_time)
return out_dict
+68 -7
View File
@@ -1,27 +1,30 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import List, Optional, Tuple, Union
from typing import Any, List, Optional, Tuple, Union
import torch
from torch import nn
from fastvideo.v1.configs.models import DiTConfig
from fastvideo.v1.platforms import _Backend
# TODO
class BaseDiT(nn.Module, ABC):
_fsdp_shard_conditions: list = []
attention_head_dim: int | None = None
_compile_conditions: list = []
_param_names_mapping: dict
hidden_size: int
num_attention_heads: int
num_channels_latents: int
# always supports torch_sdpa
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.TORCH_SDPA, )
_supported_attention_backends: Tuple[
_Backend, ...] = DiTConfig()._supported_attention_backends
def __init_subclass__(cls) -> None:
required_class_attrs = [
"_fsdp_shard_conditions", "_param_names_mapping"
"_fsdp_shard_conditions", "_param_names_mapping",
"_compile_conditions"
]
super().__init_subclass__()
for attr in required_class_attrs:
@@ -30,8 +33,9 @@ class BaseDiT(nn.Module, ABC):
f"Subclasses of BaseDiT must define '{attr}' class variable"
)
def __init__(self, *args, **kwargs) -> None:
def __init__(self, config: DiTConfig, **kwargs) -> None:
super().__init__()
self.config = config
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
@@ -49,7 +53,9 @@ class BaseDiT(nn.Module, ABC):
pass
def __post_init__(self) -> None:
required_attrs = ["hidden_size", "num_attention_heads"]
required_attrs = [
"hidden_size", "num_attention_heads", "num_channels_latents"
]
for attr in required_attrs:
if not hasattr(self, attr):
raise AttributeError(
@@ -59,3 +65,58 @@ class BaseDiT(nn.Module, ABC):
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
return self._supported_attention_backends
class CachableDiT(BaseDiT):
"""
An intermediate base class that adds TeaCache optimization functionality to DiT models.
TeaCache accelerates inference by selectively skipping redundant computation when consecutive
diffusion steps are similar enough.
"""
# These are required class attributes that should be overridden by concrete implementations
_fsdp_shard_conditions = []
_param_names_mapping = {}
# Ensure these instance attributes are properly defined in subclasses
hidden_size: int
num_attention_heads: int
num_channels_latents: int
# always supports torch_sdpa
_supported_attention_backends: Tuple[
_Backend, ...] = DiTConfig()._supported_attention_backends
def __init__(self, config: DiTConfig, **kwargs) -> None:
super().__init__(config, **kwargs)
self.cnt = 0
self.teacache_thresh = 0
self.coefficients: list[float] = []
# NOTE(will): Only wan2.1 needs these, so we are hardcoding it here
if self.config.prefix == "wan":
self.use_ret_steps = self.config.cache_config.use_ret_steps
self.is_even = False
self.previous_e0_even: torch.Tensor | None = None
self.previous_e0_odd: torch.Tensor | None = None
self.previous_residual_even: torch.Tensor | None = None
self.previous_residual_odd: torch.Tensor | None = None
self.accumulated_rel_l1_distance_even = 0
self.accumulated_rel_l1_distance_odd = 0
self.should_calc_even = True
self.should_calc_odd = True
else:
self.accumulated_rel_l1_distance = 0
self.previous_modulated_input = None
self.previous_residual = None
def maybe_cache_states(self, hidden_states: torch.Tensor,
original_hidden_states: torch.Tensor) -> None:
pass
def should_skip_forward_for_cached_states(self,
**kwargs: dict[str, Any]) -> bool:
return False
def retrieve_cached_states(self,
hidden_states: torch.Tensor) -> torch.Tensor:
raise NotImplementedError(
"maybe_retrieve_cached_states is not implemented")
+193 -210
View File
@@ -2,12 +2,16 @@
from typing import List, Optional, Tuple, Union
import numpy as np
import torch
import torch.nn as nn
from fastvideo.v1.attention import DistributedAttention, LocalAttention
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
from fastvideo.v1.configs.sample.teacache import TeaCacheParams
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.forward_context import get_forward_context
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.v1.layers.linear import ReplicatedLinear
@@ -18,7 +22,8 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder,
unpatchify)
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.models.dits.base import CachableDiT
from fastvideo.v1.models.utils import modulate
from fastvideo.v1.platforms import _Backend
@@ -417,7 +422,7 @@ class MMSingleStreamBlock(nn.Module):
return self.output_residual(x, output, mod_gate)
class HunyuanVideoTransformer3DModel(BaseDiT):
class HunyuanVideoTransformer3DModel(CachableDiT):
"""
HunyuanVideo Transformer backbone adapted for distributed training.
@@ -431,239 +436,105 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
# PY: we make the input args the same as HF config
# shard single stream, double stream blocks, and refiner_blocks
_fsdp_shard_conditions = [
lambda n, m: "double" in n and str.isdigit(n.split(".")[-1]),
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.SAGE_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\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"txt_in.t_embedder.mlp.fc_out.\1",
r"^context_embedder\.proj_in\.(.*)$":
r"txt_in.input_embedder.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"txt_in.c_embedder.fc_in.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"txt_in.c_embedder.fc_out.\1",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$":
r"txt_in.refiner_blocks.\1.norm1.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$":
r"txt_in.refiner_blocks.\1.norm2.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 0, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 1, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 2, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
_fsdp_shard_conditions = HunyuanVideoConfig()._fsdp_shard_conditions
_compile_conditions = HunyuanVideoConfig()._compile_conditions
_supported_attention_backends = HunyuanVideoConfig(
)._supported_attention_backends
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
# 3. x_embedder mapping:
r"^x_embedder\.proj\.(.*)$":
r"img_in.proj.\1",
def __init__(self, config: HunyuanVideoConfig):
super().__init__(config=config)
# 4. Top-level time_text_embed mappings:
r"^time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"time_in.mlp.fc_in.\1",
r"^time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"time_in.mlp.fc_out.\1",
r"^time_text_embed\.guidance_embedder\.linear_1\.(.*)$":
r"guidance_in.mlp.fc_in.\1",
r"^time_text_embed\.guidance_embedder\.linear_2\.(.*)$":
r"guidance_in.mlp.fc_out.\1",
r"^time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"vector_in.fc_in.\1",
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"vector_in.fc_out.\1",
# 5. transformer_blocks mapping:
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
r"double_blocks.\1.img_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
r"double_blocks.\1.txt_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"double_blocks.\1.img_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"double_blocks.\1.img_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"double_blocks.\1.img_attn_proj.\2",
# Corrected: merge attn.to_add_out into the main projection.
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
r"double_blocks.\1.txt_attn_proj.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
r"double_blocks.\1.txt_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
r"double_blocks.\1.txt_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_out.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_out.\2",
# 6. single_transformer_blocks mapping:
r"^single_transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"single_blocks.\1.q_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"single_blocks.\1.k_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"single_blocks.\1.linear1.\2", 0, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"single_blocks.\1.linear1.\2", 1, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"single_blocks.\1.linear1.\2", 2, 4),
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$":
(r"single_blocks.\1.linear1.\2", 3, 4),
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$":
r"single_blocks.\1.linear2.\2",
r"^single_transformer_blocks\.(\d+)\.norm\.linear\.(.*)$":
r"single_blocks.\1.modulation.linear.\2",
# 7. Final layers mapping:
r"^norm_out\.linear\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
r"^proj_out\.(.*)$":
r"final_layer.linear.\1",
}
def __init__(
self,
patch_size: int = 2,
patch_size_t: int = 1,
in_channels: int = 16,
out_channels: int = 16,
num_attention_heads: int = 24,
attention_head_dim: int = 128,
mlp_ratio: float = 4.0,
num_layers: int = 20,
num_single_layers: int = 40,
num_refiner_layers: int = 2,
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56),
guidance_embeds: bool = False,
dtype: Optional[torch.dtype] = None,
text_embed_dim: int = 4096,
pooled_projection_dim: int = 768,
rope_theta: int = 256,
qk_norm: str = "rms_norm", #TODO(PY)
prefix="Hunyuan",
):
super().__init__()
hidden_size = attention_head_dim * num_attention_heads
self.patch_size = [patch_size_t, patch_size, patch_size]
self.in_channels = in_channels
self.out_channels = in_channels if out_channels is None else out_channels
self.patch_size = [
config.patch_size_t, config.patch_size, config.patch_size
]
self.in_channels = config.in_channels
self.num_channels_latents = config.num_channels_latents
self.out_channels = config.in_channels if config.out_channels is None else config.out_channels
self.unpatchify_channels = self.out_channels
self.guidance_embeds = guidance_embeds
self.rope_dim_list = list(rope_axes_dim)
self.rope_theta = rope_theta
self.text_states_dim = text_embed_dim
self.text_states_dim_2 = pooled_projection_dim
self.guidance_embeds = config.guidance_embeds
self.rope_dim_list = list(config.rope_axes_dim)
self.rope_theta = config.rope_theta
self.text_states_dim = config.text_embed_dim
self.text_states_dim_2 = config.pooled_projection_dim
# TODO(will): hack?
self.dtype = dtype
self.dtype = config.dtype
if hidden_size % num_attention_heads != 0:
pe_dim = config.hidden_size // config.num_attention_heads
if sum(config.rope_axes_dim) != pe_dim:
raise ValueError(
f"Hidden size {hidden_size} must be divisible by num_attention_heads {num_attention_heads}"
f"Got {config.rope_axes_dim} but expected positional dim {pe_dim}"
)
pe_dim = hidden_size // num_attention_heads
if sum(rope_axes_dim) != pe_dim:
raise ValueError(
f"Got {rope_axes_dim} but expected positional dim {pe_dim}")
self.hidden_size = hidden_size
self.num_attention_heads = num_attention_heads
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.num_channels_latents = config.num_channels_latents
# Image projection
self.img_in = PatchEmbed(self.patch_size,
self.in_channels,
self.hidden_size,
dtype=dtype,
prefix=f"{prefix}.img_in")
dtype=config.dtype,
prefix=f"{config.prefix}.img_in")
self.txt_in = SingleTokenRefiner(self.text_states_dim,
hidden_size,
num_attention_heads,
depth=num_refiner_layers,
dtype=dtype,
prefix=f"{prefix}.txt_in")
config.hidden_size,
config.num_attention_heads,
depth=config.num_refiner_layers,
dtype=config.dtype,
prefix=f"{config.prefix}.txt_in")
# Time modulation
self.time_in = TimestepEmbedder(self.hidden_size,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.time_in")
dtype=config.dtype,
prefix=f"{config.prefix}.time_in")
# Text modulation
self.vector_in = MLP(self.text_states_dim_2,
self.hidden_size,
self.hidden_size,
act_type="silu",
dtype=dtype,
prefix=f"{prefix}.vector_in")
dtype=config.dtype,
prefix=f"{config.prefix}.vector_in")
# Guidance modulation
self.guidance_in = (TimestepEmbedder(self.hidden_size,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.guidance_in")
self.guidance_in = (TimestepEmbedder(
self.hidden_size,
act_layer="silu",
dtype=config.dtype,
prefix=f"{config.prefix}.guidance_in")
if self.guidance_embeds else None)
# Double blocks
self.double_blocks = nn.ModuleList([
MMDoubleStreamBlock(
hidden_size,
num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
config.hidden_size,
config.num_attention_heads,
mlp_ratio=config.mlp_ratio,
dtype=config.dtype,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{prefix}.double_blocks.{i}") for i in range(num_layers)
prefix=f"{config.prefix}.double_blocks.{i}")
for i in range(config.num_layers)
])
# Single blocks
self.single_blocks = nn.ModuleList([
MMSingleStreamBlock(
hidden_size,
num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
config.hidden_size,
config.num_attention_heads,
mlp_ratio=config.mlp_ratio,
dtype=config.dtype,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{prefix}.single_blocks.{i+num_layers}")
for i in range(num_single_layers)
prefix=f"{config.prefix}.single_blocks.{i+config.num_layers}")
for i in range(config.num_single_layers)
])
self.final_layer = FinalLayer(hidden_size,
self.final_layer = FinalLayer(config.hidden_size,
self.patch_size,
self.out_channels,
dtype=dtype,
prefix=f"{prefix}.final_layer")
dtype=config.dtype,
prefix=f"{config.prefix}.final_layer")
self.__post_init__()
@@ -689,6 +560,11 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
Returns:
Tuple of (output)
"""
forward_context = get_forward_context()
forward_batch = forward_context.forward_batch
assert forward_batch is not None
enable_teacache = forward_batch.enable_teacache
if guidance is None:
guidance = torch.tensor([6016.0],
device=hidden_states.device,
@@ -736,26 +612,40 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
img_seq_len = img.shape[1]
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# Process through double stream blocks
for index, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis]
img, txt = block(*double_block_args)
# Merge txt and img to pass through single stream blocks
x = torch.cat((img, txt), 1)
# Process through single stream blocks
if len(self.single_blocks) > 0:
for index, block in enumerate(self.single_blocks):
single_block_args = [
x,
vec,
txt_seq_len,
freqs_cis,
]
x = block(*single_block_args)
should_skip_forward = self.should_skip_forward_for_cached_states(
img=img, vec=vec)
if should_skip_forward:
img = self.retrieve_cached_states(img)
else:
if enable_teacache:
original_img = img.clone()
# Process through double stream blocks
for index, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis]
img, txt = block(*double_block_args)
# Merge txt and img to pass through single stream blocks
x = torch.cat((img, txt), 1)
# Process through single stream blocks
if len(self.single_blocks) > 0:
for index, block in enumerate(self.single_blocks):
single_block_args = [
x,
vec,
txt_seq_len,
freqs_cis,
]
x = block(*single_block_args)
# Extract image features
img = x[:, :img_seq_len, ...]
if enable_teacache:
self.maybe_cache_states(img, original_img)
# Extract image features
img = x[:, :img_seq_len, ...]
# Final layer processing
img = self.final_layer(img, vec)
# Unpatchify to get original shape
@@ -763,6 +653,99 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
return img
def maybe_cache_states(self, hidden_states: torch.Tensor,
original_hidden_states: torch.Tensor) -> None:
self.previous_residual = hidden_states - original_hidden_states
def should_skip_forward_for_cached_states(self, **kwargs) -> bool:
forward_context = get_forward_context()
forward_batch = forward_context.forward_batch
assert forward_batch is not None
current_timestep = forward_context.current_timestep
enable_teacache = forward_batch.enable_teacache
if not enable_teacache:
return False
raise NotImplementedError(
"teacache is not supported yet for HunyuanVideo")
teacache_params = forward_batch.teacache_params
assert teacache_params is not None, "teacache_params is not initialized"
assert isinstance(
teacache_params,
TeaCacheParams), "teacache_params is not a TeaCacheParams"
num_inference_steps = forward_batch.num_inference_steps
teache_thresh = teacache_params.teacache_thresh
coefficients = teacache_params.coefficients
if current_timestep == 0:
self.cnt = 0
inp = kwargs["img"].clone()
vec_ = kwargs["vec"].clone()
# convert to DTensor
vec_ = torch.distributed.tensor.DTensor.from_local(
vec_,
torch.distributed.DeviceMesh(
"cuda",
list(range(get_sequence_model_parallel_world_size())),
mesh_dim_names=("dp", )),
[torch.distributed.tensor.Replicate()])
inp = torch.distributed.tensor.DTensor.from_local(
inp,
torch.distributed.DeviceMesh(
"cuda",
list(range(get_sequence_model_parallel_world_size())),
mesh_dim_names=("dp", )),
[torch.distributed.tensor.Replicate()])
# txt_ = kwargs["txt"].clone()
# inp = img.clone()
# vec_ = vec.clone()
# txt_ = txt.clone()
(
img_mod1_shift,
img_mod1_scale,
img_mod1_gate,
img_mod2_shift,
img_mod2_scale,
img_mod2_gate,
) = self.double_blocks[0].img_mod(vec_).chunk(6, dim=-1)
normed_inp = self.double_blocks[0].img_attn_norm.norm(inp)
modulated_inp = modulate(normed_inp,
shift=img_mod1_shift,
scale=img_mod1_scale)
if self.cnt == 0 or self.cnt == num_inference_steps - 1:
should_calc = True
self.accumulated_rel_l1_distance = 0
else:
coefficients = [
7.33226126e+02, -4.01131952e+02, 6.75869174e+01,
-3.14987800e+00, 9.61237896e-02
]
rescale_func = np.poly1d(coefficients)
assert self.previous_modulated_input is not None, "previous_modulated_input is not initialized"
self.accumulated_rel_l1_distance += rescale_func(
((modulated_inp - self.previous_modulated_input).abs().mean() /
self.previous_modulated_input.abs().mean()).cpu().item())
if self.accumulated_rel_l1_distance < teache_thresh:
should_calc = False
else:
should_calc = True
self.accumulated_rel_l1_distance = 0
self.previous_modulated_input = modulated_inp
self.cnt += 1
return not should_calc
def retrieve_cached_states(self,
hidden_states: torch.Tensor) -> torch.Tensor:
return hidden_states + self.previous_residual
class SingleTokenRefiner(nn.Module):
"""
+682
View File
@@ -0,0 +1,682 @@
# Copyright 2025 StepFun Inc. All Rights Reserved.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
# ==============================================================================
from typing import Dict, Optional, Tuple
import torch
from einops import rearrange, repeat
from torch import nn
from fastvideo.v1.attention import DistributedAttention, LocalAttention
from fastvideo.v1.configs.models.dits import StepVideoConfig
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.layers.layernorm import LayerNormScaleShift
from fastvideo.v1.layers.linear import ReplicatedLinear
from fastvideo.v1.layers.mlp import MLP
from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.v1.layers.visual_embedding import TimestepEmbedder
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.platforms import _Backend
class PatchEmbed2D(nn.Module):
"""2D Image to Patch Embedding
Image to Patch Embedding using Conv2d
A convolution based approach to patchifying a 2D image w/ embedding projection.
Based on the impl in https://github.com/google-research/vision_transformer
Hacked together by / Copyright 2020 Ross Wightman
Remove the _assert function in forward function to be compatible with multi-resolution images.
"""
def __init__(self,
patch_size=16,
in_chans=3,
embed_dim=768,
norm_layer=None,
flatten=True,
bias=True,
dtype=None,
prefix: str = ""):
super().__init__()
# Convert patch_size to 2-tuple
if isinstance(patch_size, (list, tuple)):
if len(patch_size) == 1:
patch_size = (patch_size[0], patch_size[0])
else:
patch_size = (patch_size, patch_size)
self.patch_size = patch_size
self.flatten = flatten
self.proj = nn.Conv2d(in_chans,
embed_dim,
kernel_size=patch_size,
stride=patch_size,
bias=bias,
dtype=dtype)
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
def forward(self, x):
x = self.proj(x)
if self.flatten:
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
x = self.norm(x)
return x
class StepVideoRMSNorm(nn.Module):
def __init__(
self,
dim: int,
elementwise_affine=True,
eps: float = 1e-6,
device=None,
dtype=None,
):
"""
Initialize the RMSNorm normalization layer.
Args:
dim (int): The dimension of the input tensor.
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
Attributes:
eps (float): A small value added to the denominator for numerical stability.
weight (nn.Parameter): Learnable scaling parameter.
"""
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
def _norm(self, x) -> torch.Tensor:
"""
Apply the RMSNorm normalization to the input tensor.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The normalized tensor.
"""
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
"""
Forward pass through the RMSNorm layer.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The output tensor after applying RMSNorm.
"""
output = self._norm(x.float()).type_as(x)
if hasattr(self, "weight"):
output = output * self.weight
return output
class SelfAttention(nn.Module):
def __init__(self,
hidden_dim,
head_dim,
rope_split: Tuple[int, int, int] = (64, 32, 32),
bias: bool = False,
with_rope: bool = True,
with_qk_norm: bool = True,
attn_type: str = "torch",
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)):
super().__init__()
self.head_dim = head_dim
self.hidden_dim = hidden_dim
self.rope_split = list(rope_split)
self.n_heads = hidden_dim // head_dim
self.wqkv = ReplicatedLinear(hidden_dim, hidden_dim * 3, bias=bias)
self.wo = ReplicatedLinear(hidden_dim, hidden_dim, bias=bias)
self.with_rope = with_rope
self.with_qk_norm = with_qk_norm
if self.with_qk_norm:
self.q_norm = StepVideoRMSNorm(head_dim, elementwise_affine=True)
self.k_norm = StepVideoRMSNorm(head_dim, elementwise_affine=True)
# self.core_attention = self.attn_processor(attn_type=attn_type)
self.parallel = attn_type == 'parallel'
self.attn = DistributedAttention(
num_heads=self.n_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends)
def _apply_rope(self, x: torch.Tensor, cos: torch.Tensor,
sin: torch.Tensor):
"""
x: [B, S, H, D]
cos: [S, D/2] where D = head_dim = sum(self.rope_split)
sin: [S, D/2]
returns x with rotary applied exactly as v0 did
"""
B, S, H, D = x.shape
# 1) split cos/sin per chunk
half_splits = [c // 2
for c in self.rope_split] # [32,16,16] for [64,32,32]
cos_splits = cos.split(half_splits, dim=1)
sin_splits = sin.split(half_splits, dim=1)
outs = []
idx = 0
for (chunk_size, cos_i, sin_i) in zip(self.rope_split, cos_splits,
sin_splits):
# slice the corresponding channels
x_chunk = x[..., idx:idx + chunk_size] # [B,S,H,chunk_size]
idx += chunk_size
# flatten to [S, B*H, chunk_size]
x_flat = rearrange(x_chunk, 'b s h d -> s (b h) d')
# apply rotary on *that* chunk
out_flat = _apply_rotary_emb(x_flat,
cos_i,
sin_i,
is_neox_style=True)
# restore [B,S,H,chunk_size]
out = rearrange(out_flat, 's (b h) d -> b s h d', b=B, h=H)
outs.append(out)
# concatenate back to [B,S,H,D]
return torch.cat(outs, dim=-1)
def forward(self,
x,
cu_seqlens=None,
max_seqlen=None,
rope_positions=None,
cos_sin=None,
attn_mask=None,
mask_strategy=None):
B, S, _ = x.shape
xqkv, _ = self.wqkv(x)
xqkv = xqkv.view(*x.shape[:-1], self.n_heads, 3 * self.head_dim)
q, k, v = torch.split(xqkv, [self.head_dim] * 3, dim=-1) # [B,S,H,D]
if self.with_qk_norm:
q = self.q_norm(q)
k = self.k_norm(k)
if self.with_rope:
if rope_positions is not None:
F, Ht, W = rope_positions
assert F * Ht * W == S, "rope_positions mismatches sequence length"
cos, sin = cos_sin
cos = cos.to(x.device, dtype=x.dtype)
sin = sin.to(x.device, dtype=x.dtype)
q = self._apply_rope(q, cos, sin)
k = self._apply_rope(k, cos, sin)
output, _ = self.attn(q, k, v) # [B,heads,S,D]
output = rearrange(output, 'b s h d -> b s (h d)')
output, _ = self.wo(output)
return output
class CrossAttention(nn.Module):
def __init__(
self,
hidden_dim,
head_dim,
bias=False,
with_qk_norm=True,
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
) -> None:
super().__init__()
self.head_dim = head_dim
self.n_heads = hidden_dim // head_dim
self.wq = ReplicatedLinear(hidden_dim, hidden_dim, bias=bias)
self.wkv = ReplicatedLinear(hidden_dim, hidden_dim * 2, bias=bias)
self.wo = ReplicatedLinear(hidden_dim, hidden_dim, bias=bias)
self.with_qk_norm = with_qk_norm
if self.with_qk_norm:
self.q_norm = StepVideoRMSNorm(head_dim, elementwise_affine=True)
self.k_norm = StepVideoRMSNorm(head_dim, elementwise_affine=True)
self.attn = LocalAttention(
num_heads=self.n_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends)
def forward(self,
x: torch.Tensor,
encoder_hidden_states: torch.Tensor,
attn_mask=None) -> torch.Tensor:
xq, _ = self.wq(x)
xq = xq.view(*xq.shape[:-1], self.n_heads, self.head_dim)
xkv, _ = self.wkv(encoder_hidden_states)
xkv = xkv.view(*xkv.shape[:-1], self.n_heads, 2 * self.head_dim)
xk, xv = torch.split(xkv, [self.head_dim] * 2,
dim=-1) ## seq_len, n, dim
if self.with_qk_norm:
xq = self.q_norm(xq)
xk = self.k_norm(xk)
output = self.attn(xq, xk, xv)
output = rearrange(output, 'b s h d -> b s (h d)')
output, _ = self.wo(output)
return output
class AdaLayerNormSingle(nn.Module):
r"""
Norm layer adaptive layer norm single (adaLN-single).
As proposed in PixArt-Alpha (see: https://arxiv.org/abs/2310.00426; Section 2.3).
Parameters:
embedding_dim (`int`): The size of each embedding vector.
use_additional_conditions (`bool`): To use additional conditions for normalization or not.
"""
def __init__(self, embedding_dim: int, time_step_rescale=1000):
super().__init__()
self.emb = TimestepEmbedder(embedding_dim)
self.silu = nn.SiLU()
self.linear = ReplicatedLinear(embedding_dim,
6 * embedding_dim,
bias=True)
self.time_step_rescale = time_step_rescale ## timestep usually in [0, 1], we rescale it to [0,1000] for stability
def forward(
self,
timestep: torch.Tensor,
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
embedded_timestep = self.emb(timestep * self.time_step_rescale)
out, _ = self.linear(self.silu(embedded_timestep))
return out, embedded_timestep
class StepVideoTransformerBlock(nn.Module):
r"""
A basic Transformer block.
Parameters:
dim (`int`): The number of channels in the input and output.
num_attention_heads (`int`): The number of heads to use for multi-head attention.
attention_head_dim (`int`): The number of channels in each head.
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
num_embeds_ada_norm (:
obj: `int`, *optional*): The number of diffusion steps used during training. See `Transformer2DModel`.
attention_bias (:
obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter.
only_cross_attention (`bool`, *optional*):
Whether to use only cross-attention layers. In this case two cross attention layers are used.
double_self_attention (`bool`, *optional*):
Whether to use two self-attention layers. In this case no cross attention layers are used.
upcast_attention (`bool`, *optional*):
Whether to upcast the attention computation to float32. This is useful for mixed precision training.
norm_elementwise_affine (`bool`, *optional*, defaults to `True`):
Whether to use learnable elementwise affine parameters for normalization.
norm_type (`str`, *optional*, defaults to `"layer_norm"`):
The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
final_dropout (`bool` *optional*, defaults to False):
Whether to apply a final dropout after the last feed-forward layer.
positional_embeddings (`str`, *optional*, defaults to `None`):
The type of positional embeddings to apply to.
num_positional_embeddings (`int`, *optional*, defaults to `None`):
The maximum number of positional embeddings to apply.
"""
def __init__(self,
dim: int,
attention_head_dim: int,
norm_eps: float = 1e-5,
ff_inner_dim: Optional[int] = None,
ff_bias: bool = False,
attention_type: str = 'torch'):
super().__init__()
self.dim = dim
self.norm1 = LayerNormScaleShift(dim,
norm_type="layer",
elementwise_affine=True,
eps=norm_eps)
self.attn1 = SelfAttention(
dim,
attention_head_dim,
bias=False,
with_rope=True,
with_qk_norm=True,
)
self.norm2 = LayerNormScaleShift(dim,
norm_type="layer",
elementwise_affine=True,
eps=norm_eps)
self.attn2 = CrossAttention(dim,
attention_head_dim,
bias=False,
with_qk_norm=True)
self.ff = MLP(input_dim=dim,
mlp_hidden_dim=dim *
4 if ff_inner_dim is None else ff_inner_dim,
act_type="gelu_pytorch_tanh",
bias=ff_bias)
self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5)
@torch.no_grad()
def forward(self,
q: torch.Tensor,
kv: torch.Tensor,
t_expand: torch.LongTensor,
attn_mask=None,
rope_positions: Optional[list] = None,
cos_sin=None,
mask_strategy=None) -> torch.Tensor:
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
torch.clone(chunk)
for chunk in (self.scale_shift_table[None] +
t_expand.reshape(-1, 6, self.dim)).chunk(6, dim=1))
scale_shift_q = self.norm1(q,
scale=scale_msa.squeeze(1),
shift=shift_msa.squeeze(1))
attn_q = self.attn1(scale_shift_q,
rope_positions=rope_positions,
cos_sin=cos_sin,
mask_strategy=mask_strategy)
q = attn_q * gate_msa + q
attn_q = self.attn2(q, kv, attn_mask)
q = attn_q + q
scale_shift_q = self.norm2(q,
scale=scale_mlp.squeeze(1),
shift=shift_mlp.squeeze(1))
ff_output = self.ff(scale_shift_q)
q = ff_output * gate_mlp + q
return q
class StepVideoModel(BaseDiT):
# (Optional) Keep the same attribute for compatibility with splitting, etc.
_fsdp_shard_conditions = [
lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit(),
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
]
_param_names_mapping = StepVideoConfig()._param_names_mapping
_supported_attention_backends = StepVideoConfig(
)._supported_attention_backends
def __init__(self, config: StepVideoConfig) -> None:
super().__init__(config=config)
self.num_attention_heads = config.num_attention_heads
self.attention_head_dim = config.attention_head_dim
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_layers = config.num_layers
self.dropout = config.dropout
self.patch_size = config.patch_size
self.norm_type = config.norm_type
self.norm_elementwise_affine = config.norm_elementwise_affine
self.norm_eps = config.norm_eps
self.use_additional_conditions = config.use_additional_conditions
self.caption_channels = config.caption_channels
self.attention_type = config.attention_type
self.num_channels_latents = config.num_channels_latents
# Compute inner dimension.
self.hidden_size = config.hidden_size
# Image/video patch embedding.
self.pos_embed = PatchEmbed2D(
patch_size=self.patch_size,
in_chans=self.in_channels,
embed_dim=self.hidden_size,
)
self._rope_cache: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
# Transformer blocks.
self.transformer_blocks = nn.ModuleList([
StepVideoTransformerBlock(
dim=self.hidden_size,
attention_head_dim=self.attention_head_dim,
attention_type=self.attention_type)
for _ in range(self.num_layers)
])
# Output blocks.
self.norm_out = LayerNormScaleShift(
self.hidden_size,
norm_type="layer",
eps=self.norm_eps,
elementwise_affine=self.norm_elementwise_affine)
self.scale_shift_table = nn.Parameter(
torch.randn(2, self.hidden_size) / (self.hidden_size**0.5))
self.proj_out = ReplicatedLinear(
self.hidden_size,
self.patch_size * self.patch_size * self.out_channels)
# Time modulation via adaptive layer norm.
self.adaln_single = AdaLayerNormSingle(self.hidden_size)
# Set up caption conditioning.
if isinstance(self.caption_channels, int):
caption_channel = self.caption_channels
else:
caption_channel, clip_channel = self.caption_channels
self.clip_projection = ReplicatedLinear(clip_channel,
self.hidden_size)
self.caption_norm = nn.LayerNorm(
caption_channel,
eps=self.norm_eps,
elementwise_affine=self.norm_elementwise_affine)
self.caption_projection = MLP(input_dim=caption_channel,
mlp_hidden_dim=self.hidden_size,
act_type="gelu_pytorch_tanh")
# Flag to indicate if using parallel attention.
self.parallel = (self.attention_type == "parallel")
self.__post_init__()
def patchfy(self, hidden_states) -> torch.Tensor:
hidden_states = rearrange(hidden_states, 'b f c h w -> (b f) c h w')
hidden_states = self.pos_embed(hidden_states)
return hidden_states
def prepare_attn_mask(self, encoder_attention_mask, encoder_hidden_states,
q_seqlen) -> Tuple[torch.Tensor, torch.Tensor]:
kv_seqlens = encoder_attention_mask.sum(dim=1).int()
mask = torch.zeros([len(kv_seqlens), q_seqlen,
max(kv_seqlens)],
dtype=torch.bool,
device=encoder_attention_mask.device)
encoder_hidden_states = encoder_hidden_states[:, :max(kv_seqlens)]
for i, kv_len in enumerate(kv_seqlens):
mask[i, :, :kv_len] = 1
return encoder_hidden_states, mask
def block_forward(self,
hidden_states,
encoder_hidden_states=None,
t_expand=None,
rope_positions=None,
cos_sin=None,
attn_mask=None,
parallel=True,
mask_strategy=None) -> torch.Tensor:
for i, block in enumerate(self.transformer_blocks):
hidden_states = block(hidden_states,
encoder_hidden_states,
t_expand=t_expand,
attn_mask=attn_mask,
rope_positions=rope_positions,
cos_sin=cos_sin,
mask_strategy=mask_strategy[i])
return hidden_states
def _get_rope(self, rope_positions: tuple[int, int, int],
dtype: torch.dtype, device: torch.device):
F, Ht, W = rope_positions
key = (F, Ht, W, dtype)
if key not in self._rope_cache:
cos, sin = get_rotary_pos_embed(
rope_sizes=(F * get_sequence_model_parallel_world_size(), Ht,
W),
hidden_size=self.hidden_size,
heads_num=self.hidden_size // self.attention_head_dim,
rope_dim_list=(64, 32, 32), # same split you used
rope_theta=1.0e4,
dtype=torch.float32 # build once in fp32
)
# move & cast once
self._rope_cache[key] = (cos.to(device, dtype=dtype),
sin.to(device, dtype=dtype))
return self._rope_cache[key]
@torch.inference_mode()
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
t_expand: Optional[torch.LongTensor] = None,
encoder_hidden_states_2: Optional[torch.Tensor] = None,
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
encoder_attention_mask: Optional[torch.Tensor] = None,
fps: Optional[torch.Tensor] = None,
return_dict: bool = True,
mask_strategy=None,
guidance=None,
):
assert hidden_states.ndim == 5
"hidden_states's shape should be (bsz, f, ch, h ,w)"
frame = hidden_states.shape[2]
hidden_states = rearrange(hidden_states,
'b c f h w -> b f c h w',
f=frame)
if mask_strategy is None:
mask_strategy = [None, None]
bsz, frame, _, height, width = hidden_states.shape
height, width = height // self.patch_size, width // self.patch_size
hidden_states = self.patchfy(hidden_states)
len_frame = hidden_states.shape[1]
t_expand, embedded_timestep = self.adaln_single(t_expand)
encoder_hidden_states = self.caption_projection(
self.caption_norm(encoder_hidden_states))
if encoder_hidden_states_2 is not None and hasattr(
self, 'clip_projection'):
clip_embedding, _ = self.clip_projection(encoder_hidden_states_2)
encoder_hidden_states = torch.cat(
[clip_embedding, encoder_hidden_states], dim=1)
hidden_states = rearrange(hidden_states,
'(b f) l d-> b (f l) d',
b=bsz,
f=frame,
l=len_frame).contiguous()
encoder_hidden_states, attn_mask = self.prepare_attn_mask(
encoder_attention_mask,
encoder_hidden_states,
q_seqlen=frame * len_frame)
cos_sin = self._get_rope((frame, height, width), hidden_states.dtype,
hidden_states.device)
hidden_states = self.block_forward(
hidden_states,
encoder_hidden_states,
t_expand=t_expand,
rope_positions=[frame, height, width],
cos_sin=cos_sin,
attn_mask=attn_mask,
parallel=self.parallel,
mask_strategy=mask_strategy)
hidden_states = rearrange(hidden_states,
'b (f l) d -> (b f) l d',
b=bsz,
f=frame,
l=len_frame)
embedded_timestep = repeat(embedded_timestep, 'b d -> (b f) d',
f=frame).contiguous()
shift, scale = (self.scale_shift_table[None] +
embedded_timestep[:, None]).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states,
shift=shift.squeeze(1),
scale=scale.squeeze(1))
# Modulation
hidden_states, _ = self.proj_out(hidden_states)
# unpatchify
hidden_states = hidden_states.reshape(shape=(-1, height, width,
self.patch_size,
self.patch_size,
self.out_channels))
hidden_states = rearrange(hidden_states, 'n h w p q c -> n c h p w q')
output = hidden_states.reshape(shape=(-1, self.out_channels,
height * self.patch_size,
width * self.patch_size))
output = rearrange(output, '(b f) c h w -> b c f h w', f=frame)
return output
+158 -96
View File
@@ -3,12 +3,16 @@
import math
from typing import List, Optional, Tuple, Union
import numpy as np
import torch
import torch.nn as nn
from fastvideo.v1.attention import DistributedAttention, LocalAttention
from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.configs.sample.wan import WanTeaCacheParams
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.forward_context import get_forward_context
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
ScaleResidual,
ScaleResidualLayerNormScaleShift)
@@ -20,7 +24,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder)
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.models.dits.base import CachableDiT
from fastvideo.v1.platforms import _Backend
@@ -349,116 +353,61 @@ class WanTransformerBlock(nn.Module):
return hidden_states
class WanTransformer3DModel(BaseDiT):
_fsdp_shard_conditions = [
lambda n, m: "blocks" in n and str.isdigit(n.split(".")[-1]),
]
_supported_attention_backends = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_ATTN, _Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
_param_names_mapping = {
r"^patch_embedding\.(.*)$":
r"patch_embedding.proj.\1",
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$":
r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_out.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
r"blocks.\1.ffn.fc_out.\2",
r"blocks\.(\d+)\.norm2\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
}
class WanTransformer3DModel(CachableDiT):
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanVideoConfig()._compile_conditions
_supported_attention_backends = WanVideoConfig(
)._supported_attention_backends
_param_names_mapping = WanVideoConfig()._param_names_mapping
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__()
def __init__(self, config: WanVideoConfig) -> None:
super().__init__(config=config)
inner_dim = num_attention_heads * attention_head_dim
self.hidden_size = inner_dim
self.num_attention_heads = num_attention_heads
self.in_channels = in_channels
self.out_channels = out_channels or in_channels
self.patch_size = patch_size
self.text_len = text_len
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size
self.text_len = config.text_len
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=in_channels,
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=patch_size,
patch_size=config.patch_size,
flatten=False)
# 2. Condition embeddings
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=freq_dim,
text_embed_dim=text_dim,
image_embed_dim=image_dim,
time_freq_dim=config.freq_dim,
text_embed_dim=config.text_dim,
image_embed_dim=config.image_dim,
)
# 3. Transformer blocks
self.blocks = nn.ModuleList([
WanTransformerBlock(inner_dim,
ffn_dim,
num_attention_heads,
qk_norm,
cross_attn_norm,
eps,
added_kv_proj_dim,
config.ffn_dim,
config.num_attention_heads,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{prefix}.blocks.{i}")
for i in range(num_layers)
prefix=f"{config.prefix}.blocks.{i}")
for i in range(config.num_layers)
])
# 4. Output norm & projection
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=eps,
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32)
self.proj_out = nn.Linear(inner_dim,
out_channels * math.prod(patch_size))
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
@@ -474,6 +423,10 @@ class WanTransformer3DModel(BaseDiT):
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs) -> torch.Tensor:
forward_batch = get_forward_context().forward_batch
assert forward_batch is not None
enable_teacache = forward_batch.enable_teacache
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
@@ -517,16 +470,32 @@ class WanTransformer3DModel(BaseDiT):
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
assert encoder_hidden_states.dtype == orig_dtype
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states, timestep_proj,
freqs_cis)
# if caching is enabled, we might be able to skip the forward pass
should_skip_forward = self.should_skip_forward_for_cached_states(
timestep_proj=timestep_proj, temb=temb)
if should_skip_forward:
hidden_states = self.retrieve_cached_states(hidden_states)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis)
# if teacache is enabled, we need to cache the original hidden states
if enable_teacache:
original_hidden_states = hidden_states.clone()
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis)
# if teacache is enabled, we need to cache the original hidden states
if enable_teacache:
self.maybe_cache_states(hidden_states, original_hidden_states)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
@@ -542,3 +511,96 @@ class WanTransformer3DModel(BaseDiT):
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
def maybe_cache_states(self, hidden_states: torch.Tensor,
original_hidden_states: torch.Tensor) -> None:
if self.is_even:
self.previous_residual_even = hidden_states.squeeze(
0) - original_hidden_states
else:
self.previous_residual_odd = hidden_states.squeeze(
0) - original_hidden_states
def should_skip_forward_for_cached_states(self, **kwargs) -> bool:
forward_context = get_forward_context()
forward_batch = forward_context.forward_batch
assert forward_batch is not None
if not forward_batch.enable_teacache:
return False
teacache_params = forward_batch.teacache_params
assert teacache_params is not None, "teacache_params is not initialized"
assert isinstance(
teacache_params,
WanTeaCacheParams), "teacache_params is not a WanTeaCacheParams"
current_timestep = forward_context.current_timestep
num_inference_steps = forward_batch.num_inference_steps
# initialize the coefficients, cutoff_steps, and ret_steps
coefficients = teacache_params.coefficients
use_ret_steps = teacache_params.use_ret_steps
cutoff_steps = teacache_params.get_cutoff_steps(num_inference_steps)
ret_steps = teacache_params.ret_steps
teacache_thresh = teacache_params.teacache_thresh
if current_timestep == 0:
self.cnt = 0
timestep_proj = kwargs["timestep_proj"]
temb = kwargs["temb"]
modulated_inp = timestep_proj if use_ret_steps else temb
if self.cnt % 2 == 0: # even -> condition
self.is_even = True
if self.cnt < ret_steps or self.cnt >= cutoff_steps:
self.should_calc_even = True
self.accumulated_rel_l1_distance_even = 0
else:
assert self.previous_e0_even is not None, "previous_e0_even is not initialized"
assert self.accumulated_rel_l1_distance_even is not None, "accumulated_rel_l1_distance_even is not initialized"
rescale_func = np.poly1d(coefficients)
self.accumulated_rel_l1_distance_even += rescale_func(
((modulated_inp - self.previous_e0_even).abs().mean() /
self.previous_e0_even.abs().mean()).cpu().item())
if self.accumulated_rel_l1_distance_even < teacache_thresh:
self.should_calc_even = False
else:
self.should_calc_even = True
self.accumulated_rel_l1_distance_even = 0
self.previous_e0_even = modulated_inp.clone()
else: # odd -> unconditon
self.is_even = False
if self.cnt < ret_steps or self.cnt >= cutoff_steps:
self.should_calc_odd = True
self.accumulated_rel_l1_distance_odd = 0
else:
assert self.previous_e0_odd is not None, "previous_e0_odd is not initialized"
assert self.accumulated_rel_l1_distance_odd is not None, "accumulated_rel_l1_distance_odd is not initialized"
rescale_func = np.poly1d(coefficients)
self.accumulated_rel_l1_distance_odd += rescale_func(
((modulated_inp - self.previous_e0_odd).abs().mean() /
self.previous_e0_odd.abs().mean()).cpu().item())
if self.accumulated_rel_l1_distance_odd < teacache_thresh:
self.should_calc_odd = False
else:
self.should_calc_odd = True
self.accumulated_rel_l1_distance_odd = 0
self.previous_e0_odd = modulated_inp.clone()
self.cnt += 1
should_skip_forward = False
if self.is_even:
if not self.should_calc_even:
should_skip_forward = True
else:
if not self.should_calc_odd:
should_skip_forward = True
return should_skip_forward
def retrieve_cached_states(self,
hidden_states: torch.Tensor) -> torch.Tensor:
if self.is_even:
return hidden_states + self.previous_residual_even
else:
return hidden_states + self.previous_residual_odd
+41 -6
View File
@@ -1,22 +1,57 @@
from typing import Tuple
from abc import ABC, abstractmethod
from typing import Optional, Tuple
import torch
from torch import nn
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
ImageEncoderConfig,
TextEncoderConfig)
from fastvideo.v1.platforms import _Backend
class BaseEncoder(nn.Module):
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.TORCH_SDPA, )
class TextEncoder(nn.Module, ABC):
_supported_attention_backends: Tuple[
_Backend, ...] = TextEncoderConfig()._supported_attention_backends
def __init__(self, *args, **kwargs) -> None:
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__()
self.config = config
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
)
def forward(self, *args, **kwargs):
@abstractmethod
def forward(self,
input_ids: Optional[torch.Tensor],
position_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.Tensor] = None,
output_hidden_states: Optional[bool] = None,
**kwargs) -> BaseEncoderOutput:
pass
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
return self._supported_attention_backends
class ImageEncoder(nn.Module, ABC):
_supported_attention_backends: Tuple[
_Backend, ...] = ImageEncoderConfig()._supported_attention_backends
def __init__(self, config: ImageEncoderConfig) -> None:
super().__init__()
self.config = config
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
)
@abstractmethod
def forward(self, pixel_values: torch.Tensor,
**kwargs) -> BaseEncoderOutput:
pass
@property
+40
View File
@@ -0,0 +1,40 @@
# type: ignore
import os
import torch
import torch.nn as nn
from transformers import BertModel, BertTokenizer
class HunyuanClip(nn.Module):
"""
Hunyuan clip code copied from https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/hunyuandit/pipeline_hunyuandit.py
hunyuan's clip used BertModel and BertTokenizer, so we copy it.
"""
def __init__(self, model_dir, max_length=77):
super().__init__()
self.max_length = max_length
self.tokenizer = BertTokenizer.from_pretrained(
os.path.join(model_dir, 'tokenizer'))
self.text_encoder = BertModel.from_pretrained(
os.path.join(model_dir, 'clip_text_encoder'))
@torch.no_grad
def forward(self, prompts, with_mask=True):
self.device = next(self.text_encoder.parameters()).device
text_inputs = self.tokenizer(
prompts,
padding="max_length",
max_length=self.max_length,
truncation=True,
return_attention_mask=True,
return_tensors="pt",
)
prompt_embeds = self.text_encoder(
text_inputs.input_ids.to(self.device),
attention_mask=text_inputs.attention_mask.to(self.device)
if with_mask else None,
)
return prompt_embeds.last_hidden_state, prompt_embeds.pooler_output
+46 -104
View File
@@ -3,62 +3,32 @@
# Adapted from transformers: https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py
"""Minimal implementation of CLIPVisionModel intended to be only used
within a vision language model."""
from typing import Iterable, Optional, Set, Tuple, Union, cast
from typing import Iterable, Optional, Set, Tuple, Union
import torch
import torch.nn as nn
from transformers import CLIPTextConfig, CLIPVisionConfig
from transformers.modeling_outputs import BaseModelOutputWithPooling
from vllm.model_executor.models.interfaces import SupportsQuant
# from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask
from fastvideo.v1.attention import LocalAttention
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
CLIPTextConfig,
CLIPVisionConfig)
from fastvideo.v1.configs.quantization import QuantizationConfig
from fastvideo.v1.distributed import (divide,
get_tensor_model_parallel_world_size)
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
RowParallelLinear)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.encoders.base import BaseEncoder
from fastvideo.v1.models.encoders.vision import (VisionEncoderInfo,
resolve_visual_encoder_outputs)
from fastvideo.v1.models.encoders.base import ImageEncoder, TextEncoder
from fastvideo.v1.models.encoders.vision import resolve_visual_encoder_outputs
# TODO: support quantization
# from vllm.model_executor.layers.quantization import QuantizationConfig
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
from fastvideo.v1.platforms import _Backend
logger = init_logger(__name__)
class QuantizationConfig:
pass
class CLIPEncoderInfo(VisionEncoderInfo[CLIPVisionConfig]):
def get_num_image_tokens(
self,
*,
image_width: int,
image_height: int,
) -> int:
return self.get_patch_grid_length()**2 + 1
def get_max_image_tokens(self) -> int:
return self.get_patch_grid_length()**2 + 1
def get_image_size(self) -> int:
return cast(int, self.vision_config.image_size)
def get_patch_size(self) -> int:
return cast(int, self.vision_config.patch_size)
def get_patch_grid_length(self) -> int:
image_size, patch_size = self.get_image_size(), self.get_patch_size()
assert image_size % patch_size == 0
return image_size // patch_size
# Adapted from https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py#L164 # noqa
class CLIPVisionEmbeddings(nn.Module):
@@ -158,7 +128,7 @@ class CLIPAttention(nn.Module):
def __init__(
self,
config: CLIPVisionConfig,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
@@ -193,13 +163,13 @@ class CLIPAttention(nn.Module):
self.tp_size = get_tensor_model_parallel_world_size()
self.num_heads_per_partition = divide(self.num_heads, self.tp_size)
self.attn = LocalAttention(self.num_heads_per_partition,
self.head_dim,
self.num_heads_per_partition,
softmax_scale=self.scale,
causal=True,
supported_attention_backends=self.config.
supported_attention_backends)
self.attn = LocalAttention(
self.num_heads_per_partition,
self.head_dim,
self.num_heads_per_partition,
softmax_scale=self.scale,
causal=True,
supported_attention_backends=config._supported_attention_backends)
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
return tensor.view(bsz, seq_len, self.num_heads,
@@ -239,7 +209,7 @@ class CLIPMLP(nn.Module):
def __init__(
self,
config: CLIPVisionConfig,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
@@ -269,7 +239,7 @@ class CLIPEncoderLayer(nn.Module):
def __init__(
self,
config: CLIPTextConfig,
config: Union[CLIPTextConfig, CLIPVisionConfig],
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
@@ -314,7 +284,7 @@ class CLIPEncoder(nn.Module):
def __init__(
self,
config: CLIPVisionConfig,
config: Union[CLIPVisionConfig, CLIPTextConfig],
quant_config: Optional[QuantizationConfig] = None,
num_hidden_layers_override: Optional[int] = None,
prefix: str = "",
@@ -356,7 +326,6 @@ class CLIPTextTransformer(nn.Module):
def __init__(self,
config: CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None,
*,
num_hidden_layers_override: Optional[int] = None,
prefix: str = ""):
super().__init__()
@@ -377,27 +346,21 @@ class CLIPTextTransformer(nn.Module):
# For `pooled_output` computation
self.eos_token_id = config.eos_token_id
# For attention mask, it differs between `flash_attention_2` and other attention implementations
self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2"
def forward(
self,
input_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
input_ids: Optional[torch.Tensor],
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.Tensor] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, BaseModelOutputWithPooling]:
) -> BaseEncoderOutput:
r"""
Returns:
"""
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (output_hidden_states
if output_hidden_states is not None else
self.config.output_hidden_states)
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
if input_ids is None:
raise ValueError("You have to specify input_ids")
@@ -456,11 +419,7 @@ class CLIPTextTransformer(nn.Module):
) == self.eos_token_id).int().argmax(dim=-1),
]
if not return_dict:
return (last_hidden_state, pooled_output) + encoder_outputs[1:]
# return last_hidden_state
return BaseModelOutputWithPooling(
return BaseEncoderOutput(
last_hidden_state=last_hidden_state,
pooler_output=pooled_output,
hidden_states=encoder_outputs,
@@ -468,42 +427,34 @@ class CLIPTextTransformer(nn.Module):
)
class CLIPTextModel(BaseEncoder):
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
class CLIPTextModel(TextEncoder):
def __init__(
self,
config: CLIPTextConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
self.config.supported_attention_backends = self._supported_attention_backends
super().__init__(config)
self.text_model = CLIPTextTransformer(config=config,
quant_config=quant_config,
prefix=prefix)
quant_config=config.quant_config,
prefix=config.prefix)
def forward(
self,
input_ids: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
input_ids: Optional[torch.Tensor],
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.Tensor] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, BaseModelOutputWithPooling]:
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
**kwargs,
) -> BaseEncoderOutput:
return self.text_model(
outputs: BaseEncoderOutput = self.text_model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=None,
)
return outputs
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
@@ -548,7 +499,6 @@ class CLIPVisionTransformer(nn.Module):
self,
config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None,
*,
num_hidden_layers_override: Optional[int] = None,
require_post_norm: Optional[bool] = None,
prefix: str = "",
@@ -615,37 +565,29 @@ class CLIPVisionTransformer(nn.Module):
return encoder_outputs
class CLIPVisionModel(BaseEncoder, SupportsQuant):
class CLIPVisionModel(ImageEncoder):
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)
def __init__(
self,
config: CLIPVisionConfig,
quant_config: Optional[QuantizationConfig] = None,
*,
num_hidden_layers_override: Optional[int] = None,
require_post_norm: Optional[bool] = None,
prefix: str = "",
) -> None:
super().__init__()
self.config = config
self.config.supported_attention_backends = self._supported_attention_backends
def __init__(self, config: CLIPVisionConfig) -> None:
super().__init__(config)
self.vision_model = CLIPVisionTransformer(
config=config,
quant_config=quant_config,
num_hidden_layers_override=num_hidden_layers_override,
require_post_norm=require_post_norm,
prefix=f"{prefix}.vision_model")
quant_config=config.quant_config,
num_hidden_layers_override=config.num_hidden_layers_override,
require_post_norm=config.require_post_norm,
prefix=f"{config.prefix}.vision_model")
def forward(
self,
pixel_values: torch.Tensor,
feature_sample_layers: Optional[list[int]] = None,
) -> torch.Tensor:
return self.vision_model(pixel_values, feature_sample_layers)
**kwargs,
) -> BaseEncoderOutput:
last_hidden_state = self.vision_model(pixel_values,
feature_sample_layers)
return BaseEncoderOutput(last_hidden_state=last_hidden_state)
@property
def device(self):

Some files were not shown because too many files have changed in this diff Show More