Compare commits
86
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e4ceadb5d5 | ||
|
|
8f8ce6d9e1 | ||
|
|
e3d0cbe185 | ||
|
|
d5ec468d43 | ||
|
|
e55fa6e5dc | ||
|
|
61b6ddeee1 | ||
|
|
a9a000f45d | ||
|
|
66b8b8561e | ||
|
|
6684872616 | ||
|
|
7f654e3332 | ||
|
|
8631c1b806 | ||
|
|
5357e12b5a | ||
|
|
bdfdf1dfee | ||
|
|
d156461785 | ||
|
|
6edf113838 | ||
|
|
dcf7738cbc | ||
|
|
b2ebaaf865 | ||
|
|
7768bb80f6 | ||
|
|
a335811869 | ||
|
|
357b0533fe | ||
|
|
2ec3732758 | ||
|
|
a004408a93 | ||
|
|
007e237e69 | ||
|
|
8e18dc9f71 | ||
|
|
7ab32539af | ||
|
|
6ef8fcb61d | ||
|
|
016e24da63 | ||
|
|
85b8717545 | ||
|
|
657fd745e1 | ||
|
|
12647457a7 | ||
|
|
298f74f956 | ||
|
|
ee8babb298 | ||
|
|
60295cc03f | ||
|
|
1572e13b6e | ||
|
|
a157275b4c | ||
|
|
c4dbe7dac3 | ||
|
|
d39591108e | ||
|
|
ace6e971e5 | ||
|
|
b4f6758253 | ||
|
|
535d29b392 | ||
|
|
b4255517e0 | ||
|
|
6eeb60613f | ||
|
|
53d2c7791f | ||
|
|
53cb693dca | ||
|
|
6f72d24876 | ||
|
|
d1459e9976 | ||
|
|
59ab481eb1 | ||
|
|
0cf001986a | ||
|
|
51956369a5 | ||
|
|
4b0970cbbf | ||
|
|
94bf47a572 | ||
|
|
1a3ac9074b | ||
|
|
fb0581d5b0 | ||
|
|
6c74ab4132 | ||
|
|
9a91021c56 | ||
|
|
51c94d6a73 | ||
|
|
dba38dbc03 | ||
|
|
2034cc3c4f | ||
|
|
c69afce2f6 | ||
|
|
b08e758eb3 | ||
|
|
3f3462d7ce | ||
|
|
c9c47dd89c | ||
|
|
9c4ef7c2f1 | ||
|
|
f25eb4b905 | ||
|
|
048d55ccbb | ||
|
|
a271c55fe4 | ||
|
|
f663ae0d8a | ||
|
|
c0911aa3dd | ||
|
|
5f59687ae7 | ||
|
|
5adbc81cdc | ||
|
|
f26d5c37c1 | ||
|
|
6a4ef42378 | ||
|
|
f1098c77dc | ||
|
|
0405b618f8 | ||
|
|
eac79b753f | ||
|
|
4d58cf20d0 | ||
|
|
52c93ecc9d | ||
|
|
42d63166ac | ||
|
|
6db20345a2 | ||
|
|
ad27ea596c | ||
|
|
9aadb4bf8c | ||
|
|
bd941df271 | ||
|
|
8a73876d3b | ||
|
|
1483a1138a | ||
|
|
5e243d8292 | ||
|
|
b0c66d3200 |
@@ -4,14 +4,6 @@ title: "[Bug] "
|
||||
labels: ['Bug']
|
||||
|
||||
body:
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/env_utils.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Describe the bug
|
||||
@@ -25,5 +17,13 @@ body:
|
||||
What command or script did you run? Which **model** are you using?
|
||||
placeholder: |
|
||||
A placeholder for the command.
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
@@ -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."
|
||||
@@ -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
|
||||
@@ -8,12 +8,14 @@ on:
|
||||
- main
|
||||
paths:
|
||||
- "docs/**/*.md"
|
||||
- "fastvideo/v1/examples/**/*.py"
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
types: [opened, ready_for_review, synchronize, reopened]
|
||||
paths:
|
||||
- "docs/**/*.md"
|
||||
- "fastvideo/v1/examples/**/*.py"
|
||||
|
||||
# Allows you to run this workflow manually from the Actions tab
|
||||
workflow_dispatch:
|
||||
|
||||
+60
-170
@@ -77,199 +77,89 @@ jobs:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
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 . && 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 . && 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 . && 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 . && 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 +179,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
|
||||
|
||||
@@ -10,7 +10,7 @@ jobs:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
python-version: "3.12"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
|
||||
@@ -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
|
||||
+5
-2
@@ -27,7 +27,6 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
**.json
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
@@ -40,6 +39,7 @@ eggs/
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
docs/source/getting_started/examples/
|
||||
docs/source/inference/examples/
|
||||
|
||||
# VSCode
|
||||
.vscode/
|
||||
@@ -55,4 +55,7 @@ docs/source/getting_started/examples/
|
||||
*.pkl
|
||||
|
||||
# Reference videos
|
||||
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
|
||||
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
|
||||
|
||||
# Static images
|
||||
!docs/source/_static/images/**/*.png
|
||||
|
||||
+10
-7
@@ -19,8 +19,11 @@ exclude: |
|
||||
fastvideo/sample/.*|
|
||||
fastvideo/train\.py|
|
||||
fastvideo/utils/.*|
|
||||
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
|
||||
@@ -30,7 +33,7 @@ repos:
|
||||
args: [--in-place, --verbose]
|
||||
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.11.4
|
||||
rev: v0.11.12
|
||||
hooks:
|
||||
- id: ruff
|
||||
args: [--output-format, github, --fix]
|
||||
@@ -40,12 +43,12 @@ repos:
|
||||
- id: codespell
|
||||
additional_dependencies: ['tomli']
|
||||
args: ['--toml', 'pyproject.toml']
|
||||
# - repo: https://github.com/PyCQA/isort
|
||||
# rev: 0a0b7a830386ba6a31c2ec8316849ae4d1b8240d # 6.0.0
|
||||
# hooks:
|
||||
# - id: isort
|
||||
- repo: https://github.com/PyCQA/isort
|
||||
rev: 6.0.1
|
||||
hooks:
|
||||
- id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.29
|
||||
rev: v0.9.30
|
||||
hooks:
|
||||
- id: pymarkdown
|
||||
args: [fix]
|
||||
|
||||
@@ -2,241 +2,127 @@
|
||||
<img src=assets/logo.jpg width="30%"/>
|
||||
</div>
|
||||
|
||||
FastVideo is a lightweight framework for accelerating large video diffusion models.
|
||||
**FastVideo is a unified framework for accelerated video generation.**
|
||||
|
||||
It features a clean, consistent API that works across popular video models, making it easier for developers to author new models and incorporate system- or kernel-level optimizations.
|
||||
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
|
||||
|
||||
<p align="center">
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</b> </a> |
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</b> </a> |
|
||||
</p>
|
||||
|
||||
https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1
|
||||
<div align="center">
|
||||
<img src=assets/perf.png width="90%"/>
|
||||
</div>
|
||||
|
||||
FastVideo currently offers: (with more to come)
|
||||
## Key Features
|
||||
|
||||
- [NEW!] [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
FastVideo has the following features:
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- Cutting edge models
|
||||
- Wan2.1 T2V, I2V
|
||||
- HunyuanVideo
|
||||
- FastHunyuan: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- StepVideo T2V
|
||||
- Distillation support
|
||||
- Recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
|
||||
|
||||
Dev in progress and highly experimental.
|
||||
## Getting Started
|
||||
We recommend using an environment manager such as `Conda` to create a clean environment:
|
||||
|
||||
## Change Log
|
||||
- ```2025/02/20```: FastVideo now supports STA on [StepVideo](https://github.com/stepfun-ai/Step-Video-T2V) with 3.4X speedup!
|
||||
- ```2025/02/18```: Release the inference code and kernel for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- ```2025/01/13```: Support Lora finetuning for HunyuanVideo.
|
||||
- ```2024/12/25```: Enable single 4090 inference for `FastHunyuan`, please rerun the installation steps to update the environment.
|
||||
- ```2024/12/17```: `FastVideo` v1.0 is released.
|
||||
|
||||
## 🔧 Installation from source
|
||||
The code is tested on Python 3.10.0, CUDA 12.4 and H100.
|
||||
|
||||
```
|
||||
# Clone FastVideo
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
```bash
|
||||
# Create and activate a new conda environment
|
||||
conda create -n fastvideo python=3.12
|
||||
conda activate fastvideo
|
||||
|
||||
# Install FastVideo
|
||||
pip install -e .
|
||||
|
||||
# Install Flash Attention (optional)
|
||||
pip install flash-attn==2.7.0.post2
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
To try Sliding Tile Attention (optional), please follow the instruction in [csrc/sliding_tile_attention/README.md](csrc/sliding_tile_attention/README.md) to install STA.
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) for more detailed installation instructions.
|
||||
|
||||
You can also install the Sliding Tile Attention package using
|
||||
## Inference
|
||||
### 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:
|
||||
|
||||
```
|
||||
pip install st_attn==0.0.3
|
||||
```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()
|
||||
```
|
||||
|
||||
## 🚀 Inference
|
||||
### 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).
|
||||
Run the script with:
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_stepvideo_STA.sh # Inference stepvideo with STA
|
||||
sh scripts/inference/inference_stepvideo.sh # Inference original stepvideo
|
||||
python example.py
|
||||
```
|
||||
|
||||
### Inference HunyuanVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html).
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
|
||||
```
|
||||
### Other docs:
|
||||
|
||||
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
|
||||
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview.html)
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
|
||||
|
||||
```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.
|
||||
|
||||
### 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).
|
||||
|
||||
### 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
|
||||
```
|
||||
|
||||
## 🎯 Distill
|
||||
Our distillation recipe is based on [Phased Consistency Model](https://github.com/G-U-N/Phased-Consistency-Model). We did not find significant improvement using multi-phase distillation, so we keep the one phase setup similar to the original latent consistency model's recipe.
|
||||
We use the [MixKit](https://huggingface.co/datasets/LanguageBind/Open-Sora-Plan-v1.1.0/tree/main/all_mixkit) dataset for distillation. To avoid running the text encoder and VAE during training, we preprocess all data to generate text embeddings and VAE latents.
|
||||
Preprocessing instructions can be found [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide preprocessed data that can be downloaded directly using the following command:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/HD-Mixkit-Finetune-Hunyuan --local_dir=data/HD-Mixkit-Finetune-Hunyuan --repo_type=dataset
|
||||
```
|
||||
|
||||
Next, download the original model weights with:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model # original hunyuan
|
||||
python scripts/huggingface/download_hf.py --repo_id=genmo/mochi-1-preview --local_dir=data/mochi --repo_type=model # original mochi
|
||||
```
|
||||
|
||||
To launch the distillation process, use the following commands:
|
||||
|
||||
```
|
||||
bash scripts/distill/distill_hunyuan.sh # for hunyuan
|
||||
bash scripts/distill/distill_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
We also provide an optional script for distillation with adversarial loss, located at `fastvideo/distill_adv.py`. Although we tried adversarial loss, we did not observe significant improvements.
|
||||
## Finetune
|
||||
### ⚡ Full Finetune
|
||||
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
|
||||
```
|
||||
|
||||
Download the original model weights as specified in [Distill Section](#-distill):
|
||||
|
||||
Then you can run the finetune with:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
### ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight --local_dir=data/Hunyuan-Black-Myth-Wukong-lora-weight --repo_type=model
|
||||
```
|
||||
|
||||
#### Minimum Hardware Requirement
|
||||
- 40 GB GPU memory each for 2 GPUs with lora.
|
||||
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
|
||||
|
||||
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
|
||||
|
||||
#### Dataset Preparation
|
||||
We provide scripts to better help you get started to train on your own characters!
|
||||
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
|
||||
|
||||
```
|
||||
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
|
||||
```
|
||||
|
||||
Also, we provide script to resize your videos:
|
||||
|
||||
```
|
||||
python scripts/data_preprocess/resize_videos.py
|
||||
```
|
||||
|
||||
#### Finetuning
|
||||
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
|
||||
```
|
||||
|
||||
#### Inference
|
||||
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
|
||||
|
||||
```
|
||||
bash scripts/inference/inference_hunyuan_hf.sh
|
||||
```
|
||||
|
||||
**We also provide scripts for Mochi in the same directory.**
|
||||
|
||||
#### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_hunyuan.sh
|
||||
bash scripts/finetune/finetune_mochi_lora_mix.sh
|
||||
```
|
||||
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetuning.html)
|
||||
|
||||
## 📑 Development Plan
|
||||
|
||||
- More distillation methods
|
||||
- [ ] Add Distribution Matching Distillation
|
||||
<!-- - More distillation methods -->
|
||||
<!-- - [ ] Add Distribution Matching Distillation -->
|
||||
- More models support
|
||||
- [ ] Add CogvideoX model
|
||||
- Code update
|
||||
- [ ] fp8 support
|
||||
- [ ] faster load model and save model support
|
||||
<!-- - [ ] Add CogvideoX model -->
|
||||
- [x] Add StepVideo to V1
|
||||
- Optimization features
|
||||
- [x] Teacache in V1
|
||||
- [x] SageAttention in V1
|
||||
- Code updates
|
||||
- [x] V1 Configuration API
|
||||
- [ ] Support Training in V1
|
||||
<!-- - [ ] fp8 support -->
|
||||
<!-- - [ ] faster load model and save model support -->
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please run `bash format.sh --all` before submitting a pull request.
|
||||
|
||||
## 🔧 Testing
|
||||
Run `pytest` to verify the data preprocessing, checkpoint saving, and sequence parallel pipelines. We recommend adding corresponding test cases in the `test` folder to support your contribution.
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/developer_guide/overview.html)
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects: [PCM](https://github.com/G-U-N/Phased-Consistency-Model), [diffusers](https://github.com/huggingface/diffusers), [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan), and [xDiT](https://github.com/xdit-project/xDiT).
|
||||
We learned and reused code from the following projects:
|
||||
- [PCM](https://github.com/G-U-N/Phased-Consistency-Model)
|
||||
- [diffusers](https://github.com/huggingface/diffusers)
|
||||
- [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan)
|
||||
- [xDiT](https://github.com/xdit-project/xDiT)
|
||||
- [vLLM](https://github.com/vllm-project/vllm)
|
||||
- [SGLang](https://github.com/sgl-project/sglang)
|
||||
|
||||
We thank MBZUAI and Anyscale for their support throughout this project.
|
||||
We thank MBZUAI and [Anyscale](https://www.anyscale.com/) for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
If you use FastVideo for your research, please cite our paper:
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
|
After Width: | Height: | Size: 303 KiB |
@@ -9,7 +9,7 @@ target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "st_attn"
|
||||
VERSION = "0.0.3"
|
||||
VERSION = "0.0.4"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
|
||||
|
||||
@@ -10,7 +10,7 @@
|
||||
|
||||
#ifdef TK_COMPILE_ATTN
|
||||
extern torch::Tensor sta_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
|
||||
);
|
||||
#endif
|
||||
|
||||
|
||||
@@ -4,8 +4,13 @@ import torch
|
||||
from st_attn_cuda import sta_fwd
|
||||
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True):
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, img_latent_shape='30*48*80'):
|
||||
seq_length = q_all.shape[2]
|
||||
img_latent_shape_mapping = {
|
||||
'30x48x80':1,
|
||||
'36x48x48':2,
|
||||
'18x48x80':3,
|
||||
}
|
||||
if has_text:
|
||||
assert q_all.shape[
|
||||
2] >= 115200, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
|
||||
@@ -17,8 +22,14 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
|
||||
k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
|
||||
v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
|
||||
else:
|
||||
assert q_all.shape[2] == 82944
|
||||
if img_latent_shape == '36x48x48': # Stepvideo 204x768x68
|
||||
assert q_all.shape[2] == 82944
|
||||
elif img_latent_shape == '18x48x80': # Wan 69x768x1280
|
||||
assert q_all.shape[2] == 69120
|
||||
else:
|
||||
raise ValueError(f"Unsupported {img_latent_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
|
||||
|
||||
kernel_aspect_ratio_flag = img_latent_shape_mapping[img_latent_shape]
|
||||
hidden_states = torch.empty_like(q_all)
|
||||
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
|
||||
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
|
||||
@@ -29,7 +40,7 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
|
||||
head_index:head_index + 1],
|
||||
hidden_states[batch:batch + 1, head_index:head_index + 1])
|
||||
|
||||
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text)
|
||||
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
|
||||
if has_text:
|
||||
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True)
|
||||
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
|
||||
return hidden_states[:, :, :seq_length]
|
||||
|
||||
@@ -359,7 +359,7 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
|
||||
#include <iostream>
|
||||
|
||||
torch::Tensor
|
||||
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text)
|
||||
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
@@ -558,123 +558,267 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
|
||||
} else {
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
if (kernel_aspect_ratio_flag == 2){
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,1, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,1, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
}else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 3){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
}else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 3){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size ==6 && kernel_h_size == 3 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size ==6 && kernel_h_size == 3 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else {
|
||||
// print error
|
||||
std::cout << "Invalid kernel size" << std::endl;
|
||||
//print kernel size
|
||||
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
|
||||
}
|
||||
}
|
||||
else if (kernel_aspect_ratio_flag == 3) {
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else {
|
||||
// print error
|
||||
std::cout << "Invalid kernel size" << std::endl;
|
||||
//print kernel size
|
||||
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 2, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,1, 1, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 0, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,1, 0, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,1, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 1, 5, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,0, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 2, 3, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,0, 3, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else {
|
||||
// print error
|
||||
std::cout << "Invalid kernel size" << std::endl;
|
||||
//print kernel size
|
||||
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
|
||||
}
|
||||
}
|
||||
|
||||
else {
|
||||
std::cout << "Unsupported kernel_aspect_ratio_flag: " << kernel_aspect_ratio_flag << std::endl;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -2,7 +2,6 @@ import torch
|
||||
from flex_sta_ref import get_sliding_tile_attention_mask
|
||||
from st_attn import sliding_tile_attention
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
# from flash_attn_interface import flash_attn_func
|
||||
from tqdm import tqdm
|
||||
|
||||
flex_attention = torch.compile(flex_attention, dynamic=False)
|
||||
@@ -23,7 +22,7 @@ def h100_fwd_kernel_test(Q, K, V, kernel_size):
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
magnitude = torch.linalg.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
@@ -2,7 +2,7 @@ FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
WORKDIR /app
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
@@ -29,9 +29,20 @@ 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 . .
|
||||
|
||||
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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -22,3 +22,4 @@ help:
|
||||
clean:
|
||||
@$(SPHINXBUILD) -M clean "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
rm -rf "$(SOURCEDIR)/getting_started/examples"
|
||||
rm -rf "$(SOURCEDIR)/inference/examples"
|
||||
|
||||
@@ -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
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 303 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 18 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 27 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 40 KiB |
@@ -34,6 +34,6 @@
|
||||
}
|
||||
</style>
|
||||
|
||||
<div class="notification-bar">
|
||||
<!-- <div class="notification-bar">
|
||||
<p>You are viewing the latest developer preview docs. <a href="https://docs.vllm.ai/en/stable/">Click here</a> to view docs for the latest stable release.</p>
|
||||
</div>
|
||||
</div> -->
|
||||
|
||||
@@ -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
|
||||
```
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
(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:
|
||||
|
||||
**Image:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
|
||||
|
||||
## Starting the container
|
||||
|
||||
```bash
|
||||
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
|
||||
```
|
||||
|
||||
This will:
|
||||
|
||||
- Start the container with GPU access
|
||||
- Drop you into a shell with the `fastvideo-dev` Conda environment preconfigured
|
||||
|
||||
## Using the container
|
||||
|
||||
```bash
|
||||
# Conda environment should already be active
|
||||
# FastVideo package installed in editable mode
|
||||
|
||||
# Pull the latest changes from remote
|
||||
cd /FastVideo
|
||||
git pull
|
||||
|
||||
# Run linters and tests
|
||||
pre-commit run --all-files
|
||||
pytest tests/
|
||||
```
|
||||
@@ -0,0 +1,13 @@
|
||||
(developer-env)
|
||||
|
||||
# 🧰 Developer Environment
|
||||
|
||||
Accelerate your FastVideo development workflow by leveraging Docker images and cloud GPUs for efficient experimentation and reproducible environments.
|
||||
|
||||
:::{toctree}
|
||||
:caption: Contents
|
||||
:maxdepth: 1
|
||||
|
||||
docker
|
||||
runpod
|
||||
:::
|
||||
@@ -0,0 +1,52 @@
|
||||
(runpod)=
|
||||
|
||||
# 📦 Developing FastVideo on RunPod
|
||||
|
||||
You can easily use the FastVideo Docker image as a custom container on [RunPod](https://www.runpod.io) for development or experimentation.
|
||||
|
||||
## Creating a new pod
|
||||
|
||||
Choose a GPU that supports CUDA 12.4
|
||||
|
||||

|
||||
|
||||
When creating your pod template, use this image:
|
||||
|
||||
```
|
||||
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
|
||||
```
|
||||
|
||||
Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.io/pods/configuration/use-ssh)):
|
||||
|
||||
```bash
|
||||
bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
|
||||
```
|
||||
|
||||

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

|
||||
|
||||
## Working with the pod
|
||||
|
||||
After SSH'ing into your pod, you'll find the `fastvideo-dev` Conda environment already activated.
|
||||
|
||||
To pull in the latest changes from the GitHub repo:
|
||||
|
||||
```bash
|
||||
cd /FastVideo
|
||||
git pull
|
||||
```
|
||||
|
||||
`If you have a persistent volume and want to keep your code changes, you can move /FastVideo to /workspace/FastVideo, or simply clone the repository there.`
|
||||
|
||||
Run your development workflows as usual:
|
||||
|
||||
```bash
|
||||
# Run linters
|
||||
pre-commit run --all-files
|
||||
|
||||
# Run tests
|
||||
pytest tests/
|
||||
```
|
||||
@@ -1,6 +1,6 @@
|
||||
(developer-guide)=
|
||||
(developer-overview)=
|
||||
|
||||
# Contributing to FastVideo
|
||||
# 🛠️ Contributing to FastVideo
|
||||
|
||||
Thank you for your interest in contributing to FastVideo. We want to make the process as smooth for you as possible and this is a guide to help get you started!
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,410 @@
|
||||
# 🔍 FastVideo Overview
|
||||
|
||||
This document outlines FastVideo's architecture for developers interested in framework internals or contributions. It serves as an onboarding guide for new contributors by providing an overview of the most important directories and files within the `fastvideo/v1/` codebase.
|
||||
|
||||
## Table of Contents - V1 Directory Structure and Files
|
||||
|
||||
- [`fastvideo/v1/pipelines/`](#design-pipeline-system) - Core diffusion pipeline components
|
||||
- [`fastvideo/v1/models/`](#design-model-components) - Model implementations
|
||||
- [`dits/`](#design-transformer-models) - Transformer-based diffusion models
|
||||
- [`vaes/`](#design-vae-variational-auto-encoder) - Variational autoencoders
|
||||
- [`encoders/`](#design-text-and-image-encoders) - Text and image encoders
|
||||
- [`schedulers/`](#design-schedulers) - Diffusion schedulers
|
||||
- [`fastvideo/v1/attention/`](#design-optimized-attention) - Optimized attention implementations
|
||||
- [`fastvideo/v1/distributed/`](#design-distributed-processing) - Distributed computing utilities
|
||||
- [`fastvideo/v1/layers/`](#design-tensor-parallelism) - Custom neural network layers
|
||||
- [`fastvideo/v1/platforms/`](#design-platforms) - Hardware platform abstractions
|
||||
- [`fastvideo/v1/worker/`](#design-executor-and-worker-abstractions) - Multi-GPU process management
|
||||
- [`fastvideo/v1/fastvideo_args.py`](#design-fastvideo-args) - Argument handling
|
||||
- [`fastvideo/v1/forward_context.py`](#design-forwardcontext) - Forward pass context management
|
||||
- `fastvideo/v1/utils.py` - Utility functions
|
||||
- [`fastvideo/v1/logger.py`](#design-logger) - Logging infrastructure
|
||||
|
||||
## Core Architecture
|
||||
|
||||
FastVideo separates model components from execution logic with these principles:
|
||||
- **Component Isolation**: Models (encoders, VAEs, transformers) are isolated from execution (pipelines, stages, distributed processing)
|
||||
- **Modular Design**: Components can be independently replaced
|
||||
- **Distributed Execution**: Supports various parallelism strategies (Tensor, Sequence)
|
||||
- **Custom Attention Backends**: Components can support and use different Attention implementations
|
||||
- **Pipeline Abstraction**: Consistent interface across diffusion models
|
||||
|
||||
(design-fastvideo-args)=
|
||||
## FastVideoArgs
|
||||
|
||||
The `FastVideoArgs` class in `fastvideo/v1/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
|
||||
|
||||
Key features include:
|
||||
- **Command-line Interface**: Automatic conversion between CLI arguments and dataclass fields
|
||||
- **Configuration Groups**: Organized by functional areas (model loading, video params, optimization settings)
|
||||
- **Context Management**: Global access to current settings via `get_current_fastvideo_args()`
|
||||
- **Parameter Validation**: Ensures valid combinations of settings
|
||||
|
||||
Common configuration areas:
|
||||
- **Model paths and loading options**: `model_path`, `trust_remote_code`, `revision`
|
||||
- **Distributed execution settings**: `num_gpus`, `tp_size`, `sp_size`
|
||||
- **Video generation parameters**: `height`, `width`, `num_frames`, `num_inference_steps`
|
||||
- **Precision settings**: Control computation precision for different components
|
||||
|
||||
Example usage:
|
||||
|
||||
```python
|
||||
# Load arguments from command line
|
||||
fastvideo_args = prepare_fastvideo_args(sys.argv[1:])
|
||||
|
||||
# Access parameters
|
||||
model = load_model(fastvideo_args.model_path)
|
||||
|
||||
# Set as global context
|
||||
with set_current_fastvideo_args(fastvideo_args):
|
||||
# Code that requires access to these arguments
|
||||
result = generate_video()
|
||||
```
|
||||
|
||||
(design-pipeline-system)=
|
||||
## Pipeline System
|
||||
|
||||
### `ComposedPipelineBase`
|
||||
|
||||
This foundational class provides:
|
||||
|
||||
- **Model Loading**: Automatically loads components from HuggingFace-Diffusers-compatible model directories
|
||||
- **Stage Management**: Creates and orchestrates processing stages
|
||||
- **Data Flow Coordination**: Ensures proper state flow between stages
|
||||
|
||||
```python
|
||||
class MyCustomPipeline(ComposedPipelineBase):
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# Pipeline-specific initialization
|
||||
pass
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
self.add_stage("input_validation_stage", InputValidationStage())
|
||||
self.add_stage("text_encoding_stage", CLIPTextEncodingStage(
|
||||
text_encoder=self.get_module("text_encoder"),
|
||||
tokenizer=self.get_module("tokenizer")
|
||||
))
|
||||
# Additional stages...
|
||||
```
|
||||
|
||||
### Pipeline Stages
|
||||
Each stage handles a specific diffusion process component:
|
||||
- **Input Validation**: Parameter verification
|
||||
- **Text Encoding**: CLIP, LLaMA, or T5-based encoding
|
||||
- **Image Encoding**: Image input processing
|
||||
- **Timestep & Latent Preparation**: Setup for diffusion
|
||||
- **Denoising**: Core diffusion loop
|
||||
- **Decoding**: Latent-to-pixel conversion
|
||||
|
||||
Each stage implements a standard interface:
|
||||
|
||||
```python
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
# Process batch and update state
|
||||
return batch
|
||||
```
|
||||
|
||||
(design-forwardbatch)=
|
||||
### ForwardBatch
|
||||
|
||||
Defined in `fastvideo/v1/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
|
||||
|
||||
- **Input Data**: Prompts, images, generation parameters
|
||||
- **Intermediate State**: Embeddings, latents, timesteps, accumulated during stage execution
|
||||
- **Output Storage**: Generated results and metadata
|
||||
- **Configuration**: Sampling parameters, precision settings
|
||||
|
||||
This structure facilitates clear state transitions between stages.
|
||||
|
||||
(design-model-components)=
|
||||
## Model Components
|
||||
|
||||
The `fastvideo/v1/models/` directory contains implementations of the core neural network models used in video diffusion:
|
||||
|
||||
(design-transformer-models)=
|
||||
### Transformer Models
|
||||
|
||||
Transformer networks perform the actual denoising during diffusion:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/dits/`
|
||||
- **Examples**:
|
||||
- `WanTransformer3DModel`
|
||||
- `HunyuanVideoTransformer3DModel`
|
||||
|
||||
Features include:
|
||||
- Text/image conditioning
|
||||
- Standardized interface for model-specific optimizations
|
||||
|
||||
```python
|
||||
def forward(
|
||||
self,
|
||||
latents, # [B, T, C, H, W]
|
||||
encoder_hidden_states, # Text embeddings
|
||||
timestep, # Current diffusion timestep
|
||||
encoder_hidden_states_image=None, # Optional image embeddings
|
||||
**kwargs
|
||||
):
|
||||
# Perform denoising computation
|
||||
return noise_pred # Predicted noise residual
|
||||
```
|
||||
|
||||
(design-vae-variational-auto-encoder)=
|
||||
### VAE (Variational Auto-Encoder)
|
||||
|
||||
VAEs handle conversion between pixel space and latent space:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/vaes/`
|
||||
- **Examples**:
|
||||
- `AutoencoderKLWan`
|
||||
- `AutoencoderKLHunyuanVideo`
|
||||
|
||||
These models compress image/video data to a more efficient latent representation (typically 4x-8x smaller in each dimension).
|
||||
|
||||
FastVideo's VAE implementations include:
|
||||
- Efficient video batch processing
|
||||
- Memory optimization
|
||||
- Optional tiling for large frames
|
||||
- Distributed weight support
|
||||
|
||||
(design-text-and-image-encoders)=
|
||||
### Text and Image Encoders
|
||||
|
||||
Encoders process conditioning inputs into embeddings:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/encoders/`
|
||||
- **Text Encoders**:
|
||||
- `CLIPTextModel`
|
||||
- `LlamaModel`
|
||||
- `UMT5EncoderModel`
|
||||
- **Image Encoders**:
|
||||
- `CLIPVisionModel`
|
||||
|
||||
FastVideo implements optimizations such as:
|
||||
- Vocab parallelism for distributed processing
|
||||
- Caching for common prompts
|
||||
- Precision-tuned computation
|
||||
|
||||
(design-schedulers)=
|
||||
### Schedulers
|
||||
|
||||
Schedulers manage the diffusion sampling process:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/schedulers/`
|
||||
- **Examples**:
|
||||
- `UniPCMultistepScheduler`
|
||||
- `FlowMatchEulerDiscreteScheduler`
|
||||
|
||||
These components control:
|
||||
- Diffusion timestep sequences
|
||||
- Noise prediction to latent update conversions
|
||||
- Quality/speed trade-offs
|
||||
|
||||
```python
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
sample: torch.Tensor,
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
# Process model output and update latents
|
||||
# Return updated latents
|
||||
return prev_sample
|
||||
```
|
||||
|
||||
(design-optimized-attention)=
|
||||
## Optimized Attention
|
||||
|
||||
The `fastvideo/v1/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
|
||||
|
||||
### Attention Backends
|
||||
Multiple implementations with automatic selection:
|
||||
- **FLASH_ATTN**: Optimized for supporting hardware
|
||||
- **TORCH_SDPA**: Built-in PyTorch scaled dot-product attention
|
||||
- **SLIDING_TILE_ATTN**: For very long sequences
|
||||
|
||||
```python
|
||||
# Configure available attention backends for this layer
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
)
|
||||
|
||||
# Override via environment variable
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
```
|
||||
|
||||
### Attention Patterns
|
||||
Supports various patterns with memory optimization techniques:
|
||||
- **Cross/Self/Temporal/Global-Local Attention**
|
||||
- Chunking, progressive computation, optimized masking
|
||||
|
||||
(design-distributed-processing)=
|
||||
## Distributed Processing
|
||||
|
||||
The `fastvideo/v1/distributed/` directory contains implementations for distributed model execution:
|
||||
|
||||
(design-tensor-parallelism)=
|
||||
### Tensor Parallelism
|
||||
|
||||
Tensor parallelism splits model weights across devices:
|
||||
|
||||
- **Implementation**: Through `RowParallelLinear` and `ColumnParallelLinear` layers
|
||||
- **Use cases**: Will be used by encoder models as their sequence lengths are shorter and enables efficient sharding.
|
||||
|
||||
```python
|
||||
# Tensor-parallel layers in a transformer block
|
||||
from fastvideo.v1.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||
|
||||
# Split along output dimension
|
||||
self.qkv_proj = ColumnParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_size=3 * hidden_size,
|
||||
bias=True,
|
||||
gather_output=False
|
||||
)
|
||||
|
||||
# Split along input dimension
|
||||
self.out_proj = RowParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_size=hidden_size,
|
||||
bias=True,
|
||||
input_is_parallel=True
|
||||
)
|
||||
```
|
||||
|
||||
### Sequence Parallelism
|
||||
|
||||
Sequence parallelism splits sequences across devices:
|
||||
|
||||
- **Implementation**: Through `DistributedAttention` and sequence splitting
|
||||
- **Use cases**: Long video sequences or high-resolution processing. Used by DiT models.
|
||||
|
||||
```python
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN)
|
||||
)
|
||||
```
|
||||
|
||||
### Communication Primitives
|
||||
Efficient distributed operations via AllGather, AllReduce, and synchronization mechanisms.
|
||||
|
||||
Efficient communication primitives minimize distributed overhead:
|
||||
|
||||
- **Sequence-Parallel AllGather**: Collects sequence chunks
|
||||
- **Tensor-Parallel AllReduce**: Combines partial results
|
||||
- **Distributed Synchronization**: Coordinates execution
|
||||
|
||||
(design-forwardcontext)=
|
||||
## Forward Context Management
|
||||
|
||||
### ForwardContext
|
||||
|
||||
Defined in `fastvideo/v1/forward_context.py`, `ForwardContext` manages execution-specific state *within* a forward pass, particularly for low-level optimizations. It is accessed via `get_forward_context()`.
|
||||
|
||||
- **Attention Metadata**: Configuration for optimized attention kernels (`attn_metadata`)
|
||||
- **Profiling Data**: Potential hooks for performance metrics collection
|
||||
|
||||
This context-based approach enables:
|
||||
- Dynamic optimization based on execution state (e.g., attention backend selection)
|
||||
- Step-specific customizations within model components
|
||||
|
||||
Usage example:
|
||||
|
||||
```python
|
||||
with set_forward_context(current_timestep, attn_metadata, fastvideo_args):
|
||||
# During this forward pass, components can access context
|
||||
# through get_forward_context()
|
||||
output = model(inputs)
|
||||
```
|
||||
|
||||
(design-executor-and-worker-abstractions)=
|
||||
## Executor and Worker System
|
||||
|
||||
The `fastvideo/v1/worker/` directory contains the distributed execution framework:
|
||||
|
||||
### Executor Abstraction
|
||||
|
||||
FastVideo implements a flexible execution model for distributed processing:
|
||||
|
||||
- **Executor Base Class**: An abstract base class defining the interface for all executors
|
||||
- **MultiProcExecutor**: Primary implementation that spawns and manages worker processes
|
||||
- **GPU Workers**: Handle actual model execution on individual GPUs
|
||||
|
||||
The MultiProcExecutor implementation:
|
||||
1. Spawns worker processes for each GPU
|
||||
2. Establishes communication channels via pipes
|
||||
3. Coordinates distributed operations across workers
|
||||
4. Handles graceful startup and shutdown of the process group
|
||||
|
||||
Each GPU worker:
|
||||
1. Initializes the distributed environment
|
||||
2. Builds the pipeline for the specified model
|
||||
3. Executes requested operations on its assigned GPU
|
||||
4. Manages local resources and communicates results back to the executor
|
||||
|
||||
This design allows FastVideo to efficiently utilize multiple GPUs while providing a simple, unified interface for model execution.
|
||||
|
||||
(design-platforms)=
|
||||
## Platforms
|
||||
|
||||
The `fastvideo/v1/platforms/` directory provides hardware platform abstractions that enable FastVideo to run efficiently on different hardware configurations:
|
||||
|
||||
### Platform Abstraction
|
||||
|
||||
FastVideo's platform abstraction layer enables:
|
||||
- **Hardware Detection**: Automatic detection of available hardware
|
||||
- **Backend Selection**: Appropriate selection of compute kernels
|
||||
- **Memory Management**: Efficient utilization of hardware-specific memory features
|
||||
|
||||
The primary components include:
|
||||
- **Platform Interface**: Defines the common API for all platform implementations
|
||||
- **CUDA Platform**: Optimized implementation for NVIDIA GPUs
|
||||
- **Backend Enum**: Used throughout the codebase for feature selection
|
||||
|
||||
Usage example:
|
||||
|
||||
```python
|
||||
from fastvideo.v1.platforms import current_platform, _Backend
|
||||
|
||||
# Check hardware capabilities
|
||||
if current_platform.supports_backend(_Backend.FLASH_ATTN):
|
||||
# Use FlashAttention implementation
|
||||
else:
|
||||
# Fall back to standard implementation
|
||||
```
|
||||
|
||||
The platform system is designed to be extensible for future hardware targets.
|
||||
|
||||
(design-logger)=
|
||||
## Logger
|
||||
See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
|
||||
|
||||
*TODO*: (help wanted) Add an environment variable that disables process-aware logging.
|
||||
|
||||
## Contributing to FastVideo
|
||||
|
||||
If you're a new contributor, here are some common areas to explore:
|
||||
|
||||
1. **Adding a new model**: Implement new model types in the appropriate subdirectory of `fastvideo/v1/models/`
|
||||
2. **Optimizing performance**: Look at attention implementations or memory management
|
||||
3. **Adding a new pipeline**: Create a new pipeline subclass in `fastvideo/v1/pipelines/`
|
||||
4. **Hardware support**: Extend the `platforms` module for new hardware targets
|
||||
|
||||
When adding code, follow these practices:
|
||||
- Use type hints for better code readability
|
||||
- Add appropriate docstrings
|
||||
- Maintain the separation between model components and execution logic
|
||||
- Follow existing patterns for distributed processing
|
||||
@@ -162,52 +162,54 @@ class Example:
|
||||
return content
|
||||
|
||||
|
||||
def generate_examples():
|
||||
# Create the EXAMPLE_DOC_DIR if it doesn't exist
|
||||
if not EXAMPLE_DOC_DIR.exists():
|
||||
EXAMPLE_DOC_DIR.mkdir(parents=True)
|
||||
def generate_examples(generate_main_index=False):
|
||||
"""
|
||||
Generate example documentation.
|
||||
|
||||
Args:
|
||||
generate_main_index (bool): Whether to generate the main examples index.
|
||||
If False, only category-specific indices will be generated.
|
||||
"""
|
||||
# Create empty indices with dynamic paths
|
||||
main_index_dir = ROOT_DIR / "docs/source/examples"
|
||||
if not main_index_dir.exists():
|
||||
main_index_dir.mkdir(parents=True)
|
||||
|
||||
# Create empty indices
|
||||
examples_index = Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_index.md",
|
||||
title="Examples",
|
||||
description=
|
||||
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.", # noqa: E501
|
||||
caption="Examples",
|
||||
maxdepth=2)
|
||||
# Category indices stored in reverse order because they are inserted into
|
||||
# examples_index.documents at index 0 in order
|
||||
# Create the main examples index only if requested
|
||||
examples_index = None
|
||||
if generate_main_index:
|
||||
examples_index = Index(
|
||||
path=main_index_dir / "examples_index.md",
|
||||
title="💡 Examples",
|
||||
description=
|
||||
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.", # noqa: E501
|
||||
caption="Examples",
|
||||
maxdepth=2)
|
||||
|
||||
# Category indices with dynamic paths based on category names
|
||||
category_indices = {
|
||||
# "other":
|
||||
# Index(
|
||||
# path=EXAMPLE_DOC_DIR / "examples_other_index.md",
|
||||
# title="Other",
|
||||
# description=
|
||||
# "Other examples that don't strongly fit into the online or offline serving categories.", # noqa: E501
|
||||
# caption="Examples",
|
||||
# ),
|
||||
# "online_serving":
|
||||
# Index(
|
||||
# path=EXAMPLE_DOC_DIR / "examples_online_serving_index.md",
|
||||
# title="Online Serving",
|
||||
# description=
|
||||
# "Online serving examples demonstrate how to use FastVideo in an online setting, where the model is queried for predictions in real-time.", # noqa: E501
|
||||
# caption="Examples",
|
||||
# ),
|
||||
"inference":
|
||||
Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_inference_index.md",
|
||||
title="Inference",
|
||||
path=ROOT_DIR /
|
||||
"docs/source/inference/examples/examples_inference_index.md",
|
||||
title="🚀 Examples",
|
||||
description=
|
||||
"Inference examples demonstrate how to use FastVideo in an offline setting, where the model is queried for predictions in batches. We recommend starting with <project:basic.md>.", # noqa: E501
|
||||
caption="Examples",
|
||||
),
|
||||
}
|
||||
|
||||
# Ensure all category doc directories exist
|
||||
for category, index in category_indices.items():
|
||||
category_dir = index.path.parent
|
||||
if not category_dir.exists():
|
||||
category_dir.mkdir(parents=True)
|
||||
|
||||
examples = []
|
||||
glob_patterns = ["*.py", "*.md", "*.sh"]
|
||||
# Find categorised examples
|
||||
for category in category_indices:
|
||||
print(category)
|
||||
category_dir = EXAMPLE_DIR / category
|
||||
globs = [category_dir.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
@@ -215,33 +217,58 @@ def generate_examples():
|
||||
# Find examples in subdirectories
|
||||
for path in category_dir.glob("*/*.md"):
|
||||
examples.append(Example(path.parent, category))
|
||||
# Find uncategorised examples
|
||||
globs = [EXAMPLE_DIR.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path))
|
||||
# Find examples in subdirectories
|
||||
for path in EXAMPLE_DIR.glob("*/*.md"):
|
||||
# Skip categorised examples
|
||||
if path.parent.name in category_indices:
|
||||
continue
|
||||
examples.append(Example(path.parent))
|
||||
|
||||
# Generate the example documentation
|
||||
# Find uncategorised examples only if we're generating a main index
|
||||
if generate_main_index:
|
||||
globs = [EXAMPLE_DIR.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path))
|
||||
# Find examples in subdirectories
|
||||
for path in EXAMPLE_DIR.glob("*/*.md"):
|
||||
# Skip categorised examples
|
||||
if path.parent.name in category_indices:
|
||||
continue
|
||||
examples.append(Example(path.parent))
|
||||
|
||||
# Create document directories for each category based on category name and generate files
|
||||
for example in sorted(examples, key=lambda e: e.path.stem):
|
||||
doc_path = EXAMPLE_DOC_DIR / f"{example.path.stem}.md"
|
||||
print(example)
|
||||
|
||||
# Determine which index to use for this example
|
||||
if example.category is not None and example.category in category_indices:
|
||||
index = category_indices[example.category]
|
||||
elif generate_main_index:
|
||||
assert examples_index is not None
|
||||
index = examples_index # Default to main index if available
|
||||
else:
|
||||
# Skip examples without a category if no main index
|
||||
print(f"Skipping {example.path} (no category and no main index)")
|
||||
continue
|
||||
|
||||
# Place generated example markdown in the same directory as its index
|
||||
doc_path = index.path.parent / f"{example.path.stem}.md"
|
||||
with open(doc_path, "w+") as f:
|
||||
f.write(example.generate())
|
||||
# Add the example to the appropriate index
|
||||
assert example.category is not None
|
||||
index = category_indices.get(example.category, examples_index)
|
||||
# Add the example to the index
|
||||
index.documents.append(example.path.stem)
|
||||
|
||||
# Generate the index files
|
||||
# Generate the index files for categories
|
||||
for category_index in category_indices.values():
|
||||
if category_index.documents:
|
||||
examples_index.documents.insert(0, category_index.path.name)
|
||||
# Add to main index if it exists
|
||||
if generate_main_index:
|
||||
rel_path = category_index.path.relative_to(
|
||||
main_index_dir.parent)
|
||||
assert examples_index is not None
|
||||
examples_index.documents.insert(
|
||||
0,
|
||||
str(rel_path).replace(".md", ""))
|
||||
|
||||
# Write the category index file
|
||||
with open(category_index.path, "w+") as f:
|
||||
f.write(category_index.generate())
|
||||
|
||||
with open(examples_index.path, "w+") as f:
|
||||
f.write(examples_index.generate())
|
||||
# Write the main index file if requested
|
||||
if generate_main_index and examples_index:
|
||||
with open(examples_index.path, "w+") as f:
|
||||
f.write(examples_index.generate())
|
||||
|
||||
@@ -2,24 +2,21 @@
|
||||
|
||||
# 🔧 Installation
|
||||
|
||||
FastVideo currently only supports Linux and CUDA GPUs. The code is tested on Python 3.10.0 and CUDA 12.4, primarily with NVIDIA H100 GPUs.
|
||||
FastVideo currently only supports Linux and NVIDIA CUDA GPUs.
|
||||
|
||||
## Prerequisites
|
||||
## Requirements
|
||||
|
||||
- CUDA 12.4 installed and supported
|
||||
- Linux operating system
|
||||
- **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
|
||||
|
||||
#### 1. 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
|
||||
@@ -27,49 +24,87 @@ bash Miniconda3-latest-Linux-x86_64.sh
|
||||
source ~/.bashrc
|
||||
```
|
||||
|
||||
#### 2. 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)
|
||||
|
||||
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
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-guide)
|
||||
[Contributor Guide](#developer-overview)
|
||||
|
||||
## 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
|
||||
|
||||
@@ -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
|
||||
```
|
||||
+56
-25
@@ -9,7 +9,7 @@
|
||||
|
||||
:::{raw} html
|
||||
<p style="text-align:center">
|
||||
<strong>FastVideo is a lightweight framework for accelerating large video diffusion models.
|
||||
<strong>FastVideo is a unified framework for accelerated video generation.
|
||||
</strong>
|
||||
</p>
|
||||
|
||||
@@ -21,36 +21,64 @@
|
||||
</p>
|
||||
:::
|
||||
|
||||
FastVideo is a lightweight framework for accelerating large video diffusion models developed by the [Hao AI Lab](https://hao-ai-lab.github.io/).
|
||||
It features a clean, consistent API that works across popular video models, making it easier for developers to author new models and incorporate system- or kernel-level optimizations.
|
||||
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
|
||||
|
||||
<div style="text-align: center;">
|
||||
<video controls width="800">
|
||||
<source src="https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1" type="video/mp4">
|
||||
Your browser does not support the video tag.
|
||||
</video>
|
||||
<img src=_static/images/perf.png width="100%"/>
|
||||
</div>
|
||||
|
||||
FastVideo currently offers: (with more to come)
|
||||
## Key Features
|
||||
|
||||
- [NEW!] [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
FastVideo has the following features:
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- Cutting edge models
|
||||
- Wan2.1 T2V, I2V
|
||||
- HunyuanVideo
|
||||
- FastHunyuan: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- StepVideo T2V
|
||||
- Distillation support
|
||||
- Recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
|
||||
|
||||
Dev in progress and highly experimental.
|
||||
|
||||
## Documentation
|
||||
|
||||
% How to start using vLLM?
|
||||
% How to start using FastVideo?
|
||||
|
||||
:::{toctree}
|
||||
:caption: Getting Started
|
||||
: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/cli
|
||||
inference/add_pipeline
|
||||
inference/v0_inference
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Training
|
||||
:maxdepth: 1
|
||||
|
||||
training/data_preprocess
|
||||
training/distillation
|
||||
training/finetune
|
||||
:::
|
||||
|
||||
% What is STA Kernel?
|
||||
@@ -60,26 +88,29 @@ getting_started/examples/examples_index
|
||||
:maxdepth: 1
|
||||
|
||||
sliding_tile_attention/installation
|
||||
sliding_tile_attention/usage
|
||||
sliding_tile_attention/test
|
||||
sliding_tile_attention/demo
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Inference
|
||||
:caption: Design
|
||||
:maxdepth: 1
|
||||
|
||||
inference/stepvideo
|
||||
inference/hunyuanvideo
|
||||
inference/fasthunyuan
|
||||
inference/fastmochi
|
||||
design/overview
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Developer Guide
|
||||
:maxdepth: 1
|
||||
:maxdepth: 2
|
||||
|
||||
developer_guide/overview
|
||||
contributing/overview
|
||||
contributing/developer_env/index
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: API Reference
|
||||
:maxdepth: 2
|
||||
|
||||
<!-- api/summary -->
|
||||
api/fastvideo/fastvideo
|
||||
:::
|
||||
|
||||
## Indices and tables
|
||||
|
||||
@@ -0,0 +1,316 @@
|
||||
(add-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.
|
||||
|
||||
## Implementation Process Overview
|
||||
|
||||
1. **Port Required Modules** - Identify and implement necessary model components
|
||||
2. **Create Directory Structure** - Set up pipeline files and folders
|
||||
3. **Implement Pipeline Class** - Build the pipeline using existing or custom stages
|
||||
4. **Register Your Pipeline** - Make it discoverable by the framework
|
||||
5. **Configure Your Pipeline** - (Coming soon)
|
||||
|
||||
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
|
||||
|
||||
## Step 1: Pipeline Modules
|
||||
|
||||
### Identifying Required Modules
|
||||
|
||||
FastVideo uses the Hugging Face Diffusers format for model organization:
|
||||
|
||||
1. Examine the `model_index.json` in the HF model repository:
|
||||
|
||||
```json
|
||||
{
|
||||
"_class_name": "WanImageToVideoPipeline",
|
||||
"_diffusers_version": "0.33.0.dev0",
|
||||
"image_encoder": ["transformers", "CLIPVisionModelWithProjection"],
|
||||
"image_processor": ["transformers", "CLIPImageProcessor"],
|
||||
"scheduler": ["diffusers", "UniPCMultistepScheduler"],
|
||||
"text_encoder": ["transformers", "UMT5EncoderModel"],
|
||||
"tokenizer": ["transformers", "T5TokenizerFast"],
|
||||
"transformer": ["diffusers", "WanTransformer3DModel"],
|
||||
"vae": ["diffusers", "AutoencoderKLWan"]
|
||||
}
|
||||
```
|
||||
|
||||
1. For each component:
|
||||
- Note the originating library (`transformers` or `diffusers`)
|
||||
- Identify the class name
|
||||
- Check if it's already available in FastVideo
|
||||
|
||||
2. Review config files in each component's directory for architecture details
|
||||
|
||||
### Implementing Modules
|
||||
|
||||
Place new modules in the appropriate directories:
|
||||
- Encoders: `fastvideo/v1/models/encoders/`
|
||||
- VAEs: `fastvideo/v1/models/vaes/`
|
||||
- Transformer models: `fastvideo/v1/models/dits/`
|
||||
- Schedulers: `fastvideo/v1/models/schedulers/`
|
||||
|
||||
### Adapting Model Layers
|
||||
|
||||
#### Layer Replacements
|
||||
Replace standard PyTorch layers with FastVideo optimized versions:
|
||||
- nn.LayerNorm → fastvideo.v1.layers.layernorm.RMSNorm
|
||||
- Embedding layers → fastvideo.v1.layers.vocab_parallel_embedding modules
|
||||
- Activation functions → versions from fastvideo.v1.layers.activation
|
||||
|
||||
#### Distributed Linear Layers
|
||||
Use appropriate parallel layers for distribution:
|
||||
|
||||
```python
|
||||
# Output dimension parallelism
|
||||
from fastvideo.v1.layers.linear import ColumnParallelLinear
|
||||
self.q_proj = ColumnParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_size=head_size * num_heads,
|
||||
bias=bias,
|
||||
gather_output=False
|
||||
)
|
||||
|
||||
# Fused QKV projection
|
||||
from fastvideo.v1.layers.linear import QKVParallelLinear
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
hidden_size=hidden_size,
|
||||
head_size=attention_head_dim,
|
||||
total_num_heads=num_attention_heads,
|
||||
bias=True
|
||||
)
|
||||
|
||||
# Input dimension parallelism
|
||||
from fastvideo.v1.layers.linear import RowParallelLinear
|
||||
self.out_proj = RowParallelLinear(
|
||||
input_size=head_size * num_heads,
|
||||
output_size=hidden_size,
|
||||
bias=bias,
|
||||
input_is_parallel=True
|
||||
)
|
||||
```
|
||||
|
||||
### Attention Layers
|
||||
Replace standard attention with FastVideo's optimized attention:
|
||||
|
||||
```python
|
||||
# Local attention patterns
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
from fastvideo.v1.attention.backends.abstract import _Backend
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
dropout_rate=0.0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
)
|
||||
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
dropout_rate=0.0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
)
|
||||
```
|
||||
|
||||
#### Define supported backend selection
|
||||
|
||||
```python
|
||||
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
```
|
||||
|
||||
### Registering Models
|
||||
|
||||
Register implemented modules in the model registry:
|
||||
|
||||
```python
|
||||
# In fastvideo/v1/models/registry.py
|
||||
_TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"YourTransformerModel": ("dits", "yourmodule", "YourTransformerClass"),
|
||||
}
|
||||
|
||||
_VAE_MODELS = {
|
||||
"YourVAEModel": ("vaes", "yourvae", "YourVAEClass"),
|
||||
}
|
||||
```
|
||||
|
||||
## Step 2: Directory Structure
|
||||
|
||||
Create a new directory for your pipeline:
|
||||
|
||||
```
|
||||
fastvideo/v1/pipelines/
|
||||
├── your_pipeline/
|
||||
│ ├── __init__.py
|
||||
│ └── your_pipeline.py
|
||||
```
|
||||
|
||||
## Step 3: Implement Pipeline Class
|
||||
|
||||
Pipelines are composed of stages, each handling a specific part of the diffusion process:
|
||||
|
||||
- **InputValidationStage**: Validates input parameters
|
||||
- **Text Encoding Stages**: Handle text encoding (CLIP/Llama/T5)
|
||||
- **CLIPImageEncodingStage**: Processes image inputs
|
||||
- **TimestepPreparationStage**: Prepares diffusion timesteps
|
||||
- **LatentPreparationStage**: Manages latent representations
|
||||
- **ConditioningStage**: Processes conditioning inputs
|
||||
- **DenoisingStage**: Performs denoising diffusion
|
||||
- **DecodingStage**: Converts latents to pixels
|
||||
|
||||
### Creating Your Pipeline
|
||||
|
||||
```python
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (
|
||||
InputValidationStage, CLIPTextEncodingStage, TimestepPreparationStage,
|
||||
LatentPreparationStage, DenoisingStage, DecodingStage
|
||||
)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
import torch
|
||||
|
||||
class MyCustomPipeline(ComposedPipelineBase):
|
||||
"""Custom diffusion pipeline implementation."""
|
||||
|
||||
# Define required model components from model_index.json
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
@property
|
||||
def required_config_modules(self) -> List[str]:
|
||||
return self._required_config_modules
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""Initialize pipeline-specific components."""
|
||||
pass
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
self.add_stage(
|
||||
stage_name="input_validation_stage",
|
||||
stage=InputValidationStage()
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="prompt_encoding_stage",
|
||||
stage=CLIPTextEncodingStage(
|
||||
text_encoder=self.get_module("text_encoder"),
|
||||
tokenizer=self.get_module("tokenizer")
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae")
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")
|
||||
)
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="decoding_stage",
|
||||
stage=DecodingStage(
|
||||
vae=self.get_module("vae")
|
||||
)
|
||||
)
|
||||
|
||||
# Register the pipeline class
|
||||
EntryClass = MyCustomPipeline
|
||||
```
|
||||
|
||||
### Creating Custom Stages (Optional)
|
||||
|
||||
If existing stages don't meet your needs, create custom ones:
|
||||
|
||||
```python
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
|
||||
class MyCustomStage(PipelineStage):
|
||||
"""Custom processing stage for the pipeline."""
|
||||
|
||||
def __init__(self, custom_module, other_param=None):
|
||||
super().__init__()
|
||||
self.custom_module = custom_module
|
||||
self.other_param = other_param
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
# Access input data
|
||||
input_data = batch.some_attribute
|
||||
|
||||
# Validate inputs
|
||||
if input_data is None:
|
||||
raise ValueError("Required input is missing")
|
||||
|
||||
# Process with your module
|
||||
result = self.custom_module(input_data)
|
||||
|
||||
# Update batch with results
|
||||
batch.some_output = result
|
||||
|
||||
return batch
|
||||
```
|
||||
|
||||
Add your custom stage to the pipeline:
|
||||
|
||||
```python
|
||||
self.add_stage(
|
||||
stage_name="my_custom_stage",
|
||||
stage=MyCustomStage(
|
||||
custom_module=self.get_module("custom_module"),
|
||||
other_param="some_value"
|
||||
)
|
||||
)
|
||||
```
|
||||
|
||||
#### Stage Design Principles
|
||||
|
||||
1. **Single Responsibility**: Focus on one specific task
|
||||
2. **Functional Pattern**: Receive and return a `ForwardBatch` object
|
||||
3. **Dependency Injection**: Pass dependencies through constructor
|
||||
4. **Input Validation**: Validate inputs for clear error messages
|
||||
|
||||
## Step 4: Register Your Pipeline
|
||||
|
||||
Define `EntryClass` at the end of your pipeline file:
|
||||
|
||||
```python
|
||||
# Single pipeline class
|
||||
EntryClass = MyCustomPipeline
|
||||
|
||||
# Or multiple pipeline classes
|
||||
EntryClass = [MyCustomPipeline, MyOtherPipeline]
|
||||
```
|
||||
|
||||
The registry will automatically:
|
||||
1. Scan all packages under `fastvideo/v1/pipelines/`
|
||||
2. Look for `EntryClass` variables
|
||||
3. Register pipelines using their class names as identifiers
|
||||
|
||||
## Best Practices
|
||||
|
||||
- **Reuse Existing Components**: Leverage built-in stages and modules
|
||||
- **Follow Module Organization**: Place new modules in appropriate directories
|
||||
- **Match Model Patterns**: Follow existing code patterns and conventions
|
||||
@@ -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.
|
||||
@@ -0,0 +1,77 @@
|
||||
(inference-configuration)=
|
||||
# Configuration
|
||||
|
||||
## Multi-GPU Setup
|
||||
|
||||
FastVideo automatically distributes the generation process when multiple GPUs are specified:
|
||||
|
||||
```python
|
||||
# Will use 4 GPUs in parallel for faster generation
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=4,
|
||||
)
|
||||
```
|
||||
|
||||
## Customizing Generation
|
||||
|
||||
- `PipelineConfig`: Initialization time parameters
|
||||
- `SamplingParam`: Generation time parameters
|
||||
|
||||
You can customize various parameters when generating videos using the `PipelineConfig` and `SamplingParam` class:
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator, SamplingParam, PipelineConfig
|
||||
|
||||
def main():
|
||||
model_name = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
config = PipelineConfig.from_pretrained(model_name)
|
||||
config.vae_precision = "fp16"
|
||||
config.use_cpu_offload = True
|
||||
|
||||
# Create the generator
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
num_gpus=1,
|
||||
pipeline_config=config
|
||||
)
|
||||
|
||||
# Create and customize sampling parameters
|
||||
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
|
||||
# How many frames to generate
|
||||
sampling_param.num_frames = 45
|
||||
|
||||
# Video resolution (width, height)
|
||||
sampling_param.width = 1024
|
||||
sampling_param.height = 576
|
||||
|
||||
# How many steps we denoise the video (higher = better quality, slower generation)
|
||||
sampling_param.num_inference_steps = 30
|
||||
|
||||
# How strongly the video conforms to the prompt (higher = more faithful to prompt)
|
||||
sampling_param.guidance_scale = 7.5
|
||||
|
||||
# Random seed for reproducibility
|
||||
sampling_param.seed = 42 # Optional, leave unset for random results
|
||||
|
||||
# Generate video with custom parameters
|
||||
prompt = "A beautiful sunset over a calm ocean, with gentle waves."
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
sampling_param=sampling_param,
|
||||
output_path="my_videos/", # Controls where videos are saved
|
||||
return_frames=True, # Also return frames from this call (defaults to False)
|
||||
save_video=True
|
||||
)
|
||||
|
||||
# If return_frames=True, video contains the generated frames as a NumPy array
|
||||
print(f"Generated {len(video)} frames")
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
## Performance Optimization
|
||||
|
||||
For configuring optimizations, please see our [optimizations guide](#inference-optimizations)
|
||||
@@ -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
|
||||
@@ -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,124 @@
|
||||
# Inference Quick Start
|
||||
|
||||
This page contains step-by-step instructions to get you quickly started with video generation using FastVideo.
|
||||
|
||||
## Requirements
|
||||
- **OS**: Linux (Tested on Ubuntu 22.04+)
|
||||
- **Python**: 3.10-3.12
|
||||
- **CUDA**: 12.4
|
||||
- **GPU**: At least one NVIDIA GPU
|
||||
|
||||
## 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/`
|
||||
|
||||
More inference example scripts can be found in `scripts/inference/`
|
||||
## 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
|
||||
|
||||
def main():
|
||||
# Create the generator
|
||||
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
|
||||
|
||||
# Set up parameters with an initial image
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param.num_frames = 107
|
||||
|
||||
# Generate video based on the image
|
||||
prompt = "A photograph coming to life with gentle movement"
|
||||
generator.generate_video(prompt, sampling_param=sampling_param,
|
||||
output_path="my_videos/",
|
||||
save_video=True)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
## 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 distilled 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
|
||||
|
||||
## Next Steps
|
||||
|
||||
- Learn about [Advanced Inference Configurations](#inference-configuration)
|
||||
- Learn about using [Optimizations](#inference-optimizations)
|
||||
- See [Examples](../examples/examples_inference_index.md) for more usage scenarios
|
||||
- Join our [Community Discord](https://discord.gg/JA7cksDz86).
|
||||
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
|
||||
@@ -0,0 +1,148 @@
|
||||
(inference-optimizations)=
|
||||
# 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`.
|
||||
@@ -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
|
||||
```
|
||||
@@ -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.3B
|
||||
* `Wan-AI/Wan2.1-T2V-1.3B-Diffusers`
|
||||
* 480P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
- * Wan T2V 14B
|
||||
* `Wan-AI/Wan2.1-T2V-14B-Diffusers`
|
||||
* 480P, 720P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
- * Wan I2V 480P
|
||||
* `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers`
|
||||
* 480P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
- * Wan I2V 720P
|
||||
* `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers`
|
||||
* 720P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
- * StepVideo T2V
|
||||
* `FastVideo/stepvideo-t2v-diffusers`
|
||||
* 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.
|
||||
@@ -1,6 +1,38 @@
|
||||
(fasthunyuan)=
|
||||
(v0-inference)=
|
||||
|
||||
# [Deprecated] V0 Inference
|
||||
The following commands and APIs are deprecated but still supported until V1's API can completely replace all the features in this page.
|
||||
|
||||
## 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
|
||||
```
|
||||
|
||||
## 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.
|
||||
|
||||
# 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.
|
||||
|
||||
@@ -18,7 +50,7 @@ For more information about the VRAM requirements for BitsAndBytes quantization,
|
||||
| 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
|
||||
@@ -31,3 +63,12 @@ bash scripts/inference/inference_hunyuan.sh
|
||||
```
|
||||
|
||||
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
|
||||
|
||||
## 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
|
||||
```
|
||||
@@ -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.
|
||||
@@ -1,6 +1,6 @@
|
||||
(sta-demo)=
|
||||
|
||||
# Demo
|
||||
# 🔍 Demo
|
||||
There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
<div style="text-align: center;">
|
||||
|
||||
@@ -1,10 +1,18 @@
|
||||
(sta-installation)=
|
||||
|
||||
# Installation
|
||||
# 🔧 Installation
|
||||
You can install the Sliding Tile Attention package using
|
||||
|
||||
```
|
||||
pip install st_attn==0.0.4
|
||||
```
|
||||
|
||||
# Building from Source
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
|
||||
First, install C++20 for ThunderKittens:
|
||||
|
||||
```bash
|
||||
cd csrc/sliding_tile_attention/
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
@@ -23,3 +31,25 @@ export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
# 🧪 Test
|
||||
|
||||
```bash
|
||||
python test/test_sta.py
|
||||
```
|
||||
|
||||
# 📋 Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
|
||||
```
|
||||
|
||||
@@ -1,7 +0,0 @@
|
||||
(sta-test)=
|
||||
|
||||
# Test
|
||||
|
||||
```bash
|
||||
python test/test_sta.py
|
||||
```
|
||||
@@ -1,17 +0,0 @@
|
||||
(sta-usage)=
|
||||
|
||||
# Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
|
||||
```
|
||||
@@ -1,5 +1,6 @@
|
||||
(v0-data-preprocess)=
|
||||
|
||||
## 🧱 Data Preprocess
|
||||
# 🧱 Data Preprocess
|
||||
|
||||
To save GPU memory, we precompute text embeddings and VAE latents to eliminate the need to load the text encoder and VAE during training.
|
||||
|
||||
@@ -18,10 +19,11 @@ bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
|
||||
|
||||
The preprocessed dataset will be stored in `Image-Vid-Finetune-Mochi` or `Image-Vid-Finetune-HunYuan` correspondingly.
|
||||
|
||||
### Process your own dataset
|
||||
## Process your own dataset
|
||||
|
||||
If you wish to create your own dataset for finetuning or distillation, please structure you video dataset in the following format:
|
||||
|
||||
```
|
||||
path_to_dataset_folder/
|
||||
├── media/
|
||||
│ ├── 0.jpg
|
||||
@@ -29,6 +31,7 @@ path_to_dataset_folder/
|
||||
│ ├── 2.jpg
|
||||
├── video2caption.json
|
||||
└── merge.txt
|
||||
```
|
||||
|
||||
Format the JSON file as a list, where each item represents a media source:
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
(v0-distill)=
|
||||
# 🎯 Distill
|
||||
Our distillation recipe is based on [Phased Consistency Model](https://github.com/G-U-N/Phased-Consistency-Model). We did not find significant improvement using multi-phase distillation, so we keep the one phase setup similar to the original latent consistency model's recipe.
|
||||
We use the [MixKit](https://huggingface.co/datasets/LanguageBind/Open-Sora-Plan-v1.1.0/tree/main/all_mixkit) dataset for distillation. To avoid running the text encoder and VAE during training, we prprocess all data to generate text embeddings and VAE latents.
|
||||
Preprocessing instructions can be found [data_preprocess.md](#v0-data-preprocess). For convenience, we also provide preprocessed data that can be downloaded directly using the following command:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/HD-Mixkit-Finetune-Hunyuan --local_dir=data/HD-Mixkit-Finetune-Hunyuan --repo_type=dataset
|
||||
```
|
||||
|
||||
Next, download the original model weights with:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model # original hunyuan
|
||||
python scripts/huggingface/download_hf.py --repo_id=genmo/mochi-1-preview --local_dir=data/mochi --repo_type=model # original mochi
|
||||
```
|
||||
|
||||
To launch the distillation process, use the following commands:
|
||||
|
||||
```
|
||||
bash scripts/distill/distill_hunyuan.sh # for hunyuan
|
||||
bash scripts/distill/distill_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
We also provide an optional script for distillation with adversarial loss, located at `fastvideo/distill_adv.py`. Although we tried adversarial loss, we did not observe significant improvements.
|
||||
@@ -0,0 +1,71 @@
|
||||
(v0-finetune)=
|
||||
# 🧠 Finetune
|
||||
## ⚡ Full Finetune
|
||||
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](#v0-data-preprocess). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
|
||||
```
|
||||
|
||||
Download the original model weights as specified in [Distill Section](#v0-distill):
|
||||
|
||||
Then you can run the finetune with:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
## ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight --local_dir=data/Hunyuan-Black-Myth-Wukong-lora-weight --repo_type=model
|
||||
```
|
||||
|
||||
### Minimum Hardware Requirement
|
||||
- 40 GB GPU memory each for 2 GPUs with lora.
|
||||
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
|
||||
|
||||
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
|
||||
|
||||
### Dataset Preparation
|
||||
We provide scripts to better help you get started to train on your own characters!
|
||||
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
|
||||
|
||||
```
|
||||
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
|
||||
```
|
||||
|
||||
Also, we provide script to resize your videos:
|
||||
|
||||
```
|
||||
python scripts/data_preprocess/resize_videos.py
|
||||
```
|
||||
|
||||
### Finetuning
|
||||
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
|
||||
```
|
||||
|
||||
### Inference
|
||||
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
|
||||
|
||||
```
|
||||
bash scripts/inference/inference_hunyuan_hf.sh
|
||||
```
|
||||
|
||||
**We also provide scripts for Mochi in the same directory.**
|
||||
|
||||
### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_hunyuan.sh
|
||||
bash scripts/finetune/finetune_mochi_lora_mix.sh
|
||||
```
|
||||
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
|
||||
@@ -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()
|
||||
```
|
||||
|
||||
@@ -1 +1,43 @@
|
||||
print('Hello, world!')
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# if num_gpus > 1, FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
use_cpu_offload=False
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
# 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, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
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()
|
||||
@@ -0,0 +1,59 @@
|
||||
# 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.
|
||||
|
||||
## Overview
|
||||
|
||||
The demo uses the FastVideo framework to generate videos based on text prompts. It provides a simple web interface built with Gradio that allows users to:
|
||||
|
||||
- Enter text prompts to generate videos
|
||||
- Customize video parameters (dimensions, number of frames, etc.)
|
||||
- Use negative prompts to guide the generation process
|
||||
- Set or randomize seeds for reproducibility
|
||||
|
||||
---
|
||||
|
||||
## Usage
|
||||
|
||||
Run the demo with:
|
||||
|
||||
```bash
|
||||
python examples/inference/gradio/gradio_demo.py
|
||||
```
|
||||
|
||||
This will start a web server at `http://0.0.0.0:7860` where you can access the interface.
|
||||
|
||||
---
|
||||
|
||||
## Model Initialization
|
||||
|
||||
This demo initializes a `VideoGenerator` with the minimum required arguments for inference. Users can seamlessly adjust inference options between generations, including prompts, resolution, video length, or even the number of inference steps, *without ever needing to reload the model*.
|
||||
|
||||
## Video Generation
|
||||
|
||||
The core functionality is in the `generate_video` function, which:
|
||||
1. Processes user inputs
|
||||
2. Uses the FastVideo VideoGenerator from earlier to run inference (`generator.generate_video()`)
|
||||
3. Returns an output path that Gradio uses to display the generated video
|
||||
|
||||
## Gradio Interface
|
||||
|
||||
The interface is built with several components:
|
||||
- A text input for the prompt
|
||||
- A video display for the result
|
||||
- Inference options in a collapsible accordion:
|
||||
- Height and width sliders
|
||||
- Number of frames slider
|
||||
- Guidance scale slider
|
||||
- Inference steps slider
|
||||
- Negative prompt options
|
||||
- Seed controls
|
||||
|
||||
### Inference Options
|
||||
|
||||
- **Height/Width**: Control the resolution of the generated video
|
||||
- **Number of Frames**: Set how many frames to generate
|
||||
- **Guidance Scale**: Control how closely the generation follows the prompt
|
||||
- **Inference Steps**: More steps can improve quality but take longer
|
||||
- **Negative Prompt**: Specify what you don't want to see in the video
|
||||
- **Seed**: Control randomness for reproducible results
|
||||
@@ -0,0 +1,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,45 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora"
|
||||
def main():
|
||||
# Initialize VideoGenerator with the Wan model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=2,
|
||||
lora_path="benjamin-paine/steamboat-willie-1.3b",
|
||||
lora_nickname="steamboat"
|
||||
)
|
||||
kwargs = {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 32,
|
||||
}
|
||||
# Generate video with LoRA style
|
||||
prompt = "steamboat willie style, golden era animation, close-up of a short fluffy monster kneeling beside a melting red candle. the mood is one of wonder and curiosity, as the monster gazes at the flame with wide eyes and open mouth. Its pose and expression convey a sense of innocence and playfulness, as if it is exploring the world around it for the first time. The use of warm colors and dramatic lighting further enhances the cozy atmosphere of the image."
|
||||
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
# sampling_param=sampling_param,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
generator.set_lora_adapter(lora_nickname="flat_color", lora_path="motimalu/wan-flat-color-1.3b-v2")
|
||||
prompt = "flat color, no lineart, blending, negative space, artist:[john kafka|ponsuke kaikai|hara id 21|yoneyama mai|fuzichoco], 1girl, sakura miko, pink hair, cowboy shot, white shirt, floral print, off shoulder, outdoors, cherry blossom, tree shade, wariza, looking up, falling petals, half-closed eyes, white sky, clouds, live2d animation, upper body, high quality cinematic video of a woman sitting under a sakura tree. Dreamy and lonely, the camera close-ups on the face of the woman as she turns towards the viewer. The Camera is steady, This is a cowboy shot. The animation is smooth and fluid."
|
||||
negative_prompt = "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -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()
|
||||
@@ -0,0 +1,5 @@
|
||||
# STA Mask Search Examples
|
||||
|
||||
```bash
|
||||
bash examples/inference/sta_mask_search/inference_wan_sta.sh
|
||||
```
|
||||
@@ -0,0 +1,39 @@
|
||||
#!/bin/bash
|
||||
|
||||
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
|
||||
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
|
||||
base_port=29503
|
||||
num_gpu=$(nvidia-smi --query-gpu=gpu_name --format=csv,noheader | wc -l)
|
||||
gpu_ids=$(seq 0 $((num_gpu-1)))
|
||||
skip_time_steps=12
|
||||
|
||||
output_path="inference_results/sta/mask_search_full"
|
||||
STA_mode="STA_searching"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_extend_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode &
|
||||
sleep 1
|
||||
done
|
||||
wait
|
||||
echo "STA searching completed"
|
||||
|
||||
output_path="inference_results/sta/mask_search_sparse"
|
||||
STA_mode="STA_tuning"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_extend_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode \
|
||||
--skip_time_steps $skip_time_steps &
|
||||
sleep 1
|
||||
done
|
||||
wait
|
||||
echo "STA tuning completed"
|
||||
|
||||
echo "All jobs completed"
|
||||
@@ -0,0 +1,63 @@
|
||||
import os
|
||||
import argparse
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
def main(args):
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
num_gpus=args.num_gpus, # Adjust based on your hardware
|
||||
STA_mode=args.STA_mode,
|
||||
skip_time_steps=args.skip_time_steps
|
||||
)
|
||||
|
||||
# Prompts for your video
|
||||
prompt = args.prompt
|
||||
prompt_path = args.prompt_path
|
||||
negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
|
||||
if prompt_path is not None:
|
||||
with open(prompt_path, "r") as f:
|
||||
prompts = f.readlines()
|
||||
else:
|
||||
prompts = [prompt]
|
||||
|
||||
params = SamplingParam(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
fps=args.fps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
return_frames=True, # Also return frames from this call (defaults to False)
|
||||
output_path=args.output_path, # Controls where videos are saved
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt
|
||||
)
|
||||
|
||||
# Generate the video
|
||||
for prompt in prompts:
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
sampling_param=params,
|
||||
)
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompt", type=str, default="A man is dancing.")
|
||||
parser.add_argument("--prompt_path", type=str, default=None)
|
||||
parser.add_argument("--height", type=int, default=768)
|
||||
parser.add_argument("--width", type=int, default=1280)
|
||||
parser.add_argument("--num_frames", type=int, default=69)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=50)
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--guidance_scale", type=float, default=5.0)
|
||||
parser.add_argument("--seed", type=int, default=12345)
|
||||
parser.add_argument("--output_path", type=str, default="my_videos/")
|
||||
parser.add_argument("--num_gpus", type=int, default=1)
|
||||
parser.add_argument("--STA_mode", type=str, default="STA_searching")
|
||||
parser.add_argument("--skip_time_steps", type=int, default=12)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -1 +1,6 @@
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.version import __version__
|
||||
|
||||
__all__ = ["VideoGenerator", "PipelineConfig", "SamplingParam", "__version__"]
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import maybe_download_model, shallow_asdict
|
||||
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_i2v import PreprocessPipeline_I2V
|
||||
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_t2v import PreprocessPipeline_T2V
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
def main(args):
|
||||
args.model_path = maybe_download_model(args.model_path)
|
||||
# Assume using torchrun
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
init_distributed_environment(world_size=world_size, rank=rank, local_rank=local_rank)
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
|
||||
torch.cuda.set_device(local_rank)
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
|
||||
kwargs = {
|
||||
"use_cpu_offload": False,
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
|
||||
}
|
||||
pipeline_config_args = shallow_asdict(pipeline_config)
|
||||
pipeline_config_args.update(kwargs)
|
||||
fastvideo_args = FastVideoArgs(model_path=args.model_path,
|
||||
num_gpus=world_size,
|
||||
device_str="cuda",
|
||||
**pipeline_config_args,
|
||||
)
|
||||
fastvideo_args.check_fastvideo_args()
|
||||
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
|
||||
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
|
||||
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
|
||||
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--model_type", type=str, default="mochi")
|
||||
parser.add_argument("--data_merge_path", type=str, required=True)
|
||||
parser.add_argument("--validation_prompt_txt", type=str)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument(
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--preprocess_video_batch_size",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--preprocess_text_batch_size",
|
||||
type=int,
|
||||
default=8,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--samples_per_file",
|
||||
type=int,
|
||||
default=64
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flush_frequency",
|
||||
type=int,
|
||||
default=256,
|
||||
help="how often to save to parquet files"
|
||||
)
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument("--max_height", type=int, default=480)
|
||||
parser.add_argument("--max_width", type=int, default=848)
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--dataset", default="t2v")
|
||||
parser.add_argument("--preprocess_task", type=str, default="t2v")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
parser.add_argument("--speed_factor", type=float, default=1.0)
|
||||
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
parser.add_argument("--cfg", type=float, default=0.0)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -68,7 +68,8 @@ def main(args):
|
||||
train_dataset = T5dataset(latents_json_path, args.vae_debug)
|
||||
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
vae.enable_tiling()
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
|
||||
@@ -33,7 +33,8 @@ def main(args):
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
vae.enable_tiling()
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
|
||||
|
||||
+77
-43
@@ -12,6 +12,7 @@ import torch.distributed as dist
|
||||
import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
from diffusers.optimization import get_scheduler
|
||||
from diffusers.utils import check_min_version
|
||||
from peft import LoraConfig
|
||||
@@ -23,7 +24,7 @@ from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.utils.latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint, save_lora_checkpoint)
|
||||
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
|
||||
@@ -123,13 +124,21 @@ def distill_one_step(
|
||||
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
|
||||
# Predict the noise residual
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
if args.model_type == "wan":
|
||||
teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"return_dict": True,
|
||||
}
|
||||
else:
|
||||
teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
if hunyuan_teacher_disable_cfg:
|
||||
teacher_kwargs["guidance"] = torch.tensor([1000.0],
|
||||
device=noisy_model_input.device,
|
||||
@@ -141,47 +150,70 @@ def distill_one_step(
|
||||
with torch.no_grad():
|
||||
w = distill_cfg
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
cond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
if args.model_type == "wan":
|
||||
cond_teacher_kwargs ={
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"return_dict": True,
|
||||
}
|
||||
else:
|
||||
cond_teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
cond_teacher_output = teacher_transformer(**cond_teacher_kwargs)[0].float()
|
||||
if not_apply_cfg_solver:
|
||||
uncond_teacher_output = cond_teacher_output
|
||||
else:
|
||||
# Get teacher model prediction on noisy_latents and unconditional embedding
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
uncond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
timesteps,
|
||||
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
if args.model_type == "wan":
|
||||
uncond_teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states":uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
"timestep": timesteps,
|
||||
"return_dict": True,
|
||||
}
|
||||
else:
|
||||
uncond_teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states":uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
"return_dict": False,
|
||||
}
|
||||
|
||||
uncond_teacher_output = teacher_transformer(**uncond_teacher_kwargs)[0].float()
|
||||
|
||||
teacher_output = uncond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
|
||||
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
|
||||
|
||||
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
|
||||
with torch.no_grad():
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
if ema_transformer is not None:
|
||||
target_pred = ema_transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0]
|
||||
if args.model_type == "wan":
|
||||
target_pred_kwargs = {
|
||||
"hidden_states": x_prev.float(),
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep":timesteps_prev,
|
||||
"return_dict":True,
|
||||
}
|
||||
else:
|
||||
target_pred = transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0]
|
||||
target_pred_kwargs = {
|
||||
"hidden_states": x_prev.float(),
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep":timesteps_prev,
|
||||
"encoder_attention_mask":encoder_attention_mask,
|
||||
"return_dict":False,
|
||||
}
|
||||
if ema_transformer is not None:
|
||||
target_pred = ema_transformer(**target_pred_kwargs)[0]
|
||||
else:
|
||||
target_pred = transformer(**target_pred_kwargs)[0]
|
||||
|
||||
target, end_index = solver.euler_style_multiphase_pred(x_prev, target_pred, index, multiphase, True)
|
||||
|
||||
@@ -242,7 +274,7 @@ def main(args):
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
if rank == 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
@@ -319,7 +351,9 @@ def main(args):
|
||||
teacher_transformer.requires_grad_(False)
|
||||
if args.use_ema:
|
||||
ema_transformer.requires_grad_(False)
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
|
||||
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
|
||||
if args.scheduler_type == "pcm_linear_quadratic":
|
||||
linear_steps = int(noise_scheduler.config.num_train_timesteps * args.linear_range)
|
||||
sigmas = linear_quadratic_schedule(
|
||||
@@ -391,7 +425,7 @@ def main(args):
|
||||
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
@@ -452,7 +486,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 +516,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)
|
||||
|
||||
@@ -493,7 +527,7 @@ def main(args):
|
||||
"phases": num_phases,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss": loss,
|
||||
|
||||
@@ -23,7 +23,7 @@ from tqdm.auto import tqdm
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.distill.discriminator import Discriminator
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.utils.latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
from fastvideo.utils.checkpoint import (resume_lora_optimizer, resume_training_generator_discriminator, save_checkpoint,
|
||||
save_lora_checkpoint)
|
||||
@@ -296,7 +296,7 @@ def main(args):
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
if rank == 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
@@ -462,7 +462,7 @@ def main(args):
|
||||
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -559,7 +559,7 @@ def main(args):
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
wandb.log(
|
||||
{
|
||||
"generator_loss": generator_loss,
|
||||
|
||||
@@ -237,7 +237,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
|
||||
type=str,
|
||||
default="540p",
|
||||
choices=["540p", "720p"],
|
||||
help="Root path of all the models, including t2v models and extra models.",
|
||||
help="The resolution of the model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--load-key",
|
||||
@@ -361,7 +361,7 @@ def add_parallel_args(parser: argparse.ArgumentParser):
|
||||
"--ring-degree",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Ulysses degree.",
|
||||
help="Ring degree.",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
@@ -17,7 +17,7 @@ from fastvideo.models.hunyuan.vae import load_vae
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
|
||||
|
||||
class Inference(object):
|
||||
class Inference:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -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
|
||||
|
||||
@@ -41,7 +41,7 @@ def get_rewrite_prompt(ori_prompt, mode="Normal"):
|
||||
elif mode == "Master":
|
||||
prompt = master_mode_prompt.format(input=ori_prompt)
|
||||
else:
|
||||
raise Exception("Only supports Normal and Normal", mode)
|
||||
raise Exception("Only supports Normal and Master mode, but got {}".format(mode))
|
||||
return prompt
|
||||
|
||||
|
||||
|
||||
@@ -267,25 +267,25 @@ class Step1Model(PreTrainedModel):
|
||||
class STEP1TextEncoder(torch.nn.Module):
|
||||
|
||||
def __init__(self, model_dir, max_length=320):
|
||||
super(STEP1TextEncoder, self).__init__()
|
||||
super()
|
||||
self.max_length = max_length
|
||||
self.text_tokenizer = Wrapped_StepChatTokenizer(os.path.join(model_dir, 'step1_chat_tokenizer.model'))
|
||||
text_encoder = Step1Model.from_pretrained(model_dir)
|
||||
self.text_encoder = text_encoder.eval().to(torch.bfloat16)
|
||||
|
||||
@torch.no_grad
|
||||
@torch.autocast(device_type='cuda', dtype=torch.bfloat16)
|
||||
def forward(self, prompts, with_mask=True, max_length=None):
|
||||
self.device = next(self.text_encoder.parameters()).device
|
||||
with torch.no_grad(), torch.cuda.amp.autocast(dtype=torch.bfloat16):
|
||||
if type(prompts) is str:
|
||||
prompts = [prompts]
|
||||
if type(prompts) is str:
|
||||
prompts = [prompts]
|
||||
|
||||
txt_tokens = self.text_tokenizer(prompts,
|
||||
max_length=max_length or self.max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_tensors="pt")
|
||||
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
|
||||
txt_tokens = self.text_tokenizer(prompts,
|
||||
max_length=max_length or self.max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_tensors="pt")
|
||||
y = self.text_encoder(txt_tokens.input_ids.to(self.device),
|
||||
attention_mask=txt_tokens.attention_mask.to(self.device) if with_mask else None)
|
||||
y_mask = txt_tokens.attention_mask
|
||||
y_mask = txt_tokens.attention_mask
|
||||
return y.transpose(0, 1), y_mask
|
||||
|
||||
@@ -0,0 +1,486 @@
|
||||
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import math
|
||||
from typing import Any, Dict, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
|
||||
from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers
|
||||
from diffusers.models.attention import FeedForward
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.models.cache_utils import CacheMixin
|
||||
from diffusers.models.embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps, get_1d_rotary_pos_embed
|
||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.normalization import FP32LayerNorm
|
||||
|
||||
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
class WanAttnProcessor2_0:
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("WanAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
encoder_hidden_states_img = None
|
||||
if attn.add_k_proj is not None:
|
||||
# 512 is the context length of the text encoder, hardcoded for now
|
||||
image_context_length = encoder_hidden_states.shape[1] - 512
|
||||
encoder_hidden_states_img = encoder_hidden_states[:, :image_context_length]
|
||||
encoder_hidden_states = encoder_hidden_states[:, image_context_length:]
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
|
||||
if rotary_emb is not None:
|
||||
|
||||
def apply_rotary_emb(hidden_states: torch.Tensor, freqs: torch.Tensor):
|
||||
x_rotated = torch.view_as_complex(hidden_states.to(torch.float64).unflatten(3, (-1, 2)))
|
||||
x_out = torch.view_as_real(x_rotated * freqs).flatten(3, 4)
|
||||
return x_out.type_as(hidden_states)
|
||||
|
||||
query = apply_rotary_emb(query, rotary_emb)
|
||||
key = apply_rotary_emb(key, rotary_emb)
|
||||
|
||||
# I2V task
|
||||
hidden_states_img = None
|
||||
if encoder_hidden_states_img is not None:
|
||||
key_img = attn.add_k_proj(encoder_hidden_states_img)
|
||||
key_img = attn.norm_added_k(key_img)
|
||||
value_img = attn.add_v_proj(encoder_hidden_states_img)
|
||||
|
||||
key_img = key_img.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
value_img = value_img.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
|
||||
hidden_states_img = F.scaled_dot_product_attention(
|
||||
query, key_img, value_img, attn_mask=None, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
hidden_states_img = hidden_states_img.transpose(1, 2).flatten(2, 3)
|
||||
hidden_states_img = hidden_states_img.type_as(query)
|
||||
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
hidden_states = hidden_states.transpose(1, 2).flatten(2, 3)
|
||||
hidden_states = hidden_states.type_as(query)
|
||||
|
||||
if hidden_states_img is not None:
|
||||
hidden_states = hidden_states + hidden_states_img
|
||||
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanImageEmbedding(torch.nn.Module):
|
||||
def __init__(self, in_features: int, out_features: int, pos_embed_seq_len=None):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = FP32LayerNorm(in_features)
|
||||
self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu")
|
||||
self.norm2 = FP32LayerNorm(out_features)
|
||||
if pos_embed_seq_len is not None:
|
||||
self.pos_embed = nn.Parameter(torch.zeros(1, pos_embed_seq_len, in_features))
|
||||
else:
|
||||
self.pos_embed = None
|
||||
|
||||
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
if self.pos_embed is not None:
|
||||
batch_size, seq_len, embed_dim = encoder_hidden_states_image.shape
|
||||
encoder_hidden_states_image = encoder_hidden_states_image.view(-1, 2 * seq_len, embed_dim)
|
||||
encoder_hidden_states_image = encoder_hidden_states_image + self.pos_embed
|
||||
|
||||
hidden_states = self.norm1(encoder_hidden_states_image)
|
||||
hidden_states = self.ff(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanTimeTextImageEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
time_freq_dim: int,
|
||||
time_proj_dim: int,
|
||||
text_embed_dim: int,
|
||||
image_embed_dim: Optional[int] = None,
|
||||
pos_embed_seq_len: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0)
|
||||
self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim)
|
||||
self.act_fn = nn.SiLU()
|
||||
self.time_proj = nn.Linear(dim, time_proj_dim)
|
||||
self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh")
|
||||
|
||||
self.image_embedder = None
|
||||
if image_embed_dim is not None:
|
||||
self.image_embedder = WanImageEmbedding(image_embed_dim, dim, pos_embed_seq_len=pos_embed_seq_len)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_image: Optional[torch.Tensor] = None,
|
||||
):
|
||||
timestep = self.timesteps_proj(timestep)
|
||||
|
||||
time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype
|
||||
if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8:
|
||||
timestep = timestep.to(time_embedder_dtype)
|
||||
temb = self.time_embedder(timestep).type_as(encoder_hidden_states)
|
||||
timestep_proj = self.time_proj(self.act_fn(temb))
|
||||
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image)
|
||||
|
||||
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
|
||||
|
||||
|
||||
class WanRotaryPosEmbed(nn.Module):
|
||||
def __init__(
|
||||
self, attention_head_dim: int, patch_size: Tuple[int, int, int], max_seq_len: int, theta: float = 10000.0
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.attention_head_dim = attention_head_dim
|
||||
self.patch_size = patch_size
|
||||
self.max_seq_len = max_seq_len
|
||||
|
||||
h_dim = w_dim = 2 * (attention_head_dim // 6)
|
||||
t_dim = attention_head_dim - h_dim - w_dim
|
||||
|
||||
freqs = []
|
||||
for dim in [t_dim, h_dim, w_dim]:
|
||||
freq = get_1d_rotary_pos_embed(
|
||||
dim, max_seq_len, theta, use_real=False, repeat_interleave_real=False, freqs_dtype=torch.float64
|
||||
)
|
||||
freqs.append(freq)
|
||||
self.freqs = torch.cat(freqs, dim=1)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
ppf, pph, ppw = num_frames // p_t, height // p_h, width // p_w
|
||||
|
||||
freqs = self.freqs.to(hidden_states.device)
|
||||
freqs = freqs.split_with_sizes(
|
||||
[
|
||||
self.attention_head_dim // 2 - 2 * (self.attention_head_dim // 6),
|
||||
self.attention_head_dim // 6,
|
||||
self.attention_head_dim // 6,
|
||||
],
|
||||
dim=1,
|
||||
)
|
||||
|
||||
freqs_f = freqs[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1)
|
||||
freqs_h = freqs[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1)
|
||||
freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1)
|
||||
freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1)
|
||||
return freqs
|
||||
|
||||
|
||||
class WanTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.attn1 = Attention(
|
||||
query_dim=dim,
|
||||
heads=num_heads,
|
||||
kv_heads=num_heads,
|
||||
dim_head=dim // num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps,
|
||||
bias=True,
|
||||
cross_attention_dim=None,
|
||||
out_bias=True,
|
||||
processor=WanAttnProcessor2_0(),
|
||||
)
|
||||
|
||||
# 2. Cross-attention
|
||||
self.attn2 = Attention(
|
||||
query_dim=dim,
|
||||
heads=num_heads,
|
||||
kv_heads=num_heads,
|
||||
dim_head=dim // num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps,
|
||||
bias=True,
|
||||
cross_attention_dim=None,
|
||||
out_bias=True,
|
||||
added_kv_proj_dim=added_kv_proj_dim,
|
||||
added_proj_bias=True,
|
||||
processor=WanAttnProcessor2_0(),
|
||||
)
|
||||
self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate")
|
||||
self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
rotary_emb: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
|
||||
self.scale_shift_table + temb.float()
|
||||
).chunk(6, dim=1)
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states)
|
||||
attn_output = self.attn1(hidden_states=norm_hidden_states, rotary_emb=rotary_emb)
|
||||
hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states)
|
||||
|
||||
# 2. Cross-attention
|
||||
norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states)
|
||||
attn_output = self.attn2(hidden_states=norm_hidden_states, encoder_hidden_states=encoder_hidden_states)
|
||||
hidden_states = hidden_states + attn_output
|
||||
|
||||
# 3. Feed-forward
|
||||
norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as(
|
||||
hidden_states
|
||||
)
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin):
|
||||
r"""
|
||||
A Transformer model for video-like data used in the Wan model.
|
||||
|
||||
Args:
|
||||
patch_size (`Tuple[int]`, defaults to `(1, 2, 2)`):
|
||||
3D patch dimensions for video embedding (t_patch, h_patch, w_patch).
|
||||
num_attention_heads (`int`, defaults to `40`):
|
||||
Fixed length for text embeddings.
|
||||
attention_head_dim (`int`, defaults to `128`):
|
||||
The number of channels in each head.
|
||||
in_channels (`int`, defaults to `16`):
|
||||
The number of channels in the input.
|
||||
out_channels (`int`, defaults to `16`):
|
||||
The number of channels in the output.
|
||||
text_dim (`int`, defaults to `512`):
|
||||
Input dimension for text embeddings.
|
||||
freq_dim (`int`, defaults to `256`):
|
||||
Dimension for sinusoidal time embeddings.
|
||||
ffn_dim (`int`, defaults to `13824`):
|
||||
Intermediate dimension in feed-forward network.
|
||||
num_layers (`int`, defaults to `40`):
|
||||
The number of layers of transformer blocks to use.
|
||||
window_size (`Tuple[int]`, defaults to `(-1, -1)`):
|
||||
Window size for local attention (-1 indicates global attention).
|
||||
cross_attn_norm (`bool`, defaults to `True`):
|
||||
Enable cross-attention normalization.
|
||||
qk_norm (`bool`, defaults to `True`):
|
||||
Enable query/key normalization.
|
||||
eps (`float`, defaults to `1e-6`):
|
||||
Epsilon value for normalization layers.
|
||||
add_img_emb (`bool`, defaults to `False`):
|
||||
Whether to use img_emb.
|
||||
added_kv_proj_dim (`int`, *optional*, defaults to `None`):
|
||||
The number of channels to use for the added key and value projections. If `None`, no projection is used.
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
_skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"]
|
||||
_no_split_modules = ["WanTransformerBlock"]
|
||||
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
|
||||
_keys_to_ignore_on_load_unexpected = ["norm_added_q"]
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: Tuple[int] = (1, 2, 2),
|
||||
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: Optional[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,
|
||||
pos_embed_seq_len: Optional[int] = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
out_channels = out_channels or in_channels
|
||||
|
||||
# 1. Patch & position embedding
|
||||
self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len)
|
||||
self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size)
|
||||
|
||||
# 2. Condition embeddings
|
||||
# image_embedding_dim=1280 for I2V model
|
||||
self.condition_embedder = WanTimeTextImageEmbedding(
|
||||
dim=inner_dim,
|
||||
time_freq_dim=freq_dim,
|
||||
time_proj_dim=inner_dim * 6,
|
||||
text_embed_dim=text_dim,
|
||||
image_embed_dim=image_dim,
|
||||
pos_embed_seq_len=pos_embed_seq_len,
|
||||
)
|
||||
|
||||
# 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
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Output norm & projection
|
||||
self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False)
|
||||
self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size))
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_image: Optional[torch.Tensor] = None,
|
||||
return_dict: bool = True,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
if attention_kwargs is not None:
|
||||
attention_kwargs = attention_kwargs.copy()
|
||||
lora_scale = attention_kwargs.pop("scale", 1.0)
|
||||
else:
|
||||
lora_scale = 1.0
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# weight the lora layers by setting `lora_scale` for each PEFT layer
|
||||
scale_lora_layers(self, lora_scale)
|
||||
else:
|
||||
if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None:
|
||||
logger.warning(
|
||||
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
|
||||
)
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.config.patch_size
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
rotary_emb = self.rope(hidden_states)
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image
|
||||
)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
|
||||
# 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, rotary_emb
|
||||
)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb)
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
|
||||
|
||||
# Move the shift and scale tensors to the same device as hidden_states.
|
||||
# When using multi-GPU inference via accelerate these will be on the
|
||||
# first device rather than the last device, which hidden_states ends up
|
||||
# on.
|
||||
shift = shift.to(hidden_states.device)
|
||||
scale = scale.to(hidden_states.device)
|
||||
|
||||
hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(
|
||||
batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1
|
||||
)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# remove `lora_scale` from each PEFT layer
|
||||
unscale_lora_layers(self, lora_scale)
|
||||
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
|
||||
return Transformer2DModelOutput(sample=output)
|
||||
@@ -0,0 +1,609 @@
|
||||
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import html
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
|
||||
import regex as re
|
||||
import torch
|
||||
from transformers import AutoTokenizer, UMT5EncoderModel
|
||||
|
||||
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
||||
from diffusers.loaders import WanLoraLoaderMixin
|
||||
from diffusers.models import AutoencoderKLWan, WanTransformer3DModel
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import is_ftfy_available, is_torch_xla_available, logging, replace_example_docstring
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.pipelines.wan.pipeline_output import WanPipelineOutput
|
||||
from einops import rearrange
|
||||
from transformers import UMT5EncoderModel, T5TokenizerFast
|
||||
|
||||
from fastvideo.models.mochi_hf.modeling_wan import WanTransformer3DModel
|
||||
from fastvideo.utils.communications import all_gather
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
|
||||
XLA_AVAILABLE = True
|
||||
else:
|
||||
XLA_AVAILABLE = False
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
if is_ftfy_available():
|
||||
import ftfy
|
||||
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
Examples:
|
||||
```python
|
||||
>>> import torch
|
||||
>>> from diffusers.utils import export_to_video
|
||||
>>> from diffusers import AutoencoderKLWan, WanPipeline
|
||||
>>> from diffusers.schedulers.scheduling_unipc_multistep import UniPCMultistepScheduler
|
||||
|
||||
>>> # Available models: Wan-AI/Wan2.1-T2V-14B-Diffusers, Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
>>> model_id = "Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
>>> vae = AutoencoderKLWan.from_pretrained(model_id, subfolder="vae", torch_dtype=torch.float32)
|
||||
>>> pipe = WanPipeline.from_pretrained(model_id, vae=vae, torch_dtype=torch.bfloat16)
|
||||
>>> flow_shift = 5.0 # 5.0 for 720P, 3.0 for 480P
|
||||
>>> pipe.scheduler = UniPCMultistepScheduler.from_config(pipe.scheduler.config, flow_shift=flow_shift)
|
||||
>>> pipe.to("cuda")
|
||||
|
||||
>>> prompt = "A cat and a dog baking a cake together in a kitchen. The cat is carefully measuring flour, while the dog is stirring the batter with a wooden spoon. The kitchen is cozy, with sunlight streaming through the window."
|
||||
>>> negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
|
||||
>>> output = pipe(
|
||||
... prompt=prompt,
|
||||
... negative_prompt=negative_prompt,
|
||||
... height=720,
|
||||
... width=1280,
|
||||
... num_frames=81,
|
||||
... guidance_scale=5.0,
|
||||
... ).frames[0]
|
||||
>>> export_to_video(output, "output.mp4", fps=16)
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
def basic_clean(text):
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
|
||||
def whitespace_clean(text):
|
||||
text = re.sub(r"\s+", " ", text)
|
||||
text = text.strip()
|
||||
return text
|
||||
|
||||
|
||||
def prompt_clean(text):
|
||||
text = whitespace_clean(basic_clean(text))
|
||||
return text
|
||||
|
||||
|
||||
class WanPipeline(DiffusionPipeline, WanLoraLoaderMixin):
|
||||
r"""
|
||||
Pipeline for text-to-video generation using Wan.
|
||||
|
||||
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
|
||||
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
|
||||
|
||||
Args:
|
||||
tokenizer ([`T5Tokenizer`]):
|
||||
Tokenizer from [T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5Tokenizer),
|
||||
specifically the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant.
|
||||
text_encoder ([`T5EncoderModel`]):
|
||||
[T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
|
||||
the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant.
|
||||
transformer ([`WanTransformer3DModel`]):
|
||||
Conditional Transformer to denoise the input latents.
|
||||
scheduler ([`UniPCMultistepScheduler`]):
|
||||
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
|
||||
vae ([`AutoencoderKLWan`]):
|
||||
Variational Auto-Encoder (VAE) Model to encode and decode videos to and from latent representations.
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "text_encoder->transformer->vae"
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: AutoTokenizer,
|
||||
text_encoder: UMT5EncoderModel,
|
||||
transformer: WanTransformer3DModel,
|
||||
vae: AutoencoderKLWan,
|
||||
scheduler: FlowMatchEulerDiscreteScheduler,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
|
||||
self.vae_scale_factor_temporal = 2 ** sum(self.vae.temperal_downsample) if getattr(self, "vae", None) else 4
|
||||
self.vae_scale_factor_spatial = 2 ** len(self.vae.temperal_downsample) if getattr(self, "vae", None) else 8
|
||||
self.video_processor = VideoProcessor(vae_scale_factor=self.vae_scale_factor_spatial)
|
||||
|
||||
def _get_t5_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
num_videos_per_prompt: int = 1,
|
||||
max_sequence_length: int = 226,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
prompt = [prompt_clean(u) for u in prompt]
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_attention_mask=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), mask.to(device)).last_hidden_state
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
|
||||
prompt_embeds = torch.stack(
|
||||
[torch.cat([u, u.new_zeros(max_sequence_length - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0
|
||||
)
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
do_classifier_free_guidance: bool = True,
|
||||
num_videos_per_prompt: int = 1,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
max_sequence_length: int = 226,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
r"""
|
||||
Encodes the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
prompt to be encoded
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts not to guide the image generation. If not defined, one has to pass
|
||||
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
|
||||
less than `1`).
|
||||
do_classifier_free_guidance (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use classifier free guidance or not.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
Number of videos that should be generated per prompt. torch device to place the resulting embeddings on
|
||||
prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
negative_prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
|
||||
argument.
|
||||
device: (`torch.device`, *optional*):
|
||||
torch device
|
||||
dtype: (`torch.dtype`, *optional*):
|
||||
torch dtype
|
||||
"""
|
||||
device = device or self._execution_device
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
if prompt is not None:
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt_embeds = self._get_t5_prompt_embeds(
|
||||
prompt=prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
if do_classifier_free_guidance and negative_prompt_embeds is None:
|
||||
negative_prompt = negative_prompt or ""
|
||||
negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt
|
||||
|
||||
if prompt is not None and type(prompt) is not type(negative_prompt):
|
||||
raise TypeError(
|
||||
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
|
||||
f" {type(prompt)}."
|
||||
)
|
||||
elif batch_size != len(negative_prompt):
|
||||
raise ValueError(
|
||||
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
|
||||
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
|
||||
" the batch size of `prompt`."
|
||||
)
|
||||
|
||||
negative_prompt_embeds = self._get_t5_prompt_embeds(
|
||||
prompt=negative_prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
return prompt_embeds, negative_prompt_embeds
|
||||
|
||||
def check_inputs(
|
||||
self,
|
||||
prompt,
|
||||
negative_prompt,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds=None,
|
||||
negative_prompt_embeds=None,
|
||||
callback_on_step_end_tensor_inputs=None,
|
||||
):
|
||||
if height % 16 != 0 or width % 16 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 16 but are {height} and {width}.")
|
||||
|
||||
if callback_on_step_end_tensor_inputs is not None and not all(
|
||||
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
|
||||
)
|
||||
|
||||
if prompt is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif negative_prompt is not None and negative_prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`: {negative_prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif prompt is None and prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
|
||||
)
|
||||
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
|
||||
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||
elif negative_prompt is not None and (
|
||||
not isinstance(negative_prompt, str) and not isinstance(negative_prompt, list)
|
||||
):
|
||||
raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}")
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size: int,
|
||||
num_channels_latents: int = 16,
|
||||
height: int = 480,
|
||||
width: int = 832,
|
||||
num_frames: int = 81,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if latents is not None:
|
||||
return latents.to(device=device, dtype=dtype)
|
||||
|
||||
num_latent_frames = (num_frames - 1) // self.vae_scale_factor_temporal + 1
|
||||
shape = (
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
num_latent_frames,
|
||||
int(height) // self.vae_scale_factor_spatial,
|
||||
int(width) // self.vae_scale_factor_spatial,
|
||||
)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
return latents
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
@property
|
||||
def do_classifier_free_guidance(self):
|
||||
return self._guidance_scale > 1.0
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
@property
|
||||
def current_timestep(self):
|
||||
return self._current_timestep
|
||||
|
||||
@property
|
||||
def interrupt(self):
|
||||
return self._interrupt
|
||||
|
||||
@property
|
||||
def attention_kwargs(self):
|
||||
return self._attention_kwargs
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
negative_prompt: Union[str, List[str]] = None,
|
||||
height: int = 480,
|
||||
width: int = 832,
|
||||
num_frames: int = 81,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 5.0,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
output_type: Optional[str] = "np",
|
||||
return_dict: bool = True,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
callback_on_step_end: Optional[
|
||||
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
|
||||
] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 512,
|
||||
):
|
||||
r"""
|
||||
The call function to the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
height (`int`, defaults to `480`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, defaults to `832`):
|
||||
The width in pixels of the generated image.
|
||||
num_frames (`int`, defaults to `81`):
|
||||
The number of frames in the generated video.
|
||||
num_inference_steps (`int`, defaults to `50`):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
guidance_scale (`float`, defaults to `5.0`):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion
|
||||
Guidance](https://huggingface.co/papers/2207.12598). `guidance_scale` is defined as `w` of equation 2.
|
||||
of [Imagen Paper](https://huggingface.co/papers/2205.11487). Guidance scale is enabled by setting
|
||||
`guidance_scale > 1`. Higher guidance scale encourages to generate images that are closely linked to
|
||||
the text `prompt`, usually at the expense of lower image quality.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
|
||||
generation deterministic.
|
||||
latents (`torch.Tensor`, *optional*):
|
||||
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor is generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not
|
||||
provided, text embeddings are generated from the `prompt` input argument.
|
||||
output_type (`str`, *optional*, defaults to `"np"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`WanPipelineOutput`] instead of a plain tuple.
|
||||
attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
callback_on_step_end (`Callable`, `PipelineCallback`, `MultiPipelineCallbacks`, *optional*):
|
||||
A function or a subclass of `PipelineCallback` or `MultiPipelineCallbacks` that is called at the end of
|
||||
each denoising step during the inference. with the following arguments: `callback_on_step_end(self:
|
||||
DiffusionPipeline, step: int, timestep: int, callback_kwargs: Dict)`. `callback_kwargs` will include a
|
||||
list of all tensors as specified by `callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
autocast_dtype (`torch.dtype`, *optional*, defaults to `torch.bfloat16`):
|
||||
The dtype to use for the torch.amp.autocast.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~WanPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`, [`WanPipelineOutput`] is returned, otherwise a `tuple` is returned where
|
||||
the first element is a list with the generated images and the second element is a list of `bool`s
|
||||
indicating whether the corresponding generated image contains "not-safe-for-work" (nsfw) content.
|
||||
"""
|
||||
|
||||
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
|
||||
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
callback_on_step_end_tensor_inputs,
|
||||
)
|
||||
|
||||
if num_frames % self.vae_scale_factor_temporal != 1:
|
||||
logger.warning(
|
||||
f"`num_frames - 1` has to be divisible by {self.vae_scale_factor_temporal}. Rounding to the nearest number."
|
||||
)
|
||||
num_frames = num_frames // self.vae_scale_factor_temporal * self.vae_scale_factor_temporal + 1
|
||||
num_frames = max(num_frames, 1)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._attention_kwargs = attention_kwargs
|
||||
self._current_timestep = None
|
||||
self._interrupt = False
|
||||
|
||||
device = self._execution_device
|
||||
|
||||
# 2. Define call parameters
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
# 3. Encode input prompt
|
||||
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
do_classifier_free_guidance=self.do_classifier_free_guidance,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
)
|
||||
|
||||
transformer_dtype = self.transformer.dtype
|
||||
prompt_embeds = prompt_embeds.to(transformer_dtype)
|
||||
if negative_prompt_embeds is not None:
|
||||
negative_prompt_embeds = negative_prompt_embeds.to(transformer_dtype)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
torch.float32,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
|
||||
if get_sequence_parallel_state():
|
||||
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
|
||||
latents = latents[:, :, rank, :, :, :]
|
||||
|
||||
# 6. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
self._progress_bar_config = {"disable": nccl_info.rank_within_group != 0}
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
self._current_timestep = t
|
||||
latent_model_input = latents.to(transformer_dtype)
|
||||
timestep = t.expand(latents.shape[0])
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if self.do_classifier_free_guidance:
|
||||
noise_uncond = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=negative_prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred = noise_uncond + guidance_scale * (noise_pred - noise_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
if XLA_AVAILABLE:
|
||||
xm.mark_step()
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
latents = all_gather(latents, dim=2)
|
||||
|
||||
self._current_timestep = None
|
||||
|
||||
if not output_type == "latent":
|
||||
latents = latents.to(self.vae.dtype)
|
||||
latents_mean = (
|
||||
torch.tensor(self.vae.config.latents_mean)
|
||||
.view(1, self.vae.config.z_dim, 1, 1, 1)
|
||||
.to(latents.device, latents.dtype)
|
||||
)
|
||||
latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to(
|
||||
latents.device, latents.dtype
|
||||
)
|
||||
latents = latents / latents_std + latents_mean
|
||||
video = self.vae.decode(latents, return_dict=False)[0]
|
||||
video = self.video_processor.postprocess_video(video, output_type=output_type)
|
||||
else:
|
||||
video = latents
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (video,)
|
||||
|
||||
return WanPipelineOutput(frames=video)
|
||||
@@ -86,7 +86,7 @@ def inference(args):
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
generator=generator,
|
||||
).frames
|
||||
if nccl_info.global_rank <= 0:
|
||||
if nccl_info.global_rank == 0:
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(
|
||||
@@ -107,7 +107,7 @@ def inference(args):
|
||||
generator=generator,
|
||||
).frames
|
||||
|
||||
if nccl_info.global_rank <= 0:
|
||||
if nccl_info.global_rank == 0:
|
||||
export_to_video(videos[0], args.output_path + ".mp4", fps=24)
|
||||
|
||||
|
||||
|
||||
@@ -94,7 +94,7 @@ def main(args):
|
||||
guidance_scale=args.guidance_scale,
|
||||
generator=generator,
|
||||
).frames
|
||||
if nccl_info.global_rank <= 0:
|
||||
if nccl_info.global_rank == 0:
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(
|
||||
@@ -116,7 +116,7 @@ def main(args):
|
||||
generator=generator,
|
||||
).frames
|
||||
|
||||
if nccl_info.global_rank <= 0:
|
||||
if nccl_info.global_rank == 0:
|
||||
export_to_video(videos[0], args.output_path + ".mp4", fps=30)
|
||||
|
||||
|
||||
|
||||
+6
-6
@@ -20,7 +20,7 @@ from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.utils.latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.models.hunyuan_hf.pipeline_hunyuan import HunyuanVideoPipeline
|
||||
|
||||
@@ -185,7 +185,7 @@ def main(args):
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
if rank == 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
@@ -316,7 +316,7 @@ def main(args):
|
||||
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -393,7 +393,7 @@ def main(args):
|
||||
"grad_norm": grad_norm,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss": loss,
|
||||
|
||||
@@ -32,7 +32,7 @@ def save_checkpoint_optimizer(model, optimizer, rank, output_dir, step, discrimi
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank <= 0 and not discriminator:
|
||||
if rank == 0 and not discriminator:
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
config_dict = dict(model.config)
|
||||
@@ -60,7 +60,7 @@ def save_checkpoint(transformer, rank, output_dir, step):
|
||||
):
|
||||
cpu_state = transformer.state_dict()
|
||||
# todo move to get_state_dict
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
@@ -98,7 +98,7 @@ def save_checkpoint_generator_discriminator(
|
||||
hf_weight_dir = os.path.join(save_dir, "hf_weights")
|
||||
os.makedirs(hf_weight_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
config_dict = dict(model.config)
|
||||
config_path = os.path.join(hf_weight_dir, "config.json")
|
||||
# save dict as json
|
||||
@@ -139,7 +139,7 @@ def save_checkpoint_generator_discriminator(
|
||||
optim_state = FSDP.optim_state_dict(discriminator, discriminator_optimizer)
|
||||
model_state = discriminator.state_dict()
|
||||
state_dict = {"optimizer": optim_state, "model": model_state}
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
discriminator_fsdp_state_fil = os.path.join(discriminator_fsdp_state_dir, "discriminator_state.pt")
|
||||
torch.save(state_dict, discriminator_fsdp_state_fil)
|
||||
|
||||
@@ -178,7 +178,7 @@ def load_full_state_model(model, optimizer, checkpoint_file, rank):
|
||||
):
|
||||
discriminator_state = torch.load(checkpoint_file)
|
||||
model_state = discriminator_state["model"]
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
optim_state = discriminator_state["optimizer"]
|
||||
else:
|
||||
optim_state = None
|
||||
@@ -241,7 +241,7 @@ def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step, pipelin
|
||||
optimizer,
|
||||
)
|
||||
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
|
||||
|
||||
@@ -0,0 +1,776 @@
|
||||
# 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",
|
||||
"accelerate",
|
||||
"peft",
|
||||
"zmq",
|
||||
"nvidia",
|
||||
"pynvml",
|
||||
}
|
||||
|
||||
DEFAULT_PIP_PATTERNS = {
|
||||
"torch",
|
||||
"numpy",
|
||||
"mypy",
|
||||
"flake8",
|
||||
"triton",
|
||||
"optree",
|
||||
"onnx",
|
||||
"nccl",
|
||||
"transformers",
|
||||
"accelerate",
|
||||
"peft",
|
||||
"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,38 +0,0 @@
|
||||
import platform
|
||||
|
||||
import accelerate
|
||||
import peft
|
||||
import torch
|
||||
import transformers
|
||||
from transformers.utils import is_torch_cuda_available, is_torch_npu_available
|
||||
|
||||
VERSION = "1.2.0"
|
||||
|
||||
if __name__ == "__main__":
|
||||
info = {
|
||||
"FastVideo version": VERSION,
|
||||
"Platform": platform.platform(),
|
||||
"Python version": platform.python_version(),
|
||||
"PyTorch version": torch.__version__,
|
||||
"Transformers version": transformers.__version__,
|
||||
"Accelerate version": accelerate.__version__,
|
||||
"PEFT version": peft.__version__,
|
||||
}
|
||||
|
||||
if is_torch_cuda_available():
|
||||
info["PyTorch version"] += " (GPU)"
|
||||
info["GPU type"] = torch.cuda.get_device_name()
|
||||
|
||||
if is_torch_npu_available():
|
||||
info["PyTorch version"] += " (NPU)"
|
||||
info["NPU type"] = torch.npu.get_device_name()
|
||||
info["CANN version"] = torch.version.cann # codespell:ignore
|
||||
|
||||
try:
|
||||
import bitsandbytes
|
||||
|
||||
info["Bitsandbytes version"] = bitsandbytes.__version__
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
print("\n" + "\n".join([f"- {key}: {value}" for key, value in info.items()]) + "\n")
|
||||
@@ -0,0 +1,88 @@
|
||||
import torch
|
||||
|
||||
mochi_latents_mean = torch.tensor([
|
||||
-0.06730895953510081,
|
||||
-0.038011381506090416,
|
||||
-0.07477820912866141,
|
||||
-0.05565264470995561,
|
||||
0.012767231469026969,
|
||||
-0.04703542746246419,
|
||||
0.043896967884726704,
|
||||
-0.09346305707025976,
|
||||
-0.09918314763016893,
|
||||
-0.008729793427399178,
|
||||
-0.011931556316503654,
|
||||
-0.0321993391887285,
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
mochi_latents_std = torch.tensor([
|
||||
0.9263795028493863,
|
||||
0.9248894543193766,
|
||||
0.9393059390890617,
|
||||
0.959253732819592,
|
||||
0.8244560132752793,
|
||||
0.917259975397747,
|
||||
0.9294154431013696,
|
||||
1.3720942357788521,
|
||||
0.881393668867029,
|
||||
0.9168315692124348,
|
||||
0.9185249279345552,
|
||||
0.9274757570805041,
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
mochi_scaling_factor = 1.0
|
||||
|
||||
|
||||
wan_latents_mean = torch.tensor([
|
||||
-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,
|
||||
]).view(1, 16, 1, 1, 1)
|
||||
wan_latents_std = torch.tensor([
|
||||
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.916,
|
||||
]).view(1, 16, 1, 1, 1)
|
||||
|
||||
|
||||
def normalize_dit_input(model_type, latents):
|
||||
if model_type == "mochi":
|
||||
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
|
||||
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
|
||||
latents = (latents - latents_mean) / latents_std
|
||||
return latents
|
||||
elif model_type == "hunyuan_hf":
|
||||
return latents * 0.476986
|
||||
elif model_type == "hunyuan":
|
||||
return latents * 0.476986
|
||||
elif model_type == "wan":
|
||||
latents_mean = wan_latents_mean.to(latents.device, latents.dtype)
|
||||
latents_std = wan_latents_std.to(latents.device, latents.dtype)
|
||||
latents = (latents - latents_mean) / latents_std
|
||||
return latents
|
||||
else:
|
||||
raise NotImplementedError(f"model_type {model_type} not supported")
|
||||
+69
-2
@@ -3,9 +3,9 @@ from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi
|
||||
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi, AutoencoderKLWan
|
||||
from torch import nn
|
||||
from transformers import AutoTokenizer, T5EncoderModel
|
||||
from transformers import AutoTokenizer, T5EncoderModel, UMT5EncoderModel
|
||||
|
||||
from fastvideo.models.hunyuan.modules.models import (HYVideoDiffusionTransformer, MMDoubleStreamBlock,
|
||||
MMSingleStreamBlock)
|
||||
@@ -14,6 +14,7 @@ from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import AutoencoderKLC
|
||||
from fastvideo.models.hunyuan_hf.modeling_hunyuan import (HunyuanVideoSingleTransformerBlock,
|
||||
HunyuanVideoTransformer3DModel, HunyuanVideoTransformerBlock)
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel, MochiTransformerBlock
|
||||
from fastvideo.models.wan_hf.modeling_wan import WanTransformer3DModel, WanTransformerBlock
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
hunyuan_config = {
|
||||
@@ -200,6 +201,48 @@ class MochiTextEncoderWrapper(nn.Module):
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
class WanTextEncoderWrapper(nn.Module):
|
||||
|
||||
def __init__(self, pretrained_model_name_or_path, device):
|
||||
super().__init__()
|
||||
self.text_encoder = UMT5EncoderModel.from_pretrained(os.path.join(pretrained_model_name_or_path,
|
||||
"text_encoder")).to(device)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(os.path.join(pretrained_model_name_or_path, "tokenizer"))
|
||||
self.max_sequence_length = 256
|
||||
|
||||
def encode_prompt(self, prompt):
|
||||
device = self.text_encoder.device
|
||||
dtype = self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=self.max_sequence_length,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
prompt_attention_mask = text_inputs.attention_mask
|
||||
prompt_attention_mask = prompt_attention_mask.bool().to(device)
|
||||
|
||||
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
|
||||
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.max_sequence_length - 1:-1])
|
||||
main_print(f"Truncated text input: {prompt} to: {removed_text} for model input.")
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.view(batch_size, seq_len, -1)
|
||||
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
def load_hunyuan_state_dict(model, dit_model_name_or_path):
|
||||
load_key = "module"
|
||||
@@ -240,6 +283,20 @@ def load_transformer(
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
elif model_type == "wan":
|
||||
if dit_model_name_or_path:
|
||||
transformer = WanTransformer3DModel.from_pretrained(
|
||||
dit_model_name_or_path,
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
else:
|
||||
transformer = WanTransformer3DModel.from_pretrained(
|
||||
pretrained_model_name_or_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
elif model_type == "hunyuan_hf":
|
||||
if dit_model_name_or_path:
|
||||
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
|
||||
@@ -283,6 +340,12 @@ def load_vae(model_type, pretrained_model_name_or_path):
|
||||
torch_dtype=weight_dtype).to("cuda")
|
||||
autocast_type = torch.bfloat16
|
||||
fps = 24
|
||||
elif model_type == "wan":
|
||||
vae = AutoencoderKLWan.from_pretrained(pretrained_model_name_or_path,
|
||||
subfolder="vae",
|
||||
torch_dtype=weight_dtype).to("cuda")
|
||||
autocast_type = torch.bfloat16
|
||||
fps = 24
|
||||
elif model_type == "hunyuan":
|
||||
vae_precision = torch.float32
|
||||
vae_path = os.path.join(pretrained_model_name_or_path, "hunyuan-video-t2v-720p/vae")
|
||||
@@ -311,6 +374,8 @@ def load_vae(model_type, pretrained_model_name_or_path):
|
||||
def load_text_encoder(model_type, pretrained_model_name_or_path, device):
|
||||
if model_type == "mochi":
|
||||
text_encoder = MochiTextEncoderWrapper(pretrained_model_name_or_path, device)
|
||||
elif model_type == "wan":
|
||||
text_encoder = WanTextEncoderWrapper(pretrained_model_name_or_path, device)
|
||||
elif model_type == "hunyuan" or "hunyuan_hf":
|
||||
text_encoder = HunyuanTextEncoderWrapper(pretrained_model_name_or_path, device)
|
||||
else:
|
||||
@@ -322,6 +387,8 @@ def get_no_split_modules(transformer):
|
||||
# if of type MochiTransformer3DModel
|
||||
if isinstance(transformer, MochiTransformer3DModel):
|
||||
return (MochiTransformerBlock, )
|
||||
elif isinstance(transformer, WanTransformer3DModel):
|
||||
return (WanTransformerBlock, )
|
||||
elif isinstance(transformer, HunyuanVideoTransformer3DModel):
|
||||
return (HunyuanVideoSingleTransformerBlock, HunyuanVideoTransformerBlock)
|
||||
elif isinstance(transformer, HYVideoDiffusionTransformer):
|
||||
|
||||
@@ -129,13 +129,22 @@ def sample_validation_video(
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latent_model_input.shape[0])
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
noise_pred = transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
if model_type == "wan":
|
||||
pred_kwargs = {
|
||||
"hidden_states": latent_model_input,
|
||||
"encoder_hidden_states": prompt_embeds,
|
||||
"timestep":timestep,
|
||||
"return_dict":False,
|
||||
}
|
||||
else:
|
||||
pred_kwargs = {
|
||||
"hidden_states": latent_model_input,
|
||||
"encoder_hidden_states": prompt_embeds,
|
||||
"timestep":timestep,
|
||||
"encoder_attention_mask":prompt_attention_mask,
|
||||
"return_dict":False,
|
||||
}
|
||||
noise_pred = transformer(**pred_kwargs)[0]
|
||||
|
||||
# Mochi CFG + Sampling runs in FP32
|
||||
noise_pred = noise_pred.to(torch.float32)
|
||||
@@ -166,10 +175,12 @@ def sample_validation_video(
|
||||
# denormalize with the mean and std if available and not None
|
||||
has_latents_mean = (hasattr(vae.config, "latents_mean") and vae.config.latents_mean is not None)
|
||||
has_latents_std = (hasattr(vae.config, "latents_std") and vae.config.latents_std is not None)
|
||||
if model_type == "wan":
|
||||
vae.config.scaling_factor = 1
|
||||
if has_latents_mean and has_latents_std:
|
||||
latents_mean = (torch.tensor(vae.config.latents_mean).view(1, 12, 1, 1,
|
||||
latents_mean = (torch.tensor(vae.config.latents_mean).view(1, num_channels_latents, 1, 1,
|
||||
1).to(latents.device, latents.dtype))
|
||||
latents_std = (torch.tensor(vae.config.latents_std).view(1, 12, 1, 1, 1).to(latents.device, latents.dtype))
|
||||
latents_std = (torch.tensor(vae.config.latents_std).view(1, num_channels_latents, 1, 1, 1).to(latents.device, latents.dtype))
|
||||
latents = latents * latents_std / vae.config.scaling_factor + latents_mean
|
||||
else:
|
||||
latents = latents / vae.config.scaling_factor
|
||||
@@ -202,14 +213,15 @@ def log_validation(
|
||||
vae_spatial_scale_factor = 8
|
||||
vae_temporal_scale_factor = 6
|
||||
num_channels_latents = 12
|
||||
elif args.model_type == "hunyuan" or "hunyuan_hf":
|
||||
elif args.model_type == "hunyuan" or "hunyuan_hf" or "wan":
|
||||
vae_spatial_scale_factor = 8
|
||||
vae_temporal_scale_factor = 4
|
||||
num_channels_latents = 16
|
||||
else:
|
||||
raise ValueError(f"Model type {args.model_type} not supported")
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.pretrained_model_name_or_path)
|
||||
vae.enable_tiling()
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
if scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler(shift=shift)
|
||||
else:
|
||||
|
||||
@@ -0,0 +1,419 @@
|
||||
import json
|
||||
import os
|
||||
from collections import defaultdict
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def configure_sta(mode: str = 'STA_searching',
|
||||
layer_num: int = 40,
|
||||
time_step_num: int = 50,
|
||||
head_num: int = 40,
|
||||
**kwargs) -> List[List[List[Any]]]:
|
||||
"""
|
||||
Configure Sliding Tile Attention (STA) parameters based on the specified mode.
|
||||
|
||||
Parameters:
|
||||
----------
|
||||
mode : str
|
||||
The STA mode to use. Options are:
|
||||
- 'STA_searching': Generate a set of mask candidates for initial search
|
||||
- 'STA_tuning': Select best mask strategy based on previously saved results
|
||||
- 'STA_inference': Load and use a previously tuned mask strategy
|
||||
layer_num: int, number of layers
|
||||
time_step_num: int, number of timesteps
|
||||
head_num: int, number of heads
|
||||
|
||||
**kwargs : dict
|
||||
Mode-specific parameters:
|
||||
|
||||
For 'STA_searching':
|
||||
- mask_candidates: list of str, optional, mask candidates to use
|
||||
- mask_selected: list of int, optional, indices of selected masks
|
||||
|
||||
For 'STA_tuning':
|
||||
- mask_search_files_path: str, required, path to mask search results
|
||||
- mask_candidates: list of str, optional, mask candidates to use
|
||||
- mask_selected: list of int, optional, indices of selected masks
|
||||
- skip_time_steps: int, optional, number of time steps to use full attention (default 12)
|
||||
- save_dir: str, optional, directory to save mask strategy (default "mask_candidates")
|
||||
|
||||
For 'STA_inference':
|
||||
- load_path: str, optional, path to load mask strategy (default "mask_candidates/mask_strategy.json")
|
||||
"""
|
||||
valid_modes = [
|
||||
'STA_searching', 'STA_tuning', 'STA_inference', 'STA_tuning_cfg'
|
||||
]
|
||||
if mode not in valid_modes:
|
||||
raise ValueError(f"Mode must be one of {valid_modes}, got {mode}")
|
||||
|
||||
if mode == 'STA_searching':
|
||||
# Get parameters with defaults
|
||||
mask_candidates: Optional[List[str]] = kwargs.get('mask_candidates')
|
||||
if mask_candidates is None:
|
||||
raise ValueError(
|
||||
"mask_candidates is required for STA_searching mode")
|
||||
mask_selected: List[int] = kwargs.get('mask_selected',
|
||||
list(range(len(mask_candidates))))
|
||||
|
||||
# Parse selected masks
|
||||
selected_masks: List[List[int]] = []
|
||||
for index in mask_selected:
|
||||
mask = mask_candidates[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks.append(masks_list)
|
||||
|
||||
# Create 3D mask structure with fixed dimensions (t=50, l=60)
|
||||
masks_3d: List[List[List[List[int]]]] = []
|
||||
for i in range(time_step_num): # Fixed t dimension = 50
|
||||
row = []
|
||||
for j in range(layer_num): # Fixed l dimension = 60
|
||||
row.append(selected_masks) # Add all masks at each position
|
||||
masks_3d.append(row)
|
||||
|
||||
return masks_3d
|
||||
|
||||
elif mode == 'STA_tuning':
|
||||
# Get required parameters
|
||||
mask_search_files_path: Optional[str] = kwargs.get(
|
||||
'mask_search_files_path')
|
||||
if not mask_search_files_path:
|
||||
raise ValueError(
|
||||
"mask_search_files_path is required for STA_tuning mode")
|
||||
|
||||
# Get optional parameters with defaults
|
||||
mask_candidates_tuning: Optional[List[str]] = kwargs.get(
|
||||
'mask_candidates')
|
||||
if mask_candidates_tuning is None:
|
||||
raise ValueError("mask_candidates is required for STA_tuning mode")
|
||||
mask_selected_tuning: List[int] = kwargs.get(
|
||||
'mask_selected', list(range(len(mask_candidates_tuning))))
|
||||
skip_time_steps_tuning: Optional[int] = kwargs.get('skip_time_steps')
|
||||
save_dir_tuning: Optional[str] = kwargs.get('save_dir',
|
||||
"mask_candidates")
|
||||
|
||||
# Parse selected masks
|
||||
selected_masks_tuning: List[List[int]] = []
|
||||
for index in mask_selected_tuning:
|
||||
mask = mask_candidates_tuning[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks_tuning.append(masks_list)
|
||||
|
||||
# Read JSON results
|
||||
results = read_specific_json_files(mask_search_files_path)
|
||||
averaged_results = average_head_losses(results, selected_masks_tuning)
|
||||
|
||||
# Add full attention mask for specific cases
|
||||
full_attention_mask_tuning: Optional[List[int]] = kwargs.get(
|
||||
'full_attention_mask')
|
||||
if full_attention_mask_tuning is not None:
|
||||
selected_masks_tuning.append(full_attention_mask_tuning)
|
||||
|
||||
# Select best mask strategy
|
||||
timesteps_tuning: int = kwargs.get('timesteps', time_step_num)
|
||||
if skip_time_steps_tuning is None:
|
||||
skip_time_steps_tuning = 12
|
||||
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
|
||||
averaged_results, selected_masks_tuning, skip_time_steps_tuning,
|
||||
timesteps_tuning, head_num)
|
||||
|
||||
# Save mask strategy
|
||||
if save_dir_tuning is not None:
|
||||
os.makedirs(save_dir_tuning, exist_ok=True)
|
||||
file_path = os.path.join(
|
||||
save_dir_tuning,
|
||||
f'mask_strategy_s{skip_time_steps_tuning}.json')
|
||||
with open(file_path, 'w') as f:
|
||||
json.dump(mask_strategy, f, indent=4)
|
||||
print(f"Successfully saved mask_strategy to {file_path}")
|
||||
|
||||
# Print sparsity and strategy counts for information
|
||||
print(f"Overall sparsity: {sparsity:.4f}")
|
||||
print("\nStrategy usage counts:")
|
||||
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
|
||||
for strategy, count in strategy_counts.items():
|
||||
print(
|
||||
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
|
||||
)
|
||||
|
||||
# Convert dictionary to 3D list with fixed dimensions
|
||||
mask_strategy_3d = dict_to_3d_list(mask_strategy,
|
||||
t_max=time_step_num,
|
||||
l_max=layer_num,
|
||||
h_max=head_num)
|
||||
|
||||
return mask_strategy_3d
|
||||
elif mode == 'STA_tuning_cfg':
|
||||
# Get required parameters for both positive and negative paths
|
||||
mask_search_files_path_pos: Optional[str] = kwargs.get(
|
||||
'mask_search_files_path_pos')
|
||||
mask_search_files_path_neg: Optional[str] = kwargs.get(
|
||||
'mask_search_files_path_neg')
|
||||
save_dir_cfg: Optional[str] = kwargs.get('save_dir')
|
||||
|
||||
if not mask_search_files_path_pos or not mask_search_files_path_neg or not save_dir_cfg:
|
||||
raise ValueError(
|
||||
"mask_search_files_path_pos, mask_search_files_path_neg, and save_dir are required for STA_tuning_cfg mode"
|
||||
)
|
||||
|
||||
# Get optional parameters with defaults
|
||||
mask_candidates_cfg: Optional[List[str]] = kwargs.get('mask_candidates')
|
||||
if mask_candidates_cfg is None:
|
||||
raise ValueError(
|
||||
"mask_candidates is required for STA_tuning_cfg mode")
|
||||
mask_selected_cfg: List[int] = kwargs.get(
|
||||
'mask_selected', list(range(len(mask_candidates_cfg))))
|
||||
skip_time_steps_cfg: Optional[int] = kwargs.get('skip_time_steps')
|
||||
|
||||
# Parse selected masks
|
||||
selected_masks_cfg: List[List[int]] = []
|
||||
for index in mask_selected_cfg:
|
||||
mask = mask_candidates_cfg[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks_cfg.append(masks_list)
|
||||
|
||||
# Read JSON results for both positive and negative paths
|
||||
pos_results = read_specific_json_files(mask_search_files_path_pos)
|
||||
neg_results = read_specific_json_files(mask_search_files_path_neg)
|
||||
# Combine positive and negative results into one list
|
||||
combined_results = pos_results + neg_results
|
||||
|
||||
# Average the combined results
|
||||
averaged_results = average_head_losses(combined_results,
|
||||
selected_masks_cfg)
|
||||
|
||||
# Add full attention mask for specific cases
|
||||
full_attention_mask_cfg: Optional[List[int]] = kwargs.get(
|
||||
'full_attention_mask')
|
||||
if full_attention_mask_cfg is not None:
|
||||
selected_masks_cfg.append(full_attention_mask_cfg)
|
||||
|
||||
timesteps_cfg: int = kwargs.get('timesteps', time_step_num)
|
||||
if skip_time_steps_cfg is None:
|
||||
skip_time_steps_cfg = 12
|
||||
# Select best mask strategy using combined results
|
||||
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
|
||||
averaged_results, selected_masks_cfg, skip_time_steps_cfg,
|
||||
timesteps_cfg, head_num)
|
||||
|
||||
# Save mask strategy
|
||||
os.makedirs(save_dir_cfg, exist_ok=True)
|
||||
file_path = os.path.join(save_dir_cfg,
|
||||
f'mask_strategy_s{skip_time_steps_cfg}.json')
|
||||
with open(file_path, 'w') as f:
|
||||
json.dump(mask_strategy, f, indent=4)
|
||||
print(f"Successfully saved mask_strategy to {file_path}")
|
||||
|
||||
# Print sparsity and strategy counts for information
|
||||
print(f"Overall sparsity: {sparsity:.4f}")
|
||||
print("\nStrategy usage counts:")
|
||||
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
|
||||
for strategy, count in strategy_counts.items():
|
||||
print(
|
||||
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
|
||||
)
|
||||
|
||||
# Convert dictionary to 3D list with fixed dimensions
|
||||
mask_strategy_3d = dict_to_3d_list(mask_strategy,
|
||||
t_max=time_step_num,
|
||||
l_max=layer_num,
|
||||
h_max=head_num)
|
||||
|
||||
return mask_strategy_3d
|
||||
|
||||
else: # STA_inference
|
||||
# Get parameters with defaults
|
||||
load_path: Optional[str] = kwargs.get(
|
||||
'load_path', "mask_candidates/mask_strategy.json")
|
||||
if load_path is None:
|
||||
raise ValueError("load_path is required for STA_inference mode")
|
||||
|
||||
# Load previously saved mask strategy
|
||||
with open(load_path) as f:
|
||||
mask_strategy = json.load(f)
|
||||
|
||||
# Convert dictionary to 3D list with fixed dimensions
|
||||
mask_strategy_3d = dict_to_3d_list(mask_strategy,
|
||||
t_max=time_step_num,
|
||||
l_max=layer_num,
|
||||
h_max=head_num)
|
||||
|
||||
return mask_strategy_3d
|
||||
|
||||
|
||||
# Helper functions
|
||||
|
||||
|
||||
def read_specific_json_files(folder_path: str) -> List[Dict[str, Any]]:
|
||||
"""Read and parse JSON files containing mask search results."""
|
||||
json_contents: List[Dict[str, Any]] = []
|
||||
|
||||
# List files only in the current directory (no walk)
|
||||
files = os.listdir(folder_path)
|
||||
# Filter files
|
||||
matching_files = [f for f in files if 'mask' in f and f.endswith('.json')]
|
||||
print(f"Found {len(matching_files)} matching files: {matching_files}")
|
||||
|
||||
for file_name in matching_files:
|
||||
file_path = os.path.join(folder_path, file_name)
|
||||
with open(file_path) as file:
|
||||
data = json.load(file)
|
||||
json_contents.append(data)
|
||||
|
||||
return json_contents
|
||||
|
||||
|
||||
def average_head_losses(
|
||||
results: List[Dict[str, Any]],
|
||||
selected_masks: List[List[int]]) -> Dict[str, Dict[str, np.ndarray]]:
|
||||
"""Average losses across all prompts for each mask strategy."""
|
||||
# Initialize a dictionary to store the averaged results
|
||||
averaged_losses: Dict[str, Dict[str, np.ndarray]] = {}
|
||||
loss_type = 'L2_loss'
|
||||
# Get all loss types (e.g., 'L2_loss')
|
||||
averaged_losses[loss_type] = {}
|
||||
|
||||
for mask in selected_masks:
|
||||
mask_str = str(mask)
|
||||
data_shape = np.array(results[0][loss_type][mask_str]).shape
|
||||
accumulated_data = np.zeros(data_shape)
|
||||
|
||||
# Sum across all prompts
|
||||
for prompt_result in results:
|
||||
accumulated_data += np.array(prompt_result[loss_type][mask_str])
|
||||
|
||||
# Average by dividing by number of prompts
|
||||
averaged_data = accumulated_data / len(results)
|
||||
averaged_losses[loss_type][mask_str] = averaged_data
|
||||
|
||||
return averaged_losses
|
||||
|
||||
|
||||
def select_best_mask_strategy(
|
||||
averaged_results: Dict[str, Dict[str, np.ndarray]],
|
||||
selected_masks: List[List[int]],
|
||||
skip_time_steps: int = 12,
|
||||
timesteps: int = 50,
|
||||
head_num: int = 40
|
||||
) -> Tuple[Dict[str, List[int]], float, Dict[str, int]]:
|
||||
"""Select the best mask strategy for each head based on loss minimization."""
|
||||
best_mask_strategy: Dict[str, List[int]] = {}
|
||||
loss_type = 'L2_loss'
|
||||
# Get the shape of time steps and layers
|
||||
layers = len(averaged_results[loss_type][str(selected_masks[0])][0])
|
||||
|
||||
# Counter for sparsity calculation
|
||||
total_tokens = 0 # total number of masked tokens
|
||||
total_length = 0 # total sequence length
|
||||
|
||||
strategy_counts: Dict[str, int] = {
|
||||
str(strategy): 0
|
||||
for strategy in selected_masks
|
||||
}
|
||||
full_attn_strategy = selected_masks[-1] # Last strategy is full attention
|
||||
print(f"Strategy {full_attn_strategy}, skip first {skip_time_steps} steps ")
|
||||
|
||||
for t in range(timesteps):
|
||||
for layer_idx in range(layers):
|
||||
for h in range(head_num):
|
||||
if t < skip_time_steps: # First steps use full attention
|
||||
strategy = full_attn_strategy
|
||||
else:
|
||||
# Get losses for this head across all strategies
|
||||
head_losses = []
|
||||
for strategy in selected_masks[:
|
||||
-1]: # Exclude full attention
|
||||
head_losses.append(averaged_results[loss_type][str(
|
||||
strategy)][t][layer_idx][h])
|
||||
|
||||
# Find which strategy gives minimum loss
|
||||
best_strategy_idx = np.argmin(head_losses)
|
||||
strategy = selected_masks[best_strategy_idx]
|
||||
|
||||
best_mask_strategy[f'{t}_{layer_idx}_{h}'] = strategy
|
||||
|
||||
# Calculate sparsity
|
||||
nums = strategy # strategy is already a list of numbers
|
||||
total_tokens += nums[0] * nums[1] * nums[
|
||||
2] # masked tokens for chosen strategy
|
||||
total_length += full_attn_strategy[0] * full_attn_strategy[
|
||||
1] * full_attn_strategy[2]
|
||||
|
||||
# Count strategy usage
|
||||
strategy_counts[str(strategy)] += 1
|
||||
|
||||
overall_sparsity = 1 - total_tokens / total_length
|
||||
|
||||
return best_mask_strategy, overall_sparsity, strategy_counts
|
||||
|
||||
|
||||
def dict_to_3d_list(mask_strategy: Optional[Dict[str, List[int]]],
|
||||
t_max: int = 50,
|
||||
l_max: int = 60,
|
||||
h_max: int = 24) -> List[List[List[Optional[List[int]]]]]:
|
||||
result: List[List[List[Optional[List[int]]]]] = [[[
|
||||
None for _ in range(h_max)
|
||||
] for _ in range(l_max)] for _ in range(t_max)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
for key, value in mask_strategy.items():
|
||||
t, layer_idx, h = map(int, key.split('_'))
|
||||
result[t][layer_idx][h] = value
|
||||
return result
|
||||
|
||||
|
||||
def save_mask_search_results(
|
||||
mask_search_final_result: List[Dict[str, List[float]]],
|
||||
prompt: str,
|
||||
mask_strategies: List[str],
|
||||
output_dir: str = 'output/mask_search_result/') -> Optional[str]:
|
||||
if not mask_search_final_result:
|
||||
print("No mask search results to save")
|
||||
return None
|
||||
|
||||
# Create result dictionary with defaultdict for nested lists
|
||||
mask_search_dict: Dict[str, Dict[str, List[List[float]]]] = {
|
||||
"L2_loss": defaultdict(list),
|
||||
"L1_loss": defaultdict(list)
|
||||
}
|
||||
|
||||
mask_selected = list(range(len(mask_strategies)))
|
||||
selected_masks: List[List[int]] = []
|
||||
for index in mask_selected:
|
||||
mask = mask_strategies[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks.append(masks_list)
|
||||
|
||||
# Process each mask strategy
|
||||
for i, mask_strategy in enumerate(selected_masks):
|
||||
mask_strategy_str = str(mask_strategy)
|
||||
# Process L2 loss
|
||||
step_results: List[List[float]] = []
|
||||
for step_data in mask_search_final_result:
|
||||
if isinstance(step_data, dict) and "L2_loss" in step_data:
|
||||
layer_losses = [float(loss) for loss in step_data["L2_loss"]]
|
||||
step_results.append(layer_losses)
|
||||
mask_search_dict["L2_loss"][mask_strategy_str] = step_results
|
||||
|
||||
step_results = []
|
||||
for step_data in mask_search_final_result:
|
||||
if isinstance(step_data, dict) and "L1_loss" in step_data:
|
||||
layer_losses = [float(loss) for loss in step_data["L1_loss"]]
|
||||
step_results.append(layer_losses)
|
||||
mask_search_dict["L1_loss"][mask_strategy_str] = step_results
|
||||
|
||||
# Create the output directory if it doesn't exist
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
# Create a filename based on the first 20 characters of the prompt
|
||||
filename = prompt[:50].replace(" ", "_")
|
||||
filepath = os.path.join(output_dir, f'mask_search_{filename}.json')
|
||||
|
||||
# Save the results to a JSON file
|
||||
with open(filepath, 'w') as f:
|
||||
json.dump(mask_search_dict, f, indent=4)
|
||||
|
||||
print(f"Successfully saved mask research results to {filepath}")
|
||||
|
||||
return filepath
|
||||
@@ -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."""
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
from sageattention import sageattn
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (
|
||||
AttentionBackend) # FlashAttentionMetadata,
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionImpl,
|
||||
AttentionMetadata)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SageAttentionBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> List[int]:
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SAGE_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> Type["SageAttentionImpl"]:
|
||||
return SageAttentionImpl
|
||||
|
||||
# @staticmethod
|
||||
# def get_metadata_cls() -> Type["AttentionMetadata"]:
|
||||
# return FlashAttentionMetadata
|
||||
|
||||
|
||||
class SageAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
self.dropout = extra_impl_args.get("dropout_p", 0.0)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
output = sageattn(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
# since input is (batch_size, seq_len, head_num, head_dim)
|
||||
tensor_layout="NHD",
|
||||
is_causal=self.causal)
|
||||
return output
|
||||
@@ -1,6 +1,6 @@
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Type
|
||||
from typing import Any, Dict, List, Optional, Type
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
@@ -13,6 +13,7 @@ from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
@@ -20,20 +21,41 @@ logger = init_logger(__name__)
|
||||
|
||||
|
||||
# TODO(will-refactor): move this to a utils file
|
||||
def dict_to_3d_list(mask_strategy,
|
||||
t_max=50,
|
||||
l_max=60,
|
||||
h_max=24) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
result = [[[None for _ in range(h_max)] for _ in range(l_max)]
|
||||
for _ in range(t_max)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
def dict_to_3d_list(
|
||||
mask_strategy: Dict[str,
|
||||
Any]) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
|
||||
|
||||
max_timesteps_idx = max(
|
||||
timesteps_idx for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
max_layer_idx = max(layer_idx
|
||||
for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
max_head_idx = max(head_idx
|
||||
for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
|
||||
result = [[[None for _ in range(max_head_idx)]
|
||||
for _ in range(max_layer_idx)] for _ in range(max_timesteps_idx)]
|
||||
|
||||
for key, value in mask_strategy.items():
|
||||
t, layer, h = map(int, key.split('_'))
|
||||
result[t][layer][h] = value
|
||||
timesteps_idx, layer_idx, head_idx = map(int, key.split('_'))
|
||||
result[timesteps_idx][layer_idx][head_idx] = value
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class RangeDict(dict):
|
||||
|
||||
def __getitem__(self, item: int) -> str:
|
||||
for key in self.keys():
|
||||
if isinstance(key, tuple):
|
||||
low, high = key
|
||||
if low <= item <= high:
|
||||
return str(super().__getitem__(key))
|
||||
elif key == item:
|
||||
return str(super().__getitem__(key))
|
||||
raise KeyError(f"seq_len {item} not supported for STA")
|
||||
|
||||
|
||||
class SlidingTileAttentionBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
@@ -63,6 +85,8 @@ class SlidingTileAttentionBackend(AttentionBackend):
|
||||
@dataclass
|
||||
class SlidingTileAttentionMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
STA_param: List[List[
|
||||
Any]] # each timestep with one metadata, shape [num_layers, num_heads]
|
||||
|
||||
|
||||
class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
@@ -79,8 +103,12 @@ class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
forward_batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> SlidingTileAttentionMetadata:
|
||||
|
||||
return SlidingTileAttentionMetadata(current_timestep=current_timestep, )
|
||||
param = forward_batch.STA_param
|
||||
if param is None:
|
||||
return SlidingTileAttentionMetadata(
|
||||
current_timestep=current_timestep, STA_param=[])
|
||||
return SlidingTileAttentionMetadata(current_timestep=current_timestep,
|
||||
STA_param=param[current_timestep])
|
||||
|
||||
|
||||
class SlidingTileAttentionImpl(AttentionImpl):
|
||||
@@ -101,54 +129,75 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
if config_file is None:
|
||||
raise ValueError("FASTVIDEO_ATTENTION_CONFIG is not set")
|
||||
|
||||
# TODO(kevin): get mask strategy for different STA modes
|
||||
with open(config_file) as f:
|
||||
mask_strategy = json.load(f)
|
||||
self.mask_strategy = dict_to_3d_list(mask_strategy)
|
||||
|
||||
mask_strategy = dict_to_3d_list(mask_strategy)
|
||||
self.prefix = prefix
|
||||
self.mask_strategy = mask_strategy
|
||||
sp_group = get_sp_group()
|
||||
self.sp_size = sp_group.world_size
|
||||
# STA config
|
||||
self.STA_base_tile_size = [6, 8, 8]
|
||||
self.img_latent_shape_mapping = RangeDict({
|
||||
(115200, 115456): '30x48x80',
|
||||
82944: '36x48x48',
|
||||
69120: '18x48x80',
|
||||
})
|
||||
self.full_window_mapping = {
|
||||
'30x48x80': [5, 6, 10],
|
||||
'36x48x48': [6, 6, 6],
|
||||
'18x48x80': [3, 6, 10]
|
||||
}
|
||||
|
||||
def tile(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = rearrange(x,
|
||||
"b (sp t h w) head d -> b (t sp h w) head d",
|
||||
sp=self.sp_size,
|
||||
t=30 // self.sp_size,
|
||||
h=48,
|
||||
w=80)
|
||||
t=self.img_latent_shape_int[0] // self.sp_size,
|
||||
h=self.img_latent_shape_int[1],
|
||||
w=self.img_latent_shape_int[2])
|
||||
return rearrange(
|
||||
x,
|
||||
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
|
||||
n_t=5,
|
||||
n_h=6,
|
||||
n_w=10,
|
||||
ts_t=6,
|
||||
ts_h=8,
|
||||
ts_w=8)
|
||||
n_t=self.full_window_size[0],
|
||||
n_h=self.full_window_size[1],
|
||||
n_w=self.full_window_size[2],
|
||||
ts_t=self.STA_base_tile_size[0],
|
||||
ts_h=self.STA_base_tile_size[1],
|
||||
ts_w=self.STA_base_tile_size[2])
|
||||
|
||||
def untile(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = rearrange(
|
||||
x,
|
||||
"b (n_t n_h n_w ts_t ts_h ts_w) h d -> b (n_t ts_t n_h ts_h n_w ts_w) h d",
|
||||
n_t=5,
|
||||
n_h=6,
|
||||
n_w=10,
|
||||
ts_t=6,
|
||||
ts_h=8,
|
||||
ts_w=8)
|
||||
n_t=self.full_window_size[0],
|
||||
n_h=self.full_window_size[1],
|
||||
n_w=self.full_window_size[2],
|
||||
ts_t=self.STA_base_tile_size[0],
|
||||
ts_h=self.STA_base_tile_size[1],
|
||||
ts_w=self.STA_base_tile_size[2])
|
||||
return rearrange(x,
|
||||
"b (t sp h w) head d -> b (sp t h w) head d",
|
||||
sp=self.sp_size,
|
||||
t=30 // self.sp_size,
|
||||
h=48,
|
||||
w=80)
|
||||
t=self.img_latent_shape_int[0] // self.sp_size,
|
||||
h=self.img_latent_shape_int[1],
|
||||
w=self.img_latent_shape_int[2])
|
||||
|
||||
def preprocess_qkv(
|
||||
self,
|
||||
qkv: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
img_sequence_length = qkv.shape[1]
|
||||
self.img_latent_shape_str = self.img_latent_shape_mapping[
|
||||
img_sequence_length]
|
||||
self.full_window_size = self.full_window_mapping[
|
||||
self.img_latent_shape_str]
|
||||
self.img_latent_shape_int = list(
|
||||
map(int, self.img_latent_shape_str.split('x')))
|
||||
self.img_seq_length = self.img_latent_shape_int[
|
||||
0] * self.img_latent_shape_int[1] * self.img_latent_shape_int[2]
|
||||
return self.tile(qkv)
|
||||
|
||||
def postprocess_output(
|
||||
@@ -165,29 +214,92 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
v: torch.Tensor,
|
||||
attn_metadata: SlidingTileAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
|
||||
assert self.mask_strategy is not None, "mask_strategy cannot be None for SlidingTileAttention"
|
||||
assert self.mask_strategy[
|
||||
0] is not None, "mask_strategy[0] cannot be None for SlidingTileAttention"
|
||||
if self.mask_strategy is None:
|
||||
raise ValueError(
|
||||
"mask_strategy cannot be None for SlidingTileAttention")
|
||||
if self.mask_strategy[0] is None:
|
||||
raise ValueError(
|
||||
"mask_strategy[0] cannot be None for SlidingTileAttention")
|
||||
|
||||
timestep = attn_metadata.current_timestep
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
forward_batch = forward_context.forward_batch
|
||||
if forward_batch is None:
|
||||
raise ValueError("forward_batch cannot be None")
|
||||
# pattern:'.double_blocks.0.attn.impl' or '.single_blocks.0.attn.impl'
|
||||
layer_idx = int(self.prefix.split('.')[-3])
|
||||
# TODO: remove hardcode
|
||||
text_length = q.shape[1] - (30 * 48 * 80)
|
||||
query = q.transpose(1, 2)
|
||||
key = k.transpose(1, 2)
|
||||
value = v.transpose(1, 2)
|
||||
if attn_metadata.STA_param is None or len(
|
||||
attn_metadata.STA_param) <= layer_idx:
|
||||
raise ValueError("Invalid STA_param")
|
||||
STA_param = attn_metadata.STA_param[layer_idx]
|
||||
|
||||
text_length = q.shape[1] - self.img_seq_length
|
||||
has_text = text_length > 0
|
||||
|
||||
query = q.transpose(1, 2).contiguous()
|
||||
key = k.transpose(1, 2).contiguous()
|
||||
value = v.transpose(1, 2).contiguous()
|
||||
|
||||
head_num = query.size(1)
|
||||
sp_group = get_sp_group()
|
||||
current_rank = sp_group.rank_in_group
|
||||
start_head = current_rank * head_num
|
||||
windows = [
|
||||
self.mask_strategy[timestep][layer_idx][head_idx + start_head]
|
||||
for head_idx in range(head_num)
|
||||
]
|
||||
hidden_states = sliding_tile_attention(query, key, value, windows,
|
||||
text_length).transpose(1, 2)
|
||||
|
||||
# searching or tuning mode
|
||||
if len(STA_param) < head_num * sp_group.world_size:
|
||||
sparse_attn_hidden_states_all = []
|
||||
full_mask_window = STA_param[-1]
|
||||
for window_size in STA_param[:-1]:
|
||||
sparse_hidden_states = sliding_tile_attention(
|
||||
query, key, value, [window_size] * head_num, text_length,
|
||||
has_text, self.img_latent_shape_str).transpose(1, 2)
|
||||
sparse_attn_hidden_states_all.append(sparse_hidden_states)
|
||||
|
||||
hidden_states = sliding_tile_attention(
|
||||
query, key, value, [full_mask_window] * head_num, text_length,
|
||||
has_text, self.img_latent_shape_str).transpose(1, 2)
|
||||
|
||||
attn_L2_loss = []
|
||||
attn_L1_loss = []
|
||||
# average loss across all heads
|
||||
for sparse_attn_hidden_states in sparse_attn_hidden_states_all:
|
||||
# L2 loss
|
||||
attn_L2_loss_ = torch.mean((sparse_attn_hidden_states.float() -
|
||||
hidden_states.float())**2,
|
||||
dim=[0, 1, 3]).cpu().numpy()
|
||||
attn_L2_loss_ = [round(float(x), 6) for x in attn_L2_loss_]
|
||||
attn_L2_loss.append(attn_L2_loss_)
|
||||
# L1 loss
|
||||
attn_L1_loss_ = torch.mean(
|
||||
torch.abs(sparse_attn_hidden_states.float() -
|
||||
hidden_states.float()),
|
||||
dim=[0, 1, 3]).cpu().numpy()
|
||||
attn_L1_loss_ = [round(float(x), 6) for x in attn_L1_loss_]
|
||||
attn_L1_loss.append(attn_L1_loss_)
|
||||
|
||||
layer_loss_save = {"L2_loss": attn_L2_loss, "L1_loss": attn_L1_loss}
|
||||
|
||||
if forward_batch.is_cfg_negative:
|
||||
if forward_batch.mask_search_final_result_neg is not None:
|
||||
forward_batch.mask_search_final_result_neg[timestep].append(
|
||||
layer_loss_save)
|
||||
else:
|
||||
if forward_batch.mask_search_final_result_pos is not None:
|
||||
forward_batch.mask_search_final_result_pos[timestep].append(
|
||||
layer_loss_save)
|
||||
else:
|
||||
# windows = [
|
||||
# self.mask_strategy[timestep][layer_idx][head_idx + start_head]
|
||||
# for head_idx in range(head_num)
|
||||
# ]
|
||||
windows = [
|
||||
STA_param[head_idx + start_head] for head_idx in range(head_num)
|
||||
]
|
||||
# if has_text is False:
|
||||
# from IPython import embed
|
||||
# embed()
|
||||
hidden_states = sliding_tile_attention(
|
||||
query, key, value, windows, text_length, has_text,
|
||||
self.img_latent_shape_str).transpose(1, 2)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import List, Optional
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -13,6 +13,7 @@ from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.utils import get_compute_dtype
|
||||
|
||||
|
||||
class DistributedAttention(nn.Module):
|
||||
@@ -25,7 +26,8 @@ class DistributedAttention(nn.Module):
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[List[_Backend]] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
@@ -37,7 +39,7 @@ class DistributedAttention(nn.Module):
|
||||
if num_kv_heads is None:
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = torch.get_default_dtype()
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
@@ -83,9 +85,6 @@ class DistributedAttention(nn.Module):
|
||||
# Check input shapes
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim(
|
||||
) == 4, "Expected 4D tensors"
|
||||
# assert bs = 1
|
||||
assert q.shape[
|
||||
0] == 1, "Batch size must be 1, and there should be no padding tokens"
|
||||
batch_size, seq_len, num_heads, head_dim = q.shape
|
||||
local_rank = get_sequence_model_parallel_rank()
|
||||
world_size = get_sequence_model_parallel_world_size()
|
||||
@@ -146,7 +145,8 @@ class LocalAttention(nn.Module):
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[List[_Backend]] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
if softmax_scale is None:
|
||||
@@ -156,7 +156,7 @@ class LocalAttention(nn.Module):
|
||||
if num_kv_heads is None:
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = torch.get_default_dtype()
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
|
||||
@@ -3,7 +3,8 @@
|
||||
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from typing import Generator, List, Optional, Type, cast
|
||||
from functools import cache
|
||||
from typing import Generator, Optional, Tuple, Type, cast
|
||||
|
||||
import torch
|
||||
|
||||
@@ -81,7 +82,17 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
|
||||
def get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: Optional[List[_Backend]] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
return _cached_get_attn_backend(head_size, dtype,
|
||||
supported_attention_backends)
|
||||
|
||||
|
||||
@cache
|
||||
def _cached_get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
|
||||
@@ -1,8 +0,0 @@
|
||||
from fastvideo.v1.configs.hunyuan import HunyuanConfig, FastHunyuanConfig
|
||||
from fastvideo.v1.configs.wan import WanT2V480PConfig, WanI2V480PConfig
|
||||
from fastvideo.v1.configs.base import BaseConfig, SlidingTileAttnConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig"
|
||||
]
|
||||
|
||||
@@ -1,70 +0,0 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Dict, Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseConfig:
|
||||
"""Base configuration for all pipeline architectures."""
|
||||
|
||||
# Video parameters
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 125
|
||||
fps: int = 24
|
||||
|
||||
# Video generation parameters
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
seed: int = 1024
|
||||
guidance_rescale: float = 0.0
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
use_cpu_offload: bool = False
|
||||
disable_autocast: bool = False
|
||||
|
||||
# Model configuration
|
||||
precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = True
|
||||
vae_scale_factor: Optional[int] = None
|
||||
|
||||
# DiT configuration
|
||||
num_channels_latents: Optional[int] = None
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_len: int = -1
|
||||
hidden_state_skip_layer: int = 0
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
|
||||
# Additional parameters can be added as a dict
|
||||
extra_params: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlidingTileAttnConfig(BaseConfig):
|
||||
"""Configuration for sliding tile attention."""
|
||||
|
||||
# Override any BaseConfig defaults as needed
|
||||
# Add sliding tile specific parameters
|
||||
window_size: int = 16
|
||||
stride: int = 8
|
||||
|
||||
# You can provide custom defaults for inherited fields
|
||||
height: int = 576
|
||||
width: int = 1024
|
||||
|
||||
# Additional configuration specific to sliding tile attention
|
||||
pad_to_square: bool = False
|
||||
use_overlap_optimization: bool = True
|
||||
@@ -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
|
||||
}
|
||||
@@ -1,39 +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
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user