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