Compare commits

..
29 Commits
Author SHA1 Message Date
JerryZhou54 38ee9dc3b4 Complete Model Config design for VAEs 2025-04-23 19:08:11 +00:00
JerryZhou54 77b013fb8a Add model config for WanVAE 2025-04-23 19:01:13 +00:00
JerryZhou54 c31efe1234 Add model config for VAE 2025-04-23 18:59:25 +00:00
JerryZhou54 c056b89aea Add preliminary design for model config 2025-04-23 18:57:13 +00:00
Kevin Lin eac79b753f [V1] Worker improvements/cleanup (#361) 2025-04-22 00:32:43 -07:00
William Lin 4d58cf20d0 chore: Release FastVideo 0.0.2 and update python requirements (#360) 2025-04-21 14:10:48 -07:00
Kevin LinandWill Lin 52c93ecc9d [V1] Gradio demo with new API (#357)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-04-19 18:14:24 -07:00
William Lin 42d63166ac [V1] Process aware logging; improve logging msg (#356) 2025-04-19 15:03:29 -07:00
William Lin 6db20345a2 [V1] Worker cleanup; Logging clean up; enables isort again (#355) 2025-04-18 19:26:47 -07:00
William Lin ad27ea596c [sta] release 0.0.4 (#354) 2025-04-18 14:54:40 -07:00
William Lin 9aadb4bf8c [1/n] [v1] Add Worker abstractions for User API (#336) 2025-04-18 14:38:46 -07:00
Kevin Lin bd941df271 [Docs] Fix developer guide images (#353) 2025-04-17 22:32:19 -07:00
Yongqi Chen 8a73876d3b add STA to Wan v1 (#349) 2025-04-17 16:35:19 -07:00
Kevin Lin 1483a1138a [CLI] Fix duplicate --num-gpus (#352) 2025-04-17 13:01:48 -07:00
Wei Zhou 5e243d8292 Default to using original WanVAE's encoding/decoding algorithm (#351) 2025-04-17 13:00:25 -07:00
Kevin Lin b0c66d3200 [CI] Docker image improvements (#350) 2025-04-17 12:27:10 -07:00
Wei ZhouandWill Lin c86da2c736 [core] Pipeline config (#343)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-04-15 15:43:10 -07:00
Kevin Lin bae2a19dcf [CI] Add manual trigger to sta-publish and fastvideo-publish (#346) 2025-04-15 15:39:11 -07:00
Kevin Lin 057686f59d [CI] Free up runner disk for sta-publish (#345) 2025-04-15 15:27:09 -07:00
William Lin 67da56628b [STA] Sta release 0.0.3 (#344) 2025-04-15 13:06:03 -07:00
Zhang PeiyuanandSolitaryThinker 2325adffa2 Add STA to V1 (#312)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-04-15 02:39:58 -07:00
Kevin Lin 20cf836ef1 [CI] Support custom Docker image (#342) 2025-04-14 20:08:53 -07:00
William Lin 2bf69b6f92 [Docs] Initial examples setup and more docs (#332) 2025-04-14 16:53:57 -07:00
William Lin 13583f5ffb [Model] Remove RMSNorm's forward_native hardcode from Wan (#339) 2025-04-14 16:46:43 -07:00
William LinandJerryZhou54 008ee2099a V1 wan rebased (#335)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
2025-04-11 15:34:54 -07:00
Kevin Lin 137f61f2fe Port tests to v1 (#333) 2025-04-11 01:40:01 -07:00
William Lin ccb262974e [Docs] Add dev guide and doc building CI (#330) 2025-04-09 13:09:00 -07:00
Kevin Lin 30966e3bc9 [CI] Set allowedCudaVersions (#329) 2025-04-09 10:16:05 -07:00
William Lin 7b4272d6b7 [Docs] Fix doc lint (#325) 2025-04-09 10:14:53 -07:00
151 changed files with 407487 additions and 2086 deletions
+23 -12
View File
@@ -28,7 +28,7 @@ def parse_arguments():
parser.add_argument(
'--image',
type=str,
default='runpod/pytorch:2.4.0-py3.11-cuda12.4.1-devel-ubuntu22.04',
required=True,
help='Docker image to use')
return parser.parse_args()
@@ -46,6 +46,16 @@ HEADERS = {
def create_pod():
"""Create a RunPod instance"""
# Ensure image name is lowercase (Docker requirement)
image_name = args.image.lower()
print(f"Using specified image: {image_name}")
docker_start_cmd = [
"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"
]
print(f"Creating RunPod instance with GPU: {args.gpu_type}...")
payload = {
"name": f"fastvideo-{JOB_ID}-{RUN_ID}",
@@ -53,7 +63,9 @@ def create_pod():
"volumeInGb": args.volume_size,
"gpuTypeIds": [args.gpu_type],
"gpuCount": args.gpu_count,
"imageName": args.image
"imageName": image_name,
"allowedCudaVersions": ["12.4"],
"dockerStartCmd": docker_start_cmd
}
response = requests.post(PODS_API, headers=HEADERS, json=payload)
@@ -90,7 +102,7 @@ def wait_for_pod(pod_id):
"Timed out waiting for RunPod to reach RUNNING state")
# Wait for ports to be assigned
max_attempts = 6
max_attempts = 50
attempts = 0
while attempts < max_attempts:
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
@@ -107,7 +119,7 @@ def wait_for_pod(pod_id):
print(
f"Waiting for SSH port and public IP to be available... (attempt {attempts+1}/{max_attempts})"
)
time.sleep(10)
time.sleep(20)
attempts += 1
if attempts >= max_attempts:
@@ -144,16 +156,15 @@ def execute_command(pod_id):
]
subprocess.run(scp_command, check=True)
# For custom image, we can use the pre-configured environment
setup_steps = [
"cd /workspace",
"wget -q https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh",
"bash Miniconda3-latest-Linux-x86_64.sh -b -p $HOME/miniconda3",
"source $HOME/miniconda3/bin/activate",
"conda create --name venv python=3.10.0 -y", "conda activate venv",
"mkdir -p /workspace/repo",
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
f"cd /workspace/{repo_name}", args.test_command
f"cd /workspace/{repo_name}",
"source /opt/conda/etc/profile.d/conda.sh",
"conda activate fastvideo-dev",
args.test_command
]
remote_command = " && ".join(setup_steps)
ssh_command = [
@@ -170,7 +181,7 @@ def execute_command(pod_id):
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
universal_newlines=True,
bufsize=1)
bufsize=0)
stdout_lines = []
+78
View File
@@ -0,0 +1,78 @@
name: Build and Push Docker Image
on:
workflow_dispatch: # Only manual triggers
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."
+83
View File
@@ -0,0 +1,83 @@
# Sample workflow for building and deploying a Hugo site to GitHub Pages
name: Deploy FastVideo Docs to Pages
on:
# Runs on pushes targeting the default branch
push:
branches:
- 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:
# Sets permissions of the GITHUB_TOKEN to allow deployment to GitHub Pages
permissions:
contents: read
pages: write
id-token: write
# Allow only one concurrent deployment, skipping runs queued between the run in-progress and latest queued.
# However, do NOT cancel in-progress runs as we want to allow these production deployments to complete.
concurrency:
group: "pages"
cancel-in-progress: false
# Default to bash
defaults:
run:
shell: bash
jobs:
pre-commit:
uses: ./.github/workflows/pre-commit.yml
# Build job
build:
runs-on: ubuntu-latest
needs: pre-commit
steps:
- name: Checkout
uses: actions/checkout@v4
- name: Setup Pages
id: pages
uses: actions/configure-pages@v5
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Install dependencies
run: |
cd docs
pip install -r requirements-docs.txt
- name: Build docs
run: |
cd docs
make clean
make html
- name: Upload artifact
uses: actions/upload-pages-artifact@v3
with:
path: ./docs/build/html
# Deployment job
deploy:
environment:
name: github-pages
url: ${{ steps.deployment.outputs.page_url }}
if: ${{ github.event_name == 'push' }}
runs-on: ubuntu-latest
needs: build
steps:
- name: Deploy to GitHub Pages
id: deployment
uses: actions/deploy-pages@v4
+2 -1
View File
@@ -6,6 +6,7 @@ on:
- main
paths:
- 'pyproject.toml' # Trigger when pyproject.toml changes
workflow_dispatch:
jobs:
check-version-change:
@@ -41,7 +42,7 @@ jobs:
build-publish-main:
needs: check-version-change
if: needs.check-version-change.outputs.version-changed == 'true'
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ubuntu-latest
permissions:
id-token: write # Needed for OIDC Trusted Publishing
+133 -11
View File
@@ -14,17 +14,35 @@ on:
- ".github/workflows/pr-test.yml"
workflow_dispatch:
inputs:
custom_image:
description: "Custom image from this repository (default: fastvideo-dev:latest)"
required: false
default: "fastvideo-dev:latest"
type: string
run_encoder_test:
description: "Run encoder-test"
required: false
default: false
type: boolean
run_vae_test:
description: "Run vae-test"
required: false
default: false
type: boolean
run_transformer_test:
description: "Run transformer-test"
required: false
default: false
type: boolean
run_ssim_test:
description: "Run ssim-test"
required: false
default: false
type: boolean
env:
PYTHONUNBUFFERED: "1"
concurrency:
group: pr-test-${{ github.ref }}
cancel-in-progress: true
@@ -39,6 +57,8 @@ jobs:
if: ${{ github.event.pull_request.draft == false || github.event_name == 'workflow_dispatch' }}
outputs:
encoder-test: ${{ steps.filter.outputs.encoder-test }}
vae-test: ${{ steps.filter.outputs.vae-test }}
transformer-test: ${{ steps.filter.outputs.transformer-test }}
steps:
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
@@ -49,6 +69,14 @@ jobs:
- 'fastvideo/v1/models/encoders/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/encoders/**'
vae-test:
- 'fastvideo/v1/models/vaes/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/vaes/**'
transformer-test:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/transformers/**'
encoder-test:
needs: change-filter
@@ -87,9 +115,8 @@ jobs:
--gpu-type "NVIDIA A40"
--gpu-count 1
--volume-size 100
--test-command "pip install -e .[test] &&
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
pytest ./fastvideo/v1/tests/encoders -s"
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
@@ -99,6 +126,102 @@ jobs:
JOB_ID: "encoder-test"
run: python .github/scripts/runpod_cleanup.py
vae-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.vae-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_vae_test == 'true')
runs-on: ubuntu-latest
environment: runpod-runners
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Set up SSH key
run: |
mkdir -p ~/.ssh
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
- name: Install dependencies
run: pip install requests
- name: Run tests on RunPod
env:
JOB_ID: "vae-test"
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: 30
run: >-
python .github/scripts/runpod_api.py
--gpu-type "NVIDIA A40"
--gpu-count 1
--volume-size 100
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
env:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
JOB_ID: "vae-test"
run: python .github/scripts/runpod_cleanup.py
transformer-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.transformer-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_transformer_test == 'true')
runs-on: ubuntu-latest
environment: runpod-runners
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Set up SSH key
run: |
mkdir -p ~/.ssh
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
- name: Install dependencies
run: pip install requests
- name: Run tests on RunPod
env:
JOB_ID: "transformer-test"
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: 30
run: >-
python .github/scripts/runpod_api.py
--gpu-type "NVIDIA L40S"
--gpu-count 1
--volume-size 100
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
env:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
JOB_ID: "transformer-test"
run: python .github/scripts/runpod_cleanup.py
ssim-test:
needs: change-filter
if: >-
@@ -130,16 +253,15 @@ jobs:
JOB_ID: "ssim-test"
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: 30
timeout-minutes: 45
run: >-
python .github/scripts/runpod_api.py
--gpu-type "NVIDIA A40"
--gpu-count 2
--disk-size 100
--volume-size 100
--test-command "pip install -e .[test] &&
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
pytest ./fastvideo/v1/tests/ssim -vs"
--disk-size 200
--volume-size 200
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
- name: Terminate RunPod Instances
if: ${{ always() }}
@@ -150,7 +272,7 @@ jobs:
run: python .github/scripts/runpod_cleanup.py
runpod-cleanup:
needs: [encoder-test, ssim-test] # Add other jobs to this list as you create them
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
runs-on: ubuntu-latest
steps:
@@ -167,7 +289,7 @@ jobs:
- name: Cleanup all RunPod instances
env:
JOB_IDS: '["encoder-test", "ssim-test"]' # JSON array of job IDs
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test"]' # JSON array of job IDs
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
run: python .github/scripts/runpod_cleanup.py
+26 -2
View File
@@ -6,6 +6,7 @@ on:
- main
paths:
- "csrc/sliding_tile_attention/setup.py"
workflow_dispatch:
jobs:
check-version-change:
@@ -43,7 +44,7 @@ jobs:
build_wheels:
name: Build Wheel
needs: check-version-change
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ${{ matrix.os }}
strategy:
@@ -57,6 +58,29 @@ jobs:
cuda-version: ['12.4.1', '12.5.1', '12.6.3']
steps:
- name: Free up disk space
run: |
echo "Initial disk space:"
df -h
# Remove large directories
sudo rm -rf /usr/share/dotnet
sudo rm -rf /usr/local/lib/android
sudo rm -rf /opt/ghc
sudo rm -rf /usr/local/share/boost
sudo rm -rf /usr/share/swift
sudo rm -rf /usr/local/lib/node_modules
sudo rm -rf /usr/local/share/powershell
sudo rm -rf /usr/share/rust
sudo rm -rf /usr/local/.ghcup
# Remove cached files
sudo rm -rf /var/lib/apt/lists/*
sudo rm -rf /var/cache/apt/archives/*
echo "Disk space after cleanup:"
df -h
- name: Checkout
uses: actions/checkout@v4
@@ -145,7 +169,7 @@ jobs:
publish_package:
name: Publish package
needs: [build_wheels, check-version-change]
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ubuntu-22.04
permissions:
id-token: write # Needed for OIDC Trusted Publishing
+32 -7
View File
@@ -3,14 +3,11 @@ __pycache__
*.pth
UCF-101/
results/
build/
fastvideo.egg-info/
wandb/
*.ipynb
*.jpg
*.safetensors
*.mp4
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
*.png
*.gif
*.pth
@@ -26,11 +23,39 @@ outputs_video
sbatch.sh
*.out
env
dist/
*.o
**/build/
**.egg-info
**.pyc
**.egg
**.txt
**.json
**.json
# Distribution / packaging
build/
dist/
*.egg-info/
*.egg
eggs/
.eggs/
# Sphinx documentation
docs/_build/
docs/source/getting_started/examples/
# VSCode
.vscode/
# DS Store
.DS_Store
# vim swap files
*.swo
*.swp
# Python pickle files
*.pkl
# Reference videos
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
# Static images
!docs/source/_static/images/**/*.png
+3 -2
View File
@@ -19,6 +19,7 @@ exclude: |
fastvideo/sample/.*|
fastvideo/train\.py|
fastvideo/utils/.*|
fastvideo/v1/examples/.*|
.github/workflows/fastvideo-publish.yml|
.github/workflows/sta-publish.yml
)
@@ -41,7 +42,7 @@ repos:
additional_dependencies: ['tomli']
args: ['--toml', 'pyproject.toml']
- repo: https://github.com/PyCQA/isort
rev: 0a0b7a830386ba6a31c2ec8316849ae4d1b8240d # 6.0.0
rev: 6.0.1
hooks:
- id: isort
- repo: https://github.com/jackdewinter/pymarkdown
@@ -66,7 +67,7 @@ repos:
entry: bash
args:
- -c
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
language: system
always_run: true
pass_filenames: false
+48
View File
@@ -0,0 +1,48 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
ENV DEBIAN_FRONTEND=noninteractive
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
git \
ca-certificates \
openssh-server \
&& rm -rf /var/lib/apt/lists/*
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
rm Miniconda3-latest-Linux-x86_64.sh
ENV PATH=/opt/conda/bin:$PATH
RUN conda create --name fastvideo-dev python=3.10.0 -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.0.post2 --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
+8 -2
View File
@@ -5,7 +5,7 @@
FastVideo is a lightweight framework for accelerating large video diffusion models.
<p align="center">
🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank">FastHunyuan</a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank">FastMochi</a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> Slack </a>
| <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> |
</p>
https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1
@@ -29,7 +29,7 @@ Dev in progress and highly experimental.
- ```2024/12/17```: `FastVideo` v1.0 is released.
## 🔧 Installation from source
The code is tested on Python 3.10.0, CUDA 12.4 and H100.
The code is tested on Python 3.10-3.12, CUDA 12.4 and H100.
```
# Clone FastVideo
@@ -44,6 +44,12 @@ pip install flash-attn==2.7.0.post2
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.
You can also install the Sliding Tile Attention package using
```
pip install st_attn==0.0.4
```
## 🚀 Inference
### Inference StepVideo with Sliding Tile Attention
First, download the model:
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -9,7 +9,7 @@ target = target.lower()
# Package metadata
PACKAGE_NAME = "st_attn"
VERSION = "0.0.2"
VERSION = "0.0.4"
AUTHOR = "Hao AI Lab"
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
+1 -1
View File
@@ -10,7 +10,7 @@
#ifdef TK_COMPILE_ATTN
extern torch::Tensor sta_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
);
#endif
@@ -4,8 +4,13 @@ import torch
from st_attn_cuda import sta_fwd
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True):
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, img_latent_shape='30*48*80'):
seq_length = q_all.shape[2]
img_latent_shape_mapping = {
'30x48x80':1,
'36x48x48':2,
'18x48x80':3,
}
if has_text:
assert q_all.shape[
2] >= 115200, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
@@ -17,8 +22,14 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
else:
assert q_all.shape[2] == 82944
if img_latent_shape == '36x48x48': # Stepvideo 204x768x68
assert q_all.shape[2] == 82944
elif img_latent_shape == '18x48x80': # Wan 69x768x1280
assert q_all.shape[2] == 69120
else:
raise ValueError(f"Unsupported {img_latent_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
kernel_aspect_ratio_flag = img_latent_shape_mapping[img_latent_shape]
hidden_states = torch.empty_like(q_all)
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
@@ -29,7 +40,7 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
head_index:head_index + 1],
hidden_states[batch:batch + 1, head_index:head_index + 1])
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text)
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
if has_text:
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True)
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
return hidden_states[:, :, :seq_length]
@@ -359,7 +359,7 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
#include <iostream>
torch::Tensor
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text)
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
@@ -446,7 +446,7 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
auto threads = NUM_WORKERS * kittens::WARP_THREADS;
if (has_text) {
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4-2), qo_heads, batch);
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
dim3 grid_text(2, qo_heads, batch);
if (!process_text) {
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
@@ -558,123 +558,267 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
} else {
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
if (kernel_aspect_ratio_flag == 2){
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
}else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
}else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==6 && kernel_h_size == 3 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==6 && kernel_h_size == 3 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
}
else if (kernel_aspect_ratio_flag == 3) {
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 0, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 0, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 1, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,0, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,0, 3, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
}
else {
std::cout << "Unsupported kernel_aspect_ratio_flag: " << kernel_aspect_ratio_flag << std::endl;
}
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 18 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 27 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 40 KiB

+4 -3
View File
@@ -17,6 +17,7 @@ import inspect
import logging
import os
import sys
from typing import Optional
import requests
from sphinx.ext import autodoc
@@ -136,7 +137,8 @@ _cached_base: str = ""
_cached_branch: str = ""
def get_repo_base_and_branch(pr_number):
def get_repo_base_and_branch(
pr_number: str) -> tuple[Optional[str], Optional[str]]:
global _cached_base, _cached_branch
if _cached_base and _cached_branch:
return _cached_base, _cached_branch
@@ -158,7 +160,6 @@ def linkcode_resolve(domain, info):
return None
if not info['module']:
return None
filename = info['module'].replace('.', '/')
module = info['module']
# try to determine the correct file and line number to link to
@@ -173,7 +174,7 @@ def linkcode_resolve(domain, info):
if not (inspect.isclass(obj) or inspect.isfunction(obj)
or inspect.ismethod(obj)):
obj = obj.__class__ # Get the class of the instance
obj = obj.__class__ # type: ignore[assignment]
lineno = inspect.getsourcelines(obj)[1]
filename = (inspect.getsourcefile(obj)
+138
View File
@@ -0,0 +1,138 @@
(developer-guide)=
# 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!
Our community is open to everyone and welcomes any contributions no matter how large or small.
# Developer Environment:
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support Linux and CUDA GPUs, but we hope to support other platforms in the future.
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
Install Miniconda:
```
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
Create and activate a Conda environment for FastVideo:
```
conda create -n fastvideo python=3.10 -y
conda activate fastvideo
```
Clone the FastVideo repository and go to the FastVideo directory:
```
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
```
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
```bash
pip install -e .[dev]
# Can also install flash-attn (optional)
pip install flash-attn==2.7.0.post2 --no-build-isolation
# Linting, formatting and static type checking
pre-commit install --hook-type pre-commit --hook-type commit-msg
# You can manually run pre-commit with
pre-commit run --all-files
# Unit tests
pytest tests/
```
---
## 🐳 Using the FastVideo Docker Image
If you prefer a containerized development environment or want to avoid managing dependencies manually, you can use our prebuilt Docker image:
**Image:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
### Starting the container
```bash
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
```
This will:
- Start the container with GPU access
- Drop you into a shell with the `fastvideo-dev` Conda environment preconfigured
### Using the container
```bash
# Conda environment should already be active
# FastVideo package installed in editable mode
# Pull the latest changes from remote
cd /FastVideo
git pull
# Run linters and tests
pre-commit run --all-files
pytest tests/
```
---
## 📦 Developing FastVideo on RunPod
You can easily use the FastVideo Docker image as a custom container on [RunPod](https://www.runpod.io) for development or experimentation.
### Creating a new pod
Choose a GPU that supports CUDA 12.4
![RunPod CUDA selection](../_static/images/runpod_cuda.png)
When creating your pod template, use this image:
```
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
```
Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.io/pods/configuration/use-ssh)):
```bash
bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
```
![RunPod template configuration](../_static/images/runpod_template.png)
After deploying, the pod will take a few minutes to pull the image and start the SSH service.
![RunPod ssh](../_static/images/runpod_ssh.png)
### Working with the pod
After SSH'ing into your pod, you'll find the `fastvideo-dev` Conda environment already activated.
To pull in the latest changes from the GitHub repo:
```bash
cd /FastVideo
git pull
```
`If you have a persistent volume and want to keep your code changes, you can move /FastVideo to /workspace/FastVideo, or simply clone the repository there.`
Run your development workflows as usual:
```bash
# Run linters
pre-commit run --all-files
# Run tests
pytest tests/
```
+33 -26
View File
@@ -1,13 +1,15 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
import itertools
import re
from dataclasses import dataclass, field
from pathlib import Path
from typing import Optional
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
ROOT_DIR_RELATIVE = '../../../..'
EXAMPLE_DIR = ROOT_DIR / "examples"
EXAMPLE_DIR = ROOT_DIR / "fastvideo/v1/examples"
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/source/getting_started/examples"
@@ -30,7 +32,8 @@ def fix_case(text: str) -> str:
r"int\d+": lambda x: x.group(0).upper(), # e.g. int8, int16
}
for pattern, repl in subs.items():
text = re.sub(rf'\b{pattern}\b', repl, text, flags=re.IGNORECASE)
text = re.sub(rf'\b{pattern}\b', repl, text,
flags=re.IGNORECASE) # type: ignore[call-overload]
return text
@@ -86,7 +89,7 @@ class Example:
generate() -> str: Generates the documentation content.
""" # noqa: E501
path: Path
category: str = None
category: Optional[str] = None
main_file: Path = field(init=False)
other_files: list[Path] = field(init=False)
title: str = field(init=False)
@@ -124,7 +127,8 @@ class Example:
if self.path.is_file():
return []
is_other_file = lambda file: file.is_file() and file != self.main_file
return [file for file in self.path.rglob("*") if is_other_file(file)]
return [file for file in self.path.rglob("*")
if is_other_file(file)] # type: ignore[no-untyped-call]
def determine_title(self) -> str:
return fix_case(self.path.stem.replace("_", " ").title())
@@ -139,7 +143,7 @@ class Example:
"literalinclude"
if include == "literalinclude":
content += f"# {self.title}\n\n"
content += f":::{{{include}}} {make_relative(self.main_file)}\n"
content += f":::{{{include}}} {make_relative(self.main_file)}\n" # type: ignore[no-untyped-call]
if include == "literalinclude":
content += f":language: {self.main_file.suffix[1:]}\n"
content += ":::\n\n"
@@ -152,7 +156,7 @@ class Example:
include = "include" if file.suffix == ".md" else "literalinclude"
content += f":::{{admonition}} {file.relative_to(self.path)}\n"
content += ":class: dropdown\n\n"
content += f":::{{{include}}} {make_relative(file)}\n:::\n"
content += f":::{{{include}}} {make_relative(file)}\n:::\n" # type: ignore[no-untyped-call]
content += ":::\n\n"
return content
@@ -174,28 +178,28 @@ def generate_examples():
# Category indices stored in reverse order because they are inserted into
# examples_index.documents at index 0 in order
category_indices = {
"other":
# "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_other_index.md",
title="Other",
path=EXAMPLE_DOC_DIR / "examples_inference_index.md",
title="Inference",
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",
),
"offline_inference":
Index(
path=EXAMPLE_DOC_DIR / "examples_offline_inference_index.md",
title="Offline Inference",
description=
"Offline 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
"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",
),
}
@@ -204,6 +208,7 @@ def generate_examples():
glob_patterns = ["*.py", "*.md", "*.sh"]
# Find categorised examples
for category in category_indices:
print(category)
category_dir = EXAMPLE_DIR / category
globs = [category_dir.glob(pattern) for pattern in glob_patterns]
for path in itertools.chain(*globs):
@@ -224,10 +229,12 @@ def generate_examples():
# Generate the example documentation
for example in sorted(examples, key=lambda e: e.path.stem):
print(example)
doc_path = EXAMPLE_DOC_DIR / f"{example.path.stem}.md"
with open(doc_path, "w+") as f:
f.write(example.generate())
# Add the example to the appropriate index
assert example.category is not None
index = category_indices.get(example.category, examples_index)
index.documents.append(example.path.stem)
@@ -1,10 +0,0 @@
# Examples
A collection of examples demonstrating usage of FastVideo.
All documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.
:::{toctree}
:caption: Examples
:maxdepth: 2
:::
+89 -4
View File
@@ -1,10 +1,95 @@
(fastvideo-installation)=
# 🔧 Installation
The code is tested on Python 3.10.0, CUDA 12.4 and H100.
```
./env_setup.sh fastvideo
FastVideo currently only supports Linux and NVIDIA CUDA GPUs.
FastVideo has been tested on the following GPUs, but it should work on any GPUs that supports CUDA 12.4+, please create an issue if you discover any issues:
- RTX 4090
- A40
- L40S
- A100
- H100
## Requirements
- OS: Linux
- Python: 3.10-3.12
- CUDA 12.4+ (Untested on CUDA < 12.4)
## Installation Options
### Option 1: Quick Install
```bash
pip install fastvideo
```
To try Sliding Tile Attention (optional), please follow the instruction in [here](#sta-installation) to install STA.
### Option 2: Installation from Source
We recommend using a Python environment such as Conda.
#### 1. [Optional] Install Miniconda (if not already installed)
```bash
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
#### 2. [Optional] Create and activate a Conda environment for FastVideo
```bash
conda create -n fastvideo python=3.10 -y
conda activate fastvideo
```
#### 3. Clone the FastVideo repository
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
```
#### 4. Install FastVideo
Basic installation:
```bash
pip install -e .
```
## Optional Dependencies
### Flash Attention
```bash
pip install flash-attn==2.7.0.post2 --no-build-isolation
```
### Sliding Tile Attention (STA) (Requires CUDA 12.4+ and H100)
To try Sliding Tile Attention (optional), please follow the instructions in [csrc/sliding_tile_attention/README.md](#sta-installation) to install STA.
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
[Contributor Guide](#developer-guide)
## Hardware Requirements
### For Basic Inference
- NVIDIA GPU with CUDA support
- Minimum 20GB VRAM for quantized models (e.g., single RTX 4090)
### For Lora Finetuning
- 40GB GPU memory each for 2 GPUs with lora
- 30GB GPU memory each for 2 GPUs with CPU offload and lora
### For Full Finetuning/Distillation
- Multiple high-memory GPUs recommended (e.g., H100)
## Troubleshooting
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg) for additional support.
+8
View File
@@ -69,12 +69,20 @@ sliding_tile_attention/demo
:caption: Inference
:maxdepth: 1
inference/wanvideo
inference/stepvideo
inference/hunyuanvideo
inference/fasthunyuan
inference/fastmochi
:::
:::{toctree}
:caption: Developer Guide
:maxdepth: 1
developer_guide/overview
:::
## Indices and tables
- {ref}`genindex`
+44
View File
@@ -0,0 +1,44 @@
(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.
+3
View File
@@ -0,0 +1,3 @@
# Basic
The class provides the main python interface for using FastVideo's inference pipeline.
+1
View File
@@ -0,0 +1 @@
print('Hello, world!')
+3
View File
@@ -0,0 +1,3 @@
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
__all__ = ["VideoGenerator"]
View File
+4 -3
View File
@@ -7,7 +7,7 @@ from typing import (TYPE_CHECKING, Any, Dict, Generic, Optional, Protocol, Set,
Type, TypeVar)
if TYPE_CHECKING:
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
import torch
@@ -154,7 +154,7 @@ class AttentionMetadataBuilder(ABC, Generic[T]):
self,
current_timestep: int,
forward_batch: "ForwardBatch",
inference_args: "InferenceArgs",
fastvideo_args: "FastVideoArgs",
) -> T:
"""Build attention metadata with on-device tensors."""
raise NotImplementedError
@@ -186,9 +186,10 @@ class AttentionImpl(ABC, Generic[T]):
num_heads: int,
head_size: int,
softmax_scale: float,
dropout_rate: float = 0.0,
causal: bool = False,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
raise NotImplementedError
+18 -9
View File
@@ -3,7 +3,16 @@
from typing import List, Optional, Type
import torch
from flash_attn import flash_attn_func
from flash_attn import flash_attn_func as flash_attn_2_func
try:
from flash_attn_interface import flash_attn_func as flash_attn_3_func
# flash_attn 3 has slightly different API: it returns lse by default
flash_attn_func = lambda q, k, v, softmax_scale, causal: flash_attn_3_func(
q, k, v, softmax_scale, causal)[0]
except ImportError:
flash_attn_func = flash_attn_2_func
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
@@ -45,12 +54,12 @@ class FlashAttentionImpl(AttentionImpl):
self,
num_heads: int,
head_size: int,
dropout_rate: float,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.dropout_rate = dropout_rate
self.causal = causal
self.softmax_scale = softmax_scale
@@ -61,10 +70,10 @@ class FlashAttentionImpl(AttentionImpl):
value: torch.Tensor,
attn_metadata: AttentionMetadata,
):
output = flash_attn_func(query,
key,
value,
dropout_p=self.dropout_rate,
softmax_scale=self.softmax_scale,
causal=self.causal)
output = flash_attn_func(
query, # type: ignore[no-untyped-call]
key,
value,
softmax_scale=self.softmax_scale,
causal=self.causal)
return output
+4 -3
View File
@@ -38,14 +38,15 @@ class SDPAImpl(AttentionImpl):
self,
num_heads: int,
head_size: int,
dropout_rate: float,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.dropout_rate = dropout_rate
self.causal = causal
self.softmax_scale = softmax_scale
self.dropout = extra_impl_args.get("dropout_p", 0.0)
def forward(
self,
@@ -60,7 +61,7 @@ class SDPAImpl(AttentionImpl):
value = value.transpose(1, 2)
attn_kwargs = {
"attn_mask": None,
"dropout_p": self.dropout_rate,
"dropout_p": self.dropout,
"is_causal": self.causal,
"scale": self.softmax_scale
}
@@ -12,7 +12,7 @@ from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.v1.distributed import get_sp_group
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -20,20 +20,39 @@ logger = init_logger(__name__)
# TODO(will-refactor): move this to a utils file
def dict_to_3d_list(mask_strategy,
t_max=50,
l_max=60,
h_max=24) -> List[List[List[Optional[torch.Tensor]]]]:
result = [[[None for _ in range(h_max)] for _ in range(l_max)]
for _ in range(t_max)]
if mask_strategy is None:
return result
def dict_to_3d_list(mask_strategy) -> List[List[List[Optional[torch.Tensor]]]]:
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
max_timesteps_idx = max(
timesteps_idx for timesteps_idx, layer_idx, head_idx in indices) + 1
max_layer_idx = max(layer_idx
for timesteps_idx, layer_idx, head_idx in indices) + 1
max_head_idx = max(head_idx
for timesteps_idx, layer_idx, head_idx in indices) + 1
result = [[[None for _ in range(max_head_idx)]
for _ in range(max_layer_idx)] for _ in range(max_timesteps_idx)]
for key, value in mask_strategy.items():
t, layer, h = map(int, key.split('_'))
result[t][layer][h] = value
timesteps_idx, layer_idx, head_idx = map(int, key.split('_'))
result[timesteps_idx][layer_idx][head_idx] = value
return result
class RangeDict(dict):
def __getitem__(self, item):
for key in self.keys():
if isinstance(key, tuple):
low, high = key
if low <= item <= high:
return super().__getitem__(key)
elif key == item:
return super().__getitem__(key)
raise KeyError(f"seq_len {item} not supported for STA")
class SlidingTileAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@@ -62,7 +81,7 @@ class SlidingTileAttentionBackend(AttentionBackend):
@dataclass
class SlidingTileAttentionMetadata(AttentionMetadata):
text_length: int
current_timestep: int
class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
@@ -77,13 +96,10 @@ class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
self,
current_timestep: int,
forward_batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> SlidingTileAttentionMetadata:
return SlidingTileAttentionMetadata(
current_timestep=current_timestep,
text_length=forward_batch.attention_mask.sum(),
)
return SlidingTileAttentionMetadata(current_timestep=current_timestep, )
class SlidingTileAttentionImpl(AttentionImpl):
@@ -92,10 +108,11 @@ class SlidingTileAttentionImpl(AttentionImpl):
self,
num_heads: int,
head_size: int,
dropout_rate: float,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
# TODO(will-refactor): for now this is the mask strategy, but maybe we should
# have a more general config for STA?
@@ -105,52 +122,73 @@ class SlidingTileAttentionImpl(AttentionImpl):
with open(config_file) as f:
mask_strategy = json.load(f)
mask_strategy = dict_to_3d_list(mask_strategy)
self.prefix = prefix
self.mask_strategy = mask_strategy
sp_group = get_sp_group()
self.sp_size = sp_group.world_size
# STA config
self.STA_base_tile_size = [6, 8, 8]
self.img_latent_shape_mapping = RangeDict({
(115200, 115456): '30x48x80',
82944: '36x48x48',
69120: '18x48x80',
})
self.full_window_mapping = {
'30x48x80': [5, 6, 10],
'36x48x48': [6, 6, 6],
'18x48x80': [3, 6, 10]
}
def tile(self, x: torch.Tensor) -> torch.Tensor:
x = rearrange(x,
"b (sp t h w) head d -> b (t sp h w) head d",
sp=self.sp_size,
t=30 // self.sp_size,
h=48,
w=80)
t=self.img_latent_shape_int[0] // self.sp_size,
h=self.img_latent_shape_int[1],
w=self.img_latent_shape_int[2])
return rearrange(
x,
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
n_t=5,
n_h=6,
n_w=10,
ts_t=6,
ts_h=8,
ts_w=8)
n_t=self.full_window_size[0],
n_h=self.full_window_size[1],
n_w=self.full_window_size[2],
ts_t=self.STA_base_tile_size[0],
ts_h=self.STA_base_tile_size[1],
ts_w=self.STA_base_tile_size[2])
def untile(self, x: torch.Tensor) -> torch.Tensor:
x = rearrange(
x,
"b (n_t n_h n_w ts_t ts_h ts_w) h d -> b (n_t ts_t n_h ts_h n_w ts_w) h d",
n_t=5,
n_h=6,
n_w=10,
ts_t=6,
ts_h=8,
ts_w=8)
n_t=self.full_window_size[0],
n_h=self.full_window_size[1],
n_w=self.full_window_size[2],
ts_t=self.STA_base_tile_size[0],
ts_h=self.STA_base_tile_size[1],
ts_w=self.STA_base_tile_size[2])
return rearrange(x,
"b (t sp h w) head d -> b (sp t h w) head d",
sp=self.sp_size,
t=30 // self.sp_size,
h=48,
w=80)
t=self.img_latent_shape_int[0] // self.sp_size,
h=self.img_latent_shape_int[1],
w=self.img_latent_shape_int[2])
def preprocess_qkv(
self,
qkv: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
img_sequence_length = qkv.shape[1]
self.img_latent_shape_str = self.img_latent_shape_mapping[
img_sequence_length]
self.full_window_size = self.full_window_mapping[
self.img_latent_shape_str]
self.img_latent_shape_int = list(
map(int, self.img_latent_shape_str.split('x')))
self.img_seq_length = self.img_latent_shape_int[
0] * self.img_latent_shape_int[1] * self.img_latent_shape_int[2]
return self.tile(qkv)
def postprocess_output(
@@ -172,24 +210,32 @@ class SlidingTileAttentionImpl(AttentionImpl):
assert self.mask_strategy[
0] is not None, "mask_strategy[0] cannot be None for SlidingTileAttention"
text_length = attn_metadata.text_length
timestep = attn_metadata.current_timestep
# pattern:'.double_blocks.0.attn.impl' or '.single_blocks.0.attn.impl'
layer_idx = int(self.prefix.split('.')[-3])
query = q.transpose(1, 2)
key = k.transpose(1, 2)
value = v.transpose(1, 2)
# TODO: remove hardcode
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[head_idx + start_head]
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)
hidden_states = hidden_states.transpose(1, 2)
# if has_text is False:
# from IPython import embed
# embed()
hidden_states = sliding_tile_attention(
query, key, value, windows, text_length, has_text,
self.img_latent_shape_str).transpose(1, 2)
return hidden_states
+17 -14
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional
from typing import Optional, Tuple
import torch
import torch.nn as nn
@@ -12,6 +12,7 @@ from fastvideo.v1.distributed.communication_op import (
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
class DistributedAttention(nn.Module):
@@ -22,13 +23,13 @@ class DistributedAttention(nn.Module):
num_heads: int,
head_size: int,
num_kv_heads: Optional[int] = None,
dropout_rate: float = 0.0,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = "",
**extra_impl_args) -> None:
super().__init__()
# self.dropout_rate = dropout_rate
# self.causal = causal
if softmax_scale is None:
self.softmax_scale = head_size**-0.5
else:
@@ -38,14 +39,17 @@ class DistributedAttention(nn.Module):
num_kv_heads = num_heads
dtype = torch.get_default_dtype()
attn_backend = get_attn_backend(head_size, dtype, distributed=True)
attn_backend = get_attn_backend(
head_size,
dtype,
supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.impl = impl_cls(num_heads=num_heads,
head_size=head_size,
dropout_rate=dropout_rate,
causal=causal,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
prefix=f"{prefix}.impl",
**extra_impl_args)
self.num_heads = num_heads
self.head_size = head_size
@@ -97,7 +101,6 @@ class DistributedAttention(nn.Module):
qkv = sequence_model_parallel_all_to_all_4D(qkv,
scatter_dim=2,
gather_dim=1)
# Apply backend-specific preprocess_qkv
qkv = self.impl.preprocess_qkv(qkv, ctx_attn_metadata)
@@ -124,8 +127,7 @@ class DistributedAttention(nn.Module):
output = output[:, :seq_len * world_size]
# TODO: make this asynchronous
replicated_output = sequence_model_parallel_all_gather(
replicated_output, dim=2)
replicated_output.contiguous(), dim=2)
# Apply backend-specific postprocess_output
output = self.impl.postprocess_output(output, ctx_attn_metadata)
@@ -143,13 +145,12 @@ class LocalAttention(nn.Module):
num_heads: int,
head_size: int,
num_kv_heads: Optional[int] = None,
dropout_rate: float = 0.0,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
**extra_impl_args) -> None:
super().__init__()
# self.dropout_rate = dropout_rate
# self.causal = causal
if softmax_scale is None:
self.softmax_scale = head_size**-0.5
else:
@@ -158,11 +159,13 @@ class LocalAttention(nn.Module):
num_kv_heads = num_heads
dtype = torch.get_default_dtype()
attn_backend = get_attn_backend(head_size, dtype, distributed=False)
attn_backend = get_attn_backend(
head_size,
dtype,
supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.impl = impl_cls(num_heads=num_heads,
head_size=head_size,
dropout_rate=dropout_rate,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
causal=causal,
+10 -12
View File
@@ -4,7 +4,7 @@
import os
from contextlib import contextmanager
from functools import cache
from typing import Generator, Optional, Type, cast
from typing import Generator, Optional, Tuple, Type, cast
import torch
@@ -82,29 +82,25 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_attn_backend(
head_size: int,
dtype: torch.dtype,
distributed: bool,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
) -> Type[AttentionBackend]:
"""Selects which attention backend to use and lazily imports it."""
# Accessing envs.* behind an @lru_cache decorator can cause the wrong
# value to be returned from the cache if the value changes between calls.
return _cached_get_attn_backend(
head_size=head_size,
dtype=dtype,
distributed=distributed,
)
return _cached_get_attn_backend(head_size, dtype,
supported_attention_backends)
@cache
def _cached_get_attn_backend(
head_size: int,
dtype: torch.dtype,
distributed: bool,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
) -> Type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
#
# THIS SELECTION OVERRIDES THE FASTVIDEO_ATTENTION_BACKEND
# ENVIRONMENT VARIABLE.
if not supported_attention_backends:
raise ValueError("supported_attention_backends is empty")
selected_backend = None
backend_by_global_setting: Optional[_Backend] = (
get_global_forced_attn_backend())
@@ -117,8 +113,10 @@ def _cached_get_attn_backend(
selected_backend = backend_name_to_enum(backend_by_env_var)
# get device-specific attn_backend
if selected_backend not in supported_attention_backends:
selected_backend = None
attention_cls = current_platform.get_attn_backend_cls(
selected_backend, head_size, dtype, distributed)
selected_backend, head_size, dtype)
if not attention_cls:
raise ValueError(
f"Invalid attention backend for {current_platform.device_name}")
View File
+7
View File
@@ -0,0 +1,7 @@
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
__all__ = [
"ArchConfig", "ModelConfig",
"VAEArchConfig", "VAEConfig"
]
+47
View File
@@ -0,0 +1,47 @@
from dataclasses import dataclass, fields
from typing import Dict, Any
# 1. ArchConfig contains all fields from diffuser's/transformer's config.json (i.e. all fields related to the architecture of the model)
# 2. ArchConfig should be inherited & overriden by each model arch_config
# 3. Any field in ArchConfig is fixed upon initialization, and should be hidden away from users
@dataclass
class ArchConfig:
pass
@dataclass
class ModelConfig:
# Every model config parameter can be categorized into either ArchConfig or everything else
# Diffuser/Transformer parameters
arch_config: ArchConfig = ArchConfig()
# FastVideo-specific parameters here
# i.e. STA, quantization, teacache
# This should be used only when loading from transformers/diffusers
def update_model_arch(
self,
source_model_dict: Dict[str, Any]
) -> None:
arch_config = self.arch_config
valid_fields = {f.name for f in fields(arch_config)}
for key, value in source_model_dict.items():
if key in valid_fields:
setattr(arch_config, key, value)
else:
raise AttributeError(f"{type(arch_config).__name__} has no field '{key}'")
def update_model_config(
self,
source_model_dict: Dict[str, Any]
) -> None:
assert "arch_config" not in source_model_dict.keys(), "Source model config shouldn't contain arch_config."
valid_fields = {f.name for f in fields(self)}
for key, value in source_model_dict.items():
if key in valid_fields:
setattr(self, key, value)
else:
print(f"{type(self).__name__} does not contain field '{key}'!")
raise AttributeError(f"Invalid field: {key}")
@@ -0,0 +1,7 @@
from fastvideo.v1.configs.models.vaes.hunyuanvae import HunyuanVAEConfig, HunyuanVAEArchConfig
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig, WanVAEArchConfig
__all__ = [
"HunyuanVAEConfig", "HunyuanVAEArchConfig",
"WanVAEConfig", "WanVAEArchConfig"
]
+36
View File
@@ -0,0 +1,36 @@
from dataclasses import dataclass
from typing import Union
import torch
from fastvideo.v1.configs.models import ArchConfig, ModelConfig
@dataclass
class VAEArchConfig(ArchConfig):
scaling_factor: Union[float, torch.tensor] = 0
temporal_compression_ratio: int = 4
spatial_compression_ratio: int = 8
@dataclass
class VAEConfig(ModelConfig):
arch_config: VAEArchConfig = VAEArchConfig()
# FastVideoVAE-specific parameters
load_encoder: bool = True
load_decoder: bool = True
tile_sample_min_height: int = 256
tile_sample_min_width: int = 256
tile_sample_min_num_frames: int = 16
tile_sample_stride_height: int = 192
tile_sample_stride_width: int = 192
tile_sample_stride_num_frames: int = 12
blend_num_frames: int = 0
use_tiling: bool = True
use_temporal_tiling: bool = True
use_parallel_tiling: bool = True
def __post_init__(self):
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
@@ -0,0 +1,37 @@
from dataclasses import dataclass
from typing import Tuple
from fastvideo.v1.configs.models.vaes.base import VAEConfig, VAEArchConfig
@dataclass
class HunyuanVAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 16
down_block_types: Tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
)
up_block_types: Tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
)
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512)
layers_per_block: int = 2
act_fn: str = "silu"
norm_num_groups: int = 32
scaling_factor: float = 0.476986
spatial_compression_ratio: int = 8
temporal_compression_ratio: int = 4
mid_block_add_attention: bool = True
def __post_init__(self):
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels)-1)
@dataclass
class HunyuanVAEConfig(VAEConfig):
arch_config: VAEArchConfig = HunyuanVAEArchConfig()
@@ -0,0 +1,72 @@
from dataclasses import dataclass
from typing import Tuple
import torch
from fastvideo.v1.configs.models.vaes.base import VAEConfig, VAEArchConfig
@dataclass
class WanVAEArchConfig(VAEArchConfig):
base_dim: int = 96
z_dim: int = 16
dim_mult: Tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
attn_scales: Tuple[float, ...] = ()
temperal_downsample: Tuple[bool, ...] = (False, True, True)
dropout: float = 0.0
latents_mean: Tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
)
latents_std: Tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
)
temporal_compression_ratio = 4
spatial_compression_ratio = 8
def __post_init__(self):
self.scaling_factor: torch.tensor = 1.0 / torch.tensor(self.latents_std).view(
1, self.z_dim, 1, 1, 1)
self.shift_factor: torch.tensor = torch.tensor(self.latents_mean).view(
1, self.z_dim, 1, 1, 1)
@dataclass
class WanVAEConfig(VAEConfig):
arch_config: VAEArchConfig = WanVAEArchConfig()
use_feature_cache: bool = True
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
def __post_init__(self):
self.blend_num_frames = (self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames) * 2
@@ -0,0 +1,9 @@
from fastvideo.v1.configs.pipelines.hunyuan import HunyuanConfig, FastHunyuanConfig
from fastvideo.v1.configs.pipelines.wan import WanT2V480PConfig, WanI2V480PConfig
from fastvideo.v1.configs.pipelines.base import BaseConfig, SlidingTileAttnConfig
from fastvideo.v1.configs.pipelines.registry import get_pipeline_config_cls_for_name
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "get_pipeline_config_cls_for_name"
]
+103
View File
@@ -0,0 +1,103 @@
from dataclasses import dataclass, asdict, fields
from typing import Optional, Dict, Any
import json
from fastvideo.v1.configs.models import ModelConfig, VAEConfig
from fastvideo.v1.utils import shallow_asdict
@dataclass
class BaseConfig:
"""Base configuration for all pipeline architectures."""
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 125
fps: int = 24
# Video generation parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
seed: int = 1024
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
# vae_scale_factor: Optional[int] = None # Deprecated
vae_config: VAEConfig = VAEConfig()
# DiT configuration
num_channels_latents: Optional[int] = None # Deprecated
# Image encoder configuration
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_len: int = -1 # Deprecated
hidden_state_skip_layer: int = 0 # Deprecated
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
neg_prompt: Optional[str] = None
def dump_to_json(self, file_path: str):
output_dict = shallow_asdict(self)
for key, value in output_dict.items():
if isinstance(value, ModelConfig):
model_dict = asdict(value)
# Model Arch Config should be hidden away from the users
model_dict.pop("arch_config")
output_dict[key] = model_dict
with open(file_path, "w") as f:
json.dump(output_dict, f, indent=2)
def load_from_json(self, file_path: str):
with open(file_path, "r") as f:
input_pipeline_dict = json.load(f)
self.update_pipeline_config(input_pipeline_dict)
def update_pipeline_config(
self,
source_pipeline_dict: Dict[str, Any]
) -> None:
for f in fields(self):
key = f.name
if key in source_pipeline_dict:
current_value = getattr(self, key)
new_value = source_pipeline_dict[key]
# If it's a nested ModelConfig, update it recursively
if isinstance(current_value, ModelConfig):
current_value.update_model_config(new_value)
else:
setattr(self, key, new_value)
@dataclass
class SlidingTileAttnConfig(BaseConfig):
"""Configuration for sliding tile attention."""
# Override any BaseConfig defaults as needed
# Add sliding tile specific parameters
window_size: int = 16
stride: int = 8
# You can provide custom defaults for inherited fields
height: int = 576
width: int = 1024
# Additional configuration specific to sliding tile attention
pad_to_square: bool = False
use_overlap_optimization: bool = True
+47
View File
@@ -0,0 +1,47 @@
from dataclasses import dataclass
from fastvideo.v1.configs.pipelines.base import BaseConfig
from fastvideo.v1.configs.models import VAEConfig
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
@dataclass
class HunyuanConfig(BaseConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
# VAE
vae_config: VAEConfig = HunyuanVAEConfig()
# Denoising stage
embedded_cfg_scale: int = 6
flow_shift: int = 7
num_inference_steps: int = 50
# Text encoding stage
hidden_state_skip_layer: int = 2
text_len: int = 256
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precision: str = "fp16"
# HunyuanConfig-specific added parameters
# Secondary text encoder
text_encoder_precision_2: str = "fp16"
text_len_2: int = 77
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@dataclass
class FastHunyuanConfig(HunyuanConfig):
"""Configuration specifically optimized for FastHunyuan weights."""
# Override HunyuanConfig defaults
num_inference_steps: int = 6
flow_shift: int = 17
# No need to re-specify guidance_scale or embedded_cfg_scale as they
# already have the desired values from HunyuanConfig
@@ -0,0 +1,78 @@
"""Registry for pipeline weight-specific configurations."""
import os
from typing import Callable, Dict, Optional, Type
from fastvideo.v1.configs.pipelines.base import BaseConfig
from fastvideo.v1.configs.pipelines.hunyuan import HunyuanConfig, FastHunyuanConfig
from fastvideo.v1.configs.pipelines.wan import WanT2V480PConfig, WanI2V480PConfig
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import (maybe_download_model_index,
verify_model_config_and_directory)
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[BaseConfig]] = {
"FastVideo/FastHunyuan-Diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig
# Add other specific weight variants
}
# For determining pipeline type from model ID
PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
"hunyuan": lambda id: "hunyuan" in id.lower(),
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
# Add other pipeline architecture detectors
}
# Fallback configs when exact match isn't found but architecture is detected
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[BaseConfig]] = {
"hunyuan":
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
"wanpipeline":
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
# Other fallbacks by architecture
}
def get_pipeline_config_cls_for_name(
pipeline_name_or_path: str) -> Optional[type[BaseConfig]]:
"""Get the appropriate config class for specific pretrained weights."""
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
logger.warning(
"FastVideo may not correctly identify the optimal config for this model, as the local directory may have been renamed."
)
else:
config = maybe_download_model_index(pipeline_name_or_path)
pipeline_name = config["_class_name"]
# First try exact match for specific weights
if pipeline_name_or_path in WEIGHT_CONFIG_REGISTRY:
return WEIGHT_CONFIG_REGISTRY[pipeline_name_or_path]
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in WEIGHT_CONFIG_REGISTRY.items():
if registered_id in pipeline_name_or_path:
return config_class
# If no match, try to use the fallback config
fallback_config = None
print(pipeline_name)
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in PIPELINE_DETECTOR.items():
if detector(pipeline_name.lower()):
fallback_config = PIPELINE_FALLBACK_CONFIG.get(pipeline_type)
break
logger.warning("No match found for pipeline %s, using fallback config %s.",
pipeline_name_or_path, fallback_config)
return fallback_config
+58
View File
@@ -0,0 +1,58 @@
from dataclasses import dataclass
from fastvideo.v1.configs.pipelines.base import BaseConfig
from fastvideo.v1.configs.models import VAEConfig
from fastvideo.v1.configs.models.vaes import WanVAEConfig
@dataclass
class WanT2V480PConfig(BaseConfig):
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
# WanConfig-specific parameters with defaults
# VAE
vae_config: VAEConfig = WanVAEConfig()
vae_tiling: bool = False
vae_sp: bool = False
# Video parameters
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
use_cpu_offload: bool = True
# Denoising stage
guidance_scale: float = 3.0
neg_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
flow_shift: int = 3
num_inference_steps: int = 50
# Text encoding stage
text_len: int = 512
# Precision for each component
precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precision: str = "fp32"
# WanConfig-specific added parameters
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@dataclass
class WanI2V480PConfig(WanT2V480PConfig):
"""Base configuration for Wan I2V 14B 480P pipeline architecture."""
# WanConfig-specific parameters with defaults
# Denoising stage
guidance_scale: float = 5.0
num_inference_steps: int = 40
# Precision for each component
image_encoder_precision: str = "fp32"
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@@ -0,0 +1,41 @@
{
"height": 480,
"width": 832,
"num_frames": 81,
"fps": 16,
"num_inference_steps": 50,
"guidance_scale": 3.0,
"seed": 1024,
"guidance_rescale": 0.0,
"embedded_cfg_scale": 6.0,
"flow_shift": 3,
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
"load_encoder": false,
"load_decoder": true,
"tile_sample_min_height": 256,
"tile_sample_min_width": 256,
"tile_sample_min_num_frames": 16,
"tile_sample_stride_height": 192,
"tile_sample_stride_width": 192,
"tile_sample_stride_num_frames": 12,
"blend_num_frames": 8,
"use_tiling": false,
"use_temporal_tiling": false,
"use_parallel_tiling": false,
"use_feature_cache": true
},
"num_channels_latents": null,
"image_encoder_precision": "fp32",
"text_encoder_precision": "fp32",
"text_len": 512,
"hidden_state_skip_layer": 0,
"mask_strategy_file_path": null,
"enable_torch_compile": false,
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
}
@@ -0,0 +1,41 @@
{
"height": 480,
"width": 832,
"num_frames": 81,
"fps": 16,
"num_inference_steps": 40,
"guidance_scale": 5.0,
"seed": 1024,
"guidance_rescale": 0.0,
"embedded_cfg_scale": 6.0,
"flow_shift": 3,
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
"load_encoder": true,
"load_decoder": true,
"tile_sample_min_height": 256,
"tile_sample_min_width": 256,
"tile_sample_min_num_frames": 16,
"tile_sample_stride_height": 192,
"tile_sample_stride_width": 192,
"tile_sample_stride_num_frames": 12,
"blend_num_frames": 8,
"use_tiling": false,
"use_temporal_tiling": false,
"use_parallel_tiling": false,
"use_feature_cache": true
},
"num_channels_latents": null,
"image_encoder_precision": "fp32",
"text_encoder_precision": "fp32",
"text_len": 512,
"hidden_state_skip_layer": 0,
"mask_strategy_file_path": null,
"enable_torch_compile": false,
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
}
+16 -1
View File
@@ -1,5 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.v1.distributed.communication_op import *
from fastvideo.v1.distributed.parallel_state import *
from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size, get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size, get_world_group,
init_distributed_environment, initialize_model_parallel)
from fastvideo.v1.distributed.utils import *
__all__ = [
"init_distributed_environment",
"initialize_model_parallel",
"get_sequence_model_parallel_rank",
"get_sequence_model_parallel_world_size",
"get_tensor_model_parallel_rank",
"get_tensor_model_parallel_world_size",
"cleanup_dist_env_and_memory",
"get_world_group",
]
@@ -190,5 +190,5 @@ class DeviceCommunicatorBase:
torch.distributed.recv(tensor, self.ranks[src], self.device_group)
return tensor
def destroy(self):
def destroy(self) -> None:
pass
+2 -2
View File
@@ -845,12 +845,12 @@ def initialize_model_parallel(
group_name="sp")
def get_sequence_model_parallel_world_size():
def get_sequence_model_parallel_world_size() -> int:
"""Return world size for the sequence model parallel group."""
return get_sp_group().world_size
def get_sequence_model_parallel_rank():
def get_sequence_model_parallel_rank() -> int:
"""Return my rank for the sequence model parallel group."""
return get_sp_group().rank_in_group
+1 -2
View File
@@ -9,8 +9,7 @@ from fastvideo.v1.utils import FlexibleArgumentParser
class CLISubcommand:
"""Base class for CLI subcommands"""
def __init__(self):
self.name = ""
name: str
def cmd(self, args: argparse.Namespace) -> None:
"""Execute the command with the given arguments"""
+4 -8
View File
@@ -2,11 +2,11 @@
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
import argparse
from typing import List
from typing import List, cast
from fastvideo.v1.entrypoints.cli import utils
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.utils import FlexibleArgumentParser
@@ -73,18 +73,14 @@ class GenerateSubcommand(CLISubcommand):
required=False,
help="Read CLI options from a config YAML file.")
generate_parser.add_argument("--num-gpus",
type=int,
default=1,
help="Number of GPUs to use")
generate_parser.add_argument("--master-port",
type=int,
default=None,
help="Port for the master process")
generate_parser = InferenceArgs.add_cli_args(generate_parser)
generate_parser = FastVideoArgs.add_cli_args(generate_parser)
return generate_parser
return cast(FlexibleArgumentParser, generate_parser)
def cmd_init() -> List[CLISubcommand]:
+4 -1
View File
@@ -3,13 +3,16 @@
import os
import subprocess
import sys
from typing import List, Optional
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def launch_distributed(num_gpus=None, args=None, master_port=None):
def launch_distributed(num_gpus: int,
args: List[str],
master_port: Optional[int] = None) -> int:
"""
Launch a distributed job with the given arguments
+290
View File
@@ -0,0 +1,290 @@
# SPDX-License-Identifier: Apache-2.0
"""
VideoGenerator module for FastVideo.
This module provides a consolidated interface for generating videos using
diffusion models.
"""
import os
import time
from dataclasses import asdict
from typing import Any, Callable, Dict, List, Optional, Union
import imageio
import numpy as np
import torch
import torchvision
from einops import rearrange
from fastvideo.v1.configs.pipelines import get_pipeline_config_cls_for_name, BaseConfig
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ForwardBatch
from fastvideo.v1.utils import align_to, shallow_asdict
from fastvideo.v1.worker.executor import Executor
logger = init_logger(__name__)
class VideoGenerator:
"""
A unified class for generating videos using diffusion models.
This class provides a simple interface for video generation with rich
customization options, similar to popular frameworks like HF Diffusers.
"""
def __init__(self, fastvideo_args: FastVideoArgs,
executor_class: type[Executor], log_stats: bool):
"""
Initialize the video generator.
Args:
pipeline: The pipeline to use for inference
fastvideo_args: The inference arguments
"""
self.fastvideo_args = fastvideo_args
self.executor = executor_class(fastvideo_args)
@classmethod
def from_pretrained(cls,
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[Union[str | BaseConfig]] = None,
**kwargs) -> "VideoGenerator":
"""
Create a video generator from a pretrained model.
Args:
model_path: Path or identifier for the pretrained model
device: Device to load the model on (e.g., "cuda", "cuda:0", "cpu")
torch_dtype: Data type for model weights (e.g., torch.float16)
**kwargs: Additional arguments to customize model loading
Returns:
The created video generator
Priority level: Default pipeline config < User's pipeline config < User's kwargs
"""
config = None
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, BaseConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = kwargs
else:
config_args = shallow_asdict(config)
config_args.update(kwargs)
fastvideo_args = FastVideoArgs(
model_path=model_path,
device_str=device or "cuda" if torch.cuda.is_available() else "cpu",
**config_args)
fastvideo_args.check_fastvideo_args()
return cls.from_fastvideo_args(fastvideo_args)
@classmethod
def from_fastvideo_args(cls,
fastvideo_args: FastVideoArgs) -> "VideoGenerator":
"""
Create a video generator with the specified arguments.
Args:
fastvideo_args: The inference arguments
Returns:
The created video generator
"""
# Initialize distributed environment if needed
# initialize_distributed_and_parallelism(fastvideo_args)
executor_class = Executor.get_class(fastvideo_args)
return cls(
fastvideo_args=fastvideo_args,
executor_class=executor_class,
log_stats=False, # TODO: implement
)
def generate_video(
self,
prompt: str,
image_path: Optional[str] = None,
negative_prompt: Optional[str] = None,
output_path: Optional[str] = None,
save_video: bool = True,
return_frames: bool = False,
num_inference_steps: Optional[int] = None,
guidance_scale: Optional[float] = None,
num_frames: Optional[int] = None,
height: Optional[int] = None,
width: Optional[int] = None,
fps: Optional[int] = None,
seed: Optional[int] = None,
callback: Optional[Callable[[int, int, torch.Tensor], None]] = None,
callback_steps: int = 1,
) -> Union[Dict[str, Any], List[np.ndarray]]:
"""
Generate a video based on the given prompt.
Args:
prompt: The prompt to use for generation
negative_prompt: The negative prompt to use (overrides the one in fastvideo_args)
output_path: Path to save the video (overrides the one in fastvideo_args)
save_video: Whether to save the video to disk
return_frames: Whether to return the raw frames
num_inference_steps: Number of denoising steps (overrides fastvideo_args)
guidance_scale: Classifier-free guidance scale (overrides fastvideo_args)
num_frames: Number of frames to generate (overrides fastvideo_args)
height: Height of generated video (overrides fastvideo_args)
width: Width of generated video (overrides fastvideo_args)
fps: Frames per second for saved video (overrides fastvideo_args)
seed: Random seed for generation (overrides fastvideo_args)
callback: Callback function called after each step
callback_steps: Number of steps between each callback
Returns:
Either the output dictionary or the list of frames depending on return_frames
"""
# Create a copy of inference args to avoid modifying the original
fastvideo_args = self.fastvideo_args
# Override parameters if provided
if image_path is not None:
fastvideo_args.image_path = image_path
if negative_prompt is not None:
fastvideo_args.neg_prompt = negative_prompt
if num_inference_steps is not None:
fastvideo_args.num_inference_steps = num_inference_steps
if guidance_scale is not None:
fastvideo_args.guidance_scale = guidance_scale
if num_frames is not None:
fastvideo_args.num_frames = num_frames
if height is not None:
fastvideo_args.height = height
if width is not None:
fastvideo_args.width = width
if fps is not None:
fastvideo_args.fps = fps
if seed is not None:
fastvideo_args.seed = seed
# Validate inputs
if not isinstance(prompt, str):
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
# Process negative prompt
if fastvideo_args.neg_prompt is not None:
fastvideo_args.neg_prompt = fastvideo_args.neg_prompt.strip()
# Validate dimensions
if (fastvideo_args.height <= 0 or fastvideo_args.width <= 0
or fastvideo_args.num_frames <= 0):
raise ValueError(
f"Height, width, and num_frames must be positive integers, got "
f"height={fastvideo_args.height}, width={fastvideo_args.width}, "
f"num_frames={fastvideo_args.num_frames}")
if (fastvideo_args.num_frames - 1) % 4 != 0:
raise ValueError(
f"num_frames-1 must be a multiple of 4, got {fastvideo_args.num_frames}"
)
# Calculate sizes
target_height = align_to(fastvideo_args.height, 16)
target_width = align_to(fastvideo_args.width, 16)
# Calculate latent sizes
latents_size = [(fastvideo_args.num_frames - 1) // 4 + 1,
fastvideo_args.height // 8, fastvideo_args.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# Log parameters
debug_str = f"""
height: {target_height}
width: {target_width}
video_length: {fastvideo_args.num_frames}
prompt: {prompt}
neg_prompt: {fastvideo_args.neg_prompt}
seed: {fastvideo_args.seed}
infer_steps: {fastvideo_args.num_inference_steps}
num_videos_per_prompt: {fastvideo_args.num_videos}
guidance_scale: {fastvideo_args.guidance_scale}
n_tokens: {n_tokens}
flow_shift: {fastvideo_args.flow_shift}
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}"""
logger.info(debug_str)
# Prepare batch
device = torch.device(fastvideo_args.device_str)
batch = ForwardBatch(
prompt=prompt,
image_path=fastvideo_args.image_path,
negative_prompt=fastvideo_args.neg_prompt,
num_videos_per_prompt=fastvideo_args.num_videos,
height=fastvideo_args.height,
width=fastvideo_args.width,
num_frames=fastvideo_args.num_frames,
num_inference_steps=fastvideo_args.num_inference_steps,
guidance_scale=fastvideo_args.guidance_scale,
eta=0.0,
n_tokens=n_tokens,
data_type="video" if fastvideo_args.num_frames > 1 else "image",
device=device,
extra={},
)
# Run inference
start_time = time.time()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch
gen_time = time.time() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
# Process outputs
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
# Save video if requested
if save_video:
save_path = output_path or fastvideo_args.output_path
if save_path:
os.makedirs(os.path.dirname(save_path), exist_ok=True)
video_path = os.path.join(save_path, f"{prompt[:100]}.mp4")
imageio.mimsave(video_path, frames, fps=fastvideo_args.fps)
logger.info("Saved video to %s", video_path)
else:
logger.warning("No output path provided, video not saved")
if return_frames:
return frames
else:
return {
"samples": samples,
"prompts": prompt,
"size":
(target_height, target_width, fastvideo_args.num_frames),
"generation_time": gen_time
}
-14
View File
@@ -8,8 +8,6 @@ if TYPE_CHECKING:
FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60
FASTVIDEO_NCCL_SO_PATH: Optional[str] = None
LD_LIBRARY_PATH: Optional[str] = None
FASTVIDEO_USE_TRITON_FLASH_ATTN: bool = False
FASTVIDEO_FLASH_ATTN_VERSION: Optional[int] = None
LOCAL_RANK: int = 0
CUDA_VISIBLE_DEVICES: Optional[str] = None
FASTVIDEO_CACHE_ROOT: str = os.path.expanduser("~/.cache/fastvideo")
@@ -127,18 +125,6 @@ environment_variables: Dict[str, Callable[[], Any]] = {
"LD_LIBRARY_PATH":
lambda: os.environ.get("LD_LIBRARY_PATH", None),
# flag to control if fastvideo should use triton flash attention
"FASTVIDEO_USE_TRITON_FLASH_ATTN":
lambda:
(os.environ.get("FASTVIDEO_USE_TRITON_FLASH_ATTN", "True").lower() in
("true", "1")),
# Force fastvideo to use a specific flash-attention version (2 or 3), only valid
# when using the flash-attention backend.
"FASTVIDEO_FLASH_ATTN_VERSION":
lambda: maybe_convert_int(
os.environ.get("FASTVIDEO_FLASH_ATTN_VERSION", None)),
# Internal flag to enable Dynamo fullgraph capture
"FASTVIDEO_TEST_DYNAMO_FULLGRAPH_CAPTURE":
lambda: bool(
@@ -0,0 +1,28 @@
# Basic Video Generation Tutorial
The `VideoGenerator` class provides the primary Python interface for doing offline video generation, which is interacting with a diffusion pipeline without using a separate inference api server.
## Usage
The first script in this example shows the most basic usage of FastVideo. If you are new to Python and FastVideo, you should start here.
```bash
python fastvideo/v1/examples/inference/basic/basic.py
```
# Basic Walkthrough
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
```python
from fastvideo import VideoGenerator
generator = VideoGenerator.from_pretrained(
"FastVideo/FastHunyuan-Diffusers",
num_gpus=2,
)
prompt = "A beautiful woman in a red dress walking down a street"
video = generator.generate_video(prompt)
```
More to come! These examples and APIs are still under construction!
@@ -0,0 +1,25 @@
from fastvideo import VideoGenerator
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"FastVideo/FastHunyuan-Diffusers",
# if num_gpus > 1, FastVideo will automatically handle distributed setup
num_gpus=4,
)
# Generate videos with the same simple API, regardless of GPU count
prompt = "A beautiful woman in a red dress walking down a street"
video = generator.generate_video(prompt)
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = "A beautiful woman in a blue dress walking down a street"
video2 = generator.generate_video(prompt2)
if __name__ == "__main__":
main()
@@ -0,0 +1,81 @@
# FastVideo VideoGenerator Gradio Demo
This is a Gradio-based web interface for generating videos using the FastVideo framework. The demo allows users to create videos from text prompts with various customization options.
## Overview
The demo uses the FastVideo framework to generate videos based on text prompts. It provides a simple web interface built with Gradio that allows users to:
- Enter text prompts to generate videos
- Customize video parameters (dimensions, number of frames, etc.)
- Use negative prompts to guide the generation process
- Set or randomize seeds for reproducibility
---
## Requirements
- Linux-based OS
- Python 3.10
- Cuda 12.4
- FastVideo
## Installation
```bash
pip install fastvideo
```
## Usage
Run the demo with:
```bash
python fastvideo/v1/examples/inference/gradio/gradio_demo.py
```
This will start a web server at `http://0.0.0.0:7860` where you can access the interface.
---
## Model Initialization
```python
args = FastVideoArgs(model_path="FastVideo/FastHunyuan-Diffusers", num_gpus=2)
generator = VideoGenerator.from_pretrained(
model_path=args.model_path,
num_gpus=args.num_gpus
)
```
This demo initializes a `VideoGenerator` with the minimum required arguments for inference. Users can seamlessly adjust inference options between generations, including prompts, resolution, video length, or even the number of inference steps, *without ever needing to reload the model*.
## Video Generation
The core functionality is in the `generate_video` function, which:
1. Processes user inputs
2. Uses the FastVideo VideoGenerator from earlier to run inference (`generator.generate_video()`)
3. Returns an output path that Gradio uses to display the generated video
## Gradio Interface
The interface is built with several components:
- A text input for the prompt
- A video display for the result
- Inference options in a collapsible accordion:
- Height and width sliders
- Number of frames slider
- Guidance scale slider
- Inference steps slider
- Negative prompt options
- Seed controls
### Inference Options
- **Height/Width**: Control the resolution of the generated video
- **Number of Frames**: Set how many frames to generate
- **Guidance Scale**: Control how closely the generation follows the prompt
- **Inference Steps**: More steps can improve quality but take longer
- **Negative Prompt**: Specify what you don't want to see in the video
- **Seed**: Control randomness for reproducible results
@@ -0,0 +1,141 @@
import os
import gradio as gr
import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo import VideoGenerator
if __name__ == "__main__":
args = FastVideoArgs(model_path="FastVideo/FastHunyuan-Diffusers", num_gpus=2)
generator = VideoGenerator.from_pretrained(
model_path=args.model_path,
num_gpus=args.num_gpus
)
def generate_video(
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed=False,
):
if randomize_seed:
seed = torch.randint(0, 1000000, (1, )).item()
if not use_negative_prompt:
negative_prompt = None
generator.generate_video(
prompt=prompt,
negative_prompt=negative_prompt,
num_inference_steps=num_inference_steps,
num_frames=num_frames,
height=height,
width=width,
guidance_scale=guidance_scale,
seed=seed
)
output_path = os.path.join(args.output_path, f"{prompt[:100]}.mp4")
return output_path, seed
examples = [
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
]
with gr.Blocks() as demo:
gr.Markdown("# FastVideo Inference Demo")
with gr.Group():
with gr.Row():
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=1,
placeholder="Enter your prompt",
container=False,
)
run_button = gr.Button("Run", scale=0)
result = gr.Video(label="Result", show_label=False)
with gr.Accordion("Advanced options", open=False):
with gr.Group():
with gr.Row():
height = gr.Slider(
label="Height",
minimum=256,
maximum=1024,
step=32,
value=args.height,
)
width = gr.Slider(label="Width", minimum=256, maximum=1024, step=32, value=args.width)
with gr.Row():
num_frames = gr.Slider(
label="Number of Frames",
minimum=21,
maximum=163,
value=45,
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=args.guidance_scale,
)
num_inference_steps = gr.Slider(
label="Inference Steps",
minimum=4,
maximum=100,
value=6,
)
with gr.Row():
use_negative_prompt = gr.Checkbox(label="Use negative prompt", value=False)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=1,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed)
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
seed_output = gr.Number(label="Used Seed")
gr.Examples(examples=examples, inputs=prompt)
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=negative_prompt,
)
run_button.click(
fn=generate_video,
inputs=[
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed,
],
outputs=[result, seed_output],
)
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
@@ -4,23 +4,35 @@
import argparse
import dataclasses
from contextlib import contextmanager
from typing import List, Optional
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import FlexibleArgumentParser
from fastvideo.v1.configs.models import VAEConfig
logger = init_logger(__name__)
@dataclasses.dataclass
class InferenceArgs:
class FastVideoArgs:
# Model and path configuration
model_path: str
# Distributed executor backend
distributed_executor_backend: str = "mp"
inference_mode: bool = True # if False == training mode
# HuggingFace specific parameters
trust_remote_code: bool = False
revision: Optional[str] = None
# Parallelism
tp_size: int = 1
sp_size: int = 1
num_gpus: int = 1
tp_size: Optional[int] = None
sp_size: Optional[int] = None
dist_timeout: Optional[int] = None # timeout for torch.distributed
# Video generation parameters
@@ -31,7 +43,7 @@ class InferenceArgs:
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: int = 7
flow_shift: Optional[float] = None
output_type: str = "pil"
@@ -40,8 +52,16 @@ class InferenceArgs:
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = False
vae_tiling: bool = True # Might change in between forward passes
vae_sp: bool = False # Might change in between forward passes
# vae_scale_factor: Optional[int] = None # Deprecated
vae_config: VAEConfig = VAEConfig()
# DiT configuration
num_channels_latents: Optional[int] = None
# Image encoder configuration
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precision: str = "fp16"
@@ -54,14 +74,14 @@ class InferenceArgs:
# Flow Matching parameters
flow_solver: str = "euler"
denoise_type: str = "flow"
denoise_type: str = "flow" # Deprecated. Will use scheduler_config.json
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
# Scheduler options
scheduler_type: str = "euler"
scheduler_type: str = "euler" # Deprecated. Will use the param in scheduler_config.json
neg_prompt: Optional[str] = None
num_videos: int = 1
@@ -73,6 +93,7 @@ class InferenceArgs:
log_level: str = "info"
# Inference parameters
image_path: Optional[str] = None
prompt: Optional[str] = None
prompt_path: Optional[str] = None
output_path: str = "outputs/"
@@ -104,40 +125,55 @@ class InferenceArgs:
help="Directory containing StepVideo model",
)
# distributed_executor_backend
parser.add_argument(
"--distributed-executor-backend",
type=str,
choices=["mp"],
default=FastVideoArgs.distributed_executor_backend,
help="The distributed executor backend to use",
)
# HuggingFace specific parameters
parser.add_argument(
"--trust-remote-code",
action="store_true",
default=InferenceArgs.trust_remote_code,
default=FastVideoArgs.trust_remote_code,
help="Trust remote code when loading HuggingFace models",
)
parser.add_argument(
"--revision",
type=str,
default=InferenceArgs.revision,
default=FastVideoArgs.revision,
help=
"The specific model version to use (can be a branch name, tag name, or commit id)",
)
# Parallelism
parser.add_argument(
"--num-gpus",
type=int,
default=FastVideoArgs.num_gpus,
help="The number of GPUs to use.",
)
parser.add_argument(
"--tensor-parallel-size",
"--tp-size",
type=int,
default=InferenceArgs.tp_size,
default=FastVideoArgs.tp_size,
help="The tensor parallelism size.",
)
parser.add_argument(
"--sequence-parallel-size",
"--sp-size",
type=int,
default=InferenceArgs.sp_size,
default=FastVideoArgs.sp_size,
help="The sequence parallelism size.",
)
parser.add_argument(
"--dist-timeout",
type=int,
default=InferenceArgs.dist_timeout,
default=FastVideoArgs.dist_timeout,
help="Set timeout for torch.distributed initialization.",
)
@@ -145,56 +181,56 @@ class InferenceArgs:
parser.add_argument(
"--height",
type=int,
default=InferenceArgs.height,
default=FastVideoArgs.height,
help="Height of generated video",
)
parser.add_argument(
"--width",
type=int,
default=InferenceArgs.width,
default=FastVideoArgs.width,
help="Width of generated video",
)
parser.add_argument(
"--num-frames",
type=int,
default=InferenceArgs.num_frames,
default=FastVideoArgs.num_frames,
help="Number of frames to generate",
)
parser.add_argument(
"--num-inference-steps",
type=int,
default=InferenceArgs.num_inference_steps,
default=FastVideoArgs.num_inference_steps,
help="Number of inference steps",
)
parser.add_argument(
"--guidance-scale",
type=float,
default=InferenceArgs.guidance_scale,
default=FastVideoArgs.guidance_scale,
help="Guidance scale for classifier-free guidance",
)
parser.add_argument(
"--guidance-rescale",
type=float,
default=InferenceArgs.guidance_rescale,
default=FastVideoArgs.guidance_rescale,
help="Guidance rescale for classifier-free guidance",
)
parser.add_argument(
"--embedded-cfg-scale",
type=float,
default=InferenceArgs.embedded_cfg_scale,
default=FastVideoArgs.embedded_cfg_scale,
help="Embedded CFG scale",
)
parser.add_argument(
"--flow-shift",
"--shift",
type=int,
default=InferenceArgs.flow_shift,
type=float,
default=FastVideoArgs.flow_shift,
help="Flow shift parameter",
)
parser.add_argument(
"--output-type",
type=str,
default=InferenceArgs.output_type,
default=FastVideoArgs.output_type,
choices=["pil"],
help="Output type for the generated video",
)
@@ -202,7 +238,7 @@ class InferenceArgs:
parser.add_argument(
"--precision",
type=str,
default=InferenceArgs.precision,
default=FastVideoArgs.precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for the model",
)
@@ -211,14 +247,14 @@ class InferenceArgs:
parser.add_argument(
"--vae-precision",
type=str,
default=InferenceArgs.vae_precision,
default=FastVideoArgs.vae_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for VAE",
)
parser.add_argument(
"--vae-tiling",
action="store_true",
default=InferenceArgs.vae_tiling,
default=FastVideoArgs.vae_tiling,
help="Enable VAE tiling",
)
parser.add_argument(
@@ -230,29 +266,39 @@ class InferenceArgs:
parser.add_argument(
"--text-encoder-precision",
type=str,
default=InferenceArgs.text_encoder_precision,
default=FastVideoArgs.text_encoder_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for text encoder",
)
parser.add_argument(
"--text-len",
type=int,
default=InferenceArgs.text_len,
default=FastVideoArgs.text_len,
help="Maximum text length",
)
# Image encoder config
parser.add_argument(
"--image-encoder-precision",
type=str,
default=FastVideoArgs.image_encoder_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for image encoder",
)
# Secondary text encoder
parser.add_argument(
"--text-encoder-precision-2",
type=str,
default=InferenceArgs.text_encoder_precision_2,
default=FastVideoArgs.text_encoder_precision_2,
choices=["fp32", "fp16", "bf16"],
help="Precision for secondary text encoder",
)
parser.add_argument(
"--text-len-2",
type=int,
default=InferenceArgs.text_len_2,
default=FastVideoArgs.text_len_2,
help="Maximum secondary text length",
)
@@ -260,13 +306,13 @@ class InferenceArgs:
parser.add_argument(
"--flow-solver",
type=str,
default=InferenceArgs.flow_solver,
default=FastVideoArgs.flow_solver,
help="Solver for flow matching",
)
parser.add_argument(
"--denoise-type",
type=str,
default=InferenceArgs.denoise_type,
default=FastVideoArgs.denoise_type,
help="Denoise type for noised inputs",
)
@@ -287,7 +333,7 @@ class InferenceArgs:
parser.add_argument(
"--scheduler-type",
type=str,
default=InferenceArgs.scheduler_type,
default=FastVideoArgs.scheduler_type,
help="Type of scheduler to use",
)
@@ -295,19 +341,19 @@ class InferenceArgs:
parser.add_argument(
"--neg-prompt",
type=str,
default=InferenceArgs.neg_prompt,
default=FastVideoArgs.neg_prompt,
help="Negative prompt for sampling",
)
parser.add_argument(
"--num-videos",
type=int,
default=InferenceArgs.num_videos,
default=FastVideoArgs.num_videos,
help="Number of videos to generate per prompt",
)
parser.add_argument(
"--fps",
type=int,
default=InferenceArgs.fps,
default=FastVideoArgs.fps,
help="Frames per second for output video",
)
parser.add_argument(
@@ -326,7 +372,7 @@ class InferenceArgs:
parser.add_argument(
"--log-level",
type=str,
default=InferenceArgs.log_level,
default=FastVideoArgs.log_level,
help="The logging level of all loggers.",
)
@@ -343,23 +389,27 @@ class InferenceArgs:
help="Path to a text file containing the prompt",
)
parser.add_argument("--image-path",
type=str,
help="Path to the image for I2V generation")
parser.add_argument(
"--output-path",
type=str,
default=InferenceArgs.output_path,
default=FastVideoArgs.output_path,
help="Directory to save generated videos",
)
parser.add_argument(
"--seed",
type=int,
default=InferenceArgs.seed,
default=FastVideoArgs.seed,
help="Random seed for reproducibility",
)
return parser
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "InferenceArgs":
def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs":
args.tp_size = args.tensor_parallel_size
args.sp_size = args.sequence_parallel_size
args.flow_shift = getattr(args, "shift", args.flow_shift)
@@ -384,8 +434,20 @@ class InferenceArgs:
return cls(**kwargs)
def check_inference_args(self) -> None:
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
if self.tp_size is None:
self.tp_size = self.num_gpus
if self.sp_size is None:
self.sp_size = self.num_gpus
if self.num_gpus < max(self.tp_size, self.sp_size):
self.num_gpus = max(self.tp_size, self.sp_size)
if self.tp_size != self.sp_size:
raise ValueError(
f"tp_size ({self.tp_size}) must be equal to sp_size ({self.sp_size})"
)
# Validate VAE spatial parallelism with VAE tiling
if self.vae_sp and not self.vae_tiling:
@@ -396,10 +458,10 @@ class InferenceArgs:
raise ValueError("prompt_path must be a text file")
_inference_args = None
_current_fastvideo_args = None
def prepare_inference_args(argv: List[str]) -> InferenceArgs:
def prepare_fastvideo_args(argv: List[str]) -> FastVideoArgs:
"""
Prepare the inference arguments from the command line arguments.
@@ -411,26 +473,38 @@ def prepare_inference_args(argv: List[str]) -> InferenceArgs:
The inference arguments.
"""
parser = FlexibleArgumentParser()
InferenceArgs.add_cli_args(parser)
FastVideoArgs.add_cli_args(parser)
raw_args = parser.parse_args(argv)
inference_args = InferenceArgs.from_cli_args(raw_args)
inference_args.check_inference_args()
global _inference_args
_inference_args = inference_args
return inference_args
fastvideo_args = FastVideoArgs.from_cli_args(raw_args)
fastvideo_args.check_fastvideo_args()
global _current_fastvideo_args
_current_fastvideo_args = fastvideo_args
return fastvideo_args
def get_inference_args() -> InferenceArgs:
global _inference_args
if _inference_args is None:
raise ValueError("Inference arguments not set")
return _inference_args
@contextmanager
def set_current_fastvideo_args(fastvideo_args: FastVideoArgs):
"""
Temporarily set the current fastvideo config.
Used during model initialization.
We save the current fastvideo config in a global variable,
so that all modules can access it, e.g. custom ops
can access the fastvideo config to determine how to dispatch.
"""
global _current_fastvideo_args
old_fastvideo_args = _current_fastvideo_args
try:
_current_fastvideo_args = fastvideo_args
yield
finally:
_current_fastvideo_args = old_fastvideo_args
class DeprecatedAction(argparse.Action):
def __init__(self, option_strings, dest, nargs=0, **kwargs):
super().__init__(option_strings, dest, nargs=nargs, **kwargs)
def __call__(self, parser, namespace, values, option_string=None):
raise ValueError(self.help)
def get_current_fastvideo_args() -> FastVideoArgs:
if _current_fastvideo_args is None:
# in ci, usually when we test custom ops/modules directly,
# we don't set the fastvideo config. In that case, we set a default
# config.
# TODO(will): may need to handle this for CI.
raise ValueError("Current fastvideo args is not set.")
return _current_fastvideo_args
+2 -2
View File
@@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Optional
import torch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
if TYPE_CHECKING:
@@ -52,7 +52,7 @@ def get_forward_context() -> ForwardContext:
@contextmanager
def set_forward_context(current_timestep,
attn_metadata,
inference_args: InferenceArgs = None):
fastvideo_args: Optional[FastVideoArgs] = None):
"""A context manager that stores the current forward context,
can be attention metadata, etc.
Here we can inject common logic for every model forward pass.
+30 -28
View File
@@ -10,7 +10,7 @@ from typing import Any, Dict
import torch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
build_pipeline)
@@ -28,29 +28,29 @@ class InferenceEngine:
def __init__(
self,
pipeline: ComposedPipelineBase,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
):
"""
Initialize the inference engine.
Args:
pipeline: The pipeline to use for inference.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
default_negative_prompt: The default negative prompt to use.
"""
self.pipeline = pipeline
self.inference_args = inference_args
self.fastvideo_args = fastvideo_args
@classmethod
def create_engine(
cls,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> "InferenceEngine":
"""
Create an inference engine with the specified arguments.
Args:
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
model_loader_cls: The model loader class to use. If None, it will be
determined from the model type.
pipeline_type: The type of pipeline to create. If None, it will be
@@ -71,16 +71,16 @@ class InferenceEngine:
# this way for training we can just do pipeline_cls.from_pretrained(
# checkpoint_path) and have it handle everything.
# TODO(Peiyuan): Then maybe we should only pass in model path and device, not the entire inference args?
pipeline = build_pipeline(inference_args)
pipeline = build_pipeline(fastvideo_args)
logger.info("Pipeline Ready")
# Create the inference engine
return cls(pipeline, inference_args)
return cls(pipeline, fastvideo_args)
def run(
self,
prompt: str,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> Dict[str, Any]:
"""
Run inference with the pipeline.
@@ -96,16 +96,17 @@ class InferenceEngine:
"""
out_dict: Dict[str, Any] = dict()
num_videos_per_prompt = inference_args.num_videos
seed = inference_args.seed
height = inference_args.height
width = inference_args.width
video_length = inference_args.num_frames
negative_prompt = inference_args.neg_prompt
infer_steps = inference_args.num_inference_steps
guidance_scale = inference_args.guidance_scale
flow_shift = inference_args.flow_shift
embedded_guidance_scale = inference_args.embedded_cfg_scale
num_videos_per_prompt = fastvideo_args.num_videos
seed = fastvideo_args.seed
height = fastvideo_args.height
width = fastvideo_args.width
video_length = fastvideo_args.num_frames
negative_prompt = fastvideo_args.neg_prompt
infer_steps = fastvideo_args.num_inference_steps
guidance_scale = fastvideo_args.guidance_scale
flow_shift = fastvideo_args.flow_shift
embedded_guidance_scale = fastvideo_args.embedded_cfg_scale
image_path = fastvideo_args.image_path
# ========================================================================
# Arguments: target_width, target_height, target_video_length
@@ -160,20 +161,21 @@ class InferenceEngine:
# return
# sp_group = get_sp_group()
# local_rank = sp_group.rank
device = torch.device(inference_args.device_str)
device = torch.device(fastvideo_args.device_str)
batch = ForwardBatch(
image_path=image_path,
prompt=prompt,
negative_prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
height=inference_args.height,
width=inference_args.width,
num_frames=inference_args.num_frames,
num_inference_steps=inference_args.num_inference_steps,
guidance_scale=inference_args.guidance_scale,
height=fastvideo_args.height,
width=fastvideo_args.width,
num_frames=fastvideo_args.num_frames,
num_inference_steps=fastvideo_args.num_inference_steps,
guidance_scale=fastvideo_args.guidance_scale,
# generator=generator,
eta=0.0,
n_tokens=n_tokens,
data_type="video" if inference_args.num_frames > 1 else "image",
data_type="video" if fastvideo_args.num_frames > 1 else "image",
device=device,
extra={}, # Any additional parameters
)
@@ -182,7 +184,7 @@ class InferenceEngine:
print(batch)
print('===============================================')
print('===============================================')
print(inference_args)
print(fastvideo_args)
# ========================================================================
# Pipeline inference
@@ -190,7 +192,7 @@ class InferenceEngine:
start_time = time.time()
samples = self.pipeline.forward(
batch=batch,
inference_args=inference_args,
fastvideo_args=fastvideo_args,
).output
# TODO(will): fix and move to hunyuan stage
# out_dict["seeds"] = batch.seeds
+1 -1
View File
@@ -23,7 +23,7 @@ class SiluAndMul(CustomOp):
return: (num_tokens, d) or (batch_size, seq_len, d)
"""
def __init__(self):
def __init__(self) -> None:
super().__init__()
if current_platform.is_cuda_alike() or current_platform.is_cpu():
self.op = torch.ops._C.silu_and_mul
+3 -1
View File
@@ -113,7 +113,7 @@ class ScaleResidual(nn.Module):
Applies gated residual connection.
"""
def __init__(self):
def __init__(self, prefix: str = ""):
super().__init__()
def forward(self, residual: torch.Tensor, x: torch.Tensor,
@@ -139,6 +139,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
eps: float = 1e-6,
elementwise_affine: bool = False,
dtype: torch.dtype = torch.float32,
prefix: str = "",
):
super().__init__()
if norm_type == "rms":
@@ -189,6 +190,7 @@ class LayerNormScaleShift(nn.Module):
eps: float = 1e-6,
elementwise_affine: bool = False,
dtype: torch.dtype = torch.float32,
prefix: str = "",
):
super().__init__()
if norm_type == "rms":
+4 -3
View File
@@ -778,12 +778,13 @@ class QKVParallelLinear(ColumnParallelLinear):
# no need to narrow
is_sharded_weight = is_sharded_weight
shard_idx = 0
param_data = param_data.narrow(output_dim, shard_offset, shard_size)
if loaded_shard_id == "q":
shard_id = tp_rank
shard_idx = tp_rank
else:
shard_id = tp_rank // self.num_kv_head_replicas
start_idx = shard_id * shard_size
shard_idx = tp_rank // self.num_kv_head_replicas
start_idx = shard_idx * shard_size
if not is_sharded_weight:
loaded_weight = loaded_weight.narrow(output_dim, start_idx,
+1 -1
View File
@@ -12,7 +12,6 @@ from fastvideo.v1.layers.linear import ReplicatedLinear
class MLP(nn.Module):
"""
MLP for DiT blocks, NO gated linear units
TODO: add Tensor Parallel
"""
def __init__(
@@ -23,6 +22,7 @@ class MLP(nn.Module):
bias: bool = True,
act_type: str = "gelu_pytorch_tanh",
dtype: Optional[torch.dtype] = None,
prefix: str = "",
):
super().__init__()
self.fc_in = ReplicatedLinear(
+2 -2
View File
@@ -84,7 +84,7 @@ class RotaryEmbedding(CustomOp):
head_size: int,
rotary_dim: int,
max_position_embeddings: int,
base: int,
base: Union[int, float],
is_neox_style: bool,
dtype: torch.dtype,
) -> None:
@@ -446,7 +446,7 @@ def get_rope(
head_size: int,
rotary_dim: int,
max_position: int,
base: int,
base: Union[int, float],
is_neox_style: bool = True,
rope_scaling: Optional[Dict[str, Any]] = None,
dtype: Optional[torch.dtype] = None,
+5 -2
View File
@@ -32,7 +32,8 @@ class PatchEmbed(nn.Module):
norm_layer=None,
flatten=True,
bias=True,
dtype=None):
dtype=None,
prefix: str = ""):
super().__init__()
# Convert patch_size to 2-tuple
if isinstance(patch_size, (list, tuple)):
@@ -73,6 +74,7 @@ class TimestepEmbedder(nn.Module):
max_period=10000,
dtype=None,
freq_dtype=torch.float32,
prefix: str = "",
):
super().__init__()
self.frequency_embedding_size = frequency_embedding_size
@@ -132,6 +134,7 @@ class ModulateProjection(nn.Module):
factor: int = 2,
act_layer: str = "silu",
dtype: Optional[torch.dtype] = None,
prefix: str = "",
):
super().__init__()
self.factor = factor
@@ -148,7 +151,7 @@ class ModulateProjection(nn.Module):
return x
def unpatchify(x, t, h, w, patch_size, channels):
def unpatchify(x, t, h, w, patch_size, channels) -> torch.Tensor:
"""
Convert patched representation back to image space.
+86 -2
View File
@@ -20,6 +20,13 @@ FASTVIDEO_LOGGING_CONFIG_PATH = envs.FASTVIDEO_LOGGING_CONFIG_PATH
FASTVIDEO_LOGGING_LEVEL = envs.FASTVIDEO_LOGGING_LEVEL
FASTVIDEO_LOGGING_PREFIX = envs.FASTVIDEO_LOGGING_PREFIX
RED = '\033[91m'
GREEN = '\033[92m'
RESET = '\033[0;0m'
_warned_local_main_process = False
_warned_main_process = False
_FORMAT = (f"{FASTVIDEO_LOGGING_PREFIX}%(levelname)s %(asctime)s "
"[%(filename)s:%(lineno)d] %(message)s")
_DATE_FORMAT = "%m-%d %H:%M:%S"
@@ -68,6 +75,68 @@ def _print_warning_once(logger: Logger, msg: str) -> None:
logger.warning(msg, stacklevel=2)
# TODO(will): add env variable to control this process-aware logging behavior
def _info(logger: Logger,
msg: object,
*args: Any,
main_process_only: bool = False,
local_main_process_only: bool = True,
**kwargs: Any) -> None:
"""Process-aware INFO level logging function.
This function controls logging behavior based on the process rank, allowing for
selective logging from specific processes in a distributed environment.
Args:
logger: The logger instance to use for logging
msg: The message format string to log
*args: Format string arguments
main_process_only: If True, only log if this is the global main process (RANK=0)
local_main_process_only: If True, only log if this is the local main process (LOCAL_RANK=0)
**kwargs: Additional keyword arguments to pass to the logger.log method
- stacklevel: Defaults to 2 to show the original caller's location
Note:
- When both main_process_only and local_main_process_only are True,
the message will be logged only if both conditions are met
- When both are False, the message will be logged from all processes
- By default, only logs from processes with LOCAL_RANK=0
"""
try:
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
except Exception:
local_rank = 0
rank = 0
is_main_process = rank == 0
is_local_main_process = local_rank == 0
if (main_process_only and is_main_process) or (local_main_process_only
and is_local_main_process):
logger.log(logging.INFO, msg, *args, **kwargs)
global _warned_local_main_process, _warned_main_process
if not _warned_local_main_process and local_main_process_only:
logger.warning(
'%s is_local_main_process is set to True, logging only from the local main process.%s',
GREEN,
RESET,
)
_warned_local_main_process = True
if not _warned_main_process and main_process_only:
logger.warning(
'%s is_main_process_only is set to True, logging only from the main process.%s',
GREEN,
RESET,
)
_warned_main_process = True
if not main_process_only and not local_main_process_only:
logger.log(logging.INFO, msg, *args, **kwargs)
class _FastvideoLogger(Logger):
"""
Note:
@@ -91,6 +160,20 @@ class _FastvideoLogger(Logger):
"""
_print_warning_once(self, msg)
def info( # type: ignore[override]
self,
msg: object,
*args: Any,
main_process_only: bool = False,
local_main_process_only: bool = True,
**kwargs: Any) -> None:
_info(self,
msg,
*args,
main_process_only=main_process_only,
local_main_process_only=local_main_process_only,
**kwargs)
def _configure_fastvideo_root_logger() -> None:
logging_config = dict[str, Any]()
@@ -128,7 +211,6 @@ def _configure_fastvideo_root_logger() -> None:
dictConfig(logging_config)
# TODO: add rank_zero_only log
def init_logger(name: str) -> _FastvideoLogger:
"""The main purpose of this function is to ensure that loggers are
retrieved in such a way that we can be sure the root fastvideo logger has
@@ -139,10 +221,12 @@ def init_logger(name: str) -> _FastvideoLogger:
methods_to_patch = {
"info_once": _print_info_once,
"warning_once": _print_warning_once,
"info": _info,
}
for method_name, method in methods_to_patch.items():
setattr(logger, method_name, MethodType(method, logger))
setattr(logger, method_name,
MethodType(method, logger)) # type: ignore[arg-type]
return cast(_FastvideoLogger, logger)
+48 -2
View File
@@ -1,15 +1,61 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import List, Optional, Tuple, Union
import torch
from torch import nn
from fastvideo.v1.platforms import _Backend
# TODO
class BaseDiT(nn.Module):
class BaseDiT(nn.Module, ABC):
_fsdp_shard_conditions: list = []
attention_head_dim: int | None = None
_param_names_mapping: dict
hidden_size: int
num_attention_heads: int
# always supports torch_sdpa
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.TORCH_SDPA, )
def __init_subclass__(cls) -> None:
required_class_attrs = [
"_fsdp_shard_conditions", "_param_names_mapping"
]
super().__init_subclass__()
for attr in required_class_attrs:
if not hasattr(cls, attr):
raise AttributeError(
f"Subclasses of BaseDiT must define '{attr}' class variable"
)
def __init__(self, *args, **kwargs) -> None:
super().__init__()
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
)
def forward(self, *args, **kwargs):
@abstractmethod
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs) -> torch.Tensor:
pass
def __post_init__(self) -> None:
required_attrs = ["hidden_size", "num_attention_heads"]
for attr in required_attrs:
if not hasattr(self, attr):
raise AttributeError(
f"Subclasses of BaseDiT must define '{attr}' instance variable"
)
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
return self._supported_attention_backends
+127 -88
View File
@@ -19,6 +19,7 @@ from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder,
unpatchify)
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.platforms import _Backend
class HunyuanRMSNorm(nn.Module):
@@ -91,6 +92,8 @@ class MMDoubleStreamBlock(nn.Module):
num_attention_heads: int,
mlp_ratio: float,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
prefix: str = "",
):
super().__init__()
@@ -105,6 +108,7 @@ class MMDoubleStreamBlock(nn.Module):
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.img_mod",
)
# Fused operations for image stream
@@ -123,7 +127,8 @@ class MMDoubleStreamBlock(nn.Module):
self.img_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.img_attn_qkv")
self.img_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.img_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
@@ -131,9 +136,14 @@ class MMDoubleStreamBlock(nn.Module):
self.img_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.img_attn_proj")
self.img_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
self.img_mlp = MLP(hidden_size,
mlp_hidden_dim,
bias=True,
dtype=dtype,
prefix=f"{prefix}.img_mlp")
# Text modulation components
self.txt_mod = ModulateProjection(
@@ -141,6 +151,7 @@ class MMDoubleStreamBlock(nn.Module):
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.txt_mod",
)
# Fused operations for text stream
@@ -173,27 +184,12 @@ class MMDoubleStreamBlock(nn.Module):
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
# Distributed attention
self.attn = DistributedAttention(num_heads=num_attention_heads,
head_size=head_dim,
dropout_rate=0.0,
causal=False)
# QK norm layers for text
self.txt_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype)
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
# Distributed attention
self.attn = DistributedAttention(num_heads=num_attention_heads,
head_size=head_dim,
dropout_rate=0.0,
causal=False)
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn")
def forward(
self,
@@ -303,6 +299,8 @@ class MMSingleStreamBlock(nn.Module):
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
prefix: str = "",
):
super().__init__()
@@ -317,13 +315,15 @@ class MMSingleStreamBlock(nn.Module):
self.linear1 = ReplicatedLinear(hidden_size,
hidden_size * 3 + mlp_hidden_dim,
bias=True,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.linear1")
# Combined projection and MLP output
self.linear2 = ReplicatedLinear(hidden_size + mlp_hidden_dim,
hidden_size,
bias=True,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.linear2")
# QK norm layers
self.q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
@@ -345,13 +345,16 @@ class MMSingleStreamBlock(nn.Module):
self.modulation = ModulateProjection(hidden_size,
factor=3,
act_layer="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.modulation")
# Distributed attention
self.attn = DistributedAttention(num_heads=num_attention_heads,
head_size=head_dim,
dropout_rate=0.0,
causal=False)
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn")
def forward(
self,
@@ -433,6 +436,8 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
lambda n, m: "single" in n and str.isdigit(n.split(".")[-1]),
lambda n, m: "refiner" in n and str.isdigit(n.split(".")[-1]),
]
_supported_attention_backends = (_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
_param_names_mapping = {
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
@@ -548,24 +553,25 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
}
def __init__(
self,
patch_size: int = 2,
patch_size_t: int = 1,
in_channels: int = 16,
out_channels: int = 16,
num_attention_heads: int = 24,
attention_head_dim: int = 128,
mlp_ratio: float = 4.0,
num_layers: int = 20,
num_single_layers: int = 40,
num_refiner_layers: int = 2,
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56),
guidance_embeds: bool = False,
dtype: Optional[torch.dtype] = None,
text_embed_dim: int = 4096,
pooled_projection_dim: int = 768,
rope_theta: int = 256,
qk_norm: str = "rms_norm", #TODO(PY)
self,
patch_size: int = 2,
patch_size_t: int = 1,
in_channels: int = 16,
out_channels: int = 16,
num_attention_heads: int = 24,
attention_head_dim: int = 128,
mlp_ratio: float = 4.0,
num_layers: int = 20,
num_single_layers: int = 40,
num_refiner_layers: int = 2,
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56),
guidance_embeds: bool = False,
dtype: Optional[torch.dtype] = None,
text_embed_dim: int = 4096,
pooled_projection_dim: int = 768,
rope_theta: int = 256,
qk_norm: str = "rms_norm", #TODO(PY)
prefix="Hunyuan",
):
super().__init__()
hidden_size = attention_head_dim * num_attention_heads
@@ -598,29 +604,35 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
self.img_in = PatchEmbed(self.patch_size,
self.in_channels,
self.hidden_size,
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.img_in")
self.txt_in = SingleTokenRefiner(self.text_states_dim,
hidden_size,
num_attention_heads,
depth=num_refiner_layers,
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.txt_in")
# Time modulation
self.time_in = TimestepEmbedder(self.hidden_size,
act_layer="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.time_in")
# Text modulation
self.vector_in = MLP(self.text_states_dim_2,
self.hidden_size,
self.hidden_size,
act_type="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.vector_in")
# Guidance modulation
self.guidance_in = (TimestepEmbedder(
self.hidden_size, act_layer="silu", dtype=dtype)
self.guidance_in = (TimestepEmbedder(self.hidden_size,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.guidance_in")
if self.guidance_embeds else None)
# Double blocks
@@ -630,7 +642,8 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
) for _ in range(num_layers)
supported_attention_backends=self._supported_attention_backends,
prefix=f"{prefix}.double_blocks.{i}") for i in range(num_layers)
])
# Single blocks
@@ -640,23 +653,29 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
) for _ in range(num_single_layers)
supported_attention_backends=self._supported_attention_backends,
prefix=f"{prefix}.single_blocks.{i+num_layers}")
for i in range(num_single_layers)
])
self.final_layer = FinalLayer(hidden_size,
self.patch_size,
self.out_channels,
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.final_layer")
self.__post_init__()
# TODO: change the input the FORWAD_BACTCH Dict
# TODO: change output to a dict
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
guidance=None,
):
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs):
"""
Forward pass of the HunyuanDiT model.
@@ -758,26 +777,31 @@ class SingleTokenRefiner(nn.Module):
depth=2,
qkv_bias=True,
dtype=None,
prefix: str = "",
) -> None:
super().__init__()
# Input projection
self.input_embedder = ReplicatedLinear(in_channels,
hidden_size,
bias=True,
params_dtype=dtype)
self.input_embedder = ReplicatedLinear(
in_channels,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.input_embedder")
# Timestep embedding
self.t_embedder = TimestepEmbedder(hidden_size,
act_layer="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.t_embedder")
# Context embedding
self.c_embedder = MLP(in_channels,
hidden_size,
hidden_size,
act_type="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.c_embedder")
# Refiner blocks
self.refiner_blocks = nn.ModuleList([
@@ -786,7 +810,8 @@ class SingleTokenRefiner(nn.Module):
num_attention_heads,
qkv_bias=qkv_bias,
dtype=dtype,
) for _ in range(depth)
prefix=f"{prefix}.refiner_blocks.{i}",
) for i in range(depth)
])
def forward(self, x, t):
@@ -820,6 +845,7 @@ class IndividualTokenRefinerBlock(nn.Module):
mlp_ratio=4.0,
qkv_bias=True,
dtype=None,
prefix: str = "",
) -> None:
super().__init__()
self.num_attention_heads = num_attention_heads
@@ -834,12 +860,15 @@ class IndividualTokenRefinerBlock(nn.Module):
self.self_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=qkv_bias,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.self_attn_qkv")
self.self_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=qkv_bias,
params_dtype=dtype)
self.self_attn_proj = ReplicatedLinear(
hidden_size,
hidden_size,
bias=qkv_bias,
params_dtype=dtype,
prefix=f"{prefix}.self_attn_proj")
# MLP
self.norm2 = nn.LayerNorm(hidden_size,
@@ -850,18 +879,24 @@ class IndividualTokenRefinerBlock(nn.Module):
mlp_hidden_dim,
bias=True,
act_type="silu",
dtype=dtype)
dtype=dtype,
prefix=f"{prefix}.mlp")
# Modulation
self.adaLN_modulation = ModulateProjection(hidden_size,
factor=2,
act_layer="silu",
dtype=dtype)
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
# Scaled dot product attention
self.attn = LocalAttention(
num_heads=num_attention_heads,
head_size=hidden_size // num_attention_heads,
# TODO: remove hardcode; remove STA
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA),
)
def forward(self, x, c):
@@ -900,7 +935,8 @@ class FinalLayer(nn.Module):
hidden_size,
patch_size,
out_channels,
dtype=None) -> None:
dtype=None,
prefix: str = "") -> None:
super().__init__()
# Normalization
@@ -914,13 +950,16 @@ class FinalLayer(nn.Module):
self.linear = ReplicatedLinear(hidden_size,
output_dim,
bias=True,
params_dtype=dtype)
params_dtype=dtype,
prefix=f"{prefix}.linear")
# Modulation
self.adaLN_modulation = ModulateProjection(hidden_size,
factor=2,
act_layer="silu",
dtype=dtype)
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
def forward(self, x, c):
# What the heck HF? Why you change the scale and shift order here???
+126 -137
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any, Dict, Optional, Tuple, Union
from typing import List, Optional, Tuple, Union
import torch
import torch.nn as nn
@@ -21,6 +21,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder)
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.platforms import _Backend
class WanImageEmbedding(torch.nn.Module):
@@ -34,9 +35,10 @@ class WanImageEmbedding(torch.nn.Module):
def forward(self,
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
dtype = encoder_hidden_states_image.dtype
hidden_states = self.norm1(encoder_hidden_states_image)
hidden_states = self.ff(hidden_states)
hidden_states = self.norm2(hidden_states)
hidden_states = self.norm2(hidden_states).to(dtype)
return hidden_states
@@ -52,10 +54,7 @@ class WanTimeTextImageEmbedding(nn.Module):
super().__init__()
self.time_embedder = TimestepEmbedder(
dim,
frequency_embedding_size=time_freq_dim,
act_layer="silu",
freq_dtype=torch.float64)
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
self.time_modulation = ModulateProjection(dim,
factor=6,
act_layer="silu")
@@ -75,9 +74,8 @@ class WanTimeTextImageEmbedding(nn.Module):
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
):
with torch.cuda.amp.autocast(dtype=torch.float32):
temb = self.time_embedder(timestep.float())
timestep_proj = self.time_modulation(temb)
temb = self.time_embedder(timestep)
timestep_proj = self.time_modulation(temb)
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
if encoder_hidden_states_image is not None:
@@ -116,9 +114,14 @@ class WanSelfAttention(nn.Module):
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
# Scaled dot product attention
self.attn = LocalAttention(dropout_rate=0,
softmax_scale=None,
causal=False)
self.attn = LocalAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA))
def forward(self, x: torch.Tensor, context: torch.Tensor,
context_lens: int):
@@ -144,8 +147,8 @@ class WanT2VCrossAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
@@ -159,13 +162,17 @@ class WanT2VCrossAttention(WanSelfAttention):
class WanI2VCrossAttention(WanSelfAttention):
def __init__(self,
dim: int,
num_heads: int,
window_size=(-1, -1),
qk_norm=True,
eps=1e-6) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps)
def __init__(
self,
dim: int,
num_heads: int,
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps,
supported_attention_backends)
self.add_k_proj = ReplicatedLinear(dim, dim)
self.add_v_proj = ReplicatedLinear(dim, dim)
@@ -184,11 +191,11 @@ class WanI2VCrossAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
k_img = self.norm_added_k.forward_native(
self.add_k_proj(context_img)[0]).view(b, -1, n, d)
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
b, -1, n, d)
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
img_x = self.attn(q, k_img, v_img)
# compute attention
@@ -204,16 +211,17 @@ class WanI2VCrossAttention(WanSelfAttention):
class WanTransformerBlock(nn.Module):
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
@@ -222,10 +230,12 @@ class WanTransformerBlock(nn.Module):
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = DistributedAttention(num_heads=num_heads,
head_size=dim // num_heads,
dropout_rate=0.0,
causal=False)
self.attn1 = DistributedAttention(
num_heads=num_heads,
head_size=dim // num_heads,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1")
self.hidden_dim = dim
self.num_attention_heads = num_heads
dim_head = dim // num_heads
@@ -284,16 +294,15 @@ class WanTransformerBlock(nn.Module):
hidden_states = hidden_states.squeeze(1)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
assert temb.dtype == torch.float32
with torch.cuda.amp.autocast(dtype=torch.float32):
e = self.scale_shift_table + temb
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = self.norm1(hidden_states.float()).to(
dtype=orig_dtype) * (1 + scale_msa) + shift_msa
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
@@ -320,6 +329,8 @@ class WanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -327,10 +338,13 @@ class WanTransformerBlock(nn.Module):
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
@@ -339,6 +353,8 @@ class WanTransformer3DModel(BaseDiT):
_fsdp_shard_conditions = [
lambda n, m: "blocks" in n and str.isdigit(n.split(".")[-1]),
]
_supported_attention_backends = (_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
_param_names_mapping = {
r"^patch_embedding\.(.*)$":
r"patch_embedding.proj.\1",
@@ -378,30 +394,30 @@ class WanTransformer3DModel(BaseDiT):
r"blocks.\1.self_attn_residual_norm.norm.\2",
}
def __init__(
self,
patch_size: Tuple[int, int, int] = (1, 2, 2),
text_len=512,
num_attention_heads: int = 40,
attention_head_dim: int = 128,
in_channels: int = 16,
out_channels: int = 16,
text_dim: int = 4096,
freq_dim: int = 256,
ffn_dim: int = 13824,
num_layers: int = 40,
cross_attn_norm: bool = True,
qk_norm: str = "rms_norm_across_heads",
eps: float = 1e-6,
image_dim: Optional[int] = None,
added_kv_proj_dim: Optional[int] = None,
rope_max_seq_len: int = 1024,
) -> None:
def __init__(self,
patch_size: Tuple[int, int, int] = (1, 2, 2),
text_len=512,
num_attention_heads: int = 40,
attention_head_dim: int = 128,
in_channels: int = 16,
out_channels: int = 16,
text_dim: int = 4096,
freq_dim: int = 256,
ffn_dim: int = 13824,
num_layers: int = 40,
cross_attn_norm: bool = True,
qk_norm: str = "rms_norm_across_heads",
eps: float = 1e-6,
image_dim: Optional[int] = None,
added_kv_proj_dim: Optional[int] = None,
rope_max_seq_len: int = 1024,
prefix="Wan") -> None:
super().__init__()
inner_dim = num_attention_heads * attention_head_dim
self.inner_dim = inner_dim
self.hidden_size = inner_dim
self.num_attention_heads = num_attention_heads
self.in_channels = in_channels
self.out_channels = out_channels or in_channels
self.patch_size = patch_size
self.text_len = text_len
@@ -422,9 +438,16 @@ class WanTransformer3DModel(BaseDiT):
# 3. Transformer blocks
self.blocks = nn.ModuleList([
WanTransformerBlock(inner_dim, ffn_dim, num_attention_heads,
qk_norm, cross_attn_norm, eps,
added_kv_proj_dim) for _ in range(num_layers)
WanTransformerBlock(inner_dim,
ffn_dim,
num_attention_heads,
qk_norm,
cross_attn_norm,
eps,
added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{prefix}.blocks.{i}")
for i in range(num_layers)
])
# 4. Output norm & projection
@@ -440,19 +463,24 @@ class WanTransformer3DModel(BaseDiT):
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_hidden_states: torch.Tensor,
seq_len: Optional[int] = None,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
y: Optional[torch.Tensor] = None,
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if y is not None:
hidden_states = torch.cat([hidden_states, y], dim=1)
self.__post_init__()
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
timestep: torch.LongTensor,
encoder_hidden_states_image: Optional[Union[
torch.Tensor, List[torch.Tensor]]] = None,
guidance=None,
**kwargs) -> torch.Tensor:
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
@@ -461,40 +489,23 @@ class WanTransformer3DModel(BaseDiT):
post_patch_width = width // p_w
# Get rotary embeddings
d = self.inner_dim // self.num_attention_heads
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sequence_model_parallel_world_size(),
post_patch_height, post_patch_width),
self.inner_dim,
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float64,
rope_theta=10000)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
if seq_len is None:
seq_len = hidden_states.size(1)
hidden_states = torch.cat([
hidden_states,
hidden_states.new_zeros(1, seq_len - hidden_states.size(1),
hidden_states.size(2))
],
dim=1)
encoder_hidden_states = torch.cat([
encoder_hidden_states,
encoder_hidden_states.new_zeros(
1, self.text_len - encoder_hidden_states.size(1),
encoder_hidden_states.size(2))
],
dim=1)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
@@ -504,6 +515,7 @@ class WanTransformer3DModel(BaseDiT):
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
assert encoder_hidden_states.dtype == orig_dtype
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
@@ -516,39 +528,16 @@ class WanTransformer3DModel(BaseDiT):
timestep_proj, freqs_cis)
# 5. Output norm, projection & unpatchify
with torch.cuda.amp.autocast(dtype=torch.float32):
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(
2, dim=1)
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
hidden_states = self.proj_out(hidden_states)
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
hidden_states = self.proj_out(hidden_states)
output = self.unpatchify(hidden_states, grid_sizes)
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)
return output.float()
def unpatchify(self, x, grid_sizes) -> torch.Tensor:
r"""
Reconstruct video tensors from patch embeddings.
Args:
x (List[Tensor]):
List of patchified features, each with shape [L, C_out * prod(patch_size)]
grid_sizes (Tensor):
Original spatial-temporal grid dimensions before patching,
shape [B, 3] (3 dimensions correspond to F_patches, H_patches, W_patches)
Returns:
Tensor:
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
"""
c = self.out_channels
out = []
for u, v in zip(x, grid_sizes.tolist()):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = u.permute(6, 0, 3, 1, 4, 2, 5)
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
out.append(u)
out = torch.cat(out, dim=0)
return out
return output
+24
View File
@@ -0,0 +1,24 @@
from typing import Tuple
from torch import nn
from fastvideo.v1.platforms import _Backend
class BaseEncoder(nn.Module):
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.TORCH_SDPA, )
def __init__(self, *args, **kwargs) -> None:
super().__init__()
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
)
def forward(self, *args, **kwargs):
pass
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
return self._supported_attention_backends
+17 -3
View File
@@ -19,11 +19,13 @@ from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
RowParallelLinear)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.encoders.base import BaseEncoder
from fastvideo.v1.models.encoders.vision import (VisionEncoderInfo,
resolve_visual_encoder_outputs)
# TODO: support quantization
# from vllm.model_executor.layers.quantization import QuantizationConfig
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
from fastvideo.v1.platforms import _Backend
logger = init_logger(__name__)
@@ -195,7 +197,9 @@ class CLIPAttention(nn.Module):
self.head_dim,
self.num_heads_per_partition,
softmax_scale=self.scale,
causal=True)
causal=True,
supported_attention_backends=self.config.
supported_attention_backends)
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
return tensor.view(bsz, seq_len, self.num_heads,
@@ -464,7 +468,8 @@ class CLIPTextTransformer(nn.Module):
)
class CLIPTextModel(nn.Module):
class CLIPTextModel(BaseEncoder):
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
def __init__(
self,
@@ -475,6 +480,7 @@ class CLIPTextModel(nn.Module):
super().__init__()
self.config = config
self.config.supported_attention_backends = self._supported_attention_backends
self.text_model = CLIPTextTransformer(config=config,
quant_config=quant_config,
prefix=prefix)
@@ -598,6 +604,9 @@ class CLIPVisionTransformer(nn.Module):
inputs_embeds=hidden_states,
return_all_hidden_states=return_all_hidden_states)
if not return_all_hidden_states:
encoder_outputs = encoder_outputs[0]
# Handle post-norm (if applicable) and stacks feature layers if needed
encoder_outputs = resolve_visual_encoder_outputs(
encoder_outputs, feature_sample_layers, self.post_layernorm,
@@ -606,10 +615,11 @@ class CLIPVisionTransformer(nn.Module):
return encoder_outputs
class CLIPVisionModel(nn.Module, SupportsQuant):
class CLIPVisionModel(BaseEncoder, SupportsQuant):
config_class = CLIPVisionConfig
main_input_name = "pixel_values"
packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]}
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
def __init__(
self,
@@ -621,6 +631,8 @@ class CLIPVisionModel(nn.Module, SupportsQuant):
prefix: str = "",
) -> None:
super().__init__()
self.config = config
self.config.supported_attention_backends = self._supported_attention_backends
self.vision_model = CLIPVisionTransformer(
config=config,
quant_config=quant_config,
@@ -654,6 +666,8 @@ class CLIPVisionModel(nn.Module, SupportsQuant):
layer_count = len(self.vision_model.encoder.layers)
for name, loaded_weight in weights:
if name.startswith("visual_projection"):
continue
# post_layernorm is not needed in CLIPVisionModel
if (name.startswith("vision_model.post_layernorm")
and self.vision_model.post_layernorm is None):
+13 -8
View File
@@ -39,10 +39,11 @@ from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
QKVParallelLinear, RowParallelLinear)
from fastvideo.v1.layers.rotary_embedding import get_rope
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.v1.models.encoders.base import BaseEncoder
from fastvideo.v1.models.loader.weight_utils import (default_weight_loader,
maybe_remap_kv_scale_name)
# from ..utils import (extract_layer_index)
from fastvideo.v1.platforms import _Backend
class QuantizationConfig:
@@ -159,16 +160,18 @@ class LlamaAttention(nn.Module):
self.head_dim,
rotary_dim=self.rotary_dim,
max_position=max_position_embeddings,
base=rope_theta,
base=int(rope_theta),
rope_scaling=rope_scaling,
is_neox_style=is_neox_style,
)
self.attn = LocalAttention(self.num_heads,
self.head_dim,
self.num_kv_heads,
softmax_scale=self.scaling,
causal=True)
self.attn = LocalAttention(
self.num_heads,
self.head_dim,
self.num_kv_heads,
softmax_scale=self.scaling,
causal=True,
supported_attention_backends=config.supported_attention_backends)
def forward(
self,
@@ -276,7 +279,8 @@ class LlamaDecoderLayer(nn.Module):
return hidden_states, residual
class LlamaModel(nn.Module):
class LlamaModel(BaseEncoder):
_supported_attention_backends = (_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
def __init__(self,
config: LlamaConfig,
@@ -288,6 +292,7 @@ class LlamaModel(nn.Module):
lora_config = None
self.config = config
self.config.supported_attention_backends = self._supported_attention_backends
self.quant_config = quant_config
if lora_config is not None:
max_loras = 1
+35 -25
View File
@@ -28,11 +28,12 @@ import torch.nn.functional as F
from torch import nn
from transformers import T5Config
from fastvideo.v1.distributed import get_tensor_model_parallel_world_size
from fastvideo.v1.distributed import (get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size)
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.layers.layernorm import RMSNorm
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
RowParallelLinear)
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
QKVParallelLinear, RowParallelLinear)
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
@@ -67,7 +68,8 @@ class T5DenseActDense(nn.Module):
config: T5Config,
quant_config: Optional[QuantizationConfig] = None):
super().__init__()
self.wi = ColumnParallelLinear(config.d_model, config.d_ff, bias=False)
self.wi = MergedColumnParallelLinear(config.d_model, [config.d_ff],
bias=False)
self.wo = RowParallelLinear(config.d_ff,
config.d_model,
bias=False,
@@ -87,14 +89,12 @@ class T5DenseGatedActDense(nn.Module):
config: T5Config,
quant_config: Optional[QuantizationConfig] = None):
super().__init__()
self.wi_0 = ColumnParallelLinear(config.d_model,
config.d_ff,
bias=False,
quant_config=quant_config)
self.wi_1 = ColumnParallelLinear(config.d_model,
config.d_ff,
bias=False,
quant_config=quant_config)
self.wi_0 = MergedColumnParallelLinear(config.d_model, [config.d_ff],
bias=False,
quant_config=quant_config)
self.wi_1 = MergedColumnParallelLinear(config.d_model, [config.d_ff],
bias=False,
quant_config=quant_config)
# Should not run in fp16 unless mixed-precision is used,
# see https://github.com/huggingface/transformers/issues/20287.
self.wo = RowParallelLinear(config.d_ff,
@@ -170,6 +170,7 @@ class T5Attention(nn.Module):
config.relative_attention_max_distance
self.d_model = config.d_model
self.key_value_proj_dim = config.d_kv
self.total_num_heads = self.total_num_kv_heads = config.num_heads
# Partition heads across multiple tensor parallel GPUs.
tp_world_size = get_tensor_model_parallel_world_size()
@@ -178,29 +179,33 @@ class T5Attention(nn.Module):
self.inner_dim = self.n_heads * self.key_value_proj_dim
# No GQA in t5.
self.n_kv_heads = self.n_heads
# self.n_kv_heads = self.n_heads
self.qkv_proj = QKVParallelLinear(self.d_model,
self.d_model // self.n_heads,
self.n_heads,
self.n_kv_heads,
bias=False,
quant_config=quant_config)
self.qkv_proj = QKVParallelLinear(
self.d_model,
self.d_model // self.total_num_heads,
self.total_num_heads,
self.total_num_kv_heads,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.qkv_proj",
)
self.attn = T5MultiHeadAttention()
if self.has_relative_attention_bias:
self.relative_attention_bias = \
VocabParallelEmbedding(self.relative_attention_num_buckets,
self.n_heads,
self.total_num_heads,
org_num_embeddings=self.relative_attention_num_buckets,
padding_size=self.relative_attention_num_buckets,
quant_config=quant_config)
self.o = RowParallelLinear(
self.inner_dim,
self.d_model,
self.d_model,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.o_proj",
)
@staticmethod
@@ -295,13 +300,13 @@ class T5Attention(nn.Module):
) -> torch.Tensor:
bs, seq_len, _ = hidden_states.shape
num_seqs = bs
n, c = self.n_heads, self.d_model // self.n_heads
n, c = self.n_heads, self.d_model // self.total_num_heads
qkv, _ = self.qkv_proj(hidden_states)
# Projection of 'own' hidden state (self-attention). No GQA here.
q, k, v = qkv.split(self.inner_dim, dim=-1)
q = q.view(bs, -1, n, c)
k = k.view(bs, -1, n, c)
v = v.view(bs, -1, n, c)
q = q.reshape(bs, seq_len, n, c)
k = k.reshape(bs, seq_len, n, c)
v = v.reshape(bs, seq_len, n, c)
assert attn_metadata is not None
attn_bias = attn_metadata.attn_bias
@@ -325,6 +330,11 @@ class T5Attention(nn.Module):
-1) if attention_mask.ndim == 2 else attention_mask.unsqueeze(1)
attn_bias.masked_fill_(attention_mask == 0,
torch.finfo(q.dtype).min)
if get_tensor_model_parallel_world_size() > 1:
rank = get_tensor_model_parallel_rank()
attn_bias = attn_bias[:, rank * self.n_heads:(rank + 1) *
self.n_heads, :, :]
attn_output = self.attn(q, k, v, attn_bias)
output, _ = self.o(attn_output)
return output
+6 -4
View File
@@ -52,7 +52,6 @@ def get_hf_config(
trust_remote_code: bool,
revision: Optional[str] = None,
model_override_args: Optional[dict] = None,
inference_args: Optional[dict] = None,
**kwargs,
):
is_gguf = check_gguf_file(model)
@@ -84,20 +83,23 @@ def get_hf_config(
def get_diffusers_config(
model: str,
inference_args: Optional[dict] = None,
fastvideo_args: Optional[dict] = None,
) -> Dict[str, Any]:
"""Gets a configuration for the given diffusers model.
Args:
model: The model name or path.
inference_args: Optional inference arguments to override in the config.
fastvideo_args: Optional inference arguments to override in the config.
Returns:
The loaded configuration.
"""
config_name = "config.json"
if "scheduler" in model:
config_name = "scheduler_config.json"
# Check if the model path exists
if os.path.exists(model):
config_file = os.path.join(model, "config.json")
config_file = os.path.join(model, config_name)
if os.path.exists(config_file):
try:
# Load the config directly from the file
+92 -55
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import dataclasses
from dataclasses import asdict
import glob
import os
import time
@@ -10,10 +11,10 @@ from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
import torch
import torch.nn as nn
from safetensors.torch import load_file as safetensors_load_file
from transformers import AutoTokenizer, PretrainedConfig
from transformers import AutoImageProcessor, AutoTokenizer, PretrainedConfig
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.hf_transformer_utils import (get_diffusers_config,
get_hf_config)
@@ -23,6 +24,7 @@ from fastvideo.v1.models.loader.weight_utils import (
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
pt_weights_iterator, safetensors_weights_iterator)
from fastvideo.v1.models.registry import ModelRegistry
from fastvideo.v1.utils import PRECISION_TO_TYPE
logger = init_logger(__name__)
@@ -35,14 +37,14 @@ class ComponentLoader(ABC):
@abstractmethod
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""
Load the component based on the model path, architecture, and inference args.
Args:
model_path: Path to the component model
architecture: Architecture of the component model
inference_args: Inference arguments
fastvideo_args: Inference arguments
Returns:
The loaded component
@@ -71,6 +73,8 @@ class ComponentLoader(ABC):
"text_encoder_2": (TextEncoderLoader, "transformers"),
"tokenizer": (TokenizerLoader, "transformers"),
"tokenizer_2": (TokenizerLoader, "transformers"),
"image_processor": (ImageProcessorLoader, "transformers"),
"image_encoder": (ImageEncoderLoader, "transformers"),
}
if module_type in module_loaders:
@@ -196,24 +200,27 @@ class TextEncoderLoader(ComponentLoader):
yield from self._get_weights_iterator(source)
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the text encoders based on the model path, architecture, and inference args."""
model_config: PretrainedConfig = get_hf_config(
model=model_path,
trust_remote_code=inference_args.trust_remote_code,
revision=inference_args.revision,
trust_remote_code=fastvideo_args.trust_remote_code,
revision=fastvideo_args.revision,
model_override_args=None,
inference_args=inference_args,
)
logger.info("HF Model config: %s", model_config)
target_device = torch.device(inference_args.device_str)
target_device = torch.device(fastvideo_args.device_str)
# TODO(will): add support for other dtypes
return self.load_model(model_path, model_config, target_device)
return self.load_model(model_path, model_config, target_device,
fastvideo_args.text_encoder_precision)
def load_model(self, model_path: str, model_config,
target_device: torch.device):
with set_default_torch_dtype(torch.float16):
def load_model(self,
model_path: str,
model_config,
target_device: torch.device,
dtype: str = "fp16"):
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
with target_device:
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
@@ -240,11 +247,44 @@ class TextEncoderLoader(ComponentLoader):
return model.eval()
class ImageEncoderLoader(TextEncoderLoader):
def load(self, model_path: str, architecture: str,
fastvideo_args: FastVideoArgs):
"""Load the text encoders based on the model path, architecture, and inference args."""
model_config: PretrainedConfig = get_hf_config(
model=model_path,
trust_remote_code=fastvideo_args.trust_remote_code,
revision=fastvideo_args.revision,
model_override_args=None,
)
logger.info("HF Model config: %s", model_config)
target_device = torch.device(fastvideo_args.device_str)
# TODO(will): add support for other dtypes
return self.load_model(model_path, model_config, target_device,
fastvideo_args.image_encoder_precision)
class ImageProcessorLoader(ComponentLoader):
"""Loader for image processor."""
def load(self, model_path: str, architecture: str,
fastvideo_args: FastVideoArgs):
"""Load the image processor based on the model path, architecture, and inference args."""
logger.info("Loading image processor from %s", model_path)
image_processor = AutoImageProcessor.from_pretrained(model_path, )
logger.info("Loaded image processor: %s",
image_processor.__class__.__name__)
return image_processor
class TokenizerLoader(ComponentLoader):
"""Loader for tokenizers."""
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the tokenizer based on the model path, architecture, and inference args."""
logger.info("Loading tokenizer from %s", model_path)
@@ -262,20 +302,20 @@ class VAELoader(ComponentLoader):
"""Loader for VAE."""
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the VAE based on the model path, architecture, and inference args."""
# TODO(will): move this to a constants file
from fastvideo.v1.utils import PRECISION_TO_TYPE
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name")
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
config.pop("_diffusers_version")
vae_config = fastvideo_args.vae_config
vae_config.update_model_arch(config)
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(**config).to(inference_args.device)
vae = vae_cls(vae_config).to(fastvideo_args.device)
# Find all safetensors files
safetensors_list = glob.glob(
@@ -285,18 +325,10 @@ class VAELoader(ComponentLoader):
safetensors_list
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
loaded = safetensors_load_file(safetensors_list[0])
vae.load_state_dict(loaded)
dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
vae.load_state_dict(loaded, strict=False) # We might only load encoder or decoder
dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae = vae.eval().to(dtype)
# TODO(will): should we define hunyuan vae config class?
vae_kwargs = {
"s_ratio": config["spatial_compression_ratio"],
"t_ratio": config["temporal_compression_ratio"],
}
vae.kwargs = vae_kwargs
return vae
@@ -304,7 +336,7 @@ class TransformerLoader(ComponentLoader):
"""Loader for transformer."""
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the transformer based on the model path, architecture, and inference args."""
model_config = get_diffusers_config(model=model_path)
cls_name = model_config.pop("_class_name")
@@ -314,6 +346,10 @@ class TransformerLoader(ComponentLoader):
"Only diffusers format is supported.")
model_config.pop("_diffusers_version")
# Config from Diffusers supercedes fastvideo's model config
# dit_config = fastvideo_args.dit_config
# model_config.update(dit_config)
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
# Find all safetensors files
@@ -325,20 +361,25 @@ class TransformerLoader(ComponentLoader):
logger.info("Loading model from %s safetensors files in %s",
len(safetensors_list), model_path)
# initialize_sequence_parallel_group(inference_args.sp_size)
# initialize_sequence_parallel_group(fastvideo_args.sp_size)
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
# Load the model using FSDP loader
logger.info("Loading model from %s", cls_name)
model = load_fsdp_model(model_cls=model_cls,
init_params=model_config,
weight_dir_list=safetensors_list,
device=inference_args.device,
cpu_offload=inference_args.use_cpu_offload)
device=fastvideo_args.device,
cpu_offload=fastvideo_args.use_cpu_offload,
default_dtype=default_dtype)
total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
model.eval()
dtypes = set(param.dtype for param in model.parameters())
if len(dtypes) > 1:
model = model.to(default_dtype)
model = model.eval()
return model
@@ -346,23 +387,19 @@ class SchedulerLoader(ComponentLoader):
"""Loader for scheduler."""
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load the scheduler based on the model path, architecture, and inference args."""
if hasattr(inference_args,
'denoise_type') and inference_args.denoise_type == "flow":
# TODO(will): add schedulers to register or create a new scheduler registry
# TODO(will): default to config file but allow override through
# inference args. Currently only uses inference args.
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchDiscreteScheduler)
scheduler = FlowMatchDiscreteScheduler(
shift=inference_args.flow_shift,
solver=inference_args.flow_solver,
)
logger.info("Scheduler loaded: %s", scheduler)
else:
raise ValueError(
f"Invalid denoise type: {inference_args.denoise_type}")
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name")
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
config.pop("_diffusers_version")
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
scheduler = scheduler_cls(**config)
if fastvideo_args.flow_shift is not None:
scheduler.set_shift(fastvideo_args.flow_shift)
return scheduler
@@ -375,7 +412,7 @@ class GenericComponentLoader(ComponentLoader):
self.library = library
def load(self, model_path: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""Load a generic component based on the model path, architecture, and inference args."""
logger.warning("Using generic loader for %s with library %s",
model_path, self.library)
@@ -385,8 +422,8 @@ class GenericComponentLoader(ComponentLoader):
model = AutoModel.from_pretrained(
model_path,
trust_remote_code=inference_args.trust_remote_code,
revision=inference_args.revision,
trust_remote_code=fastvideo_args.trust_remote_code,
revision=fastvideo_args.revision,
)
logger.info("Loaded generic transformers model: %s",
model.__class__.__name__)
@@ -413,7 +450,7 @@ class PipelineComponentLoader:
@staticmethod
def load_module(module_name: str, component_model_path: str,
transformers_or_diffusers: str, architecture: str,
inference_args: InferenceArgs):
fastvideo_args: FastVideoArgs):
"""
Load a pipeline module.
@@ -422,7 +459,7 @@ class PipelineComponentLoader:
component_model_path: Path to the component model
transformers_or_diffusers: Whether the module is from transformers or diffusers
architecture: Architecture of the component model
inference_args: Inference arguments
fastvideo_args: Inference arguments
Returns:
The loaded module
@@ -439,4 +476,4 @@ class PipelineComponentLoader:
transformers_or_diffusers)
# Load the module
return loader.load(component_model_path, architecture, inference_args)
return loader.load(component_model_path, architecture, fastvideo_args)
+29 -48
View File
@@ -2,7 +2,7 @@
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/parameter.py
from fractions import Fraction
from typing import Any, Callable, Optional, Tuple, Union
from typing import Any, Callable, Tuple, Union
import torch
from torch.nn import Parameter
@@ -58,21 +58,22 @@ class BasevLLMParameter(Parameter):
cond2 = loaded_weight.ndim == 0 and loaded_weight.numel() == 1
return (cond1 and cond2)
def _assert_and_load(self, loaded_weight: torch.Tensor):
def _assert_and_load(self, loaded_weight: torch.Tensor) -> None:
assert (self.data.shape == loaded_weight.shape
or self._is_1d_and_scalar(loaded_weight))
self.data.copy_(loaded_weight)
def load_column_parallel_weight(self, loaded_weight: torch.Tensor):
def load_column_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
self._assert_and_load(loaded_weight)
def load_row_parallel_weight(self, loaded_weight: torch.Tensor):
def load_row_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
self._assert_and_load(loaded_weight)
def load_merged_column_weight(self, loaded_weight: torch.Tensor, **kwargs):
def load_merged_column_weight(self, loaded_weight: torch.Tensor,
**kwargs) -> None:
self._assert_and_load(loaded_weight)
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs):
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs) -> None:
self._assert_and_load(loaded_weight)
@@ -95,7 +96,7 @@ class _ColumnvLLMParameter(BasevLLMParameter):
def output_dim(self):
return self._output_dim
def load_column_parallel_weight(self, loaded_weight: torch.Tensor):
def load_column_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
tp_rank = get_tensor_model_parallel_rank()
shard_size = self.data.shape[self.output_dim]
loaded_weight = loaded_weight.narrow(self.output_dim,
@@ -103,10 +104,13 @@ class _ColumnvLLMParameter(BasevLLMParameter):
assert self.data.shape == loaded_weight.shape
self.data.copy_(loaded_weight)
def load_merged_column_weight(self, loaded_weight: torch.Tensor, **kwargs):
def load_merged_column_weight(self, loaded_weight: torch.Tensor,
**kwargs) -> None:
shard_offset = kwargs.get("shard_offset")
shard_size = kwargs.get("shard_size")
if shard_offset is None or shard_size is None:
raise ValueError("shard_offset and shard_size must be provided")
if isinstance(
self,
(PackedColumnParameter,
@@ -124,13 +128,18 @@ class _ColumnvLLMParameter(BasevLLMParameter):
assert param_data.shape == loaded_weight.shape
param_data.copy_(loaded_weight)
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs):
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs) -> None:
shard_offset = kwargs.get("shard_offset")
shard_size = kwargs.get("shard_size")
shard_id = kwargs.get("shard_id")
num_heads = kwargs.get("num_heads")
assert shard_offset is not None
assert shard_size is not None
assert shard_id is not None
assert num_heads is not None
if isinstance(
self,
(PackedColumnParameter,
@@ -166,7 +175,7 @@ class RowvLLMParameter(BasevLLMParameter):
def input_dim(self):
return self._input_dim
def load_row_parallel_weight(self, loaded_weight: torch.Tensor):
def load_row_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
tp_rank = get_tensor_model_parallel_rank()
shard_size = self.data.shape[self.input_dim]
loaded_weight = loaded_weight.narrow(self.input_dim,
@@ -233,16 +242,16 @@ class PerTensorScaleParameter(BasevLLMParameter):
# For row parallel layers, no sharding needed
# load weight into parameter as is
def load_row_parallel_weight(self, *args, **kwargs):
def load_row_parallel_weight(self, *args, **kwargs) -> None:
super().load_row_parallel_weight(*args, **kwargs)
def load_merged_column_weight(self, *args, **kwargs):
def load_merged_column_weight(self, *args, **kwargs) -> None:
self._load_into_shard_id(*args, **kwargs)
def load_qkv_weight(self, *args, **kwargs):
def load_qkv_weight(self, *args, **kwargs) -> None:
self._load_into_shard_id(*args, **kwargs)
def load_column_parallel_weight(self, *args, **kwargs):
def load_column_parallel_weight(self, *args, **kwargs) -> None:
super().load_row_parallel_weight(*args, **kwargs)
def _load_into_shard_id(self, loaded_weight: torch.Tensor,
@@ -273,14 +282,10 @@ class PackedColumnParameter(_ColumnvLLMParameter):
for more details on the packed properties.
"""
def __init__(self,
packed_factor: Union[int, Fraction],
packed_dim: int,
marlin_tile_size: Optional[int] = None,
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
**kwargs):
self._packed_factor = packed_factor
self._packed_dim = packed_dim
self._marlin_tile_size = marlin_tile_size
super().__init__(**kwargs)
@property
@@ -291,17 +296,12 @@ class PackedColumnParameter(_ColumnvLLMParameter):
def packed_factor(self):
return self._packed_factor
@property
def marlin_tile_size(self):
return self._marlin_tile_size
def adjust_shard_indexes_for_packing(self, shard_size,
shard_offset) -> Tuple[Any, Any]:
return _adjust_shard_indexes_for_packing(
shard_size=shard_size,
shard_offset=shard_offset,
packed_factor=self.packed_factor,
marlin_tile_size=self.marlin_tile_size)
packed_factor=self.packed_factor)
class PackedvLLMParameter(ModelWeightParameter):
@@ -315,14 +315,10 @@ class PackedvLLMParameter(ModelWeightParameter):
by accounting for packing and optionally, marlin tile size.
"""
def __init__(self,
packed_factor: Union[int, Fraction],
packed_dim: int,
marlin_tile_size: Optional[int] = None,
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
**kwargs):
self._packed_factor = packed_factor
self._packed_dim = packed_dim
self._marlin_tile_size = marlin_tile_size
super().__init__(**kwargs)
@property
@@ -333,16 +329,11 @@ class PackedvLLMParameter(ModelWeightParameter):
def packed_factor(self):
return self._packed_factor
@property
def marlin_tile_size(self):
return self._marlin_tile_size
def adjust_shard_indexes_for_packing(self, shard_size, shard_offset):
return _adjust_shard_indexes_for_packing(
shard_size=shard_size,
shard_offset=shard_offset,
packed_factor=self.packed_factor,
marlin_tile_size=self.marlin_tile_size)
packed_factor=self.packed_factor)
class BlockQuantScaleParameter(_ColumnvLLMParameter, RowvLLMParameter):
@@ -412,18 +403,8 @@ def permute_param_layout_(param: BasevLLMParameter, input_dim: int,
return param
def _adjust_shard_indexes_for_marlin(shard_size, shard_offset,
marlin_tile_size) -> Tuple[Any, Any]:
return shard_size * marlin_tile_size, shard_offset * marlin_tile_size
def _adjust_shard_indexes_for_packing(shard_size, shard_offset, packed_factor,
marlin_tile_size) -> Tuple[Any, Any]:
def _adjust_shard_indexes_for_packing(shard_size, shard_offset,
packed_factor) -> Tuple[Any, Any]:
shard_size = shard_size // packed_factor
shard_offset = shard_offset // packed_factor
if marlin_tile_size is not None:
return _adjust_shard_indexes_for_marlin(
shard_size=shard_size,
shard_offset=shard_offset,
marlin_tile_size=marlin_tile_size)
return shard_size, shard_offset
+10
View File
@@ -38,6 +38,7 @@ _TEXT_ENCODER_MODELS = {
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
# "HunyuanVideoTransformer3DModel": ("image_encoder", "hunyuanvideo", "HunyuanVideoImageEncoder"),
"CLIPVisionModelWithProjection": ("encoders", "clip", "CLIPVisionModel"),
}
_VAE_MODELS = {
@@ -46,12 +47,21 @@ _VAE_MODELS = {
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
}
_SCHEDULERS = {
"FlowMatchEulerDiscreteScheduler":
("schedulers", "scheduling_flow_match_euler_discrete",
"FlowMatchDiscreteScheduler"),
"UniPCMultistepScheduler":
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
}
_FAST_VIDEO_MODELS = {
**_TEXT_TO_VIDEO_DIT_MODELS,
**_IMAGE_TO_VIDEO_DIT_MODELS,
**_TEXT_ENCODER_MODELS,
**_IMAGE_ENCODER_MODELS,
**_VAE_MODELS,
**_SCHEDULERS,
}
_SUBPROCESS_COMMAND = [
+46
View File
@@ -0,0 +1,46 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import Optional, Tuple, Union
import torch
from diffusers.utils import BaseOutput
class BaseScheduler(ABC):
timesteps: torch.Tensor
order: int
def __init__(self, *args, **kwargs) -> None:
# Check if subclass has defined all required properties
required_attributes = ['timesteps', 'order']
for attr in required_attributes:
if not hasattr(self, attr):
raise AttributeError(
f"Subclasses of BaseScheduler must define '{attr}' property"
)
@abstractmethod
def set_shift(self, shift: float) -> None:
pass
@abstractmethod
def set_timesteps(self, *args, **kwargs) -> None:
pass
@abstractmethod
def scale_model_input(self,
sample: torch.Tensor,
timestep: Optional[int] = None) -> torch.Tensor:
pass
@abstractmethod
def step(
self,
model_output: torch.Tensor,
timestep: Union[int, torch.Tensor],
sample: torch.Tensor,
return_dict: bool = True,
) -> Union[BaseOutput, Tuple]:
pass
@@ -27,6 +27,8 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
from fastvideo.v1.models.schedulers.base import BaseScheduler
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -44,7 +46,7 @@ class FlowMatchDiscreteSchedulerOutput(BaseOutput):
prev_sample: torch.FloatTensor
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
"""
Euler scheduler.
@@ -74,6 +76,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
reverse: bool = True,
solver: str = "euler",
n_tokens: Optional[int] = None,
**kwargs,
):
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
@@ -94,6 +97,8 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
)
BaseScheduler.__init__(self)
@property
def step_index(self):
"""
@@ -170,6 +175,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
return idx
def set_shift(self, shift: float) -> None:
self.config.shift = shift
def _init_step_index(self, timestep) -> None:
if self.begin_index is None:
if isinstance(timestep, torch.Tensor):
@@ -192,6 +200,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
return_dict: bool = True,
**kwargs,
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
File diff suppressed because it is too large Load Diff
-202
View File
@@ -1,202 +0,0 @@
from dataclasses import dataclass
from typing import Any, Optional
import torch
import torch.nn as nn
from transformers.utils import ModelOutput
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def use_default(value, default) -> Any:
return value if value is not None else default
@dataclass
class TextEncoderModelOutput(ModelOutput):
"""
Base class for model's outputs that also contains a pooling of the last hidden states.
Args:
hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
Sequence of hidden-states at the output of the last layer of the model.
attention_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
Mask to avoid performing attention on padding token indices. Mask values selected in ``[0, 1]``:
hidden_states_list (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed):
Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
text_outputs (`list`, *optional*, returned when `return_texts=True` is passed):
List of decoded texts.
"""
hidden_state: torch.FloatTensor = None
attention_mask: Optional[torch.LongTensor] = None
text_outputs: Optional[list] = None
class TextEncoder(nn.Module):
def __init__(
self,
text_encoder,
tokenizer,
max_length: int,
text_encoder_precision: Optional[str] = None,
text_encoder_path: Optional[str] = None,
output_key: Optional[str] = None,
use_attention_mask: bool = True,
prompt_template: Optional[dict] = None,
prompt_template_video: Optional[dict] = None,
hidden_state_skip_layer: Optional[int] = None,
apply_final_norm: bool = False,
device=None,
):
super().__init__()
# TODO(will): check if there's a cleaner way to do this
self.text_encoder_type = text_encoder.config.architectures[0]
self.max_length = max_length
self.precision = text_encoder_precision
self.model_path = text_encoder_path
self.use_attention_mask = use_attention_mask
if prompt_template_video is not None:
assert (use_attention_mask is True
), "Attention mask is True required when training videos."
self.prompt_template = prompt_template
self.prompt_template_video = prompt_template_video
self.hidden_state_skip_layer = hidden_state_skip_layer
self.apply_final_norm = apply_final_norm
if "T5" in self.text_encoder_type:
self.output_key = output_key or "last_hidden_state"
elif "CLIPTextModel" in self.text_encoder_type:
self.output_key = output_key or "pooler_output"
elif "LlamaModel" in self.text_encoder_type or "glm" in self.text_encoder_type:
self.output_key = output_key or "last_hidden_state"
else:
raise ValueError(
f"Unsupported text encoder type: {self.text_encoder_type}")
self.model = text_encoder
# self.dtype = self.model.dtype
self.device = device
self.tokenizer = tokenizer
def __repr__(self):
return f"{self.text_encoder_type} ({self.precision} - {self.model_path})"
@staticmethod
def apply_text_to_template(text, template, prevent_empty_text=True) -> str:
"""
Apply text to template.
Args:
text (str): Input text.
template (str or list): Template string or list of chat conversation.
prevent_empty_text (bool): If True, we will prevent the user text from being empty
by adding a space. Defaults to True.
"""
if isinstance(template, str):
# Will send string to tokenizer. Used for llm
return template.format(text)
else:
raise TypeError(f"Unsupported template type: {type(template)}")
def text2tokens(self, text) -> dict:
"""
Tokenize the input text.
Args:
text (str or list): Input text.
"""
if self.prompt_template_video is not None:
prompt_template = self.prompt_template_video["template"]
text = self.apply_text_to_template(text, prompt_template)
kwargs = dict(
truncation=True,
max_length=self.max_length,
return_tensors="pt",
)
batch_encoding: dict = self.tokenizer(
text,
return_length=False,
return_overflowing_tokens=False,
return_attention_mask=True,
**kwargs,
)
return batch_encoding
def encode(
self,
batch_encoding,
use_attention_mask=None,
hidden_state_skip_layer=None,
device=None,
) -> TextEncoderModelOutput:
"""
Args:
batch_encoding (dict): Batch encoding from tokenizer.
use_attention_mask (bool): Whether to use attention mask. If None, use self.use_attention_mask.
Defaults to None.
output_hidden_states (bool): Whether to output hidden states. If False, return the value of
self.output_key. If True, return the entire output. If set self.hidden_state_skip_layer,
output_hidden_states will be set True. Defaults to False.
hidden_state_skip_layer (int): Number of hidden states to hidden_state_skip_layer. 0 means the last layer.
If None, self.output_key will be used. Defaults to None.
return_texts (bool): Whether to return the decoded texts. Defaults to False.
"""
device = self.model.device if device is None else device
use_attention_mask = use_default(use_attention_mask,
self.use_attention_mask)
hidden_state_skip_layer = use_default(hidden_state_skip_layer,
self.hidden_state_skip_layer)
# note: clip will need attention mask
# TODO(will): unify interface with dit
# TODO (peiyuan): why clip need attention mask?
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs = self.model(
input_ids=batch_encoding["input_ids"].to(device),
output_hidden_states=hidden_state_skip_layer is not None,
)
if hidden_state_skip_layer is not None:
last_hidden_state = outputs.hidden_states[-(
hidden_state_skip_layer + 1)]
# Real last hidden state already has layer norm applied. So here we only apply it
# for intermediate layers.
if hidden_state_skip_layer > 0 and self.apply_final_norm:
last_hidden_state = self.model.final_layer_norm(
last_hidden_state)
else:
last_hidden_state = outputs[self.output_key]
# Remove hidden states of instruction tokens, only keep prompt tokens.
if self.prompt_template_video is not None:
crop_start = self.prompt_template_video.get("crop_start", -1)
last_hidden_state = last_hidden_state[:, crop_start:]
return TextEncoderModelOutput(last_hidden_state)
def forward(
self,
text,
use_attention_mask=None,
output_hidden_states=False,
hidden_state_skip_layer=None,
return_texts=False,
):
batch_encoding = self.text2tokens(text)
return self.encode(
batch_encoding,
use_attention_mask=use_attention_mask,
hidden_state_skip_layer=hidden_state_skip_layer,
)
+50 -26
View File
@@ -2,14 +2,16 @@
from abc import ABC, abstractmethod
from math import prod
from typing import Iterator, Optional, Tuple
from typing import Iterator, Optional, Tuple, Union
import numpy as np
import torch
import torch.distributed as dist
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size)
from fastvideo.v1.configs.models import VAEConfig
class ParallelTiledVAE(ABC):
@@ -19,25 +21,36 @@ class ParallelTiledVAE(ABC):
tile_sample_stride_height: int
tile_sample_stride_width: int
tile_sample_stride_num_frames: int
blend_num_frames: int
use_tiling: bool
temporal_compression_ratio: int
spatial_compression_ratio: int
use_temporal_tiling: bool
use_parallel_tiling: bool
def __init__(self, *args, **kwargs) -> None:
# Check if subclass has defined all required properties
required_attributes = [
'tile_sample_min_height', 'tile_sample_min_width',
'tile_sample_min_num_frames', 'tile_sample_stride_height',
'tile_sample_stride_width', 'tile_sample_stride_num_frames',
'spatial_compression_ratio', 'temporal_compression_ratio',
'use_tiling'
]
def __init__(self, config: VAEConfig, **kwargs) -> None:
self.config = config
self.arch_config = config.arch_config
self.tile_sample_min_height = config.tile_sample_min_height
self.tile_sample_min_width = config.tile_sample_min_width
self.tile_sample_min_num_frames = config.tile_sample_min_num_frames
self.tile_sample_stride_height = config.tile_sample_stride_height
self.tile_sample_stride_width = config.tile_sample_stride_width
self.tile_sample_stride_num_frames = config.tile_sample_stride_num_frames
self.blend_num_frames = config.blend_num_frames
self.use_tiling = config.use_tiling
self.use_temporal_tiling = config.use_temporal_tiling
self.use_parallel_tiling = config.use_parallel_tiling
for attr in required_attributes:
if not hasattr(self, attr):
raise AttributeError(
f"Subclasses of ParallelVAE must define '{attr}' property")
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
@property
def temporal_compression_ratio(self) -> int:
return self.arch_config.temporal_compression_ratio
@property
def spatial_compression_ratio(self) -> int:
return self.arch_config.spatial_compression_ratio
@property
def scaling_factor(self) -> Union[float, torch.tensor]:
return self.arch_config.scaling_factor
@abstractmethod
def _encode(self, *args, **kwargs) -> torch.Tensor:
@@ -52,13 +65,13 @@ class ParallelTiledVAE(ABC):
latent_num_frames = (num_frames -
1) // self.temporal_compression_ratio + 1
if self.use_tiling and num_frames > self.tile_sample_min_num_frames:
if self.use_tiling and self.use_temporal_tiling and num_frames > self.tile_sample_min_num_frames:
latents = self.tiled_encode(x)[:, :, :latent_num_frames]
elif self.use_tiling and (width > self.tile_sample_min_width
or height > self.tile_sample_min_height):
latents = self.spatial_tiled_encode(x)
latents = self.spatial_tiled_encode(x)[:, :, :latent_num_frames]
else:
latents = self._encode(x)
latents = self._encode(x)[:, :, :latent_num_frames]
return DiagonalGaussianDistribution(latents)
def decode(self, z: torch.Tensor) -> torch.Tensor:
@@ -69,16 +82,17 @@ class ParallelTiledVAE(ABC):
num_sample_frames = (num_frames -
1) * self.temporal_compression_ratio + 1
if self.use_tiling and get_sequence_model_parallel_world_size() > 1:
if self.use_tiling and self.use_parallel_tiling and get_sequence_model_parallel_world_size(
) > 1:
return self.parallel_tiled_decode(z)[:, :, :num_sample_frames]
if self.use_tiling and num_frames > tile_latent_min_num_frames:
if self.use_tiling and self.use_temporal_tiling and num_frames > tile_latent_min_num_frames:
return self.tiled_decode(z)[:, :, :num_sample_frames]
if self.use_tiling and (width > tile_latent_min_width
or height > tile_latent_min_height):
return self.spatial_tiled_decode(z)
return self.spatial_tiled_decode(z)[:, :, :num_sample_frames]
return self._decode(z)
return self._decode(z)[:, :, :num_sample_frames]
def blend_v(self, a: torch.Tensor, b: torch.Tensor,
blend_extent: int) -> torch.Tensor:
@@ -402,6 +416,10 @@ class ParallelTiledVAE(ABC):
tile_sample_stride_height: Optional[int] = None,
tile_sample_stride_width: Optional[int] = None,
tile_sample_stride_num_frames: Optional[int] = None,
blend_num_frames: Optional[int] = None,
use_tiling: Optional[bool] = None,
use_temporal_tiling: Optional[bool] = None,
use_parallel_tiling: Optional[bool] = None,
) -> None:
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
@@ -433,7 +451,13 @@ class ParallelTiledVAE(ABC):
self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height
self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width
self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
if blend_num_frames is not None:
self.blend_num_frames = blend_num_frames
else:
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
self.use_tiling = use_tiling or self.use_tiling
self.use_temporal_tiling = use_temporal_tiling or self.use_temporal_tiling
self.use_parallel_tiling = use_parallel_tiling or self.use_parallel_tiling
def disable_tiling(self) -> None:
r"""
@@ -462,7 +486,7 @@ class DiagonalGaussianDistribution:
def sample(self,
generator: Optional[torch.Generator] = None) -> torch.Tensor:
# make sure sample is on the same device as the parameters and has same dtype
sample = torch.randn(
sample = randn_tensor(
self.mean.shape,
generator=generator,
device=self.parameters.device,
+33 -72
View File
@@ -15,7 +15,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Optional, Tuple, Union
from typing import Optional, Tuple, Union, cast
import numpy as np
import torch
@@ -24,8 +24,8 @@ import torch.nn.functional as F
import torch.utils.checkpoint
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.models.utils import auto_attributes
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig, HunyuanVAEArchConfig
def prepare_causal_attention_mask(
@@ -773,92 +773,53 @@ class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
_supports_gradient_checkpointing = True
@auto_attributes
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
latent_channels: int = 16,
down_block_types: Tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
),
up_block_types: Tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
),
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
layers_per_block: int = 2,
act_fn: str = "silu",
norm_num_groups: int = 32,
scaling_factor: float = 0.476986,
spatial_compression_ratio: int = 8,
temporal_compression_ratio: int = 4,
mid_block_add_attention: bool = True,
load_encoder: bool = True,
load_decoder: bool = True,
config: HunyuanVAEConfig,
) -> None:
super().__init__()
ParallelTiledVAE.__init__(self, config)
arch_config: HunyuanVAEArchConfig = cast(HunyuanVAEArchConfig, config.arch_config)
# TODO(will): only pass in config. We do this by manually defining a
# config for hunyuan vae
self.block_out_channels = block_out_channels
if load_encoder:
self.block_out_channels = arch_config.block_out_channels
if config.load_encoder:
self.encoder = HunyuanVideoEncoder3D(
in_channels=in_channels,
out_channels=latent_channels,
down_block_types=down_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
in_channels=arch_config.in_channels,
out_channels=arch_config.latent_channels,
down_block_types=arch_config.down_block_types,
block_out_channels=arch_config.block_out_channels,
layers_per_block=arch_config.layers_per_block,
norm_num_groups=arch_config.norm_num_groups,
act_fn=arch_config.act_fn,
double_z=True,
mid_block_add_attention=mid_block_add_attention,
temporal_compression_ratio=temporal_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=arch_config.mid_block_add_attention,
temporal_compression_ratio=arch_config.temporal_compression_ratio,
spatial_compression_ratio=arch_config.spatial_compression_ratio,
)
self.quant_conv = nn.Conv3d(2 * latent_channels,
2 * latent_channels,
self.quant_conv = nn.Conv3d(2 * arch_config.latent_channels,
2 * arch_config.latent_channels,
kernel_size=1)
if load_decoder:
if config.load_decoder:
self.decoder = HunyuanVideoDecoder3D(
in_channels=latent_channels,
out_channels=out_channels,
up_block_types=up_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
time_compression_ratio=temporal_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=mid_block_add_attention,
in_channels=arch_config.latent_channels,
out_channels=arch_config.out_channels,
up_block_types=arch_config.up_block_types,
block_out_channels=arch_config.block_out_channels,
layers_per_block=arch_config.layers_per_block,
norm_num_groups=arch_config.norm_num_groups,
act_fn=arch_config.act_fn,
time_compression_ratio=arch_config.temporal_compression_ratio,
spatial_compression_ratio=arch_config.spatial_compression_ratio,
mid_block_add_attention=arch_config.mid_block_add_attention,
)
self.post_quant_conv = nn.Conv3d(latent_channels,
latent_channels,
self.post_quant_conv = nn.Conv3d(arch_config.latent_channels,
arch_config.latent_channels,
kernel_size=1)
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
# intermediate tiles together, the memory requirement can be lowered.
self.use_tiling = True
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 16
# The minimal distance between two spatial tiles
self.tile_sample_stride_height = 192
self.tile_sample_stride_width = 192
self.tile_sample_stride_num_frames = 12
ParallelTiledVAE.__init__(self)
def _encode(self, x: torch.Tensor) -> torch.Tensor:
x = self.encoder(x)
enc = self.quant_conv(x)
+332 -113
View File
@@ -1,3 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
@@ -12,19 +14,39 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Optional, Tuple, Union
import contextvars
from contextlib import contextmanager
from typing import Optional, Tuple, Union, cast
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.checkpoint
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.utils import auto_attributes
from fastvideo.v1.models.vaes.common import ParallelTiledVAE, DiagonalGaussianDistribution
from fastvideo.v1.configs.models.vaes import WanVAEConfig, WanVAEArchConfig
CACHE_T = 2
is_first_frame = contextvars.ContextVar("is_first_frame", default=False)
feat_cache = contextvars.ContextVar("feat_cache", default=None)
feat_idx = contextvars.ContextVar("feat_idx", default=0)
@contextmanager
def forward_context(first_frame_arg=False,
feat_cache_arg=None,
feat_idx_arg=None):
is_first_frame_token = is_first_frame.set(first_frame_arg)
feat_cache_token = feat_cache.set(feat_cache_arg)
feat_idx_token = feat_idx.set(feat_idx_arg)
try:
yield
finally:
is_first_frame.reset(is_first_frame_token)
feat_cache.reset(feat_cache_token)
feat_idx.reset(feat_idx_token)
class WanCausalConv3d(nn.Conv3d):
r"""
@@ -58,12 +80,17 @@ class WanCausalConv3d(nn.Conv3d):
)
self.padding: Tuple[int, int, int]
# Set up causal padding
self._padding = (self.padding[2], self.padding[2], self.padding[1],
self.padding[1], 2 * self.padding[0], 0)
self._padding: Tuple[int, ...] = (self.padding[2], self.padding[2],
self.padding[1], self.padding[1],
2 * self.padding[0], 0)
self.padding = (0, 0, 0)
def forward(self, x):
def forward(self, x, cache_x=None):
padding = list(self._padding)
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
x = torch.cat([cache_x, x], dim=2)
padding[4] -= cache_x.shape[2]
x = F.pad(x, padding)
return super().forward(x)
@@ -155,28 +182,82 @@ class WanResample(nn.Module):
self.time_conv = WanCausalConv3d(dim,
dim, (3, 1, 1),
stride=(2, 1, 1),
padding=(1, 0, 0))
padding=(0, 0, 0))
else:
self.resample = nn.Identity()
def forward(self, x, first_frame=False):
def forward(self, x):
b, c, t, h, w = x.size()
first_frame = is_first_frame.get()
if first_frame:
assert t == 1
if self.mode == "upsample3d" and not first_frame and hasattr(
self, "time_conv"):
x = self.time_conv(x)
x = x.reshape(b, 2, c, t, h, w)
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3)
x = x.reshape(b, c, t * 2, h, w)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if self.mode == "upsample3d":
if _feat_cache is not None:
idx = _feat_idx
if _feat_cache[idx] is None:
_feat_cache[idx] = "Rep"
_feat_idx += 1
else:
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[
idx] is not None and _feat_cache[idx] != "Rep":
# cache last frame of last two chunk
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
if cache_x.shape[2] < 2 and _feat_cache[
idx] is not None and _feat_cache[idx] == "Rep":
cache_x = torch.cat([
torch.zeros_like(cache_x).to(cache_x.device),
cache_x
],
dim=2)
if _feat_cache[idx] == "Rep":
x = self.time_conv(x)
else:
x = self.time_conv(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
x = x.reshape(b, 2, c, t, h, w)
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]),
3)
x = x.reshape(b, c, t * 2, h, w)
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
elif not first_frame and hasattr(self, "time_conv"):
x = self.time_conv(x)
x = x.reshape(b, 2, c, t, h, w)
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3)
x = x.reshape(b, c, t * 2, h, w)
t = x.shape[2]
x = x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)
x = self.resample(x)
x = x.view(b, t, x.size(1), x.size(2), x.size(3)).permute(0, 2, 1, 3, 4)
if self.mode == "downsample3d" and not first_frame and hasattr(
self, "time_conv"):
x = self.time_conv(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if self.mode == "downsample3d":
if _feat_cache is not None:
idx = _feat_idx
if _feat_cache[idx] is None:
_feat_cache[idx] = x.clone()
_feat_idx += 1
else:
cache_x = x[:, :, -1:, :, :].clone()
x = self.time_conv(
torch.cat([_feat_cache[idx][:, :, -1:, :, :], x], 2))
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
elif not first_frame and hasattr(self, "time_conv"):
x = self.time_conv(x)
return x
@@ -220,7 +301,25 @@ class WanResidualBlock(nn.Module):
x = self.norm1(x)
x = self.nonlinearity(x)
x = self.conv1(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv1(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv1(x)
# Second normalization and activation
x = self.norm2(x)
@@ -229,7 +328,25 @@ class WanResidualBlock(nn.Module):
# Dropout
x = self.dropout(x)
x = self.conv2(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv2(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv2(x)
# Add residual connection
return x + h
@@ -398,15 +515,30 @@ class WanEncoder3d(nn.Module):
self.gradient_checkpointing = False
def forward(self, x, first_frame=False):
x = self.conv_in(x)
def forward(self, x):
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv_in(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv_in(x)
## downsamples
for layer in self.down_blocks:
if isinstance(layer, WanResample):
x = layer(x, first_frame=first_frame)
else:
x = layer(x)
x = layer(x)
## middle
x = self.mid_block(x)
@@ -414,7 +546,26 @@ class WanEncoder3d(nn.Module):
## head
x = self.norm_out(x)
x = self.nonlinearity(x)
x = self.conv_out(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv_out(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv_out(x)
return x
@@ -463,7 +614,7 @@ class WanUpBlock(nn.Module):
self.gradient_checkpointing = False
def forward(self, x, first_frame=False):
def forward(self, x):
"""
Forward pass through the upsampling block.
@@ -479,7 +630,7 @@ class WanUpBlock(nn.Module):
x = resnet(x)
if self.upsamplers is not None:
x = self.upsamplers[0](x, first_frame=first_frame)
x = self.upsamplers[0](x)
return x
@@ -567,21 +718,57 @@ class WanDecoder3d(nn.Module):
self.gradient_checkpointing = False
def forward(self, x, first_frame=False):
def forward(self, x):
## conv1
x = self.conv_in(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv_in(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv_in(x)
## middle
x = self.mid_block(x)
## upsamples
for up_block in self.up_blocks:
x = up_block(x, first_frame=first_frame)
x = up_block(x)
## head
x = self.norm_out(x)
x = self.nonlinearity(x)
x = self.conv_out(x)
_feat_cache = feat_cache.get()
_feat_idx = feat_idx.get()
if _feat_cache is not None:
idx = _feat_idx
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv_out(x, _feat_cache[idx])
_feat_cache[idx] = cache_x
_feat_idx += 1
feat_cache.set(_feat_cache)
feat_idx.set(_feat_idx)
else:
x = self.conv_out(x)
return x
@@ -593,90 +780,89 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
_supports_gradient_checkpointing = False
@auto_attributes
def __init__(self,
base_dim: int = 96,
z_dim: int = 16,
dim_mult: Tuple[int, ...] = (1, 2, 4, 4),
num_res_blocks: int = 2,
attn_scales: Tuple[float, ...] = (),
temperal_downsample: Tuple[bool, ...] = (False, True, True),
dropout: float = 0.0,
latents_mean: Tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
),
latents_std: Tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
),
load_encoder: bool = True,
load_decoder: bool = True) -> None:
config: WanVAEConfig,
) -> None:
super().__init__()
ParallelTiledVAE.__init__(self, config)
self.z_dim = z_dim
self.temperal_downsample = list(temperal_downsample)
self.temperal_upsample = list(temperal_downsample)[::-1]
self.latents_mean = list(latents_mean)
self.latents_std = list(latents_std)
self.arch_config = cast(WanVAEArchConfig, self.arch_config)
self.z_dim = self.arch_config.z_dim
self.temperal_downsample = list(self.arch_config.temperal_downsample)
self.temperal_upsample = list(self.arch_config.temperal_downsample)[::-1]
self.latents_mean = list(self.arch_config.latents_mean)
self.latents_std = list(self.arch_config.latents_std)
self.shift_factor = self.arch_config.shift_factor
if load_encoder:
self.encoder = WanEncoder3d(base_dim, z_dim * 2, dim_mult,
num_res_blocks, attn_scales,
self.temperal_downsample, dropout)
self.quant_conv = WanCausalConv3d(z_dim * 2, z_dim * 2, 1)
self.post_quant_conv = WanCausalConv3d(z_dim, z_dim, 1)
if config.load_encoder:
self.encoder = WanEncoder3d(self.arch_config.base_dim, self.z_dim * 2, self.arch_config.dim_mult,
self.arch_config.num_res_blocks, self.arch_config.attn_scales,
self.temperal_downsample, self.arch_config.dropout)
self.quant_conv = WanCausalConv3d(self.z_dim * 2, self.z_dim * 2, 1)
self.post_quant_conv = WanCausalConv3d(self.z_dim, self.z_dim, 1)
if load_decoder:
self.decoder = WanDecoder3d(base_dim, z_dim, dim_mult,
num_res_blocks, attn_scales,
self.temperal_upsample, dropout)
if config.load_decoder:
self.decoder = WanDecoder3d(self.arch_config.base_dim, self.z_dim, self.arch_config.dim_mult,
self.arch_config.num_res_blocks, self.arch_config.attn_scales,
self.temperal_upsample, self.arch_config.dropout)
self.use_tiling = True
self.spatial_compression_ratio = 8
self.temporal_compression_ratio = 4
self.use_feature_cache = config.use_feature_cache
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 16
def clear_cache(self) -> None:
# The minimal distance between two spatial tiles
self.tile_sample_stride_height = 192
self.tile_sample_stride_width = 192
self.tile_sample_stride_num_frames = 12
ParallelTiledVAE.__init__(self)
def _count_conv3d(model) -> int:
count = 0
for m in model.modules():
if isinstance(m, WanCausalConv3d):
count += 1
return count
if self.config.load_decoder:
self._conv_num = _count_conv3d(self.decoder)
self._conv_idx = 0
self._feat_map = [None] * self._conv_num
# cache encode
if self.config.load_encoder:
self._enc_conv_num = _count_conv3d(self.encoder)
self._enc_conv_idx = 0
self._enc_feat_map = [None] * self._enc_conv_num
def encode(self, x: torch.Tensor) -> torch.Tensor:
if self.use_feature_cache:
self.clear_cache()
with forward_context(feat_cache_arg=self._enc_feat_map,
feat_idx_arg=self._enc_conv_idx):
t = x.shape[2]
iter_ = 1 + (t - 1) // 4
for i in range(iter_):
feat_idx.set(0)
if i == 0:
out = self.encoder(x[:, :, :1, :, :])
else:
out_ = self.encoder(x[:, :,
1 + 4 * (i - 1):1 + 4 * i, :, :])
out = torch.cat([out, out_], 2)
enc = self.quant_conv(out)
mu, logvar = enc[:, :self.z_dim, :, :, :], enc[:,
self.z_dim:, :, :, :]
enc = torch.cat([mu, logvar], dim=1)
enc = DiagonalGaussianDistribution(enc)
self.clear_cache()
else:
for block in self.encoder.down_blocks:
if isinstance(block,
WanResample) and block.mode == "downsample3d":
_padding = list(block.time_conv._padding)
_padding[4] = 2
block.time_conv._padding = tuple(_padding)
enc = ParallelTiledVAE.encode(self, x)
return enc
def _encode(self, x: torch.Tensor, first_frame=False) -> torch.Tensor:
out = self.encoder(x, first_frame=first_frame)
with forward_context(first_frame_arg=first_frame):
out = self.encoder(x)
enc = self.quant_conv(out)
mu, logvar = enc[:, :self.z_dim, :, :, :], enc[:, self.z_dim:, :, :, :]
enc = torch.cat([mu, logvar], dim=1)
@@ -691,14 +877,41 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
enc = torch.cat([first_frame, enc], dim=2)
return enc
def spatial_tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
first_frame = x[:, :, 0, :, :].unsqueeze(2)
first_frame = self._encode(first_frame, first_frame=True)
enc = ParallelTiledVAE.spatial_tiled_encode(self, x)
enc = enc[:, :, 1:]
enc = torch.cat([first_frame, enc], dim=2)
return enc
def decode(self, z: torch.Tensor) -> torch.Tensor:
if self.use_feature_cache:
self.clear_cache()
iter_ = z.shape[2]
x = self.post_quant_conv(z)
with forward_context(feat_cache_arg=self._feat_map,
feat_idx_arg=self._conv_idx):
for i in range(iter_):
feat_idx.set(0)
if i == 0:
out = self.decoder(x[:, :, i:i + 1, :, :])
else:
out_ = self.decoder(x[:, :, i:i + 1, :, :])
out = torch.cat([out, out_], 2)
out = torch.clamp(out, min=-1.0, max=1.0)
self.clear_cache()
else:
out = ParallelTiledVAE.decode(self, z)
return out
def _decode(self, z: torch.Tensor, first_frame=False) -> torch.Tensor:
latents_mean = (torch.tensor(self.latents_mean).view(
1, self.z_dim, 1, 1, 1).to(z.device, z.dtype))
latents_std = 1.0 / torch.tensor(self.latents_std).view(
1, self.z_dim, 1, 1, 1).to(z.device, z.dtype)
z = z / latents_std + latents_mean
x = self.post_quant_conv(z)
out = self.decoder(x, first_frame=first_frame)
with forward_context(first_frame_arg=first_frame):
out = self.decoder(x)
out = torch.clamp(out, min=-1.0, max=1.0)
@@ -711,6 +924,12 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
dec = dec[:, :, start_frame_idx:]
return dec
def spatial_tiled_decode(self, z: torch.Tensor) -> torch.Tensor:
dec = ParallelTiledVAE.spatial_tiled_decode(self, z)
start_frame_idx = self.temporal_compression_ratio - 1
dec = dec[:, :, start_frame_idx:]
return dec
def parallel_tiled_decode(self, z: torch.FloatTensor) -> torch.FloatTensor:
self.blend_num_frames *= 2
dec = ParallelTiledVAE.parallel_tiled_decode(self, z)
+220
View File
@@ -0,0 +1,220 @@
# SPDX-License-Identifier: Apache-2.0
import os
from typing import Callable, List, Optional, Tuple, Union
import numpy as np
import PIL.Image
import PIL.ImageOps
import requests
import torch
from packaging import version
if version.parse(version.parse(
PIL.__version__).base_version) >= version.parse("9.1.0"):
PIL_INTERPOLATION = {
"linear": PIL.Image.Resampling.BILINEAR,
"bilinear": PIL.Image.Resampling.BILINEAR,
"bicubic": PIL.Image.Resampling.BICUBIC,
"lanczos": PIL.Image.Resampling.LANCZOS,
"nearest": PIL.Image.Resampling.NEAREST,
}
else:
PIL_INTERPOLATION = {
"linear": PIL.Image.LINEAR,
"bilinear": PIL.Image.BILINEAR,
"bicubic": PIL.Image.BICUBIC,
"lanczos": PIL.Image.LANCZOS,
"nearest": PIL.Image.NEAREST,
}
def pil_to_numpy(
images: Union[List[PIL.Image.Image], PIL.Image.Image]) -> np.ndarray:
r"""
Convert a PIL image or a list of PIL images to NumPy arrays.
Args:
images (`PIL.Image.Image` or `List[PIL.Image.Image]`):
The PIL image or list of images to convert to NumPy format.
Returns:
`np.ndarray`:
A NumPy array representation of the images.
"""
if not isinstance(images, list):
images = [images]
images = [np.array(image).astype(np.float32) / 255.0 for image in images]
images_arr: np.ndarray = np.stack(images, axis=0)
return images_arr
def numpy_to_pt(images: np.ndarray) -> torch.Tensor:
r"""
Convert a NumPy image to a PyTorch tensor.
Args:
images (`np.ndarray`):
The NumPy image array to convert to PyTorch format.
Returns:
`torch.Tensor`:
A PyTorch tensor representation of the images.
"""
if images.ndim == 3:
images = images[..., None]
images = torch.from_numpy(images.transpose(0, 3, 1, 2))
return images
def normalize(
images: Union[np.ndarray,
torch.Tensor]) -> Union[np.ndarray, torch.Tensor]:
r"""
Normalize an image array to [-1,1].
Args:
images (`np.ndarray` or `torch.Tensor`):
The image array to normalize.
Returns:
`np.ndarray` or `torch.Tensor`:
The normalized image array.
"""
return 2.0 * images - 1.0
def load_image(
image: Union[str, PIL.Image.Image],
convert_method: Optional[Callable[[PIL.Image.Image],
PIL.Image.Image]] = None
) -> PIL.Image.Image:
"""
Loads `image` to a PIL Image.
Args:
image (`str` or `PIL.Image.Image`):
The image to convert to the PIL Image format.
convert_method (Callable[[PIL.Image.Image], PIL.Image.Image], *optional*):
A conversion method to apply to the image after loading it. When set to `None` the image will be converted
"RGB".
Returns:
`PIL.Image.Image`:
A PIL Image.
"""
if isinstance(image, str):
if image.startswith("http://") or image.startswith("https://"):
image = PIL.Image.open(requests.get(image, stream=True).raw)
elif os.path.isfile(image):
image = PIL.Image.open(image)
else:
raise ValueError(
f"Incorrect path or URL. URLs must start with `http://` or `https://`, and {image} is not a valid path."
)
elif isinstance(image, PIL.Image.Image):
image = image
else:
raise ValueError(
"Incorrect format used for the image. Should be a URL linking to an image, a local path, or a PIL image."
)
image = PIL.ImageOps.exif_transpose(image)
if convert_method is not None:
image = convert_method(image)
else:
image = image.convert("RGB")
return image
def get_default_height_width(
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
vae_scale_factor: int,
height: Optional[int] = None,
width: Optional[int] = None,
) -> Tuple[int, int]:
r"""
Returns the height and width of the image, downscaled to the next integer multiple of `vae_scale_factor`.
Args:
image (`Union[PIL.Image.Image, np.ndarray, torch.Tensor]`):
The image input, which can be a PIL image, NumPy array, or PyTorch tensor. If it is a NumPy array, it
should have shape `[batch, height, width]` or `[batch, height, width, channels]`. If it is a PyTorch
tensor, it should have shape `[batch, channels, height, width]`.
height (`Optional[int]`, *optional*, defaults to `None`):
The height of the preprocessed image. If `None`, the height of the `image` input will be used.
width (`Optional[int]`, *optional*, defaults to `None`):
The width of the preprocessed image. If `None`, the width of the `image` input will be used.
Returns:
`Tuple[int, int]`:
A tuple containing the height and width, both resized to the nearest integer multiple of
`vae_scale_factor`.
"""
if height is None:
if isinstance(image, PIL.Image.Image):
height = image.height
elif isinstance(image, torch.Tensor):
height = image.shape[2]
else:
height = image.shape[1]
if width is None:
if isinstance(image, PIL.Image.Image):
width = image.width
elif isinstance(image, torch.Tensor):
width = image.shape[3]
else:
width = image.shape[2]
width, height = (x - x % vae_scale_factor for x in (width, height)
) # resize to integer multiple of vae_scale_factor
return height, width
def resize(
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
height: int,
width: int,
resize_mode: str = "default", # "default", "fill", "crop"
resample: str = "lanczos",
) -> Union[PIL.Image.Image, np.ndarray, torch.Tensor]:
"""
Resize image.
Args:
image (`PIL.Image.Image`, `np.ndarray` or `torch.Tensor`):
The image input, can be a PIL image, numpy array or pytorch tensor.
height (`int`):
The height to resize to.
width (`int`):
The width to resize to.
resize_mode (`str`, *optional*, defaults to `default`):
The resize mode to use, can be one of `default` or `fill`. If `default`, will resize the image to fit
within the specified width and height, and it may not maintaining the original aspect ratio. If `fill`,
will resize the image to fit within the specified width and height, maintaining the aspect ratio, and
then center the image within the dimensions, filling empty with data from image. If `crop`, will resize
the image to fit within the specified width and height, maintaining the aspect ratio, and then center
the image within the dimensions, cropping the excess. Note that resize_mode `fill` and `crop` are only
supported for PIL image input.
Returns:
`PIL.Image.Image`, `np.ndarray` or `torch.Tensor`:
The resized image.
"""
if resize_mode != "default" and not isinstance(image, PIL.Image.Image):
raise ValueError(
f"Only PIL image input is supported for resize_mode {resize_mode}")
assert isinstance(image, PIL.Image.Image)
if resize_mode == "default":
image = image.resize((width, height),
resample=PIL_INTERPOLATION[resample])
else:
raise ValueError(f"resize_mode {resize_mode} is not supported")
return image
+5 -5
View File
@@ -40,7 +40,7 @@ from fastvideo.v1.pipelines.stages import (
ConditioningStage,
# Import other required stages
)
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
class YourCustomPipeline(ComposedPipelineBase):
@@ -53,7 +53,7 @@ class YourCustomPipeline(ComposedPipelineBase):
# Add other required modules
]
def create_pipeline_stages(self, inference_args: InferenceArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
# Add and configure pipeline stages
self.add_stage(
stage_name="input_validation_stage",
@@ -61,14 +61,14 @@ class YourCustomPipeline(ComposedPipelineBase):
)
# Add more stages as needed
def initialize_pipeline(self, inference_args: InferenceArgs):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Initialize pipeline-specific components
pass
@torch.no_grad()
def forward(self, batch: ForwardBatch, inference_args: InferenceArgs) -> ForwardBatch:
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
# Implement your pipeline's forward pass
batch = self.input_validation_stage(batch, inference_args)
batch = self.input_validation_stage(batch, fastvideo_args)
# Add more stage executions
return batch
+5 -5
View File
@@ -5,7 +5,7 @@ Diffusion pipelines for fastvideo.v1.
This package contains diffusion pipelines for generating videos and images.
"""
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -16,7 +16,7 @@ from fastvideo.v1.utils import (maybe_download_model,
logger = init_logger(__name__)
def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
def build_pipeline(fastvideo_args: FastVideoArgs) -> ComposedPipelineBase:
"""
Only works with valid hf diffusers configs. (model_index.json)
We want to build a pipeline based on the inference args mode_path:
@@ -25,9 +25,9 @@ def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
3. based on the config, determine the pipeline class
"""
# Get pipeline type
model_path = inference_args.model_path
model_path = fastvideo_args.model_path
model_path = maybe_download_model(model_path)
# inference_args.downloaded_model_path = model_path
# fastvideo_args.downloaded_model_path = model_path
logger.info("Model path: %s", model_path)
config = verify_model_config_and_directory(model_path)
@@ -41,7 +41,7 @@ def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
pipeline_architecture)
# instantiate the pipeline
pipeline = pipeline_cls(model_path, inference_args, config)
pipeline = pipeline_cls(model_path, fastvideo_args, config)
logger.info("Pipeline instantiated")
# pipeline is now initialized and ready to use
@@ -12,7 +12,7 @@ from typing import Any, Dict, List, Optional, cast
import torch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -38,7 +38,7 @@ class ComposedPipelineBase(ABC):
# TODO(will): args should support both inference args and training args
def __init__(self,
model_path: str,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
config: Optional[Dict[str, Any]] = None):
"""
Initialize the pipeline. After __init__, the pipeline should be ready to
@@ -61,12 +61,12 @@ class ComposedPipelineBase(ABC):
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
self.modules = self.load_modules(inference_args)
self.modules = self.load_modules(fastvideo_args)
self.initialize_pipeline(inference_args)
self.initialize_pipeline(fastvideo_args)
logger.info("Creating pipeline stages...")
self.create_pipeline_stages(inference_args)
self.create_pipeline_stages(fastvideo_args)
def get_module(self, module_name: str) -> Any:
return self.modules[module_name]
@@ -77,7 +77,7 @@ class ComposedPipelineBase(ABC):
def _load_config(self, model_path: str) -> Dict[str, Any]:
model_path = maybe_download_model(self.model_path)
self.model_path = model_path
# inference_args.downloaded_model_path = model_path
# fastvideo_args.downloaded_model_path = model_path
logger.info("Model path: %s", model_path)
config = verify_model_config_and_directory(model_path)
return cast(Dict[str, Any], config)
@@ -108,20 +108,20 @@ class ComposedPipelineBase(ABC):
return self._stages
@abstractmethod
def create_pipeline_stages(self, inference_args: InferenceArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""
Create the pipeline stages.
"""
raise NotImplementedError
@abstractmethod
def initialize_pipeline(self, inference_args: InferenceArgs):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
raise NotImplementedError
def load_modules(self, inference_args: InferenceArgs) -> Dict[str, Any]:
def load_modules(self, fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
"""
Load the modules from the config.
"""
@@ -156,7 +156,7 @@ class ComposedPipelineBase(ABC):
component_model_path=component_model_path,
transformers_or_diffusers=transformers_or_diffusers,
architecture=architecture,
inference_args=inference_args,
fastvideo_args=fastvideo_args,
)
logger.info("Loaded module %s from %s", module_name,
component_model_path)
@@ -185,14 +185,14 @@ class ComposedPipelineBase(ABC):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Generate a video or image using the pipeline.
Args:
batch: The batch to generate from.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
ForwardBatch: The batch with the generated video or image.
"""
@@ -201,7 +201,7 @@ class ComposedPipelineBase(ABC):
self._stage_name_mapping.keys())
logger.info("Batch: %s", batch)
for stage in self.stages:
batch = stage(batch, inference_args)
batch = stage(batch, fastvideo_args)
# Return the output
return batch
@@ -6,9 +6,7 @@ This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
from diffusers.image_processor import VaeImageProcessor
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.stages import (CLIPTextEncodingStage,
@@ -30,7 +28,7 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
"transformer", "scheduler"
]
def create_pipeline_stages(self, inference_args: InferenceArgs):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
@@ -67,20 +65,12 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def initialize_pipeline(self, inference_args: InferenceArgs):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""
Initialize the pipeline.
"""
vae_scale_factor = 2**(len(self.get_module("vae").block_out_channels) -
1)
inference_args.vae_scale_factor = vae_scale_factor
self.image_processor = VaeImageProcessor(
vae_scale_factor=vae_scale_factor)
self.add_module("image_processor", self.image_processor)
num_channels_latents = self.get_module("transformer").in_channels
inference_args.num_channels_latents = num_channels_latents
fastvideo_args.num_channels_latents = num_channels_latents
EntryClass = HunyuanVideoPipeline
@@ -22,13 +22,17 @@ class ForwardBatch:
execution, allowing methods to update specific components without needing
to manage numerous individual parameters.
"""
# TODO(will): double check that args are separate from inference_args
# TODO(will): double check that args are separate from fastvideo_args
# properly. Also maybe think about providing an abstraction for pipeline
# specific arguments.
data_type: str
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None
# Image inputs
image_path: Optional[str] = None
image_embeds: List[torch.Tensor] = field(default_factory=list)
# Text inputs
prompt: Optional[Union[str, List[str]]] = None
negative_prompt: Optional[Union[str, List[str]]] = None
@@ -36,8 +40,6 @@ class ForwardBatch:
# Primary encoder embeddings
prompt_embeds: List[torch.Tensor] = field(default_factory=list)
negative_prompt_embeds: Optional[List[torch.Tensor]] = None
attention_mask: List[torch.Tensor] = field(default_factory=list)
negative_attention_mask: List[torch.Tensor] = field(default_factory=list)
# Additional text-related parameters
max_sequence_length: Optional[int] = None
@@ -55,6 +57,7 @@ class ForwardBatch:
# Latent tensors
latents: Optional[torch.Tensor] = None
noise_pred: Optional[torch.Tensor] = None
image_latent: Optional[torch.Tensor] = None
# Latent dimensions
num_channels_latents: Optional[int] = None
@@ -100,3 +103,4 @@ class ForwardBatch:
# Set do_classifier_free_guidance based on guidance scale and negative prompt
if self.guidance_scale > 1.0:
self.do_classifier_free_guidance = True
self.negative_prompt_embeds = []
@@ -7,15 +7,19 @@ complete diffusion pipelines.
"""
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.clip_image_encoding import (
CLIPImageEncodingStage)
from fastvideo.v1.pipelines.stages.clip_text_encoding import (
CLIPTextEncodingStage)
from fastvideo.v1.pipelines.stages.conditioning import ConditioningStage
from fastvideo.v1.pipelines.stages.decoding import DecodingStage
from fastvideo.v1.pipelines.stages.denoising import DenoisingStage
from fastvideo.v1.pipelines.stages.encoding import EncodingStage
from fastvideo.v1.pipelines.stages.input_validation import InputValidationStage
from fastvideo.v1.pipelines.stages.latent_preparation import (
LatentPreparationStage)
from fastvideo.v1.pipelines.stages.llama_encoding import LlamaEncodingStage
from fastvideo.v1.pipelines.stages.t5_encoding import T5EncodingStage
from fastvideo.v1.pipelines.stages.timestep_preparation import (
TimestepPreparationStage)
@@ -26,7 +30,10 @@ __all__ = [
"LatentPreparationStage",
"ConditioningStage",
"DenoisingStage",
"EncodingStage",
"DecodingStage",
"LlamaEncodingStage",
"T5EncodingStage",
"CLIPTextEncodingStage",
"CLIPImageEncodingStage",
]
+8 -8
View File
@@ -12,7 +12,7 @@ from abc import ABC, abstractmethod
import torch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
@@ -45,7 +45,7 @@ class PipelineStage(ABC):
def __call__(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Execute the stage's processing on the batch with optional logging.
@@ -53,7 +53,7 @@ class PipelineStage(ABC):
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The updated batch information after this stage's processing.
@@ -65,7 +65,7 @@ class PipelineStage(ABC):
try:
# Call the actual implementation
result = self._call_implementation(batch, inference_args)
result = self._call_implementation(batch, fastvideo_args)
execution_time = time.time() - start_time
self._logger.info("[%s] Execution completed in %s ms",
@@ -85,13 +85,13 @@ class PipelineStage(ABC):
else:
# Just call the implementation directly if logging is disabled
# TODO(will): Also handle backward
return self.forward(batch, inference_args)
return self.forward(batch, fastvideo_args)
@abstractmethod
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Forward pass of the stage's processing.
@@ -101,7 +101,7 @@ class PipelineStage(ABC):
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The updated batch information after this stage's processing.
@@ -111,6 +111,6 @@ class PipelineStage(ABC):
def backward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
raise NotImplementedError
@@ -0,0 +1,71 @@
# SPDX-License-Identifier: Apache-2.0
"""
Image encoding stages for I2V diffusion pipelines.
This module contains implementations of image encoding stages for diffusion pipelines.
"""
import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vision_utils import load_image
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
logger = init_logger(__name__)
class CLIPImageEncodingStage(PipelineStage):
"""
Stage for encoding image prompts into embeddings for diffusion models.
This stage handles the encoding of image prompts into the embedding space
expected by the diffusion model.
"""
def __init__(self, image_encoder, image_processor) -> None:
"""
Initialize the prompt encoding stage.
Args:
enable_logging: Whether to enable logging for this stage.
is_secondary: Whether this is a secondary image encoder.
"""
super().__init__()
self.image_processor = image_processor
self.image_encoder = image_encoder
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode the prompt into image encoder hidden states.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded prompt embeddings.
"""
if fastvideo_args.use_cpu_offload:
self.image_encoder = self.image_encoder.to(batch.device)
image = load_image(batch.image_path)
image_inputs = self.image_processor(
images=image, return_tensors="pt").to(batch.device)
with set_forward_context(current_timestep=0, attn_metadata=None):
image_embeds = self.image_encoder(**image_inputs)
batch.image_embeds.append(image_embeds)
if fastvideo_args.use_cpu_offload:
self.image_encoder.to('cpu')
torch.cuda.empty_cache()
return batch
@@ -5,8 +5,10 @@ Prompt encoding stages for diffusion pipelines.
This module contains implementations of prompt encoding stages for diffusion pipelines.
"""
import torch
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
@@ -37,18 +39,20 @@ class CLIPTextEncodingStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode the prompt into text encoder hidden states.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded prompt embeddings.
"""
if fastvideo_args.use_cpu_offload:
self.text_encoder = self.text_encoder.to(batch.device)
text_inputs = self.tokenizer(
batch.prompt,
@@ -64,4 +68,25 @@ class CLIPTextEncodingStage(PipelineStage):
batch.prompt_embeds.append(prompt_embeds)
if batch.do_classifier_free_guidance:
negative_text_inputs = self.tokenizer(
batch.negative_prompt,
truncation=True,
# better way to handle this?
max_length=77,
return_tensors="pt",
)
with set_forward_context(current_timestep=0, attn_metadata=None):
negative_outputs = self.text_encoder(
input_ids=negative_text_inputs["input_ids"].to(
batch.device), )
negative_prompt_embeds = negative_outputs["pooler_output"]
assert batch.negative_prompt_embeds is not None
batch.negative_prompt_embeds.append(negative_prompt_embeds)
if fastvideo_args.use_cpu_offload:
self.text_encoder.to('cpu')
torch.cuda.empty_cache()
return batch
@@ -5,7 +5,7 @@ Conditioning stage for diffusion pipelines.
import torch
from fastvideo.v1.inference_args import InferenceArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
@@ -24,14 +24,14 @@ class ConditioningStage(PipelineStage):
def forward(
self,
batch: ForwardBatch,
inference_args: InferenceArgs,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Apply conditioning to the diffusion process.
Args:
batch: The current batch information.
inference_args: The inference arguments.
fastvideo_args: The inference arguments.
Returns:
The batch with applied conditioning.
@@ -39,8 +39,7 @@ class ConditioningStage(PipelineStage):
if not batch.do_classifier_free_guidance:
return batch
else:
raise NotImplementedError(
"Classifier-free guidance is not supported yet")
return batch
logger.info("batch.negative_prompt_embeds: %s",
batch.negative_prompt_embeds)

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